Merge pull request #35285 from BerriAI/litellm_internal_staging
Some checks failed
CodeQL / Analyze (actions) (push) Waiting to run
CodeQL / Analyze (javascript-typescript) (push) Waiting to run
CodeQL / Analyze (python) (push) Waiting to run
CodSpeed Benchmarks / benchmarks (push) Waiting to run
Helm unit test / unit-test (push) Waiting to run
Scorecard supply-chain security / Scorecard analysis (push) Waiting to run
GitHub Actions Security Analysis / zizmor (push) Waiting to run
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

chore(ci): promote internal staging to main
This commit is contained in:
yuneng-jiang 2026-07-30 16:08:39 -07:00 • committed by GitHub
commit 122f9359ce
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
340 changed files with 22822 additions and 5466 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

@ -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

@ -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

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

@ -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

@ -1,12 +1,12 @@
{
"reportAny": {
"limit": 34906
"limit": 31903
},
"reportArgumentType": {
"limit": 2701
"limit": 2645
},
"reportAssignmentType": {
"limit": 330
"limit": 329
},
"reportAttributeAccessIssue": {
"limit": 516
@ -18,13 +18,13 @@
"limit": 59
},
"reportDeprecated": {
"limit": 326
"limit": 325
},
"reportDuplicateImport": {
"limit": 42
},
"reportExplicitAny": {
"limit": 10230
"limit": 10214
},
"reportFunctionMemberAccess": {
"limit": 11
@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5893
"limit": 5869
},
"reportMissingTypeArgument": {
"limit": 15886
"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": 45870
"limit": 45366
},
"reportUnknownLambdaType": {
"limit": 113
},
"reportUnknownMemberType": {
"limit": 40525
"limit": 40477
},
"reportUnknownParameterType": {
"limit": 20384
"limit": 20338
},
"reportUnknownVariableType": {
"limit": 32099
"limit": 32047
},
"reportUnnecessaryCast": {
"limit": 177
},
"reportUnnecessaryComparison": {
"limit": 1023
"limit": 1021
},
"reportUnnecessaryContains": {
"limit": 7
},
"reportUnnecessaryIsInstance": {
"limit": 1206
"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

@ -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

@ -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

@ -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,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

@ -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

@ -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

@ -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",
)
@ -2980,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,
)
@ -3014,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)
@ -3032,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

@ -281,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(
@ -304,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

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

@ -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

@ -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

@ -69,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")
@ -1245,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-"):
@ -4098,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

@ -4680,6 +4680,7 @@ class StandardLoggingPayloadSetup:
applied_guardrails=applied_guardrails,
mcp_tool_call_metadata=mcp_tool_call_metadata,
vector_store_request_metadata=vector_store_request_metadata,
routing_decision=None,
usage_object=usage_object,
requester_custom_headers=None,
cold_storage_object_key=None,
@ -5519,6 +5520,7 @@ def get_standard_logging_metadata(
applied_guardrails=None,
mcp_tool_call_metadata=None,
vector_store_request_metadata=None,
routing_decision=None,
usage_object=None,
requester_custom_headers=None,
user_api_key_request_route=None,

View file

@ -1267,7 +1267,7 @@ def _get_dummy_thought_signature() -> str:
def convert_to_gemini_tool_call_invoke(
message: ChatCompletionAssistantMessage,
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
forward_function_call_id: bool = False,
) -> List[VertexPartType]:
"""
OpenAI tool invokes:
@ -1317,16 +1317,12 @@ def convert_to_gemini_tool_call_invoke(
VertexGeminiConfig,
)
forward_tool_call_id = bool(
model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider)
)
if tool_calls is not None:
for idx, tool in enumerate(tool_calls):
if "function" in tool:
gemini_function_call: Optional[VertexFunctionCall] = _gemini_tool_call_invoke_helper(
function_call_params=tool["function"],
tool_call_id=(tool.get("id") if forward_tool_call_id else None),
tool_call_id=(tool.get("id") if forward_function_call_id else None),
)
if gemini_function_call is not None:
part_dict: VertexPartType = {"function_call": gemini_function_call}
@ -1378,8 +1374,7 @@ def convert_to_gemini_tool_call_invoke(
def convert_to_gemini_tool_call_result(
message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage],
last_message_with_tool_calls: Optional[dict],
model: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
forward_function_call_id: bool = False,
) -> Union[VertexPartType, List[VertexPartType]]:
"""
OpenAI message with a tool result looks like:
@ -1501,14 +1496,8 @@ def convert_to_gemini_tool_call_result(
name = tool.get("function", {}).get("name", "")
# Echo the OpenAI tool_call_id on functionResponse (strip thought-signature suffix).
# Only Google AI Studio Gemini 3+ accepts `id` on function_response parts.
# Vertex AI and older Gemini models reject the field with HTTP 400.
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
gemini_call_id: Optional[str] = None
if model and VertexGeminiConfig._forward_gemini_function_call_id(model, custom_llm_provider):
if forward_function_call_id:
raw_tool_call_id = message.get("tool_call_id")
if raw_tool_call_id and isinstance(raw_tool_call_id, str):
stripped_id = raw_tool_call_id.split(THOUGHT_SIGNATURE_SEPARATOR, 1)[0]

View file

@ -393,24 +393,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if compaction_event is not None:
return compaction_event
if self.sent_content_block_start is False:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": {"type": "text", "text": ""},
}
)
return self.chunk_queue.popleft()
for chunk in self.completion_stream:
if chunk == "None" or chunk is None:
raise Exception
should_start_new_block = self._should_start_new_content_block(chunk)
if should_start_new_block:
is_opening_first_block = self.sent_content_block_start is False
if is_opening_first_block and self._is_blank_delta(chunk):
continue
if is_opening_first_block:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": self.current_content_block_start,
}
)
elif should_start_new_block:
self._increment_content_block_index()
# applied_edits only needs to flow to the final message_delta
@ -447,7 +448,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
# ``not self.queued_usage_chunk``.
continue
if should_start_new_block and not self.sent_content_block_finish:
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# -> (optionally) the trigger chunk's delta.
#
@ -615,25 +616,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
if compaction_event is not None:
return compaction_event
if self.sent_content_block_start is False:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": {"type": "text", "text": ""},
}
)
return self.chunk_queue.popleft()
async for chunk in self.completion_stream:
if chunk == "None" or chunk is None:
raise Exception
# Check if we need to start a new content block
should_start_new_block = self._should_start_new_content_block(chunk)
if should_start_new_block:
is_opening_first_block = self.sent_content_block_start is False
if is_opening_first_block and self._is_blank_delta(chunk):
continue
if is_opening_first_block:
self.sent_content_block_start = True
self.sent_content_block_finish = False
self.chunk_queue.append(
{
"type": "content_block_start",
"index": self.current_content_block_index,
"content_block": self.current_content_block_start,
}
)
elif should_start_new_block:
self._increment_content_block_index()
# applied_edits only needs to flow to the final message_delta
@ -664,7 +665,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
# Check if this processed chunk has a stop_reason - hold it for next chunk
if not self.queued_usage_chunk:
if should_start_new_block and not self.sent_content_block_finish:
if should_start_new_block and not is_opening_first_block and not self.sent_content_block_finish:
# Queue the sequence: content_block_stop -> content_block_start
# -> (optionally) the trigger chunk's delta.
#
@ -875,6 +876,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
return False
return bool(delta.get(_delta_payload_field(delta_type)))
@staticmethod
def _is_blank_delta(chunk: "ModelResponseStream") -> bool:
choice = chunk.choices[0]
if choice.finish_reason is not None:
return False
delta = choice.delta
if getattr(delta, "tool_calls", None):
return False
if getattr(delta, "content", None):
return False
if getattr(delta, "reasoning_content", None):
return False
if getattr(delta, "thinking_blocks", None):
return False
return True
def _should_start_new_content_block(self, chunk: "ModelResponseStream") -> bool:
"""
Determine if we should start a new content block based on the processed chunk.

View file

@ -331,6 +331,7 @@ class LiteLLMAnthropicMessagesAdapter:
"thinking",
"output_format",
"output_config",
"stop_sequences",
]
def _is_web_search_tool(self, tool: Dict[str, Any]) -> bool:
@ -615,7 +616,7 @@ class LiteLLMAnthropicMessagesAdapter:
thinking_type = thinking.get("type", "disabled")
if thinking_type == "disabled":
return None
return "none"
elif thinking_type == "enabled":
return reasoning_effort_from_thinking_budget(thinking.get("budget_tokens", 0))
elif thinking_type == "adaptive":
@ -683,25 +684,37 @@ class LiteLLMAnthropicMessagesAdapter:
thinking
)
if reasoning_effort:
summary = thinking.get("summary") if isinstance(thinking, dict) else None
auto_summary = is_reasoning_auto_summary_enabled()
if summary:
return {
"reasoning_effort": {
"effort": reasoning_effort,
"summary": summary,
}
}
elif auto_summary:
return {
"reasoning_effort": {
"effort": reasoning_effort,
"summary": "detailed",
}
}
return {"reasoning_effort": reasoning_effort}
return {
"reasoning_effort": LiteLLMAnthropicMessagesAdapter._apply_reasoning_summary_wrapping(
reasoning_effort, thinking
)
}
return {}
@staticmethod
def _apply_reasoning_summary_wrapping(
reasoning_effort: str,
thinking: Dict[str, Any],
) -> Any:
"""
Apply the reasoning_effort/summary wrapping rules shared by every
thinking->reasoning_effort translation path.
Disabled thinking always stays a plain string - there's no reasoning
trace to summarize, and non-Claude providers (e.g. Fireworks) expect
reasoning_effort as a plain string, not a summary dict.
"""
thinking_type = thinking.get("type") if isinstance(thinking, dict) else None
if thinking_type == "disabled":
return reasoning_effort
summary = thinking.get("summary") if isinstance(thinking, dict) else None
if summary:
return {"effort": reasoning_effort, "summary": summary}
if is_reasoning_auto_summary_enabled():
return {"effort": reasoning_effort, "summary": "detailed"}
return reasoning_effort
def translate_anthropic_tool_choice_to_openai(
self, tool_choice: AnthropicMessagesToolChoice
) -> ChatCompletionToolChoiceValues:
@ -919,6 +932,18 @@ class LiteLLMAnthropicMessagesAdapter:
tool_choice=cast(AnthropicMessagesToolChoice, tool_choice)
)
def _translate_stop_sequences_to_openai(
self,
anthropic_message_request: AnthropicMessagesRequest,
new_kwargs: ChatCompletionRequest,
) -> None:
if "stop_sequences" not in anthropic_message_request:
return
stop_sequences = anthropic_message_request["stop_sequences"]
if not stop_sequences:
return
new_kwargs["stop"] = stop_sequences
def _translate_tools_to_openai(
self,
anthropic_message_request: AnthropicMessagesRequest,
@ -976,32 +1001,17 @@ class LiteLLMAnthropicMessagesAdapter:
if not reasoning_effort:
return
thinking_type = thinking.get("type") if isinstance(thinking, dict) else None
# For adaptive thinking, override with output_config.effort if available
if isinstance(thinking, dict) and thinking.get("type") == "adaptive":
if thinking_type == "adaptive":
output_config = anthropic_message_request.get("output_config")
if isinstance(output_config, dict) and output_config.get("effort"):
reasoning_effort = output_config["effort"]
summary = thinking.get("summary") if isinstance(thinking, dict) else None
auto_summary = is_reasoning_auto_summary_enabled()
if summary:
new_kwargs["reasoning_effort"] = cast(
Any,
{
"effort": reasoning_effort,
"summary": summary,
},
)
elif auto_summary:
new_kwargs["reasoning_effort"] = cast(
Any,
{
"effort": reasoning_effort,
"summary": "detailed",
},
)
else:
new_kwargs["reasoning_effort"] = reasoning_effort
new_kwargs["reasoning_effort"] = self._apply_reasoning_summary_wrapping(
reasoning_effort, cast(Dict[str, Any], thinking)
)
def _translate_output_format_to_openai(
self,
@ -1098,6 +1108,11 @@ class LiteLLMAnthropicMessagesAdapter:
anthropic_message_request=anthropic_message_request,
new_kwargs=new_kwargs,
)
## CONVERT STOP_SEQUENCES
self._translate_stop_sequences_to_openai(
anthropic_message_request=anthropic_message_request,
new_kwargs=new_kwargs,
)
## CONVERT OUTPUT_FORMAT to RESPONSE_FORMAT
self._translate_output_format_to_openai(
anthropic_message_request=anthropic_message_request,

View file

@ -5,7 +5,6 @@ from .invoke_handler import (
AmazonAnthropicClaudeStreamDecoder,
AmazonDeepSeekR1StreamDecoder,
AWSEventStreamDecoder,
BedrockLLM,
)

View file

@ -1,19 +1,10 @@
"""
TODO: DELETE FILE. Bedrock LLM is no longer used. Goto `litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py`
"""
import copy
import time
import types
from functools import partial
from typing import (
AsyncIterator,
Callable,
Iterator,
Optional,
Tuple,
cast,
get_args,
)
import httpx # type: ignore
@ -25,16 +16,6 @@ from litellm.caching.caching import InMemoryCache
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.logging_utils import track_llm_api_timing
from litellm.litellm_core_utils.prompt_templates.factory import (
cohere_message_pt,
construct_tool_use_system_prompt,
contains_tag,
custom_prompt,
extract_between_tags,
parse_xml_params,
prompt_factory,
)
from litellm.llms.anthropic.chat.handler import (
ModelResponseIterator as AnthropicModelResponseIterator,
)
@ -64,12 +45,9 @@ from litellm.types.utils import (
StreamingChoices,
Usage,
)
from litellm.utils import CustomStreamWrapper, get_secret
from ..base_aws_llm import BaseAWSLLM
from ..common_utils import (
BedrockError,
ModelResponseIterator,
build_bedrock_stream_error,
get_bedrock_response_stream_shape,
get_bedrock_tool_name,
@ -77,9 +55,6 @@ from ..common_utils import (
bedrock_tool_name_mappings: InMemoryCache = InMemoryCache(max_size_in_memory=50, default_ttl=600)
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
from litellm.llms.bedrock.chat.invoke_transformations.amazon_openai_transformation import (
AmazonBedrockOpenAIConfig,
)
converse_config = AmazonConverseConfig()
@ -351,932 +326,6 @@ def make_sync_call(
raise BedrockError(status_code=500, message=str(e))
class BedrockLLM(BaseAWSLLM):
"""
Example call
```
curl --location --request POST 'https://bedrock-runtime.{aws_region_name}.amazonaws.com/model/{bedrock_model_name}/invoke' \
--header 'Content-Type: application/json' \
--header 'Accept: application/json' \
--user "$AWS_ACCESS_KEY_ID":"$AWS_SECRET_ACCESS_KEY" \
--aws-sigv4 "aws:amz:us-east-1:bedrock" \
--data-raw '{
"prompt": "Hi",
"temperature": 0,
"p": 0.9,
"max_tokens": 4096
}'
```
"""
def __init__(self) -> None:
super().__init__()
@staticmethod
def is_claude_messages_api_model(model: str) -> bool:
"""
Check if the model uses the Claude Messages API (Claude 3+).
Handles:
- Regional prefixes: eu.anthropic.claude-*, us.anthropic.claude-*
- Claude 3 models: claude-3-haiku, claude-3-sonnet, claude-3-opus, claude-3-5-*, claude-3-7-*
- Claude 4 models: claude-opus-4, claude-sonnet-4, claude-haiku-4
"""
# Normalize model string to lowercase for matching
model_lower = model.lower()
# Claude 3+ indicators (all use Messages API)
messages_api_indicators = [
"claude-3", # Claude 3.x models
"claude-opus-4", # Claude Opus 4
"claude-sonnet-4", # Claude Sonnet 4
"claude-haiku-4", # Claude Haiku 4
]
return any(indicator in model_lower for indicator in messages_api_indicators)
def convert_messages_to_prompt(self, model, messages, provider, custom_prompt_dict) -> Tuple[str, Optional[list]]:
# handle anthropic prompts and amazon titan prompts
prompt = ""
chat_history: Optional[list] = None
## CUSTOM PROMPT
if model in custom_prompt_dict:
# check if the model has a registered custom prompt
model_prompt_details = custom_prompt_dict[model]
prompt = custom_prompt(
role_dict=model_prompt_details["roles"],
initial_prompt_value=model_prompt_details.get("initial_prompt_value", ""),
final_prompt_value=model_prompt_details.get("final_prompt_value", ""),
messages=messages,
)
return prompt, None
## ELSE
if provider == "anthropic" or provider == "amazon":
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
elif provider == "mistral":
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
elif provider == "meta" or provider == "llama":
prompt = prompt_factory(model=model, messages=messages, custom_llm_provider="bedrock")
elif provider == "openai":
# OpenAI uses messages directly, no prompt conversion needed
# Return empty prompt as it won't be used
prompt = ""
elif provider == "cohere":
prompt, chat_history = cohere_message_pt(messages=messages)
else:
prompt = ""
for message in messages:
if "role" in message:
if message["role"] == "user":
prompt += f"{message['content']}"
else:
prompt += f"{message['content']}"
else:
prompt += f"{message['content']}"
return prompt, chat_history # type: ignore
def process_response(
self,
model: str,
response: httpx.Response,
model_response: ModelResponse,
stream: Optional[bool],
logging_obj: Logging,
optional_params: dict,
api_key: str,
data: Union[dict, str],
messages: List,
print_verbose,
encoding,
) -> Union[ModelResponse, CustomStreamWrapper]:
provider = self.get_bedrock_invoke_provider(model)
## LOGGING
logging_obj.post_call(
input=messages,
api_key=api_key,
original_response=response.text,
additional_args={"complete_input_dict": data},
)
print_verbose(f"raw model_response: {response.text}")
## RESPONSE OBJECT
try:
completion_response = response.json()
except Exception:
raise BedrockError(message=response.text, status_code=422)
outputText: Optional[str] = None
try:
if provider == "cohere":
if "text" in completion_response:
outputText = completion_response["text"] # type: ignore
elif "generations" in completion_response:
outputText = completion_response["generations"][0]["text"]
model_response.choices[0].finish_reason = map_finish_reason(
completion_response["generations"][0]["finish_reason"]
)
elif provider == "anthropic":
if self.is_claude_messages_api_model(model):
json_schemas: dict = {}
_is_function_call = False
## Handle Tool Calling
if "tools" in optional_params:
_is_function_call = True
for tool in optional_params["tools"]:
json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
outputText = completion_response.get("content")[0].get("text", None)
if outputText is not None and contains_tag("invoke", outputText): # OUTPUT PARSE FUNCTION CALL
function_name = extract_between_tags("tool_name", outputText)[0]
function_arguments_str = extract_between_tags("invoke", outputText)[0].strip()
function_arguments_str = f"<invoke>{function_arguments_str}</invoke>"
function_arguments = parse_xml_params(
function_arguments_str,
json_schema=json_schemas.get(
function_name, None
), # check if we have a json schema for this function name)
)
_message = litellm.Message(
tool_calls=[
{
"id": f"call_{uuid.uuid4()}",
"type": "function",
"function": {
"name": function_name,
"arguments": json.dumps(function_arguments),
},
}
],
content=None,
)
model_response.choices[0].message = _message # type: ignore
model_response._hidden_params["original_response"] = (
outputText # allow user to access raw anthropic tool calling response
)
if _is_function_call is True and stream is not None and stream is True:
print_verbose("INSIDE BEDROCK STREAMING TOOL CALLING CONDITION BLOCK")
# return an iterator
streaming_model_response = ModelResponseStream()
streaming_model_response.choices[0].finish_reason = getattr(
model_response.choices[0], "finish_reason", "stop"
)
# streaming_model_response.choices = [litellm.utils.StreamingChoices()]
streaming_choice = litellm.utils.StreamingChoices()
streaming_choice.index = model_response.choices[0].index
_tool_calls = []
print_verbose(f"type of model_response.choices[0]: {type(model_response.choices[0])}")
print_verbose(f"type of streaming_choice: {type(streaming_choice)}")
if isinstance(model_response.choices[0], litellm.Choices):
if getattr(
model_response.choices[0].message, "tool_calls", None
) is not None and isinstance(model_response.choices[0].message.tool_calls, list):
for tool_call in model_response.choices[0].message.tool_calls:
_tool_call = {**tool_call.dict(), "index": 0}
_tool_calls.append(_tool_call)
delta_obj = Delta(
content=getattr(model_response.choices[0].message, "content", None),
role=model_response.choices[0].message.role,
tool_calls=_tool_calls,
)
streaming_choice.delta = delta_obj
streaming_model_response.choices = [streaming_choice]
completion_stream = ModelResponseIterator(model_response=streaming_model_response)
print_verbose(
"Returns anthropic CustomStreamWrapper with 'cached_response' streaming object"
)
return litellm.CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="cached_response",
logging_obj=logging_obj,
)
model_response.choices[0].finish_reason = map_finish_reason(
completion_response.get("stop_reason", "")
)
_usage = litellm.Usage(
prompt_tokens=completion_response["usage"]["input_tokens"],
completion_tokens=completion_response["usage"]["output_tokens"],
total_tokens=completion_response["usage"]["input_tokens"]
+ completion_response["usage"]["output_tokens"],
)
setattr(model_response, "usage", _usage)
else:
outputText = completion_response["completion"]
model_response.choices[0].finish_reason = completion_response["stop_reason"]
elif provider == "ai21":
outputText = completion_response.get("completions")[0].get("data").get("text")
elif provider == "meta" or provider == "llama":
outputText = completion_response["generation"]
elif provider == "openai":
# OpenAI imported models use OpenAI Chat Completions format
if "choices" in completion_response and len(completion_response["choices"]) > 0:
choice = completion_response["choices"][0]
if "message" in choice:
outputText = choice["message"].get("content")
elif "text" in choice: # fallback for completion format
outputText = choice["text"]
# Set finish reason
if "finish_reason" in choice:
model_response.choices[0].finish_reason = map_finish_reason(choice["finish_reason"])
# Set usage if available
if "usage" in completion_response:
usage = completion_response["usage"]
_usage = litellm.Usage(
prompt_tokens=usage.get("prompt_tokens", 0),
completion_tokens=usage.get("completion_tokens", 0),
total_tokens=usage.get("total_tokens", 0),
)
setattr(model_response, "usage", _usage)
elif provider == "mistral":
outputText = completion_response["outputs"][0]["text"]
model_response.choices[0].finish_reason = completion_response["outputs"][0]["stop_reason"]
else: # amazon titan
outputText = completion_response.get("results")[0].get("outputText")
except Exception as e:
raise BedrockError(
message="Error processing={}, Received error={}".format(response.text, str(e)),
status_code=422,
)
try:
if (
outputText is not None
and len(outputText) > 0
and hasattr(model_response.choices[0], "message")
and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
is None
):
model_response.choices[0].message.content = outputText # type: ignore
elif (
hasattr(model_response.choices[0], "message")
and getattr(model_response.choices[0].message, "tool_calls", None) # type: ignore
is not None
):
pass
else:
raise Exception()
except Exception as e:
raise BedrockError(
message="Error parsing received text={}.\nError-{}".format(outputText, str(e)),
status_code=response.status_code,
)
if stream and provider == "ai21":
streaming_model_response = ModelResponseStream()
streaming_model_response.choices[0].finish_reason = model_response.choices[ # type: ignore
0
].finish_reason
# streaming_model_response.choices = [litellm.utils.StreamingChoices()]
streaming_choice = litellm.utils.StreamingChoices()
streaming_choice.index = model_response.choices[0].index
delta_obj = litellm.utils.Delta(
content=getattr(model_response.choices[0].message, "content", None), # type: ignore
role=model_response.choices[0].message.role, # type: ignore
)
streaming_choice.delta = delta_obj
streaming_model_response.choices = [streaming_choice]
mri = ModelResponseIterator(model_response=streaming_model_response)
return CustomStreamWrapper(
completion_stream=mri,
model=model,
custom_llm_provider="cached_response",
logging_obj=logging_obj,
)
## CALCULATING USAGE - bedrock returns usage in the headers
# Skip if usage was already set (e.g., from JSON response for OpenAI provider)
if not hasattr(model_response, "usage") or getattr(model_response, "usage", None) is None:
bedrock_input_tokens = response.headers.get("x-amzn-bedrock-input-token-count", None)
bedrock_output_tokens = response.headers.get("x-amzn-bedrock-output-token-count", None)
prompt_tokens = int(bedrock_input_tokens or litellm.token_counter(messages=messages))
completion_tokens = int(
bedrock_output_tokens
or litellm.token_counter(
text=model_response.choices[0].message.content, # type: ignore
count_response_tokens=True,
)
)
model_response.created = int(time.time())
model_response.model = model
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
setattr(model_response, "usage", usage)
else:
# Ensure created and model are set even if usage was already set
model_response.created = int(time.time())
model_response.model = model
return model_response
def completion(
self,
model: str,
messages: list,
api_base: Optional[str],
custom_prompt_dict: dict,
model_response: ModelResponse,
print_verbose: Callable,
encoding,
logging_obj: Logging,
optional_params: dict,
acompletion: bool,
timeout: Optional[Union[float, httpx.Timeout]],
litellm_params=None,
logger_fn=None,
extra_headers: Optional[dict] = None,
client: Optional[Union[AsyncHTTPHandler, HTTPHandler]] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
try:
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
## SETUP ##
stream = optional_params.pop("stream", None)
stream_chunk_size = optional_params.pop("stream_chunk_size", None)
provider = self.get_bedrock_invoke_provider(model)
modelId = self.get_bedrock_model_id(
model=model,
provider=provider,
optional_params=optional_params,
)
## CREDENTIALS ##
# pop aws_secret_access_key, aws_access_key_id, aws_session_token, aws_region_name from kwargs, since completion calls fail with them
aws_secret_access_key = optional_params.pop("aws_secret_access_key", None)
aws_access_key_id = optional_params.pop("aws_access_key_id", None)
aws_session_token = optional_params.pop("aws_session_token", None)
aws_region_name = optional_params.pop("aws_region_name", None)
aws_role_name = optional_params.pop("aws_role_name", None)
aws_session_name = optional_params.pop("aws_session_name", None)
aws_profile_name = optional_params.pop("aws_profile_name", None)
aws_bedrock_runtime_endpoint = optional_params.pop(
"aws_bedrock_runtime_endpoint", None
) # https://bedrock-runtime.{region_name}.amazonaws.com
aws_web_identity_token = optional_params.pop("aws_web_identity_token", None)
aws_sts_endpoint = optional_params.pop("aws_sts_endpoint", None)
ssl_verify = optional_params.pop("ssl_verify", None)
### SET REGION NAME ###
if aws_region_name is None:
# check env #
litellm_aws_region_name = get_secret("AWS_REGION_NAME", None)
if litellm_aws_region_name is not None and isinstance(litellm_aws_region_name, str):
aws_region_name = litellm_aws_region_name
standard_aws_region_name = get_secret("AWS_REGION", None)
if standard_aws_region_name is not None and isinstance(standard_aws_region_name, str):
aws_region_name = standard_aws_region_name
if aws_region_name is None:
aws_region_name = "us-west-2"
credentials: Credentials = self.get_credentials(
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
aws_region_name=aws_region_name,
aws_session_name=aws_session_name,
aws_profile_name=aws_profile_name,
aws_role_name=aws_role_name,
aws_web_identity_token=aws_web_identity_token,
aws_sts_endpoint=aws_sts_endpoint,
ssl_verify=ssl_verify,
)
### SET RUNTIME ENDPOINT ###
endpoint_url, proxy_endpoint_url = self.get_runtime_endpoint(
api_base=api_base,
aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint,
aws_region_name=aws_region_name,
)
if (stream is not None and stream is True) and provider != "ai21":
endpoint_url = f"{endpoint_url}/model/{modelId}/invoke-with-response-stream"
proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke-with-response-stream"
else:
endpoint_url = f"{endpoint_url}/model/{modelId}/invoke"
proxy_endpoint_url = f"{proxy_endpoint_url}/model/{modelId}/invoke"
if acompletion and provider == "anthropic" and self.is_claude_messages_api_model(model):
if isinstance(client, HTTPHandler):
client = None
return self._async_anthropic_messages_completion(
model=model,
messages=messages,
endpoint_url=endpoint_url,
proxy_endpoint_url=proxy_endpoint_url,
credentials=credentials,
aws_region_name=aws_region_name,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
litellm_params=litellm_params,
logger_fn=logger_fn,
extra_headers=extra_headers,
timeout=timeout,
client=client,
stream_chunk_size=stream_chunk_size,
) # type: ignore[return-value]
prompt, chat_history = self.convert_messages_to_prompt(model, messages, provider, custom_prompt_dict)
inference_params = copy.deepcopy(optional_params)
json_schemas: dict = {}
if provider == "cohere":
if model.startswith("cohere.command-r"):
## LOAD CONFIG
config = litellm.AmazonCohereChatConfig().get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
_data = {"message": prompt, **inference_params}
if chat_history is not None:
_data["chat_history"] = chat_history
data = json.dumps(_data)
else:
## LOAD CONFIG
config = litellm.AmazonCohereConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
if stream is True:
inference_params["stream"] = True # cohere requires stream = True in inference params
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "anthropic":
if self.is_claude_messages_api_model(model):
# Separate system prompt from rest of message
system_prompt_idx: list[int] = []
system_messages: list[str] = []
for idx, message in enumerate(messages):
if message["role"] == "system":
system_messages.append(message["content"])
system_prompt_idx.append(idx)
if len(system_prompt_idx) > 0:
inference_params["system"] = "\n".join(system_messages)
messages = [i for j, i in enumerate(messages) if j not in system_prompt_idx]
# Format rest of message according to anthropic guidelines
messages = prompt_factory(model=model, messages=messages, custom_llm_provider="anthropic_xml") # type: ignore
## LOAD CONFIG
config = litellm.AmazonAnthropicClaudeConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
## Handle Tool Calling
if "tools" in inference_params:
_is_function_call = True
for tool in inference_params["tools"]:
json_schemas[tool["function"]["name"]] = tool["function"].get("parameters", None)
tool_calling_system_prompt = construct_tool_use_system_prompt(tools=inference_params["tools"])
inference_params["system"] = (
inference_params.get("system", "\n") + tool_calling_system_prompt
) # add the anthropic tool calling prompt to the system prompt
inference_params.pop("tools")
data = json.dumps({"messages": messages, **inference_params})
else:
## LOAD CONFIG
config = litellm.AmazonAnthropicConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "ai21":
## LOAD CONFIG
config = litellm.AmazonAI21Config.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "mistral":
## LOAD CONFIG
config = litellm.AmazonMistralConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "amazon": # amazon titan
## LOAD CONFIG
config = litellm.AmazonTitanConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps(
{
"inputText": prompt,
"textGenerationConfig": inference_params,
}
)
elif provider == "meta" or provider == "llama":
## LOAD CONFIG
config = litellm.AmazonLlamaConfig.get_config()
for k, v in config.items():
if (
k not in inference_params
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
inference_params[k] = v
data = json.dumps({"prompt": prompt, **inference_params})
elif provider == "openai":
## OpenAI imported models use OpenAI Chat Completions format (messages-based)
# Use AmazonBedrockOpenAIConfig for proper OpenAI transformation
openai_config = AmazonBedrockOpenAIConfig()
supported_params = openai_config.get_supported_openai_params(model=model)
# Filter to only supported OpenAI params
filtered_params = {k: v for k, v in inference_params.items() if k in supported_params}
# OpenAI uses messages format, not prompt
data = json.dumps({"messages": messages, **filtered_params})
else:
## LOGGING
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": inference_params,
},
)
raise BedrockError(
status_code=404,
message="Bedrock Invoke HTTPX: Unknown provider={}, model={}. Try calling via converse route - `bedrock/converse/<model>`.".format(
provider, model
),
)
## COMPLETION CALL
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=data,
headers=headers,
)
## LOGGING
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
if isinstance(client, HTTPHandler):
client = None
if stream is True and provider != "ai21":
return self.async_streaming(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=True,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
stream_chunk_size=stream_chunk_size,
) # type: ignore
### ASYNC COMPLETION
return self.async_completion(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream, # type: ignore
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
) # type: ignore
if client is None or isinstance(client, AsyncHTTPHandler):
_params = {}
if timeout is not None:
if isinstance(timeout, float) or isinstance(timeout, int):
timeout = httpx.Timeout(timeout)
_params["timeout"] = timeout
self.client = _get_httpx_client(_params) # type: ignore
else:
self.client = client
if (stream is not None and stream is True) and provider != "ai21":
response = self.client.post(
url=proxy_endpoint_url,
headers=prepped.headers, # type: ignore
data=data,
stream=stream,
logging_obj=logging_obj,
)
if response.status_code != 200:
raise BedrockError(status_code=response.status_code, message=str(response.read()))
decoder = AWSEventStreamDecoder(model=model)
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
streaming_response = CustomStreamWrapper(
completion_stream=completion_stream,
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
)
## LOGGING
logging_obj.post_call(
input=messages,
api_key="",
original_response=streaming_response,
additional_args={"complete_input_dict": data},
)
return streaming_response
try:
response = self.client.post(
url=proxy_endpoint_url,
headers=dict(prepped.headers),
data=data,
logging_obj=logging_obj,
)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code
raise BedrockError(status_code=error_code, message=err.response.text)
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return self.process_response(
model=model,
response=response,
model_response=model_response,
stream=stream,
logging_obj=logging_obj,
optional_params=optional_params,
api_key="",
data=data,
messages=messages,
print_verbose=print_verbose,
encoding=encoding,
)
async def _async_anthropic_messages_completion(
self,
model: str,
messages: list,
endpoint_url: str,
proxy_endpoint_url: str,
credentials,
aws_region_name: str,
model_response: ModelResponse,
print_verbose: Callable,
encoding,
logging_obj: Logging,
optional_params: dict,
stream,
litellm_params=None,
logger_fn=None,
extra_headers: Optional[dict] = None,
timeout: Optional[Union[float, httpx.Timeout]] = None,
client: Optional[AsyncHTTPHandler] = None,
stream_chunk_size: Optional[int] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
transformed_request = await litellm.AmazonAnthropicClaudeConfig().async_transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params or {},
headers=extra_headers or {},
)
data = json.dumps(transformed_request)
headers = {"Content-Type": "application/json"}
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
endpoint_url=endpoint_url,
data=data,
headers=headers,
)
logging_obj.pre_call(
input=messages,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": proxy_endpoint_url,
"headers": prepped.headers,
},
)
if stream is True:
return await self.async_streaming(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=True,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
stream_chunk_size=stream_chunk_size,
)
return await self.async_completion(
model=model,
messages=messages,
data=data,
api_base=proxy_endpoint_url,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream, # type: ignore
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=prepped.headers,
timeout=timeout,
client=client,
)
async def async_completion(
self,
model: str,
messages: list,
api_base: str,
model_response: ModelResponse,
print_verbose: Callable,
data: str,
timeout: Optional[Union[float, httpx.Timeout]],
encoding,
logging_obj: Logging,
stream,
optional_params: dict,
litellm_params=None,
logger_fn=None,
headers={},
client: Optional[AsyncHTTPHandler] = None,
) -> Union[ModelResponse, CustomStreamWrapper]:
if client is None:
_params = {}
if timeout is not None:
if isinstance(timeout, float) or isinstance(timeout, int):
timeout = httpx.Timeout(timeout)
_params["timeout"] = timeout
client = get_async_httpx_client(params=_params, llm_provider=litellm.LlmProviders.BEDROCK) # type: ignore
else:
client = client # type: ignore
try:
response = await client.post(
api_base,
headers=headers,
data=data,
timeout=timeout,
logging_obj=logging_obj,
)
response.raise_for_status()
except httpx.HTTPStatusError as err:
error_code = err.response.status_code
raise BedrockError(status_code=error_code, message=err.response.text)
except httpx.TimeoutException:
raise BedrockError(status_code=408, message="Timeout error occurred.")
return self.process_response(
model=model,
response=response,
model_response=model_response,
stream=stream if isinstance(stream, bool) else False,
logging_obj=logging_obj,
api_key="",
data=data,
messages=messages,
print_verbose=print_verbose,
optional_params=optional_params,
encoding=encoding,
)
@track_llm_api_timing() # for streaming, we need to instrument the function calling the wrapper
async def async_streaming(
self,
model: str,
messages: list,
api_base: str,
model_response: ModelResponse,
print_verbose: Callable,
data: str,
timeout: Optional[Union[float, httpx.Timeout]],
encoding,
logging_obj: Logging,
stream,
optional_params: dict,
litellm_params=None,
logger_fn=None,
headers={},
client: Optional[AsyncHTTPHandler] = None,
stream_chunk_size: Optional[int] = None,
) -> CustomStreamWrapper:
# The call is not made here; instead, we prepare the necessary objects for the stream.
streaming_response = CustomStreamWrapper(
completion_stream=None,
make_call=partial(
make_call,
client=client,
api_base=api_base,
headers=headers,
data=data, # type: ignore
model=model,
messages=messages,
logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False,
stream_chunk_size=stream_chunk_size,
),
model=model,
custom_llm_provider="bedrock",
logging_obj=logging_obj,
)
return streaming_response
@staticmethod
def _get_provider_from_model_path(
model_path: str,
) -> Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL]:
"""
Helper function to get the provider from a model path with format: provider/model-name
Args:
model_path (str): The model path (e.g., 'llama/arn:aws:bedrock:us-east-1:086734376398:imported-model/r4c4kewx2s0n' or 'anthropic/model-name')
Returns:
Optional[str]: The provider name, or None if no valid provider found
"""
parts = model_path.split("/")
if len(parts) >= 1:
provider = parts[0]
if provider in get_args(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL):
return cast(litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL, provider)
return None
class AWSEventStreamDecoder:
def __init__(self, model: str, json_mode: Optional[bool] = False) -> None:
from botocore.parsers import EventStreamJSONParser

View file

@ -1109,8 +1109,10 @@ def get_bedrock_chat_config(model: str):
Returns:
The appropriate Bedrock config class instance
"""
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
bedrock_route = BedrockModelInfo.get_bedrock_route(model)
bedrock_invoke_provider = litellm.BedrockLLM.get_bedrock_invoke_provider(model=model)
bedrock_invoke_provider = BaseAWSLLM.get_bedrock_invoke_provider(model=model)
base_model = BedrockModelInfo.get_base_model(model)
# Handle explicit routes first

View file

@ -143,13 +143,26 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
client: Union[ClientSession, Callable[[], ClientSession]],
ssl_verify: Optional[Union[bool, ssl.SSLContext]] = None,
owns_session: bool = True,
session_factory: Callable[[], ClientSession] | None = None,
):
self.client = client
self._ssl_verify = ssl_verify # Store for per-request SSL override
super().__init__(client=client, owns_session=owns_session)
# Store the client factory for recreating sessions when needed
if callable(client):
self._client_factory = client
default_factory: Callable[[], ClientSession] = client if callable(client) else ClientSession
self._client_factory: Callable[[], ClientSession] = session_factory or default_factory
def _rebuild_session(self) -> ClientSession:
"""
Build a replacement session from the configured factory.
The replacement is reachable only from this transport, so the transport
owns it from here on even when it was originally handed a session it did
not own (the proxy's shared session).
"""
session = self._client_factory()
self._owns_session = True
return session
def _get_valid_client_session(self) -> ClientSession:
"""
@ -158,24 +171,16 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
This handles the case where the session was created in a different
event loop that may have been closed (common in CI/CD environments).
"""
from aiohttp.client import ClientSession
# If we don't have a client or it's not a ClientSession, create one
if not isinstance(self.client, ClientSession):
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:
self.client = ClientSession()
self.client = self._rebuild_session()
# Don't return yet - check if the newly created session is valid
# Check if the session itself is closed
if self.client.closed:
verbose_logger.debug("Session is closed, creating new session")
# Create a new session
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:
self.client = ClientSession()
self.client = self._rebuild_session()
return self.client
# Check if the existing session is still valid for the current event loop
@ -188,7 +193,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
# Close old session to prevent leaks
old_session = self.client
try:
if not old_session.closed:
if self._owns_session and not old_session.closed:
try:
asyncio.create_task(old_session.close())
except RuntimeError:
@ -198,17 +203,11 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
verbose_logger.debug(f"Error closing old session: {e}")
# Create a new session in the current event loop
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:
self.client = ClientSession()
self.client = self._rebuild_session()
except (RuntimeError, AttributeError):
# If we can't check the loop or session is invalid, recreate it
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:
self.client = ClientSession()
self.client = self._rebuild_session()
return self.client
@ -303,10 +302,7 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
if "Session is closed" in str(e):
verbose_logger.debug(f"Session closed during request, retrying with new session: {e}")
# Force creation of a new session
if hasattr(self, "_client_factory") and callable(self._client_factory):
self.client = self._client_factory()
else:
self.client = ClientSession()
self.client = self._rebuild_session()
client_session = self.client
# Retry the request with the new session

View file

@ -1013,17 +1013,6 @@ class AsyncHTTPHandler:
verbose_logger.debug("Creating AiohttpTransport...")
# Use shared session if provided and valid
if shared_session is not None and not shared_session.closed:
verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})")
return LiteLLMAiohttpTransport(
client=shared_session,
ssl_verify=ssl_for_transport,
owns_session=False,
)
# Create new session only if none provided or existing one is invalid
verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)")
transport_connector_kwargs = {
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
@ -1041,11 +1030,26 @@ class AsyncHTTPHandler:
if socket_factory is not None:
transport_connector_kwargs["socket_factory"] = socket_factory
return LiteLLMAiohttpTransport(
client=lambda: ClientSession(
def session_factory() -> ClientSession:
return ClientSession(
connector=TCPConnector(**transport_connector_kwargs),
trust_env=trust_env,
),
)
# Use shared session if provided and valid
if shared_session is not None and not shared_session.closed:
verbose_logger.debug(f"SHARED SESSION: Reusing existing ClientSession (ID: {id(shared_session)})")
return LiteLLMAiohttpTransport(
client=shared_session,
ssl_verify=ssl_for_transport,
owns_session=False,
session_factory=session_factory,
)
# Create new session only if none provided or existing one is invalid
verbose_logger.debug("NEW SESSION: Creating new ClientSession (no shared session provided)")
return LiteLLMAiohttpTransport(
client=session_factory,
ssl_verify=ssl_for_transport,
)

View file

@ -1626,7 +1626,7 @@ class OpenAIFilesAPI(BaseLLM):
openai_client: AsyncOpenAI,
) -> OpenAIFileObject:
response = await openai_client.files.create(**create_file_data) # type: ignore[arg-type]
return OpenAIFileObject(**response.model_dump())
return OpenAIFileObject.model_validate(response.model_dump())
def create_file(
self,
@ -1662,7 +1662,7 @@ class OpenAIFilesAPI(BaseLLM):
create_file_data=create_file_data, openai_client=openai_client
)
response = cast(OpenAI, openai_client).files.create(**create_file_data) # type: ignore[arg-type]
return OpenAIFileObject(**response.model_dump())
return OpenAIFileObject.model_validate(response.model_dump())
async def afile_content(
self,
@ -1986,7 +1986,7 @@ class OpenAIBatchesAPI(BaseLLM):
openai_client: AsyncOpenAI,
) -> LiteLLMBatch:
response = await openai_client.batches.create(**create_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
def create_batch(
self,
@ -2023,7 +2023,7 @@ class OpenAIBatchesAPI(BaseLLM):
)
response = cast(OpenAI, openai_client).batches.create(**create_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
async def aretrieve_batch(
self,
@ -2032,7 +2032,7 @@ class OpenAIBatchesAPI(BaseLLM):
) -> LiteLLMBatch:
verbose_logger.debug("retrieving batch, args= %s", retrieve_batch_data)
response = await openai_client.batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
def retrieve_batch(
self,
@ -2068,7 +2068,7 @@ class OpenAIBatchesAPI(BaseLLM):
retrieve_batch_data=retrieve_batch_data, openai_client=openai_client
)
response = cast(OpenAI, openai_client).batches.retrieve(**retrieve_batch_data) # type: ignore[arg-type]
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
async def acancel_batch(
self,
@ -2077,7 +2077,7 @@ class OpenAIBatchesAPI(BaseLLM):
) -> LiteLLMBatch:
verbose_logger.debug("async cancelling batch, args= %s", cancel_batch_data)
response = await openai_client.batches.cancel(**cancel_batch_data)
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
def cancel_batch(
self,
@ -2117,7 +2117,7 @@ class OpenAIBatchesAPI(BaseLLM):
if not isinstance(openai_client, OpenAI):
raise ValueError("OpenAI client is not an instance of OpenAI. Make sure you passed a sync OpenAI client.")
response = openai_client.batches.cancel(**cancel_batch_data)
return LiteLLMBatch(**response.model_dump())
return LiteLLMBatch.model_validate(response.model_dump())
async def alist_batches(
self,
@ -2477,9 +2477,9 @@ class OpenAIAssistantsAPI(BaseLLM):
response_obj: Optional[OpenAIMessage] = None
if getattr(thread_message, "status", None) is None:
thread_message.status = "completed"
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
else:
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
return response_obj
# fmt: off
@ -2556,9 +2556,9 @@ class OpenAIAssistantsAPI(BaseLLM):
response_obj: Optional[OpenAIMessage] = None
if getattr(thread_message, "status", None) is None:
thread_message.status = "completed"
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
else:
response_obj = OpenAIMessage(**thread_message.dict())
response_obj = OpenAIMessage.model_validate(thread_message.dict())
return response_obj
async def async_get_messages(

View file

@ -280,7 +280,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
try:
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
except Exception:
verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct")
response = ResponsesAPIResponse.model_construct(**raw_response_json)
@ -506,7 +506,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raise OpenAIError(message=raw_response.text, status_code=raw_response.status_code)
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
response._hidden_params["additional_headers"] = processed_headers
response._hidden_params["headers"] = raw_response_headers
@ -588,7 +588,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
raw_response_headers = dict(raw_response.headers)
processed_headers = process_response_headers(raw_response_headers)
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
response._hidden_params["additional_headers"] = processed_headers
response._hidden_params["headers"] = raw_response_headers
@ -647,7 +647,7 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
processed_headers = process_response_headers(raw_response_headers)
try:
response = ResponsesAPIResponse(**raw_response_json)
response = ResponsesAPIResponse.model_validate(raw_response_json)
except Exception:
verbose_logger.debug(f"Error constructing ResponsesAPIResponse: {raw_response_json}, using model_construct")
response = ResponsesAPIResponse.model_construct(**raw_response_json)

View file

@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
"""
import re
from typing import List, Optional, Tuple, Literal
from typing import List, Optional, Sequence, Tuple, Literal
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai import CachedContentRequestBody
@ -152,6 +152,20 @@ def separate_cached_messages(
return cached_messages, non_cached_messages
def cached_messages_end_on_supported_turn(cached_messages: Sequence[AllMessageValues]) -> bool:
"""
The cachedContents API rejects contents ending on a model turn, which is how it
classifies both assistant messages and tool results, with HTTP 400
"Requests ending with a model turn are not supported". System messages are
extracted into system_instruction before contents are built, so the terminal
turn is the last non-system message.
"""
non_system_messages = tuple(message for message in cached_messages if message.get("role") != "system")
if not non_system_messages:
return bool(cached_messages)
return non_system_messages[-1].get("role") not in ("assistant", "tool", "function")
def transform_openai_messages_to_gemini_context_caching(
model: str,
messages: List[AllMessageValues],

View file

@ -22,6 +22,7 @@ from litellm.types.llms.vertex_ai import (
from ..common_utils import VertexAIError, get_vertex_base_url
from ..vertex_llm_base import VertexBase
from .transformation import (
cached_messages_end_on_supported_turn,
separate_cached_messages,
transform_openai_messages_to_gemini_context_caching,
)
@ -308,6 +309,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
return messages, optional_params, None
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(
@ -459,6 +468,14 @@ class ContextCachingEndpoints(VertexBase):
if len(cached_messages) == 0:
return messages, optional_params, None
if not cached_messages_end_on_supported_turn(cached_messages):
verbose_logger.debug(
"Vertex AI context caching: cached message block ends on a model turn once "
"system messages are extracted, which the cachedContents API rejects. "
"Skipping context caching."
)
return messages, optional_params, None
# Gemini requires a minimum of 1024 tokens for context caching.
# Skip caching if the cached content is too small to avoid API errors.
if not is_prompt_caching_valid_prompt(

View file

@ -1,7 +1,9 @@
import asyncio
import json
import os
import time
from urllib.parse import unquote
from typing import Any, Coroutine, Optional, Tuple, Union
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union
import httpx
@ -10,6 +12,7 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import (
GCSBucketBase,
GCSLoggingConfig,
)
from litellm.types.utils import StandardCallbackDynamicParams
from litellm.litellm_core_utils.cloud_storage_security import (
VERTEX_AI_MANAGED_GCS_PREFIX,
should_allow_legacy_cloud_file_ids,
@ -39,6 +42,35 @@ class VertexAIFilesHandler(GCSBucketBase):
llm_provider=LlmProviders.VERTEX_AI,
)
def _resolve_read_gcs_config(
self,
litellm_params: Mapping[str, object] | None,
vertex_credentials: VERTEX_CREDENTIALS_TYPES | None,
) -> tuple[str | None, str | None]:
"""
Resolve the GCS bucket and service-account credentials for the read/content path.
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
run entirely at the model-group level, so output written to a per-model bucket is
readable without setting the global env vars.
"""
params: Mapping[str, object] = litellm_params or {}
bucket_candidate = params.get("gcs_bucket_name") or params.get("bucket_name")
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
credentials = params.get("vertex_credentials") or vertex_credentials
if isinstance(credentials, dict):
path_service_account: str | None = json.dumps(credentials)
elif isinstance(credentials, str):
path_service_account = credentials
else:
path_service_account = os.getenv("GCS_PATH_SERVICE_ACCOUNT")
return configured_bucket_name, path_service_account
def _extract_bucket_and_object_from_file_id(
self,
file_id: str,
@ -91,7 +123,17 @@ class VertexAIFilesHandler(GCSBucketBase):
if not file_id:
raise ValueError("file_id is required in file_content_request")
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(kwargs={})
configured_bucket_name, path_service_account = self._resolve_read_gcs_config(
litellm_params=litellm_params,
vertex_credentials=vertex_credentials,
)
dynamic_params = StandardCallbackDynamicParams(
gcs_bucket_name=configured_bucket_name,
gcs_path_service_account=path_service_account,
)
gcs_logging_config: GCSLoggingConfig = await self.get_gcs_logging_config(
kwargs={"standard_callback_dynamic_params": dynamic_params}
)
bucket_name, object_path = self._extract_bucket_and_object_from_file_id(
file_id=file_id,
configured_bucket_name=gcs_logging_config["bucket_name"],

View file

@ -661,6 +661,10 @@ def _gemini_convert_messages_with_history(
vertex_project = litellm_params.get("vertex_project") or litellm_params.get("vertex_ai_project")
vertex_credentials = litellm_params.get("vertex_credentials") or litellm_params.get("vertex_ai_credentials")
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
forward_function_call_id = VertexGeminiConfig._forward_gemini_function_call_id(model or "")
try:
while msg_i < len(messages):
user_content: List[PartType] = []
@ -910,7 +914,7 @@ def _gemini_convert_messages_with_history(
gemini_tool_call_parts = convert_to_gemini_tool_call_invoke(
assistant_msg,
model=model,
custom_llm_provider=custom_llm_provider,
forward_function_call_id=forward_function_call_id,
)
## check if gemini_tool_call already exists in assistant_content
for gemini_tool_call_part in gemini_tool_call_parts:
@ -973,8 +977,7 @@ def _gemini_convert_messages_with_history(
_part = convert_to_gemini_tool_call_result(
messages[msg_i], # type: ignore
last_message_with_tool_calls, # type: ignore
model=model,
custom_llm_provider=custom_llm_provider,
forward_function_call_id=forward_function_call_id,
)
msg_i += 1
# Handle both single part and list of parts (for Computer Use with images)

View file

@ -289,15 +289,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return False
@staticmethod
def _forward_gemini_function_call_id(model: str, custom_llm_provider: Optional[str] = None) -> bool:
def _forward_gemini_function_call_id(model: str) -> bool:
"""
Whether to include `id` on function_call / function_response parts.
Gemini 3+ on Google AI Studio accepts (and returns) `id` for strict
tool-call matching. Vertex AI rejects the field with HTTP 400.
Gemini 3+ accepts (and returns) `id` for strict tool-call matching, on Vertex AI and
Google AI Studio alike. Older Gemini models reject the field with HTTP 400.
"""
if custom_llm_provider != "gemini":
return False
return VertexGeminiConfig._is_gemini_3_or_newer(model)
def _supports_penalty_parameters(self, model: str) -> bool:

View file

@ -207,7 +207,7 @@ from .llms.azure.chat.o_series_handler import AzureOpenAIO1ChatCompletion
from .llms.azure.completion.handler import AzureTextCompletion
from .llms.azure_ai.anthropic.handler import AzureAnthropicChatCompletion
from .llms.azure_ai.embed import AzureAIEmbedding
from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM
from .llms.bedrock.chat import BedrockConverseLLM
from .llms.bedrock.embed.embedding import BedrockEmbedding
from .llms.bedrock.image_edit.handler import BedrockImageEdit
from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration

View file

@ -3454,22 +3454,16 @@
},
"azure_ai/gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_above_272k_tokens": 1.5e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"cache_read_input_token_cost_above_272k_tokens_priority": 3e-07,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_above_272k_tokens": 1.5e-06,
"input_cost_per_token_priority": 1.5e-06,
"input_cost_per_token_above_272k_tokens_priority": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_above_272k_tokens": 6.75e-06,
"output_cost_per_token_priority": 9e-06,
"output_cost_per_token_above_272k_tokens_priority": 1.35e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-mini",
"supported_endpoints": [
"/v1/chat/completions",
@ -3500,22 +3494,16 @@
},
"azure_ai/gpt-5.4-mini-2026-03-17": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_above_272k_tokens": 1.5e-07,
"cache_read_input_token_cost_priority": 1.5e-07,
"cache_read_input_token_cost_above_272k_tokens_priority": 3e-07,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_above_272k_tokens": 1.5e-06,
"input_cost_per_token_priority": 1.5e-06,
"input_cost_per_token_above_272k_tokens_priority": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_above_272k_tokens": 6.75e-06,
"output_cost_per_token_priority": 9e-06,
"output_cost_per_token_above_272k_tokens_priority": 1.35e-05,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-mini",
"supported_endpoints": [
"/v1/chat/completions",
@ -3546,22 +3534,16 @@
},
"azure_ai/gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_above_272k_tokens": 4e-08,
"cache_read_input_token_cost_priority": 4e-08,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_272k_tokens": 4e-07,
"input_cost_per_token_priority": 4e-07,
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_above_272k_tokens": 1.875e-06,
"output_cost_per_token_priority": 2.5e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.75e-06,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-nano",
"supported_endpoints": [
"/v1/chat/completions",
@ -3592,22 +3574,16 @@
},
"azure_ai/gpt-5.4-nano-2026-03-17": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_above_272k_tokens": 4e-08,
"cache_read_input_token_cost_priority": 4e-08,
"cache_read_input_token_cost_above_272k_tokens_priority": 8e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_above_272k_tokens": 4e-07,
"input_cost_per_token_priority": 4e-07,
"input_cost_per_token_above_272k_tokens_priority": 8e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 400000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_above_272k_tokens": 1.875e-06,
"output_cost_per_token_priority": 2.5e-06,
"output_cost_per_token_above_272k_tokens_priority": 3.75e-06,
"source": "https://ai.azure.com/catalog/models/gpt-5.4-nano",
"supported_endpoints": [
"/v1/chat/completions",
@ -7201,7 +7177,7 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7236,7 +7212,7 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7271,7 +7247,7 @@
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -7306,7 +7282,7 @@
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
@ -13567,6 +13543,56 @@
}
]
},
"dashscope/qwen3.7-max": {
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 2.5e-06,
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 7.5e-06,
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"dashscope/qwen3.7-plus": {
"litellm_provider": "dashscope",
"max_input_tokens": 991808,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"source": "https://www.alibabacloud.com/help/en/model-studio/models",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"tiered_pricing": [
{
"cache_read_input_token_cost": 8e-08,
"input_cost_per_token": 4e-07,
"output_cost_per_token": 1.6e-06,
"range": [
0,
256000.0
]
},
{
"cache_read_input_token_cost": 2.4e-07,
"input_cost_per_token": 1.2e-06,
"output_cost_per_token": 4.8e-06,
"range": [
256000.0,
1000000.0
]
}
]
},
"dashscope/qwq-plus": {
"input_cost_per_token": 8e-07,
"litellm_provider": "dashscope",

View file

@ -1,6 +1,6 @@
import re
from datetime import datetime, timezone
from typing import Dict, List, Optional, Set, Tuple, cast
from typing import TYPE_CHECKING, Dict, List, Optional, Sequence, Set, Tuple, cast
from fastapi import HTTPException
from starlette.datastructures import Headers
@ -30,6 +30,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credent
)
from litellm.proxy._types import (
UI_TEAM_ID,
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamTable,
ProxyException,
SpecialHeaders,
@ -43,13 +44,27 @@ from litellm.proxy.auth.user_api_key_auth import (
user_api_key_auth,
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.common_utils.user_api_key_cache import get_management_object_ttl
from litellm.proxy.common_utils.user_api_key_cache import (
USER_NO_MCP_PERMISSION_SENTINEL,
get_management_object_ttl,
user_object_permission_id_cache_key,
)
from litellm.repositories.table_repositories import (
AgentsRepository,
MCPServerRepository,
)
from litellm.repositories.user_repository import UserRepository
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
def _as_list(values: Sequence[str] | None) -> list[str] | None: # mutable-ok: resolver returns a list
"""Widen a read-only allowlist back to the mutable list the resolver's own contract returns,
preserving the ``None`` that means "no restriction"."""
return None if values is None else list(values)
def _parse_mcp_server_names_from_path(path: str, mcp_servers_header: Optional[List[str]] = None) -> Optional[List[str]]:
"""Resolve the single MCP server name a cold-start passthrough bypass may
@ -1408,6 +1423,15 @@ class MCPRequestHandler:
f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}"
)
#########################################################
# Apply the internal user's own ceiling (the entitlement attached to the human)
#########################################################
capped, user_restricts = await MCPRequestHandler._apply_user_server_ceiling(
allowed_mcp_servers, user_api_key_auth, keyless_source=keyless_source
)
allowed_mcp_servers = list(capped)
has_lower_level_mcp_restrictions = has_lower_level_mcp_restrictions or user_restricts
#########################################################
# Apply org-level ceiling if org_id is set
#########################################################
@ -1831,6 +1855,12 @@ class MCPRequestHandler:
# No team restrictions → use key restrictions
allowed_tools = cast(List[str], key_tools)
allowed_tools = _as_list(
await MCPRequestHandler._apply_user_tool_ceiling(
allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source
)
)
return await MCPRequestHandler._apply_agent_and_org_tool_ceilings(
allowed_tools, server_id, user_api_key_auth, keyless_source=keyless_source
)
@ -2376,6 +2406,203 @@ class MCPRequestHandler:
verbose_logger.warning(f"Failed to get allowed MCP servers for end_user: {str(e)}")
return []
@staticmethod
async def _get_user_object_permission(
user_api_key_auth: UserAPIKeyAuth | None = None,
) -> LiteLLM_ObjectPermissionTable | None:
"""The internal user's OWN object_permission: the entitlement attached to the HUMAN rather
than to the credential they authenticated with.
A key's object_permission is the credential's scope and a team's is the group's; this one
answers "which MCP servers and tools is this person entitled to", independent of how many keys
they hold. Caches the ``user_id -> object_permission_id`` mapping (with a sentinel for "no
entitlement") exactly as the agent path does, then reuses the shared ``object_permission_id``
cache, so a warm request reads no rows.
``None`` means the human places NO ceiling: no user row, or a row naming no permission. The
two fault classes are deliberately NOT collapsed into that: a user row we cannot read leaves
us unable to say whether they are entitled at all, which is exactly the state before this
level existed, so it places no ceiling; a row that NAMES a permission we cannot read is a
KNOWN entitlement with unknown contents, so it raises and the caller denies.
"""
from litellm.proxy.auth.auth_checks import get_object_permission
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if not user_api_key_auth or not user_api_key_auth.user_id:
return None
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
user_id = user_api_key_auth.user_id
object_permission_id = await MCPRequestHandler._user_object_permission_id(user_id, prisma_client)
if object_permission_id is None:
return None
object_permission = await get_object_permission(
object_permission_id=object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if object_permission is None:
raise ValueError(
f"user {user_id!r} names object_permission_id {object_permission_id!r} which could not be loaded"
)
return object_permission
@staticmethod
async def _user_object_permission_id(user_id: str, prisma_client: "PrismaClient") -> str | None:
"""The permission row this human's user row links to, or None when they link none.
Caches the link (with a sentinel for "links none") so a human without an entitlement costs no
DB read per MCP request. Anything other than an id string is treated as a cache MISS rather
than carried into the permission lookup, and a read that fails answers None: not knowing
whether someone is entitled is the state that existed before this level, so it places no
ceiling. Only a link we DID resolve can make the caller deny.
"""
from litellm.proxy.proxy_server import user_api_key_cache
cache_key = user_object_permission_id_cache_key(user_id)
try:
cached: object = await user_api_key_cache.async_get_cache(key=cache_key)
if cached == USER_NO_MCP_PERMISSION_SENTINEL:
return None
if isinstance(cached, str) and cached:
return cached
user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id})
linked: object = getattr(user_row, "object_permission_id", None) if user_row is not None else None
object_permission_id = linked if isinstance(linked, str) and linked else None
await user_api_key_cache.async_set_cache(
key=cache_key,
value=object_permission_id or USER_NO_MCP_PERMISSION_SENTINEL,
ttl=get_management_object_ttl(user_api_key_cache),
)
return object_permission_id
except Exception as e: # noqa: BLE001 # unknown whether entitled at all: no ceiling, as before
verbose_logger.warning(f"MCP user entitlement: link for {user_id!r} unresolved, no ceiling: {str(e)}")
return None
@staticmethod
async def _get_allowed_mcp_servers_for_user(
user_api_key_auth: UserAPIKeyAuth | None = None,
) -> Sequence[str] | None:
"""The MCP servers the internal user is entitled to, as server ids.
``[]`` means this human places no restriction (allow-all from this level); ``None`` means the
ceiling is UNRESOLVED, which the caller denies on. Servers named only under
``mcp_tool_permissions`` count as entitled, exactly as they do for a key or a team, so
granting one tool never requires naming its server twice.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
try:
object_permissions = await MCPRequestHandler._get_user_object_permission(user_api_key_auth)
if object_permissions is None:
return []
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(object_permissions.mcp_servers or [])
access_group_servers = await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
)
tool_perm_servers = list(
global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).keys()
)
return list(set(direct_mcp_servers + access_group_servers + tool_perm_servers))
except Exception as e: # noqa: BLE001 # any resolution fault is an unresolved ceiling, never "no ceiling"
verbose_logger.warning(f"Failed to get allowed MCP servers for user: {str(e)}")
return None
@staticmethod
async def _apply_user_server_ceiling(
allowed_mcp_servers: Sequence[str],
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
) -> tuple[tuple[str, ...], bool]:
"""Narrow a resolved server list by the internal user's own entitlement.
Returns the capped list and whether this human restricted it at all; the caller needs the
second value because an org list may only CAP a lower-level restriction, never replace one, so
a user ceiling has to be visible to the org step.
RAISES when the entitlement is known but unreadable, which the resolver's own handler turns
into deny-all. That is the point of the level: dropping a ceiling we know exists is exactly the
silent widening it is there to prevent.
"""
if keyless_source:
return tuple(allowed_mcp_servers), False
entitled = await MCPRequestHandler._get_allowed_mcp_servers_for_user(user_api_key_auth)
if entitled is None:
raise ValueError(
f"MCP user ceiling unresolvable for user_id="
f"{user_api_key_auth.user_id if user_api_key_auth else None!r}"
)
if not entitled:
return tuple(allowed_mcp_servers), False
capped = tuple(server for server in allowed_mcp_servers if server in set(entitled))
verbose_logger.debug(f"Applied user ceiling filter. Final allowed servers: {capped}")
return capped, True
@staticmethod
async def _user_places_mcp_ceiling(user_api_key_auth: UserAPIKeyAuth | None = None) -> bool:
"""Whether this human's own entitlement bounds their MCP access at all.
True when they are entitled to a specific set of servers, and also when that entitlement is
UNRESOLVED — a caller uses this to decide whether it may skip the resolver, and skipping it on
a transient fault would widen access.
"""
entitled_servers = await MCPRequestHandler._get_allowed_mcp_servers_for_user(user_api_key_auth)
return entitled_servers is None or len(entitled_servers) > 0
@staticmethod
async def _apply_user_tool_ceiling(
allowed_tools: Sequence[str] | None,
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
) -> Sequence[str] | None:
"""Narrow a key/team tool allowlist by the internal user's own tool entitlement.
The human's entitlement can only ever narrow: a user naming tools on ``server_id`` intersects
(and becomes the allowlist when no lower level restricts), while a user naming none places no
restriction. Returns ``[]`` (deny every tool on this server) when the entitlement cannot be
resolved, because the caller's own except-handler treats a raise as allow-all for key auth.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
if keyless_source:
return allowed_tools
try:
object_permissions = await MCPRequestHandler._get_user_object_permission(user_api_key_auth)
except Exception as e: # noqa: BLE001 # an unresolved human entitlement must deny, not widen
verbose_logger.warning(f"MCP user tool ceiling unresolvable, denying tools on {server_id!r}: {str(e)}")
return []
if object_permissions is None or not object_permissions.mcp_tool_permissions:
return allowed_tools
user_tools = global_mcp_server_manager.expand_tool_permissions(object_permissions.mcp_tool_permissions).get(
server_id
)
if user_tools is None:
return allowed_tools
if allowed_tools is None:
return list(user_tools)
return list(set(allowed_tools) & set(user_tools))
# Sentinel stored in cache when an agent has no object_permission, so we
# don't re-query the DB on every MCP request for that agent.
_AGENT_NO_PERMISSION_SENTINEL = "__agent_no_mcp_permission__"

View file

@ -1778,10 +1778,12 @@ async def token_endpoint(
@router.post("/authorize/complete")
async def authorize_complete(request: Request, flow: str = Form(...)):
async def authorize_complete(request: Request, flow: str = Form(...), delivery: str | None = Form(None)):
"""Finish an aggregate connect flow: mint the gateway authorization code for the
signed-in user and redirect back to the DCR client. POST plus the per-flow HttpOnly
cookie set at /authorize; an anonymous or bad-flow request just 400s."""
signed-in user and hand it back to the DCR client, by 303 redirect (default) or, for
a loopback client on a different machine, as a copyable callback URL
(``delivery=manual``). POST plus the per-flow HttpOnly cookie set at /authorize; an
anonymous or bad-flow request just 400s."""
from litellm.proxy.proxy_server import user_api_key_cache # noqa: PLC0415 # circular import at module load
return await complete_connect_flow(
@ -1789,6 +1791,7 @@ async def authorize_complete(request: Request, flow: str = Form(...)):
flow_handle=flow,
session_user_id=_session_cookie_user_id(request),
cache=user_api_key_cache,
delivery=delivery,
)

View file

@ -39,6 +39,7 @@ from __future__ import annotations
import hashlib
import hmac
import html
import secrets
from base64 import urlsafe_b64encode
from collections.abc import Mapping
@ -47,7 +48,7 @@ from typing import Awaitable, Callable, Literal, TypeVar
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
from fastapi import HTTPException, Request
from fastapi.responses import JSONResponse, RedirectResponse, Response
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response
from pydantic import BaseModel, ConfigDict, Field, ValidationError
from typing_extensions import assert_never
@ -94,6 +95,13 @@ server-side session store, and the sealed value never appears in a URL)."""
CONNECT_FLOW_TTL_SECONDS = 600
GATEWAY_AUTH_CODE_TTL_SECONDS = 120
MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS = 300
"""Lifetime of a code the user delivers by hand (headless/remote client, LIT-4863 class):
copy-pasting a callback URL from a laptop browser to an SSH session is slower than a
browser redirect, so manual-delivery codes get 5 minutes instead of 2, still well under
the 10-minute ceiling RFC 6749 section 4.1.2 recommends. Single-use and PKCE binding are
unchanged, so the longer window only extends how long the legitimate holder has to paste
it, not what an observer could do with it."""
_CLAIM_TTL_BUFFER_SECONDS = 60
_USED_CODE_CACHE_PREFIX = "mcp_gateway_dcr_code_used:"
_USED_FLOW_CACHE_PREFIX = "mcp_gateway_dcr_flow_used:"
@ -390,6 +398,7 @@ async def complete_connect_flow(
flow_handle: str,
session_user_id: str | None,
cache: DualCache,
delivery: str | None = None,
) -> Response:
"""The deliberate finish step of the connect flow: mint the gateway authorization
code and send the browser back to the client.
@ -399,7 +408,24 @@ async def complete_connect_flow(
into the flow: a link crafted by another party dies here with ``access_denied``
instead of minting a code for the victim's identity. The flow is single-use (an atomic
claim on its ``jti``), so a double-submit cannot mint two codes from one sign-in.
``delivery`` chooses how the code reaches the client. Default (absent or
``"redirect"``) is the 303 to the client's registered redirect URI. ``"manual"``
renders the callback URL on a page instead, for a client whose redirect URI is a
loopback host but which runs on a DIFFERENT machine than the browser (EC2/SSH box,
container): the 303 would dereference the browser machine's loopback and the code
would never arrive, so the user carries it over by pasting the URL into the client or
fetching it from the client machine's terminal. Manual delivery is honored only for
loopback redirect URIs; a routable redirect URI works from any browser by
construction, so those flows always redirect. The user who sees the page is exactly
the user the 303 would have carried the code to, and the same user already sees the
code today in the dead redirect's address bar, so the page exposes the code to no new
party. Unknown ``delivery`` values are rejected rather than defaulted: a client that
asked for manual delivery and got a dead redirect instead would silently lose its
code.
"""
if delivery not in (None, "redirect", "manual"):
return _oauth_error(400, "invalid_request", "delivery must be 'redirect' or 'manual'")
sealed_flow = request.cookies.get(_flow_cookie_name(flow_handle))
if sealed_flow is None:
return _oauth_error(400, "invalid_request", "unknown or expired connect flow")
@ -417,6 +443,8 @@ async def complete_connect_flow(
f"{_USED_FLOW_CACHE_PREFIX}{flow.jti}", CONNECT_FLOW_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
):
return _oauth_error(400, "invalid_request", "this connect flow was already completed; restart the connection")
manual_delivery = delivery == "manual" and is_loopback_redirect_host(urlparse(flow.redirect_uri))
code_ttl = MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS if manual_delivery else GATEWAY_AUTH_CODE_TTL_SECONDS
code = _seal(
GATEWAY_AUTH_CODE_PREFIX,
_GatewayAuthCode(
@ -426,16 +454,46 @@ async def complete_connect_flow(
code_challenge=flow.code_challenge,
jti=secrets.token_urlsafe(24),
iat=int(now.timestamp()),
exp=int(now.timestamp()) + GATEWAY_AUTH_CODE_TTL_SECONDS,
exp=int(now.timestamp()) + code_ttl,
),
)
params = {"code": code, **({"state": flow.state} if flow.state else {})}
response = RedirectResponse(_append_query_params(flow.redirect_uri, params), status_code=303)
callback_url = _append_query_params(flow.redirect_uri, params)
response: Response = (
_manual_delivery_response(callback_url) if manual_delivery else RedirectResponse(callback_url, status_code=303)
)
path, secure = _cookie_path_and_secure(request)
response.delete_cookie(key=_flow_cookie_name(flow_handle), path=path, secure=secure, httponly=True, samesite="lax")
return response
def _manual_delivery_response(callback_url: str) -> Response:
"""The manual code-delivery page: the callback URL the 303 would have followed,
rendered for the user to carry to the machine the client actually runs on (paste into
the client's prompt, or fetch with curl from that machine's terminal). Served
no-store because the body holds a live single-use code, and the URL is HTML-escaped
because it is client-influenced. The page renders the URL as data only, never as a
ready-to-paste shell command: no single quoting of an attacker-influenced string is
correct across POSIX shells, cmd.exe, and PowerShell (cmd.exe ignores single quotes
and percent-expands inside double quotes), so any command string this page suggested
would be wrong for some shell the user might paste it into."""
safe_url = html.escape(callback_url, quote=True)
minutes = MANUAL_DELIVERY_AUTH_CODE_TTL_SECONDS // 60
body = (
"<html><head><title>Finish connecting</title></head><body>"
"<h2>Almost done</h2>"
"<p>Your MCP client runs on a different machine, so this browser cannot deliver the"
" authorization code to it. On the machine where the client runs, paste this URL into"
" the client's prompt (Claude Code accepts the pasted callback URL), or pass it as the"
" quoted argument of a curl command from that machine's terminal:</p>"
f'<p><input type="text" value="{safe_url}" readonly size="100" onclick="this.select()"></p>'
f"<p>The code is single-use and expires in {minutes} minutes. You can close this window"
" once the client confirms it is connected.</p>"
"</body></html>"
)
return HTMLResponse(body, headers=TOKEN_NO_CACHE_HEADERS)
def _pkce_verifier_matches(code_verifier: str, code_challenge: str) -> bool:
"""RFC 7636 S256 verification, total over hostile input. The comparison is over bytes
so a non-ASCII ``code_challenge`` (which reaches here unvalidated from the client's
@ -601,9 +659,11 @@ async def _authorization_code_grant(
if failure is not None:
return _reload_failure_response(failure)
# Atomic single-use claim is the gate: on a concurrent double-redeem exactly one caller
# wins, and a claim that cannot be recorded fails closed.
# wins, and a claim that cannot be recorded fails closed. The marker's TTL derives from
# the code's own remaining lifetime so it outlives whichever lifetime the code was minted with.
if not await guard.claim(
f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}", GATEWAY_AUTH_CODE_TTL_SECONDS + _CLAIM_TTL_BUFFER_SECONDS
f"{_USED_CODE_CACHE_PREFIX}{parsed.jti}",
parsed.exp - int(now.timestamp()) + _CLAIM_TTL_BUFFER_SECONDS,
):
return _oauth_error(400, "invalid_grant", "the authorization code was already used")
return _session_token_pair(SessionPrincipal(user_id=parsed.user_id, client_id=client_id), keys, now)

View file

@ -13,11 +13,22 @@ payload (name + arguments) so we just build the tool_call.
from typing import TYPE_CHECKING, Any, Dict, Optional
from fastapi import HTTPException
from mcp.types import Tool as MCPTool
from litellm._logging import verbose_proxy_logger
from litellm.experimental_mcp_client.tools import transform_mcp_tool_to_openai_tool
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.proxy._experimental.mcp_server.utils import (
json_string_leaves,
json_unrewritable_labels,
mcp_content_item_text,
mcp_tool_result_content_list,
mcp_tool_result_structured_content,
set_mcp_tool_result_structured_content,
with_json_string_leaves,
with_mcp_content_item_text,
)
from litellm.types.llms.openai import (
ChatCompletionToolParam,
ChatCompletionToolParamFunctionChunk,
@ -92,7 +103,93 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
user_api_key_dict: Optional[Any] = None,
request_data: Optional[dict] = None,
) -> Any:
verbose_proxy_logger.debug(
"MCP Guardrail: Output processing not implemented for MCP tools",
"""Scan the text content of an MCP tool result and write masked text back.
The content list is rewritten in place (only the entries the guardrail
actually changed) rather than returned as a new result: the same object is
already referenced by the logging payload captured before this hook runs,
so a copy would leave the unmasked text in the spend log / span. A
guardrail that rejects the result raises, and the exception propagates to
the caller.
``structuredContent`` is scanned and masked too, in the same
``apply_guardrail`` call: it is serialized to the client alongside
``content``, so a value living only there would otherwise reach the
client unscanned.
"""
content = mcp_tool_result_content_list(response)
text_blocks = (
tuple(
(index, text) for index, item in enumerate(content) if (text := mcp_content_item_text(item)) is not None
)
if content is not None
else ()
)
structured = mcp_tool_result_structured_content(response)
structured_leaves = json_string_leaves(structured) if structured is not None else ()
structured_labels = json_unrewritable_labels(structured) if structured is not None else ()
if structured_leaves is None or structured_labels is None:
raise HTTPException(
status_code=400,
detail={
"error": (
"Content blocked: MCP tool result structuredContent is nested too deeply to be scanned "
"by the configured guardrail"
)
},
)
if not text_blocks and not structured_leaves and not structured_labels:
verbose_proxy_logger.debug("MCP Guardrail: tool result has no scannable text, nothing to do")
return response
originals = (
tuple(text for _, text in text_blocks) + tuple(text for _, text in structured_leaves) + structured_labels
)
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=GenericGuardrailAPIInputs(texts=list(originals)),
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
masked_texts = guardrailed_inputs.get("texts") if guardrailed_inputs else None
if masked_texts is None:
return response
if len(masked_texts) != len(originals):
verbose_proxy_logger.warning(
"MCP Guardrail: guardrail returned %d texts for %d tool result texts; leaving the result unmasked",
len(masked_texts),
len(originals),
)
return response
split = len(text_blocks)
if content is not None:
for (index, original), masked in zip(text_blocks, masked_texts[:split]):
if masked != original:
content[index] = with_mcp_content_item_text(content[index], masked)
label_start = split + len(structured_leaves)
if any(masked != original for original, masked in zip(structured_labels, masked_texts[label_start:])):
raise HTTPException(
status_code=400,
detail={
"error": (
"Content blocked: MCP tool result matched a masking rule on a non-rewritable field "
"(a structuredContent key or numeric value), which cannot be redacted without changing "
"the payload contract"
)
},
)
structured_replacements = {
path: masked
for (path, original), masked in zip(structured_leaves, masked_texts[split:label_start])
if masked != original
}
if structured_replacements:
set_mcp_tool_result_structured_content(
response, with_json_string_leaves(structured, structured_replacements)
)
return response

View file

@ -200,6 +200,13 @@ _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES: tuple[MCPAuth, ...] = (
)
# OAuth discovery retry cooldown for servers whose endpoints stay unresolved. The base is one
# reload cadence so a transient upstream failure recovers immediately; the cap bounds the request
# amplification and log volume of a permanently broken configuration.
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS = 30.0
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS = 900.0
def _blank_to_none(value: str | None) -> str | None:
"""Collapse an absent, empty, or whitespace-only string to ``None``.
@ -247,6 +254,7 @@ def _endpoints_yield_to_issuer(
authorization_url: str | None,
token_url: str | None,
registration_url: str | None,
server_ref: str,
) -> tuple[str | None, str | None, str | None]:
"""The single rule that makes an admin-configured ``issuer`` the sole authoritative endpoint
source (RFC 8414 §3.3): when it is set for a discovery auth type, the stored/manual
@ -256,9 +264,29 @@ def _endpoints_yield_to_issuer(
i.e. all ``None`` when issuer-anchored, else the inputs unchanged. Called at every resolution site
so the invariant holds in one place instead of being re-derived per merge.
"""
if issuer is not None and is_discovery_auth_type:
return None, None, None
return authorization_url, token_url, registration_url
if issuer is None or not is_discovery_auth_type:
return authorization_url, token_url, registration_url
discarded = sorted(
label
for label, value in (
("authorization_url", authorization_url),
("token_url", token_url),
("registration_url", registration_url),
)
if value
)
if discarded:
verbose_logger.warning(
"MCP server %s has a pinned Issuer, so its stored %s %s not used: an anchored issuer is the "
"sole endpoint source (RFC 8414 section 3.3) and a failed issuer fetch fails closed rather "
"than falling back to them. To use manually configured endpoints instead, clear the Issuer "
"field and re-enter the endpoint urls (clearing the Issuer also clears endpoints that may "
"have been resolved under it), or clear the Issuer alone to re-discover from the server url.",
server_ref,
", ".join(discarded),
"is" if len(discarded) == 1 else "are",
)
return None, None, None
def _normalized_authorize_endpoint(url: str) -> str:
@ -280,6 +308,68 @@ def _issuer_matches(claimed_issuer: object, configured_issuer: str) -> bool:
return _normalized_authorize_endpoint(claimed_issuer) == _normalized_authorize_endpoint(configured_issuer)
def _flow_endpoints_missing(
auth_type: MCPAuthType | None,
oauth2_flow: str | None,
authorization_url: str | None,
token_url: str | None,
token_exchange_endpoint: str | None = None,
) -> bool:
"""Whether a built server is missing an endpoint its flow needs to run at all.
Used by the reload fast-path exemption: discovery runs at build time only, and the fast path
reuses an unchanged row's registry entry verbatim, so a server whose discovery came back empty
(transient upstream failure, rate limiting) would stay broken until some unrelated config write
bumps ``updated_at``, serving its 400 the whole time. Rebuilding just these entries retries
discovery on the normal reload cadence. It costs no extra fetch for servers that resolved, and
none for those with no discovery source, since the build skips discovery for both.
"""
if auth_type == MCPAuth.oauth2_token_exchange:
# A configured exchange endpoint replaces discovery entirely; only a server that must
# discover its token endpoint and still has none is unresolved.
return token_exchange_endpoint is None and token_url is None
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return False
if oauth2_flow == "client_credentials":
return token_url is None
return authorization_url is None or token_url is None
def _oauth_endpoints_unresolved(server: MCPServer) -> bool:
"""``_flow_endpoints_missing`` over a built registry entry, for the reload fast-path check.
The flow comes from ``effective_oauth2_flow``, the one column-first, shape-fallback judge every
flow decision uses, not from the raw column: a legacy row the startup backfill deliberately left
unstamped (the ambiguous M2M shape) serves M2M at request time, and reading the bare column here
would classify it as interactive-missing-endpoints and re-run discovery on every reload.
"""
if (
server.auth_type == MCPAuth.oauth2_token_exchange
and server.token_exchange_profile == "entra_obo"
and not server.scopes
):
# entra_obo fails closed at exchange time without a scope (token_exchanger.py), and scopes
# can come from resource discovery, so a server that resolved its endpoints but no scopes is
# still unresolved for its flow.
return True
if server.is_dcr_bridge and not server.client_id and server.registration_url is None:
# A DCR bridge with no admin-configured client can only register callers through the
# upstream's registration endpoint, so a build that resolved the authorize and token
# endpoints but not registration_endpoint (partial metadata) is still unresolved for its
# flow and must keep retrying; without this it silently degrades to the short-circuit arm
# until an unrelated config write. Scopes are deliberately NOT part of completeness: they
# are a request hint the authorization server bounds at consent (RFC 6749 section 3.3),
# and a server without them is fully functional.
return True
return _flow_endpoints_missing(
server.auth_type,
MCPServerManager.effective_oauth2_flow(server),
server.authorization_url,
server.token_url,
server.token_exchange_endpoint,
)
def _endpoints_corroborate_authorization_url(
source_authorization_url: str | None,
trusted_authorization_url: str | None,
@ -311,11 +401,10 @@ def _carry_forward_resolved_oauth_endpoints(new_server: MCPServer, previous_serv
during re-discovery downgrades a working server (``authorization_url`` set) to a broken one
(``None``, /authorize 400s) with no configuration change. Mirrors the ``short_prefix``
carry-forward. Skipped when the server's ``url`` or ``auth_type`` changed, since the previous
endpoints may then belong to a different upstream. ``registration_url`` IS carried even though
``_persist_discovered_oauth_endpoints`` refuses to write it to the row: carrying only restores
the same in-memory value the previous build already ran with, while persisting it would flip
``_dcr_bridge_relays_client_registration`` (which keys off the stored column) for dcr_bridge
servers that never had one configured.
endpoints may then belong to a different upstream. Discovery results live only on the in-memory
registry entry; the gateway never writes them to the row, whose OAuth columns carry admin intent
alone, so this carry is the sole last-known-good mechanism and restores exactly the values the
previous build already ran with.
Carry-forward is a non-manual endpoint source, so the same trust rule as discovery applies: the
previous ``token_url``/``registration_url``/``scopes`` are carried only when the previous
@ -1182,6 +1271,40 @@ class MCPServerManager:
# empty result, or failure). Used to throttle re-probes for servers that do
# not return instructions, and to apply a short cooldown after failures.
self._upstream_initialize_instructions_probed_at: dict[str, float] = {}
# Per-server (consecutive failures, monotonic timestamp) for OAuth discovery retries, so a
# server whose endpoints never resolve backs off instead of re-running the full
# RFC 9728 -> 8414 chain, and re-logging its warning, on every reload forever.
self._oauth_discovery_retry_state: dict[
str, tuple[int, float]
] = {} # mutable-ok: retry cooldown cache, keyed per server and pruned on success
def _oauth_discovery_retry_due(self, server_id: str) -> bool:
"""Whether an unresolved server is due for another discovery attempt.
The reload fast-path exemption is what retries a failed discovery, so without a cooldown a
permanently unresolvable server re-runs the whole RFC 9728 -> RFC 8414 -> origin-fallback
chain and re-emits its unresolved-endpoints warning on every reload, per server, forever.
Delay doubles per consecutive failure from ``_OAUTH_DISCOVERY_RETRY_BASE_SECONDS`` up to
``_OAUTH_DISCOVERY_RETRY_MAX_SECONDS``, so a transient outage still recovers on the next
reload while a broken configuration settles to one attempt per cap.
"""
state = self._oauth_discovery_retry_state.get(server_id)
if state is None:
return True
failures, attempted_at = state
delay = min(
_OAUTH_DISCOVERY_RETRY_BASE_SECONDS * (2 ** max(failures - 1, 0)),
_OAUTH_DISCOVERY_RETRY_MAX_SECONDS,
)
return (time.monotonic() - attempted_at) >= delay
def _record_oauth_discovery_outcome(self, server: MCPServer) -> None:
"""Advance or clear a server's retry cooldown after a rebuild resolved it or did not."""
if not _oauth_endpoints_unresolved(server):
self._oauth_discovery_retry_state.pop(server.server_id, None)
return
failures, _ = self._oauth_discovery_retry_state.get(server.server_id, (0, 0.0))
self._oauth_discovery_retry_state[server.server_id] = (failures + 1, time.monotonic())
def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None:
raw = getattr(client, "_last_initialize_instructions", None)
@ -1357,6 +1480,7 @@ class MCPServerManager:
manual_authorization_url,
manual_token_url,
manual_registration_url,
server_name or server_id,
)
should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and (
is_discovery_auth_type or obo_needs_discovery
@ -1834,7 +1958,6 @@ class MCPServerManager:
*,
credentials_are_encrypted: bool = True,
env_vars_are_encrypted: Optional[bool] = None,
persist_discovered_endpoints: bool = True,
) -> MCPServer:
_mcp_info: MCPInfo = mcp_server.mcp_info or {}
env_dict = _deserialize_json_dict(getattr(mcp_server, "env", None))
@ -1925,7 +2048,12 @@ class MCPServerManager:
or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url),
)
manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer(
manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url
manual_issuer,
is_discovery_auth_type,
manual_authorization_url,
manual_token_url,
manual_registration_url,
mcp_server.alias or mcp_server.server_name or mcp_server.server_id,
)
gated_oauth_metadata = await self._resolve_table_oauth_metadata(
mcp_server=mcp_server,
@ -2033,143 +2161,8 @@ class MCPServerManager:
max_concurrent_requests=getattr(mcp_server, "max_concurrent_requests", None),
)
_warn_internal_delegate_pkce_if_applicable(new_server, source="database")
if persist_discovered_endpoints:
await self._persist_discovered_obo_token_url(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_token_url=manual_token_url,
discovered_token_url=new_server.token_url,
)
await self._persist_discovered_oauth_endpoints(
server_id=mcp_server.server_id,
auth_type=auth_type,
existing_issuer=manual_issuer,
existing_authorization_url=manual_authorization_url,
existing_token_url=manual_token_url,
existing_scopes=scopes,
metadata=gated_oauth_metadata,
is_issuer_anchored=use_issuer_anchor,
)
return new_server
async def _persist_discovered_obo_token_url(
self,
*,
server_id: str,
auth_type: Optional[MCPAuthType],
existing_token_url: Optional[str],
discovered_token_url: Optional[str],
) -> None:
"""Write a freshly discovered OBO token endpoint back onto the DB row.
``build_mcp_server_from_table`` resolves ``token_url`` via RFC 9728 -> RFC 8414 for an
``oauth2_token_exchange`` server that has none configured, but that resolved value otherwise
lives only on the returned in-memory object; the row keeps ``token_url=None`` so every rebuild
re-runs discovery, and a transient upstream outage during a rebuild leaves the server with no
endpoint until discovery next succeeds. Persisting it makes ``_obo_needs_endpoint_discovery``
return False on the next build. Fires at most once per server (skipped once the row has a
value), and is best-effort: a write failure just means discovery runs again next time.
"""
if auth_type != MCPAuth.oauth2_token_exchange:
return
if existing_token_url or not discovered_token_url:
return
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415
if prisma_client is None:
return
try:
await MCPServerRepository(prisma_client).table.update(
where={"server_id": server_id},
data={"token_url": discovered_token_url},
)
verbose_logger.debug("Persisted discovered OBO token_url for MCP server %s", server_id)
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
verbose_logger.warning("Failed to persist discovered OBO token_url for MCP server %s: %s", server_id, exc)
async def _persist_discovered_oauth_endpoints(
self,
*,
server_id: str,
auth_type: MCPAuthType | None,
existing_issuer: str | None,
existing_authorization_url: str | None,
existing_token_url: str | None,
existing_scopes: list[str] | None,
metadata: MCPOAuthMetadata | None,
is_issuer_anchored: bool = False,
) -> None:
"""Write freshly discovered OAuth endpoints back onto the DB row.
Same rationale as ``_persist_discovered_obo_token_url`` but for the interactive oauth2
family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on
the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path
calls ``update_server``) and on every post-write DB reload, so one failed re-discovery
serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds.
Only fills row fields that are currently empty, never persists origin-fallback guesses
(RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url``
because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a
failed write re-discovers on the next build. Scopes go through ``update_mcp_server`` so
they merge into the credentials blob without touching the stored client credentials.
For an issuer-anchored server (``is_issuer_anchored``) the endpoints are re-derived from the
§3.3-validated issuer document on every build, so they are NOT persisted into the endpoint
columns: persisting them would make the next build see populated endpoints and treat them as
authoritative stored values, defeating the "endpoints come solely from the issuer" invariant.
Only the resource-driven scopes are persisted for such servers.
"""
if auth_type not in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES:
return
if metadata is None or metadata.from_origin_fallback:
return
issuer_update = (
{"issuer": metadata.discovered_issuer} if metadata.discovered_issuer and not existing_issuer else {}
)
authorization_url_update = (
{"authorization_url": metadata.authorization_url}
if metadata.authorization_url and not existing_authorization_url and not is_issuer_anchored
else {}
)
token_url_update = (
{"token_url": metadata.token_url}
if metadata.token_url and not existing_token_url and not is_issuer_anchored
else {}
)
scopes_update = {"credentials": {"scopes": metadata.scopes}} if metadata.scopes and not existing_scopes else {}
updates: dict[str, object] = {
**issuer_update,
**authorization_url_update,
**token_url_update,
**scopes_update,
}
if not updates:
return
from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # db.py imports this module at load
update_mcp_server,
)
from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415 # heavy module; import at call time
from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime value, set after startup
if prisma_client is None:
return
try:
await update_mcp_server(
prisma_client=prisma_client,
data=UpdateMCPServerRequest.model_validate({"server_id": server_id, **updates}),
touched_by="mcp_oauth_discovery",
)
verbose_logger.info(
"Persisted discovered OAuth endpoints for MCP server %s: %s",
server_id,
sorted(updates),
)
except Exception as exc: # noqa: BLE001 - best-effort; a failed write re-discovers next build
verbose_logger.warning(
"Failed to persist discovered OAuth endpoints for MCP server %s: %s",
server_id,
exc,
)
async def _maybe_register_openapi_tools(self, server: MCPServer, *, initialize_mapping: bool = True):
"""Register OpenAPI tools if the server has a spec_path configured."""
if server.spec_path:
@ -2418,6 +2411,11 @@ class MCPServerManager:
and not is_admitted_subject
and _user_has_admin_view(user_api_key_auth)
and not has_explicit_object_permission
# An entitlement attached to the HUMAN binds them whatever their role: it is the
# person's scope, not the credential's, so an admin role is not a waiver of it. An
# UNRESOLVED entitlement also skips the shortcut, so the resolver denies rather than
# handing over the whole registry on a transient fault.
and not await MCPRequestHandler._user_places_mcp_ceiling(user_api_key_auth)
):
verbose_logger.debug("Admin user without explicit object_permission - returning all servers")
return list(self.get_registry().keys())
@ -5329,7 +5327,7 @@ class MCPServerManager:
]
}
)
db_mcp_servers = [LiteLLM_MCPServerTable(**r.model_dump()) for r in raw_rows]
db_mcp_servers = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in raw_rows]
verbose_logger.info(f"Found {len(db_mcp_servers)} MCP servers in database")
previous_registry = self.registry
@ -5347,6 +5345,10 @@ class MCPServerManager:
and existing_server.updated_at is not None
and server.updated_at is not None
and existing_server.updated_at == server.updated_at
and not (
_oauth_endpoints_unresolved(existing_server)
and self._oauth_discovery_retry_due(server.server_id)
)
):
# Re-use existing server instance to avoid re-running build_mcp_server_from_table()
# which can perform network discovery for OAuth2 servers.
@ -5364,6 +5366,7 @@ class MCPServerManager:
# already-decrypted records add_server/update_server are handed.
# Decrypt them while building the registry entry.
new_server = await self.build_mcp_server_from_table(server, env_vars_are_encrypted=True)
self._record_oauth_discovery_outcome(new_server)
# Carry the cached short_prefix from the previous registry entry
# (if any) so the prefix is stable across reloads.
if existing_server is not None and existing_server.short_prefix:

View file

@ -0,0 +1,148 @@
"""One-time heal for MCP server rows whose ``issuer`` a released version wrote by itself.
Until the write was removed, OAuth discovery stamped the issuer it discovered onto the ``issuer``
column trust-on-first-use. That column means "the admin pinned this trust anchor", so the next
registry build read the gateway's own output back as admin intent: the server turned issuer-anchored
(RFC 8414 section 3.3), its stored authorization/token/registration URLs stopped applying, and a
failed issuer-document fetch left it with no authorize endpoint (GH #34985).
Deleting the write fixes every row created afterwards but cannot fix a row already stamped, which
still reads as pinned. This heals those rows by clearing the stamp so their configured endpoints
apply again.
The signal is a heuristic, and deliberately a narrow one. ``updated_by`` records only the most recent
writer, and no audit trail says which field that writer touched, so "discovery wrote this issuer" is
not directly knowable. Two independent clauses bound it, and each rules out a different way of
destroying a pin an admin meant.
Configured endpoints must be present. A deliberately pinned row very often has none, both because the
Issuer field is documented as overriding them and because ``update_mcp_server`` clears them when an
issuer changes, so "issuer set, endpoints empty" is the canonical shape of a real pin and must never
be cleared on this evidence. Skipping those rows costs little: with nothing configured to restore, the
anchored and resource-rooted paths resolve from the same upstream document, and the row still gets the
unresolved-endpoint retry and the anchored-discard warning.
The configured endpoints must also share the issuer's origin. A stamped issuer is by construction the
one self-attested by the authorization-server document discovery reached from this very server, so
endpoints typed alongside it address that same authority. An admin who pinned an issuer and typed
endpoints for a different authority is expressing an intent that clearing the issuer would discard, so
that row is warned about and never healed.
What survives both clauses is a row whose configured endpoints and stamped issuer share an origin,
which is exactly the GH #34985 shape. An admin who pinned that same origin by hand lands here too, and
for them the clear is close to a no-op: their typed endpoints keep serving and still anchor the
RFC 9700 corroboration gate, with only the stricter section 3.3 anchoring lost. Every heal logs the
cleared value so it can be restored, and the clear is recorded under this module's actor so the heal
runs at most once per row.
"""
from typing import Protocol
from urllib.parse import urlparse
from litellm._logging import verbose_proxy_logger
from litellm.proxy._experimental.mcp_server.oauth_utils import canonicalize_url_identity
from litellm.proxy.utils import PrismaClient
# The actor the removed discovery write-back stamped rows with.
_DISCOVERY_ACTOR = "mcp_oauth_discovery"
# The actor recorded on a healed row, which also makes the heal idempotent: once a row is cleared it
# no longer matches ``updated_by == _DISCOVERY_ACTOR`` and is never reconsidered.
_BACKFILL_ACTOR = "mcp_oauth_issuer_stamp_backfill"
_AUTH_TYPES_WITH_ISSUER_ANCHORING = ("oauth2", "true_passthrough", "oauth_delegate")
def _origin(url: str) -> str | None:
"""The scheme-and-authority identity of ``url``, or ``None`` when it has none.
Built on the shared URL canonicalizer so the lowercase-host and default-port rules match the
RFC 8414 issuer comparison the resolution path uses, instead of being re-derived here.
"""
parsed = urlparse(canonicalize_url_identity(url))
if not parsed.scheme or not parsed.netloc:
return None
return f"{parsed.scheme}://{parsed.netloc}"
class _MCPServerRow(Protocol):
"""The MCP server row fields this heal reads, so the untyped DB record is narrowed once here."""
server_id: str
alias: str | None
server_name: str | None
auth_type: str | None
issuer: str | None
authorization_url: str | None
token_url: str | None
registration_url: str | None
updated_by: str | None
def _is_stamped_issuer_row(row: _MCPServerRow) -> bool:
"""Whether this row carries the full signature of a gateway-written issuer stamp.
The whole rule lives here, including the writer check the query also filters on, so the decision
to clear an admin-visible field is auditable in one place rather than split between a predicate
and a query.
"""
if getattr(row, "updated_by", None) != _DISCOVERY_ACTOR:
return False
if not (getattr(row, "issuer", None) or "").strip():
return False
if getattr(row, "auth_type", None) not in _AUTH_TYPES_WITH_ISSUER_ANCHORING:
return False
configured = tuple(
value.strip()
for value in (row.authorization_url, row.token_url, row.registration_url)
if value and value.strip()
)
if not configured:
return False
issuer_origin = _origin(row.issuer or "")
return issuer_origin is not None and all(_origin(endpoint) == issuer_origin for endpoint in configured)
async def backfill_discovery_stamped_issuers(prisma_client: PrismaClient) -> int:
"""Clear gateway-written issuer stamps, returning the number of rows healed."""
candidate_rows: list[_MCPServerRow] = await prisma_client.db.litellm_mcpservertable.find_many(
where={
"updated_by": _DISCOVERY_ACTOR,
"auth_type": {"in": list(_AUTH_TYPES_WITH_ISSUER_ANCHORING)},
},
)
stamped = tuple(row for row in candidate_rows if _is_stamped_issuer_row(row))
if not stamped:
return 0
healed = 0
for row in stamped:
try:
await prisma_client.db.litellm_mcpservertable.update(
where={"server_id": row.server_id},
data={"issuer": None, "updated_by": _BACKFILL_ACTOR},
)
except Exception as exc: # noqa: BLE001 - per-row best effort; the next boot retries
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: could not heal server_id=%s: %s", row.server_id, exc
)
continue
healed += 1
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: cleared issuer %r on server_id=%s (alias=%s). OAuth discovery "
"had written that value onto the Issuer column, which made the server issuer-anchored and "
"fail-closed, and its configured Authorization/Token/Registration URLs were being ignored "
"as a result; those now apply again. If you pinned this issuer deliberately, set it again "
"via the dashboard or PUT /v1/mcp/server to restore RFC 8414 section 3.3 anchoring.",
row.issuer,
row.server_id,
row.alias or row.server_name,
)
if healed:
verbose_proxy_logger.warning(
"MCP issuer stamp backfill: healed %d server(s) whose Issuer had been written by OAuth "
"discovery rather than by an admin",
healed,
)
return healed

View file

@ -12,6 +12,11 @@ import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from litellm._logging import verbose_logger
from litellm.exceptions import (
BlockedPiiEntityError,
GuardrailRaisedException,
ModifyResponseException,
)
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPServerListError,
MCPUpstreamAuthError,
@ -33,6 +38,8 @@ from litellm.proxy.auth.ip_address_utils import IPAddressUtils
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
if TYPE_CHECKING:
from mcp.types import CallToolResult
from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
from litellm.types.mcp import MCPAuth
@ -51,6 +58,13 @@ router = APIRouter(
tags=["mcp"],
)
_MCP_GUARDRAIL_REJECTIONS = (
BlockedPiiEntityError,
GuardrailRaisedException,
ModifyResponseException,
HTTPException,
)
def _connection_error_message(exc: BaseException) -> str:
if isinstance(exc, httpx.LocalProtocolError):
@ -99,9 +113,17 @@ if MCP_AVAILABLE:
end_time: datetime,
user_api_key_auth: UserAPIKeyAuth | None = None,
request_data: Mapping[str, object] | None = None,
) -> None:
) -> "CallToolResult":
"""Fire post-call logging, returning the tool result to send to the client.
``post_mcp_call`` guardrails already ran on ``execute_mcp_tool``'s return
path, so the result arriving here is the guardrailed one. A guardrail
rejection raised by a native ``async_post_mcp_tool_call_hook`` is still
re-raised rather than swallowed as a logging failure, which would return
the unguarded result.
"""
if logging_obj is None:
return
return result
logging_results = await asyncio.gather(
_fire_mcp_tool_call_logging(
logging_obj,
@ -113,11 +135,13 @@ if MCP_AVAILABLE:
),
return_exceptions=True,
)
logging_error = logging_results[0]
if isinstance(logging_error, asyncio.CancelledError):
raise logging_error
if isinstance(logging_error, BaseException):
verbose_logger.warning("MCP tool call logging failed (continuing): %s", logging_error)
outcome = logging_results[0]
if isinstance(outcome, (asyncio.CancelledError, *_MCP_GUARDRAIL_REJECTIONS)):
raise outcome
if isinstance(outcome, BaseException):
verbose_logger.warning("MCP tool call logging failed (continuing): %s", outcome)
return result
return outcome
def _relay_upstream_auth_http_exception(e: MCPUpstreamAuthError, request: Request) -> HTTPException:
"""Convert a client-forwarded pass-through upstream 401 into an HTTPException that preserves the
@ -196,7 +220,7 @@ if MCP_AVAILABLE:
raw_headers=virtual_raw_headers,
litellm_logging_obj=virtual_logging_obj,
)
await _safe_fire_mcp_tool_call_logging(
return await _safe_fire_mcp_tool_call_logging(
virtual_logging_obj,
result,
_tool_start_time,
@ -204,7 +228,6 @@ if MCP_AVAILABLE:
user_api_key_auth=user_api_key_dict,
request_data=data,
)
return result
def _get_server_auth_header(
server,
@ -998,7 +1021,7 @@ if MCP_AVAILABLE:
litellm_logging_obj=data.get("litellm_logging_obj"),
requested_server_id=canonical_server_id,
)
await _safe_fire_mcp_tool_call_logging(
return await _safe_fire_mcp_tool_call_logging(
logging_obj,
result,
_tool_start_time,
@ -1006,7 +1029,6 @@ if MCP_AVAILABLE:
user_api_key_auth=user_api_key_dict,
request_data=data,
)
return result
except MCPMissingUserEnvVarsError as e:
verbose_logger.info(
"MCP tool call missing per-user env vars: server_id=%s missing=%s",

View file

@ -2910,7 +2910,38 @@ if MCP_AVAILABLE:
local_content = await _handle_local_mcp_tool(original_tool_name, arguments)
response = CallToolResult(content=cast(Any, local_content), isError=False)
return response
return await _run_post_mcp_call_guardrails(
result=response,
litellm_logging_obj=litellm_logging_obj,
user_api_key_auth=user_api_key_auth,
request_data=kwargs,
)
async def _run_post_mcp_call_guardrails(
result: CallToolResult,
litellm_logging_obj: LiteLLMLoggingObj | None,
user_api_key_auth: UserAPIKeyAuth | None,
request_data: Mapping[str, object],
) -> CallToolResult:
"""Run ``post_mcp_call`` guardrails over an executed tool result.
Lives on ``execute_mcp_tool``'s return path rather than inside
``_fire_mcp_tool_call_logging`` so enforcement never depends on logging
being configured, and so every dispatch route gets it: the MCP protocol
handler, the REST endpoint, and tool search all funnel through here.
A guardrail that rejects the result raises, matching ``pre_mcp_call``.
"""
from litellm.proxy.proxy_server import proxy_logging_obj
if proxy_logging_obj is None:
return result
return await proxy_logging_obj.post_mcp_call_hook(
response=result,
request_data=(
litellm_logging_obj.model_call_details if litellm_logging_obj is not None else dict(request_data)
),
user_api_key_dict=user_api_key_auth,
)
_MCP_CREDENTIAL_REQUEST_FIELDS = frozenset(
{
@ -2929,8 +2960,14 @@ if MCP_AVAILABLE:
end_time: datetime,
user_api_key_auth: UserAPIKeyAuth | None = None,
request_data: Mapping[str, object] | None = None,
) -> None:
"""Fire post-call logging for an executed MCP tool call.
) -> CallToolResult:
"""Fire post-call logging for an executed MCP tool call, returning the result to send.
The returned result is what the caller must forward to the client: a
``post_mcp_call`` guardrail may rewrite the tool output (e.g. mask
sensitive values) or reject it, in which case its exception propagates.
Guardrails run before the success/failure logging so the masked text, not
the raw one, is what gets logged.
A result with ``isError=True`` is logged as a failure (``status="failure"``
payload, so OTel marks the span ERROR) while the HTTP wire behavior stays
@ -2946,6 +2983,8 @@ if MCP_AVAILABLE:
stripped before the dict is handed to ``post_call_failure_hook``
callbacks.
"""
from litellm.proxy.proxy_server import proxy_logging_obj
logging_obj.post_call(original_response=result)
await logging_obj.async_post_mcp_tool_call_hook(
kwargs=logging_obj.model_call_details,
@ -2957,7 +2996,7 @@ if MCP_AVAILABLE:
error_message = extract_mcp_tool_result_error_message(result)
if error_message is None:
await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
return
return result
logging_obj.has_run_logging(event_type="sync_success")
logging_obj.has_run_logging(event_type="async_success")
@ -2966,8 +3005,7 @@ if MCP_AVAILABLE:
await logging_obj.async_failure_handler(tool_error, "", start_time, end_time)
if user_api_key_auth is None:
return
from litellm.proxy.proxy_server import proxy_logging_obj
return result
if proxy_logging_obj:
sanitized_request_data = {
@ -2979,6 +3017,7 @@ if MCP_AVAILABLE:
user_api_key_dict=user_api_key_auth,
route="/mcp/call_tool",
)
return result
@client
async def call_mcp_tool(
@ -3062,7 +3101,7 @@ if MCP_AVAILABLE:
raise
if litellm_logging_obj:
await _fire_mcp_tool_call_logging(
response = await _fire_mcp_tool_call_logging(
logging_obj=litellm_logging_obj,
result=response,
start_time=start_time,

View file

@ -4,6 +4,7 @@ MCP Server Utilities
import json
import re
from collections.abc import MutableMapping, MutableSequence
from typing import (
Any,
Dict,
@ -434,6 +435,56 @@ def extract_mcp_tool_result_error_message(result: object) -> Optional[str]:
return "MCP tool call returned isError=true"
def mcp_tool_result_content_list(result: object) -> MutableSequence[object] | None: # mutable-ok: see below
"""The mutable content list of an MCP tool result, or ``None`` when it has none.
Deliberately mutable: a guardrail masking the result rewrites entries in place,
because the logging payload captured before the guardrail runs references this
same list, so handing back a copy would leave the unmasked text in the spend log
and the OTel span.
Accepts both ``mcp.types.CallToolResult`` objects and their dict
equivalents, duck-typed so the ``mcp`` package is not required.
"""
content: object = result.get("content") if isinstance(result, Mapping) else getattr(result, "content", None)
if isinstance(content, MutableSequence):
return content
return None
def mcp_content_item_text(item: object) -> str | None:
"""The ``text`` of a rewritable MCP content item, or ``None``.
Only mappings and Pydantic-style models report a text, because those are the
only shapes ``with_mcp_content_item_text`` can rewrite; a caller therefore
never reads text it would be unable to write back (e.g. masked by a
guardrail). Non-text content (images, embedded resources) has no ``text``
and is reported as ``None``.
"""
text: object
if isinstance(item, Mapping):
text = item.get("text")
elif callable(getattr(item, "model_copy", None)):
text = getattr(item, "text", None)
else:
return None
return text if isinstance(text, str) else None
def with_mcp_content_item_text(item: object, text: str) -> object:
"""A copy of an MCP content item carrying ``text`` instead of its own.
Only meaningful for items ``mcp_content_item_text`` returned a text for; any
other item is returned unchanged.
"""
if isinstance(item, Mapping):
return {**item, "text": text}
model_copy = getattr(item, "model_copy", None)
if callable(model_copy):
return model_copy(update={"text": text})
return item
TOOL_DISPLAY_NAME_PATTERN = re.compile(r"^[a-zA-Z0-9_-]+$")
@ -618,3 +669,112 @@ def merge_mcp_headers(
merged.update({str(k): str(v) for k, v in static_headers.items()})
return merged or None
# Local rather than litellm.constants: this module deliberately imports no litellm
# package, so pulling one in for a single integer would drag in litellm/__init__.
MAX_STRUCTURED_CONTENT_SCAN_DEPTH = 100
JSONLeafPath = tuple[str | int, ...]
def _flatten_leaf_groups(
groups: Iterable[tuple[tuple[JSONLeafPath, str], ...] | None],
) -> tuple[tuple[JSONLeafPath, str], ...] | None:
"""Concatenate child leaf groups, propagating the too-deep sentinel."""
materialized = tuple(groups)
if any(group is None for group in materialized):
return None
return tuple(leaf for group in materialized if group is not None for leaf in group)
def json_string_leaves(value: object, path: JSONLeafPath = ()) -> tuple[tuple[JSONLeafPath, str], ...] | None:
"""Depth-first, deterministically ordered string leaves of a JSON value.
Returns ``None`` when the value is nested past ``MAX_STRUCTURED_CONTENT_SCAN_DEPTH``,
so the caller blocks rather than letting deeper values through unscanned; an
empty tuple means there was simply nothing to scan. A sentinel rather than an
exception because this module is reloaded by tests (see the note above the
environment-backed constants), which would give a custom exception class a new
identity and let it escape a caller's ``except``.
"""
if len(path) > MAX_STRUCTURED_CONTENT_SCAN_DEPTH:
return None
if isinstance(value, str):
return ((path, value),)
if isinstance(value, dict):
return _flatten_leaf_groups(json_string_leaves(item, (*path, key)) for key, item in value.items())
if isinstance(value, list):
return _flatten_leaf_groups(json_string_leaves(item, (*path, index)) for index, item in enumerate(value))
return ()
def with_json_string_leaves(
value: object,
replacements: Mapping[JSONLeafPath, str],
path: JSONLeafPath = (),
) -> object:
"""Rebuild a JSON value with the guardrail's rewritten string leaves."""
if isinstance(value, str):
return replacements.get(path, value)
if isinstance(value, dict):
return {key: with_json_string_leaves(item, replacements, (*path, key)) for key, item in value.items()}
if isinstance(value, list):
return [with_json_string_leaves(item, replacements, (*path, index)) for index, item in enumerate(value)]
return value
def json_unrewritable_labels(value: object, path_depth: int = 0) -> tuple[str, ...] | None:
"""Strings in a JSON value that carry meaning but cannot be rewritten.
Dictionary keys and non-string scalars: masking either would change the
payload's contract rather than redact a value, so a caller scans these and
blocks on a match instead of rewriting, matching what the content filter
already does for MCP tool call arguments. ``None`` means the value is nested
past the scan depth, same contract as ``json_string_leaves``.
"""
if path_depth > MAX_STRUCTURED_CONTENT_SCAN_DEPTH:
return None
if isinstance(value, bool) or value is None or isinstance(value, str):
return ()
if isinstance(value, (int, float)):
return (str(value),)
if isinstance(value, dict):
own = tuple(key for key in value if isinstance(key, str))
nested = tuple(json_unrewritable_labels(item, path_depth + 1) for item in value.values())
if any(group is None for group in nested):
return None
return own + tuple(label for group in nested if group is not None for label in group)
if isinstance(value, list):
nested = tuple(json_unrewritable_labels(item, path_depth + 1) for item in value)
if any(group is None for group in nested):
return None
return tuple(label for group in nested if group is not None for label in group)
return ()
def mcp_tool_result_structured_content(result: object) -> object:
"""The ``structuredContent`` of an MCP tool result, or ``None`` when it has none."""
if isinstance(result, Mapping):
return result.get("structuredContent")
return getattr(result, "structuredContent", None)
def set_mcp_tool_result_structured_content(result: object, value: object) -> bool:
"""Replace ``structuredContent`` in place; ``False`` when the shape does not carry it.
In place for the same reason the content list is: the logging payload captured
before the guardrail ran references this object, so a copy would leave the
unmasked value in the spend log and the OTel span.
"""
if isinstance(result, MutableMapping):
result["structuredContent"] = value
return True
if not hasattr(result, "structuredContent"):
return False
try:
setattr(result, "structuredContent", value) # attribute name is fixed by the MCP result shape
return True
except (AttributeError, TypeError, ValueError):
return False

View file

@ -53,6 +53,7 @@ from litellm.types.utils import (
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingRoutingDecision,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse,
@ -1169,7 +1170,6 @@ class GenerateKeyResponse(KeyRequestBase):
class UpdateKeyRequest(KeyRequestBase):
# Note: the defaults of all Params here MUST BE NONE
# else they will get overwritten
key: str # type: ignore
duration: Optional[str] = None
spend: Optional[float] = None
metadata: Optional[dict] = None
@ -1186,6 +1186,12 @@ class UpdateKeyRequest(KeyRequestBase):
raise ValueError("temp_budget_increase and temp_budget_expiry must be set together")
return self
@model_validator(mode="after")
def validate_key_identifier(self) -> "UpdateKeyRequest":
if self.key is None and self.key_alias is None:
raise ValueError("either key or key_alias must be provided")
return self
class RegenerateKeyRequest(GenerateKeyRequest):
# This needs to be different from UpdateKeyRequest, because "key" is optional for this
@ -2455,16 +2461,6 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
"is active as a reminder that hard enforcement is relaxed."
),
)
skip_user_budget_on_team_key: bool | None = Field(
None,
description=(
"If True, restores the legacy behavior where a user's personal "
"max_budget is NOT enforced when their key belongs to a team; only "
"the team (and team-member) budgets apply. Defaults to False, meaning "
"the user's personal max_budget is always enforced regardless of "
"whether the key belongs to a team (see GitHub issue #12905)."
),
)
user_url_validation: Optional[bool] = Field(
None,
description=(
@ -2777,6 +2773,7 @@ class UserInfoV2Response(LiteLLMPydanticObjectBase):
updated_at: Optional[datetime] = None
sso_user_id: Optional[str] = None
teams: List[str] = [] # Just team IDs, not full team objects
object_permission: LiteLLM_ObjectPermissionTable | None = None
from litellm.models.config import LiteLLM_Config as LiteLLM_Config # noqa: E402
@ -3306,6 +3303,7 @@ class SpendLogsMetadata(TypedDict):
applied_guardrails: Optional[List[str]]
mcp_tool_call_metadata: Optional[StandardLoggingMCPToolCall]
vector_store_request_metadata: Optional[List[StandardLoggingVectorStoreRequest]]
routing_decision: StandardLoggingRoutingDecision | None
guardrail_information: Optional[List[StandardLoggingGuardrailInformation]]
eval_information: Optional[Any]
status: StandardLoggingPayloadStatus
@ -4330,7 +4328,13 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
team_id_upsert: bool = False
team_ids_jwt_field: Optional[str] = None
upsert_sso_user_to_team: bool = False
team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"]
team_allowed_routes: List[str] = [
"openai_routes",
"info_routes",
"mcp_routes",
"/v1/messages",
"/v1/messages/count_tokens",
]
team_id_default: Optional[str] = Field(
default=None,
description="If no team_id given, default permissions/spend-tracking to this team.s",

View file

@ -1,105 +1,61 @@
#### Analytics Endpoints #####
from datetime import datetime, timezone
from typing import List, Optional
from typing import Annotated
import fastapi
from fastapi import APIRouter, Depends, HTTPException, status
from litellm.proxy._types import *
from litellm.proxy.analytics_endpoints.cache_activity import CacheActivityResponse, get_cache_activity
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
def _parse_date(value: str, param_name: str) -> datetime:
try:
return datetime.strptime(value, "%Y-%m-%d").replace(tzinfo=timezone.utc)
except ValueError:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"{param_name} must be a YYYY-MM-DD date, got {value!r}"},
)
@router.get(
"/global/activity/cache_hits",
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
responses={
200: {"model": List[LiteLLM_SpendLogs]},
},
response_model=CacheActivityResponse,
include_in_schema=False,
)
async def get_global_activity(
start_date: Optional[str] = fastapi.Query(
default=None,
description="Time from which to start viewing spend",
),
end_date: Optional[str] = fastapi.Query(
default=None,
description="Time till which to view spend",
),
):
start_date: Annotated[str, fastapi.Query(description="Time from which to start viewing spend")],
end_date: Annotated[str, fastapi.Query(description="Time till which to view spend")],
key_aliases: Annotated[
list[str] | None, fastapi.Query(description="Only include spend from these key aliases")
] = None,
models: Annotated[list[str] | None, fastapi.Query(description="Only include spend for these models")] = None,
) -> CacheActivityResponse:
"""
Get number of cache hits, vs misses
{
"daily_data": [
const chartdata = [
{
date: 'Jan 22',
cache_hits: 10,
llm_api_calls: 2000
},
{
date: 'Jan 23',
cache_hits: 10,
llm_api_calls: 12
},
],
"sum_cache_hits": 20,
"sum_llm_api_calls": 2012
}
Cache activity for the Admin UI cache dashboard, aggregated per call_type:
cache hits vs successful LLM API requests vs failed requests, plus totals
for the stat cards and the available key-alias/model filter options.
"""
if start_date is None or end_date is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Please provide start_date and end_date"},
)
start_date_obj = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
end_date_obj = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=timezone.utc)
from litellm.proxy.proxy_server import prisma_client
try:
if prisma_client is None:
raise ValueError(
"Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
)
sql_query = """
SELECT
CASE
WHEN vt."key_alias" IS NOT NULL THEN vt."key_alias"
ELSE 'Unnamed Key'
END AS api_key,
sl."call_type",
sl."model",
COUNT(*) AS total_rows,
SUM(CASE WHEN sl."cache_hit" = 'True' THEN 1 ELSE 0 END) AS cache_hit_true_rows,
SUM(CASE WHEN sl."cache_hit" = 'True' THEN sl."completion_tokens" ELSE 0 END) AS cached_completion_tokens,
SUM(CASE WHEN sl."cache_hit" != 'True' THEN sl."completion_tokens" ELSE 0 END) AS generated_completion_tokens
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
GROUP BY
vt."key_alias",
sl."call_type",
sl."model"
"""
db_response = await prisma_client.db.query_raw(sql_query, start_date_obj, end_date_obj)
if db_response is None:
return []
return db_response
except Exception as e:
if prisma_client is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": str(e)},
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={
"error": "Database not connected. Connect a database to your proxy - https://docs.litellm.ai/docs/simple_proxy#managing-auth---virtual-keys"
},
)
return await get_cache_activity(
prisma_client=prisma_client,
start_date=_parse_date(start_date, "start_date"),
end_date=_parse_date(end_date, "end_date"),
key_aliases=key_aliases or [],
models=models or [],
)

View file

@ -0,0 +1,137 @@
import asyncio
import json
from datetime import datetime
from typing import TYPE_CHECKING, Sequence
from pydantic import BaseModel, TypeAdapter
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
UNKNOWN_CALL_TYPE = "Unknown"
class CacheActivityGroup(BaseModel):
call_type: str
api_requests: int
cache_hits: int
failed_requests: int
cached_completion_tokens: int
generated_completion_tokens: int
class CacheActivityTotals(BaseModel):
api_requests: int
cache_hits: int
failed_requests: int
cached_completion_tokens: int
cache_hit_ratio: float
class CacheActivityFilterOptions(BaseModel):
key_aliases: list[str]
models: list[str]
class CacheActivityResponse(BaseModel):
groups: list[CacheActivityGroup]
totals: CacheActivityTotals
filter_options: CacheActivityFilterOptions
GROUPS_SQL = """
SELECT
CASE WHEN sl."call_type" = '' THEN 'Unknown' ELSE sl."call_type" END AS call_type,
(COUNT(*)
- SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)
- SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END))::int AS api_requests,
SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN 1 ELSE 0 END)::int AS cache_hits,
SUM(CASE WHEN sl."status" = 'failure' THEN 1 ELSE 0 END)::int AS failed_requests,
SUM(CASE WHEN COALESCE(sl."cache_hit", '') = 'True' THEN sl."completion_tokens" ELSE 0 END)::int
AS cached_completion_tokens,
SUM(CASE WHEN COALESCE(sl."cache_hit", '') != 'True' THEN sl."completion_tokens" ELSE 0 END)::int
AS generated_completion_tokens
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
AND ($3::jsonb = '[]'::jsonb
OR COALESCE(vt."key_alias", 'Unnamed Key') IN (SELECT jsonb_array_elements_text($3::jsonb)))
AND ($4::jsonb = '[]'::jsonb
OR sl."model" IN (SELECT jsonb_array_elements_text($4::jsonb)))
GROUP BY 1
ORDER BY (COUNT(*)) DESC
"""
KEY_ALIAS_OPTIONS_SQL = """
SELECT DISTINCT COALESCE(vt."key_alias", 'Unnamed Key') AS key_alias
FROM "LiteLLM_SpendLogs" sl
LEFT JOIN "LiteLLM_VerificationToken" vt ON sl."api_key" = vt."token"
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
ORDER BY 1
"""
MODEL_OPTIONS_SQL = """
SELECT DISTINCT sl."model" AS model
FROM "LiteLLM_SpendLogs" sl
WHERE
sl."startTime" >= ($1::timestamptz AT TIME ZONE 'UTC')
AND sl."startTime" < (($2::timestamptz + INTERVAL '1 day') AT TIME ZONE 'UTC')
AND sl."model" != ''
ORDER BY 1
"""
class _KeyAliasRow(BaseModel):
key_alias: str
class _ModelRow(BaseModel):
model: str
_groups_adapter = TypeAdapter(list[CacheActivityGroup])
_key_alias_rows_adapter = TypeAdapter(list[_KeyAliasRow])
_model_rows_adapter = TypeAdapter(list[_ModelRow])
def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals:
api_requests = sum(group.api_requests for group in groups)
cache_hits = sum(group.cache_hits for group in groups)
failed_requests = sum(group.failed_requests for group in groups)
all_requests = api_requests + cache_hits + failed_requests
return CacheActivityTotals(
api_requests=api_requests,
cache_hits=cache_hits,
failed_requests=failed_requests,
cached_completion_tokens=sum(group.cached_completion_tokens for group in groups),
cache_hit_ratio=(cache_hits / all_requests) * 100 if all_requests > 0 else 0.0,
)
async def get_cache_activity(
prisma_client: "PrismaClient",
start_date: datetime,
end_date: datetime,
key_aliases: Sequence[str],
models: Sequence[str],
) -> CacheActivityResponse:
group_rows, key_alias_rows, model_rows = await asyncio.gather(
prisma_client.db.query_raw(
GROUPS_SQL, start_date, end_date, json.dumps(list(key_aliases)), json.dumps(list(models))
),
prisma_client.db.query_raw(KEY_ALIAS_OPTIONS_SQL, start_date, end_date),
prisma_client.db.query_raw(MODEL_OPTIONS_SQL, start_date, end_date),
)
groups = _groups_adapter.validate_python(group_rows or [])
return CacheActivityResponse(
groups=groups,
totals=compute_totals(groups),
filter_options=CacheActivityFilterOptions(
key_aliases=[row.key_alias for row in _key_alias_rows_adapter.validate_python(key_alias_rows or [])],
models=[row.model for row in _model_rows_adapter.validate_python(model_rows or [])],
),
)

View file

@ -74,6 +74,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
from litellm.proxy.common_utils.user_api_key_cache import (
UserApiKeyCache,
get_management_object_ttl,
object_permission_cache_key,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.guardrails.tool_name_extraction import (
@ -485,6 +486,14 @@ MODEL_DISCOVERY_ROUTES = frozenset(
}
)
BUDGET_ENFORCED_SIDE_EFFECT_ROUTES = frozenset(
{
"/health",
"/health/services",
"/health/test_connection",
}
)
async def common_checks(
request_body: dict,
@ -531,8 +540,10 @@ async def common_checks(
request=request,
)
if route in MODEL_DISCOVERY_ROUTES:
skip_budget_checks = True
skip_all_budget_checks = skip_budget_checks or (
route not in BUDGET_ENFORCED_SIDE_EFFECT_ROUTES
and (route in MODEL_DISCOVERY_ROUTES or not RouteChecks.is_llm_api_route(route=route))
)
# 1. If team is blocked
if team_object is not None and team_object.blocked is True:
@ -606,7 +617,7 @@ async def common_checks(
project_object=project_object,
_model=_model,
llm_router=llm_router,
skip_budget_checks=skip_budget_checks,
skip_budget_checks=skip_all_budget_checks,
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
@ -615,7 +626,7 @@ async def common_checks(
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
# If this is a free model, skip all budget checks
if not skip_budget_checks:
if not skip_all_budget_checks:
# Key metadata.tags are injected into request_body here so the tag budget
# check can read them; this mutation must run before the gathered checks.
if valid_token is not None:
@ -632,31 +643,28 @@ async def common_checks(
)
async def _user_max_budget_check() -> None:
if user_object is None or user_object.max_budget is None:
return
skip_for_team = (
general_settings.get("skip_user_budget_on_team_key") is True
and team_object is not None
and team_object.team_id is not None
)
if skip_for_team:
return
from litellm.proxy.proxy_server import get_current_spend
# 4.1 personal budget, if personal key
if (
(team_object is None or team_object.team_id is None)
and user_object is not None
and user_object.max_budget is not None
):
from litellm.proxy.proxy_server import get_current_spend
user_budget = user_object.max_budget
user_spend = await get_current_spend(
counter_key=f"spend:user:{user_object.user_id}",
fallback_spend=user_object.spend or 0.0,
max_budget=user_budget,
)
if math.isfinite(user_budget) and user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_spend,
user_budget = user_object.max_budget
user_spend = await get_current_spend(
counter_key=f"spend:user:{user_object.user_id}",
fallback_spend=user_object.spend or 0.0,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
entity_type=Litellm_EntityType.USER.value,
entity_id=user_object.user_id,
)
if math.isfinite(user_budget) and user_spend >= user_budget:
raise litellm.BudgetExceededError(
current_cost=user_spend,
max_budget=user_budget,
message=f"ExceededBudget: User={user_object.user_id} over budget. Spend={user_spend}, Budget={user_budget}",
entity_type=Litellm_EntityType.USER.value,
entity_id=user_object.user_id,
)
# Each scope reads a distinct counter key with no cross-scope ordering
# dependency, so the per-scope Redis-first reads run concurrently instead
@ -715,7 +723,7 @@ async def common_checks(
raise budget_error
_enforce_user_param_check(general_settings, request, request_body, route)
_global_proxy_budget_check(global_proxy_spend, skip_budget_checks, route)
_global_proxy_budget_check(global_proxy_spend, skip_all_budget_checks, route)
_guardrail_modification_check(request_body, team_object)
# 10 [OPTIONAL] Organization RBAC checks
@ -953,7 +961,7 @@ async def get_default_end_user_budget(
)
return None
_budget_obj = LiteLLM_BudgetTable(**budget_record.dict())
_budget_obj = LiteLLM_BudgetTable.model_validate(budget_record.dict())
# Cache the budget for 60 seconds
await user_api_key_cache.async_set_cache(
key=cache_key,
@ -999,7 +1007,7 @@ async def get_team_member_default_budget(
if isinstance(cached_budget, LiteLLM_BudgetTable):
return cached_budget
if isinstance(cached_budget, dict):
return LiteLLM_BudgetTable(**cached_budget)
return LiteLLM_BudgetTable.model_validate(cached_budget)
try:
budget_record = await BudgetRepository(prisma_client).table.find_unique(where={"budget_id": budget_id})
@ -1014,7 +1022,7 @@ async def get_team_member_default_budget(
ttl=get_management_object_ttl(user_api_key_cache),
)
return LiteLLM_BudgetTable(**budget_record.dict())
return LiteLLM_BudgetTable.model_validate(budget_record.dict())
except Exception:
verbose_proxy_logger.exception(f"Error fetching team-default member budget {budget_id}")
@ -1168,7 +1176,7 @@ async def get_end_user_object(
raise Exception
# Convert to LiteLLM_EndUserTable object
_response = LiteLLM_EndUserTable(**response.dict())
_response = LiteLLM_EndUserTable.model_validate(response.dict())
# Apply default budget if needed
_response = await _apply_default_budget_to_end_user(
@ -1360,7 +1368,7 @@ async def get_tag_objects_batch(
for db_tag in db_tags:
tag_name = db_tag.tag_name
cache_key = f"tag:{tag_name}"
_tag_obj = LiteLLM_TagTable(**db_tag.dict())
_tag_obj = LiteLLM_TagTable.model_validate(db_tag.dict())
await user_api_key_cache.async_set_cache(
key=cache_key,
value=_tag_obj,
@ -1453,7 +1461,7 @@ async def get_team_membership(
if response is None:
return None
_response = LiteLLM_TeamMembership(**response.dict())
_response = LiteLLM_TeamMembership.model_validate(response.dict())
await user_api_key_cache.async_set_cache(
key=_key,
value=_response,
@ -1719,13 +1727,13 @@ async def get_user_object(
if response.organization_memberships is not None and len(response.organization_memberships) > 0:
# dump each organization membership to type LiteLLM_OrganizationMembershipTable
_dumped_memberships = [
LiteLLM_OrganizationMembershipTable(**membership.model_dump())
LiteLLM_OrganizationMembershipTable.model_validate(membership.model_dump())
for membership in response.organization_memberships
if membership is not None
]
response.organization_memberships = _dumped_memberships
_response = LiteLLM_UserTable(**dict(response))
_response = LiteLLM_UserTable.model_validate(dict(response))
response_dict = _response.model_dump()
# save the user object to cache
@ -1781,9 +1789,22 @@ async def _cache_team_object(
## CACHE REFRESH TIME!
team_table.last_refreshed_at = time.time()
key = "team_id:{}".format(team_id)
if proxy_logging_obj is not None:
try:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=key)
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the write
verbose_proxy_logger.warning(
"Failed to invalidate internal usage cache entry %s; "
"a stale team object may be served until its TTL expires: %s",
key,
e,
)
# team_id is the table primary key — guaranteed unique, safe to write.
await _cache_management_object(
key="team_id:{}".format(team_id),
key=key,
value=team_table,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
@ -1805,9 +1826,17 @@ async def _cache_team_object(
# the cache from a verified single row.
if team_table.team_alias:
alias_key = "team_alias:{}".format(team_table.team_alias)
user_api_key_cache.delete_cache(key=alias_key)
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key)
try:
user_api_key_cache.delete_cache(key=alias_key)
if proxy_logging_obj is not None:
await proxy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache(key=alias_key)
except Exception as e: # noqa: BLE001 # best-effort invalidation: any cache backend error must not fail the mutation
verbose_proxy_logger.warning(
"Failed to invalidate cached team alias entry %s; "
"a stale team object may be served until its TTL expires: %s",
alias_key,
e,
)
async def _cache_key_object(
@ -1862,7 +1891,7 @@ async def _get_team_db_check(team_id: str, prisma_client: PrismaClient, team_id_
http_request=mock_request,
user_api_key_dict=system_admin_user,
)
response = LiteLLM_TeamTable(**created_team_dict)
response = LiteLLM_TeamTable.model_validate(created_team_dict)
return response
@ -1894,7 +1923,7 @@ async def _get_team_object_from_user_api_key_cache(
if response is None:
raise Exception
_response = LiteLLM_TeamTableCachedObj(**response.dict())
_response = LiteLLM_TeamTableCachedObj.model_validate(response.dict())
# Load object_permission if object_permission_id exists but object_permission is not loaded
if _response.object_permission_id and not _response.object_permission:
@ -2085,7 +2114,7 @@ async def get_access_object(
detail={"error": f"Access group doesn't exist in db. Access group={access_group_id}."},
)
_response = LiteLLM_AccessGroupTable(**response.dict())
_response = LiteLLM_AccessGroupTable.model_validate(response.dict())
# Save to cache
await _cache_access_object(
@ -2170,7 +2199,7 @@ async def get_team_object_by_alias(
)
team = teams[0]
team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump())
team_obj = LiteLLM_TeamTableCachedObj.model_validate(team.model_dump())
# Load object_permission if object_permission_id exists but object_permission is not loaded
if team_obj.object_permission_id and not team_obj.object_permission:
@ -2272,7 +2301,7 @@ async def get_org_object_by_alias(
)
org = orgs[0]
org_obj = LiteLLM_OrganizationTable(**org.model_dump())
org_obj = LiteLLM_OrganizationTable.model_validate(org.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
@ -2413,7 +2442,7 @@ class ExperimentalUIJWTToken:
if decrypted_token is None:
return None
try:
return UserAPIKeyAuth(**json.loads(decrypted_token))
return UserAPIKeyAuth.model_validate(json.loads(decrypted_token))
except Exception as e:
raise Exception(f"Invalid hash key. Hash key={hashed_token}. Decrypted token={decrypted_token}. Error: {e}")
@ -2532,7 +2561,7 @@ async def get_key_object(
code=status.HTTP_401_UNAUTHORIZED,
)
_response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True))
_response = UserAPIKeyAuth.model_validate(_valid_token.model_dump(exclude_none=True))
# Load object_permission if object_permission_id exists but object_permission is not loaded
if _response.object_permission_id and not _response.object_permission:
@ -2588,7 +2617,7 @@ async def get_object_permission(
raise Exception("No DB Connected. See - https://docs.litellm.ai/docs/proxy/virtual_keys")
# check if in cache
key = "object_permission_id:{}".format(object_permission_id)
key = object_permission_cache_key(object_permission_id)
deserialized_perm = await user_api_key_cache.async_get_cache(
key=key,
model_type=LiteLLM_ObjectPermissionTable,
@ -2605,7 +2634,7 @@ async def get_object_permission(
if response is None:
return None
_perm_obj = LiteLLM_ObjectPermissionTable(**response.dict())
_perm_obj = LiteLLM_ObjectPermissionTable.model_validate(response.dict())
await user_api_key_cache.async_set_cache(
key=key,
value=_perm_obj,
@ -2665,7 +2694,7 @@ async def get_managed_vector_store_rows_by_uuids(
row_dict = dict(row) if hasattr(row, "__dict__") else {}
if not row_dict:
continue
cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict)
cached_obj = LiteLLM_ManagedVectorStoresTable.model_validate(row_dict)
key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id)
await user_api_key_cache.async_set_cache(
key=key,
@ -2746,7 +2775,7 @@ async def get_org_object(
f"Organization doesn't exist in db. Organization={org_id}. Create organization via `/organization/new` call."
)
_org_obj = LiteLLM_OrganizationTable(**response.model_dump())
_org_obj = LiteLLM_OrganizationTable.model_validate(response.model_dump())
# Cache the result
await user_api_key_cache.async_set_cache(
key=cache_key,
@ -4221,7 +4250,7 @@ async def get_project_object(
if project_row is None:
return None
project_obj = LiteLLM_ProjectTableCachedObj(**project_row.model_dump())
project_obj = LiteLLM_ProjectTableCachedObj.model_validate(project_row.model_dump())
# Cache with TTL following _cache_management_object pattern
project_obj.last_refreshed_at = time.time()

View file

@ -1432,7 +1432,7 @@ def _extract_models_from_managed_resource_id(
)
_append_model_candidates(
candidates=candidates,
value=get_model_id_from_unified_batch_id(unified_file_id),
value=_resolve_model_id_with_router(get_model_id_from_unified_batch_id(unified_file_id), llm_router),
)
except Exception as e:
verbose_proxy_logger.debug("Unable to extract model from managed file/batch ID: %s", str(e))
@ -1442,7 +1442,10 @@ def _extract_models_from_managed_resource_id(
parsed_id = parse_unified_id(resource_id)
if parsed_id:
_append_model_candidates(candidates=candidates, value=parsed_id.get("model_id"))
_append_model_candidates(
candidates=candidates,
value=_resolve_model_id_with_router(parsed_id.get("model_id"), llm_router),
)
_append_model_candidates(candidates=candidates, value=parsed_id.get("target_model_names"))
except Exception as e:
verbose_proxy_logger.debug("Unable to extract model from unified managed resource ID: %s", str(e))

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