merge(litellm_internal_staging): resolve vendor strategy e2e conflicts

Keep PR coverage for images edits/input validation, stream disconnect
handling in e2e_http, and gemini-backed vertex embeddings path
This commit is contained in:
mubashir1osmani 2026-07-31 17:25:18 -07:00
commit 4ec82d2e0d
1802 changed files with 69893 additions and 22279 deletions

View file

@ -40,6 +40,7 @@ If you're seeing a delay in your PR being merged, ping the LiteLLM Team on [Slac
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
If the change applies to all three LLM endpoints (/v1/responses, /v1/chat/completions, /v1/messages), include proof for every single one of them, not just one
For new features: show the feature working end-to-end
For UI changes: include before/after screenshots -->

View file

@ -80,7 +80,7 @@ jobs:
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'

View file

@ -24,9 +24,12 @@ jobs:
- name: Update JSON Data
run: |
uv run --frozen --with 'aiohttp==3.13.3' python ".github/workflows/auto_update_price_and_context_window_file.py"
- name: Regenerate JSON Schema
run: |
uv run --frozen python ci_cd/generate_model_prices_schema.py
- name: Create Pull Request
run: |
git add model_prices_and_context_window.json
git add model_prices_and_context_window.json model_prices_and_context_window.schema.json
git commit -m "Update model_prices_and_context_window.json file: $(date +'%Y-%m-%d')"
gh pr create --title "Update model_prices_and_context_window.json file" \
--body "Automated update for model_prices_and_context_window.json" \

View file

@ -55,7 +55,7 @@ jobs:
- name: Install dependencies
run: |
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
- name: Generate Prisma client
env:

View file

@ -104,6 +104,7 @@ jobs:
- name: Check basedpyright budget (delta vs base)
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
NODE_OPTIONS: --max-old-space-size=12288
run: |
(uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --base "$BASE_SHA"

View file

@ -29,7 +29,15 @@ jobs:
id: changed
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
HEAD_SHA: ${{ github.event.pull_request.head.sha }}
run: |
# base.sha is the base branch tip from when the PR was opened, while
# actions/checkout leaves HEAD on a merge of the PR into the *current*
# base tip. "$BASE_SHA"...HEAD therefore spans every base-branch commit
# landed since, so a PR that touches no UI file still gets linted
# against hundreds of other people's files. Diff the PR head against its
# own merge base instead, which is exactly what this PR changed.
merge_base=$(git merge-base "$BASE_SHA" "$HEAD_SHA")
: > "$RUNNER_TEMP/prettier_files.txt"
: > "$RUNNER_TEMP/eslint_files.txt"
while IFS= read -r f; do
@ -41,7 +49,7 @@ jobs:
*.json | *.css | *.scss | *.md | *.mdx | *.yml | *.yaml | *.html)
printf '%s\n' "$f" >> "$RUNNER_TEMP/prettier_files.txt" ;;
esac
done < <(git diff --name-only --diff-filter=ACMR --relative "$BASE_SHA"...HEAD -- .)
done < <(git diff --name-only --diff-filter=ACMR --relative "$merge_base" "$HEAD_SHA" -- .)
if [ -s "$RUNNER_TEMP/prettier_files.txt" ] || [ -s "$RUNNER_TEMP/eslint_files.txt" ]; then
echo "has_files=true" >> "$GITHUB_OUTPUT"
else

View file

@ -22,3 +22,12 @@ jobs:
- name: Validate model_prices_and_context_window.json
run: |
jq empty model_prices_and_context_window.json
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Check model_prices_and_context_window.schema.json is in sync
run: |
uv run --frozen python ci_cd/generate_model_prices_schema.py --check

View file

@ -27,6 +27,7 @@ jobs:
tests/test_litellm/a2a_protocol
tests/test_litellm/anthropic_interface
tests/test_litellm/completion_extras
tests/test_litellm/compression
tests/test_litellm/containers
tests/test_litellm/experimental_mcp_client
tests/test_litellm/models

View file

@ -1,151 +0,0 @@
name: Test Proxy SERVER_ROOT_PATH Routing
permissions:
contents: read
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
jobs:
test-server-root-path:
runs-on: ubuntu-latest
timeout-minutes: 30
strategy:
fail-fast: false
matrix:
root_path: ["/api/v1", "/llmproxy"]
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Free up disk space
run: |
sudo rm -rf /usr/local/lib/android /usr/share/dotnet /opt/ghc /usr/local/share/boost
sudo apt-get clean
df -h /
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
- name: Build Docker image
uses: docker/build-push-action@0adf9959216b96bec444f325f1e493d4aa344497 # v6.14.0
with:
context: .
file: ./docker/Dockerfile.non_root
tags: litellm-test:${{ github.sha }}
load: true
push: false
- name: Start LiteLLM container with SERVER_ROOT_PATH
run: |
docker run -d \
--name litellm-test \
-p 4000:4000 \
-e SERVER_ROOT_PATH="${{ matrix.root_path }}" \
-e LITELLM_MASTER_KEY="sk-1234" \
litellm-test:${{ github.sha }} \
--detailed_debug
- name: Wait for container to be healthy
run: |
echo "Waiting for LiteLLM to start..."
max_attempts=30
attempt=0
while [ $attempt -lt $max_attempts ]; do
if docker logs litellm-test 2>&1 | grep -q "Uvicorn running"; then
echo "LiteLLM started successfully"
break
fi
attempt=$((attempt + 1))
echo "Attempt $attempt/$max_attempts - waiting for server to start..."
sleep 2
done
if [ $attempt -eq $max_attempts ]; then
echo "Server failed to start within timeout"
docker logs litellm-test
exit 1
fi
sleep 5
- name: Show container logs
if: always()
run: docker logs litellm-test
- name: Test UI endpoint with root path
run: |
ROOT_PATH="${{ matrix.root_path }}"
echo "Testing UI at: http://localhost:4000${ROOT_PATH}/ui/"
for i in 1 2 3; do
content=$(curl -sL --max-time 5 -H "Authorization: Bearer sk-1234" "http://localhost:4000${ROOT_PATH}/ui/")
if echo "$content" | grep -q -E "(html|<!DOCTYPE|<head|<body)"; then
echo "UI page contains valid HTML content"
exit 0
fi
echo "Attempt $i/3 - no valid HTML, retrying in 5s..."
sleep 5
done
echo "UI page does not contain expected HTML content"
echo "Response: $content"
docker logs litellm-test
exit 1
- name: Setup Node for Playwright
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4.4.0
with:
node-version: "20"
- name: Install e2e deps and Chromium
working-directory: tests/e2e/ui
run: |
retry() {
local attempt=1
local max_attempts=4
until "$@"; do
if [ "$attempt" -ge "$max_attempts" ]; then
echo "Command failed after $attempt attempts: $*"
return 1
fi
echo "Attempt $attempt failed: $*. Retrying in $((attempt * 15))s..."
sleep $((attempt * 15))
attempt=$((attempt + 1))
done
}
npm config set fetch-retries 5
npm config set fetch-retry-mintimeout 20000
npm config set fetch-retry-maxtimeout 120000
retry npm ci
retry npx playwright install --with-deps chromium
- name: Run SERVER_ROOT_PATH redirect e2e
working-directory: tests/e2e/ui
env:
SERVER_ROOT_PATH: ${{ matrix.root_path }}
run: npx playwright test --config=serverRootPath.config.ts
- name: Upload Playwright artifacts on failure
if: failure()
uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2
with:
name: playwright-trace-${{ strategy.job-index }}
path: tests/e2e/ui/test-results/
retention-days: 7
- name: Cleanup
if: always()
run: |
docker stop litellm-test || true
docker rm litellm-test || true

1
.gitignore vendored
View file

@ -141,3 +141,4 @@ crash.*.log
.coverage
ui/litellm-dashboard/out/
litellm.log

View file

@ -61,6 +61,8 @@ Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages
When working on a PR, keep the PR description in sync with new commits being made
Replies/rebuttals to AI PR review bots must be 15-25 word human-readable replies
Monkeypatching attributes of a class to do testing is an anti-pattern. Prefer dependency-injecting things into classes. That way, at unit test time, you can pass a mocked dependency in
Do not put names of customers or customer company names in code, PR descriptions, issue bodies, etc. This means never mention literally any company name. Especially if you're about to say a sentence mentioning that the reason the PR exists was a feature/model/bug fix/etc. requested by a company. That's the indication that you should replace that company name with "the customer". e.g. not "Model request from Acme (Pylon #1234)" but "Model request from a customer (Pylon #1234)". This is because the codebase is public. The only exception is for publicly known providers or vendors such as OpenAI, Anthropic, AWS Bedrock, etc. only IF we're adding support for that provider/vendor in general and NOT if that PR or whatnot was a request by one of them, and they're actually one of our customers.

View file

@ -64,6 +64,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
# Copy full source tree
@ -84,6 +85,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -176,6 +176,8 @@ lint-ruff-FULL-dev: install-dev
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
else echo "No changed .py files to check."; fi
lint-basedpyright lint-basedpyright-budget-update: export NODE_OPTIONS := --max-old-space-size=12288
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging

View file

@ -70,6 +70,10 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/project/",
"/memory/",
"/mcp/",
# Control plane (see the List Endpoints + Tables standard). Every resource
# eventually moves under this prefix, so allowlist it once rather than
# per-resource.
"/management/v1/",
# Spend / analytics
"/spend/",
"/analytics/",

View file

@ -1,30 +1,30 @@
{
"reportAny": {
"limit": 37484
"limit": 31903
},
"reportArgumentType": {
"limit": 2704
"limit": 2645
},
"reportAssignmentType": {
"limit": 330
"limit": 329
},
"reportAttributeAccessIssue": {
"limit": 516
},
"reportCallIssue": {
"limit": 124
"limit": 123
},
"reportConstantRedefinition": {
"limit": 59
},
"reportDeprecated": {
"limit": 326
"limit": 325
},
"reportDuplicateImport": {
"limit": 42
},
"reportExplicitAny": {
"limit": 10389
"limit": 10214
},
"reportFunctionMemberAccess": {
"limit": 11
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5900
"limit": 5869
},
"reportMissingTypeArgument": {
"limit": 15903
"limit": 15861
},
"reportMissingTypeStubs": {
"limit": 41
@ -72,7 +72,7 @@
"limit": 0
},
"reportOptionalMemberAccess": {
"limit": 1085
"limit": 1079
},
"reportOptionalOperand": {
"limit": 0
@ -84,13 +84,13 @@
"limit": 77
},
"reportPrivateUsage": {
"limit": 2438
"limit": 2437
},
"reportRedeclaration": {
"limit": 12
},
"reportReturnType": {
"limit": 225
"limit": 219
},
"reportTypedDictNotRequiredAccess": {
"limit": 27
@ -99,31 +99,31 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 45894
"limit": 45366
},
"reportUnknownLambdaType": {
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40539
"limit": 40477
},
"reportUnknownParameterType": {
"limit": 20403
"limit": 20338
},
"reportUnknownVariableType": {
"limit": 32141
"limit": 32047
},
"reportUnnecessaryCast": {
"limit": 177
},
"reportUnnecessaryComparison": {
"limit": 1025
"limit": 1021
},
"reportUnnecessaryContains": {
"limit": 7
},
"reportUnnecessaryIsInstance": {
"limit": 1209
"limit": 1205
},
"reportUntypedBaseClass": {
"limit": 165
@ -135,10 +135,10 @@
"limit": 33
},
"reportUnusedFunction": {
"limit": 206
"limit": 204
},
"reportUnusedImport": {
"limit": 1005
"limit": 1003
},
"reportUnusedVariable": {
"limit": 1297

View file

@ -0,0 +1,325 @@
from __future__ import annotations
import json
import sys
from pathlib import Path
from typing import Optional
import jsonschema
REPO_ROOT = Path(__file__).parent.parent
PRICES_PATH = REPO_ROOT / "model_prices_and_context_window.json"
SCHEMA_PATH = REPO_ROOT / "model_prices_and_context_window.schema.json"
SPECIAL_ROOT_KEYS = frozenset({"sample_spec", "fallback_generalizations"})
JsonSchema = dict
NONNEG_NUMBER: JsonSchema = {"type": "number", "minimum": 0}
NONNEG_INTEGER: JsonSchema = {"type": "integer", "minimum": 0}
BOOLEAN: JsonSchema = {"type": "boolean"}
STRING: JsonSchema = {"type": "string"}
EXTRA_BOOLEAN_KEYS = frozenset(
{
"gemini_native_audio",
"gemini_audio_only_live",
"uses_embed_content",
"use_openai_responses_path",
"bedrock_converse_supports_strict_tools",
}
)
OBJECT_KEYS: dict[str, JsonSchema] = {
"search_context_cost_per_query": {
"type": "object",
"description": "USD cost per web search query, keyed by search context size.",
"properties": {
"search_context_size_low": NONNEG_NUMBER,
"search_context_size_medium": NONNEG_NUMBER,
"search_context_size_high": NONNEG_NUMBER,
},
"additionalProperties": False,
},
"metadata": {
"type": "object",
"description": "Free-form notes about the entry (e.g. pricing derivation).",
},
"provider_specific_entry": {
"type": "object",
"description": "Provider-internal routing hints (e.g. bedrock_invocation_schema).",
},
}
ARRAY_KEYS: dict[str, JsonSchema] = {
"supported_endpoints": {
"type": "array",
"description": "OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.",
"items": STRING,
},
"supported_modalities": {
"type": "array",
"description": "Input modalities the model accepts.",
"items": {"type": "string", "enum": ["text", "image", "audio", "video"]},
},
"supported_output_modalities": {
"type": "array",
"description": "Output modalities the model can produce.",
"items": {"type": "string", "enum": ["text", "image", "audio", "video", "code"]},
},
"supported_regions": {
"type": "array",
"description": "Cloud regions the model is available in ('global' or region ids).",
"items": STRING,
},
"tiered_pricing": {
"type": "array",
"description": "Context-length or result-count tiered rates; each tier's costs apply within its range.",
"items": {
"type": "object",
"properties": {
"range": {
"type": "array",
"description": "[min, max] prompt-token span this tier applies to.",
"items": NONNEG_NUMBER,
"minItems": 2,
"maxItems": 2,
},
"max_results_range": {
"type": "array",
"description": "[min, max] result-count span this tier applies to (search models).",
"items": NONNEG_NUMBER,
"minItems": 2,
"maxItems": 2,
},
"input_cost_per_token": NONNEG_NUMBER,
"output_cost_per_token": NONNEG_NUMBER,
"output_cost_per_reasoning_token": NONNEG_NUMBER,
"cache_read_input_token_cost": NONNEG_NUMBER,
"input_cost_per_query": NONNEG_NUMBER,
},
"additionalProperties": False,
},
},
}
INTEGER_KEYS: dict[str, JsonSchema] = {
"max_tokens": {
**NONNEG_INTEGER,
"description": "Legacy field: max output tokens if the provider specifies it, else max input tokens.",
},
"max_input_tokens": {
**NONNEG_INTEGER,
"description": "Maximum prompt/context tokens the model accepts.",
},
"max_output_tokens": {
**NONNEG_INTEGER,
"description": "Maximum tokens the model can generate in one response.",
},
"output_vector_size": {
**NONNEG_INTEGER,
"description": "Embedding dimension for embedding models.",
},
"prompt_cache_min_tokens": {
**NONNEG_INTEGER,
"description": "Smallest prefix the provider will actually cache; absent means the provider default applies.",
},
"tpm": {**NONNEG_INTEGER, "description": "Provider default tokens-per-minute limit."},
"rpm": {**NONNEG_INTEGER, "description": "Provider default requests-per-minute limit."},
}
NUMBER_KEYS: dict[str, JsonSchema] = {
"regional_processing_uplift_multiplier_eu": {
"type": "number",
"minimum": 1,
"description": "Multiplier applied to all token costs for EU data residency (e.g. 1.10 = +10%).",
},
"regional_processing_uplift_multiplier_us": {
"type": "number",
"minimum": 1,
"description": "Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%).",
},
}
COST_DESCRIPTIONS: dict[str, str] = {
"input_cost_per_token": "USD per prompt token.",
"output_cost_per_token": "USD per generated token.",
"output_cost_per_reasoning_token": "USD per reasoning/thinking token, when billed separately.",
"cache_creation_input_token_cost": "USD per token written to the provider's prompt cache.",
"cache_read_input_token_cost": "USD per prompt token served from the provider's prompt cache.",
"input_cost_per_token_batches": "USD per prompt token via the provider's batch API.",
"output_cost_per_token_batches": "USD per generated token via the provider's batch API.",
}
def cost_description(key: str) -> Optional[str]:
if key in COST_DESCRIPTIONS:
return COST_DESCRIPTIONS[key]
if key.endswith("_flex"):
return "Flex service-tier rate for the same-named base field."
if key.endswith("_priority"):
return "Priority service-tier rate for the same-named base field."
if "_above_" in key:
return "Rate applied once the prompt exceeds the token threshold in the field name."
return None
def cost_schema(key: str) -> JsonSchema:
description = cost_description(key)
return {**NONNEG_NUMBER, "description": description} if description else dict(NONNEG_NUMBER)
def string_key_schemas(modes: tuple) -> dict[str, JsonSchema]:
return {
"litellm_provider": {
"type": "string",
"description": "LiteLLM provider slug; one of https://docs.litellm.ai/docs/providers.",
},
"mode": {
"type": "string",
"description": "Primary API surface / task type of the model.",
"enum": list(modes),
},
"source": {
"type": "string",
"description": "URL of the provider pricing/model page this entry was taken from.",
},
"deprecation_date": {
"type": "string",
"description": "Date the provider deprecates the model, YYYY-MM-DD.",
"format": "date",
"pattern": "^\\d{4}-(0[1-9]|1[0-2])-(0[1-9]|[12]\\d|3[01])$",
},
"web_search_billing_unit": {
"type": "string",
"description": "Whether web search is billed per query or per prompt.",
"enum": ["per_query", "per_prompt"],
},
"bedrock_output_config_effort_ceiling": {
"type": "string",
"description": "Highest reasoning effort the Bedrock output_config accepts for this model.",
"enum": ["low", "medium", "high", "max", "xhigh"],
},
"comment": STRING,
"audio_transcription_config": STRING,
}
def classify(key: str, modes: tuple) -> Optional[JsonSchema]:
curated = {**OBJECT_KEYS, **ARRAY_KEYS, **string_key_schemas(modes), **INTEGER_KEYS, **NUMBER_KEYS}
if key in curated:
return curated[key]
if key.startswith("supports_") or key in EXTRA_BOOLEAN_KEYS:
return BOOLEAN
if "cost" in key:
return cost_schema(key)
return None
def build_schema(prices: dict) -> JsonSchema:
entries = {name: entry for name, entry in prices.items() if name not in SPECIAL_ROOT_KEYS}
all_keys = tuple(sorted({key for entry in entries.values() for key in entry}))
modes = tuple(sorted({entry["mode"] for entry in entries.values() if "mode" in entry}))
unclassified = tuple(key for key in all_keys if classify(key, modes) is None)
if unclassified:
raise SystemExit(
f"Unclassified keys in {PRICES_PATH.name}: {', '.join(unclassified)}. "
f"Add them to the key tables in {Path(__file__).name} and rerun it."
)
entry_properties = {key: classify(key, modes) for key in all_keys}
return {
"$schema": "https://json-schema.org/draft/2020-12/schema",
"title": "LiteLLM model_prices_and_context_window.json",
"description": (
"Schema for LiteLLM's model price and context window registry "
"(https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). "
"Every top-level key except 'sample_spec' and 'fallback_generalizations' is a model id, "
"optionally prefixed with its provider (e.g. 'azure/gpt-5.4'), mapping to a model entry. "
"All costs are USD per unit. New optional fields are added regularly, so consumers should "
"ignore unknown fields rather than reject them."
),
"type": "object",
"properties": {
"sample_spec": {
"type": "object",
"description": (
"Documentation placeholder illustrating the entry shape; not a real model and not "
"schema-conformant (several values are prose)."
),
},
"fallback_generalizations": {
"type": "object",
"description": "Regex rules that generalize unknown model ids to known families; not a model entry.",
"properties": {
"rules": {
"type": "array",
"items": {
"type": "object",
"properties": {
"name": STRING,
"pattern": STRING,
"description": STRING,
},
"required": ["name", "pattern"],
"additionalProperties": True,
},
}
},
"additionalProperties": False,
},
},
"additionalProperties": {"$ref": "#/$defs/modelEntry"},
"$defs": {
"modelEntry": {
"type": "object",
"description": (
"Pricing, limits, and capability flags for one model. Fields other than litellm_provider "
"are optional; boolean capability flags are simply omitted when unknown or false."
),
"required": ["litellm_provider"],
"properties": entry_properties,
"additionalProperties": True,
}
},
}
def render(schema: JsonSchema) -> str:
return json.dumps(schema, indent=2) + "\n"
def validation_errors(prices: dict, schema: JsonSchema) -> tuple:
validator = jsonschema.Draft202012Validator(
schema, format_checker=jsonschema.Draft202012Validator.FORMAT_CHECKER
)
return tuple(
f"{'.'.join(str(part) for part in error.absolute_path)}: {error.message}"
for error in validator.iter_errors(prices)
)
def main() -> int:
check = "--check" in sys.argv[1:]
prices = json.loads(PRICES_PATH.read_text())
rendered = render(build_schema(prices))
errors = validation_errors(prices, json.loads(rendered))
if errors:
print(f"{PRICES_PATH.name} does not validate against the generated schema:")
print("\n".join(errors[:20]))
return 1
if not check:
SCHEMA_PATH.write_text(rendered)
print(f"wrote {SCHEMA_PATH}")
return 0
if not SCHEMA_PATH.exists() or SCHEMA_PATH.read_text() != rendered:
print(
f"{SCHEMA_PATH.name} is out of sync with {PRICES_PATH.name}. "
f"Run `python {Path(__file__).relative_to(REPO_ROOT)}` and commit the result."
)
return 1
print(f"{SCHEMA_PATH.name} is in sync and {PRICES_PATH.name} validates against it")
return 0
if __name__ == "__main__":
sys.exit(main())

View file

@ -0,0 +1,523 @@
{
"annotations": {
"list": []
},
"editable": true,
"fiscalYearStartMonth": 0,
"graphTooltip": 0,
"links": [],
"panels": [
{
"type": "stat",
"title": "Requests",
"gridPos": {
"h": 4,
"w": 6,
"x": 0,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "short",
"decimals": 0,
"color": {
"mode": "fixed",
"fixedColor": "blue"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "sum(increase(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
}
],
"id": 1
},
{
"type": "stat",
"title": "Spend",
"description": "LiteLLM's computed cost for the selected window, from gen_ai.usage.cost",
"gridPos": {
"h": 4,
"w": 6,
"x": 6,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "currencyUSD",
"decimals": 4,
"color": {
"mode": "fixed",
"fixedColor": "green"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "sum(increase(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
}
],
"id": 2
},
{
"type": "stat",
"title": "Tokens",
"gridPos": {
"h": 4,
"w": 6,
"x": 12,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "short",
"decimals": 0,
"color": {
"mode": "fixed",
"fixedColor": "purple"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "sum(increase(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range]))"
}
],
"id": 3
},
{
"type": "stat",
"title": "p95 request duration",
"gridPos": {
"h": 4,
"w": 6,
"x": 18,
"y": 0
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"decimals": 2,
"color": {
"mode": "fixed",
"fixedColor": "orange"
}
},
"overrides": []
},
"options": {
"reduceOptions": {
"calcs": [
"lastNotNull"
],
"fields": "",
"values": false
},
"colorMode": "background",
"graphMode": "none"
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"instant": true,
"expr": "histogram_quantile(0.95, sum by (le) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__range])))"
}
],
"id": 4
},
{
"type": "timeseries",
"title": "Request rate by model",
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 4
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "reqpm",
"custom": {
"lineWidth": 2,
"fillOpacity": 8,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "sum by (gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60"
}
],
"id": 5
},
{
"type": "timeseries",
"title": "Spend rate by model",
"description": "USD per hour, derived from the gen_ai.usage.cost histogram",
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 4
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "currencyUSD",
"custom": {
"lineWidth": 2,
"fillOpacity": 8,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "sum by (gen_ai_request_model) (rate(gen_ai_usage_cost_USD_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 3600"
}
],
"id": 6
},
{
"type": "timeseries",
"title": "Tokens per minute by model and type",
"description": "gen_ai.client.token.usage split by the gen_ai.token.type attribute",
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 12
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "short",
"custom": {
"lineWidth": 2,
"fillOpacity": 8,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}} {{gen_ai_token_type}}",
"expr": "sum by (gen_ai_request_model, gen_ai_token_type) (rate(gen_ai_client_token_usage_sum{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])) * 60"
}
],
"id": 7
},
{
"type": "timeseries",
"title": "p95 request duration by model",
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 12
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"custom": {
"lineWidth": 2,
"fillOpacity": 0,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_operation_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
}
],
"id": 8
},
{
"type": "timeseries",
"title": "p95 time to first token (streaming)",
"description": "gen_ai.server.time_to_first_token, recorded only for streaming requests",
"gridPos": {
"h": 8,
"w": 12,
"x": 0,
"y": 20
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"custom": {
"lineWidth": 2,
"fillOpacity": 0,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_server_time_to_first_token_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
}
],
"id": 9
},
{
"type": "timeseries",
"title": "p95 provider generation time",
"description": "gen_ai.client.response.duration, upstream generation time excluding LiteLLM overhead",
"gridPos": {
"h": 8,
"w": 12,
"x": 12,
"y": 20
},
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"fieldConfig": {
"defaults": {
"unit": "s",
"custom": {
"lineWidth": 2,
"fillOpacity": 0,
"showPoints": "never"
}
},
"overrides": []
},
"options": {
"legend": {
"displayMode": "list",
"placement": "bottom"
},
"tooltip": {
"mode": "multi",
"sort": "desc"
}
},
"targets": [
{
"refId": "A",
"editorMode": "code",
"legendFormat": "{{gen_ai_request_model}}",
"expr": "histogram_quantile(0.95, sum by (le, gen_ai_request_model) (rate(gen_ai_client_response_duration_seconds_bucket{service_name=~\"$service\", gen_ai_request_model=~\"$model\"}[$__rate_interval])))"
}
],
"id": 10
}
],
"preload": false,
"refresh": "30s",
"schemaVersion": 42,
"tags": [
"litellm",
"genai",
"opentelemetry"
],
"templating": {
"list": [
{
"name": "datasource",
"label": "Prometheus",
"type": "datasource",
"query": "prometheus",
"current": {},
"hide": 0
},
{
"name": "service",
"label": "Service",
"type": "query",
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"query": "label_values(gen_ai_client_operation_duration_seconds_count, service_name)",
"refresh": 2,
"includeAll": true,
"multi": true,
"current": {
"text": "All",
"value": "$__all"
}
},
{
"name": "model",
"label": "Model",
"type": "query",
"datasource": {
"type": "prometheus",
"uid": "${datasource}"
},
"query": "label_values(gen_ai_client_operation_duration_seconds_count{service_name=~\"$service\"}, gen_ai_request_model)",
"refresh": 2,
"includeAll": true,
"multi": true,
"current": {
"text": "All",
"value": "$__all"
}
}
]
},
"time": {
"from": "now-1h",
"to": "now"
},
"timepicker": {},
"timezone": "browser",
"title": "LiteLLM GenAI (OpenTelemetry)",
"uid": "litellm-genai-otel",
"weekStart": ""
}

View file

@ -0,0 +1,35 @@
# LiteLLM GenAI dashboard (OpenTelemetry metrics)
Dashboard for the `gen_ai.*` metrics the OpenTelemetry v2 integration emits, as opposed to the `litellm_*` Prometheus metrics the other dashboards in this folder chart.
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source. Panels: request count, spend, token count, p95 duration, request rate by model, spend rate per hour by model, tokens per minute split by input and output, p95 duration by model, p95 time to first token, and p95 provider generation time. Template variables for data source, service, and model.
## Pre-requisites
Metrics are off by default. In the proxy environment:
```shell
LITELLM_OTEL_V2=true
LITELLM_OTEL_INTEGRATION_ENABLE_METRICS=true
OTEL_EXPORTER="otlp_http"
OTEL_ENDPOINT="<your OTLP endpoint>"
```
You also need the metric attribute filter, or the panels will plot flat lines at zero. LiteLLM's default attribute set includes per-request fields, so nearly every request lands in its own time series with a single sample, and `rate()` has nothing to compute over:
```yaml title="config.yaml"
callback_settings:
otel:
attributes:
include_list:
- gen_ai.operation.name
- gen_ai.system
- gen_ai.request.model
- gen_ai.framework
```
See [Grafana Cloud](https://docs.litellm.ai/docs/observability/grafana_cloud) for the full setup, and [OpenTelemetry v2](https://docs.litellm.ai/docs/observability/opentelemetry_v2#metrics) for the metric reference.
## Note on Grafana's AI Observability integration
Grafana Cloud ships prebuilt GenAI dashboards that query these same metric names, so they look like a drop-in alternative to this one. They are not: twenty of their twenty-two panels filter on `telemetry_sdk_name="openlit"`, a label LiteLLM does not carry and cannot be configured to add, so those panels stay empty.

View file

@ -2,6 +2,10 @@
This folder contains the `json` for creating Grafana Dashboards
## [LiteLLM GenAI Dashboard (OpenTelemetry)](./dashboard_genai_otel)
Charts the `gen_ai.*` metrics from the OpenTelemetry v2 integration: spend, tokens, request rate, and latency percentiles by model. Separate from the dashboards below, which chart the `litellm_*` Prometheus metrics.
## [LiteLLM v2 Dashboard](./dashboard_v2)
<img width="1316" alt="grafana_1" src="https://github.com/user-attachments/assets/d0df802d-0cb9-4906-a679-941c547789ab">

View file

@ -0,0 +1,46 @@
-- One-shot backfill of the LiteLLM_DailyToolSpend rollup from the per-request
-- LiteLLM_SpendLogToolIndex x LiteLLM_SpendLogs tables.
--
-- This is an opt-in, manual operation. New deployments do not need it: the
-- rollup is written at request time from the moment the release is deployed.
-- Run it only if you want the Cost Optimization "Spend by tool" card to show
-- history from before the deploy, and only once.
--
-- IMPORTANT caveats before running:
--
-- 1. Pre-deploy index rows may include tools that were merely DECLARED in a
-- request body but never invoked (the release this ships with stops
-- recording those). For agentic clients that declare many tools per
-- request, backfilled history attributes each request's full spend to
-- every declared tool, overstating per-tool spend. Post-deploy rows do not
-- have this problem. If your traffic is mostly such clients, consider not
-- backfilling.
--
-- 2. Coverage is bounded by spend-log retention: rows older than
-- maximum_spend_logs_retention_period are already gone.
--
-- 3. Replace the cutover timestamp below with the time you deployed the
-- release, so backfilled per-request rows cannot double-count on top of
-- rollup rows the new writer already created. ON CONFLICT DO NOTHING is a
-- second guard for (date, tool_name) buckets the writer already touched:
-- such buckets keep the writer's numbers and skip the backfill's.
--
-- Usage:
-- psql "$DATABASE_URL" -v cutover="'2026-07-25T00:00:00Z'" -f db_scripts/backfill_daily_tool_spend.sql
SET TIME ZONE 'UTC';
INSERT INTO "LiteLLM_DailyToolSpend" (date, tool_name, spend, total_tokens, request_count, created_at, updated_at)
SELECT
to_char(ti.start_time, 'YYYY-MM-DD') AS date,
ti.tool_name,
COALESCE(SUM(sl.spend), 0) AS spend,
COALESCE(SUM(sl.total_tokens), 0) AS total_tokens,
COUNT(*) AS request_count,
now() AS created_at,
now() AS updated_at
FROM "LiteLLM_SpendLogToolIndex" ti
JOIN "LiteLLM_SpendLogs" sl ON sl.request_id = ti.request_id
WHERE ti.start_time < :cutover::timestamptz
GROUP BY 1, 2
ON CONFLICT (date, tool_name) DO NOTHING;

View file

@ -62,6 +62,7 @@ RUN uv sync --frozen --no-install-project --no-install-workspace --no-default-gr
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
# Copy full source tree
@ -82,6 +83,7 @@ RUN uv sync --frozen --no-default-groups --no-editable \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \

View file

@ -68,6 +68,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3
# Copy full source tree
@ -94,6 +95,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3 \
--no-sources-package litellm-proxy-extras; \
else \
@ -102,6 +104,7 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \
--extra proxy-runtime \
--extra extra_proxy \
--extra semantic-router \
--extra saml \
--python python3; \
fi

View file

@ -316,26 +316,34 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
where_clause: Dict[str, Any] = {"file_purpose": "batch", **owner_filter}
if after:
where_clause["id"] = {"gt": after}
cursor_row = (
await self.prisma_client.db.litellm_managedobjecttable.find_first(
where={**where_clause, "unified_object_id": after}
)
)
if cursor_row is None:
raise HTTPException(
status_code=400,
detail=f"Invalid 'after' cursor: no batch found with id '{after}'.",
)
fetch_limit = limit or 20
if target_model_names:
# Oversample so post-fetch model-name filtering still has enough rows.
fetch_limit = max(fetch_limit * 3, 100)
page_size = limit or 20
cursor_args: Dict[str, Any] = (
{"cursor": {"unified_object_id": after}, "skip": 1} if after else {}
)
batches = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where=where_clause,
take=fetch_limit,
order={"created_at": "desc"},
take=page_size + 1,
order=[{"created_at": "desc"}, {"unified_object_id": "desc"}],
**cursor_args,
)
batch_objects: List[LiteLLMBatch] = []
for batch in batches:
try:
# Stop once we have enough after filtering
if len(batch_objects) >= (limit or 20):
break
has_more = len(batches) > page_size
batch_objects: List[LiteLLMBatch] = []
for batch in batches[:page_size]:
try:
batch_data = (
json.loads(batch.file_object)
if isinstance(batch.file_object, str)
@ -351,9 +359,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
continue
return build_list_page(
batch_objects, has_more=len(batch_objects) == (limit or 20)
)
return build_list_page(batch_objects, has_more=has_more)
async def get_user_created_file_ids(
self, user_api_key_dict: UserAPIKeyAuth, model_object_ids: List[str]

View file

@ -11,7 +11,8 @@ Endpoints for /project operations
#### PROJECT MANAGEMENT ####
import json
from typing import List, Optional, Union
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING
from fastapi import APIRouter, Depends, HTTPException, Request
@ -25,15 +26,24 @@ from litellm.proxy.management_helpers.utils import (
)
from litellm.proxy.utils import PrismaClient, handle_exception_on_proxy
if TYPE_CHECKING:
from prisma import models as prisma_models
from prisma.actions import LiteLLM_TeamTableActions
router = APIRouter()
def _team_table(prisma_client: PrismaClient) -> "LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable]":
team_table: LiteLLM_TeamTableActions[prisma_models.LiteLLM_TeamTable] = prisma_client.db.litellm_teamtable
return team_table
async def _check_user_permission_for_project(
user_api_key_dict: UserAPIKeyAuth,
team_id: Optional[str],
team_id: str | None,
prisma_client: PrismaClient,
require_admin: bool = False,
team_object: Optional[LiteLLM_TeamTable] = None,
team_object: LiteLLM_TeamTable | None = None,
) -> bool:
"""
Check if user has permission to manage a project.
@ -57,9 +67,7 @@ async def _check_user_permission_for_project(
team = team_object
if team is None:
team = await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": team_id}
)
team = await _team_table(prisma_client).find_unique(where={"team_id": team_id})
if team and team.admins:
return user_api_key_dict.user_id in team.admins
@ -70,9 +78,9 @@ async def _check_user_permission_for_project(
async def _validate_team_exists(
team_id: str,
prisma_client: PrismaClient,
):
) -> "prisma_models.LiteLLM_TeamTable":
"""Validate that a team exists. Returns the team row."""
team = await prisma_client.db.litellm_teamtable.find_unique(
team = await _team_table(prisma_client).find_unique(
where={"team_id": team_id},
)
@ -89,7 +97,7 @@ async def _validate_team_exists(
def _check_team_project_limits(
team_object: LiteLLM_TeamTable,
data: Union[NewProjectRequest, UpdateProjectRequest],
data: NewProjectRequest | UpdateProjectRequest,
) -> None:
"""
Check that project limits respect its parent Team's limits.
@ -108,16 +116,12 @@ def _check_team_project_limits(
if data.max_budget is not None and data.max_budget < 0:
raise HTTPException(
status_code=400,
detail={
"error": f"max_budget cannot be negative. Received: {data.max_budget}"
},
detail={"error": f"max_budget cannot be negative. Received: {data.max_budget}"},
)
if data.soft_budget is not None and data.soft_budget < 0:
raise HTTPException(
status_code=400,
detail={
"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"
},
detail={"error": f"soft_budget cannot be negative. Received: {data.soft_budget}"},
)
# --- soft_budget < max_budget ---
@ -131,7 +135,7 @@ def _check_team_project_limits(
)
# --- Validate project models are a subset of team models ---
project_models = getattr(data, "models", None)
project_models = data.models
team_models = team_object.models or []
if project_models and len(team_models) > 0:
# If team has 'all-proxy-models', skip validation as it allows all models
@ -148,11 +152,7 @@ def _check_team_project_limits(
# --- Validate project max_budget <= team max_budget ---
# Team stores budget fields directly (max_budget, tpm_limit, rpm_limit)
# unlike Project which uses a separate LiteLLM_BudgetTable relation
if (
data.max_budget is not None
and team_object.max_budget is not None
and data.max_budget > team_object.max_budget
):
if data.max_budget is not None and team_object.max_budget is not None and data.max_budget > team_object.max_budget:
raise HTTPException(
status_code=400,
detail={
@ -161,11 +161,7 @@ def _check_team_project_limits(
)
# --- Validate project tpm_limit <= team tpm_limit ---
if (
data.tpm_limit is not None
and team_object.tpm_limit is not None
and data.tpm_limit > team_object.tpm_limit
):
if data.tpm_limit is not None and team_object.tpm_limit is not None and data.tpm_limit > team_object.tpm_limit:
raise HTTPException(
status_code=400,
detail={
@ -174,11 +170,7 @@ def _check_team_project_limits(
)
# --- Validate project rpm_limit <= team rpm_limit ---
if (
data.rpm_limit is not None
and team_object.rpm_limit is not None
and data.rpm_limit > team_object.rpm_limit
):
if data.rpm_limit is not None and team_object.rpm_limit is not None and data.rpm_limit > team_object.rpm_limit:
raise HTTPException(
status_code=400,
detail={
@ -189,19 +181,19 @@ def _check_team_project_limits(
async def _create_budget_for_project(
data: NewProjectRequest,
user_id: Optional[str],
user_id: str | None,
litellm_proxy_admin_name: str,
prisma_client: PrismaClient,
) -> str:
"""Create a budget for the project and return budget_id."""
budget_params = LiteLLM_BudgetTable.model_fields.keys()
_json_data = data.json(exclude_none=True)
_json_data: Mapping[str, object] = data.json(exclude_none=True)
_budget_data = {k: v for k, v in _json_data.items() if k in budget_params}
budget_row = LiteLLM_BudgetTable(**_budget_data)
budget_row = LiteLLM_BudgetTable.model_validate(_budget_data)
new_budget = prisma_client.jsonify_object(budget_row.json(exclude_none=True))
_budget = await prisma_client.db.litellm_budgettable.create(
_budget: prisma_models.LiteLLM_BudgetTable = await prisma_client.db.litellm_budgettable.create(
data={
**new_budget,
"created_by": user_id or litellm_proxy_admin_name,
@ -214,8 +206,8 @@ async def _create_budget_for_project(
async def _set_project_object_permission(
data: NewProjectRequest,
prisma_client: Optional[PrismaClient],
) -> Optional[str]:
prisma_client: PrismaClient | None,
) -> str | None:
"""
Creates the LiteLLM_ObjectPermissionTable record for the project.
Returns the object_permission_id if created, otherwise None.
@ -224,7 +216,7 @@ async def _set_project_object_permission(
return None
if data.object_permission is not None:
created_object_permission = (
created_object_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
await prisma_client.db.litellm_objectpermissiontable.create(
data=data.object_permission.model_dump(exclude_none=True),
)
@ -344,8 +336,7 @@ async def new_project(
raise HTTPException(
status_code=403,
detail={
"error": "Only premium users can add tags to projects. "
+ CommonProxyErrors.not_premium_user.value
"error": "Only premium users can add tags to projects. " + CommonProxyErrors.not_premium_user.value
},
)
@ -353,8 +344,7 @@ async def new_project(
raise HTTPException(
status_code=403,
detail={
"error": "Project management is an enterprise feature. "
+ CommonProxyErrors.not_premium_user.value
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
},
)
@ -375,13 +365,11 @@ async def new_project(
)
# Validate team exists and get team object with budget
team_object = await _validate_team_exists(
team_id=data.team_id, prisma_client=prisma_client
)
team_object = await _validate_team_exists(team_id=data.team_id, prisma_client=prisma_client)
# Validate project limits against team limits
_check_team_project_limits(
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
team_object=LiteLLM_TeamTable.model_validate(team_object.model_dump()),
data=data,
)
@ -391,7 +379,7 @@ async def new_project(
user_api_key_dict=user_api_key_dict,
team_id=data.team_id,
prisma_client=prisma_client,
team_object=LiteLLM_TeamTable(**team_object.model_dump()),
team_object=LiteLLM_TeamTable.model_validate(team_object.model_dump()),
)
if not has_permission:
@ -449,17 +437,13 @@ async def new_project(
value=getattr(data, field),
)
new_project_row = prisma_client.jsonify_object(
project_row.json(exclude_none=True)
)
new_project_row = prisma_client.jsonify_object(project_row.json(exclude_none=True))
# Remove budget fields (following organization_endpoints.py pattern)
new_project_row = _remove_budget_fields_from_project_data(new_project_row)
verbose_proxy_logger.info(
f"new_project_row: {json.dumps(new_project_row, indent=2)}"
)
response = await prisma_client.db.litellm_projecttable.create(
verbose_proxy_logger.info(f"new_project_row: {json.dumps(new_project_row, indent=2)}")
response: prisma_models.LiteLLM_ProjectTable = await prisma_client.db.litellm_projecttable.create(
data={
**new_project_row, # type: ignore
},
@ -469,9 +453,7 @@ async def new_project(
return response
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(
str(e)
)
"litellm.proxy.management_endpoints.project_endpoints.new_project(): Exception occured - {}".format(str(e))
)
raise handle_exception_on_proxy(e)
@ -539,8 +521,7 @@ async def update_project(
raise HTTPException(
status_code=403,
detail={
"error": "Only premium users can add tags to projects. "
+ CommonProxyErrors.not_premium_user.value
"error": "Only premium users can add tags to projects. " + CommonProxyErrors.not_premium_user.value
},
)
@ -548,8 +529,7 @@ async def update_project(
raise HTTPException(
status_code=403,
detail={
"error": "Project management is an enterprise feature. "
+ CommonProxyErrors.not_premium_user.value
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
},
)
@ -576,9 +556,9 @@ async def update_project(
)
# Fetch existing project
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
where={"project_id": data.project_id}
)
existing_project: (
prisma_models.LiteLLM_ProjectTable | None
) = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": data.project_id})
if existing_project is None:
raise ProxyException(
@ -595,9 +575,7 @@ async def update_project(
target_team_id = data.team_id or existing_project.team_id
target_team_obj = None
if target_team_id is not None:
target_team_obj = await _validate_team_exists(
team_id=target_team_id, prisma_client=prisma_client
)
target_team_obj = await _validate_team_exists(team_id=target_team_id, prisma_client=prisma_client)
has_permission = await _check_user_permission_for_project(
user_api_key_dict=user_api_key_dict,
@ -620,32 +598,26 @@ async def update_project(
team_id=data.team_id,
prisma_client=prisma_client,
team_object=(
LiteLLM_TeamTable(**target_team_obj.model_dump())
if target_team_obj
else None
LiteLLM_TeamTable.model_validate(target_team_obj.model_dump()) if target_team_obj else None
),
)
if not can_assign_to_target:
raise HTTPException(
status_code=403,
detail={
"error": "Cannot reassign project to a team you are not an admin of"
},
detail={"error": "Cannot reassign project to a team you are not an admin of"},
)
# Validate project limits against team limits
if target_team_obj is not None:
_check_team_project_limits(
team_object=LiteLLM_TeamTable(**target_team_obj.model_dump()),
team_object=LiteLLM_TeamTable.model_validate(target_team_obj.model_dump()),
data=data,
)
# Prepare update data
update_data = data.json(exclude_none=True, exclude={"project_id"})
update_data = prisma_client.jsonify_object(update_data)
update_data["updated_by"] = (
user_api_key_dict.user_id or litellm_proxy_admin_name
)
update_data["updated_by"] = user_api_key_dict.user_id or litellm_proxy_admin_name
# Handle budget updates
budget_fields = LiteLLM_BudgetTable.model_fields.keys()
@ -671,21 +643,17 @@ async def update_project(
if existing_project.object_permission_id:
# Update existing permission
await prisma_client.db.litellm_objectpermissiontable.update(
where={
"object_permission_id": existing_project.object_permission_id
},
where={"object_permission_id": existing_project.object_permission_id},
data=object_permission_data,
)
else:
# Create new permission
created_permission = (
created_permission: prisma_models.LiteLLM_ObjectPermissionTable = (
await prisma_client.db.litellm_objectpermissiontable.create(
data=object_permission_data,
)
)
update_data["object_permission_id"] = (
created_permission.object_permission_id
)
update_data["object_permission_id"] = created_permission.object_permission_id
# Handle metadata fields
for field in LiteLLM_ManagementEndpoint_MetadataFields:
@ -698,7 +666,7 @@ async def update_project(
update_data = _remove_budget_fields_from_project_data(update_data)
# Update project
updated_project = await prisma_client.db.litellm_projecttable.update(
updated_project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.update(
where={"project_id": data.project_id},
data=update_data,
include={"litellm_budget_table": True, "object_permission": True},
@ -718,7 +686,7 @@ async def update_project(
"/project/delete",
tags=["project management"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[LiteLLM_ProjectTable],
response_model=list[LiteLLM_ProjectTable],
)
@management_endpoint_wrapper
async def delete_project(
@ -749,8 +717,7 @@ async def delete_project(
raise HTTPException(
status_code=403,
detail={
"error": "Project management is an enterprise feature. "
+ CommonProxyErrors.not_premium_user.value
"error": "Project management is an enterprise feature. " + CommonProxyErrors.not_premium_user.value
},
)
@ -778,9 +745,7 @@ async def delete_project(
for project_id in data.project_ids:
# Check if project exists
existing_project = await prisma_client.db.litellm_projecttable.find_unique(
where={"project_id": project_id}
)
existing_project = await prisma_client.db.litellm_projecttable.find_unique(where={"project_id": project_id})
if existing_project is None:
raise ProxyException(
@ -791,11 +756,9 @@ async def delete_project(
)
# Check if there are any keys associated with this project
associated_keys = (
await prisma_client.db.litellm_verificationtoken.find_many(
where={"project_id": project_id}
)
)
associated_keys: Sequence[
prisma_models.LiteLLM_VerificationToken
] = await prisma_client.db.litellm_verificationtoken.find_many(where={"project_id": project_id})
if len(associated_keys) > 0:
raise ProxyException(
@ -806,9 +769,9 @@ async def delete_project(
)
# Delete the project
deleted_project = await prisma_client.db.litellm_projecttable.delete(
where={"project_id": project_id}
)
deleted_project: (
prisma_models.LiteLLM_ProjectTable | None
) = await prisma_client.db.litellm_projecttable.delete(where={"project_id": project_id})
deleted_projects.append(deleted_project)
@ -854,7 +817,7 @@ async def project_info(
)
# Fetch project
project = await prisma_client.db.litellm_projecttable.find_unique(
project: prisma_models.LiteLLM_ProjectTable | None = await prisma_client.db.litellm_projecttable.find_unique(
where={"project_id": project_id},
include={"litellm_budget_table": True, "object_permission": True},
)
@ -872,17 +835,11 @@ async def project_info(
is_team_member = False
if project.team_id and user_api_key_dict.user_id:
team = await prisma_client.db.litellm_teamtable.find_unique(
where={"team_id": project.team_id}
)
team = await _team_table(prisma_client).find_unique(where={"team_id": project.team_id})
if team:
caller_user_id = user_api_key_dict.user_id
for m in team.members_with_roles or []:
m_user_id = (
m.get("user_id")
if isinstance(m, dict)
else getattr(m, "user_id", None)
)
m_user_id = m.get("user_id") if isinstance(m, dict) else getattr(m, "user_id", None)
if m_user_id == caller_user_id:
is_team_member = True
break
@ -896,9 +853,7 @@ async def project_info(
return project
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(
str(e)
)
"litellm.proxy.management_endpoints.project_endpoints.project_info(): Exception occured - {}".format(str(e))
)
raise handle_exception_on_proxy(e)
@ -907,7 +862,7 @@ async def project_info(
"/project/list",
tags=["project management"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[LiteLLM_ProjectTable],
response_model=list[LiteLLM_ProjectTable],
)
async def list_projects(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
@ -932,21 +887,19 @@ async def list_projects(
# If proxy admin, get all projects
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
projects = await prisma_client.db.litellm_projecttable.find_many(
projects: Sequence[
prisma_models.LiteLLM_ProjectTable
] = await prisma_client.db.litellm_projecttable.find_many(
include={"litellm_budget_table": True, "object_permission": True}
)
else:
# Look up the user's team memberships via the reverse-index on
# LiteLLM_UserTable.teams (maintained by team_member_add alongside
# members_with_roles). This avoids a full scan of all team rows.
user_record = await prisma_client.db.litellm_usertable.find_unique(
user_record: prisma_models.LiteLLM_UserTable | None = await prisma_client.db.litellm_usertable.find_unique(
where={"user_id": user_api_key_dict.user_id},
)
user_team_ids = (
user_record.teams
if user_record is not None and user_record.teams
else []
)
user_team_ids: Sequence[str] = user_record.teams if user_record is not None and user_record.teams else []
projects = await prisma_client.db.litellm_projecttable.find_many(
where={"team_id": {"in": user_team_ids}},

View file

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

View file

@ -54,6 +54,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/messages",
"/v1/skills",
"/v1/a2a/",
"/a2a/",
# LiteLLM-native LLM surface
"/v1/rerank",
"/v2/rerank",

View file

@ -5,5 +5,5 @@ dependencies:
- name: redis
repository: oci://registry-1.docker.io/bitnamicharts
version: 18.19.1
digest: sha256:8660fe6287f9941d08c0902f3f13731079b8cecd2a5da2fbc54e5b7aae4a6f62
generated: "2024-03-10T02:28:52.275022+05:30"
digest: sha256:38962e231f6596b93f82a8412bbe4cf5de696caecf5775dfbbd163383eb1c009
generated: "2026-07-28T10:21:22.511401-07:00"

View file

@ -18,7 +18,7 @@ type: application
# This is the chart version. This version number should be incremented each time you make changes
# to the chart and its templates, including the app version.
# Versions are expected to follow Semantic Versioning (https://semver.org/)
version: 1.1.0
version: 1.1.1
# This is the version number of the application being deployed. This version number should be
# incremented each time you make changes to the application. Versions are not expected to
@ -32,10 +32,10 @@ annotations:
dependencies:
- name: "postgresql"
version: ">=13.3.0"
version: "14.3.1"
repository: oci://registry-1.docker.io/bitnamicharts
condition: db.deployStandalone
- name: redis
version: ">=18.0.0"
version: "18.19.1"
repository: oci://registry-1.docker.io/bitnamicharts
condition: redis.enabled

View file

@ -130,6 +130,16 @@ Set `billingMetrics.caSecretName` only when the collector is a private or test o
| `db.deployStandalone` | Deploy a standalone, single instance deployment of Postgres, using the Bitnami postgresql chart. This is useful for getting started but doesn't provide HA or (by default) data backups. | `true` |
| `postgresql.*` | If `db.deployStandalone` is `true`, configuration passed to the Bitnami postgresql chart. See the [Bitnami Documentation](https://github.com/bitnami/charts/tree/main/bitnami/postgresql) for full configuration details. See [values.yaml](./values.yaml) for the default configuration. | See [values.yaml](./values.yaml) |
| `postgresql.auth.*` | If `db.deployStandalone` is `true`, care should be taken to ensure the default `password` and `postgres-password` values are **NOT** used. | `NoTaGrEaTpAsSwOrD` |
| `postgresql.image.*` | If `db.deployStandalone` is `true`, the image for the bundled Postgres. Pinned to a `docker.io/bitnamilegacy` build because Bitnami retired the versioned tags under `docker.io/bitnami`. | `bitnamilegacy/postgresql:16.2.0-debian-12-r6` |
| `redis.image.*` | If `redis.enabled` is `true`, the image for the bundled Redis. Pinned to a `docker.io/bitnamilegacy` build for the same reason. | `bitnamilegacy/redis:7.2.4-debian-12-r9` |
#### Bundled Postgres image
Bitnami removed the versioned tags from `docker.io/bitnami` and republished the archived builds under `docker.io/bitnamilegacy`, so the image defaults that ship inside the `postgresql` and `redis` subcharts no longer pull. The chart pins both to the `bitnamilegacy` copies of the exact builds those subchart versions were released with, which keeps the on-disk data directory layout unchanged for existing installs.
Keep `postgresql.image.tag` pinned. `docker.io/bitnami/postgresql` still publishes a floating `latest`, and pointing the bundled Postgres at a different major version starts the server against a data directory it cannot read (`database files are incompatible with server`). There is no in-place way back, so crossing a major version means dumping the database with the old image and restoring it into the new one. The chart refuses to render when the tag is empty or `latest`.
Those images no longer receive updates. For anything beyond getting started, run Postgres outside the chart and point at it with `db.useExisting`.
#### Example Postgres `db.useExisting` Secret

View file

@ -146,3 +146,18 @@ Get redis service port
{{ .Values.redis.master.service.ports.redis }}
{{- end -}}
{{- end -}}
{{/*
Reject an unpinned image tag for the bundled PostgreSQL.
A floating tag lets a chart upgrade start a newer PostgreSQL major against the
existing PersistentVolumeClaim. The server then refuses to start on a data
directory written by another major version, and the only way back is a dump
taken before the change, which by that point no longer exists.
*/}}
{{- define "litellm.validateBundledPostgresImageTag" -}}
{{- $tag := .Values.postgresql.image.tag | default "" | toString -}}
{{- $digest := .Values.postgresql.image.digest | default "" | toString -}}
{{- if and (eq $digest "") (or (eq $tag "") (eq $tag "latest")) -}}
{{- fail (printf "postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got %q). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore." $tag) -}}
{{- end -}}
{{- end -}}

View file

@ -1,4 +1,5 @@
{{- if .Values.db.deployStandalone -}}
{{- include "litellm.validateBundledPostgresImageTag" . -}}
apiVersion: v1
kind: Secret
metadata:

View file

@ -10,7 +10,7 @@ metadata:
spec:
containers:
- name: test
image: bitnami/kubectl:latest
image: docker.io/bitnamilegacy/kubectl:1.29.2-debian-12-r3
command: ['sh', '-c']
args:
- |

View file

@ -0,0 +1,94 @@
suite: test bundled database images
templates:
- charts/postgresql/templates/primary/statefulset.yaml
- charts/redis/templates/master/application.yaml
- charts/redis/templates/configmap.yaml
- charts/redis/templates/health-configmap.yaml
- charts/redis/templates/scripts-configmap.yaml
- charts/redis/templates/secret.yaml
- secret-dbcredentials.yaml
- templates/tests/test-servicemonitor.yaml
tests:
- it: should pull the bundled postgres from a repository that still publishes the pinned tag
template: charts/postgresql/templates/primary/statefulset.yaml
set:
db.deployStandalone: true
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: docker.io/bitnamilegacy/postgresql:16.2.0-debian-12-r6
- it: should pull the bundled postgres metrics exporter from the same repository
template: charts/postgresql/templates/primary/statefulset.yaml
set:
db.deployStandalone: true
postgresql.metrics.enabled: true
asserts:
- equal:
path: spec.template.spec.containers[1].image
value: docker.io/bitnamilegacy/postgres-exporter:0.15.0-debian-12-r14
- it: should run the bundled postgres init container from the same repository
template: charts/postgresql/templates/primary/statefulset.yaml
set:
db.deployStandalone: true
postgresql.volumePermissions.enabled: true
asserts:
- equal:
path: spec.template.spec.initContainers[0].image
value: docker.io/bitnamilegacy/os-shell:12-debian-12-r16
- it: should pull the bundled redis from a repository that still publishes the pinned tag
template: charts/redis/templates/master/application.yaml
set:
redis.enabled: true
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: docker.io/bitnamilegacy/redis:7.2.4-debian-12-r9
- it: should reject a floating postgres tag that could cross a major version on an existing volume
template: secret-dbcredentials.yaml
set:
db.deployStandalone: true
postgresql.image.tag: latest
asserts:
- failedTemplate:
errorMessage: 'postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got "latest"). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore.'
- it: should reject an empty postgres tag
template: secret-dbcredentials.yaml
set:
db.deployStandalone: true
postgresql.image.tag: ""
asserts:
- failedTemplate:
errorMessage: 'postgresql.image.tag must be pinned to an explicit version when db.deployStandalone is true (got ""). An unpinned tag can start a different PostgreSQL major against the existing data directory, which makes the database unreadable and is not recoverable in place. Crossing a major version requires a dump and restore.'
- it: should accept an empty postgres tag when the image is pinned by digest
template: secret-dbcredentials.yaml
set:
db.deployStandalone: true
postgresql.image.tag: ""
postgresql.image.digest: sha256:0d0e2f1a5b3c4d6e7f8091a2b3c4d5e6f708192a3b4c5d6e7f8091a2b3c4d5e6
asserts:
- hasDocuments:
count: 1
- it: should run the servicemonitor test pod from a pinned image
template: templates/tests/test-servicemonitor.yaml
set:
serviceMonitor.enabled: true
asserts:
- equal:
path: spec.containers[0].image
value: docker.io/bitnamilegacy/kubectl:1.29.2-debian-12-r3
- it: should not constrain the postgres tag when the bundled database is not deployed
template: secret-dbcredentials.yaml
set:
db.deployStandalone: false
postgresql.image.tag: latest
asserts:
- hasDocuments:
count: 0

View file

@ -328,8 +328,32 @@ lifecycle: {}
# Settings for Bitnami postgresql chart (if db.deployStandalone is true, ignored
# otherwise)
#
# Bitnami retired the versioned tags under docker.io/bitnami and republished the
# archived builds under docker.io/bitnamilegacy, so the subchart's own image
# defaults no longer resolve. The repository below points at the same build the
# subchart was released with, which keeps the on-disk data directory layout
# identical for existing installs.
#
# Keep the tag pinned. docker.io/bitnami still publishes a floating `latest`,
# and starting a newer PostgreSQL major against an existing data directory
# leaves the server refusing to boot ("database files are incompatible with
# server") with no way back other than a dump taken beforehand. Crossing a major
# version is a dump-and-restore, not an image bump. The chart refuses to render
# an unpinned tag for this reason
postgresql:
architecture: standalone
image:
repository: bitnamilegacy/postgresql
tag: 16.2.0-debian-12-r6
volumePermissions:
image:
repository: bitnamilegacy/os-shell
tag: 12-debian-12-r16
metrics:
image:
repository: bitnamilegacy/postgres-exporter
tag: 0.15.0-debian-12-r14
auth:
username: litellm
database: litellm
@ -359,9 +383,36 @@ postgresql:
# When `redis.sentinel.enabled` is set, the coordination block is rendered with
# `sentinel_nodes` and `service_name` (from `redis.sentinel.masterSet`) instead
# of host/port, because a plain Redis client cannot talk to the sentinel port
#
# The image repositories carry the same bitnamilegacy repoint as postgresql
# above; the versioned tags the subchart ships with are gone from
# docker.io/bitnami
redis:
enabled: false
architecture: standalone
image:
repository: bitnamilegacy/redis
tag: 7.2.4-debian-12-r9
sentinel:
image:
repository: bitnamilegacy/redis-sentinel
tag: 7.2.4-debian-12-r7
metrics:
image:
repository: bitnamilegacy/redis-exporter
tag: 1.58.0-debian-12-r4
volumePermissions:
image:
repository: bitnamilegacy/os-shell
tag: 12-debian-12-r16
sysctl:
image:
repository: bitnamilegacy/os-shell
tag: 12-debian-12-r16
kubectl:
image:
repository: bitnamilegacy/kubectl
tag: 1.29.2-debian-12-r3
coordination:
# Set to false to keep the bundled Redis for response caching only and leave
# `general_settings.coordination_redis` out of the rendered config. A

View file

@ -19,7 +19,7 @@
"/v1/fine-tuning" "/fine-tuning" "/v1/responses" "/responses" "/v1/threads" "/threads"
"/v1/assistants" "/assistants" "/v1/vector_stores" "/vector_stores" "/v1/indexes"
"/v1/models" "/models" "/openai" "/engines"
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a"
"/v1/messages" "/messages" "/v1/skills" "/v1/a2a" "/a2a"
"/v1/rerank" "/v2/rerank" "/rerank" "/v1/ocr" "/ocr" "/v1/rag" "/rag"
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"

View file

@ -0,0 +1,2 @@
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogToolIndex_start_time_idx" ON "LiteLLM_SpendLogToolIndex"("start_time");

View file

@ -0,0 +1,12 @@
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_DailyToolSpend" (
"date" TEXT NOT NULL,
"tool_name" TEXT NOT NULL,
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"total_tokens" BIGINT NOT NULL DEFAULT 0,
"request_count" BIGINT NOT NULL DEFAULT 0,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_DailyToolSpend_pkey" PRIMARY KEY ("date","tool_name")
);

View file

@ -0,0 +1,106 @@
"""Optional post-migration step that raises Postgres REPLICA IDENTITY to FULL.
Logical-replication consumers (Neon / lakehouse sync and similar) need FULL
replica identity to reconstruct the old row of an UPDATE or DELETE. Prisma
leaves every table it creates at the Postgres default, so the setting has to be
re-applied by hand after each migration run. Setting
``LITELLM_SET_REPLICA_IDENTITY_FULL`` makes every migration run re-assert it.
The statement goes through the Prisma CLI rather than a Postgres driver because
``litellm-proxy-extras`` has no runtime dependencies, while the CLI is already
required for the migrations themselves.
"""
import subprocess
import tempfile
from pathlib import Path
from litellm_proxy_extras._logging import logger
REPLICA_IDENTITY_FULL_ENV_VAR = "LITELLM_SET_REPLICA_IDENTITY_FULL"
REPLICA_IDENTITY_FULL_SQL = r"""
DO $$
DECLARE
target regclass;
BEGIN
SET LOCAL lock_timeout = '5s';
FOR target IN
SELECT c.oid::regclass
FROM pg_class c
JOIN pg_namespace n ON n.oid = c.relnamespace
WHERE c.relkind = 'r'
AND c.relreplident <> 'f'
AND n.nspname = ANY (current_schemas(false))
AND c.relname LIKE 'LiteLLM\_%'
LOOP
BEGIN
EXECUTE format('ALTER TABLE %s REPLICA IDENTITY FULL', target);
EXCEPTION WHEN lock_not_available THEN
RAISE WARNING 'REPLICA IDENTITY FULL skipped for %: table busy, retrying next run', target;
END;
END LOOP;
END
$$;
"""
def apply_replica_identity_full(
schema_path: str,
prisma_command: str,
prisma_env: dict[str, str],
) -> bool:
"""Set REPLICA IDENTITY FULL on every LiteLLM table that is not already FULL.
Never raises. Replication metadata is not needed to serve requests, so
every failure mode is reported and stepped over rather than taking down a
migration run that already succeeded: a database that refuses the ALTER
(most often because the runtime user does not own the tables), a missing
or unrunnable Prisma CLI, a read-only temp directory, or a timeout.
Returns True when the statement was applied, False when it failed.
"""
logger.info("Applying REPLICA IDENTITY FULL to LiteLLM tables")
try:
with tempfile.TemporaryDirectory(prefix="litellm_replica_identity_") as tmp_dir:
sql_path = Path(tmp_dir) / "replica_identity_full.sql"
sql_path.write_text(REPLICA_IDENTITY_FULL_SQL)
subprocess.run(
[
prisma_command,
"db",
"execute",
"--file",
str(sql_path),
"--schema",
schema_path,
],
timeout=60,
check=True,
capture_output=True,
text=True,
env=prisma_env,
)
except subprocess.CalledProcessError as e:
logger.error(
"Failed to set REPLICA IDENTITY FULL. Logical replication "
"consumers may reject updates to these tables. Grant table "
"ownership to the migration user, or apply "
"`ALTER TABLE ... REPLICA IDENTITY FULL` by hand. Error: %s",
e.stderr,
)
return False
except subprocess.TimeoutExpired:
logger.error("Timed out setting REPLICA IDENTITY FULL on LiteLLM tables")
return False
except OSError as e:
logger.error(
"Could not run the REPLICA IDENTITY FULL statement. Logical "
"replication consumers may reject updates to these tables. "
"Error: %s",
e,
)
return False
logger.info("REPLICA IDENTITY FULL applied to LiteLLM tables")
return True

View file

@ -1094,6 +1094,20 @@ model LiteLLM_SpendLogToolIndex {
@@id([request_id, tool_name])
@@index([tool_name, start_time])
@@index([start_time])
}
// Daily tool spend rollup (one row per tool per day) – the Cost Optimization card reads this, never SpendLogs
model LiteLLM_DailyToolSpend {
date String
tool_name String
spend Float @default(0.0)
total_tokens BigInt @default(0)
request_count BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@id([date, tool_name])
}
// Prompt table for storing prompt configurations

View file

@ -10,6 +10,10 @@ from pathlib import Path
from typing import Optional
from litellm_proxy_extras._logging import logger
from litellm_proxy_extras.replica_identity import (
REPLICA_IDENTITY_FULL_ENV_VAR,
apply_replica_identity_full,
)
def str_to_bool(value: Optional[str]) -> bool:
@ -676,6 +680,39 @@ class ProxyExtrasDBManager:
finally:
os.chdir(original_dir)
@staticmethod
def apply_replica_identity_full_if_requested() -> bool:
"""
Re-assert REPLICA IDENTITY FULL on LiteLLM's tables when the operator
opted in via LITELLM_SET_REPLICA_IDENTITY_FULL.
Prisma leaves new tables at the Postgres default, which logical
replication consumers reject, so the setting has to be re-applied after
every migration run rather than once by hand.
Returns:
bool: True if the setting was applied, False if it was not
requested or could not be applied.
"""
if not str_to_bool(os.getenv(REPLICA_IDENTITY_FULL_ENV_VAR)):
return False
try:
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
prisma_command = _get_prisma_command()
prisma_env = _get_prisma_env()
except OSError as e:
logger.error(
"Could not resolve the migrations directory for the REPLICA "
"IDENTITY FULL step, skipping it. Error: %s",
e,
)
return False
return apply_replica_identity_full(
schema_path=schema_path,
prisma_command=prisma_command,
prisma_env=prisma_env,
)
@staticmethod
def setup_database(
use_migrate: bool = False, use_v2_resolver: bool = False
@ -694,6 +731,15 @@ class ProxyExtrasDBManager:
Returns:
bool: True if setup was successful, False otherwise
"""
migrated = ProxyExtrasDBManager._run_migrations(
use_migrate=use_migrate, use_v2_resolver=use_v2_resolver
)
if migrated:
ProxyExtrasDBManager.apply_replica_identity_full_if_requested()
return migrated
@staticmethod
def _run_migrations(use_migrate: bool, use_v2_resolver: bool) -> bool:
if use_v2_resolver:
logger.info("Using v2 migration resolver (--use_v2_migration_resolver)")
return ProxyExtrasDBManager._setup_database_v2(use_migrate=use_migrate)

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.80"
version = "0.4.81"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.80"
version = "0.4.81"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -1,10 +1,11 @@
# Adding a provider / route to litellm-rust
Three layers, same for every route (see `ocr` and `realtime` as references):
Everything for a route lives in `crates/core/src/<route>/`; `crates/core/src/messages` is the reference. A host (the axum gateway, the Python bridge) only calls the route's entrypoint.
1. **Transform contract (pure)** — `crates/core/src/<route>/transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) + types in `types.rs`. No network, env, or auth.
2. **Provider config (pure)** — `crates/providers/src/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
3. **HTTP / transport (the host)** — `crates/providers/src/<route>.rs` (e.g. `ocr.rs`, `realtime.rs`): the callable fn (`run_ocr`, `realtime`). It resolves the key, builds the auth header, builds URL + transforms via the config, then does the network call. This is the only layer allowed to do I/O.
1. **Entrypoint** — `mod.rs`: `pub async fn <route>(request) -> CoreResult<Response>`, the Rust equivalent of `litellm.<route>()`, plus a `<route>_stream` variant when the route streams. It is the only thing a host touches.
2. **Transform contract** — `transformation.rs`: a `…ProviderConfig` trait (URL build + request/response transforms) with types in `types.rs`.
3. **Provider config** — `crates/core/src/providers/<provider>/<route>/transformation.rs`: implement that trait as a `const <PROVIDER>_<ROUTE>_CONFIG`, mirroring the Python provider tree. Add parity unit tests.
4. **Prepare + handler** — `prepare.rs` resolves provider/model, credentials, auth headers, and URL, then transforms the request; `handler.rs` performs the provider call through the shared client in `client.rs` and transforms the response.
## Coding standards
@ -25,4 +26,4 @@ variants of it. The test for a good abstraction is that adding the next provider
is a few declarative lines, not a new file of duplicated flow. Only diverge from
the base when behavior is genuinely different, and say so explicitly in the PR.
**Calling:** the host invokes the route fn — the Python bridge calls `run_ocr`; the `ai-gateway` server calls `realtime`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.
**Calling:** hosts invoke the core entrypoint — the Python bridge and the `ai-gateway` route service both call `litellm_core::messages::messages`. Never add a provider handler to `ai-gateway`. Register new modules in `lib.rs` / `mod.rs`, then run `cargo fmt && cargo clippy --workspace -- -D warnings && cargo test --workspace`.

View file

@ -4,14 +4,30 @@ litellm-rust has exactly THREE crates. A crate is a LAYER, not a route. Routes (
## Crates
| Crate | Role | Pure / I/O |
|-------|------|------------|
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust. One public entrypoint per top-level call (`messages::messages()`), owning types, transforms, provider resolution, auth, and the provider HTTP call. Call it, get a typed response. |
| litellm-ai-gateway | The axum server (behind the `server` feature) plus the WebSocket hosts. Translates HTTP/WS to core entrypoints; owns no provider logic and no handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
## Where a route lives
A top-level LiteLLM call is a module under `crates/core/src/<route>/`, shaped like `messages`:
```
core/src/messages/
mod.rs # pub async fn messages(..) -> CoreResult<..> (+ messages_stream for SSE)
types.rs # request/response types, MessagesRequest
transformation.rs # the provider template trait
prepare.rs # provider resolution, auth headers, URL
handler.rs # the provider call
client.rs # the shared reqwest client
```
Handlers never live in `ai-gateway`. `ocr`, `audio_transcription`, and `realtime` are still hosted there from before this rule; they move to `core` as they are touched.
Adding a crate: default to a MODULE. New crate ONLY on a real trigger — separate artifact (binary/cdylib), proc-macro, shared foundation, or publishable standalone. A new provider or route is none of these.
Adding a crate fails crates/core/tests/workspace_crate_allowlist.rs until you update its allowlist and this file — intentional.

View file

@ -23,21 +23,34 @@ the base when behavior is genuinely different, and say so explicitly in the PR.
## Crates (exactly three — see AGENTS.md)
`litellm-core` describes work; `litellm-ai-gateway` executes it; `litellm-python-bridge`
exposes it to the Python SDK. A crate is a **layer**, not a route — add modules, not crates.
`litellm-core` **is** the LiteLLM SDK in Rust: it makes the LLM call.
`litellm-ai-gateway` is an HTTP/WebSocket server in front of it, and
`litellm-python-bridge` exposes it to the Python SDK. A crate is a **layer**, not
a route — add modules, not crates.
## Core Boundary
`litellm-core` is the pure translation layer; the `litellm-ai-gateway` host executes work.
`litellm-core` owns the whole call. The Rust equivalent of `litellm.messages()`
is `litellm_core::messages::messages(request).await`: you call it, it does the
provider call, and you get a typed non-streaming response back.
Route-level Rust structure mirrors LiteLLM's Python responsibilities:
- `core/src/<route>/` owns the route contract, shared types, and provider
template traits. For OCR, this means `core/src/ocr`.
- `core/src/<route>/` owns the route end to end: the public entrypoint fn named
after the route in `mod.rs`, the request/response types (`types.rs`), the
provider template trait (`transformation.rs`), the provider/auth/URL
resolution (`prepare.rs`), the HTTP client (`client.rs`), and the handler that
performs the call (`handler.rs`). `core/src/messages` is the reference.
- `core/src/providers/<provider>/<route>/transformation.rs` owns the
provider-specific transform. For Mistral OCR, this means
`core/src/providers/mistral/ocr/transformation.rs`.
- Network execution lives in the host crate `ai-gateway` (`ai-gateway/src/io/`),
never inside `core`.
provider-specific transform. For Anthropic Messages, this means
`core/src/providers/anthropic/messages/transformation.rs`.
- Handlers live in `core`, never in a host. `ai-gateway` must not contain a
route handler that talks to a provider; its axum route reads the HTTP request,
picks a deployment, and calls the `core` entrypoint. `python-bridge` marshals
Python objects and calls the same entrypoint.
Streaming keeps the same shape: the route entrypoint has a `<route>_stream`
variant in `core` that returns the upstream response so a host can splice it to
its own caller; the host still owns no provider logic.
Call-hook and lifecycle instrumentation, including phase timing, usage
accumulation, and callback payload construction, always lives in `core`.
@ -45,21 +58,31 @@ Hosts feed observed events into core and dispatch the completed payloads through
their I/O logger; hosts must not own callback orchestration.
Allowed in `core`:
- Pure request transforms
- Pure response transforms
- Pure stream chunk normalization
- The public entrypoint for a top-level LiteLLM call
- Request/response transforms and stream chunk normalization
- Provider resolution, auth header construction, and URL building
- The provider HTTP call itself, through a shared reused client with connect and
request timeouts
- Shared data types and validation errors
- Deterministic token/cost helper logic
Not allowed in `core`:
- Network calls
- Environment variable or secret reads
- Serving HTTP: axum routes, extractors, and transport concerns stay in the host
- Filesystem access
- Database or cache access
- Provider SDK signing or auth flows
- Database access
- Config file reading and rollout state
- Logging callbacks, spend writes, or custom callbacks
- Global mutable runtime state
Env reads in `core` are limited to credential fallback inside a route's
`prepare.rs` (the `env_lookup` closure), mirroring what the Python SDK does when
no key is passed. Everything else config-shaped is resolved by the host and
passed in.
Routes still hosted in `ai-gateway` (`ocr`, `audio_transcription`, `realtime`)
predate this rule and are being moved into `core` route modules; do not add new
ones there, and prefer moving one when you touch it.
Python owns rollout state and fallback while Rust is being introduced. Rust
paths must be off by default until parity tests prove equivalence with Python.
A new provider/route may instead be implemented rust-only with no Python
@ -93,10 +116,10 @@ the first PR:
- Preserve Python output shape intentionally. If a field is always serialized as
`null` for Python parity, leave a short comment explaining that parity choice.
## Host I/O Rules
## Network I/O Rules
These rules apply when adding future crates or modules that execute network I/O,
such as `ai-gateway`, router hosts, or standalone servers:
These rules apply to every module that executes network I/O, whether it is a
`core` route handler or a host such as `ai-gateway`:
- Set connect and full-request timeouts. No unbounded waits.
- Reuse HTTP clients; do not construct clients per request.

View file

@ -2,18 +2,31 @@
This workspace contains the staged Rust implementation for LiteLLM.
Rust starts as a pure transform core used by the existing Python host. Python
continues to own auth, configuration, network I/O, retries, routing, logging,
`litellm-core` is the LiteLLM SDK in Rust: one entrypoint per top-level call
that makes the LLM call and hands back a typed response, the same shape as
`litellm.messages()` in Python.
```rust
let response = litellm_core::messages::messages(MessagesRequest {
model: "claude-sonnet-4-5",
body,
api_key: Some(key),
..
})
.await?;
```
Python continues to own configuration, retries, routing policy, logging,
callbacks, spend tracking, and customer plugins until each Rust path has parity
coverage and production evidence.
## Crates
| Crate | Role | Pure / I/O |
|-------|------|------------|
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under providers/), and the router. Builds requests/responses; no network. | Pure |
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under io/) plus the axum server binary (behind the `server` feature). | I/O |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
| Crate | Role |
|-------|------|
| litellm-core | The SDK. Per-route entrypoints (`messages::messages()`), types, provider transforms (modules under `providers/`), provider resolution, auth, the provider HTTP call, and the router. |
| litellm-ai-gateway | The axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.
@ -21,16 +34,16 @@ Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-
```text
crates/
core/ Route contracts, shared pure types, errors, and templates.
src/ocr/
providers/ Provider-specific pure transforms.
src/mistral/ocr/transformation.rs
core/ The SDK: route modules + provider transforms.
src/messages/ mod.rs (entrypoint), types, transformation, prepare, handler, client
src/providers/anthropic/messages/transformation.rs
ai-gateway/ Axum server + WebSocket hosts; calls core entrypoints.
python-bridge/ PyO3 bridge for Python LiteLLM.
```
The folder shape should follow the Python provider tree:
`providers/src/<provider>/<route>/transformation.rs`. The bridge should expose
one function per top-level route, starting with `ocr(payload)`.
The folder shape follows the Python provider tree:
`core/src/providers/<provider>/<route>/transformation.rs`. The bridge exposes one
function per top-level route, mirroring the core entrypoints.
## Checks

View file

@ -1,6 +1,6 @@
# Provider coding standards (litellm-rust)
Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MISTRAL_OCR_CONFIG`) is the reference; `messages` (`ANTHROPIC_MESSAGES_CONFIG`) is the next port.
Rules for adding or changing an LLM provider/route in `litellm-rust`. `messages` (`core/src/messages`, `ANTHROPIC_MESSAGES_CONFIG`) is the reference: a route is a `core` module with a public entrypoint that makes the call and returns a typed response.
## Provider resolution
@ -16,10 +16,10 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
## Boundaries
7. Layers never cross: `core` = pure transforms/types (no network, env, secrets, auth, logging, global mutable state); `ai-gateway` = all I/O, auth headers, HTTP/SSE, lifecycle hooks; `python-bridge` = thin PyO3 adapter.
7. Layers never cross: `core` = the call itself (entrypoint, types, transforms, provider resolution, auth headers, provider HTTP, lifecycle hooks); `ai-gateway` = serving HTTP/WS (routing, extractors, auth of *our* callers, streaming to the client); `python-bridge` = thin PyO3 adapter. Hosts call the core entrypoint; they never build a provider request.
8. Generic/route files contain zero provider-specific branches. A provider is one module under `core/src/providers/<provider>/<route>/`; a route is a module, never a new crate.
9. Route entry point stays thin: `<route>()` -> `prepare_*` -> `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing. Handlers validate and delegate; no business logic in them.
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Env reads happen only at the host/config layer, with the `DEFAULT_*` fallback defined in `constants.rs`.
9. Route entry point stays thin: `core::<route>::<route>()` -> `prepare_*` -> handler (or `CallLifecycle::run_request`, which owns the pre_call -> during_call -> provider call -> success/failure order and phase timing). Axum handlers validate and delegate to a service that calls the entrypoint; no business logic in them.
10. Constants (URLs, env-var names, API versions, error messages) live in a crate `constants.rs`, never inline. Config-shaped env reads happen at the host/config layer with the `DEFAULT_*` fallback defined in `constants.rs`; the only env read in `core` is the credential fallback in a route's `prepare.rs`.
## Types and errors
@ -33,7 +33,7 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
16. Never log request/response bodies, base64 payloads, document contents, or secrets. Truncate and bound any upstream body before it crosses a host boundary.
17. Treat empty/whitespace credentials, URLs, and config values as absent at the host resolution layer.
18. Host I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
18. Network I/O sets connect + request timeouts (no unbounded waits), reuses a shared HTTP client, and prefers rustls TLS.
## Tests and rollout

View file

@ -1,7 +1,9 @@
# ai-gateway — folder architecture
The Axum server that fronts the Rust gateway. It owns transport + config + auth
only; deployment selection lives in `core::router`, transforms in `core`/`providers`.
only; deployment selection lives in `core::router`, and the LLM call itself
(transforms, auth headers, provider HTTP) lives behind a `core` route entrypoint
such as `litellm_core::messages::messages`. No provider handler lives here.
```
src/
@ -32,6 +34,11 @@ src/
args; it runs during extraction. Never re-implement the check per route.
- **Handlers are thin.** A handler validates and delegates to its `service`. No
business logic, no provider calls, no transforms in handlers.
- **Services call `core`, they don't reimplement it.** A `service` picks the
deployment and calls the `core` route entrypoint. Provider resolution, auth
headers, URL building, and the HTTP call are `core`'s job; a service that
builds a provider request itself is a bug (`routes/messages/service.rs` is
the reference).
- **State is shared and cheap to clone.** Long-lived handles live behind `Arc` in
`state.rs`; read env/config only in `main.rs` when building state.

View file

@ -8,11 +8,11 @@ dials OpenAI upstream, and splices the two sockets frame-by-frame.
`litellm-rust` is exactly three crates (a crate is a **layer**, not a route):
| Crate | Role | Pure / I/O |
|-------|------|------------|
| litellm-core | Translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. Builds requests/responses; no network. | Pure |
| litellm-ai-gateway | Routes + host — the only crate that touches the network. HTTP/WebSocket I/O (modules under `io/`) plus the Axum server binary (behind the `server` feature). | I/O |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — a thin adapter over litellm-ai-gateway's I/O. | Binding |
| Crate | Role |
|-------|------|
| litellm-core | The LiteLLM SDK in Rust — per-route entrypoints (`messages::messages()`) that resolve the provider, transform, and make the call; plus types, provider transforms, and the router. |
| litellm-ai-gateway | The Axum server (behind the `server` feature) and WebSocket hosts. Translates HTTP/WS to core entrypoints; no provider handlers. |
| litellm-python-bridge | PyO3 cdylib exposing Rust to the litellm Python SDK — marshals Python objects and calls core entrypoints. |
Dependency direction (acyclic): litellm-core ← litellm-ai-gateway ← litellm-python-bridge.

View file

@ -29,18 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
/// Full-request timeout ceiling for Anthropic Messages provider calls, in
/// seconds. Mirrors the Python Anthropic Messages default. The per-request
/// timeout from `litellm_params` still overrides this on the request builder.
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for Anthropic Messages provider calls, in seconds.
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Max characters of an upstream error body echoed across the host boundary
/// before truncation, so provider bodies are bounded and data-minimized.
pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256;
pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
@ -48,10 +36,6 @@ pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
#[cfg(feature = "server")]
pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";
/// Provider name used by the Anthropic Messages route when a deployment's
/// provider model does not carry an explicit provider prefix.
pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
/// Request headers owned by the gateway and never forwarded upstream.
#[cfg(feature = "server")]
pub(crate) const MESSAGES_HEADERS_NOT_FORWARDED: &[&str] =

View file

@ -1 +0,0 @@
pub use crate::messages::{MessagesRequest, messages};

View file

@ -1,5 +1,4 @@
pub mod audio_transcription;
pub mod messages;
pub mod ocr;
pub mod realtime;
pub mod realtime_pool;

View file

@ -4,7 +4,9 @@
//! without pulling in the HTTP server:
//!
//! - Call-type modules such as [`ocr`]: provider transforms, lifecycle hooks,
//! and provider I/O. Always available — no feature required.
//! and provider I/O. Always available — no feature required. These predate the
//! rule that a route's entrypoint and handler live in `litellm-core` (see
//! `litellm_core::messages`) and move there as they are touched.
//! - [`io`]: compatibility exports and realtime WebSocket splice helpers.
//! - The server modules ([`auth`], [`routes`], [`state`]) and anything pulling
//! `axum` are gated behind the `server` feature, which the `litellm-ai-gateway`
@ -14,7 +16,6 @@
pub mod audio_transcription;
mod client;
pub mod io;
pub mod messages;
pub mod ocr;
/// GIL-activity tracking. Pure (atomics only); shared by the `server` routes and

View file

@ -1,49 +0,0 @@
use litellm_core::CoreResult;
use serde_json::Value;
mod client;
mod common_utils;
mod handler;
mod prepare;
mod types;
pub use types::MessagesRequest;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
match execute_messages(request, false).await? {
MessagesResponse::Json(body) => Ok(body),
MessagesResponse::Stream(response) => {
drop(response);
Err(litellm_core::CoreError::InvalidResponse(
"non-streaming messages execution returned a stream".to_string(),
))
}
}
}
pub(crate) enum MessagesResponse {
Json(Value),
Stream(reqwest::Response),
}
pub(crate) async fn execute_messages(
request: MessagesRequest<'_>,
stream: bool,
) -> CoreResult<MessagesResponse> {
let prepared = prepare_messages_call(request)?;
if stream {
execute_messages_provider_stream(prepared)
.await
.map(MessagesResponse::Stream)
} else {
execute_messages_provider_call(prepared)
.await
.map(MessagesResponse::Json)
}
}
#[cfg(test)]
mod tests;

View file

@ -1,24 +0,0 @@
use std::time::Duration;
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
use serde_json::{Map, Value};
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
pub(crate) struct ProviderMessagesRequest {
pub(crate) provider: String,
pub(crate) model: String,
pub(crate) config: &'static dyn AnthropicMessagesProviderConfig,
pub(crate) url: String,
pub(crate) body: Value,
pub(crate) upstream_headers: Vec<(String, String)>,
pub(crate) timeout: Option<Duration>,
}

View file

@ -19,7 +19,10 @@ async fn handle(...) -> impl IntoResponse { ... }
When a route has business logic worth testing without axum, put it in a sibling
`service` (a file, or a folder if the route grows). The route file stays the
**axum surface** (router + handler + any socket/SSE adapter); `service` is plain
Rust with **no axum types**. `realtime/` is the example:
Rust with **no axum types**, and its job is to pick the deployment and call the
`core` route entrypoint (see `messages/service.rs` calling
`litellm_core::messages::messages`). Never build a provider request, resolve a
key, or perform the provider call here. `realtime/` is the older example:
```
realtime/
mod.rs # axum surface: router() + handler + the WS<->events adapter
@ -33,6 +36,8 @@ genuinely gets hard to read.
`crate::auth::RequireMasterKey` to its arguments; it runs during extraction.
Never re-implement the check per route.
- **Handlers contain no business logic; `service` contains no axum types.**
- **No provider handlers in this crate.** Transforms, auth headers, and the
provider HTTP call live in `core/src/<route>/`.
- A route owns its paths in its own `router()`; `mod.rs` only merges.
- Cross-cutting concerns (logging, CORS, timeouts) → Tower layers in `mod.rs`,
not duplicated in handlers.

View file

@ -1,12 +1,12 @@
use std::sync::Arc;
use litellm_core::constants::ANTHROPIC_MESSAGES_PROVIDER;
use litellm_core::messages::types::MessagesRequest;
use litellm_core::messages::{messages, messages_stream};
use litellm_core::router::Router;
use litellm_core::{CoreError, CoreResult};
use serde_json::{Map, Value};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::messages::{MessagesRequest, execute_messages};
pub(crate) enum MessagesResponse {
Json(Value),
Stream(reqwest::Response),
@ -52,13 +52,14 @@ pub async fn run(
extra_headers,
timeout: None,
};
let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
execute_messages(request, stream)
.await
.map(|response| match response {
crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body),
crate::messages::MessagesResponse::Stream(upstream) => {
MessagesResponse::Stream(upstream)
}
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
return messages_stream(request).await.map(MessagesResponse::Stream);
}
let response = messages(request).await?;
serde_json::to_value(response)
.map(MessagesResponse::Json)
.map_err(|err| {
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
})
}

View file

@ -1,3 +1,7 @@
litellm-core is the PURE translation layer — types, route contracts (traits), provider transforms (modules under `providers/`), and the router. No network, no I/O, no env reads.
litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-level call is a module under `src/<route>/` exposing a public entrypoint named after the route (`messages::messages()`, the Rust equivalent of `litellm.messages()`): you call it and get a typed non-streaming response back.
Routes (ocr, realtime) and providers (mistral, openai) are modules, not crates.
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.

View file

@ -4,20 +4,28 @@ Rules for `litellm-rust/crates/core`.
## Responsibility
`core` owns shared data types, typed errors, and deterministic helper contracts.
It must stay pure and host-independent.
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
LiteLLM call has a public entrypoint here, named after the route
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
calling it returns a typed non-streaming response.
Allowed:
- The public entrypoint for a route, plus its `<route>_stream` variant when the
route supports streaming.
- Provider resolution, auth header construction, URL building, and the provider
HTTP call (shared reused client, connect + request timeouts).
- Shared request/response structs.
- Typed errors with stable, non-sensitive messages.
- Deterministic validation helpers.
- Serialization helpers that intentionally mirror Python output shape.
- Route templates that match Python base config responsibilities, such as
`ocr::transformation::OcrProviderConfig`.
`messages::transformation::AnthropicMessagesProviderConfig`.
Not allowed:
- Network, filesystem, database, cache, or environment access.
- Secret reads or auth/header construction.
- Serving HTTP: axum routers, extractors, and other transport concerns.
- Filesystem, database, or cache access.
- Config file reading or rollout state; the host resolves those and passes them
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
- Provider-specific branching that belongs in `providers`.
- Panics for user/provider-controlled input.
@ -33,10 +41,21 @@ typed field on a struct, not a raw string threaded through the API.
## Structure
Use route names directly under `src/`: `ocr`, future `messages`,
Use route names directly under `src/`: `messages`, `ocr`, future
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
invent broad names like `engine` for route contracts.
`src/messages` is the reference shape for a route module:
```
mod.rs pub async fn messages(..) (+ messages_stream)
types.rs request/response types
transformation.rs the provider template trait
prepare.rs provider resolution, auth headers, URL
handler.rs the provider call
client.rs the shared reqwest client
```
## Parity Rules
- Every shared type used by a provider transform needs unit tests for

View file

@ -7,6 +7,7 @@ repository.workspace = true
[dependencies]
rand.workspace = true
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
@ -30,5 +31,4 @@ bedrock-auth = [
]
[dev-dependencies]
reqwest.workspace = true
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }

View file

@ -1,3 +1,19 @@
pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com";
pub const OPENAI_RESPONSES_DEFAULT_API_BASE: &str = "https://api.openai.com/v1";
pub const OPENAI_RESPONSES_PATH: &str = "/responses";
/// Full-request timeout ceiling for Anthropic Messages provider calls, in
/// seconds. Mirrors the Python Anthropic Messages default. The per-request
/// timeout from the caller still overrides this on the request builder.
pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for Anthropic Messages provider calls, in seconds.
pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Max characters of an upstream error body echoed across the call boundary
/// before truncation, so provider bodies are bounded and data-minimized.
pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256;
/// Provider name used for Anthropic Messages when a deployment's provider model
/// does not carry an explicit provider prefix.
pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";

View file

@ -1,11 +1,11 @@
use litellm_core::CoreResult;
use litellm_core::error::{CoreError, json_type_name};
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
use serde_json::{Map, Value};
use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS;
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
use super::transformation::AnthropicMessagesProviderConfig;
pub(super) fn truncate_error_body(body: &str) -> String {
if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS {

View file

@ -1,15 +1,13 @@
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use serde_json::Value;
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::error::{CoreError, CoreResult};
use super::client::http_client;
use super::common_utils::truncate_error_body;
use super::types::ProviderMessagesRequest;
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use super::types::{AnthropicMessagesResponse, ProviderMessagesRequest};
pub(super) async fn execute_messages_provider_call(
request: ProviderMessagesRequest,
) -> CoreResult<Value> {
) -> CoreResult<AnthropicMessagesResponse> {
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -39,12 +37,7 @@ pub(super) async fn execute_messages_provider_call(
let response = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
})?;
let transformed = request
.config
.transform_response(&request.model, response)?;
serde_json::to_value(transformed).map_err(|err| {
CoreError::InvalidResponse(format!("failed to serialize messages response: {err}"))
})
request.config.transform_response(&request.model, response)
}
pub(super) async fn execute_messages_provider_stream(

View file

@ -1,2 +1,32 @@
//! The Anthropic Messages call, the Rust equivalent of Python's
//! `litellm.messages()`.
//!
//! [`messages`] is the top-level entrypoint: give it a model, a body, and
//! credentials, and it resolves the provider, transforms the request, calls the
//! provider, and returns a typed non-streaming response. [`messages_stream`]
//! is the streaming variant; it hands the raw upstream response back so a host
//! can splice the event stream to its own caller.
mod client;
mod common_utils;
mod handler;
mod prepare;
pub mod transformation;
pub mod types;
use crate::error::CoreResult;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use prepare::prepare_messages_call;
use types::{AnthropicMessagesResponse, MessagesRequest};
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<AnthropicMessagesResponse> {
execute_messages_provider_call(prepare_messages_call(request)?).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> CoreResult<reqwest::Response> {
execute_messages_provider_stream(prepare_messages_call(request)?).await
}
#[cfg(test)]
mod tests;

View file

@ -1,9 +1,8 @@
use litellm_core::CoreError;
use litellm_core::CoreResult;
use litellm_core::messages::transformation::MessagesAuthStrategy;
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use crate::error::{CoreError, CoreResult};
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
use super::transformation::MessagesAuthStrategy;
use super::types::{MessagesRequest, ProviderMessagesRequest};
pub(super) fn prepare_messages_call(

View file

@ -1,14 +1,16 @@
use std::time::Duration;
use litellm_core::error::CoreError;
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use crate::error::CoreError;
use super::common_utils::{
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
};
use super::{MessagesRequest, messages};
use super::messages;
use super::types::MessagesRequest;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
@ -152,8 +154,8 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
.await
.expect("messages request succeeds");
assert_eq!(response["content"][0]["text"], "hi");
assert_eq!(response["stop_reason"], "end_turn");
assert_eq!(response.content[0]["text"], "hi");
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
let request = server.await.expect("server task completes");
let (head, body) = request.split_once("\r\n\r\n").expect("has body");
@ -208,8 +210,8 @@ async fn messages_round_trip_builds_native_anthropic_request() {
.await
.expect("messages request succeeds");
assert_eq!(response["content"][0]["text"], "hi");
assert_eq!(response["stop_reason"], "end_turn");
assert_eq!(response.content[0]["text"], "hi");
assert_eq!(response.stop_reason.as_deref(), Some("end_turn"));
let request = server.await.expect("server task completes");
let (head, _) = request.split_once("\r\n\r\n").expect("has body");

View file

@ -1,6 +1,30 @@
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::transformation::AnthropicMessagesProviderConfig;
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
pub(super) struct ProviderMessagesRequest {
pub(super) provider: String,
pub(super) model: String,
pub(super) config: &'static dyn AnthropicMessagesProviderConfig,
pub(super) url: String,
pub(super) body: Value,
pub(super) upstream_headers: Vec<(String, String)>,
pub(super) timeout: Option<Duration>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub enum SystemPrompt {

View file

@ -1,3 +1,3 @@
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over litellm-ai-gateway.
litellm-python-bridge is the PyO3 cdylib that exposes Rust to the litellm Python SDK — a thin adapter (Python objects → Rust calls → Python results) over the litellm-core route entrypoints (e.g. `litellm_core::messages::messages`).
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call into litellm-ai-gateway.
Keep it thin: no business logic, no transforms, no I/O orchestration — just marshal in/out and call the core entrypoint.

View file

@ -11,11 +11,11 @@ Python-compatible dictionaries.
## Bridge Shape
- Prefer one stable method per top-level LiteLLM route, for example
`ocr(payload)`.
`messages(...)`, calling the matching `litellm-core` entrypoint.
- Do not add one exported PyO3 function per provider helper unless there is a
measured reason.
- Provider dispatch belongs in Rust route modules such as
`litellm_providers::ocr`, not in this PyO3 crate.
- Provider dispatch belongs in the `litellm-core` route module (e.g.
`litellm_core::messages`), not in this PyO3 crate.
- Python owns rollout state and fallback. Rust should return errors; Python
decides whether to raise or fall back. For a rust-only provider/route (no
Python reference), the Python side is a thin dispatch that calls Rust and

View file

@ -4,10 +4,11 @@ use std::time::Duration;
use litellm_ai_gateway::io::audio_transcription::{
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
};
use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages};
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::error::CoreError;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
@ -35,6 +36,15 @@ fn json_to_py(py: Python<'_>, value: Value) -> PyResult<Py<PyAny>> {
Ok(json.call_method1("loads", (encoded,))?.unbind())
}
fn messages_response_to_py(
py: Python<'_>,
response: AnthropicMessagesResponse,
) -> PyResult<Py<PyAny>> {
let value =
serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?;
json_to_py(py, value)
}
fn core_error_to_pyerr(err: CoreError) -> PyErr {
match err {
CoreError::Auth(message) => PyValueError::new_err(message),
@ -382,7 +392,7 @@ fn messages(
});
match result {
Ok(value) => json_to_py(py, value),
Ok(response) => messages_response_to_py(py, response),
Err(err) => Err(core_error_to_pyerr(err)),
}
}
@ -404,7 +414,7 @@ fn amessages(
marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let value = run_messages(MessagesRequest {
let response = run_messages(MessagesRequest {
model: &model,
body,
api_key: api_key.as_deref(),
@ -416,7 +426,7 @@ fn amessages(
.await
.map_err(core_error_to_pyerr)?;
Python::attach(|py| json_to_py(py, value))
Python::attach(|py| messages_response_to_py(py, response))
})
}

View file

@ -61,23 +61,51 @@ def _get_redis_kwargs():
return available_args
def _get_redis_url_kwargs(client=None):
def _init_arg_names(cls: type) -> frozenset[str]:
"""Every ``__init__`` parameter accepted anywhere in a class's MRO.
Keyword-only parameters are included, and the MRO is walked because redis-py splits a
connection's parameters between ``AbstractConnection`` and its concrete subclasses.
"""
return frozenset(
name
for klass in inspect.getmro(cls)
if klass is not object
for spec in (inspect.getfullargspec(klass.__init__),)
for name in spec.args + spec.kwonlyargs
)
def _get_redis_url_kwargs(client: Optional[type] = None) -> tuple[str, ...]:
"""Connection kwargs that redis-py forwards from ``from_url`` down to the connection.
``from_url`` is declared as ``(cls, url, **kwargs)``, so introspecting it yields no
connection kwargs at all. What it really does is hand its kwargs to the connection
class, so that class's signature is the allowlist.
Taking the client's signature instead would be wrong in both directions: it omits
nothing useful, but it admits client-only parameters such as
``single_connection_client`` and ``auto_close_connection_pool``, plus the ``ssl_*``
family that only ``SSLConnection`` accepts. Those reach ``AbstractConnection`` and
raise ``TypeError`` the first time a connection is created. TLS on a url config is
selected by the ``rediss://`` scheme, which picks ``SSLConnection`` on its own.
"""
if client is None:
client = redis.Redis.from_url
arg_spec = inspect.getfullargspec(redis.Redis.from_url)
client = redis.Redis
connection_cls = async_redis.Connection if client is async_redis.Redis else redis.Connection
exclude_args = frozenset(
{
"self",
"connection_pool",
"retry",
}
)
# Only allow primitive arguments
exclude_args = {
"self",
"connection_pool",
"retry",
}
include_args = ("url", "max_connections")
include_args = ["url"]
available_args = [x for x in arg_spec.args if x not in exclude_args] + include_args
return available_args
return tuple(x for x in _init_arg_names(connection_cls) if x not in exclude_args) + include_args
def _get_redis_cluster_kwargs(client=None):
@ -614,7 +642,7 @@ def get_redis_async_client(
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
if connection_pool is not None:
return async_redis.Redis(connection_pool=connection_pool)
args = _get_redis_url_kwargs(client=async_redis.Redis.from_url)
args = _get_redis_url_kwargs(client=async_redis.Redis)
url_kwargs = {}
for arg in redis_kwargs:
if arg in args:
@ -662,10 +690,10 @@ def get_redis_connection_pool(
return None
if "url" in redis_kwargs and redis_kwargs["url"] is not None:
pool_kwargs = {
"timeout": REDIS_CONNECTION_POOL_TIMEOUT,
"url": redis_kwargs["url"],
}
allowed_args = _get_redis_url_kwargs(client=async_redis.Redis)
pool_kwargs = {k: v for k, v in redis_kwargs.items() if k in allowed_args and k != "max_connections"}
pool_kwargs["timeout"] = REDIS_CONNECTION_POOL_TIMEOUT
pool_kwargs["url"] = redis_kwargs["url"]
if "max_connections" in redis_kwargs:
try:
pool_kwargs["max_connections"] = int(redis_kwargs["max_connections"])

View file

@ -568,6 +568,7 @@ def _build_streaming_logging_obj(
logging_obj.custom_llm_provider = "a2a_agent"
logging_obj.model_call_details["model"] = model
logging_obj.model_call_details["custom_llm_provider"] = "a2a_agent"
logging_obj.model_call_details["call_type"] = logging_obj.call_type
if agent_id:
logging_obj.model_call_details["agent_id"] = agent_id

View file

@ -1,5 +1,6 @@
import json
from typing import Any, Iterator, List, Literal, Optional, Tuple
from dataclasses import dataclass
from typing import Any, Iterable, Iterator, List, Literal, Optional, Tuple
import litellm
from litellm._logging import verbose_logger
@ -24,20 +25,20 @@ async def calculate_batch_cost_and_usage(
deployment-specific pricing (e.g. input_cost_per_token_batches)
is used instead of the global cost map.
"""
batch_cost = _batch_cost_calculator(
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
):
batch_cost, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
return batch_cost, batch_usage, [model_name]
return _aggregate_batch_cost_usage_models(
entries=file_content_dictionary,
custom_llm_provider=custom_llm_provider,
file_content_dictionary=file_content_dictionary,
model_name=model_name,
model_info=model_info,
)
batch_usage = _get_batch_job_total_usage_from_file_content(
file_content_dictionary=file_content_dictionary,
custom_llm_provider=custom_llm_provider,
model_name=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
async def _handle_completed_batch(
@ -46,7 +47,9 @@ async def _handle_completed_batch(
model_name: Optional[str] = None,
litellm_params: Optional[dict] = None,
) -> Tuple[float, Usage, List[str]]:
"""Helper function to process a completed batch and handle logging
"""Fetch a completed batch's output file and aggregate its cost, usage, and
models in a single pass over the JSONL lines, so the parsed file content is
never materialized in memory.
Args:
batch: The batch object
@ -54,75 +57,109 @@ async def _handle_completed_batch(
model_name: Optional model name
litellm_params: Optional litellm parameters containing credentials (api_key, api_base, etc.)
"""
# Get batch results
file_content_dictionary = await _get_batch_output_file_content_as_dictionary(
batch, custom_llm_provider, litellm_params=litellm_params
)
file_content = await _fetch_batch_output_file_content(batch, custom_llm_provider, litellm_params=litellm_params)
# Calculate costs and usage
batch_cost = _batch_cost_calculator(
custom_llm_provider=custom_llm_provider,
file_content_dictionary=file_content_dictionary,
model_name=model_name,
)
batch_usage = _get_batch_job_total_usage_from_file_content(
file_content_dictionary=file_content_dictionary,
custom_llm_provider=custom_llm_provider,
model_name=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
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
"""
if model_name:
return [model_name]
batch_models = []
for _item in file_content_dictionary:
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)
return batch_models
def _batch_cost_calculator(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
model_name: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> float:
"""
Calculate the cost of a batch based on the output file id
"""
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
):
batch_cost, _ = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
verbose_logger.debug("vertex_ai_total_cost=%s", batch_cost)
return batch_cost
batch_cost, batch_usage = calculate_vertex_ai_batch_cost_and_usage(
_get_file_content_as_dictionary(file_content), model_name
)
return batch_cost, batch_usage, [model_name]
# For other providers, use the existing logic
total_cost = _get_batch_job_cost_from_file_content(
file_content_dictionary=file_content_dictionary,
return _aggregate_batch_cost_usage_models(
entries=_iter_batch_input_entries(file_content),
custom_llm_provider=custom_llm_provider,
model_name=model_name,
model_info=model_info,
)
verbose_logger.debug("total_cost=%s", total_cost)
return total_cost
@dataclass(frozen=True, slots=True)
class _BatchOutputLineStats:
cost: float
prompt_tokens: int
completion_tokens: int
total_tokens: int
cache_read_tokens: int
cache_creation_tokens: int
model: Optional[str]
def _iter_successful_output_line_stats(
entries: Iterable[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: Optional[str],
model_info: Optional[ModelInfo],
) -> Iterator[_BatchOutputLineStats]:
from litellm.cost_calculator import batch_cost_calculator
for entry in entries:
if not _batch_response_was_successful(entry, custom_llm_provider):
continue
response_body = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
usage = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
prompt_details = _parse_prompt_tokens_details(usage)
raw_model = response_body.get("model")
response_model = raw_model if isinstance(raw_model, str) and raw_model else None
if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"):
if custom_llm_provider == "bedrock" and model_name:
cost_model = model_name
else:
cost_model = response_model or model_name or ""
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=cost_model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
line_cost = prompt_cost + completion_cost
else:
line_cost = litellm.completion_cost(
completion_response=response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
yield _BatchOutputLineStats(
cost=line_cost,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
cache_read_tokens=prompt_details["cache_hit_tokens"],
cache_creation_tokens=prompt_details["cache_creation_tokens"],
model=response_model,
)
def _aggregate_batch_cost_usage_models(
entries: Iterable[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> Tuple[float, Usage, List[str]]:
"""Aggregate cost, usage, and models from batch output entries in a single
pass, holding one small stats record per line instead of the parsed file."""
line_stats = tuple(_iter_successful_output_line_stats(entries, custom_llm_provider, model_name, model_info))
cache_token_params = {
key: tokens
for key, tokens in (
("cache_read_input_tokens", sum(stats.cache_read_tokens for stats in line_stats)),
("cache_creation_input_tokens", sum(stats.cache_creation_tokens for stats in line_stats)),
)
if tokens > 0
}
batch_usage = Usage(
total_tokens=sum(stats.total_tokens for stats in line_stats),
prompt_tokens=sum(stats.prompt_tokens for stats in line_stats),
completion_tokens=sum(stats.completion_tokens for stats in line_stats),
**cache_token_params,
)
batch_models = [model_name] if model_name else [stats.model for stats in line_stats if stats.model]
total_cost = sum((stats.cost for stats in line_stats), 0.0)
verbose_logger.debug("batch output aggregate: cost=%s usage=%s models=%s", total_cost, batch_usage, batch_models)
return total_cost, batch_usage, batch_models
def calculate_vertex_ai_batch_cost_and_usage(
@ -193,13 +230,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
)
async def _get_batch_output_file_content_as_dictionary(
async def _fetch_batch_output_file_content(
batch: Batch,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
litellm_params: Optional[dict] = None,
) -> List[dict]:
) -> bytes:
"""
Get the batch output file content as a list of dictionaries
Fetch the batch output file and return its raw JSONL bytes
Args:
batch: The batch object
@ -212,9 +249,6 @@ async def _get_batch_output_file_content_as_dictionary(
_is_base64_encoded_unified_file_id,
)
if custom_llm_provider == "vertex_ai":
raise ValueError("Vertex AI does not support file content retrieval")
if batch.output_file_id is None:
raise ValueError("Output file id is None cannot retrieve file content")
@ -240,7 +274,7 @@ async def _get_batch_output_file_content_as_dictionary(
file_content_kwargs.update(credentials)
_file_content = await afile_content(**file_content_kwargs) # type: ignore[reportArgumentType]
return _get_file_content_as_dictionary(_file_content.content)
return _file_content.content
def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict:
@ -270,6 +304,8 @@ def _extract_file_access_credentials(litellm_params: Optional[dict]) -> dict:
"vertex_project",
"vertex_location",
"vertex_credentials",
"gcs_bucket_name",
"bucket_name",
"timeout",
"max_retries",
]
@ -284,17 +320,7 @@ def _get_file_content_as_dictionary(file_content: bytes) -> List[dict]:
"""
Get the file content as a list of dictionaries from JSON Lines format
"""
try:
_file_content_str = file_content.decode("utf-8")
# Split by newlines and parse each line as a separate JSON object
json_objects = []
for line in _file_content_str.strip().split("\n"):
if line: # Skip empty lines
json_objects.append(json.loads(line))
verbose_logger.debug("json_objects=%s", json.dumps(json_objects, indent=4))
return json_objects
except Exception as e:
raise e
return list(_iter_batch_input_entries(file_content))
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
@ -361,101 +387,6 @@ def _count_entry_tokens(
return 0
def _get_batch_job_cost_from_file_content(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
model_name: Optional[str] = None,
model_info: Optional[ModelInfo] = None,
) -> float:
"""
Get the cost of a batch job from the file content
"""
from litellm.cost_calculator import batch_cost_calculator
try:
total_cost: float = 0.0
# 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, 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 in ("anthropic", "bedrock"):
usage = _get_batch_job_usage_from_response_body(_response_body, custom_llm_provider)
# Bedrock batch output lines report a short internal model id
# (e.g. "claude-sonnet-4-6") that is not in the cost map; use the
# deployment model name for pricing when available.
if custom_llm_provider == "bedrock" and model_name:
model = model_name
else:
model = _response_body.get("model") or model_name or ""
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
total_cost += prompt_cost + completion_cost
else:
total_cost += litellm.completion_cost(
completion_response=_response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
verbose_logger.debug("total_cost=%s", total_cost)
return total_cost
except Exception as e:
verbose_logger.error("error in _get_batch_job_cost_from_file_content", e)
raise e
def _get_batch_job_total_usage_from_file_content(
file_content_dictionary: List[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"] = "openai",
model_name: Optional[str] = None,
) -> Usage:
"""
Get the tokens of a batch job from the file content
"""
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
):
_, batch_usage = calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
return batch_usage
# For other providers, use the existing logic
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, 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,
)
def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
"""Token-count a ``prompt`` / ``input`` field that the OpenAI batch
schema allows in four shapes:

View file

@ -517,6 +517,7 @@ class LLMCachingHandler:
cached_result=final_embedding_cached_response,
is_async=True,
is_embedding=True,
custom_llm_provider=custom_llm_provider,
)
self._async_log_cache_hit_on_callbacks(
logging_obj=logging_obj,

View file

@ -17,7 +17,8 @@ import json
import time
from collections.abc import Awaitable, Callable, Sequence
from datetime import timedelta
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast
from contextvars import ContextVar
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, TypeVar, Union, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
@ -168,24 +169,97 @@ class RedisCircuitBreaker:
self._state = self.CLOSED
_RedisCallResult = TypeVar("_RedisCallResult")
_swallowed_redis_failures: ContextVar[int] = ContextVar("litellm_swallowed_redis_failures", default=0)
@functools.lru_cache(maxsize=1)
def _redis_health_error_types() -> tuple[type, ...]:
"""Exception types that mean the Redis backend itself is unhealthy.
Command and data errors say nothing about connectivity: an INCR against a non-numeric
value or an undecodable cached entry is a request problem, and counting those would let
a caller trip the shared breaker on demand, dropping rate limiting to per-process
counters that spreading traffic across replicas can outrun.
Imported lazily because this module is reachable from a base ``import litellm`` while
redis is not a base dependency.
"""
from redis.exceptions import BusyLoadingError, ClusterDownError
from redis.exceptions import ConnectionError as RedisConnectionError
from redis.exceptions import TimeoutError as RedisTimeoutError
return (RedisConnectionError, RedisTimeoutError, BusyLoadingError, ClusterDownError, OSError, asyncio.TimeoutError)
def _is_redis_health_failure(exc: BaseException) -> bool:
"""True when ``exc`` indicates Redis is unreachable rather than the request being bad."""
try:
return isinstance(exc, _redis_health_error_types())
except ImportError:
return True
def _record_swallowed_redis_failure(breaker: RedisCircuitBreaker, exc: BaseException) -> None:
"""Record a Redis failure that the calling method is about to swallow.
The marker is a ContextVar rather than a counter on the breaker because breakers are
shared by every concurrent caller. A plain shared counter cannot tell "my call failed"
from "some other in-flight call failed", so a success overlapping someone else's
failure would be discarded and a Redis that is answering would still be evicted.
asyncio gives each task its own copy of the context, so this is per-call.
"""
if not _is_redis_health_failure(exc):
return
breaker.record_failure()
_swallowed_redis_failures.set(_swallowed_redis_failures.get() + 1)
async def _run_under_circuit_breaker(
breaker: RedisCircuitBreaker,
name: str,
call: Callable[[], Awaitable[_RedisCallResult]],
) -> _RedisCallResult:
"""Run one Redis coroutine under a circuit breaker.
Shared by the method decorator and the Lua script executor so both feed the same
health signal. Success is recorded only when nothing failed while ``call`` ran,
because several Redis methods catch their own connection errors and return a default.
"""
if breaker.is_open():
raise Exception(f"Redis circuit breaker is open — skipping {name}")
swallowed_before = _swallowed_redis_failures.get()
try:
result = await call()
except Exception as e:
if _is_redis_health_failure(e):
breaker.record_failure()
raise
if _swallowed_redis_failures.get() == swallowed_before:
breaker.record_success()
return result
def _redis_circuit_breaker_guard(method): # type: ignore
"""
Decorator for RedisCache async methods.
Checks the circuit breaker before each call; records success/failure after.
Does not apply to ping/disconnect/test_connection (health/teardown must always run).
A returning method is not proof of a healthy Redis: several methods catch their own
connection errors and return a default so callers degrade rather than fail. Counting
those as successes reset the failure streak on every request, so the breaker could
never open and Redis was never taken out of the pool. Success is therefore recorded
only when no failure was registered while the method ran.
"""
@functools.wraps(method)
async def wrapper(self, *args, **kwargs): # type: ignore
if self._circuit_breaker.is_open():
raise Exception(f"Redis circuit breaker is open — skipping {method.__name__}")
try:
result = await method(self, *args, **kwargs)
self._circuit_breaker.record_success()
return result
except Exception:
self._circuit_breaker.record_failure()
raise
return await _run_under_circuit_breaker(
self._circuit_breaker, method.__name__, lambda: method(self, *args, **kwargs)
)
return wrapper
@ -551,13 +625,16 @@ class RedisCache(BaseCache):
)
async def run_script(keys: Sequence[str], args: Sequence[Any], client: Any = None) -> Any:
executor: Optional[Callable[..., Awaitable[Any]]] = litellm.in_memory_llm_clients_cache.get_cache(
key=script_cache_key
)
if executor is None:
executor = self._register_script_for_current_loop(script)
litellm.in_memory_llm_clients_cache.set_cache(key=script_cache_key, value=executor)
return await executor(keys=keys, args=args, client=client)
async def execute() -> object:
executor: Optional[Callable[..., Awaitable[Any]]] = litellm.in_memory_llm_clients_cache.get_cache(
key=script_cache_key
)
if executor is None:
executor = self._register_script_for_current_loop(script)
litellm.in_memory_llm_clients_cache.set_cache(key=script_cache_key, value=executor)
return await executor(keys=keys, args=args, client=client)
return await _run_under_circuit_breaker(self._circuit_breaker, "run_script", execute)
return run_script
@ -674,6 +751,7 @@ class RedisCache(BaseCache):
str(e),
value,
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
async def _pipeline_helper(
self,
@ -758,6 +836,7 @@ class RedisCache(BaseCache):
str(e),
cache_value,
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
async def _set_cache_sadd_helper(
self,
@ -842,6 +921,7 @@ class RedisCache(BaseCache):
str(e),
value,
)
_record_swallowed_redis_failure(self._circuit_breaker, e)
@_redis_circuit_breaker_guard
async def batch_cache_write(self, key, value, **kwargs):
@ -1106,6 +1186,7 @@ class RedisCache(BaseCache):
)
)
print_verbose(f"litellm.caching.caching: async get() - Got exception from REDIS: {str(e)}")
_record_swallowed_redis_failure(self._circuit_breaker, e)
@_redis_circuit_breaker_guard
async def async_batch_get_cache(
@ -1177,6 +1258,7 @@ class RedisCache(BaseCache):
)
)
verbose_logger.error(f"Error occurred in async batch get cache - {str(e)}")
_record_swallowed_redis_failure(self._circuit_breaker, e)
return key_value_dict
def sync_ping(self) -> bool:
@ -1432,6 +1514,7 @@ class RedisCache(BaseCache):
return ttl
except Exception as e:
verbose_logger.debug(f"Redis TTL Error: {e}")
_record_swallowed_redis_failure(self._circuit_breaker, e)
return None
@_redis_circuit_breaker_guard

View file

@ -209,7 +209,15 @@ class ResponsesToCompletionBridgeHandler:
json_mode=kwargs.get("json_mode"),
)
elif isinstance(result, ModelResponse):
return result
if not stream:
return result
return self._completed_response_as_stream(
response=result,
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
json_mode=kwargs.get("json_mode"),
)
elif not stream:
responses_api_response = self._collect_response_from_stream(result)
return self.transformation_handler.transform_response(
@ -299,7 +307,15 @@ class ResponsesToCompletionBridgeHandler:
json_mode=kwargs.get("json_mode"),
)
elif isinstance(result, ModelResponse):
return result
if not stream:
return result
return self._completed_response_as_stream(
response=result,
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
json_mode=kwargs.get("json_mode"),
)
elif not stream:
responses_api_response = await self._collect_response_from_stream_async(result)
return self.transformation_handler.transform_response(
@ -331,6 +347,25 @@ class ResponsesToCompletionBridgeHandler:
)
return self._apply_post_stream_processing(streamwrapper, model, custom_llm_provider)
def _completed_response_as_stream(
self,
response: "ModelResponse",
model: str,
custom_llm_provider: str,
logging_obj: "LiteLLMLoggingObj",
json_mode: bool | None,
) -> "CustomStreamWrapper":
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
streamwrapper = CustomStreamWrapper(
completion_stream=MockResponseIterator(model_response=response, json_mode=json_mode),
model=model,
custom_llm_provider=custom_llm_provider,
logging_obj=logging_obj,
)
return self._apply_post_stream_processing(streamwrapper, model, custom_llm_provider)
@staticmethod
def _apply_post_stream_processing(
stream: "CustomStreamWrapper",

View file

@ -35,6 +35,7 @@ from litellm.responses.sse_output_recovery import (
record_output_item_chunk,
record_output_text_chunk,
)
from litellm.responses.utils import normalize_responses_api_stream_options
from litellm.types.llms.openai import (
ChatCompletionAnnotation,
ChatCompletionReasoningItem,
@ -320,6 +321,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
responses_api_request["tool_choice"] = ( # type: ignore[assignment]
self._normalize_tool_choice_for_responses_api(value)
)
elif key == "stream_options":
stream_options = normalize_responses_api_stream_options(value)
if stream_options is not None:
responses_api_request["stream_options"] = stream_options
elif key in ResponsesAPIOptionalRequestParams.__annotations__.keys():
responses_api_request[key] = value # type: ignore
elif key == "previous_response_id":
@ -360,8 +365,6 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
continue
if key == "instructions" and instructions:
request_data["instructions"] = instructions
elif key == "stream_options" and isinstance(value, dict):
request_data["stream_options"] = value.get("include_obfuscation")
elif key == "user" and isinstance(value, str):
# OpenAI API requires user param to be max 64 chars - truncate if longer
if len(value) <= 64:
@ -1074,6 +1077,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False):
super().__init__(streaming_response, sync_stream, json_mode)
self._chat_completion_id: str | None = None
def _handle_string_chunk(
self, str_line: Union[str, "BaseModel"]
@ -1381,4 +1385,13 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator):
ModelResponseStream: OpenAI-formatted streaming chunk
"""
verbose_logger.debug(f"Chat provider: transform_streaming_response called with chunk: {chunk}")
return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
return self._with_stream_scoped_id(
OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk)
)
def _with_stream_scoped_id(self, chunk: "ModelResponseStream") -> "ModelResponseStream":
if self._chat_completion_id is None:
self._chat_completion_id = chunk.id
else:
chunk.id = self._chat_completion_id
return chunk

View file

@ -3,6 +3,7 @@ Main compress() function — normalizes input messages, orchestrates BM25/embedd
scoring, message stubbing, and retrieval tool injection.
"""
from collections.abc import Mapping, Sequence
from typing import Any, Dict, List, Optional, Set, Tuple, Union, cast
from litellm.caching.dual_cache import DualCache
@ -204,33 +205,21 @@ def _extract_anthropic_tool_exchange_spans(
return spans, None
def _get_protected_indices(messages: List[dict]) -> List[int]:
def get_protected_indices(messages: Sequence[Mapping[str, object]]) -> tuple[int, ...]:
"""
Return indices of messages that must never be compressed:
- All system messages
- The last user message
- The last assistant message
The last user message is what the model is being asked to act on right now,
so compressing it replaces the live instruction with a marker. Compression
guardrails share this policy; see the Headroom guardrail.
"""
protected: List[int] = []
last_user_idx = None
last_assistant_idx = None
for i, msg in enumerate(messages):
role = msg.get("role", "")
if role == "system":
protected.append(i)
elif role == "user":
last_user_idx = i
elif role == "assistant":
last_assistant_idx = i
if last_user_idx is not None:
protected.append(last_user_idx)
if last_assistant_idx is not None:
protected.append(last_assistant_idx)
return protected
system_indices = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "system")
last_user = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "user")[-1:]
last_assistant = tuple(index for index, msg in enumerate(messages) if msg.get("role", "") == "assistant")[-1:]
return system_indices + last_user + last_assistant
def _combine_scores(
@ -432,7 +421,7 @@ def compress(
combined_scores = bm25_scores
# Protected messages are never compressed
protected_indices = _get_protected_indices(normalized_messages)
protected_indices = get_protected_indices(normalized_messages)
kept_indices: Set[int] = set(protected_indices)
tool_exchange_spans: List[Set[int]] = []

View file

@ -1297,6 +1297,7 @@ X_LITELLM_DISABLE_CALLBACKS = "x-litellm-disable-callbacks"
LITELLM_METADATA_FIELD = "litellm_metadata"
OLD_LITELLM_METADATA_FIELD = "metadata"
RETURN_RAW_MODEL_NAME_METADATA_KEY = "_complexity_router_return_raw_model_name"
INTERNAL_CALL_ORIGIN_METADATA_KEY = "internal_call_origin"
LITELLM_TRUNCATED_PAYLOAD_FIELD = "litellm_truncated"
LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE = (
"Truncation is a DB storage safeguard. "
@ -1418,6 +1419,8 @@ LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE = int(
os.getenv("LITELLM_EXPIRED_UI_SESSION_KEY_CLEANUP_BATCH_SIZE", 1000)
)
LITELLM_PROXY_ADMIN_NAME = "default_user_id"
LITELLM_PROXY_BUDGET_NAME = "litellm-proxy-budget"
GLOBAL_PROXY_SPEND_CACHE_KEY = f"{LITELLM_PROXY_ADMIN_NAME}:spend"
########################### CLI SSO AUTHENTICATION CONSTANTS ###########################
LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli"
@ -1455,6 +1458,7 @@ SPEND_LOG_CLEANUP_MAX_CONSECUTIVE_BATCH_FAILURES = int(os.getenv("SPEND_LOG_CLEA
SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS = float(
os.getenv("SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS", 0.5)
)
TOOL_SPEND_TOP_TOOLS = 100
SPEND_LOG_PARTITION_INTERVAL = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day")
SPEND_LOG_PARTITION_PRECREATE_AHEAD = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7))
SPEND_LOG_QUEUE_SIZE_THRESHOLD = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100))

View file

@ -1,3 +1,4 @@
import contextvars
import hashlib
import os
import secrets
@ -16,7 +17,11 @@ from typing import (
)
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.core_helpers import redact_nested_match_and_regex_keys
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
get_or_create_metadata_bucket,
redact_nested_match_and_regex_keys,
)
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
from litellm.secret_managers.main import str_to_bool
@ -64,6 +69,12 @@ from litellm.exceptions import (
# proxy's metadata sanitizer.
_PRE_CALL_EXECUTED_TOKEN = secrets.token_hex(16)
_GUARDRAIL_BLOCK_STATUS_CODES = frozenset({400, 403, 422})
_guardrail_self_recorded: contextvars.ContextVar[bool] = contextvars.ContextVar(
"litellm_guardrail_self_recorded", default=False
)
def _strict_guardrail_modes_enabled() -> bool:
"""Whether guardrail-mode validation raises (default) or logs a warning.
@ -102,6 +113,8 @@ class CustomGuardrail(CustomLogger):
# If True, during_call runs async_moderation_hook instead of the unified apply_guardrail path.
use_native_during_call_hook: ClassVar[bool] = False
records_own_guardrail_information: ClassVar[bool] = False
def __init__(
self,
guardrail_name: Optional[str] = None,
@ -117,6 +130,7 @@ class CustomGuardrail(CustomLogger):
on_sensitive_data: Optional[str] = None,
sensitive_data_route_to_model: Optional[str] = None,
sticky_session_routing: bool = True,
run_in_parallel: bool = False,
only_scan_new_messages: bool = False,
**kwargs,
):
@ -136,6 +150,9 @@ class CustomGuardrail(CustomLogger):
on_sensitive_data: Action when sensitive data is detected. 'block' (default) or 'route'
sensitive_data_route_to_model: Model to route to when on_sensitive_data='route'
sticky_session_routing: When True, all subsequent requests in the session use the same model
run_in_parallel: When True, this pre_call or post_call guardrail runs concurrently with
other opted-in guardrails of the same hook. Only safe for block-only guardrails that
do not mutate the request or response.
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -150,6 +167,7 @@ class CustomGuardrail(CustomLogger):
self.on_sensitive_data: Optional[str] = on_sensitive_data
self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model
self.sticky_session_routing: bool = sticky_session_routing
self.run_in_parallel: bool = run_in_parallel
self.only_scan_new_messages: bool = only_scan_new_messages
if supported_event_hooks:
@ -944,17 +962,10 @@ class CustomGuardrail(CustomLogger):
# should not happen
container[key] = [existing, slg]
if "metadata" in request_data:
if request_data["metadata"] is None:
request_data["metadata"] = {}
_append_guardrail_info(request_data["metadata"])
elif "litellm_metadata" in request_data:
_append_guardrail_info(request_data["litellm_metadata"])
else:
# Ensure guardrail info is always logged (e.g. proxy may not have set
# metadata yet). Attach to "metadata" so spend log / standard logging see it.
request_data["metadata"] = {}
_append_guardrail_info(request_data["metadata"])
_, metadata_bucket = get_or_create_metadata_bucket(request_data)
_append_guardrail_info(metadata_bucket)
_guardrail_self_recorded.set(True)
# Emit the otel guardrail span here, where every guardrail execution lands,
# rather than relying on a post-call hook that does not fire on every path
@ -1046,8 +1057,15 @@ class CustomGuardrail(CustomLogger):
- GuardrailRaisedException (generic guardrail API, tool permission)
- BlockedPiiEntityError (Presidio PII detection)
- SensitiveDataRouteException (sensitive-data reroute to on-premise model)
- HTTPException with status 400 (content policy violation)
- HTTPException with a block-signalling status (400, 403, 422)
- ModifyResponseException (passthrough mode violation)
Only the statuses guardrails use in-tree to signal a deliberate rejection
count as an intervention: 400 (content policy), 403 (e.g. akto) and 422
(e.g. llm_as_a_judge). Other 4xx codes are commonly propagated from an
upstream guardrail provider response (401 bad key, 408 timeout, 429 rate
limit, or a raw upstream status), which are technical failures, not
blocks, so they stay guardrail_failed_to_respond.
"""
if isinstance(e, ModifyResponseException):
return True
@ -1060,7 +1078,11 @@ class CustomGuardrail(CustomLogger):
),
):
return True
if HTTPException is not None and isinstance(e, HTTPException) and e.status_code == 400:
if (
HTTPException is not None
and isinstance(e, HTTPException)
and e.status_code in _GUARDRAIL_BLOCK_STATUS_CODES
):
return True
return False
@ -1160,6 +1182,7 @@ class CustomGuardrail(CustomLogger):
call_type == CallTypes.completion.value
or call_type == CallTypes.acompletion.value
or call_type == CallTypes.anthropic_messages.value
or call_type == CallTypes.call_mcp_tool.value
):
return data.get("messages")
@ -1211,7 +1234,7 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object)
"""
if logging_obj is None:
return
meta_src = request_data.get("metadata") or request_data.get("litellm_metadata") or {}
meta_src = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
slg_info = meta_src.get("standard_logging_guardrail_information")
if not slg_info:
return
@ -1238,8 +1261,20 @@ def log_guardrail_information(func):
(structured detections, tracing detail) than this decorator's
"allow"/"mask"/raw-response default. To avoid double-recording in that
case (which would emit two spans, two Datadog records, two spend-log
entries, etc.), snapshot the entry count before invocation: if the
wrapped function already appended its own entry, skip the auto-record.
entries, etc.), a context-local flag records whether the wrapped function
appended its own entry; if so, the auto-record is skipped. The flag is a
``ContextVar`` rather than a count of entries in the shared ``request_data``
so it stays correct when guardrails run concurrently (asyncio copies the
context into each gathered task): counting shared entries would let one
guardrail's append hide another guardrail's missing record.
A guardrail that only records an entry when it actually runs (e.g.
``HeadroomGuardrail``, which returns the inputs untouched on an endpoint
whose payload it cannot act on) sets ``records_own_guardrail_information =
True`` so the auto-record is skipped even on the return paths where it
recorded nothing; otherwise a no-op early return would be logged as an
"allow"/"success" run even though the guardrail did nothing. The exception
branch below still records so a genuine failure is not lost.
"""
import functools
import inspect
@ -1259,16 +1294,6 @@ def log_guardrail_information(func):
return GuardrailEventHooks.post_call
return None
def _count_recorded_guardrail_entries(request_data: dict) -> int:
total = 0
for container_key in ("metadata", "litellm_metadata"):
container = request_data.get(container_key)
if isinstance(container, dict):
entries = container.get("standard_logging_guardrail_information")
if isinstance(entries, list):
total += len(entries)
return total
@functools.wraps(func)
async def async_wrapper(*args, **kwargs):
start_time = datetime.now() # Move start_time inside the wrapper
@ -1282,10 +1307,10 @@ def log_guardrail_information(func):
original_inputs = kwargs.get("inputs")
logging_obj = kwargs.get("logging_obj")
entries_before = _count_recorded_guardrail_entries(request_data)
self_recorded_token = _guardrail_self_recorded.set(False)
try:
response = await func(*args, **kwargs)
if _count_recorded_guardrail_entries(request_data) > entries_before:
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
return response
return self._process_response(
response=response,
@ -1297,7 +1322,7 @@ def log_guardrail_information(func):
original_inputs=original_inputs,
)
except Exception as e:
if _count_recorded_guardrail_entries(request_data) > entries_before:
if _guardrail_self_recorded.get():
raise
return self._process_error(
e=e,
@ -1308,6 +1333,7 @@ def log_guardrail_information(func):
event_type=event_type,
)
finally:
_guardrail_self_recorded.reset(self_recorded_token)
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
@functools.wraps(func)
@ -1323,10 +1349,10 @@ def log_guardrail_information(func):
original_inputs = kwargs.get("inputs")
logging_obj = kwargs.get("logging_obj")
entries_before = _count_recorded_guardrail_entries(request_data)
self_recorded_token = _guardrail_self_recorded.set(False)
try:
response = func(*args, **kwargs)
if _count_recorded_guardrail_entries(request_data) > entries_before:
if self.records_own_guardrail_information or _guardrail_self_recorded.get():
return response
return self._process_response(
response=response,
@ -1336,7 +1362,7 @@ def log_guardrail_information(func):
original_inputs=original_inputs,
)
except Exception as e:
if _count_recorded_guardrail_entries(request_data) > entries_before:
if _guardrail_self_recorded.get():
raise
return self._process_error(
e=e,
@ -1345,6 +1371,7 @@ def log_guardrail_information(func):
event_type=event_type,
)
finally:
_guardrail_self_recorded.reset(self_recorded_token)
_sync_guardrail_info_to_logging_obj(request_data, logging_obj)
@functools.wraps(func)

View file

@ -519,7 +519,13 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac
"""
This log gets called after the MCP tool call is made.
Useful if you want to modiy the standard logging payload after the MCP tool call is made.
Useful if you want to modify the standard logging payload after the MCP tool call is made.
To change what the caller sends back to the MCP client, mutate ``response_obj``
in place: every call site discards the returned object, because the
dispatcher unwraps it to ``mcp_tool_call_response`` (a raw content list, not
a ``CallToolResult``) which the tool-call paths cannot forward. Guardrails
that mask or reject tool output should use ``post_mcp_call`` instead.
"""
return None

View file

@ -133,6 +133,15 @@ class LangsmithLogger(CustomBatchLogger):
"dotted_order": metadata.get("dotted_order", None),
}
def _redact_metadata(self, metadata: dict) -> dict:
# helper is shallow; also scrub nested requester_metadata since
# LangSmith forwards the whole dict into the run
redacted = redact_user_api_key_info(metadata=dict(metadata))
nested = redacted.get("requester_metadata")
if isinstance(nested, dict):
redacted["requester_metadata"] = redact_user_api_key_info(metadata=nested)
return redacted
def _build_extra_metadata(self, metadata: Dict):
extra_metadata = dict(metadata)
requester_metadata = extra_metadata.get("requester_metadata")
@ -141,13 +150,7 @@ class LangsmithLogger(CustomBatchLogger):
if key in requester_metadata and key not in extra_metadata:
extra_metadata[key] = requester_metadata[key]
# helper is shallow; also scrub nested requester_metadata since
# LangSmith forwards the whole dict into `extra`
extra_metadata = redact_user_api_key_info(metadata=extra_metadata)
nested = extra_metadata.get("requester_metadata")
if isinstance(nested, dict):
extra_metadata["requester_metadata"] = redact_user_api_key_info(metadata=nested)
return extra_metadata
return self._redact_metadata(extra_metadata)
def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> Dict[str, Any]:
response = payload["response"]
@ -200,12 +203,13 @@ class LangsmithLogger(CustomBatchLogger):
metadata = payload["metadata"]
extra_metadata = self._build_extra_metadata(dict(metadata))
inputs = {**payload, "metadata": self._redact_metadata(dict(metadata))}
outputs = self._build_outputs_with_usage(payload)
data = {
"name": fields["run_name"],
"run_type": "llm",
"inputs": payload,
"inputs": inputs,
"outputs": outputs,
"session_name": fields["project_name"],
"start_time": payload["startTime"],

View file

@ -27,6 +27,7 @@ from litellm.integrations.opentelemetry_utils.gen_ai_semconv import (
OTELSemconvCategory,
parse_semconv_opt_in,
)
from litellm.integrations.otel.model.semconv import Metric
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.secret_redaction import redact_string
from litellm.secret_managers.main import get_secret_bool, str_to_bool
@ -117,6 +118,7 @@ TOKEN_TYPE_ATTRIBUTE: str = "gen_ai.token.type"
VALID_METRIC_ATTRIBUTE_NAMES: FrozenSet[str] = frozenset(
(
"gen_ai.operation.name",
"gen_ai.provider.name",
"gen_ai.system",
"gen_ai.request.model",
"gen_ai.framework",
@ -597,32 +599,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
meter = meter_provider.get_meter(__name__)
self._operation_duration_histogram = meter.create_histogram(
name="gen_ai.client.operation.duration", # Replace with semconv constant in otel 1.38
name=Metric.OPERATION_DURATION,
description="GenAI operation duration",
unit="s",
)
self._token_usage_histogram = meter.create_histogram(
name="gen_ai.client.token.usage", # Replace with semconv constant in otel 1.38
name=Metric.TOKEN_USAGE,
description="GenAI token usage",
unit="{token}",
)
self._cost_histogram = meter.create_histogram(
name="gen_ai.client.token.cost",
name=Metric.TOKEN_COST,
description="GenAI request cost",
unit="USD",
)
self._time_to_first_token_histogram = meter.create_histogram(
name="gen_ai.client.response.time_to_first_token",
name=Metric.TIME_TO_FIRST_TOKEN,
description="Time to first token for streaming requests",
unit="s",
)
self._time_per_output_token_histogram = meter.create_histogram(
name="gen_ai.client.response.time_per_output_token",
name=Metric.TIME_PER_OUTPUT_TOKEN,
description="Average time per output token (generation time / completion tokens)",
unit="s",
)
self._response_duration_histogram = meter.create_histogram(
name="gen_ai.client.response.duration",
name=Metric.RESPONSE_DURATION,
description="Total LLM API generation time (excludes LiteLLM overhead)",
unit="s",
)
@ -883,8 +885,8 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
request_data: dict,
parent_span: Optional[Any],
) -> None:
"""Emit ``guardrail`` spans from ``request_data["metadata"]
["standard_logging_guardrail_information"]``.
"""Emit ``guardrail`` spans from the request's proxy-internal metadata bucket
(``standard_logging_guardrail_information``).
Routed through ``_create_guardrail_span`` so the dedupe state in
``_otel_internal`` is honoured — if ``_handle_failure`` already
@ -892,7 +894,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
"""
from opentelemetry import trace as _trace
metadata = (request_data or {}).get("metadata") or {}
from litellm.litellm_core_utils.core_helpers import (
get_metadata_variable_name_from_kwargs,
)
request_data = request_data or {}
metadata = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) or {}
guardrail_information = metadata.get("standard_logging_guardrail_information")
if not guardrail_information:
return
@ -2975,10 +2982,12 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
def _get_metric_reader(self):
"""
Get the appropriate metric reader based on the configuration.
Histograms keep the SDK's default cumulative temporality: Prometheus-backed
OTLP receivers reject delta histograms and drop the whole batch, while
backends that prefer delta still accept cumulative.
"""
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import (
AggregationTemporality,
ConsoleMetricExporter,
PeriodicExportingMetricReader,
)
@ -3009,7 +3018,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
exporter = OTLPMetricExporter(
endpoint=normalized_endpoint,
headers=_split_otel_headers,
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)
@ -3027,7 +3035,6 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
exporter = OTLPMetricExporter(
endpoint=normalized_endpoint,
headers=_split_otel_headers,
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
return PeriodicExportingMetricReader(exporter, export_interval_millis=5000)

View file

@ -222,7 +222,29 @@ lives in [`plumbing/`](./plumbing):
otherwise the operator's globally configured `MeterProvider` is reused so its
readers/exporters receive them alongside the server metrics, and one is built
and registered as the global only when none is set (mirroring how V2 owns trace
export).
export). A **failed** call records `gen_ai.client.operation.duration` too,
carrying the semconv `error.type` (the mapped provider exception's class name),
so the histogram covers the whole traffic and failure-rate panels are buildable;
the other five instruments describe a completed generation and are skipped
rather than filled with a fabricated zero. `error.type` is stamped after the
cardinality filter, so an `otel.attributes` list cannot merge the failure series
back into the success series. A proxy-gate rejection (auth / rate limit) records
nothing, for the same reason it gets no span: no upstream call happened.
Both paths cap their attributes at `METRIC_ATTRIBUTE_CEILING` before the
operator's own `otel.attributes` filter runs, so the filter can narrow the set
but never widen it. The ceiling is what keeps series count bounded by the
deployment's own key/team/user/deployment count instead of by its traffic: a
label value that moves per request mints a time series per request, which both
bills per request on a hosted backend and leaves a histogram that cannot be
aggregated. So client-supplied and per-request metadata (`requester_metadata`,
`spend_logs_metadata`, `user_api_key_end_user_id`, `requester_ip_address`) is
metric-ineligible and stays on the span, where cardinality is free, and the
`hidden_params` label carries only `model_id`, the deployment identity a
per-deployment panel joins on. `api_base` is excluded despite naming the same
deployment, because it is a documented per-call parameter and so is caller-chosen
in SDK use. Because the shared validator accepts every span attribute name, a
filter that names a metric-ineligible one logs a warning once when the filter
resolves rather than silently emitting nothing for it.
- [`events.py`](./plumbing/events.py) — GenAI client events. Gated on
`enable_events` (`LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS`), a failed LLM call
records the semconv `gen_ai.client.operation.exception` log event at severity
@ -267,7 +289,10 @@ lives in [`plumbing/`](./plumbing):
- **A new attribute vocabulary for a backend**: add a mapper in `mappers/`
(a class with a `map(data) -> AttributeMap` method, typically built from
`key -> extractor` tables) and register it in `mappers/__init__._MAPPER_BY_NAME`.
`key -> extractor` tables) and register it in `mappers/__init__._PLAIN_MAPPERS`.
If it spells declared tool definitions out per index, register it in
`_TOOL_DEFINITION_MAPPERS` instead and take the shared attribute budget in its
constructor, so the family stays bounded span-wide rather than per vocabulary.
- **A new integration**: add a preset in `presets/` that returns an
`OpenTelemetryV2Config`, and register it in `presets/__init__.PRESET_BY_CALLBACK`.
If it supports dynamic credentials, add a header builder to

View file

@ -17,6 +17,7 @@ from litellm.integrations.otel.model.baggage import promoted_baggage
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
from litellm.integrations.otel.plumbing.context import (
is_recordable_span,
mcp_message_transport_span,
request_root_span,
resolve_mcp_span_context,
resolve_parent_context,
@ -280,13 +281,29 @@ class OpenTelemetryV2(CustomLogger):
self._record_metrics(kwargs, response_obj, start_time, end_time)
def _record_metrics(self, kwargs, response_obj, start_time, end_time) -> None:
"""Record the GenAI metrics for a successful LLM call. Best-effort: a
recording failure (e.g. a malformed payload) must never break the span
close or the request itself."""
"""Record the GenAI metrics for a successful LLM call."""
self._guarded_record(lambda recorder: recorder.record(kwargs, response_obj, start_time, end_time))
def _record_failure_metrics(self, kwargs, start_time, end_time) -> None:
"""Record the GenAI metrics for a failed LLM call, so the duration
histogram covers the whole traffic rather than only what survived.
A synthetic proxy-gate log (auth / rate-limit rejection) is skipped for the
same reason it gets no span: no upstream call happened, so its duration is
not a GenAI operation's duration and would pull the histogram down."""
if LLMCallEvent.from_dict(kwargs).is_no_upstream_call:
return
self._guarded_record(lambda recorder: recorder.record_failure(kwargs, start_time, end_time))
def _guarded_record(self, record: "Callable[[GenAIMetricRecorder], None]") -> None:
"""Run one metric recording. Best-effort: a recording failure (e.g. a
malformed payload) must never break the span close or the request itself. A
misconfigured attribute filter is operator-fixable, so it is surfaced once
at ERROR instead of being swallowed."""
if self._metrics_recorder is None:
return
try:
self._metrics_recorder.record(kwargs, response_obj, start_time, end_time)
record(self._metrics_recorder)
except ValueError as exc:
if not self._metric_filter_error_logged:
verbose_logger.error(
@ -303,6 +320,7 @@ class OpenTelemetryV2(CustomLogger):
if self._emit_mcp_list_tools(kwargs, start_time, end_time):
return
self._close_llm_call(kwargs, start_time, end_time)
self._record_failure_metrics(kwargs, start_time, end_time)
def _seed_identity_baggage(self, identity: RequestIdentity, model: str | None, context: Context) -> Context:
"""Seed authenticated request-identity Baggage onto ``context`` so the Baggage
@ -641,8 +659,14 @@ class OpenTelemetryV2(CustomLogger):
endpoint, auth failure), so the failed request carries the same error keys
a failed LLM call does. v1's ``OpenTelemetry`` implemented this same hook;
v2 lost it when it stopped subclassing ``OpenTelemetry``, which is the
LIT-4179 regression for pre-call failures."""
span = request_root_span() or user_api_key_dict.parent_otel_span
LIT-4179 regression for pre-call failures.
An MCP message is handled on the session's task, where the request-root
anchor is whatever request opened the session, so prefer the transport the
gateway published for this specific message. Without that, a failed tool
call aimed its error at the ``initialize`` request's finished span and the
SDK dropped it, leaving the POST that actually failed unmarked."""
span = mcp_message_transport_span() or request_root_span() or user_api_key_dict.parent_otel_span
if span is None or not is_recordable_span(span):
return None
stamp_error(span, _span_error_from_exception(original_exception, traceback_str=traceback_str))

View file

@ -18,13 +18,19 @@ from litellm.integrations.otel.mappers.langfuse import LangfuseMapper
from litellm.integrations.otel.mappers.langtrace import LangtraceMapper
from litellm.integrations.otel.mappers.legacy import LegacyMapper
from litellm.integrations.otel.mappers.openinference import OpenInferenceMapper
from litellm.integrations.otel.mappers.utils import tool_attr_budget
from litellm.integrations.otel.mappers.weave import WeaveMapper
# Registry keyed by ``config.mapper_names`` entries.
_MAPPER_BY_NAME: dict[str, Callable[[], AttributeMapper]] = {
# Registries keyed by ``config.mapper_names`` entries, split by whether the
# vocabulary spells declared tool definitions out per index. Those share one
# span-wide attribute ceiling, so resolution has to know how many of them are
# active before it can build them.
_TOOL_DEFINITION_MAPPERS: dict[str, Callable[[int], AttributeMapper]] = {
"genai": GenAIMapper,
"legacy": LegacyMapper,
"openinference": OpenInferenceMapper,
}
_PLAIN_MAPPERS: dict[str, Callable[[], AttributeMapper]] = {
"langfuse": LangfuseMapper,
"weave": WeaveMapper,
"langtrace": LangtraceMapper,
@ -33,13 +39,19 @@ _MAPPER_BY_NAME: dict[str, Callable[[], AttributeMapper]] = {
def resolve_mappers(names: Iterable[str]) -> list[AttributeMapper]:
"""Resolve mapper names to instances. Unknown names raise ``ValueError``."""
out: list[AttributeMapper] = []
for name in names:
factory = _MAPPER_BY_NAME.get(name)
if factory is None:
raise ValueError(f"unknown mapper name {name!r}; known: {sorted(_MAPPER_BY_NAME)}")
out.append(factory())
return out
ordered = tuple(names)
for name in ordered:
if name not in _TOOL_DEFINITION_MAPPERS and name not in _PLAIN_MAPPERS:
known = sorted((*_TOOL_DEFINITION_MAPPERS, *_PLAIN_MAPPERS))
raise ValueError(f"unknown mapper name {name!r}; known: {known}")
# Distinct vocabularies each write the tool family under their own keys, so
# the ceiling is split by how many of them are configured. Repeating a name
# rewrites the same keys, so only distinct ones count.
budget = tool_attr_budget(len({*ordered} & _TOOL_DEFINITION_MAPPERS.keys()))
return [
_TOOL_DEFINITION_MAPPERS[name](budget) if name in _TOOL_DEFINITION_MAPPERS else _PLAIN_MAPPERS[name]()
for name in ordered
]
__all__ = [

View file

@ -11,10 +11,11 @@ from typing import Callable
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData
from litellm.integrations.otel.mappers.utils import (
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
collect,
drop_none,
output_messages,
serialize_messages,
tool_definition_attrs,
)
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
@ -135,6 +136,9 @@ class GenAIMapper:
LiteLLM.SERVICE_CALL_TYPE: lambda d: d.call_type,
}
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
self._tool_attr_budget = tool_attr_budget
def map(self, data: SpanData) -> AttributeMap:
match data:
case LLMCallSpanData():
@ -150,18 +154,18 @@ class GenAIMapper:
case _:
return {}
@classmethod
def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap:
attrs = collect(cls._LLM_CALL_ATTRS, data)
attrs.update(
drop_none(
{
f"gen_ai.tool.{idx}.{suffix}": extract(tool)
for idx, tool in enumerate(data.tools)
for suffix, extract in cls._TOOL_ATTRS.items()
}
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
attrs = collect(self._LLM_CALL_ATTRS, data)
if data.tools:
attrs[LiteLLM.TOOLS_DECLARED] = len(data.tools)
attrs.update(
tool_definition_attrs(
lambda idx, suffix: f"gen_ai.tool.{idx}.{suffix}",
data.tools,
self._TOOL_ATTRS,
self._tool_attr_budget,
)
)
)
return attrs
@classmethod

View file

@ -12,7 +12,11 @@ Like ``GenAIMapper``, each span kind declares its schema as a flat
from typing import Callable, Final
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, SpanData
from litellm.integrations.otel.mappers.utils import collect, drop_none
from litellm.integrations.otel.mappers.utils import (
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
collect,
tool_definition_attrs,
)
from litellm.integrations.otel.model.payloads import (
LLMCallSpanData,
ServiceSpanData,
@ -63,6 +67,9 @@ class LegacyMapper:
_LEGACY_ERROR: lambda d: d.error.message if d.error is not None and d.error.message else None,
}
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
self._tool_attr_budget = tool_attr_budget
def map(self, data: SpanData) -> AttributeMap:
match data:
case LLMCallSpanData():
@ -72,16 +79,14 @@ class LegacyMapper:
case _:
return {}
@classmethod
def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap:
attrs = collect(cls._LLM_CALL_ATTRS, data)
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
attrs = collect(self._LLM_CALL_ATTRS, data)
attrs.update(
drop_none(
{
f"llm.request.functions.{idx}.{suffix}": extract(tool)
for idx, tool in enumerate(data.tools)
for suffix, extract in cls._TOOL_ATTRS.items()
}
tool_definition_attrs(
lambda idx, suffix: f"llm.request.functions.{idx}.{suffix}",
data.tools,
self._TOOL_ATTRS,
self._tool_attr_budget,
)
)
return attrs

View file

@ -13,9 +13,11 @@ from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue, Span
from litellm.integrations.otel.mappers.utils import (
collect,
drop_none,
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN,
json_if,
message_content,
output_messages,
tool_definition_attrs,
)
from litellm.integrations.otel.model.payloads import (
LLMCallSpanData,
@ -70,6 +72,9 @@ class OpenInferenceMapper:
),
}
def __init__(self, tool_attr_budget: int = MAX_TOOL_DEFINITION_ATTRS_PER_SPAN) -> None:
self._tool_attr_budget = tool_attr_budget
def map(self, data: SpanData) -> AttributeMap:
match data:
case LLMCallSpanData():
@ -77,14 +82,13 @@ class OpenInferenceMapper:
case _:
return {}
@classmethod
def _llm_call(cls, data: LLMCallSpanData) -> AttributeMap:
def _llm_call(self, data: LLMCallSpanData) -> AttributeMap:
return {
**collect(cls._LLM_CALL_ATTRS, data),
**collect(cls._BLOB_ATTRS, data),
**cls._messages("llm.input_messages", "input.value", data.messages_in),
**cls._messages("llm.output_messages", "output.value", output_messages(data)),
**cls._tools(data),
**collect(self._LLM_CALL_ATTRS, data),
**collect(self._BLOB_ATTRS, data),
**self._messages("llm.input_messages", "input.value", data.messages_in),
**self._messages("llm.output_messages", "output.value", output_messages(data)),
**self._tools(data),
}
@staticmethod
@ -108,12 +112,10 @@ class OpenInferenceMapper:
attrs[value_key] = json.dumps([{"role": role, "content": content} for role, content in parsed])
return attrs
@classmethod
def _tools(cls, data: LLMCallSpanData) -> AttributeMap:
return drop_none(
{
f"llm.tools.{idx}.{suffix}": extract(tool)
for idx, tool in enumerate(data.tools)
for suffix, extract in cls._TOOL_ATTRS.items()
}
def _tools(self, data: LLMCallSpanData) -> AttributeMap:
return tool_definition_attrs(
lambda idx, suffix: f"llm.tools.{idx}.{suffix}",
data.tools,
self._TOOL_ATTRS,
self._tool_attr_budget,
)

View file

@ -6,10 +6,34 @@ they live in one place.
"""
import json
from typing import Callable, Mapping, Sequence
from typing import Callable, Final, Mapping, Sequence
from litellm.integrations.otel.mappers.base import AttributeMap, AttrValue
from litellm.integrations.otel.model.payloads import LLMCallSpanData
from litellm.integrations.otel.model.payloads import LLMCallSpanData, ToolDefinition
DEFAULT_SPAN_ATTRIBUTE_LIMIT: Final = 128
"""The OTel SDK's default per-span attribute count limit."""
MAX_TOOL_DEFINITION_ATTRS_PER_SPAN: Final = DEFAULT_SPAN_ATTRIBUTE_LIMIT // 4
"""Span-wide ceiling on attributes spent spelling out declared tool definitions.
Tool definitions are an unbounded attribute family: one entry per declared
tool, per field, per active vocabulary. Agentic clients declare hundreds, which
overruns the span attribute limit. That limit evicts oldest-first, so an
uncapped family silently destroys the core ``gen_ai.*`` attributes written
before it.
The ceiling is span-wide rather than per-mapper because several vocabularies
can be active at once and each spells the same tools out under its own keys, so
a per-mapper allowance multiplies by the number of vocabularies and reaches the
limit again. Reserving a quarter of the span for tool detail leaves the rest to
core telemetry no matter how many vocabularies are configured.
"""
def tool_attr_budget(vocabularies: int) -> int:
"""Split the span-wide tool-definition ceiling across active vocabularies."""
return MAX_TOOL_DEFINITION_ATTRS_PER_SPAN // max(vocabularies, 1)
def drop_none(values: Mapping[str, AttrValue | None]) -> AttributeMap:
@ -17,6 +41,29 @@ def drop_none(values: Mapping[str, AttrValue | None]) -> AttributeMap:
return {k: v for k, v in values.items() if v is not None}
def tool_definition_attrs(
key_for: Callable[[int, str], str],
tools: Sequence[ToolDefinition],
extractors: Mapping[str, Callable[[ToolDefinition], AttrValue | None]],
attr_budget: int,
) -> AttributeMap:
"""Per-index attributes for as many tools as ``attr_budget`` affords.
``key_for`` builds a vocabulary's key from the tool's index and the field
name, so each mapper keeps its own naming while sharing the budget. One tool
always keeps its detail, so the family stays legible even when many
vocabularies split the ceiling.
"""
max_tools = max(attr_budget // max(len(extractors), 1), 1)
return drop_none(
{
key_for(idx, suffix): extract(tool)
for idx, tool in enumerate(tools[:max_tools])
for suffix, extract in extractors.items()
}
)
def collect(table: Mapping[str, Callable], source: object) -> AttributeMap:
"""Apply an extractor table to ``source``, dropping ``None`` results."""
return drop_none({key: extract(source) for key, extract in table.items()})

View file

@ -6,17 +6,30 @@ without a semconv equivalent lives under the ``litellm.*`` vendor namespace.
from enum import Enum
from typing import Final
from litellm._logging import verbose_logger
class GenAIOperation(str, Enum):
"""Values for ``gen_ai.operation.name``."""
"""Values for ``gen_ai.operation.name``.
The first block is the convention's own vocabulary. The ``LITELLM_`` members
are vendor values for operations the convention names nothing for; its note
on this attribute directs instrumentation to use a system-specific name in
exactly that case, the same allowance :func:`resolve_provider` relies on for
unmapped providers. They stay under the ``litellm.`` prefix so a value the
convention adds later can never collide with one of ours.
"""
CHAT = "chat"
TEXT_COMPLETION = "text_completion"
EMBEDDINGS = "embeddings"
GENERATE_CONTENT = "generate_content"
RETRIEVAL = "retrieval" # vector-store search / RAG query spans
CREATE_AGENT = "create_agent" # reserved for future agent spans
INVOKE_AGENT = "invoke_agent" # reserved for future agent spans
INVOKE_AGENT = "invoke_agent" # agent (A2A) message spans
EXECUTE_TOOL = "execute_tool" # MCP tool-call spans
LITELLM_VECTOR_STORE_MANAGEMENT = "litellm.vector_store_management"
LITELLM_VECTOR_STORE_FILE_MANAGEMENT = "litellm.vector_store_file_management"
class GenAIProvider(str, Enum):
@ -49,11 +62,17 @@ class MCPMethod(str, Enum):
class GenAI:
"""Canonical OTel GenAI span-attribute keys."""
"""Canonical OTel GenAI attribute keys.
``SYSTEM`` is the one exception: the convention deprecated it in favor of
``PROVIDER_NAME``, and it survives here only so already-shipped series keep
resolving for consumers that query it. Nothing new should use it.
"""
# request
OPERATION_NAME: Final = "gen_ai.operation.name"
PROVIDER_NAME: Final = "gen_ai.provider.name"
SYSTEM: Final = "gen_ai.system"
REQUEST_MODEL: Final = "gen_ai.request.model"
REQUEST_TEMPERATURE: Final = "gen_ai.request.temperature"
REQUEST_TOP_P: Final = "gen_ai.request.top_p"
@ -233,6 +252,7 @@ class LiteLLM:
# ``litellm_params.model``), distinct from the user-facing ``gen_ai.request.model``.
PROVIDER_MODEL: Final = "litellm.provider.model"
REQUEST_STREAMING: Final = "litellm.request.streaming"
TOOLS_DECLARED: Final = "litellm.request.tools.declared"
GUARDRAIL_NAME: Final = "litellm.guardrail.name"
GUARDRAIL_MODE: Final = "litellm.guardrail.mode"
GUARDRAIL_STATUS: Final = "litellm.guardrail.status"
@ -257,13 +277,27 @@ class LiteLLM:
class Metric:
"""GenAI metric instrument names."""
"""GenAI metric instrument names.
Every name here that a convention or a backend defines uses that name, so a
consumer charting GenAI telemetry finds litellm's series where it looks for
them. ``TOKEN_USAGE``, ``OPERATION_DURATION``, ``TIME_TO_FIRST_TOKEN`` and
``TIME_PER_OUTPUT_TOKEN`` are semconv instruments, defined in the GenAI
conventions; the ``gen_ai.client.response.*`` spellings litellm used for the
latter two are not conventions at all, so nothing downstream could chart
them. Cost has no semconv instrument, so it takes ``gen_ai.usage.cost``, the
name backends already query for spend.
``RESPONSE_DURATION`` keeps its vendor spelling deliberately: the closest
convention, ``gen_ai.server.request.duration``, would collide in meaning with
``OPERATION_DURATION``, which litellm already emits for the whole operation.
"""
TOKEN_USAGE: Final = "gen_ai.client.token.usage"
OPERATION_DURATION: Final = "gen_ai.client.operation.duration"
TOKEN_COST: Final = "gen_ai.client.token.cost"
TIME_TO_FIRST_TOKEN: Final = "gen_ai.client.response.time_to_first_token"
TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.client.response.time_per_output_token"
TOKEN_COST: Final = "gen_ai.usage.cost"
TIME_TO_FIRST_TOKEN: Final = "gen_ai.server.time_to_first_token"
TIME_PER_OUTPUT_TOKEN: Final = "gen_ai.server.time_per_output_token"
RESPONSE_DURATION: Final = "gen_ai.client.response.duration"
@ -301,6 +335,35 @@ _OPERATION_BY_CALL_TYPE: dict[str, GenAIOperation] = {
"responses": GenAIOperation.CHAT,
"aresponses": GenAIOperation.CHAT,
"call_mcp_tool": GenAIOperation.EXECUTE_TOOL,
"vector_store_search": GenAIOperation.RETRIEVAL,
"avector_store_search": GenAIOperation.RETRIEVAL,
"query": GenAIOperation.RETRIEVAL,
"aquery": GenAIOperation.RETRIEVAL,
"send_message": GenAIOperation.INVOKE_AGENT,
"asend_message": GenAIOperation.INVOKE_AGENT,
"asend_message_streaming": GenAIOperation.INVOKE_AGENT,
"vector_store_create": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"avector_store_create": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"vector_store_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"avector_store_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"vector_store_list": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"avector_store_list": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"vector_store_update": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"avector_store_update": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"vector_store_delete": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"avector_store_delete": GenAIOperation.LITELLM_VECTOR_STORE_MANAGEMENT,
"vector_store_file_create": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"avector_store_file_create": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"vector_store_file_list": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"avector_store_file_list": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"vector_store_file_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"avector_store_file_retrieve": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"vector_store_file_content": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"avector_store_file_content": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"vector_store_file_update": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"avector_store_file_update": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"vector_store_file_delete": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
"avector_store_file_delete": GenAIOperation.LITELLM_VECTOR_STORE_FILE_MANAGEMENT,
}
@ -317,7 +380,21 @@ def resolve_provider(custom_llm_provider: str | None) -> str:
def resolve_operation(call_type: str | None) -> GenAIOperation:
"""Map a litellm ``call_type`` to a ``gen_ai.operation.name`` value."""
"""Map a litellm ``call_type`` to a ``gen_ai.operation.name`` value.
An unmapped call type still falls back to ``chat`` so every series keeps an
operation label, but it logs at debug rather than falling through silently:
a new call type mislabelled as ``chat`` mixes its latency and cost into
everyone's chat charts, which is invisible until someone reads the numbers.
"""
if not call_type:
return GenAIOperation.CHAT
return _OPERATION_BY_CALL_TYPE.get(call_type.lower(), GenAIOperation.CHAT)
mapped = _OPERATION_BY_CALL_TYPE.get(call_type.lower())
if mapped is not None:
return mapped
verbose_logger.debug(
"otel: call_type %r has no gen_ai.operation.name mapping; labelling it %r. Add it to _OPERATION_BY_CALL_TYPE.",
call_type,
GenAIOperation.CHAT.value,
)
return GenAIOperation.CHAT

View file

@ -18,12 +18,14 @@ before the LLM call even starts), so a guardrail is a sibling of the LLM call,
not a child of it. The emitter parents every span to the ambient OTel context
(the active server span), which matches this.
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) are intentionally NOT in this
tree. Per the OTel GenAI MCP semconv, MCP and the HTTP transport are independent
contexts, so an MCP span parents to the trace context the client propagated in
``params._meta`` (or starts its own root when none is propagated) and records the
``PROXY_REQUEST`` transport span as a span *link*, never a parent. The registry
encodes this as ``parent=None, links=PROXY_REQUEST``.
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) have two shapes, chosen at emit
time by :func:`resolve_mcp_span_context`. When the client propagates trace context
in ``params._meta`` MCP and the HTTP transport are independent contexts per the
OTel GenAI MCP semconv, so the span parents to that propagated context and records
the ``PROXY_REQUEST`` transport span as a span *link*, never a parent — the shape
this registry's ``parent=None, links=PROXY_REQUEST`` entry encodes. When nothing is
propagated (the common case) the span nests under the transport span of the request
carrying that message, so the tool call stays in one trace.
Not every service call becomes a span — :func:`span_role_for_service` decides:
@ -89,12 +91,13 @@ class SpanSpec:
SPAN_REGISTRY: dict[SpanRole, SpanSpec] = {
SpanRole.PROXY_REQUEST: SpanSpec(SpanRole.PROXY_REQUEST, LiteLLMSpanKind.SERVER, parent=None),
SpanRole.LLM_CALL: SpanSpec(SpanRole.LLM_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
# MCP and the HTTP transport are independent contexts (OTel GenAI MCP semconv),
# so an MCP span does not nest under the transport span. The proxy is an MCP
# client to the upstream server, so it's a CLIENT span; it parents to the trace
# context the client propagated in ``params._meta`` (or starts its own root when
# none is propagated) and records the PROXY_REQUEST transport span as a span
# *link*, never a parent — hence ``parent=None, links=PROXY_REQUEST``.
# The proxy is an MCP client to the upstream server, so MCP spans are CLIENT
# spans. With trace context propagated in ``params._meta``, MCP and the HTTP
# transport are independent contexts (OTel GenAI MCP semconv): the span parents
# to the propagated context and records the PROXY_REQUEST transport span as a
# span *link*, never a parent — the shape ``parent=None, links=PROXY_REQUEST``
# encodes. With nothing propagated, ``resolve_mcp_span_context`` nests the span
# under that message's transport span instead, keeping the call in one trace.
SpanRole.MCP_TOOL_CALL: SpanSpec(
SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=None, links=SpanRole.PROXY_REQUEST
),

View file

@ -1,9 +1,11 @@
"""Shared, OpenTelemetry-free helpers for the otel integration.
Generic value coercion (for reading heterogeneous logging-payload dicts), time
conversion, and header parsing — pulled out of the individual modules so they
live in one place. Deliberately free of any ``opentelemetry`` import so the
OTel-free sources of truth (payloads, semconv, spans, config) can use it too.
Generic value coercion (for reading heterogeneous logging-payload dicts) and
time conversion — pulled out of the individual modules so they live in one
place. Deliberately free of any ``opentelemetry`` import so the OTel-free
sources of truth (payloads, semconv, spans, config) can use it too. OTLP header
parsing lives in :mod:`litellm.integrations.otel.plumbing.providers` instead,
because it delegates to the OTel SDK's own W3C Baggage parser.
"""
from datetime import datetime
@ -89,15 +91,3 @@ def to_seconds(value: datetime | float | int | str | None) -> float | None:
except ValueError:
continue
return None
def parse_headers(raw: str | None) -> dict[str, str]:
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict."""
headers: dict[str, str] = {}
if not raw:
return headers
for pair in raw.split(","):
if "=" in pair:
key, _, value = pair.partition("=")
headers[key.strip()] = value.strip()
return headers

View file

@ -5,7 +5,14 @@ from typing import Mapping
from opentelemetry import baggage
from opentelemetry.context import Context, get_current
from opentelemetry.trace import Link, Span, get_current_span, set_span_in_context
from opentelemetry.trace import (
Link,
NonRecordingSpan,
Span,
SpanContext,
get_current_span,
set_span_in_context,
)
from opentelemetry.trace.propagation.tracecontext import (
TraceContextTextMapPropagator,
)
@ -72,6 +79,86 @@ def reset_mcp_message_trace_carrier(token: "Token[Mapping[str, str] | None]") ->
_mcp_message_trace_carrier.reset(token)
# The transport span of the HTTP request carrying the CURRENT MCP message.
#
# ``_request_root_span`` above cannot be used for MCP: a *stateful* streamable-HTTP
# session runs every message on the single task spawned by that session's
# ``initialize`` POST, so the ContextVar the ASGI request task writes at auth time
# is frozen at ``initialize`` there and never sees the later ``tools/call`` POSTs.
# Reading it from the message handler parents every tool call in the session to the
# first request's server span and aims that call's ``error.*`` at it — a span that
# ended long ago, so the SDK drops the write and the failure reaches no request at
# all. The gateway instead resolves the current message's transport span on the
# request task and hands it over the same way it hands over per-request auth, and
# the handler publishes it here for the span emitter and the failure hook.
_mcp_message_transport_span: "ContextVar[Span | None]" = ContextVar(
"litellm_otel_mcp_message_transport_span", default=None
)
def set_mcp_message_transport_span(span: object) -> "Token[Span | None]":
"""Publish the transport span of the request carrying the current MCP message.
Also re-anchors the request root, so everything else the message emits or stamps
— the identity attributes seeded onto the server span, a guardrail span, a
proxy-level failure — lands on this request instead of on the one that opened
the session. The MCP SDK dispatches each message on its own task, so the anchor
is scoped to this message; the handler re-publishes it for the next one either
way. Only a transport still open for writes is anchored: replacing the anchor
with a request that already answered would just move the dropped writes from one
finished span to another.
Takes ``object`` because the gateway reads it back out of the ASGI scope, whose
values are untyped; anything that is not a usable span is stored as ``None``
rather than trusted.
Returns the reset token; the caller must reset it once the message is handled
so the transport never leaks to the next message on the same session task.
"""
transport = span if isinstance(span, Span) and is_recordable_span(span) else None
if transport is not None and transport.is_recording():
set_request_root_span(transport)
return _mcp_message_transport_span.set(transport)
def reset_mcp_message_transport_span(token: "Token[Span | None]") -> None:
_mcp_message_transport_span.reset(token)
def mcp_message_transport_span() -> "Span | None":
"""The published transport span, only while it is still open for writes.
Recording — not merely valid — is the bar here because this span is the target
of ``error.*`` stamping from another task, and the publisher's validity check
cannot speak for a span that has since ended. A finished span keeps a valid
context forever, so it would otherwise be handed back for a write the SDK then
refuses. The POST carrying a ``tools/call`` stays open until the result is
written, so it is recording for the life of the call; a notification POST can
answer first, and this returns ``None`` for it rather than writing into the void.
"""
span = _mcp_message_transport_span.get()
if span is None or not span.is_recording():
return None
return span
def _mcp_transport_span_context() -> "SpanContext | None":
"""The transport span an MCP message span should attach to.
Prefers the transport the gateway published for this specific message; falls
back to the ambient request anchor for paths that emit an MCP span on the
request task itself (the REST MCP endpoints, the SDK). Parenting and linking
only need the immutable context, and unlike ``mcp_message_transport_span`` they
stay correct against a transport that has already finished, so this does not
require the span to still be recording.
"""
published = _mcp_message_transport_span.get()
if published is not None:
return published.get_span_context()
span = request_root_span()
return span.get_span_context() if span is not None else None
def set_request_baggage(values: Mapping[str, str], context: Context | None = None) -> Context:
"""Return a context with ``values`` written into Baggage."""
ctx = context
@ -132,33 +219,44 @@ def resolve_request_span_context() -> Context:
def resolve_mcp_span_context(
carrier: "Mapping[str, str] | None" = None,
) -> "tuple[Context, tuple[Link, ...]]":
"""Parent context + links for an MCP message span, per the OTel GenAI MCP semconv.
"""Parent context + links for an MCP message span.
MCP and the underlying transport (HTTP) are independent lifecycles — one
streamable-HTTP session multiplexes many messages, so nesting the message span
under the HTTP/session span is wrong (it renders the message at the session's
start, skewed by however long the session has been open). Instead:
When the client propagates W3C trace context in the request's ``params._meta``
(SEP-414), MCP and the underlying transport are independent lifecycles — one
streamable-HTTP session multiplexes many messages, and the client's own span is
the truthful parent. So, per the OTel GenAI MCP semconv:
* parent to the trace context the client propagated in the request's
``params._meta`` (a *remote* parent), and
* record the transport/session span as a *link*, never the parent.
* parent to the trace context the client propagated (a *remote* parent), and
* record the transport span as a *link*, never the parent.
Almost no client implements SEP-414 yet, so in practice nothing is propagated.
Rooting the span there splits a single tool call into two disconnected traces
joined only by a link, which is how it surfaces in APM: the ``POST`` transaction
and the ``tools/call`` span share no trace. With no remote parent to honor,
parent to the transport span of the request carrying this message instead, so
the call stays in one trace; no link is added since the transport is now the
real parent. The transport comes from :func:`_mcp_transport_span_context`, which
is the *current message's* POST rather than whatever request happened to open
the session, so a long-lived session does not glue every message under its
first request. With neither a remote parent nor a transport the returned context
carries no span and the span legitimately starts its own root trace.
Only trace context (``traceparent``/``tracestate``) is extracted, never the
client's W3C Baggage: ``params._meta`` is caller-controlled, and the otel
baggage processor stamps allowlisted baggage keys (``litellm.team.id``,
``litellm.metadata.*``, ...) onto the span as attributes, so honoring remote
baggage would let a client spoof a span's identity attribution.
With no propagated context the returned context carries no span, so the span
starts its own root trace (still linked to the transport). The base context is
explicitly empty so an absent ``traceparent`` can never fall through to the
ambient (stale session) span.
baggage would let a client spoof a span's identity attribution. The base context
for extraction is explicitly empty so an absent or malformed ``traceparent`` can
never fall through to the ambient (stale session) span.
"""
source = carrier if carrier is not None else _mcp_message_trace_carrier.get()
parent = _PROPAGATOR.extract(dict(source or {}), context=Context())
transport = request_root_span()
links = (Link(transport.get_span_context()),) if transport is not None else ()
return parent, links
transport = _mcp_transport_span_context()
if is_recordable_span(get_current_span(parent)):
return parent, (Link(transport),) if transport is not None else ()
if transport is not None:
return context_from_span(NonRecordingSpan(transport)), ()
return parent, ()
def is_recordable_span(obj: object) -> bool:

View file

@ -1,6 +1,6 @@
"""GenAI client metrics: the six ``gen_ai.client.*`` histograms plus the
recorder that builds attributes, applies the shared cardinality filter, and
records a request's metrics in the success path.
records a request's metrics on both the success and the failure path.
The instrument names/units/descriptions and the recording + timing math mirror
the v1 :mod:`litellm.integrations.opentelemetry` integration so both engines emit
@ -10,11 +10,12 @@ identical metrics. The attribute cardinality filter is reused from v1 by import
from dataclasses import dataclass
from datetime import datetime
from typing import Any, FrozenSet, Mapping, Optional
from typing import Any, Final, FrozenSet, Mapping, Optional, TypeAlias
from opentelemetry.metrics import Histogram, Meter
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.opentelemetry import (
METRIC_METADATA_KEYS,
TOKEN_TYPE_ATTRIBUTE,
@ -22,11 +23,34 @@ from litellm.integrations.opentelemetry import (
_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.semconv import (
Error,
GenAI,
Metric,
resolve_operation,
resolve_provider,
)
from litellm.integrations.otel.model.utils import to_seconds
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
def _provider_attributes(custom_llm_provider: object) -> Mapping[str, str]:
"""The provider labels for one call's metrics.
``gen_ai.provider.name`` carries the semconv-mapped value; the deprecated
``gen_ai.system`` spelling is dual-emitted with the raw litellm provider
string it has always carried, so a dashboard already querying it keeps
matching. A call with no provider gets neither label: a placeholder value
would mint a permanent series that no operator can act on.
"""
if not isinstance(custom_llm_provider, str) or not custom_llm_provider:
return {}
return {
GenAI.PROVIDER_NAME: resolve_provider(custom_llm_provider),
GenAI.SYSTEM: custom_llm_provider,
}
@dataclass(frozen=True)
class GenAIMetrics:
operation_duration: Histogram
@ -72,8 +96,83 @@ def create_genai_metrics(meter: Meter) -> GenAIMetrics:
)
# A metric datapoint's attributes. Values are the strings the recorder builds, except
# the request model, which is whatever the caller passed and may be absent.
MetricAttributes: TypeAlias = Mapping[str, "str | None"]
ERROR_TYPE_FALLBACK: Final = "_OTHER"
# Every attribute a metric datapoint may carry, on either path. A label value that
# is unique per request is a new time series that will never be written to again, so
# this set is what keeps the series count bounded by the deployment's own
# key/team/user/deployment count rather than by its traffic. Each entry is a fixed
# enum or an operator-provisioned identifier.
#
# Deliberately excluded is everything the *client* supplies or that moves per
# request: ``metadata.requester_metadata`` and ``metadata.spend_logs_metadata`` (both
# free-form from the request body), ``metadata.user_api_key_end_user_id`` (the body's
# ``user`` field), and ``metadata.requester_ip_address``. Those stay on the span,
# where cardinality is free and where they already are.
# ``metadata.user_api_key_user_email`` is left out too: it is bounded, but it is PII
# duplicating the user id already here.
#
# This is a CEILING, applied before the operator's own include/exclude filter, so an
# operator can narrow it but never widen it back to an unbounded attribute.
METRIC_ATTRIBUTE_CEILING: Final[frozenset[str]] = frozenset(
(
"gen_ai.operation.name",
"gen_ai.provider.name",
"gen_ai.system",
"gen_ai.request.model",
"gen_ai.framework",
"metadata.user_api_key_hash",
"metadata.user_api_key_alias",
"metadata.user_api_key_team_id",
"metadata.user_api_key_team_alias",
"metadata.user_api_key_org_id",
"metadata.user_api_key_user_id",
"hidden_params",
)
)
# The only ``hidden_params`` field that becomes part of the ``hidden_params`` label.
# The object as a whole is per-request by construction -- ``response_cost``,
# ``litellm_overhead_time_ms``, ``cache_key``, ``usage_object`` and the provider's
# ``additional_headers`` rate-limit counters all move on every call -- so dumping it
# whole made one series per request out of every instrument.
#
# ``model_id`` is the router's own deployment id, so it is bounded by the deployment
# list and is what a per-deployment panel joins on. ``api_base`` is deliberately NOT
# here even though it names the same thing: it is a documented per-call parameter, so
# in SDK use it is chosen by the caller rather than provisioned by the operator, and a
# caller varying it would put the per-request cardinality straight back.
BOUNDED_HIDDEN_PARAM_KEYS: Final[tuple[str, ...]] = ("model_id",)
def resolve_error_type(kwargs: Mapping[str, Any]) -> str:
"""The ``error.type`` value for a failed request.
Bounded by construction: the mapped provider exception's class name (the same
``error_information.error_class`` the failure span stamps), else the provider
status code, else the raw exception's class name, else ``_OTHER`` — the value
the convention reserves for a failure the instrumentation cannot classify. The
exception *message* is unbounded and never becomes a label; it stays on the
span and its exception event, where high cardinality is free.
"""
std_log = kwargs.get("standard_logging_object")
info = getattr(std_log, "error_information", None) or (std_log or {}).get("error_information") or {}
error_class = info.get("error_class") or info.get("error_code")
if error_class:
return str(error_class)
exception = kwargs.get("exception")
if exception is not None:
return type(exception).__name__
return ERROR_TYPE_FALLBACK
class GenAIMetricRecorder:
"""Records the six GenAI histograms for one successful LLM call.
"""Records the six GenAI histograms for one successful LLM call, and the
duration histogram alone for one failed LLM call (see :meth:`record_failure`).
The cardinality filter is resolved lazily on the first record: the proxy
populates ``callback_settings.otel.attributes`` after the logger is built, so
@ -96,7 +195,7 @@ class GenAIMetricRecorder:
start_time: datetime,
end_time: datetime,
) -> None:
common_attrs = self._filter_attributes(self._common_attributes(kwargs))
common_attrs = self._filter_attributes(self._bounded_attributes(kwargs))
duration_s = (end_time - start_time).total_seconds()
self._metrics.operation_duration.record(duration_s, attributes=common_attrs)
@ -110,17 +209,48 @@ class GenAIMetricRecorder:
self._record_time_per_output_token(kwargs, response_obj, end_time, duration_s, common_attrs)
self._record_response_duration(kwargs, end_time, common_attrs)
def record_failure(
self,
kwargs: Mapping[str, Any],
start_time: datetime,
end_time: datetime,
) -> None:
"""Record the one metric a failed request can honestly report: the
operation's duration, tagged with ``error.type``.
The other five instruments all describe a completed generation and have
nothing to measure here. litellm hands the failure callback no
``response_obj`` at all, so there is no usage to split into input/output
tokens and no completion-token count to divide generation time by; it also
zeroes ``response_cost`` on failure. Recording them anyway would put a
fabricated zero into series that dashboards average.
The attribute set is :data:`METRIC_ATTRIBUTE_CEILING`, the same cap the
success path uses. A failure needs no provider spend, so a caller who can put
a unique value into a client-supplied attribute could mint one histogram
series per request for free; the cap is what makes that impossible on either
path.
``error.type`` is stamped after both filters, exactly like
``gen_ai.token.type``, so an operator's include/exclude list cannot strip
the discriminator and silently merge failures back into the success series.
"""
attributes = {
**self._filter_attributes(self._bounded_attributes(kwargs)),
Error.TYPE: resolve_error_type(kwargs),
}
self._metrics.operation_duration.record((end_time - start_time).total_seconds(), attributes=attributes)
# ------------------------------------------------------------------ #
# Attribute building + cardinality filter
# ------------------------------------------------------------------ #
def _common_attributes(self, kwargs: Mapping[str, Any]) -> dict:
params = kwargs.get("litellm_params") or {}
provider = params.get("custom_llm_provider", "Unknown")
common_attrs: dict = {
"gen_ai.operation.name": resolve_operation(kwargs.get("call_type")).value,
"gen_ai.system": provider,
"gen_ai.request.model": kwargs.get("model"),
GenAI.OPERATION_NAME: resolve_operation(kwargs.get("call_type")).value,
**_provider_attributes(params.get("custom_llm_provider")),
GenAI.REQUEST_MODEL: kwargs.get("model"),
"gen_ai.framework": "litellm",
}
@ -136,11 +266,25 @@ class GenAIMetricRecorder:
common_attrs[f"metadata.{key}"] = str(value)
hidden_params = getattr(std_log, "hidden_params", None) or (std_log or {}).get("hidden_params", {})
if hidden_params:
common_attrs["hidden_params"] = safe_dumps(hidden_params)
bounded_hidden_params = {
key: hidden_params[key]
for key in BOUNDED_HIDDEN_PARAM_KEYS
if isinstance(hidden_params, Mapping) and hidden_params.get(key) is not None
}
if bounded_hidden_params:
common_attrs["hidden_params"] = safe_dumps(bounded_hidden_params)
return common_attrs
def _bounded_attributes(self, kwargs: Mapping[str, Any]) -> MetricAttributes:
"""The datapoint attributes, capped at :data:`METRIC_ATTRIBUTE_CEILING`.
The cap runs BEFORE the operator's include/exclude filter so the filter can
only narrow it. An operator who names an excluded attribute in an include
list gets nothing for it rather than reintroducing an unbounded label.
"""
return {k: v for k, v in self._common_attributes(kwargs).items() if k in METRIC_ATTRIBUTE_CEILING}
def _ensure_filter(self) -> None:
if self._filter_resolved:
return
@ -157,8 +301,29 @@ class GenAIMetricRecorder:
# without reconstructing the recorder.
self._include, self._exclude = _resolve_metric_attribute_filter(attributes)
self._filter_resolved = True
self._warn_about_metric_ineligible_names()
def _filter_attributes(self, attrs: dict) -> dict:
def _warn_about_metric_ineligible_names(self) -> None:
"""Say so when the operator's filter names an attribute the ceiling removes.
The shared validator accepts every span attribute name, so a name that is
legal on a span but metric-ineligible would otherwise be a silent no-op: an
``include_list`` naming it emits nothing for it and an ``exclude_list`` naming
it looks like it worked. Logged once, when the filter resolves, rather than
per request.
"""
named = (self._include or frozenset()) | (self._exclude or frozenset())
ineligible = sorted(named - METRIC_ATTRIBUTE_CEILING - {TOKEN_TYPE_ATTRIBUTE})
if ineligible:
verbose_logger.warning(
"OTel metrics: %s cannot be a metric attribute and is being ignored; it varies "
"per request or is client-supplied, so it would make one time series per request. "
"It is still on the span. Metric attributes are limited to: %s",
", ".join(ineligible),
", ".join(sorted(METRIC_ATTRIBUTE_CEILING)),
)
def _filter_attributes(self, attrs: MetricAttributes) -> MetricAttributes:
self._ensure_filter()
if self._include is not None:
return {k: v for k, v in attrs.items() if k in self._include}

View file

@ -29,15 +29,13 @@ from opentelemetry.sdk.trace.export.in_memory_span_exporter import (
InMemorySpanExporter,
)
from opentelemetry.trace import Span, SpanKind, Tracer
from opentelemetry.util.re import parse_env_headers
from litellm._version import version as litellm_version
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
from litellm.integrations.otel.model.semconv import LiteLLM
from litellm.integrations.otel.model.spans import LiteLLMSpanKind
# Re-exported so ``providers.parse_headers`` remains a stable entry point.
from litellm.integrations.otel.model.utils import parse_headers as parse_headers
if TYPE_CHECKING:
from opentelemetry.metrics import Meter
from opentelemetry.sdk.metrics.export import MetricReader
@ -119,6 +117,23 @@ def _otlp_traces_endpoint(endpoint: str | None) -> str | None:
return endpoint + "/v1/traces"
def parse_headers(raw: str | None) -> dict[str, str]:
"""Parse an OTLP ``"k=v,k=v"`` header string into a dict.
``OTEL_EXPORTER_OTLP_HEADERS`` is W3C Baggage encoded per the OTLP spec, so
values are percent-decoded: a vendor that documents
``Authorization=Basic%20<token>`` (Grafana Cloud does, because a bare space
is not representable there) has to reach the exporter as ``Basic <token>``,
not with a literal ``%20`` that the backend rejects as malformed. The SDK's
own parser is used so litellm decodes exactly what the OTLP exporters do
when they read the env var themselves; ``liberal`` keeps values that are not
percent-encoded (``Authorization=Bearer <token>``) working unchanged.
"""
if not raw:
return {}
return dict(parse_env_headers(raw, liberal=True))
def _exporter_from_spec(spec: ExporterSpec) -> SpanExporter:
kind = (spec.kind or "console").lower()
factory = _EXPORTER_FACTORIES.get(kind)
@ -191,6 +206,13 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
``console`` (and any unrecognized kind) exports to the console; ``otlp_http``
and ``otlp_grpc`` export over OTLP with the configured endpoint/headers. The
reader exports on a 5s period, matching v1.
Histograms keep the SDK's default cumulative temporality. Prometheus-backed
OTLP receivers (Grafana Cloud / Mimir, and the Prometheus OTLP endpoint)
reject delta histograms outright with ``invalid temporality and type
combination``, which drops the whole metric batch, while backends that
prefer delta still accept cumulative. The enterprise billing exporter
already relies on the same default.
"""
from opentelemetry.sdk.metrics.export import (
ConsoleMetricExporter,
@ -202,18 +224,12 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
from opentelemetry.exporter.otlp.proto.http.metric_exporter import (
OTLPMetricExporter as HTTPMetricExporter,
)
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import AggregationTemporality
exporter: Any = HTTPMetricExporter(
endpoint=_otlp_metrics_endpoint(config.endpoint),
headers=parse_headers(config.headers),
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
elif kind in ("otlp_grpc", "grpc"):
from opentelemetry.sdk.metrics import Histogram
from opentelemetry.sdk.metrics.export import AggregationTemporality
try:
from opentelemetry.exporter.otlp.proto.grpc.metric_exporter import (
OTLPMetricExporter as GRPCMetricExporter,
@ -227,7 +243,6 @@ def build_metric_reader(config: OpenTelemetryV2Config) -> "MetricReader":
exporter = GRPCMetricExporter(
endpoint=config.endpoint,
headers=parse_headers(config.headers),
preferred_temporality={Histogram: AggregationTemporality.DELTA},
)
else:
exporter = ConsoleMetricExporter()

View file

@ -16,6 +16,7 @@ from typing import (
Dict,
List,
Literal,
Mapping,
Optional,
Sequence,
Tuple,
@ -68,6 +69,15 @@ else:
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0
# Tiers a caller may name in a request, across the providers that accept the
# parameter: OpenAI ("auto", "default", "flex", "priority", "scale"), Bedrock and
# Groq (subsets of those), Anthropic ("auto", "standard_only") and Vertex, which
# maps "default" to "standard". Used to bound the caller-controlled fallback in
# ``get_service_tier_from_standard_logging_payload``.
KNOWN_REQUEST_SERVICE_TIERS = frozenset(
{"auto", "batch", "default", "flex", "priority", "scale", "standard", "standard_only"}
)
def _get_budget_metrics_per_request_timeout() -> float:
raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT")
@ -1244,6 +1254,7 @@ class PrometheusLogger(CustomLogger):
client_ip=standard_logging_payload["metadata"].get("requester_ip_address"),
user_agent=standard_logging_payload["metadata"].get("user_agent"),
stream=(str(standard_logging_payload.get("stream")) if litellm.prometheus_emit_stream_label else None),
service_tier=get_service_tier_from_standard_logging_payload(standard_logging_payload),
)
if user_api_key is not None and isinstance(user_api_key, str) and user_api_key.startswith("sk-"):
@ -1449,6 +1460,8 @@ class PrometheusLogger(CustomLogger):
prompt_details = usage_object.get("prompt_tokens_details") or {}
completion_details = usage_object.get("completion_tokens_details") or {}
cache_creation_detail_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
detail_metrics: List[Tuple[Any, DEFINED_PROMETHEUS_METRICS, Any]] = [
(
self.litellm_input_cached_tokens_metric,
@ -1458,7 +1471,7 @@ class PrometheusLogger(CustomLogger):
(
self.litellm_input_cache_creation_tokens_metric,
"litellm_input_cache_creation_tokens_metric",
(prompt_details.get("cache_creation_tokens") if isinstance(prompt_details, dict) else None),
cache_creation_detail_tokens,
),
(
self.litellm_input_audio_tokens_metric,
@ -1597,27 +1610,12 @@ class PrometheusLogger(CustomLogger):
)
# Provider prompt caching metrics are independent of LiteLLM cache_hit.
provider_cache_read_tokens = 0
provider_cache_creation_tokens = 0
usage_obj = (standard_logging_payload.get("metadata", {}) or {}).get("usage_object")
if isinstance(usage_obj, dict):
# Prefer explicit provider cache fields when available.
_read = usage_obj.get("cache_read_input_tokens")
_write = usage_obj.get("cache_creation_input_tokens")
if isinstance(_read, int):
provider_cache_read_tokens = _read
if isinstance(_write, int):
provider_cache_creation_tokens = _write
# Fallback to prompt_tokens_details.cached_tokens (common normalization point).
# Only fallback when the explicit field is genuinely absent (None).
if _read is None:
prompt_details = usage_obj.get("prompt_tokens_details")
if isinstance(prompt_details, dict):
cached_tokens = prompt_details.get("cached_tokens")
if isinstance(cached_tokens, int):
provider_cache_read_tokens = cached_tokens
(
provider_cache_read_tokens,
provider_cache_creation_tokens,
) = PrometheusLogger._resolve_provider_cache_tokens(usage_obj)
if provider_cache_read_tokens > 0:
PrometheusLogger._inc_labeled_counter(
@ -1639,6 +1637,40 @@ class PrometheusLogger(CustomLogger):
amount=float(provider_cache_creation_tokens),
)
@staticmethod
def _resolve_provider_cache_tokens(usage_obj: Mapping[str, object]) -> tuple[int, int]:
# Prefer explicit provider cache fields when available.
_read = usage_obj.get("cache_read_input_tokens")
_write = usage_obj.get("cache_creation_input_tokens")
provider_cache_read_tokens = _read if isinstance(_read, int) else 0
provider_cache_creation_tokens = _write if isinstance(_write, int) else 0
# Fallback to prompt_tokens_details (common normalization point).
# Only fallback when the explicit field is genuinely absent (None).
prompt_details = usage_obj.get("prompt_tokens_details")
if _read is None and isinstance(prompt_details, dict):
cached_tokens = prompt_details.get("cached_tokens")
if isinstance(cached_tokens, int):
provider_cache_read_tokens = cached_tokens
if _write is None:
write_tokens = PrometheusLogger._resolve_cache_write_tokens(prompt_details)
if write_tokens is not None:
provider_cache_creation_tokens = write_tokens
return provider_cache_read_tokens, provider_cache_creation_tokens
@staticmethod
def _resolve_cache_write_tokens(prompt_details: object) -> int | None:
if not isinstance(prompt_details, dict):
return None
for key in ("cache_write_tokens", "cache_creation_tokens"):
value = prompt_details.get(key)
if isinstance(value, int) and not isinstance(value, bool):
return value
return None
def _increment_mcp_tool_call_metrics(
self,
standard_logging_payload: StandardLoggingPayload,
@ -4076,6 +4108,44 @@ def get_custom_labels_from_metadata(metadata: dict) -> Dict[str, str]:
return result
def get_service_tier_from_standard_logging_payload(
standard_logging_payload: StandardLoggingPayload,
) -> str | None:
"""
Resolve the service tier a request ran on, for the ``service_tier`` label.
The tier the provider actually served wins over the tier the caller asked for,
so latency and spend stay segmentable when the request said ``auto`` and the
provider picked the concrete tier. Providers report the served tier either at
the top level of the response (OpenAI, Bedrock, Groq) or on the usage object
(Anthropic).
Streaming responses carry no served tier, so the requested tier is the
fallback. That value is caller-controlled and survives param mapping even
where the provider then ignores it (Bedrock and Groq accept the request and
drop an unrecognized tier), so it is only labelled when it names a known
tier; otherwise one caller could mint a Prometheus series per string. Values
the provider itself reports are not caller-controlled and stay unrestricted,
so a tier a provider adds later is still labelled correctly.
"""
response = standard_logging_payload.get("response")
usage_object = standard_logging_payload.get("metadata", {}).get("usage_object")
served_candidates: tuple[object, ...] = (
response.get("service_tier") if isinstance(response, dict) else None,
usage_object.get("service_tier") if isinstance(usage_object, dict) else None,
)
served_tier = next((tier for tier in served_candidates if isinstance(tier, str) and tier), None)
if served_tier is not None:
return served_tier
model_parameters = standard_logging_payload.get("model_parameters")
requested_tier = model_parameters.get("service_tier") if isinstance(model_parameters, dict) else None
if isinstance(requested_tier, str) and requested_tier in KNOWN_REQUEST_SERVICE_TIERS:
return requested_tier
return None
def _get_combined_custom_metadata_from_standard_logging_payload(
standard_logging_payload: Optional[dict],
) -> Dict[str, Any]:

View file

@ -24,6 +24,8 @@ class S3Logger:
s3_aws_secret_access_key=None,
s3_aws_session_token=None,
s3_config=None,
s3_server_side_encryption: str | None = None,
s3_sse_kms_key_id: str | None = None,
**kwargs,
):
import boto3
@ -50,11 +52,16 @@ class S3Logger:
s3_aws_session_token = litellm.s3_callback_params.get("s3_aws_session_token")
s3_config = litellm.s3_callback_params.get("s3_config")
s3_path = litellm.s3_callback_params.get("s3_path")
s3_server_side_encryption = litellm.s3_callback_params.get("s3_server_side_encryption")
s3_sse_kms_key_id = litellm.s3_callback_params.get("s3_sse_kms_key_id")
# done reading litellm.s3_callback_params
s3_use_team_prefix = bool(litellm.s3_callback_params.get("s3_use_team_prefix", False))
self.s3_use_team_prefix = s3_use_team_prefix
self.bucket_name = s3_bucket_name
self.s3_path = s3_path
self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params(
s3_server_side_encryption, s3_sse_kms_key_id
)
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
# Create an S3 client with custom endpoint URL
self.s3_client = boto3.client(
@ -136,6 +143,15 @@ class S3Logger:
print_verbose(f"\ns3 Logger - Logging payload = {payload_str}")
sse_params = {
key: value
for key, value in {
"ServerSideEncryption": self.s3_server_side_encryption,
"SSEKMSKeyId": self.s3_sse_kms_key_id,
}.items()
if value
}
response = self.s3_client.put_object(
Bucket=self.bucket_name,
Key=s3_object_key,
@ -144,6 +160,7 @@ class S3Logger:
ContentLanguage="en",
ContentDisposition=f'inline; filename="{s3_object_download_filename}"',
CacheControl="private, immutable, max-age=31536000, s-maxage=0",
**sse_params,
)
print_verbose(f"Response from s3:{str(response)}")
@ -155,6 +172,33 @@ class S3Logger:
pass
def _validated_sse_value(name: str, value: str | None) -> str | None:
if value is None or isinstance(value, str):
return value
verbose_logger.warning(
f"s3 logging: ignoring {name} because it has invalid type {type(value).__name__}; expected a string"
)
return None
def resolve_sse_params(
server_side_encryption: str | None,
sse_kms_key_id: str | None,
) -> tuple[str | None, str | None]:
valid_sse = _validated_sse_value("s3_server_side_encryption", server_side_encryption)
valid_key_id = _validated_sse_value("s3_sse_kms_key_id", sse_kms_key_id)
algorithm = valid_sse or ("aws:kms" if valid_key_id else None)
if algorithm is None:
return None, None
if valid_key_id and not algorithm.startswith("aws:kms"):
verbose_logger.warning(
f"s3 logging: ignoring s3_sse_kms_key_id because s3_server_side_encryption is {algorithm}; "
"set it to aws:kms to encrypt with the KMS key"
)
return algorithm, None
return algorithm, valid_key_id
def get_s3_object_key(
s3_path: str,
prefix: str,

View file

@ -8,13 +8,14 @@ NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to uplo
import asyncio
import time
from collections.abc import Mapping
from datetime import datetime
from typing import List, Optional, cast
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
from litellm.integrations.s3 import get_s3_object_key
from litellm.integrations.s3 import get_s3_object_key, resolve_sse_params
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
@ -55,6 +56,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
s3_server_side_encryption: Optional[str] = None,
s3_sse_kms_key_id: str | None = None,
s3_callback_params_override: Optional[dict] = None,
**kwargs,
):
@ -94,6 +96,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_key_prefix=s3_use_key_prefix,
s3_use_virtual_hosted_style=s3_use_virtual_hosted_style,
s3_server_side_encryption=s3_server_side_encryption,
s3_sse_kms_key_id=s3_sse_kms_key_id,
)
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
@ -148,6 +151,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
s3_server_side_encryption: Optional[str] = None,
s3_sse_kms_key_id: str | None = None,
params_source: Optional[dict] = None,
):
"""
@ -197,10 +201,20 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style
)
self.s3_server_side_encryption = params.get("s3_server_side_encryption") or s3_server_side_encryption
self.s3_server_side_encryption, self.s3_sse_kms_key_id = resolve_sse_params(
params.get("s3_server_side_encryption") or s3_server_side_encryption,
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
)
return
def _sse_headers(self) -> Mapping[str, str]:
candidates = {
"x-amz-server-side-encryption": self.s3_server_side_encryption,
"x-amz-server-side-encryption-aws-kms-key-id": self.s3_sse_kms_key_id,
}
return {key: value for key, value in candidates.items() if value}
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
await self._async_log_event_base(
kwargs=kwargs,
@ -335,11 +349,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
"Content-Language": "en",
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**(
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
if self.s3_server_side_encryption
else {}
),
**self._sse_headers(),
}
req = requests.Request("PUT", url, data=json_string, headers=headers)
prepped = req.prepare()
@ -510,11 +520,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
"Content-Language": "en",
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**(
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
if self.s3_server_side_encryption
else {}
),
**self._sse_headers(),
}
req = requests.Request("PUT", url, data=json_string, headers=headers)
prepped = req.prepare()

View file

@ -195,6 +195,25 @@ def get_metadata_variable_name_from_kwargs(
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
def get_or_create_metadata_bucket(
request_data: dict,
) -> tuple[Literal["metadata", "litellm_metadata"], dict]:
"""
Return the proxy-internal metadata bucket for this request, creating it if absent.
Batch/file routes store proxy state in ``litellm_metadata`` so the OpenAI
``metadata`` field can remain provider-safe (string values only). Every writer and
reader of proxy-internal metadata resolves the bucket through here, so a caller that
supplies its own ``metadata`` field cannot split them across two dicts.
"""
metadata_key = get_metadata_variable_name_from_kwargs(request_data)
metadata_bucket = request_data.get(metadata_key)
if not isinstance(metadata_bucket, dict):
metadata_bucket = {}
request_data[metadata_key] = metadata_bucket
return metadata_key, metadata_bucket
def get_litellm_metadata_from_kwargs(kwargs: dict):
"""
Helper to get litellm metadata from all litellm request kwargs

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