Merge branch 'litellm_internal_staging' into litellm_fix_mcp_streaming_tool_call_leak

This commit is contained in:
fangkangmi 2026-07-08 13:19:25 +01:00
commit 1af14ae4bf
292 changed files with 12782 additions and 5012 deletions

View file

@ -122,10 +122,28 @@ jobs:
makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false";
}
try {
await github.rest.git.createRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `refs/tags/${tag}`,
sha: commitHash,
});
} catch (error) {
if (error.status !== 422) throw error;
const existing = await github.rest.git.getRef({
owner: context.repo.owner,
repo: context.repo.repo,
ref: `tags/${tag}`,
});
if (existing.data.object.sha !== commitHash) {
throw new Error(`Tag ${tag} already exists at ${existing.data.object.sha}, expected ${commitHash}`);
}
}
const response = await github.rest.repos.createRelease({
draft: true,
generate_release_notes: true,
target_commitish: commitHash,
name: tag,
owner: context.repo.owner,
prerelease: isPrerelease,
@ -138,11 +156,21 @@ jobs:
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
body: updatedBody,
draft: false,
make_latest: makeLatest,
});
if (!isPrerelease) {
await github.rest.repos.updateRelease({
owner: context.repo.owner,
repo: context.repo.repo,
release_id: response.data.id,
tag_name: tag,
make_latest: makeLatest,
});
}
} catch (error) {
core.setFailed(error.message);
}

2
.gitignore vendored
View file

@ -130,3 +130,5 @@ crash.*.log
# pytest coverage data
.coverage
ui/litellm-dashboard/out/

View file

@ -17,6 +17,8 @@ Same thing for bug fixes. The tests should make it so that this specific bug can
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
Always use @.github/pull_request_template.md as a guide for your PR body

View file

@ -4,7 +4,7 @@
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
info lint lint-dev format \
info lint lint-dev lint-checks format \
lint-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
install-dev install-proxy-dev install-test-deps install-hooks \
@ -53,6 +53,11 @@ help:
UV := uv
UV_RUN := $(UV) run --no-sync
LINT_DEP_INSTALL ?= install-dev
LINT_DEP_BASE ?= lint-fetch-base
LINT_JOBS := $(shell sysctl -n hw.ncpu 2>/dev/null || nproc 2>/dev/null || echo 4)
LINT_OUTPUT_SYNC := $(if $(filter output-sync,$(.FEATURES)),--output-sync=target,)
# Show info
info:
@echo "UV: $(UV)"
@ -107,12 +112,12 @@ lint-fetch-base:
# running proxy need.
lint-install:
$(UV) sync --inexact --frozen --group proxy-dev
$(UV_RUN) prisma generate --schema litellm/proxy/schema.prisma
$(UV_RUN) python scripts/prisma_generate_if_needed.py
# Diff-scoped format check, identical to test-linting.yml's "Check ruff format" step:
# only the litellm Python files changed vs the base are checked, so a pre-existing
# format issue elsewhere doesn't block an unrelated commit.
lint-format-check-changed: install-dev lint-fetch-base
lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
@files=$$(git diff --name-only origin/litellm_internal_staging...HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' || true); \
if [ -z "$$files" ]; then \
echo "No changed litellm Python files to format-check."; \
@ -121,7 +126,7 @@ lint-format-check-changed: install-dev lint-fetch-base
fi
# Linting targets
lint-ruff: install-dev
lint-ruff: $(LINT_DEP_INSTALL)
cd litellm && $(UV_RUN) ruff check . && cd ..
# faster linter for developing ...
@ -156,12 +161,12 @@ 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: install-dev lint-fetch-base
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
# Type-discipline budget (mutable collections / casts / type guards / kwargs /
# unexplained suppressions), the test-linting.yml step `make lint` used to omit.
lint-type-discipline: install-dev lint-fetch-base
lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
$(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging
# --update lowers each limit by what this branch fixed since its branch point, so
@ -176,7 +181,7 @@ lint-ruff-budget: install-dev
# Strict gate, invoked the same way CI does in test-linting.yml so a local pass
# means the CI check will pass too.
lint-gate: install-dev lint-fetch-base
lint-gate: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
$(UV_RUN) python scripts/ruff_strict_gate.py --base origin/litellm_internal_staging
lint-ruff-budget-update: install-dev lint-fetch-base
@ -188,10 +193,10 @@ lint-type-discipline-budget-update: install-dev lint-fetch-base
# Ratchet all budgets in one shot (ruff strict + type-discipline + basedpyright)
lint-budget-update: lint-ruff-budget-update lint-type-discipline-budget-update lint-basedpyright-budget-update
check-circular-imports: install-dev
check-circular-imports: $(LINT_DEP_INSTALL)
cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd ..
check-import-safety: install-dev
check-import-safety: $(LINT_DEP_INSTALL)
@$(UV_RUN) python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
# Combined linting, isomorphic to test-linting.yml's lint job so a local pass means a
@ -199,9 +204,13 @@ check-import-safety: install-dev
# runs the diff-scoped ruff format check, whole-tree ruff check, the strict-rule /
# type-discipline / basedpyright budgets as a delta vs the base, then the circular-import
# and import-safety checks. Steps that compare against the base resolve it the same way CI
# does (merge-base with origin/litellm_internal_staging). lint-install is first so the
# Prisma client exists before basedpyright runs.
lint: lint-install lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright check-circular-imports check-import-safety
# does (merge-base with origin/litellm_internal_staging). Setup (env sync, Prisma client,
# base fetch) runs once up front; the checks themselves are independent, so a sub-make
# fans them out with -j and the fast ones finish under basedpyright's shadow.
lint: lint-install lint-fetch-base
$(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_DEP_BASE= lint-checks
lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright check-circular-imports check-import-safety
# Faster linting for local development (only checks changed code)
lint-dev: lint-format-changed check-circular-imports check-import-safety

View file

@ -141,6 +141,6 @@
"limit": 1005
},
"reportUnusedVariable": {
"limit": 1298
"limit": 1297
}
}

View file

@ -57,8 +57,6 @@ source ~/.nvm/nvm.sh
nvm install v18.17.0
nvm use v18.17.0
# copy _enterprise.json from this directory to /ui/litellm-dashboard, and rename it to ui_colors.json
cp enterprise/enterprise_ui/enterprise_colors.json ui/litellm-dashboard/ui_colors.json
# cd in to /ui/litellm-dashboard
cd ui/litellm-dashboard

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.45"
version = "0.1.46"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.45"
version = "0.1.46"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "max_concurrent_requests" INTEGER;

View file

@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
byok_api_key_help_url String?
source_url String?
timeout Float?
max_concurrent_requests Int?
// BYOM submission lifecycle
approval_status String? @default("active")
submitted_by String?

View file

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

View file

@ -588,6 +588,7 @@ gemini_models: Set = set()
xai_models: Set = set()
zai_models: Set = set()
deepseek_models: Set = set()
tencent_models: Set = set()
runwayml_models: Set = set()
azure_ai_models: Set = set()
jina_ai_models: Set = set()
@ -801,6 +802,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
fal_ai_models.add(key)
elif value.get("litellm_provider") == "deepseek":
deepseek_models.add(key)
elif value.get("litellm_provider") == "tencent":
tencent_models.add(key)
elif value.get("litellm_provider") == "runwayml":
runwayml_models.add(key)
elif value.get("litellm_provider") == "meta_llama":
@ -1093,6 +1096,7 @@ models_by_provider: dict = {
"zai": zai_models,
"fal_ai": fal_ai_models,
"deepseek": deepseek_models,
"tencent": tencent_models,
"runwayml": runwayml_models,
"mistral": mistral_chat_models,
"azure_ai": azure_ai_models,
@ -1804,6 +1808,9 @@ if TYPE_CHECKING:
from .llms.deepseek.chat.transformation import (
DeepSeekChatConfig as _DeepSeekChatConfig,
)
from .llms.tencent.chat.transformation import (
TencentChatConfig as _TencentChatConfig,
)
from .llms.sap.chat.transformation import (
GenAIHubOrchestrationConfig as _GenAIHubOrchestrationConfig,
)
@ -1846,6 +1853,7 @@ if TYPE_CHECKING:
# Type stubs for lazy-loaded config classes (to help mypy understand types)
VLLMConfig: Type[_VLLMConfig]
DeepSeekChatConfig: Type[_DeepSeekChatConfig]
TencentChatConfig: Type[_TencentChatConfig]
GenAIHubOrchestrationConfig: Type[_GenAIHubOrchestrationConfig]
GenAIHubEmbeddingConfig: Type[_GenAIHubEmbeddingConfig]
AzureOpenAIO1Config: Type[_AzureOpenAIO1Config]

View file

@ -284,6 +284,7 @@ LLM_CONFIG_NAMES = (
"LiteLLMProxyChatConfig",
"VLLMConfig",
"DeepSeekChatConfig",
"TencentChatConfig",
"LMStudioChatConfig",
"LmStudioEmbeddingConfig",
"NscaleConfig",
@ -1096,6 +1097,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
),
"VLLMConfig": (".llms.vllm.completion.transformation", "VLLMConfig"),
"DeepSeekChatConfig": (".llms.deepseek.chat.transformation", "DeepSeekChatConfig"),
"TencentChatConfig": (".llms.tencent.chat.transformation", "TencentChatConfig"),
"LMStudioChatConfig": (".llms.lm_studio.chat.transformation", "LMStudioChatConfig"),
"LmStudioEmbeddingConfig": (
".llms.lm_studio.embed.transformation",

View file

@ -129,6 +129,33 @@ def _set_agent_id_on_logging_obj(
litellm_logging_obj.model_call_details["agent_id"] = agent_id
_A2A_COST_PARAM_KEYS = ("cost_per_query", "input_cost_per_token", "output_cost_per_token")
def _set_litellm_params_on_logging_obj(
kwargs: dict[str, Any],
litellm_params: dict[str, Any],
) -> None:
"""
Merge the agent's pricing params into model_call_details["litellm_params"]
so A2ACostCalculator can read them.
The non-streaming path reuses the proxy-built logging object, whose
litellm_params already carries metadata / proxy_server_request / user-key
context, so merge the pricing keys in rather than replacing the dict.
"""
logging_obj = kwargs.get("litellm_logging_obj")
if logging_obj is None:
return
cost_params = {key: litellm_params[key] for key in _A2A_COST_PARAM_KEYS if litellm_params.get(key) is not None}
if not cost_params:
return
existing = logging_obj.model_call_details.get("litellm_params") or {}
logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params}
def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str:
"""
Extract agent info and set model/custom_llm_provider for cost tracking.
@ -477,6 +504,9 @@ async def asend_message(
completion_tokens=completion_tokens,
)
# Merge agent pricing params into the logging obj so cost is calculated
_set_litellm_params_on_logging_obj(kwargs=kwargs, litellm_params=litellm_params)
# Set agent_id on logging obj for SpendLogs tracking
_set_agent_id_on_logging_obj(kwargs=kwargs, agent_id=agent_id)

View file

@ -121,8 +121,13 @@ class A2ARequestUtils:
Returns:
Tuple of (prompt_tokens, completion_tokens, total_tokens)
"""
# Count input tokens
# Count input tokens. Dump the message to a dict first so extraction hits
# the dict branch — request-side parts are a2a-sdk Part RootModels whose
# kind/text live on part.root, which the object branch cannot read. This
# mirrors how the response side already works (it operates on model_dump).
input_message = A2ARequestUtils.get_input_message_from_request(request)
if input_message is not None and hasattr(input_message, "model_dump"):
input_message = input_message.model_dump(mode="json")
input_text = A2ARequestUtils.extract_text_from_message(input_message)
prompt_tokens = A2ARequestUtils.count_tokens(input_text)

View file

@ -508,6 +508,7 @@ LITELLM_CHAT_PROVIDERS = [
"text-completion-codestral",
"text-completion-inception",
"deepseek",
"tencent",
"sambanova",
"maritalk",
"cloudflare",
@ -729,6 +730,7 @@ openai_compatible_providers: List = [
"volcengine",
"codestral",
"deepseek",
"tencent",
"deepinfra",
"perplexity",
"xinference",

View file

@ -52,6 +52,9 @@ from litellm.llms.databricks.cost_calculator import (
from litellm.llms.deepseek.cost_calculator import (
cost_per_token as deepseek_cost_per_token,
)
from litellm.llms.tencent.cost_calculator import (
cost_per_token as tencent_cost_per_token,
)
from litellm.llms.fireworks_ai.cost_calculator import (
cost_per_token as fireworks_ai_cost_per_token,
)
@ -625,6 +628,8 @@ def cost_per_token(
return gemini_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
elif custom_llm_provider == "deepseek":
return deepseek_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "tencent":
return tencent_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "perplexity":
return perplexity_cost_per_token(model=model, usage=usage_block)
elif custom_llm_provider == "xai":

View file

@ -1165,12 +1165,18 @@ class ModifyResponseException(Exception):
request_data: Dict[str, Any],
guardrail_name: Optional[str] = None,
detection_info: Optional[Dict[str, Any]] = None,
original_response: Optional[Any] = None,
):
self.message = message
self.model = model
self.request_data = request_data
self.guardrail_name = guardrail_name
self.detection_info = detection_info or {}
# The LLM response that was blocked (post-call). Carries the real token
# usage the upstream call consumed, so the synthetic block response can
# report it instead of discarding it. None for pre-call blocks (the LLM
# was never invoked).
self.original_response = original_response
super().__init__(message)

View file

@ -4,6 +4,7 @@
from __future__ import annotations
import asyncio
import math
import os
import sys
from datetime import datetime, timedelta
@ -65,6 +66,26 @@ if TYPE_CHECKING:
else:
AsyncIOScheduler = Any
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0
def _get_budget_metrics_per_request_timeout() -> float:
raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT")
if raw is None:
return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
try:
parsed = float(raw)
except ValueError:
parsed = None
if parsed is None or not math.isfinite(parsed) or parsed <= 0:
verbose_logger.debug(
"[Non-Blocking] Prometheus: invalid PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT=%r; using default %ss.",
raw,
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT,
)
return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
return parsed
class PrometheusLogger(CustomLogger):
# Class variables or attributes
@ -1607,7 +1628,15 @@ class PrometheusLogger(CustomLogger):
_user_spend = _metadata.get("user_api_key_user_spend", None)
_user_max_budget = _metadata.get("user_api_key_user_max_budget", None)
results = await asyncio.gather(
# Bound the per-request budget-metric emission so that slow Redis/DB
# lookups under load cannot consume the whole LoggingWorker watchdog
# (LOGGING_WORKER_MAX_TIME_PER_COROUTINE, default 20s) and get the entire
# success-logging event cancelled. Budget gauges are also refreshed by the
# periodic cron every PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES,
# so dropping one slow per-request emission only loses sub-cron real-time
# detail, not correctness.
budget_metrics_timeout = _get_budget_metrics_per_request_timeout()
gather_coro = asyncio.gather(
self._set_api_key_budget_metrics_after_api_request(
user_api_key=user_api_key,
user_api_key_alias=user_api_key_alias,
@ -1634,6 +1663,16 @@ class PrometheusLogger(CustomLogger):
),
return_exceptions=True,
)
try:
results = await asyncio.wait_for(gather_coro, timeout=budget_metrics_timeout)
except asyncio.TimeoutError:
verbose_logger.debug(
"[Non-Blocking] Prometheus: per-request budget metric emission "
"exceeded %ss under load; skipping (values are refreshed by the "
"periodic budget-metrics cron job).",
budget_metrics_timeout,
)
return
for i, r in enumerate(results):
if isinstance(r, Exception):
verbose_logger.debug(

View file

@ -54,6 +54,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_strip_base64_files: bool = False,
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
s3_server_side_encryption: Optional[str] = None,
s3_callback_params_override: Optional[dict] = None,
**kwargs,
):
@ -92,6 +93,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_strip_base64_files=s3_strip_base64_files,
s3_use_key_prefix=s3_use_key_prefix,
s3_use_virtual_hosted_style=s3_use_virtual_hosted_style,
s3_server_side_encryption=s3_server_side_encryption,
)
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
@ -145,6 +147,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_strip_base64_files: bool = False,
s3_use_key_prefix: bool = False,
s3_use_virtual_hosted_style: bool = False,
s3_server_side_encryption: Optional[str] = None,
params_source: Optional[dict] = None,
):
"""
@ -194,6 +197,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style
)
self.s3_server_side_encryption = params.get("s3_server_side_encryption") or s3_server_side_encryption
return
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
@ -273,6 +278,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
try:
import base64
import hashlib
import requests
@ -317,14 +323,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Calculate SHA256 hash of the content
content_hash = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
content_md5 = base64.b64encode(
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
).decode()
# Prepare the request
headers = {
"Content-Type": "application/json",
"Content-MD5": content_md5,
"x-amz-content-sha256": content_hash,
"Content-Language": "en",
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**(
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
if self.s3_server_side_encryption
else {}
),
}
req = requests.Request("PUT", url, data=json_string, headers=headers)
prepped = req.prepare()
@ -447,6 +462,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
def upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
try:
import base64
import hashlib
import requests
@ -482,14 +498,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Calculate SHA256 hash of the content
content_hash = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
content_md5 = base64.b64encode(
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
).decode()
# Prepare the request
headers = {
"Content-Type": "application/json",
"Content-MD5": content_md5,
"x-amz-content-sha256": content_hash,
"Content-Language": "en",
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
**(
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
if self.s3_server_side_encryption
else {}
),
}
req = requests.Request("PUT", url, data=json_string, headers=headers)
prepped = req.prepare()

View file

@ -652,6 +652,10 @@ def _get_openai_compatible_provider_info(
api_base = api_base or get_secret("DEEPSEEK_API_BASE") or "https://api.deepseek.com/beta" # type: ignore
dynamic_api_key = api_key or get_secret_str("DEEPSEEK_API_KEY")
elif custom_llm_provider == "tencent":
api_base = api_base or get_secret("TENCENT_API_BASE") or "https://tokenhub-intl.tencentcloudmaas.com/v1"
dynamic_api_key = api_key or get_secret_str("TENCENT_API_KEY")
elif custom_llm_provider == "fireworks_ai":
# fireworks is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.fireworks.ai/inference/v1
(

View file

@ -106,6 +106,8 @@ def get_supported_openai_params(
return litellm.VLLMConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "deepseek":
return litellm.DeepSeekChatConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "tencent":
return litellm.TencentChatConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere":
return litellm.CohereChatConfig().get_supported_openai_params(model=model)
elif custom_llm_provider == "maritalk":

View file

@ -1,7 +1,23 @@
from typing import Dict, Optional
from typing import Any, Dict, Iterator, Optional
from litellm.types.utils import StandardCallbackDynamicParams
_CLIENT_CALLBACK_METADATA_SLOTS: tuple[str, ...] = ("litellm_metadata", "metadata")
def iter_client_callback_metadata_dicts(
kwargs: dict[str, Any],
) -> Iterator[tuple[str, dict[str, Any]]]:
litellm_params = kwargs.get("litellm_params")
if isinstance(litellm_params, dict):
nested = litellm_params.get("metadata")
if isinstance(nested, dict):
yield "litellm_params.metadata", nested
for key in _CLIENT_CALLBACK_METADATA_SLOTS:
candidate = kwargs.get(key)
if isinstance(candidate, dict):
yield key, candidate
def _is_env_reference(value: object) -> bool:
return isinstance(value, str) and "os.environ/" in value
@ -55,6 +71,7 @@ _supported_callback_params = [
"dd_site",
"dd_agent_host",
"dd_agent_port",
"turn_off_message_logging",
]
_request_blocked_callback_params = {
@ -87,19 +104,13 @@ def initialize_standard_callback_dynamic_params(
validate_no_callback_env_reference(param, _param_value, source="request body")
standard_callback_dynamic_params[param] = _param_value # type: ignore
# 2. Fallback: check "metadata" or "litellm_params" -> "metadata"
metadata = (kwargs.get("metadata") or {}).copy()
litellm_params = kwargs.get("litellm_params") or {}
if isinstance(litellm_params, dict):
metadata.update(litellm_params.get("metadata") or {})
if isinstance(metadata, dict):
for slot_label, metadata in iter_client_callback_metadata_dicts(kwargs):
for param in _supported_callback_params:
if param in _request_blocked_callback_params:
continue
if param not in standard_callback_dynamic_params and param in metadata:
_param_value = metadata.get(param)
validate_no_callback_env_reference(param, _param_value, source="metadata")
validate_no_callback_env_reference(param, _param_value, source=slot_label)
standard_callback_dynamic_params[param] = _param_value # type: ignore
return standard_callback_dynamic_params

View file

@ -10,7 +10,7 @@ import httpx
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, Message, ModelResponse
from litellm.types.utils import Choices, Message, ModelResponse, Usage
from ..common_utils import (
A2AError,
@ -312,6 +312,25 @@ class A2AConfig(BaseConfig):
# Set ID from response
model_response.id = response_json.get("id", str(uuid.uuid4()))
# A2A agents don't return token usage; estimate it so per-token pricing
# produces real cost and callers don't receive usage of 0/0/0.
try:
from litellm.utils import token_counter
prompt_tokens = token_counter(model="gpt-3.5-turbo", messages=messages)
completion_tokens = token_counter(model="gpt-3.5-turbo", text=text, count_response_tokens=True)
setattr(
model_response,
"usage",
Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
),
)
except Exception: # noqa: BLE001 - best-effort estimate; a tokenizer hiccup must not break the response
pass
return model_response
def get_model_response_iterator(

View file

@ -48,7 +48,10 @@ from litellm.types.utils import (
)
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
@ -70,6 +73,170 @@ class AnthropicMessagesHandler(BaseTranslation):
super().__init__()
self.adapter = LiteLLMAnthropicMessagesAdapter()
@staticmethod
def _build_streaming_usage_response(
responses_so_far: list[Any],
request_data: Optional[dict],
) -> Optional[ModelResponse]:
chunks = tuple(response for response in responses_so_far if isinstance(response, (str, bytes)))
if not chunks:
return None
try:
return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
all_chunks=chunks,
model=str((request_data or {}).get("model") or ""),
)
except (AttributeError, TypeError, ValueError):
return None
def build_block_sse_chunks(
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: Optional[list[Any]] = None,
) -> list[bytes]:
"""
Build an Anthropic SSE sequence delivering the guardrail block message
and terminating the stream cleanly.
- ``stream_started`` False (buffered / pre-stream): nothing has been
sent, so emit a complete standalone message (message_start ->
content_block_* -> message_delta -> message_stop) via
FakeAnthropicMessagesStreamIterator, the same converter the
/v1/messages pre-stream block handler uses.
- ``stream_started`` True (sampling / detect-only end-of-stream): real
chunks were already sent, so *continue* the in-progress message --
close the open content block, append the block message as a new text
block, then end the message. Emitting a second ``message_start`` here
would make Anthropic clients reject the stream.
"""
if stream_started:
return self._block_continuation_chunks(exc, responses_so_far or [])
return self._standalone_block_chunks(exc)
def _standalone_block_chunks(self, exc: "ModifyResponseException") -> list[bytes]:
import uuid
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_response_usage,
)
from litellm.types.utils import AnthropicMessagesResponse
block_response = AnthropicMessagesResponse(
id=f"msg_{uuid.uuid4()}",
type="message",
role="assistant",
content=[{"type": "text", "text": exc.message}],
model=exc.model,
stop_reason="end_turn",
usage=blocked_response_usage(getattr(exc, "original_response", None)),
)
return list(FakeAnthropicMessagesStreamIterator(response=block_response))
def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[Any]) -> list[bytes]:
"""Continue an already-started message: close the open content block,
append the block message as a new text block, then end the message --
without a second message_start."""
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_response_usage,
)
def _sse(event_type: str, payload: dict) -> bytes:
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
output_tokens = blocked_response_usage(getattr(exc, "original_response", None))["output_tokens"]
open_index, max_index = self._content_block_state(responses_so_far)
new_index = (max_index + 1) if max_index is not None else 0
chunks: list[bytes] = []
if open_index is not None:
chunks.append(_sse("content_block_stop", {"type": "content_block_stop", "index": open_index}))
chunks += [
_sse(
"content_block_start",
{
"type": "content_block_start",
"index": new_index,
"content_block": {"type": "text", "text": ""},
},
),
_sse(
"content_block_delta",
{
"type": "content_block_delta",
"index": new_index,
"delta": {"type": "text_delta", "text": exc.message},
},
),
_sse("content_block_stop", {"type": "content_block_stop", "index": new_index}),
_sse(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": output_tokens},
},
),
_sse("message_stop", {"type": "message_stop"}),
]
return chunks
@staticmethod
def _content_block_state(
responses_so_far: list[Any],
) -> tuple[Optional[int], Optional[int]]:
"""From the SSE chunks already sent to the client, return (open
content-block index or None, highest content-block index seen or None).
A single streamed item may bundle multiple SSE events (raw bytes) or be
an already-parsed event dict, so every event across every item is
considered -- matching how ``get_streaming_string_so_far`` reads the
same stream."""
open_indices: set[int] = set()
max_index: Optional[int] = None
for item in responses_so_far:
for data in AnthropicMessagesHandler._iter_sse_events(item):
event_type = data.get("type")
index = data.get("index")
if not isinstance(index, int):
continue
if event_type == "content_block_start":
open_indices.add(index)
max_index = index if max_index is None else max(max_index, index)
elif event_type == "content_block_stop":
open_indices.discard(index)
open_index = max(open_indices) if open_indices else None
return open_index, max_index
@staticmethod
def _iter_sse_events(item: Any) -> list[dict]:
"""Yield the event-data dicts in one stream chunk.
Handles both formats this stream can carry (see
``get_streaming_string_so_far``): raw SSE ``bytes`` -- which may bundle
several events separated by a blank line -- and an already-parsed event
``dict``."""
if isinstance(item, dict):
return [item]
if not isinstance(item, (bytes, bytearray)):
return []
events: list[dict] = []
for block in item.decode("utf-8", errors="replace").split("\n\n"):
for line in block.split("\n"):
line = line.strip()
if not line.startswith("data:"):
continue
try:
parsed = json.loads(line[len("data:") :].strip())
except json.JSONDecodeError:
continue
if isinstance(parsed, dict):
events.append(parsed)
return events
def _translate_to_openai(self, data: dict) -> ChatCompletionRequest:
"""Translate Anthropic request to OpenAI chat completion format."""
(
@ -406,6 +573,8 @@ class AnthropicMessagesHandler(BaseTranslation):
Get the string so far, check the apply guardrail to the string so far, and return the list of responses so far.
"""
from litellm.integrations.custom_guardrail import ModifyResponseException
has_ended = self._check_streaming_has_ended(responses_so_far)
if has_ended:
# build the model response from the responses_so_far
@ -430,25 +599,35 @@ class AnthropicMessagesHandler(BaseTranslation):
if tool_calls_list:
guardrail_inputs["tool_calls"] = tool_calls_list
_guardrailed_inputs = (
await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
try:
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=guardrail_inputs,
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
)
except ModifyResponseException as e:
if e.original_response is None:
e.original_response = built_response or self._build_streaming_usage_response(
responses_so_far, request_data
)
raise
else:
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
return responses_so_far
string_so_far = self.get_streaming_string_so_far(responses_so_far)
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
inputs={"texts": [string_so_far]},
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
try:
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [string_so_far]},
request_data=request_data if request_data is not None else {},
input_type="response",
logging_obj=litellm_logging_obj,
)
except ModifyResponseException as e:
if e.original_response is None:
e.original_response = self._build_streaming_usage_response(responses_so_far, request_data)
raise
return responses_so_far
def _prepare_request_data(

View file

@ -15,6 +15,7 @@ from typing import Any, Dict
from urllib.parse import quote
import httpx
from pydantic import BaseModel
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
@ -38,6 +39,30 @@ from litellm.secret_managers.main import get_secret_str
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
class AzureDocumentIntelligenceLine(BaseModel):
content: str | None = None
class AzureDocumentIntelligencePage(BaseModel):
pageNumber: int | None = None
width: float | None = None
height: float | None = None
unit: str | None = None
lines: tuple[AzureDocumentIntelligenceLine, ...] = ()
class AzureDocumentIntelligenceAnalyzeResult(BaseModel):
content: str | None = None
pages: tuple[AzureDocumentIntelligencePage, ...] = ()
tables: list[dict[str, object]] | None = None
keyValuePairs: list[dict[str, object]] | None = None
class AzureDocumentIntelligenceOperation(BaseModel):
status: str | None = None
analyzeResult: AzureDocumentIntelligenceAnalyzeResult | None = None
class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"""
Azure Document Intelligence OCR transformation configuration.
@ -67,11 +92,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
(1-based, e.g. "1-3,5,7-9"). To keep the public request shape
aligned with Mistral OCR, callers pass `pages` using Mistral
semantics — a list of 0-based integers — or a pre-formatted
Azure-style string. Other Mistral-specific params (e.g.
Azure-style string. Azure DI also exposes a `features` query
parameter enabling add-on capabilities (e.g. "keyValuePairs",
"languages"), passed as a list of feature names or a
comma-separated string. Other Mistral-specific params (e.g.
`include_image_base64`) are not supported by Azure DI and are
ignored during transformation.
"""
return ["pages"]
return ["pages", "features"]
def map_ocr_params(
self,
@ -85,16 +113,18 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
Translates Mistral-style `pages` (list[int], 0-based) into Azure's
`pages` query string (1-based, e.g. "1,2,3" or "1-3,5"). A raw
string that already matches Azure's format is passed through
unchanged.
unchanged. `features` (list[str] or comma-separated string) is
normalized into Azure's comma-joined `features` query string.
"""
pages = non_default_params.get("pages")
if pages is None:
return optional_params
normalized = self._normalize_pages_param(pages)
if normalized:
optional_params["pages"] = normalized
return optional_params
features = non_default_params.get("features")
normalized_pages = self._normalize_pages_param(pages) if pages is not None else ""
normalized_features = self._normalize_features_param(features) if features is not None else ""
return {
**optional_params,
**({"pages": normalized_pages} if normalized_pages else {}),
**({"features": normalized_features} if normalized_features else {}),
}
@staticmethod
def _normalize_pages_param(pages: Any) -> str:
@ -140,6 +170,39 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
raise ValueError("`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'.")
@staticmethod
def _normalize_features_param(features: object) -> str:
"""
Convert a caller-provided `features` value to Azure DI's query-string
form (comma-joined feature names, e.g. "keyValuePairs,languages").
Accepted inputs:
- list[str]: feature names like ["keyValuePairs", "languages"].
- str: a single feature name or comma-separated names.
"""
invalid_features_error = ValueError(
f"Invalid `features` for Azure Document Intelligence: {features!r}. "
f"Expected a list of feature names or a comma-separated string like "
f"'keyValuePairs' or 'keyValuePairs,languages'."
)
if isinstance(features, str):
raw_tokens = features.split(",")
elif isinstance(features, list):
if len(features) == 0:
return ""
raw_tokens = [feature for feature in features if isinstance(feature, str)]
if len(raw_tokens) != len(features):
raise invalid_features_error
else:
raise invalid_features_error
tokens = tuple(token.strip() for token in raw_tokens)
feature_pattern = re.compile(r"^[A-Za-z][A-Za-z0-9]*$")
if not all(feature_pattern.match(token) for token in tokens):
raise invalid_features_error
return ",".join(tokens)
def validate_environment(
self,
headers: Dict,
@ -228,13 +291,15 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
)
# Azure DI accepts `pages` as a query param (1-based, e.g. "1-3,5").
# Azure DI accepts `pages` (1-based, e.g. "1-3,5") and `features`
# (comma-joined names, e.g. "keyValuePairs") as query params.
# `optional_params` has already been normalized in `map_ocr_params`.
pages = optional_params.get("pages") if optional_params else None
if pages:
url += f"&pages={quote(str(pages), safe=',-')}"
features = optional_params.get("features") if optional_params else None
pages_query = f"&pages={quote(str(pages), safe=',-')}" if pages else ""
features_query = f"&features={quote(str(features), safe=',')}" if features else ""
return url
return f"{url}{pages_query}{features_query}"
def _extract_base64_from_data_uri(self, data_uri: str) -> str:
"""
@ -328,27 +393,15 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
return OCRRequestData(data=data, files=None)
def _extract_page_markdown(self, page_data: Dict[str, Any]) -> str:
"""
Extract text from Azure DI page and format as markdown.
Azure DI provides text in 'lines' array. We concatenate them with newlines.
Args:
page_data: Azure DI page object
Returns:
Markdown-formatted text
"""
lines = page_data.get("lines", [])
if not lines:
return ""
# Extract text content from each line
text_lines = [line.get("content", "") for line in lines]
# Join with newlines to preserve structure
return "\n".join(text_lines)
def _transform_azure_page(self, azure_page: AzureDocumentIntelligencePage) -> OCRPage:
page_number = azure_page.pageNumber if azure_page.pageNumber is not None else 1
markdown = "\n".join(line.content or "" for line in azure_page.lines)
dimensions = self._convert_dimensions(
width=azure_page.width if azure_page.width is not None else 8.5,
height=azure_page.height if azure_page.height is not None else 11,
unit=azure_page.unit if azure_page.unit is not None else "inch",
)
return OCRPage(index=page_number - 1, markdown=markdown, dimensions=dimensions)
def _convert_dimensions(self, width: float, height: float, unit: str) -> OCRPageDimensions:
"""
@ -526,6 +579,52 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
retry_after = self._get_retry_after(response=response)
await asyncio.sleep(retry_after)
def _get_polling_target(self, raw_response: httpx.Response) -> tuple[str, Dict[str, str]]:
operation_url = raw_response.headers.get("Operation-Location")
if not operation_url:
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
# Reject cross-origin polling URLs — the auth headers
# below would otherwise leak to whatever URL the upstream
# (or an attacker-controlled upstream) returns. VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
poll_headers = {"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")}
return operation_url, poll_headers
def _transform_completed_response(self, model: str, raw_response: httpx.Response) -> OCRResponse:
"""
Transform a completed Azure Document Intelligence analyze operation
into the Mistral OCR response shape, preserving Azure-native
`analyzeResult` fields (`content`, `tables`, `keyValuePairs`) as
top-level response fields.
"""
operation = AzureDocumentIntelligenceOperation.model_validate(raw_response.json())
verbose_logger.debug(f"Azure Document Intelligence response status: {operation.status}")
if operation.status != "succeeded":
raise ValueError(f"Azure Document Intelligence analysis failed with status: {operation.status}")
analyze_result = (
operation.analyzeResult if operation.analyzeResult is not None else AzureDocumentIntelligenceAnalyzeResult()
)
mistral_pages = [self._transform_azure_page(azure_page) for azure_page in analyze_result.pages]
usage_info = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
return OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
object="ocr",
content=analyze_result.content,
tables=analyze_result.tables,
keyValuePairs=analyze_result.keyValuePairs,
)
def transform_ocr_response(
self,
model: str,
@ -552,11 +651,13 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"unit": "inch",
"lines": [{"content": "text", "boundingBox": [...]}]
}
]
],
"tables": [...],
"keyValuePairs": [...]
}
}
Mistral OCR format:
Mistral OCR format (with Azure-native fields preserved):
{
"pages": [
{
@ -567,7 +668,10 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
],
"model": "azure_ai/doc-intelligence/prebuilt-layout",
"usage_info": {"pages_processed": 1},
"object": "ocr"
"object": "ocr",
"content": "Full document text...",
"tables": [...],
"keyValuePairs": [...]
}
Args:
@ -578,86 +682,17 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
Returns:
OCRResponse in Mistral format
"""
try:
# Check if we got 202 Accepted (async operation started)
if raw_response.status_code == 202:
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
if raw_response.status_code != 202:
return self._transform_completed_response(model=model, raw_response=raw_response)
# Get Operation-Location header
operation_url = raw_response.headers.get("Operation-Location")
if not operation_url:
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
# Reject cross-origin polling URLs — the auth headers
# below would otherwise leak to whatever URL the upstream
# (or an attacker-controlled upstream) returns. VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
# Get headers for polling (need auth)
poll_headers = {
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")
}
# Get timeout from kwargs or use default
timeout_secs = AZURE_OPERATION_POLLING_TIMEOUT
# Poll until operation completes
raw_response = self._poll_operation_sync(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=timeout_secs,
)
# Now parse the completed response
response_json = raw_response.json()
verbose_logger.debug(f"Azure Document Intelligence response status: {response_json.get('status')}")
# Check if request succeeded
status = response_json.get("status")
if status != "succeeded":
raise ValueError(f"Azure Document Intelligence analysis failed with status: {status}")
# Extract analyze result
analyze_result = response_json.get("analyzeResult", {})
azure_pages = analyze_result.get("pages", [])
# Transform pages to Mistral format
mistral_pages = []
for azure_page in azure_pages:
page_number = azure_page.get("pageNumber", 1)
index = page_number - 1 # Convert to 0-based index
# Extract markdown text
markdown = self._extract_page_markdown(azure_page)
# Convert dimensions
width = azure_page.get("width", 8.5)
height = azure_page.get("height", 11)
unit = azure_page.get("unit", "inch")
dimensions = self._convert_dimensions(width=width, height=height, unit=unit)
# Build OCR page
ocr_page = OCRPage(index=index, markdown=markdown, dimensions=dimensions)
mistral_pages.append(ocr_page)
# Build usage info
usage_info = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
# Return Mistral OCR response
return OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
object="ocr",
)
except Exception as e:
verbose_logger.error(f"Error parsing Azure Document Intelligence response: {e}")
raise e
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
operation_url, poll_headers = self._get_polling_target(raw_response)
completed_response = self._poll_operation_sync(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
)
return self._transform_completed_response(model=model, raw_response=completed_response)
async def async_transform_ocr_response(
self,
@ -680,81 +715,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
Returns:
OCRResponse in Mistral format
"""
try:
# Check if we got 202 Accepted (async operation started)
if raw_response.status_code == 202:
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
if raw_response.status_code != 202:
return self._transform_completed_response(model=model, raw_response=raw_response)
# Get Operation-Location header
operation_url = raw_response.headers.get("Operation-Location")
if not operation_url:
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
# Reject cross-origin polling URLs (see sync path). VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
# Get headers for polling (need auth)
poll_headers = {
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")
}
# Get timeout from kwargs or use default
timeout_secs = AZURE_OPERATION_POLLING_TIMEOUT
# Poll until operation completes (async)
raw_response = await self._poll_operation_async(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=timeout_secs,
)
# Now parse the completed response
response_json = raw_response.json()
verbose_logger.debug(f"Azure Document Intelligence response status: {response_json.get('status')}")
# Check if request succeeded
status = response_json.get("status")
if status != "succeeded":
raise ValueError(f"Azure Document Intelligence analysis failed with status: {status}")
# Extract analyze result
analyze_result = response_json.get("analyzeResult", {})
azure_pages = analyze_result.get("pages", [])
# Transform pages to Mistral format
mistral_pages = []
for azure_page in azure_pages:
page_number = azure_page.get("pageNumber", 1)
index = page_number - 1 # Convert to 0-based index
# Extract markdown text
markdown = self._extract_page_markdown(azure_page)
# Convert dimensions
width = azure_page.get("width", 8.5)
height = azure_page.get("height", 11)
unit = azure_page.get("unit", "inch")
dimensions = self._convert_dimensions(width=width, height=height, unit=unit)
# Build OCR page
ocr_page = OCRPage(index=index, markdown=markdown, dimensions=dimensions)
mistral_pages.append(ocr_page)
# Build usage info
usage_info = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
# Return Mistral OCR response
return OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
object="ocr",
)
except Exception as e:
verbose_logger.error(f"Error parsing Azure Document Intelligence response (async): {e}")
raise e
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
operation_url, poll_headers = self._get_polling_target(raw_response)
completed_response = await self._poll_operation_async(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
)
return self._transform_completed_response(model=model, raw_response=completed_response)

View file

@ -2,7 +2,10 @@ from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Dict, List, Optional
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
ModifyResponseException,
)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import AllMessageValues
@ -98,6 +101,30 @@ class BaseTranslation(ABC):
"""
return responses_so_far
def build_block_sse_chunks(
self,
exc: "ModifyResponseException",
stream_started: bool = False,
responses_so_far: Optional[list[Any]] = None,
) -> Optional[list[bytes]]:
"""
Build the streaming chunks that deliver a guardrail block message and
cleanly terminate the stream in this provider's wire format.
``stream_started`` is True when real chunks were already sent to the
client: the result must *continue* the in-progress message (e.g. close
the open content block and append the block message) rather than start
a new one, which clients reject. ``responses_so_far`` provides the prior
chunks needed to do so. When False, nothing has been sent and a
standalone block message is emitted.
Returns None when the format has no safe terminator; the caller then
re-raises ``exc`` so the proxy can surface a clean error instead.
Override in provider subclasses that support synthesizing a block
stream.
"""
return None
def get_structured_messages(self, data: dict) -> Optional[List["AllMessageValues"]]:
"""
Convert request data to OpenAI-spec structured messages.

View file

@ -1,10 +1,100 @@
from __future__ import annotations
from typing import Any, List
import json
from typing import Any, List, Optional
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
from litellm.types.llms.openai import AllMessageValues
def _anthropic_stream_chunk_events(item: Any) -> list[dict]:
if isinstance(item, dict):
return [item]
if isinstance(item, bytes):
chunk = item.decode("utf-8", errors="replace")
elif isinstance(item, str):
chunk = item
else:
return []
events: list[dict] = []
for block in chunk.split("\n\n"):
for line in block.splitlines():
stripped = line.strip()
if not stripped.startswith("data:"):
continue
payload = stripped[len("data:") :].strip()
if not payload or payload == "[DONE]":
continue
try:
parsed = json.loads(payload)
except json.JSONDecodeError:
continue
if isinstance(parsed, dict):
events.append(parsed)
return events
def _usage_from_anthropic_stream_chunks(original_response: list[Any]) -> Optional[AnthropicUsage]:
input_tokens = 0
output_tokens = 0
found_usage = False
for item in original_response:
for event in _anthropic_stream_chunk_events(item):
event_type = event.get("type")
if event_type == "message_start":
message = event.get("message") or {}
usage_obj = message.get("usage") or {}
elif event_type == "message_delta":
usage_obj = event.get("usage") or {}
else:
usage_obj = {}
if not isinstance(usage_obj, dict):
continue
if usage_obj.get("input_tokens") is not None:
input_tokens = int(usage_obj.get("input_tokens") or 0)
found_usage = True
if usage_obj.get("output_tokens") is not None:
output_tokens = int(usage_obj.get("output_tokens") or 0)
found_usage = True
if not found_usage:
return None
return AnthropicUsage(input_tokens=input_tokens, output_tokens=output_tokens)
def blocked_response_usage(original_response: Optional[Any]) -> AnthropicUsage:
"""
Token usage for a synthetic guardrail-blocked response.
A post-call block replaces the LLM's response with the violation message,
but the upstream call already consumed tokens -- report that real usage
(carried on ``ModifyResponseException.original_response``) rather than
discarding it. Pre-call blocks never invoked the LLM (no original_response),
so usage is zero.
"""
usage_obj: Any = None
if isinstance(original_response, list):
stream_usage = _usage_from_anthropic_stream_chunks(original_response)
if stream_usage is not None:
return stream_usage
elif isinstance(original_response, dict):
usage_obj = original_response.get("usage")
elif original_response is not None:
usage_obj = getattr(original_response, "usage", None)
def _tokens(key: str, fallback_key: str) -> int:
if isinstance(usage_obj, dict):
return int(usage_obj.get(key, usage_obj.get(fallback_key, 0)) or 0)
return int(getattr(usage_obj, key, getattr(usage_obj, fallback_key, 0)) or 0)
return AnthropicUsage(
input_tokens=_tokens("input_tokens", "prompt_tokens"),
output_tokens=_tokens("output_tokens", "completion_tokens"),
)
def effective_skip_system_message_for_guardrail(guardrail_to_apply: Any) -> bool:
per = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None)
if per is not None:

View file

@ -70,6 +70,9 @@ class OCRResponse(LiteLLMPydanticObjectBase):
model: str
document_annotation: Any | None = None
usage_info: OCRUsageInfo | None = None
content: str | None = None
tables: list[dict[str, object]] | None = None
keyValuePairs: list[dict[str, object]] | None = None
object: str = "ocr"
model_config = {"extra": "allow"}

View file

@ -76,6 +76,7 @@ from litellm.utils import (
from ..common_utils import (
BedrockError,
BedrockModelInfo,
bedrock_converse_supports_parallel_tool_use_config,
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
is_claude_4_5_on_bedrock,
@ -1106,18 +1107,28 @@ class AmazonConverseConfig(BaseConfig):
if cache_control is None:
return None
cache_point = CachePointBlock(type="default")
if isinstance(cache_control, dict) and "ttl" in cache_control:
ttl = cache_control["ttl"]
if ttl in ["5m", "1h"] and model is not None:
if is_claude_4_5_on_bedrock(model):
cache_point["ttl"] = ttl
cache_point = self._build_cache_point_block(cache_control, model)
if block_type == "system":
return SystemContentBlock(cachePoint=cache_point)
else:
return ContentBlock(cachePoint=cache_point)
@staticmethod
def _build_cache_point_block(control: Optional[dict], model: Optional[str] = None) -> CachePointBlock:
"""Build a Bedrock ``cachePoint`` block from an OpenAI-style ``cache_control``/``control`` dict.
``type`` is always ``"default"`` (the only value Bedrock's Converse API
accepts). ``ttl`` is only honored for models that support extended TTL
caching (Claude 4.5 family on Bedrock).
"""
cache_point = CachePointBlock(type="default")
if isinstance(control, dict) and "ttl" in control:
ttl = control["ttl"]
if ttl in ["5m", "1h"] and model is not None and is_claude_4_5_on_bedrock(model):
cache_point["ttl"] = ttl
return cache_point
def _transform_system_message(
self, messages: List[AllMessageValues], model: Optional[str] = None
) -> Tuple[List[AllMessageValues], List[SystemContentBlock]]:
@ -1241,7 +1252,7 @@ class AmazonConverseConfig(BaseConfig):
# Handle parallel_tool_calls configuration
parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None)
if parallel_tool_use_config is not None and is_claude_4_5_on_bedrock(model):
if parallel_tool_use_config is not None and bedrock_converse_supports_parallel_tool_use_config(model):
for key, value in parallel_tool_use_config.items():
if (
key in additional_request_params
@ -1526,7 +1537,8 @@ class AmazonConverseConfig(BaseConfig):
if cache_injection_points and len(bedrock_tools) > 0:
for point in cache_injection_points:
if point.get("location") == "tool_config":
bedrock_tools.append({"cachePoint": {"type": "default"}})
cache_point = self._build_cache_point_block(point.get("control"), model)
bedrock_tools.append(ToolBlock(cachePoint=cache_point))
break
bedrock_tool_config: Optional[ToolConfigBlock] = None

View file

@ -5,6 +5,7 @@ from functools import partial
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast, get_args
import httpx
from pydantic import TypeAdapter, ValidationError
import litellm
from litellm._logging import verbose_logger
@ -24,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import (
HTTPHandler,
_get_httpx_client,
)
from litellm.types.llms.bedrock import GuardrailConfigBlock
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import ModelResponse, Usage
from litellm.utils import CustomStreamWrapper
@ -37,6 +39,38 @@ else:
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
_GUARDRAIL_CONFIG_VALIDATOR: "TypeAdapter[GuardrailConfigBlock]" = TypeAdapter(GuardrailConfigBlock)
_GUARDRAIL_CONFIG_EXPECTED_FORMAT = (
"{'guardrailIdentifier': str, 'guardrailVersion': str, 'trace': 'enabled'|'disabled'|'enabled_full'}"
)
def _bedrock_invoke_guardrail_headers(raw_guardrail_config: object) -> "dict[str, str]":
try:
guardrail_config = _GUARDRAIL_CONFIG_VALIDATOR.validate_python(raw_guardrail_config)
except ValidationError as e:
raise BedrockError(
status_code=400,
message="Invalid guardrailConfig={}. Expected format: {}. Error: {}".format(
raw_guardrail_config, _GUARDRAIL_CONFIG_EXPECTED_FORMAT, e
),
)
if "guardrailIdentifier" not in guardrail_config:
raise BedrockError(
status_code=400,
message="guardrailConfig={} is missing 'guardrailIdentifier'. Expected format: {}".format(
raw_guardrail_config, _GUARDRAIL_CONFIG_EXPECTED_FORMAT
),
)
trace = guardrail_config.get("trace")
candidate_headers = {
"X-Amzn-Bedrock-GuardrailIdentifier": guardrail_config.get("guardrailIdentifier"),
"X-Amzn-Bedrock-GuardrailVersion": guardrail_config.get("guardrailVersion"),
"X-Amzn-Bedrock-Trace": trace.upper() if trace is not None else None,
}
return {name: value for name, value in candidate_headers.items() if value is not None}
class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
def __init__(self, **kwargs):
@ -390,7 +424,16 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
return headers
raw_guardrail_config = optional_params.pop("guardrailConfig", None)
if raw_guardrail_config is None:
return headers
existing_header_names = frozenset(name.lower() for name in headers)
guardrail_headers = {
name: value
for name, value in _bedrock_invoke_guardrail_headers(raw_guardrail_config).items()
if name.lower() not in existing_header_names
}
return {**headers, **guardrail_headers}
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]

View file

@ -685,39 +685,27 @@ def get_bedrock_base_model(model: str) -> str:
return model
def bedrock_converse_supports_parallel_tool_use_config(model: str) -> bool:
return any(
(litellm.model_cost.get(candidate) or {}).get("supports_parallel_tool_use_config") is True
for candidate in (model, get_bedrock_base_model(model))
)
def is_claude_4_5_on_bedrock(model: str) -> bool:
"""
Check if the model is a Claude 4.5 model on Bedrock.
Claude 4.5 models support prompt caching with '5m' and '1h' TTL on Bedrock.
Check if the model supports Bedrock prompt caching with an extended '1h' TTL
(in addition to the default 5m TTL).
Backed by the ``cache_creation_input_token_cost_above_1hr`` field in
``model_prices_and_context_window.json`` instead of a hardcoded list of
model-name patterns, so newly released models pick up support as soon as
their pricing entry ships, with no code change required here.
"""
model_lower = model.lower()
claude_4_5_patterns = [
"sonnet-4.5",
"sonnet_4.5",
"sonnet-4-5",
"sonnet_4_5",
"haiku-4.5",
"haiku_4.5",
"haiku-4-5",
"haiku_4_5",
"opus-4.5",
"opus_4.5",
"opus-4-5",
"opus_4_5",
"sonnet-4.6",
"sonnet_4.6",
"sonnet-4-6",
"sonnet_4_6",
"opus-4.6",
"opus_4.6",
"opus-4-6",
"opus_4_6",
"opus-4.7",
"opus_4.7",
"opus-4-7",
"opus_4_7",
]
return any(pattern in model_lower for pattern in claude_4_5_patterns)
return any(
(litellm.model_cost.get(candidate) or {}).get("cache_creation_input_token_cost_above_1hr") is not None
for candidate in (model, get_bedrock_base_model(model))
)
_BEDROCK_MODEL_VERSION_SUFFIX_RE = re.compile(r"-v\d+(?::\d+)?$")

View file

@ -45,6 +45,7 @@ from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolCallChunk,
ChatCompletionToolParam,
ResponsesAPIStreamEvents,
)
from litellm.types.responses.main import (
GenericResponseOutputItem,
@ -586,7 +587,14 @@ class OpenAIResponsesHandler(BaseTranslation):
"""
Check if the streaming has ended.
"""
return all(response.choices[0].finish_reason is not None for response in responses_so_far)
if not responses_so_far:
return False
terminal_types = {
ResponsesAPIStreamEvents.RESPONSE_COMPLETED.value,
ResponsesAPIStreamEvents.RESPONSE_FAILED.value,
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value,
}
return responses_so_far[-1].get("type") in terminal_types
def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str:
"""

View file

View file

View file

@ -0,0 +1,68 @@
"""
Translates from OpenAI's `/v1/chat/completions` to Tencent TokenHub's
OpenAI-compatible endpoint.
"""
from typing import Optional
from litellm.secret_managers.main import get_secret_str
from litellm.utils import supports_reasoning
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
class TencentChatConfig(OpenAIGPTConfig):
def get_supported_openai_params(self, model: str) -> list:
params = super().get_supported_openai_params(model)
if supports_reasoning(model, custom_llm_provider="tencent"):
params.extend(["thinking", "reasoning_effort"])
return params
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
optional_params = super().map_openai_params(non_default_params, optional_params, model, drop_params)
thinking_value = optional_params.pop("thinking", None)
reasoning_effort = optional_params.pop("reasoning_effort", None)
if thinking_value is not None:
if isinstance(thinking_value, dict):
optional_params["thinking"] = thinking_value
elif reasoning_effort is not None and reasoning_effort != "none":
optional_params["thinking"] = {"type": "enabled"}
return optional_params
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> tuple[Optional[str], Optional[str]]:
api_base = api_base or get_secret_str("TENCENT_API_BASE") or "https://tokenhub-intl.tencentcloudmaas.com/v1"
dynamic_api_key = api_key or get_secret_str("TENCENT_API_KEY")
return api_base, dynamic_api_key
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if not api_base:
api_base = "https://tokenhub-intl.tencentcloudmaas.com/v1"
api_base = api_base.rstrip("/")
if api_base.endswith("/chat/completions"):
return api_base
if not api_base.endswith("/v1"):
api_base = f"{api_base}/v1"
return f"{api_base}/chat/completions"

View file

@ -0,0 +1,6 @@
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.types.utils import Usage
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
return generic_cost_per_token(model=model, usage=usage, custom_llm_provider="tencent")

View file

@ -0,0 +1,85 @@
"""
Tencent Anthropic-compatible messages transformation config.
Tencent TokenHub exposes an Anthropic-compatible Messages API endpoint
alongside its standard OpenAI-compatible chat completions endpoint.
"""
from typing import Any, Optional
import litellm
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
from litellm.secret_managers.main import get_secret_str
class TencentAnthropicMessagesConfig(AnthropicMessagesConfig):
"""
Tencent TokenHub exposes an Anthropic-compatible Messages API.
Unlike the chat completions endpoint (which uses /v1), the Anthropic
endpoint may use a different base URL. Configure via
TENCENT_ANTHROPIC_API_BASE or TENCENT_API_BASE.
"""
@property
def custom_llm_provider(self) -> Optional[str]:
return "tencent"
def should_strip_billing_metadata(self) -> bool:
return True
@staticmethod
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
return api_key or get_secret_str("TENCENT_API_KEY") or litellm.api_key
@staticmethod
def get_api_base(api_base: Optional[str] = None) -> str:
return (
api_base
or get_secret_str("TENCENT_ANTHROPIC_API_BASE")
or get_secret_str("TENCENT_API_BASE")
or "https://tokenhub-intl.tencentcloudmaas.com"
)
def validate_anthropic_messages_environment(
self,
headers: dict,
model: str,
messages: list[Any],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> tuple[dict, Optional[str]]:
return super().validate_anthropic_messages_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=self.get_api_key(api_key=api_key),
api_base=api_base,
)
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
base_url = self.get_api_base(api_base=api_base).rstrip("/")
if base_url.endswith("/v1/messages"):
return base_url
if base_url.endswith("/v1/chat/completions"):
base_url = base_url[: -len("/v1/chat/completions")]
elif base_url.endswith("/v1"):
base_url = base_url[: -len("/v1")]
return f"{base_url}/v1/messages"

View file

@ -1,5 +1,5 @@
import os
from typing import TYPE_CHECKING, Any, Dict, List, Optional
from typing import TYPE_CHECKING, Any, Optional
from litellm._logging import verbose_logger
@ -175,7 +175,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
messages: list[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
@ -217,10 +217,10 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
contents = [{"role": "user", "parts": [{"text": prompt}]}]
# Prepare generation config
generation_config: Dict[str, Any] = {"responseModalities": ["IMAGE"]}
generation_config: dict[str, Any] = {"responseModalities": ["IMAGE"]}
# Seed from user-supplied imageConfig dict; flat params are overlaid for backward compat.
image_config: Dict[str, Any] = dict(optional_params.get("imageConfig") or {})
image_config: dict[str, Any] = dict(optional_params.get("imageConfig") or {})
if "aspectRatio" in optional_params:
image_config["aspectRatio"] = optional_params["aspectRatio"]
@ -241,7 +241,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
elif "n" in optional_params:
generation_config["candidateCount"] = optional_params["n"]
request_body: Dict[str, Any] = {
request_body: dict[str, Any] = {
"contents": contents,
"generationConfig": generation_config,
}

File diff suppressed because it is too large Load diff

View file

@ -93,6 +93,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
has_user_credential: Optional[bool] = None
source_url: Optional[str] = None
timeout: Optional[float] = None
max_concurrent_requests: Optional[int] = None
approval_status: Optional[str] = Field(
default="active",
description="Approval status: 'pending_review', 'active', 'rejected'",

File diff suppressed because it is too large Load diff

View file

@ -1261,6 +1261,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
byok_api_key_help_url: Optional[str] = None
source_url: Optional[str] = None
timeout: Optional[float] = None
max_concurrent_requests: Optional[int] = None
# BYOM submission fields — set by the endpoint, not by the caller.
# Any caller-provided values are silently overridden before persistence.
approval_status: Optional[str] = Field(
@ -1346,6 +1347,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
byok_api_key_help_url: Optional[str] = None
source_url: Optional[str] = None
timeout: Optional[float] = None
max_concurrent_requests: Optional[int] = None
@model_validator(mode="before")
@classmethod

View file

@ -12,6 +12,9 @@ from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.anthropic.experimental_pass_through.context_management import (
AnthropicContextManagementError,
)
from litellm.llms.base_llm.guardrail_translation.utils import (
blocked_response_usage as _blocked_response_usage,
)
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_request_processing import (
@ -134,6 +137,10 @@ async def anthropic_response(
from litellm.types.utils import AnthropicMessagesResponse
# Report the blocked LLM response's real token usage (carried on the
# exception) instead of discarding it; zero for pre-call blocks.
_usage = _blocked_response_usage(e.original_response)
_anthropic_response = AnthropicMessagesResponse(
id=f"msg_{str(uuid.uuid4())}",
type="message",
@ -141,7 +148,7 @@ async def anthropic_response(
content=[{"type": "text", "text": e.message}],
model=e.model,
stop_reason="end_turn",
usage={"input_tokens": 0, "output_tokens": 0},
usage=_usage,
)
if data.get("stream", None) is not None and data["stream"] is True:

View file

@ -3017,7 +3017,9 @@ async def can_key_call_resolved_model(
)
skip_key_model_check = valid_token.config or (
isinstance(valid_token.models, list) and SpecialModelNames.all_team_models.value in valid_token.models
isinstance(valid_token.models, list)
and SpecialModelNames.all_team_models.value in valid_token.models
and valid_token.team_id is not None
)
if not skip_key_model_check:
await can_key_call_model(

View file

@ -375,6 +375,16 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
metadata = _coerce_metadata_to_dict(request_body.get(metadata_key))
if metadata is not None:
_check_banned_params(metadata, general_settings, llm_router, model)
litellm_params = _coerce_metadata_to_dict(request_body.get("litellm_params"))
if litellm_params is not None:
litellm_params_metadata = _coerce_metadata_to_dict(litellm_params.get("metadata"))
if litellm_params_metadata is not None:
_check_banned_params(
litellm_params_metadata,
general_settings,
llm_router,
model,
)
return True

View file

@ -2676,7 +2676,11 @@ async def _enforce_key_and_fallback_model_access(
model_list = config.get("model_list", [])
new_model_list = model_list
verbose_proxy_logger.debug(f"\n new llm router model list {new_model_list}")
elif isinstance(valid_token.models, list) and "all-team-models" in valid_token.models:
elif (
isinstance(valid_token.models, list)
and "all-team-models" in valid_token.models
and valid_token.team_id is not None
):
pass
else:
model = _get_model_from_request_context(

View file

@ -8,7 +8,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
import copy
import json
from typing import Any, AsyncGenerator, List, Optional, Union
from typing import TYPE_CHECKING, Any, AsyncGenerator, List, Optional, Union
from fastapi import HTTPException
@ -23,6 +23,11 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import CallTypes, CallTypesLiteral
if TYPE_CHECKING:
# Imported lazily at runtime (inside the streaming hook) to avoid a
# module-level cyclic import with litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
# Call types that use NDJSON streaming (A2A); guardrail HTTPException is emitted as in-stream error
A2A_CALL_TYPES = (CallTypes.asend_message, CallTypes.send_message)
@ -197,6 +202,10 @@ class UnifiedLLMGuardrails(CustomLogger):
)
from litellm.types.guardrails import GuardrailEventHooks
# Local import avoids a module-level cyclic import with
# litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
guardrail_to_apply: CustomGuardrail = data.pop("guardrail_to_apply", None)
if guardrail_to_apply is None:
@ -238,18 +247,51 @@ class UnifiedLLMGuardrails(CustomLogger):
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
response = await endpoint_translation.process_output_response(
response=response, # type: ignore
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=data,
)
try:
response = await endpoint_translation.process_output_response(
response=response, # type: ignore
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=data.get("litellm_logging_obj"),
user_api_key_dict=user_api_key_dict,
request_data=data,
)
except ModifyResponseException as e:
# The guardrail blocked the response. Attach the original LLM
# response so the endpoint handler can report its real token usage
# instead of discarding it (the block replaces the content, but the
# upstream call already consumed those tokens).
if e.original_response is None:
e.original_response = response
raise
# Add guardrail to applied guardrails header
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name)
return response
async def _handle_streaming_block(
self,
exc: "ModifyResponseException",
endpoint_translation: Any,
stream_started: bool,
responses_so_far: list[Any],
) -> AsyncGenerator[Any, None]:
"""
Terminate a streamed response cleanly when a guardrail blocks it.
Format-agnostic routing: delegates to the provider translation handler's
``build_block_sse_chunks`` (see ``BaseTranslation.build_block_sse_chunks``
for the ``stream_started`` / ``responses_so_far`` contract). When the
format has no safe terminator the handler returns None and we re-raise
``exc`` so the proxy can surface a clean error.
"""
block_chunks = endpoint_translation.build_block_sse_chunks(
exc, stream_started=stream_started, responses_so_far=responses_so_far
)
if block_chunks is None:
raise exc
for chunk in block_chunks:
yield chunk
async def async_post_call_streaming_iterator_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
@ -271,26 +313,53 @@ class UnifiedLLMGuardrails(CustomLogger):
global endpoint_guardrail_translation_mappings
# Local import avoids a module-level cyclic import with
# litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
guardrail_to_apply: CustomGuardrail = request_data.pop("guardrail_to_apply", None)
# Get streaming configuration from guardrail or optional_params
sampling_rate = 5
end_of_stream_only = False # If True, only apply guardrail at end of stream
# Get streaming configuration. Resolution order (later wins): default
# < guardrail attribute < guardrail_config dict < this callback's
# optional_params.
def _streaming_flag(name: str, default: Any) -> Any:
value = default
if guardrail_to_apply is not None:
value = getattr(guardrail_to_apply, name, value)
config = getattr(guardrail_to_apply, "guardrail_config", {})
if isinstance(config, dict):
value = config.get(name, value)
return self.optional_params.get(name, value)
if guardrail_to_apply is not None:
# Check direct attributes on guardrail first
sampling_rate = getattr(guardrail_to_apply, "streaming_sampling_rate", sampling_rate)
end_of_stream_only = getattr(guardrail_to_apply, "streaming_end_of_stream_only", end_of_stream_only)
sampling_rate = _streaming_flag("streaming_sampling_rate", 5)
# Only apply the guardrail at end of stream (not per chunk).
end_of_stream_only = _streaming_flag("streaming_end_of_stream_only", False)
# Withhold every chunk until end-of-stream moderation passes, then
# release the original chunks (clean) or only the block message
# (blocked) -- moderating the whole response *before* any content
# reaches the client. Only safe for allow/block guardrails: on
# release the original chunks are replayed as-is, so a
# content-rewriting guardrail (e.g. PII masking) would leak
# unredacted content. Guarded below via mask_response_content.
buffer_until_moderated = _streaming_flag("streaming_buffer_until_moderated", False)
# Also check guardrail_config dict if present
guardrail_config = getattr(guardrail_to_apply, "guardrail_config", {})
if isinstance(guardrail_config, dict):
sampling_rate = guardrail_config.get("streaming_sampling_rate", sampling_rate)
end_of_stream_only = guardrail_config.get("streaming_end_of_stream_only", end_of_stream_only)
if (
buffer_until_moderated
and guardrail_to_apply is not None
and getattr(guardrail_to_apply, "mask_response_content", False)
):
verbose_proxy_logger.warning(
"UnifiedLLMGuardrails: streaming_buffer_until_moderated is disabled for %s "
"because mask_response_content=True -- buffered replay would release "
"unredacted original chunks instead of the moderated output.",
guardrail_to_apply.guardrail_name,
)
buffer_until_moderated = False
# Also check optional_params as fallback
sampling_rate = self.optional_params.get("streaming_sampling_rate", sampling_rate)
end_of_stream_only = self.optional_params.get("streaming_end_of_stream_only", end_of_stream_only)
# Buffering can only moderate the assembled response, so it always
# defers to end-of-stream.
if buffer_until_moderated:
end_of_stream_only = True
if guardrail_to_apply is None:
async for item in response:
@ -315,6 +384,12 @@ class UnifiedLLMGuardrails(CustomLogger):
call_type = None
chunk_counter = 0
responses_so_far: List[Any] = []
responses_yielded: list[Any] = []
pending_end_of_stream_items: list[Any] = []
# Whether any real response chunk has been forwarded to the client.
# Drives how a block terminates the stream: continue the in-progress
# message (True) vs emit a standalone block message (False, buffered).
chunks_yielded = False
async for item in response:
chunk_counter += 1
@ -336,9 +411,22 @@ class UnifiedLLMGuardrails(CustomLogger):
yield remaining_item
return
# If end_of_stream_only mode, yield chunks without processing
# If end_of_stream_only mode, yield chunks without processing.
# When buffering, withhold them instead -- they are released (or
# replaced by the block message) only after end-of-stream
# moderation runs below.
if end_of_stream_only:
yield item
if not buffer_until_moderated:
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
stream_has_ended = hasattr(
endpoint_translation, "_check_streaming_has_ended"
) and endpoint_translation._check_streaming_has_ended(responses_so_far)
if pending_end_of_stream_items or stream_has_ended:
pending_end_of_stream_items.append(item)
else:
chunks_yielded = True
responses_yielded.append(item)
yield item
continue
# Process chunk based on sampling rate
@ -368,6 +456,26 @@ class UnifiedLLMGuardrails(CustomLogger):
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
except ModifyResponseException as e:
if e.original_response is None:
e.original_response = responses_so_far
# Guardrail blocked the response mid-stream. Emit a clean
# terminating SSE sequence delivering the block message
# instead of letting the exception propagate into a bare
# `data: {"error": ...}` blob (which truncates the stream).
# Chunks have already been forwarded here, so the block
# continues the in-progress message (stream_started=True).
# The current chunk was appended to responses_so_far but not
# yet yielded, so exclude it: the continuation must reflect
# only what the client has actually received.
async for block_chunk in self._handle_streaming_block(
e,
endpoint_translation,
stream_started=chunks_yielded,
responses_so_far=responses_yielded,
):
yield block_chunk
return
except HTTPException as e:
# Response already started (we already yielded chunks); cannot send 400.
# For A2A (NDJSON), yield an in-stream JSON-RPC error so the client sees it.
@ -394,8 +502,12 @@ class UnifiedLLMGuardrails(CustomLogger):
yield error_chunk
return
raise
chunks_yielded = True
responses_yielded.append(original_item)
yield original_item
else:
chunks_yielded = True
responses_yielded.append(item)
yield item
# Stream has ended - do final processing with all collected chunks
@ -408,6 +520,15 @@ class UnifiedLLMGuardrails(CustomLogger):
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
# When buffering, snapshot the original chunks before moderation.
# A shallow copy suffices: end-of-stream
# process_output_streaming_response builds a separate assembled
# response (it does not mutate the individual chunks in place), and
# the chunks themselves are replayed verbatim -- so we only need to
# preserve the list, not clone every chunk (deepcopy would double
# peak memory for large responses).
buffered_items = list(responses_so_far) if buffer_until_moderated else None
try:
await endpoint_translation.process_output_streaming_response(
responses_so_far=responses_so_far,
@ -416,6 +537,28 @@ class UnifiedLLMGuardrails(CustomLogger):
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
# Moderation passed: release the withheld original chunks.
if buffered_items is not None:
for buffered_item in buffered_items:
yield buffered_item
for pending_item in pending_end_of_stream_items:
responses_yielded.append(pending_item)
yield pending_item
except ModifyResponseException as e:
if e.original_response is None:
e.original_response = responses_so_far
# Block detected during end-of-stream processing. Emit a clean
# terminating SSE sequence with the block message rather than
# propagating into a bare error blob that truncates the stream.
# The withheld original chunks are never released.
async for block_chunk in self._handle_streaming_block(
e,
endpoint_translation,
stream_started=bool(responses_yielded),
responses_so_far=responses_yielded,
):
yield block_chunk
return
except HTTPException as e:
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
request_id = _get_a2a_request_id(responses_so_far, request_data)

View file

@ -15,6 +15,9 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm._service_logger import ServiceLogging
from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
iter_client_callback_metadata_dicts,
)
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
from litellm.proxy._types import (
@ -302,6 +305,24 @@ def _key_or_team_allows_client_pricing_override(
)
def _strip_client_message_redaction_opt_out(data: dict[str, Any]) -> None:
stripped: list[str] = []
if "turn_off_message_logging" in data and _is_false_like(data["turn_off_message_logging"]):
stripped.append("turn_off_message_logging")
data.pop("turn_off_message_logging", None)
for slot_label, metadata in iter_client_callback_metadata_dicts(data):
if "turn_off_message_logging" in metadata and _is_false_like(metadata["turn_off_message_logging"]):
stripped.append(f"{slot_label}.turn_off_message_logging")
metadata.pop("turn_off_message_logging", None)
if stripped:
verbose_proxy_logger.debug(
"Stripped client-supplied message-redaction opt-out fields from request body: %s. "
"Set `allow_client_message_redaction_opt_out: true` on the key or team metadata "
"to keep these values.",
", ".join(stripped),
)
def _strip_client_pricing_overrides(data: Dict[str, Any]) -> None:
"""Drop pricing overrides from the request body and any metadata variant.
@ -1308,13 +1329,6 @@ async def add_litellm_data_to_request(
_headers,
allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out,
)
if (
not _allow_client_message_redaction_opt_out
and litellm.turn_off_message_logging is True
and "turn_off_message_logging" in data
and _is_false_like(data["turn_off_message_logging"])
):
data.pop("turn_off_message_logging", None)
verbose_proxy_logger.debug(f"Request Headers: {_headers}")
verbose_proxy_logger.debug(f"Raw Headers: {_raw_headers}")
@ -1466,6 +1480,9 @@ async def add_litellm_data_to_request(
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
_strip_client_pricing_overrides(data)
if not _allow_client_message_redaction_opt_out and litellm.turn_off_message_logging is True:
_strip_client_message_redaction_opt_out(data)
# Fill in the proxy_server_request body snapshot now that metadata has
# been parsed. Consumers (standard_logging_payload, lago,
# spend_tracking_utils, streaming_iterator) read `body` to audit the
@ -1649,6 +1666,16 @@ async def add_litellm_data_to_request(
tags_to_add=tags,
)
if _metadata_variable_name != "metadata":
_user_metadata = data.get("metadata")
if isinstance(_user_metadata, dict):
_user_tags = _user_metadata.get("tags")
if isinstance(_user_tags, list) and _user_tags:
data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags(
request_tags=data[_metadata_variable_name].get("tags"),
tags_to_add=_user_tags,
)
# Team Callbacks controls
callback_settings_obj = _get_dynamic_logging_metadata(
user_api_key_dict=user_api_key_dict, proxy_config=proxy_config

View file

@ -39,6 +39,7 @@ from litellm.proxy.management_endpoints.common_utils import (
validate_finite_spend,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_permissions_caller_permission,
generate_key_helper_fn,
prepare_metadata_fields,
)
@ -440,6 +441,11 @@ async def new_user(
detail=f"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). Attempted to create user with role: {data.user_role}. Your role: {user_api_key_dict.user_role}",
)
_check_permissions_caller_permission(
data=data,
user_api_key_dict=user_api_key_dict,
)
data_json = data.json() # type: ignore
data_json = _update_internal_new_user_params(data_json, data)
_hash_password_in_dict(data_json)
@ -1198,6 +1204,11 @@ async def _update_single_user_helper(
if not user_request.user_id and not user_request.user_email:
raise ValueError("Either user_id or user_email must be provided")
_check_permissions_caller_permission(
data=user_request,
user_api_key_dict=user_api_key_dict,
)
data_json: dict = user_request.model_dump(exclude_unset=True)
non_default_values = _update_internal_user_params(data_json=data_json, data=user_request)
_hash_password_in_dict(non_default_values)

View file

@ -530,26 +530,33 @@ def _check_allowed_routes_caller_permission(
allowed_routes: Optional[list],
user_api_key_dict: UserAPIKeyAuth,
*,
allowed_routes_was_provided: bool = False,
allow_safe_presets: bool = False,
) -> None:
"""
Only proxy admins may set `allowed_routes` on a key.
Require PROXY_ADMIN when `allowed_routes` is present in the request body,
unless the caller went through the `key_type` preset flow.
`allowed_routes` overrides the standard role-based route gate in
RouteChecks.non_proxy_admin_allowed_routes_check, so the field is
restricted to admins. Non-admins must instead use `key_type` to pick a
preset bucket — that path goes through `handle_key_type` and re-enters
this function with `allow_safe_presets=True`, which lets the derived
`llm_api_routes` / `info_routes` values through. Raw-body call sites
leave `allow_safe_presets=False` so non-admins can't write those values
directly.
Raw-body call sites pass
`allowed_routes_was_provided="allowed_routes" in data.model_fields_set` so a
caller that omits the field (model default flows through) is distinct from
one that sends any explicit value.
Post-`handle_key_type` call sites pass `allow_safe_presets=True` with the
values derived by `handle_key_type`; those values are not from the request
body, so `allowed_routes_was_provided` stays False and the safe-preset
carve-out below accepts any list of tokens in
`_NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS`.
"""
# Empty list is the default on GenerateKeyRequest — treat as "not set".
if not allowed_routes:
if not allowed_routes_was_provided and not allowed_routes:
return
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
return
if allow_safe_presets and all(r in _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS for r in allowed_routes):
if (
allow_safe_presets
and allowed_routes
and all(r in _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS for r in allowed_routes)
):
return
raise HTTPException(
status_code=403,
@ -580,7 +587,7 @@ def _check_permissions_caller_permission(
return
raise HTTPException(
status_code=403,
detail={"error": "Only proxy admins can set `permissions` on a key."},
detail={"error": "Only proxy admins can set `permissions`."},
)
@ -1566,6 +1573,7 @@ async def generate_key_fn(
_check_allowed_routes_caller_permission(
allowed_routes=data.allowed_routes,
user_api_key_dict=user_api_key_dict,
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
)
_check_passthrough_routes_caller_permission(
data=data,
@ -1736,6 +1744,7 @@ async def generate_service_account_key_fn(
_check_allowed_routes_caller_permission(
allowed_routes=data.allowed_routes,
user_api_key_dict=user_api_key_dict,
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
)
_check_passthrough_routes_caller_permission(
data=data,
@ -2060,6 +2069,11 @@ async def _process_single_key_update(
# Validate max_budget
_validate_max_budget(update_key_request.max_budget)
_check_permissions_caller_permission(
data=update_key_request,
user_api_key_dict=user_api_key_dict,
)
# Get and validate existing key
if existing_key_row is None:
existing_key_row = await _get_and_validate_existing_key(
@ -2231,6 +2245,7 @@ async def _validate_update_key_data(
_check_allowed_routes_caller_permission(
allowed_routes=data.allowed_routes,
user_api_key_dict=user_api_key_dict,
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
)
_check_passthrough_routes_caller_permission(
data=data,
@ -4551,6 +4566,7 @@ async def regenerate_key_fn(
_check_allowed_routes_caller_permission(
allowed_routes=data.allowed_routes,
user_api_key_dict=user_api_key_dict,
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
)
_check_passthrough_routes_caller_permission(
data=data,

View file

@ -654,6 +654,7 @@ if MCP_AVAILABLE:
allow_all_keys=payload.allow_all_keys,
available_on_public_internet=payload.available_on_public_internet,
timeout=payload.timeout,
max_concurrent_requests=payload.max_concurrent_requests,
)
def get_prisma_client_or_throw(message: str):

View file

@ -8361,6 +8361,22 @@ async def model_info(
)
def _blocked_response_usage(original_response: Optional[Any]) -> "litellm.Usage":
"""
Token usage for a synthetic guardrail-blocked response.
A post-call block replaces the LLM's response with the violation message,
but the upstream call already consumed tokens -- report that real usage
(carried on ``ModifyResponseException.original_response``) rather than
discarding it. Pre-call blocks never invoked the LLM (no original_response),
so usage is zero.
"""
usage = getattr(original_response, "usage", None) if original_response is not None else None
if isinstance(usage, litellm.Usage):
return usage
return litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
@router.post(
"/v1/chat/completions",
dependencies=[Depends(user_api_key_auth)],
@ -8467,6 +8483,9 @@ async def chat_completion(
_chat_response.model = e.model # type: ignore
_chat_response.choices[0].message.content = e.message # type: ignore
_chat_response.choices[0].finish_reason = "content_filter" # type: ignore
# Report the blocked LLM response's real usage (set before the stream
# branch so both paths carry it); zero for pre-call blocks.
_chat_response.usage = _blocked_response_usage(e.original_response) # type: ignore
if data.get("stream", None) is not None and data["stream"] is True:
_iterator = litellm.utils.ModelResponseIterator(model_response=_chat_response, convert_to_delta=True)
@ -8488,8 +8507,6 @@ async def chat_completion(
media_type="text/event-stream",
status_code=200, # Return 200 for passthrough mode
)
_usage = litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
_chat_response.usage = _usage # type: ignore
return _chat_response
except RejectedRequestError as e:
_data = e.request_data
@ -8618,11 +8635,7 @@ async def completion(
# Set text attribute dynamically for text completion format
setattr(_text_response.choices[0], "text", e.message)
_text_response.model = e.model # type: ignore[assignment]
_usage = litellm.Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=0,
)
_usage = _blocked_response_usage(e.original_response)
# Set usage attribute dynamically (ModelResponse accepts usage in __init__ but it's not in type definition)
setattr(_text_response, "usage", _usage)
_iterator = litellm.utils.ModelResponseIterator(model_response=_text_response, convert_to_delta=True)
@ -8647,11 +8660,7 @@ async def completion(
_response = litellm.TextCompletionResponse()
_response.choices[0].text = e.message
_response.model = e.model # type: ignore
_usage = litellm.Usage(
prompt_tokens=0,
completion_tokens=0,
total_tokens=0,
)
_usage = _blocked_response_usage(e.original_response)
_response.usage = _usage # type: ignore
return _response
except RejectedRequestError as e:

View file

@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
byok_api_key_help_url String?
source_url String?
timeout Float?
max_concurrent_requests Int?
// BYOM submission lifecycle
approval_status String? @default("active")
submitted_by String?

View file

@ -18,6 +18,7 @@ class CacheControlToolConfigInjectionPoint(TypedDict):
"""Type for tool_config-level injection points (Bedrock)."""
location: Literal["tool_config"]
control: Optional[ChatCompletionCachedContent]
CacheControlInjectionPoint = Union[

View file

@ -325,7 +325,7 @@ class ToolConfigBlock(TypedDict, total=False):
class GuardrailConfigBlock(TypedDict, total=False):
guardrailIdentifier: str
guardrailVersion: str
trace: Literal["enabled", "disabled"]
trace: Literal["enabled", "disabled", "enabled_full"]
class InferenceConfig(TypedDict, total=False):

View file

@ -118,6 +118,9 @@ class MCPServer(BaseModel):
# MCP_PER_USER_TOKEN_DEFAULT_TTL when expires_in is absent.
token_storage_ttl_seconds: Optional[int] = None
timeout: Optional[float] = None
# Max concurrent outbound tool calls to this server; excess calls queue.
# None or a value <= 0 means unlimited.
max_concurrent_requests: Optional[int] = None
# Resolved short-ID tool prefix when LITELLM_USE_SHORT_MCP_TOOL_PREFIX is
# enabled. Set by ``MCPServerManager._assign_unique_short_prefix`` at
# registration time so that natural-hash collisions between two

View file

@ -3311,6 +3311,7 @@ class LlmProviders(str, Enum):
CUSTOM = "custom"
LITELLM_PROXY = "litellm_proxy"
HOSTED_VLLM = "hosted_vllm"
TENCENT = "tencent"
LLAMAFILE = "llamafile"
LM_STUDIO = "lm_studio"
GALADRIEL = "galadriel"

View file

@ -4204,6 +4204,13 @@ def get_optional_params(
model=model,
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
)
elif custom_llm_provider == "tencent":
optional_params = litellm.TencentChatConfig().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
)
elif custom_llm_provider == "openrouter":
optional_params = litellm.OpenrouterConfig().map_openai_params(
non_default_params=non_default_params,
@ -6017,6 +6024,11 @@ def validate_environment(
keys_in_environment = True
else:
missing_keys.append("DEEPSEEK_API_KEY")
elif custom_llm_provider == "tencent":
if "TENCENT_API_KEY" in os.environ:
keys_in_environment = True
else:
missing_keys.append("TENCENT_API_KEY")
elif custom_llm_provider == "mistral":
if "MISTRAL_API_KEY" in os.environ:
keys_in_environment = True
@ -7558,6 +7570,7 @@ class ProviderConfigManager:
),
# Simple provider mappings (no model parameter needed)
LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False),
LlmProviders.TENCENT: (lambda: litellm.TencentChatConfig(), False),
LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False),
LlmProviders.BEDROCK_MANTLE: (
lambda: litellm.BedrockMantleChatConfig(),
@ -7996,6 +8009,12 @@ class ProviderConfigManager:
)
return DeepSeekAnthropicMessagesConfig()
elif litellm.LlmProviders.TENCENT == provider:
from litellm.llms.tencent.messages.transformation import (
TencentAnthropicMessagesConfig,
)
return TencentAnthropicMessagesConfig()
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
if "claude" in model_lower:
from litellm.llms.github_copilot.messages.transformation import (

File diff suppressed because it is too large Load diff

View file

@ -2295,6 +2295,24 @@
"text_completion": true
}
},
"tencent": {
"display_name": "Tencent TokenHub (`tencent`)",
"url": "https://docs.litellm.ai/docs/providers/tencent",
"endpoints": {
"chat_completions": true,
"messages": true,
"responses": true,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false,
"text_completion": false
}
},
"text-completion-codestral": {
"display_name": "Text Completion Codestral (`text-completion-codestral`)",
"url": "https://docs.litellm.ai/docs/providers/codestral",

View file

@ -62,8 +62,8 @@ proxy = [
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.26.0,<2.0",
"litellm-proxy-extras==0.4.74",
"litellm-enterprise==0.1.45",
"litellm-proxy-extras==0.4.75",
"litellm-enterprise==0.1.46",
"RestrictedPython>=8.1,<9.0",
"rich>=13.9.4,<14.0",
"polars>=1.38.1,<2.0",

View file

@ -222,7 +222,7 @@
"limit": 38
},
"RET504": {
"limit": 722
"limit": 721
},
"RUF010": {
"limit": 874
@ -249,7 +249,7 @@
"limit": 6
},
"RUF059": {
"limit": 74
"limit": 73
},
"RUF100": {
"limit": 480
@ -264,7 +264,7 @@
"limit": 63
},
"SIM102": {
"limit": 326
"limit": 325
},
"SIM103": {
"limit": 129
@ -306,7 +306,7 @@
"limit": 9
},
"TID251": {
"limit": 2714
"limit": 2710
},
"TRY002": {
"limit": 548
@ -324,7 +324,7 @@
"limit": 883
},
"UP006": {
"limit": 13041
"limit": 12870
},
"UP007": {
"limit": 2570
@ -354,7 +354,7 @@
"limit": 4
},
"UP035": {
"limit": 2300
"limit": 2295
},
"UP036": {
"limit": 4

View file

@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
byok_api_key_help_url String?
source_url String?
timeout Float?
max_concurrent_requests Int?
// BYOM submission lifecycle
approval_status String? @default("active")
submitted_by String?

View file

@ -111,7 +111,7 @@ if [ -n "$spec_files" ]; then
# and an up-to-date Prisma client; check-ui-api-types.yml installs those and runs
# prisma generate before gen:api, so mirror that here or a stale client can mask
# drift that CI will still flag.
if ! uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma; then
if ! uv run --no-sync python scripts/prisma_generate_if_needed.py; then
echo "✗ Could not regenerate Prisma client (prisma generate failed)." >&2
status=1
elif ( cd ui/litellm-dashboard && LITELLM_PYTHON="uv run --no-sync python" npm run gen:api ); then

View file

@ -0,0 +1,69 @@
#!/usr/bin/env python3
"""Run ``prisma generate`` only when its inputs changed since the last run.
The generated client is a pure function of ``litellm/proxy/schema.prisma`` and
the installed prisma package version, so a stamp of those two written next to
the venv is enough to prove the client is current. The stamp lives under
``sys.prefix`` so recreating the venv discards it, and a missing generated
client (a fresh or reinstalled prisma package) forces a regenerate even when
the stamp matches. The prisma package itself is never imported here: once
generated it re-exports the whole client on import, which costs more than the
generate this script exists to skip.
"""
import hashlib
import importlib.metadata
import importlib.util
import subprocess
import sys
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parent.parent
SCHEMA = REPO_ROOT / "litellm" / "proxy" / "schema.prisma"
STAMP = Path(sys.prefix) / "litellm-prisma-schema.stamp"
def stamp_value(schema_bytes: bytes, prisma_version: str) -> str:
return f"{hashlib.sha256(schema_bytes).hexdigest()}:{prisma_version}"
def should_skip(stamp: Path, expected: str, client_generated: bool) -> bool:
if not client_generated:
return False
try:
return stamp.read_text() == expected
except OSError:
return False
def client_is_generated() -> bool:
spec = importlib.util.find_spec("prisma")
if spec is None or not spec.submodule_search_locations:
return False
return any(
(Path(location) / "client.py").exists()
for location in spec.submodule_search_locations
)
def main() -> int:
version = importlib.metadata.version("prisma")
expected = stamp_value(SCHEMA.read_bytes(), version)
if should_skip(STAMP, expected, client_is_generated()):
print(
f"Prisma client already generated for {SCHEMA.relative_to(REPO_ROOT)} "
f"(prisma {version}); skipping prisma generate"
)
return 0
result = subprocess.run(
[sys.executable, "-m", "prisma", "generate", "--schema", str(SCHEMA)],
cwd=REPO_ROOT,
)
if result.returncode != 0:
return result.returncode
STAMP.write_text(expected)
return 0
if __name__ == "__main__":
raise SystemExit(main())

View file

@ -89,6 +89,18 @@ def base_counts(ref: str) -> dict:
shutil.rmtree(parent, ignore_errors=True)
def over_ceiling(head: dict, budget: dict) -> frozenset:
"""Rules whose head count already exceeds their limit.
A rule can only breach when it is over its limit, so when none are the base
comparison cannot change the verdict and the base worktree scan can be skipped.
"""
return frozenset(
rule for rule, spec in budget.items()
if head.get(rule, 0) > spec["limit"]
)
def evaluate(head: dict, base: dict, budget: dict) -> list:
breaches = []
for rule, spec in budget.items():
@ -119,8 +131,12 @@ def introduced(violations: list, changed: dict) -> list:
def cmd_check(base: str) -> None:
budget = json.loads(BUDGET_PATH.read_text())
head = head_violations()
head_counts = count_by_rule(head)
if not over_ceiling(head_counts, budget):
print(f"OK: every strict rule is within its codebase ceiling (base {base})")
return
base_point = _run(["git", "merge-base", base, "HEAD"]).strip() or base
breaches = evaluate(count_by_rule(head), base_counts(base_point), budget)
breaches = evaluate(head_counts, base_counts(base_point), budget)
if not breaches:
print(f"OK: every strict rule is within its codebase ceiling (base {base})")
return

View file

@ -13,9 +13,13 @@ bystander's count equals its base, so it is spared, while any PR that actually
grows the rule past its limit still fails.
Head counts are read from stdin (the caller runs basedpyright once and pipes
``--outputjson`` in); the base count is a second basedpyright pass over a
detached worktree at the merge-base, run under the same environment so import
resolution matches. ``--update`` ratchets each rule's ``limit`` down by the
``--outputjson`` in). The base count only matters once some rule is over its
limit, so when none is the base pass is skipped outright. When it is needed, it
is a second basedpyright pass over a detached worktree at the merge-base, run
under the same environment so import resolution matches, and its per-rule
counts are cached under the repo's git common dir keyed by merge-base commit,
``pyrightconfig.json``, and ``uv.lock``, so re-runs against the same branch
point pay for it once. ``--update`` ratchets each rule's ``limit`` down by the
number of errors this branch fixed relative to its branch point (the merge-base),
so the headroom you were granted shrinks by exactly what you cleared and never
grows.
@ -28,20 +32,24 @@ carries an unambiguous ``rule`` field.
import argparse
import contextlib
import hashlib
import json
import os
import shutil
import subprocess
import sys
import tempfile
from collections import Counter
from collections.abc import Iterator, Mapping
from collections.abc import Callable, Iterator, Mapping
from pathlib import Path
from typing import NamedTuple
REPO_ROOT = Path(__file__).resolve().parent.parent
BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json"
UV_LOCK = REPO_ROOT / "uv.lock"
DEFAULT_BASE = "origin/litellm_internal_staging"
CACHE_FILE_PREFIX = "basedpyright-base-"
# Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated.
UNCODED = "<uncoded>"
@ -129,6 +137,109 @@ def base_counts(ref: str) -> dict[str, int]:
return count_basedpyright(proc.stdout, root=worktree)
def over_ceiling(
head: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]
) -> frozenset[str]:
"""Rules whose head count already exceeds their limit.
A rule can only breach when it is over its limit, so when none are the base
comparison cannot change the verdict and the base worktree pass can be skipped.
"""
return frozenset(
code
for code, total in head.items()
if total > (budget[code]["limit"] if code in budget else DEFAULT_LIMIT)
)
def environment_fingerprints() -> tuple[str, ...]:
return tuple(
hashlib.sha256(path.read_bytes()).hexdigest()
for path in (PYRIGHT_CONFIG, UV_LOCK)
if path.exists()
)
def cache_key(base_point: str, fingerprints: tuple[str, ...]) -> str:
return hashlib.sha256("|".join((base_point, *fingerprints)).encode()).hexdigest()[
:16
]
def cache_path(
directory: Path, base_point: str, fingerprints: tuple[str, ...]
) -> Path:
return directory / f"{CACHE_FILE_PREFIX}{cache_key(base_point, fingerprints)}.json"
def default_cache_dir() -> Path:
common = Path(_run(["git", "rev-parse", "--git-common-dir"]).strip())
resolved = common if common.is_absolute() else REPO_ROOT / common
return resolved / "litellm-lint-cache"
def load_cached_counts(path: Path) -> dict[str, int] | None:
try:
data = json.loads(path.read_text())
except (OSError, json.JSONDecodeError):
return None
counts = data.get("counts") if isinstance(data, dict) else None
if not isinstance(counts, dict):
return None
if not all(
isinstance(code, str) and isinstance(total, int) and not isinstance(total, bool)
for code, total in counts.items()
):
return None
return counts
def scratch_path(path: Path) -> Path:
"""In-flight scratch for the tmp+rename write. Dot-prefixed so the prune
glob in `store_counts` can never match it (a concurrent run would otherwise
unlink it between write and rename), and pid-suffixed so two concurrent
writers of the same entry never share a scratch."""
return path.with_name(f".{path.name}.{os.getpid()}.tmp")
def store_counts(
directory: Path, path: Path, base_point: str, counts: Mapping[str, int]
) -> None:
directory.mkdir(parents=True, exist_ok=True)
for stale in directory.glob(f"{CACHE_FILE_PREFIX}*.json"):
if stale != path:
stale.unlink(missing_ok=True)
scratch = scratch_path(path)
scratch.write_text(
json.dumps(
{"base_point": base_point, "counts": dict(sorted(counts.items()))},
indent=2,
)
+ "\n"
)
scratch.replace(path)
def base_counts_cached(
base_point: str,
cache_dir: Path | None = None,
compute: Callable[[str], dict[str, int]] = base_counts,
) -> dict[str, int]:
"""`base_counts` memoized on disk. The base tree at a given commit is
immutable, so its counts are a pure function of the merge-base plus the
environment fingerprints in the cache key; an empty result is never stored
because it is the signature of a crashed pass, not a clean tree."""
directory = default_cache_dir() if cache_dir is None else cache_dir
path = cache_path(directory, base_point, environment_fingerprints())
cached = load_cached_counts(path)
if cached is not None:
return cached
counts = compute(base_point)
if counts:
store_counts(directory, path, base_point, counts)
return counts
def evaluate(
head: Mapping[str, int],
base: Mapping[str, int],
@ -185,7 +296,7 @@ def cmd_update(current: Mapping[str, int], base_ref: str = DEFAULT_BASE) -> None
"""
budget = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {}
base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref
updated = ratcheted_budget(budget, current, base_counts(base_point))
updated = ratcheted_budget(budget, current, base_counts_cached(base_point))
BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n")
cleared = sum(budget[code]["limit"] - updated[code]["limit"] for code in updated)
print(
@ -205,8 +316,13 @@ def cmd_check(base_ref: str) -> None:
f"nothing; refusing to certify a vacuous run."
)
raise SystemExit(1)
if not over_ceiling(head, budget):
print(
f"OK: every rule is within its basedpyright limit ({sum(head.values())} errors total)"
)
return
base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref
base = base_counts(base_point)
base = base_counts_cached(base_point)
if is_vacuous_run(base, budget):
print(
f"FAIL: basedpyright produced no errors for the base tree at "

View file

@ -9,15 +9,18 @@ When contributing to this directory, please first discuss the change you wish to
## Setup
The suites run against a live proxy, so bring one up first. `docker-compose.yml` here starts that proxy with its Postgres and Redis, serving `gateway/litellm-config.yml`; add any model, pricing override, or guardrail your test needs to that file and read it back in the test rather than hardcoding values. `gateway/` holds proxy configuration only, so never put tests there
The suites run against a live proxy, so bring one up first. `docker-compose.yml` here starts that proxy with a throwaway Postgres and Redis; `docker compose down -v` resets everything, so no state leaks between runs. The proxy config is inlined in the compose file under `configs`, prewired with example models (`gpt-5.5`, `claude-haiku-4-5`, `gemini-2.5-flash`, `openai-text-embedding-3-small`) whose keys come from your `.env`. If your test needs another model, a pricing override, or a guardrail declared up front, add it to that inline config and read it back in the test rather than hardcoding values
## Running the tests locally
1. Create a .env file and add provider keys:
1. Create a `.env` file in this directory with the provider keys the example models use:
```bash
OPENAI_API_KEY="sk-..."
ANTHROPIC_API_KEY="sk-..."
GEMINI_API_KEY="..."
```
2. Bring the stack up from this directory:
```bash
@ -123,7 +126,17 @@ Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure cover
Before you push
- Run basedpyright over your changes; the harness is fully typed and new code must not add `Any` or widen the budgets
- Bring the stack up with docker-compose from this directory and run your suite locally against it, so you exercise the same skip-vs-fail path CI does
- Use the config at `tests/e2e/gateway/litellm-config.yml` if your feature needs a model, pricing override, guardrail, or other proxy setting declared up front; add the deployment there and read it back in the test rather than hardcoding values
- Capture screenshots of the tests passing and attach them to the PR as proof of fix
1. Run basedpyright over your changes; the harness is fully typed and new code must not add `Any` or widen the budgets
2. Add the models your test needs to the inline config in `docker-compose.yml`
3. Bring the stack up and run your suite against it:
```bash
docker compose up -d
uv run pytest tests/e2e/<your_suite>/ -v
```
4. Capture screenshots of the test run and attach them to the PR as proof
5. If a test fails because it surfaced a real issue in the product, flag that explicitly in the PR rather than reworking the test until it passes

View file

@ -0,0 +1,87 @@
# local setup to run e2e tests
configs:
litellm_config:
content: |
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
database_url: os.environ/DATABASE_URL
store_prompts_in_spend_logs: true
litellm_settings:
drop_params: true
num_retries: 3
request_timeout: 600
cache: true
cache_params:
type: redis
host: redis
port: 6379
router_settings:
routing_strategy: simple-shuffle
num_retries: 3
allowed_fails: 5
cooldown_time: 30
fallbacks:
- gemini-2.5-flash: ["gpt-5.5", "claude-haiku-4-5"]
model_list:
- model_name: gpt-5.5
litellm_params:
model: openai/gpt-5.5
api_key: os.environ/OPENAI_API_KEY
- model_name: claude-haiku-4-5
litellm_params:
model: anthropic/claude-haiku-4-5
api_key: os.environ/ANTHROPIC_API_KEY
- model_name: gemini-2.5-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
- model_name: openai-text-embedding-3-small
litellm_params:
model: openai/text-embedding-3-small
api_key: os.environ/OPENAI_API_KEY
services:
litellm:
image: ghcr.io/berriai/litellm:main-latest
depends_on:
db:
condition: service_healthy
redis:
condition: service_healthy
env_file: .env
environment:
LITELLM_MASTER_KEY: sk-1234
DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm
ports:
- "4000:4000"
configs:
- source: litellm_config
target: /app/config.yaml
command: ["--config", "/app/config.yaml", "--port", "4000"]
# throwaway db
db:
image: postgres:16
environment:
POSTGRES_USER: litellm
POSTGRES_PASSWORD: litellm
POSTGRES_DB: litellm
healthcheck:
test: ["CMD-SHELL", "pg_isready -U litellm"]
interval: 3s
timeout: 3s
retries: 20
redis:
image: redis:7
healthcheck:
test: ["CMD", "redis-cli", "ping"]
interval: 3s
timeout: 3s
retries: 20

View file

@ -1,287 +0,0 @@
# This default config file aims to support most popular model providers out of the box
#In general, the model name used by the client will be the same as the ones from the provider (For example, you will use "anthropic.claude-3-5-sonnet-20240620-v1:0" when you're calling LiteLLM just like you would when calling Amazon Bedrock directly)
#In the case where there are model name conflicts, a prefix will be used (For example, the Azure and the openAI model names conflict, so when you are using Azure, you will use "azure/gpt-4o-realtime-preview-2024-10-01")
#Some model providers require additional user-specific configuration (such as Azure which requires you to specify your own api_base with your resource name, and your api_version).
#In this case, the provider is commented out, and you should uncomment it and provide your specific info
#For more detailed information about each provider, refer to the docs: https://docs.litellm.ai/docs/providers
#If you are not interested in a particular provider, just remove it from your config.yaml, and redeploy, and it will no longer show up in your LiteLLM deployment
#If a particular provider is not working, double check your .env file, and make sure you have provided a valid api key for that provider, and then redeploy
#Full details on guardrails here: https://docs.litellm.ai/docs/proxy/guardrails/bedrock
general_settings:
store_prompts_in_spend_logs: true
master_key: os.environ/LITELLM_MASTER_KEY
proxy_batch_write_at: 60
database_connection_pool_limit: 10
# disable_error_logs: True
forward_client_headers_to_llm_api: false
maximum_spend_logs_retention_period: "60d" # GSE-13389: Cleanup logs older than 60 days
maximum_spend_logs_cleanup_cron: "0 1 * * *" # 01:00 UTC daily = 18:00 PDT
database_url: os.environ/DATABASE_URL
control_plane_url: os.environ/CONTROL_PLANE_URL
alerts: ["email"]
proxy_budget_rescheduler_min_time: 15
proxy_budget_rescheduler_max_time: 20
# fallbacks: [{"gpt-4": ["anthropic.claude-3-5-sonnet-20240620-v1:0"]}] #Configure fallbacks for context window exeeded errors (In this example, we will fall back to Claude Sonnet if over 8000 tokens, which is gpt-4's limit)
# default_fallbacks: ["anthropic.claude-3-haiku-20240307-v1:0"] #Configure fallbacks for any error for every model (the above fallback configurations override this one)
# environment_variables:
# STORE_MODEL_IN_DB: 'True'
# LITELLM_LOG: "DEBUG"
litellm_settings:
drop_params: True
# Spend counters inherit this as their Redis TTL, so an idle counter goes cold and
# the next request reseeds it from the DB; kept short to exercise the cross-pod
# reseed path in test_spend_counter_reseed_e2e. Response-cache writes pass their own
# ttl and are unaffected.
default_redis_ttl: 20
request_timeout: 600
num_retries: 3
json_logs: true
store_audit_logs: True
cache: true
cache_params:
type: redis
host: redis
port: 6379
password: os.environ/REDIS_PASSWORD
namespace: litellm.caching
ttl: 16600
# max_budget: 1000000000.0 # (float) sets max budget in dollars across the entire proxy across all API keys. Note, the budget does not apply to the master key. That is the only exception.
# budget_duration: 1mo # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
# max_internal_user_budget: 1000000000.0 # (float) sets default budget in dollars for each internal user. (Doesn't apply to Admins. Doesn't apply to Teams. Doesn't apply to master key)
# internal_user_budget_duration: "1mo" # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
# success_callback: ["s3_v2"]
# failure_callback: ["s3_v2"]
# service_callback: ["datadog"]
callbacks: ["arize_phoenix", "datadog", "smtp_email", "prometheus", "otel"]
require_auth_for_metrics_endpoint: false
#type: redis-semantic
#similarity_threshold: 0.8 # similarity threshold for semantic cache
#redis_semantic_cache_embedding_model: text-embedding-ada-002 # only works with text-embedding-ada-002 for now... https://github.com/BerriAI/litellm/issues/4001
router_settings:
routing_strategy: simple-shuffle
num_retries: 3
allowed_fails: 5
cooldown_time: 30
# When gemini deployments are exhausted (provider 429 / auth), cross over to
# working models. Exercised by tests/e2e/router/test_rate_limiter.py.
fallbacks:
- gemini-2.5-flash: ["gpt-5.5", "claude-haiku-4-5"]
#ttl: Optional[float]
#default_in_memory_ttl: Optional[float]
#default_in_redis_ttl: Optional[float]
model_list:
- model_name: gpt-5.5
litellm_params:
model: openai/gpt-5.5
api_key: os.environ/OPENAI_API_KEY
- model_name: claude-haiku-4-5
litellm_params:
model: anthropic/claude-haiku-4-5
api_key: os.environ/ANTHROPIC_API_KEY
# Same underlying model via Vertex AI — distinct routing/auth path
# # (service-account JSON), so it gets its own model_name.
- model_name: gemini-2.5-flash-vertex
litellm_params:
model: vertex_ai/gemini-2.5-flash
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
- model_name: gemini-2.5-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
# load balancing to a different deployment, if gemini gets rate limited.
- model_name: gemini-2.5-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
# Custom per-token pricing exercised by llm_translation/test_custom_pricing_e2e.py.
# Rates deliberately exceed canonical gemini-2.5-flash (input 3e-7 / output 2.5e-6)
# so an override that is ignored or under-applied reports spend at the base rate
# and fails that test. The test reads these same rates back from this file.
- model_name: custom-priced-flash
litellm_params:
model: gemini/gemini-2.5-flash
api_key: os.environ/GEMINI_API_KEY
input_cost_per_token: 0.00005
output_cost_per_token: 0.0001
# embedding models
- model_name: openai-text-embedding-3-small
litellm_params:
model: openai/text-embedding-3-small
api_key: os.environ/OPENAI_API_KEY
- model_name: gemini-2-embedding
litellm_params:
model: gemini/gemini-2-embedding
api_key: os.environ/GEMINI_API_KEY
- model_name: openai-realtime
litellm_params:
model: openai/gpt-realtime
api_key: os.environ/OPENAI_API_KEY
model_info:
mode: realtime
- model_name: azure-realtime
litellm_params:
model: azure/gpt-realtime-2
api_key: os.environ/AZURE_API_KEY
api_base: os.environ/AZURE_API_BASE
api_version: "2025-08-28"
realtime_protocol: GA # Possible values: "GA"/ "v1", "beta"
model_info:
mode: realtime
- model_name: gemini-realtime
litellm_params:
model: gemini/gemini-3.1-flash-live-preview
api_key: os.environ/GEMINI_API_KEY
model_info:
mode: realtime
- model_name: vertex-realtime
litellm_params:
model: vertex_ai/gemini-live-2.5-flash-preview-native-audio-09-2025
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
model_info:
mode: realtime
- model_name: bedrock-realtime
litellm_params:
model: bedrock/amazon.nova-sonic-v1:0
aws_region_name: us-east-1
model_info:
mode: realtime
- model_name: xai-realtime
litellm_params:
model: xai/grok-voice-latest
api_key: os.environ/XAI_API_KEY
model_info:
mode: realtime
- model_name: rust-ocr-mistral
litellm_params:
model: mistral/mistral-ocr-latest
api_key: os.environ/MISTRAL_API_KEY
- model_name: rust-ocr-azure-ai
litellm_params:
model: azure_ai/mistral-document-ai-2505
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
- model_name: rust-ocr-azure-document-intelligence
litellm_params:
model: azure_ai/doc-intelligence/prebuilt-layout
api_base: os.environ/AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT
api_key: os.environ/AZURE_DOCUMENT_INTELLIGENCE_API_KEY
- model_name: rust-ocr-vertex-mistral
litellm_params:
model: vertex_ai/mistral-ocr-2505
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
- model_name: rust-ocr-vertex-deepseek
litellm_params:
model: vertex_ai/deepseek-ocr-maas
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
# batch models exercised by tests/e2e/batches/
- model_name: openai-batch
litellm_params:
model: openai/gpt-4o-mini
api_key: os.environ/OPENAI_API_KEY
model_info:
mode: batch
- model_name: azure-batch
litellm_params:
model: azure/gpt-4.1-mini-batch
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2024-07-01-preview"
model_info:
mode: batch
- model_name: vertex-batch
litellm_params:
model: vertex_ai/gemini-2.5-flash
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
bucket_name: os.environ/GCS_BUCKET_NAME
model_info:
mode: batch
- model_name: bedrock-batch
litellm_params:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
s3_bucket_name: os.environ/AWS_BATCH_S3_BUCKET
s3_region_name: us-west-2
s3_access_key_id: os.environ/AWS_ACCESS_KEY_ID
s3_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_batch_role_arn: os.environ/AWS_BATCH_ROLE_ARN
model_info:
mode: batch
files_settings:
- custom_llm_provider: openai
api_key: os.environ/OPENAI_API_KEY
- custom_llm_provider: azure
api_base: os.environ/AZURE_API_BASE
api_key: os.environ/AZURE_API_KEY
api_version: "2024-07-01-preview"
- custom_llm_provider: vertex_ai
vertex_project: os.environ/VERTEXAI_PROJECT
vertex_location: us-central1
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
bucket_name: os.environ/GCS_BUCKET_NAME
mcp_servers:
deepwiki_mcp:
url: "https://mcp.deepwiki.com/mcp"
auth_type: none
description: "just a test"
atlassian:
url: "https://mcp.atlassian.com/v1/mcp"
auth_type: oauth2
authorization_url: https://auth.atlassian.com/authorize
guardrails:
- guardrail_name: "presidio-pii"
litellm_params:
guardrail: presidio
mode: pre_call
presidio_analyzer_api_base: os.environ/PRESIDIO_ANALYZER_API_BASE
presidio_anonymizer_api_base: os.environ/PRESIDIO_ANONYMIZER_API_BASE
default_on: false
pii_entities_config:
EMAIL_ADDRESS: BLOCK
CREDIT_CARD: BLOCK
US_SSN: BLOCK
PHONE_NUMBER: BLOCK

View file

@ -28,7 +28,7 @@ Status: `covered` / `partial` / `gap`.
|----------|---------------|-----------|------------|-------------|--------|
| Gemini (`/gemini/v1beta/models/{m}:generateContent` / `:streamGenerateContent`) | live | live | live | live | **covered** |
| Anthropic (`/anthropic/v1/messages`) | live | live | live | live | **covered** |
| Vertex AI (`/vertex_ai/...`) | - | - | - | - | gap (gcloud auth) |
| Vertex AI (`/vertex_ai/v1/projects/{p}/locations/{loc}/.../models/{m}:generateContent`) | live | - | - | live | **partial** |
| OpenAI / Bedrock / Cohere / Mistral / VLLM | - | - | - | - | gap |
Each covered cell asserts: `call_type == "pass_through_endpoint"`, `spend > 0`,
@ -60,11 +60,21 @@ most likely to silently break and the one a mock can't prove works.
| `test_anthropic_passthrough_nonstreaming_logs_cost` | anthropic native, non-stream, cost |
| `test_anthropic_passthrough_streaming_logs_cost` | anthropic native, stream, cost |
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
| `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost |
Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is
added at runtime instead of declared in the gateway config: the test POSTs `/model/new`
with `use_in_pass_through`, so the proxy registers that deployment's service account for
the `/vertex_ai` route, then deletes it on teardown. The passthrough call sends only its
litellm virtual key (`x-litellm-api-key`), no upstream bearer, and the proxy mints the
Vertex token itself. Credentials (`VERTEXAI_PROJECT` / `VERTEXAI_CREDENTIALS`) are read
from the same env the proxy uses, so the test never mints a token.
## Gaps
- Vertex / OpenAI / Bedrock / Cohere passthrough (same shape; add once the
provider credential is configured; Vertex is closest - route exists, auth stale).
- Vertex streaming / tool-call passthrough (non-streaming + cost now covered).
- OpenAI / Bedrock / Cohere passthrough (same shape; add once the provider
credential is configured).
- Non-passthrough tool calls over `/chat/completions` end to end with cost.
- Image / audio / rerank / responses / realtime translation + cost.
- Streaming cost-injection (`include_cost_in_streaming_usage`); passthrough on

View file

@ -48,6 +48,16 @@ class AnthropicHeaders(Headers):
tags: str | None = None
class VertexHeaders(Headers):
# Only the litellm virtual key; the /vertex_ai passthrough mints the Vertex token
# from the proxy's own service account (the deployment marked use_in_pass_through),
# so no upstream Authorization bearer is sent from the client.
x_litellm_api_key: str = Field(serialization_alias="x-litellm-api-key")
content_type: str = Field(
default="application/json", serialization_alias="Content-Type"
)
class AltSseParams(BaseModel):
alt: str = "sse"
@ -132,6 +142,23 @@ class PassthroughClient:
stream=True,
)
# ---- Vertex AI native passthrough (/vertex_ai/v1/projects/...) -------
def vertex_generate(
self, key: str, project: str, location: str, model: str, text: str
) -> StreamingResponse:
path = (
f"/vertex_ai/v1/projects/{project}/locations/{location}"
f"/publishers/google/models/{model}:generateContent"
)
return self.gateway.transport.send(
path,
headers=VertexHeaders(x_litellm_api_key=key),
json=GeminiGenerateBody(
contents=[GeminiContent(parts=[GeminiPart(text=text)])]
),
)
# ---- Anthropic native passthrough (/anthropic/v1/messages) ----------
def anthropic_message(

View file

@ -0,0 +1,166 @@
"""Live e2e: a native Vertex AI generateContent call over the proxy's /vertex_ai
passthrough is forwarded to Vertex and still logged as a costed SpendLogs row.
Ports the de-flake of the SDK-based spend test (#31689). That test configured the
vertexai SDK with an api_endpoint override pointing at the proxy, but the SDK
intermittently ignored the override and billed the public Vertex endpoint directly,
so the request never reached LiteLLM and no spend was recorded; the bypass, not
logging lag, was the flake. Driving raw HTTP through the shared transport always
reaches the proxy, which the harness already guarantees, so the only residual
nondeterminism is the ~60s async spend flush the poll absorbs.
The vertex deployment is added at runtime through the management endpoint rather than
declared in the gateway config: the test POSTs /model/new with use_in_pass_through so
the proxy registers that deployment's service account for the /vertex_ai route, then
deletes it on teardown. The credential is the one the proxy already holds (read from
the same VERTEXAI_CREDENTIALS the deployment uses), so the passthrough call sends only
its litellm virtual key in x-litellm-api-key and no upstream bearer, and the proxy
mints the Vertex token itself. The test never mints a token.
Asserts both sides of the promise: the forward succeeds (2xx with a candidate) and
the costed row lands (call_type pass_through_endpoint, vertex_ai provider, a gemini
model, spend > 0), correlated by the x-litellm-call-id header.
"""
import os
import pytest
from pydantic import BaseModel
from e2e_config import unique_marker
from e2e_http import NoBody, require_successful_call, unwrap
from lifecycle import ResourceManager
from models import SpendLogRow
from passthrough_client import PassthroughClient
pytestmark = pytest.mark.e2e
VERTEX_MODEL = "gemini-2.5-flash"
# The added deployment's region and the passthrough URL's region are the same constant,
# so they always agree; the proxy registers passthrough credentials per project+region.
VERTEX_LOCATION = os.environ.get("VERTEXAI_LOCATION", "us-central1")
@pytest.fixture(scope="session")
def vertex_project() -> str:
"""The Vertex project to bill, read from the same VERTEXAI_PROJECT the proxy uses.
Skip when unset, since that is an environment gap rather than a behavior failure."""
project = os.environ.get("VERTEXAI_PROJECT")
if not project:
pytest.skip("set VERTEXAI_PROJECT (the project the vertex deployment bills)")
return project
@pytest.fixture(scope="session")
def vertex_credentials() -> str:
"""The service-account JSON the added deployment authenticates with, the same
VERTEXAI_CREDENTIALS the proxy holds. Skip when unset."""
credentials = os.environ.get("VERTEXAI_CREDENTIALS")
if not credentials:
pytest.skip("set VERTEXAI_CREDENTIALS (the vertex service-account JSON)")
return credentials
class _VertexDeploymentParams(BaseModel):
model: str
vertex_project: str
vertex_location: str
vertex_credentials: str
use_in_pass_through: bool
class _ModelInfoId(BaseModel):
id: str
class _ModelNewBody(BaseModel):
model_name: str
litellm_params: _VertexDeploymentParams
model_info: _ModelInfoId
class _ModelNewResponse(BaseModel):
model_id: str
class _ModelDeleteBody(BaseModel):
id: str
def _add_vertex_passthrough_model(
client: PassthroughClient, model_name: str, project: str, credentials: str
) -> str:
return unwrap(
client.gateway.transport.post(
"/model/new",
headers=client.gateway.transport.master,
json=_ModelNewBody(
model_name=model_name,
litellm_params=_VertexDeploymentParams(
model=f"vertex_ai/{VERTEX_MODEL}",
vertex_project=project,
vertex_location=VERTEX_LOCATION,
vertex_credentials=credentials,
use_in_pass_through=True,
),
model_info=_ModelInfoId(id=model_name),
),
response_type=_ModelNewResponse,
)
).model_id
def _delete_model(client: PassthroughClient, model_id: str) -> None:
_ = client.gateway.transport.post(
"/model/delete",
headers=client.gateway.transport.master,
json=_ModelDeleteBody(id=model_id),
response_type=NoBody,
)
def _costed_row(client: PassthroughClient, call_id: str | None) -> SpendLogRow:
"""The passthrough call's SpendLogs row, polled until it carries a cost.
A 2xx passthrough call that produced no costed row is a hard failure, not a skip:
a billed Vertex call that LiteLLM did not track is the exact regression #31689
guards against."""
assert call_id, "vertex passthrough response had no x-litellm-call-id header"
rows = client.gateway.poll_logs_for_request_id(
call_id,
predicate=lambda rs: (rs[0].spend or 0) > 0,
)
assert rows, f"no SpendLogs row for vertex passthrough call_id {call_id}"
row = rows[0]
assert row.call_type == "pass_through_endpoint", f"unexpected call_type: {row}"
assert (row.spend or 0) > 0, f"vertex passthrough call was not costed: {row}"
assert row.status == "success", f"unexpected status: {row}"
return row
class TestVertexPassthroughSpendTracking:
def test_vertex_passthrough_via_managed_model_logs_cost(
self,
client: PassthroughClient,
scoped_key: str,
resources: ResourceManager,
vertex_project: str,
vertex_credentials: str,
) -> None:
model_name = f"e2e-vertex-pt-{unique_marker()}"
model_id = _add_vertex_passthrough_model(client, model_name, vertex_project, vertex_credentials)
resources.defer(lambda: _delete_model(client, model_id))
result = client.vertex_generate(
key=scoped_key,
project=vertex_project,
location=VERTEX_LOCATION,
model=VERTEX_MODEL,
text=f"reply with one word {unique_marker()}",
)
require_successful_call(result)
assert '"candidates"' in result.body, f"vertex passthrough returned no candidates: {result.body[:300]}"
row = _costed_row(client, result.call_id)
assert row.custom_llm_provider == "vertex_ai", f"passthrough spend logged under the wrong provider: {row}"
assert "gemini" in (row.model or ""), f"unexpected model in spend log: {row}"

View file

@ -56,69 +56,56 @@ async def test_global_redaction_on():
)
@pytest.mark.parametrize("turn_off_message_logging", [True, False])
@pytest.mark.parametrize(
"dynamic_turn_off, expect_redacted",
[(True, True), (False, False)],
)
@pytest.mark.asyncio
async def test_global_redaction_ignores_dynamic_param(turn_off_message_logging):
"""
Request-body `turn_off_message_logging` is no longer honored as a dynamic
callback param — global setting (or admin-configured key/team config) wins.
With global redaction ON, the caller cannot disable redaction via the
request body.
"""
async def test_dynamic_turn_off_message_logging_overrides_global_on(dynamic_turn_off, expect_redacted):
litellm.turn_off_message_logging = True
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
response = await litellm.acompletion(
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
turn_off_message_logging=turn_off_message_logging,
turn_off_message_logging=dynamic_turn_off,
mock_response="hello",
)
await asyncio.sleep(1)
standard_logging_payload = test_custom_logger.logged_standard_logging_payload
assert standard_logging_payload is not None
print(
"logged standard logging payload",
json.dumps(standard_logging_payload, indent=2),
)
response = standard_logging_payload["response"]
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
expected_response_content = "redacted-by-litellm" if expect_redacted else "hello"
expected_message_content = "redacted-by-litellm" if expect_redacted else "hi"
assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content
assert standard_logging_payload["messages"][0]["content"] == expected_message_content
@pytest.mark.parametrize("turn_off_message_logging", [True, False])
@pytest.mark.parametrize(
"dynamic_turn_off, expect_redacted",
[(True, True), (False, False)],
)
@pytest.mark.asyncio
async def test_global_redaction_off_ignores_dynamic_param(turn_off_message_logging):
"""
Request-body `turn_off_message_logging` is no longer honored as a dynamic
callback param — global setting (or admin-configured key/team config) wins.
With global redaction OFF, the caller cannot enable redaction via the
request body.
"""
async def test_dynamic_turn_off_message_logging_overrides_global_off(dynamic_turn_off, expect_redacted):
litellm.turn_off_message_logging = False
test_custom_logger = TestCustomLogger()
litellm.callbacks = [test_custom_logger]
response = await litellm.acompletion(
await litellm.acompletion(
model="gpt-5-mini",
messages=[{"role": "user", "content": "hi"}],
turn_off_message_logging=turn_off_message_logging,
turn_off_message_logging=dynamic_turn_off,
mock_response="hello",
)
await asyncio.sleep(1)
standard_logging_payload = test_custom_logger.logged_standard_logging_payload
assert standard_logging_payload is not None
print(
"logged standard logging payload",
json.dumps(standard_logging_payload, indent=2),
)
assert (
standard_logging_payload["response"]["choices"][0]["message"]["content"]
== "hello"
)
assert standard_logging_payload["messages"][0]["content"] == "hi"
expected_response_content = "redacted-by-litellm" if expect_redacted else "hello"
expected_message_content = "redacted-by-litellm" if expect_redacted else "hi"
assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content
assert standard_logging_payload["messages"][0]["content"] == expected_message_content
@pytest.mark.asyncio

View file

@ -61,8 +61,8 @@ class TestAzureDocumentIntelligencePagesParam:
def cfg(self) -> AzureDocumentIntelligenceOCRConfig:
return AzureDocumentIntelligenceOCRConfig()
def test_get_supported_ocr_params_includes_pages(self, cfg):
assert cfg.get_supported_ocr_params("prebuilt-layout") == ["pages"]
def test_get_supported_ocr_params_includes_pages_and_features(self, cfg):
assert cfg.get_supported_ocr_params("prebuilt-layout") == ["pages", "features"]
def test_map_ocr_params_mistral_zero_based_int_list(self, cfg):
mapped = cfg.map_ocr_params({"pages": [0, 1, 2]}, {}, "prebuilt-layout")

View file

@ -103,9 +103,7 @@ async def _mock_stream_messages(a2a_client: Any, request: Any) -> AsyncIterator[
kind="message",
)
for _ in range(2):
yield SendStreamingMessageResponse(
root=SendStreamingMessageSuccessResponse(id=request.id, result=msg)
)
yield SendStreamingMessageResponse(root=SendStreamingMessageSuccessResponse(id=request.id, result=msg))
class CostLogger(CustomLogger):
@ -119,9 +117,7 @@ class CostLogger(CustomLogger):
slp = kwargs.get("standard_logging_object")
if slp:
self.response_cost = (
slp.get("response_cost")
if isinstance(slp, dict)
else getattr(slp, "response_cost", None)
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
)
@ -160,6 +156,43 @@ async def test_asend_message_uses_cost_per_query():
assert cost_logger.response_cost == 0.05
@pytest.mark.asyncio
async def test_asend_message_uses_cost_per_query_from_litellm_params_dict():
"""
Proxy passes agent pricing as the litellm_params dict param (not top-level
kwargs). Regression for cost_per_query landing at $0 on the native path.
"""
from litellm.a2a_protocol import asend_message
litellm.logging_callback_manager._reset_all_callbacks()
cost_logger = CostLogger()
litellm.callbacks = [cost_logger]
mock_client = MagicMock()
mock_client._litellm_agent_card = MagicMock()
mock_client._litellm_agent_card.name = "test-agent"
mock_request = _make_send_message_request("test-123")
with patch(
"litellm.a2a_protocol.main._execute_a2a_send_with_retry",
new=_mock_execute_a2a_send,
):
await asend_message(
a2a_client=mock_client,
request=mock_request,
litellm_params={
"cost_per_query": 0.5,
"input_cost_per_token": 0.099999,
"output_cost_per_token": 0.1,
},
)
await asyncio.sleep(0.1)
assert cost_logger.response_cost == 0.5
class TokenAndCostLogger(CustomLogger):
"""Custom logger to capture both token counts and cost."""
@ -173,19 +206,13 @@ class TokenAndCostLogger(CustomLogger):
slp = kwargs.get("standard_logging_object")
if slp:
self.response_cost = (
slp.get("response_cost")
if isinstance(slp, dict)
else getattr(slp, "response_cost", None)
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
)
self.prompt_tokens = (
slp.get("prompt_tokens")
if isinstance(slp, dict)
else getattr(slp, "prompt_tokens", None)
slp.get("prompt_tokens") if isinstance(slp, dict) else getattr(slp, "prompt_tokens", None)
)
self.completion_tokens = (
slp.get("completion_tokens")
if isinstance(slp, dict)
else getattr(slp, "completion_tokens", None)
slp.get("completion_tokens") if isinstance(slp, dict) else getattr(slp, "completion_tokens", None)
)
@ -207,9 +234,7 @@ async def test_asend_message_uses_input_output_cost_per_token():
mock_client._litellm_agent_card = MagicMock()
mock_client._litellm_agent_card.name = "test-agent"
mock_request = _make_send_message_request(
"test-123", user_text="Hello, what can you do?"
)
mock_request = _make_send_message_request("test-123", user_text="Hello, what can you do?")
# Define specific cost per token values
input_cost_per_token = 0.00001 # $0.01 per 1000 tokens
@ -246,15 +271,11 @@ async def test_asend_message_uses_input_output_cost_per_token():
assert response_cost is not None, "response_cost should be captured"
# Calculate expected cost
expected_cost = (prompt_tokens * input_cost_per_token) + (
completion_tokens * output_cost_per_token
)
expected_cost = (prompt_tokens * input_cost_per_token) + (completion_tokens * output_cost_per_token)
print(f"expected_cost: {expected_cost}")
# Verify exact cost calculation
assert (
response_cost == expected_cost
), f"response_cost {response_cost} should equal expected {expected_cost}"
assert response_cost == expected_cost, f"response_cost {response_cost} should equal expected {expected_cost}"
class AgentIdLogger(CustomLogger):
@ -305,9 +326,9 @@ async def test_asend_message_passes_agent_id_to_callback():
await asyncio.sleep(0.1)
# Verify agent_id was passed to callback
assert (
agent_id_logger.agent_id == test_agent_id
), f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'"
assert agent_id_logger.agent_id == test_agent_id, (
f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'"
)
class MetadataLogger(CustomLogger):
@ -418,9 +439,7 @@ async def test_asend_message_streaming_triggers_callbacks():
assert len(chunks) == 2
# Verify callbacks WERE triggered after stream completed
assert (
callback_logger.kwargs is not None
), "Streaming should trigger callbacks after completion"
assert (
callback_logger.agent_id == test_agent_id
), f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
assert callback_logger.kwargs is not None, "Streaming should trigger callbacks after completion"
assert callback_logger.agent_id == test_agent_id, (
f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
)

View file

@ -0,0 +1,41 @@
"""Tests for litellm/a2a_protocol/utils.py token/usage extraction."""
import pytest
pytest.importorskip("a2a.compat.v0_3.types")
from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest
from litellm.a2a_protocol.utils import A2ARequestUtils
def _request(user_text: str) -> SendMessageRequest:
return SendMessageRequest(
id="r1",
params=MessageSendParams(
message={
"messageId": "m1",
"role": "user",
"parts": [{"kind": "text", "text": user_text}],
}
),
)
def test_calculate_usage_counts_input_tokens_from_request_object():
"""Regression: request-side Part is a RootModel; input tokens must be counted."""
request = _request("count these input tokens please")
response_dict = {
"result": {
"kind": "message",
"parts": [{"kind": "text", "text": "ok"}],
}
}
prompt_tokens, completion_tokens, total_tokens = A2ARequestUtils.calculate_usage_from_request_response(
request=request, response_dict=response_dict
)
assert prompt_tokens > 0
assert completion_tokens > 0
assert total_tokens == prompt_tokens + completion_tokens

View file

@ -0,0 +1,168 @@
"""
Unit tests for the per-request budget-metric emission timeout in
PrometheusLogger._increment_remaining_budget_metrics.
A slow Redis/DB lookup in one of the budget branches must not let the gather run
unbounded; it is wrapped in asyncio.wait_for so the success-logging coroutine
cannot exceed the LoggingWorker watchdog and get the whole event cancelled.
"""
import asyncio
from unittest.mock import AsyncMock, patch
import pytest
from prometheus_client import REGISTRY
from litellm.integrations.prometheus import (
PrometheusLogger,
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT,
_get_budget_metrics_per_request_timeout,
)
TIMEOUT_ENV = "PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT"
@pytest.fixture(autouse=True)
def cleanup_prometheus_registry():
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
yield
collectors = list(REGISTRY._collector_to_names.keys())
for collector in collectors:
try:
REGISTRY.unregister(collector)
except Exception:
pass
@pytest.fixture
def prometheus_logger():
return PrometheusLogger()
def _call_increment(logger: PrometheusLogger):
return logger._increment_remaining_budget_metrics(
user_api_team="team-1",
user_api_team_alias="team-alias",
user_api_key="key-1",
user_api_key_alias="key-alias",
litellm_params={"metadata": {}},
response_cost=0.01,
user_id="user-1",
user_api_key_org_id="org-1",
)
def _skip_logged(debug_mock) -> bool:
return any("skipping" in str(call.args[0]) for call in debug_mock.call_args_list if call.args)
@pytest.mark.asyncio
async def test_budget_metric_emission_skips_on_timeout(prometheus_logger, monkeypatch):
"""A branch slower than the timeout is skipped without propagating, and the
skip is logged instead of cancelling the success-logging event."""
monkeypatch.setenv(TIMEOUT_ENV, "0.05")
async def _slow_branch(**kwargs):
await asyncio.sleep(30)
prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger:
await _call_increment(prometheus_logger)
assert _skip_logged(mock_logger.debug)
@pytest.mark.asyncio
async def test_budget_metric_emission_completes_within_timeout(prometheus_logger, monkeypatch):
"""With a generous timeout every branch is awaited and no skip is logged."""
monkeypatch.setenv(TIMEOUT_ENV, "5.0")
prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger:
await _call_increment(prometheus_logger)
assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1
assert prometheus_logger._set_team_budget_metrics_after_api_request.await_count == 1
assert prometheus_logger._set_user_budget_metrics_after_api_request.await_count == 1
assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1
assert not _skip_logged(mock_logger.debug)
@pytest.mark.asyncio
async def test_invalid_timeout_env_falls_back_to_default(prometheus_logger, monkeypatch):
"""A malformed timeout env value must not raise (which would recreate the
failure mode); it falls back to the default and every branch still runs."""
monkeypatch.setenv(TIMEOUT_ENV, "not-a-number")
prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
await _call_increment(prometheus_logger)
assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1
assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1
@pytest.mark.parametrize("value", ["not-a-number", "0", "-1", "nan", "inf", "-inf"])
def test_unusable_timeout_env_falls_back_to_default(value, monkeypatch):
"""Values that parse but disable or unbound the timeout (0, negative, nan,
inf) must fall back to the default instead of being used; otherwise they
either skip every emission or recreate the unbounded-wait failure mode."""
monkeypatch.setenv(TIMEOUT_ENV, value)
assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
@pytest.mark.parametrize("value,expected", [("0.05", 0.05), ("5.0", 5.0), ("30", 30.0)])
def test_valid_timeout_env_is_used(value, expected, monkeypatch):
"""A finite positive value is parsed and returned unchanged."""
monkeypatch.setenv(TIMEOUT_ENV, value)
assert _get_budget_metrics_per_request_timeout() == expected
def test_missing_timeout_env_uses_default(monkeypatch):
"""With the env unset the default is returned."""
monkeypatch.delenv(TIMEOUT_ENV, raising=False)
assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
@pytest.mark.asyncio
async def test_outer_cancellation_still_propagates(prometheus_logger, monkeypatch):
"""Only asyncio.TimeoutError is swallowed; an outer cancellation (cooperative
shutdown / watchdog) injected while awaiting must still propagate."""
monkeypatch.setenv(TIMEOUT_ENV, "30")
started = asyncio.Event()
async def _slow_branch(**kwargs):
started.set()
await asyncio.sleep(30)
prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
task = asyncio.create_task(_call_increment(prometheus_logger))
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task

View file

@ -1194,3 +1194,156 @@ def test_s3_callback_params_override_empty_dict_is_opt_in():
assert logger.s3_bucket_name is None
finally:
litellm.s3_callback_params = original
def _expected_content_md5(payload: dict) -> str:
import base64
import hashlib
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
json_string = safe_dumps(payload)
return base64.b64encode(
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
).decode()
def _require_non_security_md5(monkeypatch):
import hashlib
original_md5 = hashlib.md5
def fips_md5(data=b"", *, usedforsecurity=True):
if usedforsecurity:
raise ValueError("MD5 blocked for security use")
return original_md5(data, usedforsecurity=usedforsecurity)
monkeypatch.setattr(hashlib, "md5", fips_md5)
@pytest.mark.asyncio
async def test_async_upload_sets_content_md5_header(monkeypatch):
"""
Object Lock buckets reject PUTs without a Content-MD5 header (AWS spec).
The async upload must send a base64 md5 of the exact signed body.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
)
payload = {"test": "content-md5"}
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-md5.json",
payload=payload,
s3_object_download_filename="test-md5.json",
)
_require_non_security_md5(monkeypatch)
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put.return_value = response
await logger.async_upload_data_to_s3(test_element)
headers = logger.async_httpx_client.put.call_args.kwargs["headers"]
assert headers["Content-MD5"] == _expected_content_md5(payload)
assert "x-amz-server-side-encryption" not in headers
def test_sync_upload_sets_content_md5_header(monkeypatch):
"""The sync upload path must also send Content-MD5 for Object Lock buckets."""
from unittest.mock import MagicMock
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
)
payload = {"test": "sync-content-md5"}
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-sync-md5.json",
payload=payload,
s3_object_download_filename="test-sync-md5.json",
)
_require_non_security_md5(monkeypatch)
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
mock_sync_client = MagicMock()
mock_sync_client.put.return_value = response
with patch(
"litellm.integrations.s3_v2._get_httpx_client",
return_value=mock_sync_client,
):
logger.upload_data_to_s3(test_element)
headers = mock_sync_client.put.call_args.kwargs["headers"]
assert headers["Content-MD5"] == _expected_content_md5(payload)
assert "x-amz-server-side-encryption" not in headers
@pytest.mark.asyncio
async def test_async_upload_sets_server_side_encryption_header_when_configured():
"""
When s3_server_side_encryption is set (e.g. buckets with a KMS default
encryption policy), the PUT must carry x-amz-server-side-encryption.
"""
from unittest.mock import AsyncMock, MagicMock
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_server_side_encryption="aws:kms",
)
test_element = s3BatchLoggingElement(
s3_object_key="2025-09-14/test-sse.json",
payload={"test": "sse"},
s3_object_download_filename="test-sse.json",
)
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put.return_value = response
await logger.async_upload_data_to_s3(test_element)
headers = logger.async_httpx_client.put.call_args.kwargs["headers"]
assert headers["x-amz-server-side-encryption"] == "aws:kms"
def test_s3_server_side_encryption_read_from_callback_params():
"""s3_server_side_encryption can be configured via s3_callback_params."""
import litellm
original = litellm.s3_callback_params
litellm.s3_callback_params = {
"s3_bucket_name": "from-global",
"s3_server_side_encryption": "aws:kms",
}
try:
logger = S3Logger()
assert logger.s3_server_side_encryption == "aws:kms"
finally:
litellm.s3_callback_params = original

View file

@ -1,5 +1,6 @@
import base64
import json
import os
from unittest.mock import MagicMock, patch
import pytest
@ -2744,102 +2745,132 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5():
TTL ordering constraint (tools -> system -> messages).
Ref: https://github.com/BerriAI/litellm/issues/XXXXX
Forces the bundled local cost map so ttl eligibility (driven by
`cache_creation_input_token_cost_above_1hr` in litellm.model_cost) reads
this branch's pricing data rather than the network-fetched `main` copy,
which lacks the fix until merge.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
add_cache_point_tool_block,
)
tool_with_1h = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
tool_with_1h = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
# Claude 4.5 model: ttl should be preserved
result = add_cache_point_tool_block(
tool_with_1h, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
assert result is not None
assert result["cachePoint"]["type"] == "default"
assert result["cachePoint"]["ttl"] == "1h"
# Claude 4.5 model: ttl should be preserved
result = add_cache_point_tool_block(
tool_with_1h, model="jp.anthropic.claude-opus-4-7"
)
assert result is not None
assert result["cachePoint"]["type"] == "default"
assert result["cachePoint"]["ttl"] == "1h"
# Claude 4.5 model with 5m ttl: also preserved
tool_with_5m = {
"cache_control": {"type": "ephemeral", "ttl": "5m"},
}
result_5m = add_cache_point_tool_block(
tool_with_5m, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
assert result_5m is not None
assert result_5m["cachePoint"]["ttl"] == "5m"
# Claude 4.5 model with 5m ttl: also preserved
tool_with_5m = {
"cache_control": {"type": "ephemeral", "ttl": "5m"},
}
result_5m = add_cache_point_tool_block(
tool_with_5m, model="jp.anthropic.claude-opus-4-7"
)
assert result_5m is not None
assert result_5m["cachePoint"]["ttl"] == "5m"
# Older model: ttl should be stripped
result_old = add_cache_point_tool_block(
tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
assert result_old is not None
assert result_old["cachePoint"]["type"] == "default"
assert "ttl" not in result_old["cachePoint"]
# Older model: ttl should be stripped
result_old = add_cache_point_tool_block(
tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
assert result_old is not None
assert result_old["cachePoint"]["type"] == "default"
assert "ttl" not in result_old["cachePoint"]
# No model provided: ttl should be stripped (safe default)
result_no_model = add_cache_point_tool_block(tool_with_1h, model=None)
assert result_no_model is not None
assert "ttl" not in result_no_model["cachePoint"]
# No model provided: ttl should be stripped (safe default)
result_no_model = add_cache_point_tool_block(tool_with_1h, model=None)
assert result_no_model is not None
assert "ttl" not in result_no_model["cachePoint"]
# No cache_control: returns None (unchanged behavior)
tool_no_cache = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
}
assert add_cache_point_tool_block(tool_no_cache) is None
# No cache_control: returns None (unchanged behavior)
tool_no_cache = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
}
assert add_cache_point_tool_block(tool_no_cache) is None
# cache_control without ttl: returns default cachePoint (unchanged behavior)
tool_no_ttl = {"cache_control": {"type": "ephemeral"}}
result_no_ttl = add_cache_point_tool_block(
tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
assert result_no_ttl is not None
assert result_no_ttl["cachePoint"]["type"] == "default"
assert "ttl" not in result_no_ttl["cachePoint"]
# cache_control without ttl: returns default cachePoint (unchanged behavior)
tool_no_ttl = {"cache_control": {"type": "ephemeral"}}
result_no_ttl = add_cache_point_tool_block(
tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0"
)
assert result_no_ttl is not None
assert result_no_ttl["cachePoint"]["type"] == "default"
assert "ttl" not in result_no_ttl["cachePoint"]
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
"""
End-to-end: _bedrock_tools_pt should produce cachePoint blocks with ttl
for Claude 4.5+ models when tools have cache_control with ttl.
Forces the bundled local cost map so ttl eligibility (driven by
`cache_creation_input_token_cost_above_1hr` in litellm.model_cost) reads
this branch's pricing data rather than the network-fetched `main` copy,
which lacks the fix until merge.
"""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
# Claude 4.5: cachePoint should have ttl
result = _bedrock_tools_pt(
tools, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
cache_blocks = [b for b in result if "cachePoint" in b]
assert len(cache_blocks) == 1
assert cache_blocks[0]["cachePoint"]["ttl"] == "1h"
# Claude 4.5: cachePoint should have ttl
result = _bedrock_tools_pt(tools, model="jp.anthropic.claude-opus-4-7")
cache_blocks = [b for b in result if "cachePoint" in b]
assert len(cache_blocks) == 1
assert cache_blocks[0]["cachePoint"]["ttl"] == "1h"
# Older model: cachePoint should not have ttl
result_old = _bedrock_tools_pt(
tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
assert len(cache_blocks_old) == 1
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
# Older model: cachePoint should not have ttl
result_old = _bedrock_tools_pt(
tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
assert len(cache_blocks_old) == 1
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document():

View file

@ -7,9 +7,53 @@ sys.path.insert(0, os.path.abspath("../../.."))
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
initialize_standard_callback_dynamic_params,
iter_client_callback_metadata_dicts,
)
def test_iter_client_callback_metadata_dicts_covers_all_read_paths():
md = {"m": 1}
lm = {"lm": 1}
lp_md = {"lp": 1}
slots = dict(
iter_client_callback_metadata_dicts(
{
"metadata": md,
"litellm_metadata": lm,
"litellm_params": {"metadata": lp_md},
}
)
)
assert slots == {
"metadata": md,
"litellm_metadata": lm,
"litellm_params.metadata": lp_md,
}
def test_iter_client_callback_metadata_dicts_skips_non_dict_slots():
slots = list(
iter_client_callback_metadata_dicts(
{
"metadata": "not-a-dict",
"litellm_metadata": None,
"litellm_params": {"metadata": []},
}
)
)
assert slots == []
def test_extractor_reads_turn_off_message_logging_from_every_slot():
for kwargs in (
{"metadata": {"turn_off_message_logging": True}},
{"litellm_metadata": {"turn_off_message_logging": True}},
{"litellm_params": {"metadata": {"turn_off_message_logging": True}}},
):
params = initialize_standard_callback_dynamic_params(kwargs)
assert params.get("turn_off_message_logging") is True, kwargs
def test_resolves_plain_values_at_top_level():
kwargs = {
"langfuse_public_key": "pk-test",
@ -36,6 +80,33 @@ def test_resolves_plain_values_from_metadata():
assert params.get("langfuse_host") == "https://test.langfuse.com"
def test_litellm_params_metadata_overrides_metadata():
kwargs = {
"metadata": {
"langfuse_public_key": "pk-meta",
},
"litellm_params": {
"metadata": {
"langfuse_public_key": "pk-litellm-params",
}
},
}
params = initialize_standard_callback_dynamic_params(kwargs)
assert params.get("langfuse_public_key") == "pk-litellm-params"
def test_top_level_kwargs_overrides_metadata_slots():
kwargs = {
"langfuse_public_key": "from-top-level",
"metadata": {"langfuse_public_key": "from-metadata"},
"litellm_params": {"metadata": {"langfuse_public_key": "from-litellm-params"}},
}
params = initialize_standard_callback_dynamic_params(kwargs)
assert params.get("langfuse_public_key") == "from-top-level"
def test_env_reference_at_top_level_raises_with_guidance():
kwargs = {"langfuse_public_key": "os.environ/LANGFUSE_PUBLIC_KEY"}
@ -100,11 +171,17 @@ def test_non_string_values_are_not_flagged():
assert params.get("langsmith_sampling_rate") == 0.5
def test_turn_off_message_logging_not_extracted_from_request():
"""turn_off_message_logging is admin-only — must not be settable via request."""
kwargs = {"turn_off_message_logging": True}
@pytest.mark.parametrize(
"kwargs,expected",
[
({"turn_off_message_logging": False}, False),
({"turn_off_message_logging": "False"}, "False"),
({"metadata": {"turn_off_message_logging": True}}, True),
],
)
def test_turn_off_message_logging_extracted_from_kwargs(kwargs, expected):
params = initialize_standard_callback_dynamic_params(kwargs)
assert params.get("turn_off_message_logging") is None
assert params.get("turn_off_message_logging") == expected
def test_empty_kwargs_returns_empty_params():

View file

@ -0,0 +1,42 @@
"""Tests for litellm/llms/a2a/chat/transformation.py response transform."""
from unittest.mock import MagicMock
from litellm.llms.a2a.chat.transformation import A2AConfig
from litellm.types.utils import ModelResponse
def _raw_response(text: str) -> MagicMock:
raw = MagicMock()
raw.status_code = 200
raw.headers = {}
raw.json.return_value = {
"jsonrpc": "2.0",
"id": "resp-1",
"result": {
"kind": "message",
"parts": [{"kind": "text", "text": text}],
},
}
return raw
def test_transform_response_sets_usage():
"""Regression: A2AConfig.transform_response must populate usage so per-token
pricing computes real cost and callers don't get usage 0/0/0."""
result = A2AConfig().transform_response(
model="a2a/test-agent",
raw_response=_raw_response("hello from the agent"),
model_response=ModelResponse(),
logging_obj=MagicMock(),
request_data={},
messages=[{"role": "user", "content": "hi there agent"}],
optional_params={},
litellm_params={},
encoding=None,
)
assert result.usage is not None
assert result.usage.prompt_tokens > 0
assert result.usage.completion_tokens > 0
assert result.usage.total_tokens == (result.usage.prompt_tokens + result.usage.completion_tokens)

View file

@ -1,3 +1,6 @@
from unittest.mock import MagicMock
import httpx
import pytest
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
@ -31,3 +34,217 @@ def test_should_reject_dot_segment_azure_document_intelligence_model_id():
optional_params={},
litellm_params={},
)
AZURE_TABLES = [
{
"rowCount": 2,
"columnCount": 2,
"cells": [
{"kind": "columnHeader", "rowIndex": 0, "columnIndex": 0, "content": "Item"},
{"kind": "columnHeader", "rowIndex": 0, "columnIndex": 1, "content": "Price"},
{"rowIndex": 1, "columnIndex": 0, "content": "Widget"},
{"rowIndex": 1, "columnIndex": 1, "content": "$100.00"},
],
},
{
"rowCount": 1,
"columnCount": 1,
"cells": [{"rowIndex": 0, "columnIndex": 0, "content": "Totals"}],
},
]
AZURE_KEY_VALUE_PAIRS = [
{"key": {"content": "Invoice No"}, "value": {"content": "INV-12345"}, "confidence": 0.98},
{"key": {"content": "Total"}, "value": {"content": "$100.00"}, "confidence": 0.95},
]
AZURE_ANALYZE_SUCCEEDED = {
"status": "succeeded",
"createdDateTime": "2026-07-02T00:00:00Z",
"lastUpdatedDateTime": "2026-07-02T00:00:05Z",
"analyzeResult": {
"apiVersion": "2024-11-30",
"modelId": "prebuilt-layout",
"content": "Invoice\nInvoice No: INV-12345\nTotal: $100.00",
"pages": [
{
"pageNumber": 1,
"width": 8.5,
"height": 11,
"unit": "inch",
"lines": [
{"content": "Invoice"},
{"content": "Invoice No: INV-12345"},
{"content": "Total: $100.00"},
],
}
],
"tables": AZURE_TABLES,
"keyValuePairs": AZURE_KEY_VALUE_PAIRS,
},
}
def _completed_response(payload: dict) -> httpx.Response:
return httpx.Response(
status_code=200,
json=payload,
request=httpx.Request("GET", "https://example.cognitiveservices.azure.com/analyzeResults/xyz"),
)
def _assert_native_fields_preserved(serialized: dict) -> None:
assert serialized["content"] == "Invoice\nInvoice No: INV-12345\nTotal: $100.00"
assert serialized["tables"] == AZURE_TABLES
assert serialized["keyValuePairs"] == AZURE_KEY_VALUE_PAIRS
assert serialized["object"] == "ocr"
assert serialized["usage_info"]["pages_processed"] == 1
assert serialized["pages"][0]["index"] == 0
assert serialized["pages"][0]["markdown"] == "Invoice\nInvoice No: INV-12345\nTotal: $100.00"
assert serialized["pages"][0]["dimensions"] == {"width": 816, "height": 1056, "dpi": 96}
def test_transform_ocr_response_preserves_azure_native_fields():
config = AzureDocumentIntelligenceOCRConfig()
result = config.transform_ocr_response(
model="azure_ai/doc-intelligence/prebuilt-layout",
raw_response=_completed_response(AZURE_ANALYZE_SUCCEEDED),
logging_obj=MagicMock(),
)
_assert_native_fields_preserved(result.model_dump())
@pytest.mark.asyncio
async def test_async_transform_ocr_response_preserves_azure_native_fields():
config = AzureDocumentIntelligenceOCRConfig()
result = await config.async_transform_ocr_response(
model="azure_ai/doc-intelligence/prebuilt-layout",
raw_response=_completed_response(AZURE_ANALYZE_SUCCEEDED),
logging_obj=MagicMock(),
)
_assert_native_fields_preserved(result.model_dump())
def test_transform_ocr_response_tolerates_missing_native_fields():
config = AzureDocumentIntelligenceOCRConfig()
payload = {
"status": "succeeded",
"analyzeResult": {
"pages": [
{
"pageNumber": 1,
"width": 8.5,
"height": 11,
"unit": "inch",
"lines": [{"content": "hello"}],
}
],
},
}
result = config.transform_ocr_response(
model="azure_ai/doc-intelligence/prebuilt-read",
raw_response=_completed_response(payload),
logging_obj=MagicMock(),
)
serialized = result.model_dump()
assert serialized["pages"][0]["markdown"] == "hello"
assert serialized["content"] is None
assert serialized["tables"] is None
assert serialized["keyValuePairs"] is None
def test_transform_ocr_response_non_succeeded_status_raises():
config = AzureDocumentIntelligenceOCRConfig()
with pytest.raises(ValueError, match="failed with status: failed"):
config.transform_ocr_response(
model="azure_ai/doc-intelligence/prebuilt-layout",
raw_response=_completed_response({"status": "failed"}),
logging_obj=MagicMock(),
)
def test_get_supported_ocr_params_includes_features():
config = AzureDocumentIntelligenceOCRConfig()
assert config.get_supported_ocr_params("prebuilt-layout") == ["pages", "features"]
@pytest.mark.parametrize(
"features,expected",
[
(["keyValuePairs"], "keyValuePairs"),
(["keyValuePairs", "languages"], "keyValuePairs,languages"),
("keyValuePairs", "keyValuePairs"),
("keyValuePairs,languages", "keyValuePairs,languages"),
("keyValuePairs, languages", "keyValuePairs,languages"),
],
)
def test_map_ocr_params_features(features, expected):
config = AzureDocumentIntelligenceOCRConfig()
mapped = config.map_ocr_params({"features": features}, {}, "prebuilt-layout")
assert mapped == {"features": expected}
def test_map_ocr_params_empty_features_list_omitted():
config = AzureDocumentIntelligenceOCRConfig()
assert config.map_ocr_params({"features": []}, {}, "prebuilt-layout") == {}
@pytest.mark.parametrize(
"features",
[
"keyValuePairs&pages=9",
"key value pairs",
"",
[1, 2],
[["keyValuePairs"]],
{"feature": "keyValuePairs"},
5,
],
)
def test_map_ocr_params_invalid_features_raises(features):
config = AzureDocumentIntelligenceOCRConfig()
with pytest.raises(ValueError, match="Invalid `features`"):
config.map_ocr_params({"features": features}, {}, "prebuilt-layout")
def test_get_complete_url_appends_features_query():
config = AzureDocumentIntelligenceOCRConfig()
url = config.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="azure_ai/doc-intelligence/prebuilt-layout",
optional_params={"features": "keyValuePairs"},
)
assert "&features=keyValuePairs" in url
def test_get_complete_url_combines_pages_and_features():
config = AzureDocumentIntelligenceOCRConfig()
optional_params = config.map_ocr_params(
{"pages": [0, 1, 2], "features": ["keyValuePairs", "languages"]},
{},
"prebuilt-layout",
)
url = config.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="prebuilt-layout",
optional_params=optional_params,
)
assert "&pages=1,2,3" in url
assert "&features=keyValuePairs,languages" in url

View file

@ -14,6 +14,7 @@ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transfor
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
AmazonInvokeConfig,
)
from litellm.llms.bedrock.common_utils import BedrockError
@pytest.mark.parametrize(
@ -39,3 +40,145 @@ def test_transform_request_drops_stream_chunk_size(config, model):
)
assert "stream_chunk_size" not in json.dumps(request_body)
def test_validate_environment_maps_guardrail_config_to_invoke_headers():
"""The InvokeModel API takes the guardrail identifier/version/trace as
X-Amzn-Bedrock-* request headers, unlike Converse which takes them in the
body. https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InvokeModel.html"""
optional_params = {
"guardrailConfig": {
"guardrailIdentifier": "ff6ujrregl1q",
"guardrailVersion": "DRAFT",
"trace": "enabled",
},
"max_tokens": 10,
}
headers = AmazonInvokeConfig().validate_environment(
headers={},
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
optional_params=optional_params,
litellm_params={},
)
assert headers["X-Amzn-Bedrock-GuardrailIdentifier"] == "ff6ujrregl1q"
assert headers["X-Amzn-Bedrock-GuardrailVersion"] == "DRAFT"
assert headers["X-Amzn-Bedrock-Trace"] == "ENABLED"
assert "guardrailConfig" not in optional_params
def test_validate_environment_without_guardrail_config_leaves_headers_untouched():
headers = AmazonInvokeConfig().validate_environment(
headers={"foo": "bar"},
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
optional_params={"max_tokens": 10},
litellm_params={},
)
assert headers == {"foo": "bar"}
def test_validate_environment_skips_absent_guardrail_fields():
headers = AmazonInvokeConfig().validate_environment(
headers={},
model="amazon.titan-text-express-v1",
messages=[{"role": "user", "content": "hi"}],
optional_params={"guardrailConfig": {"guardrailIdentifier": "gr-id", "guardrailVersion": "1"}},
litellm_params={},
)
assert headers == {
"X-Amzn-Bedrock-GuardrailIdentifier": "gr-id",
"X-Amzn-Bedrock-GuardrailVersion": "1",
}
def test_validate_environment_does_not_clobber_explicit_guardrail_headers():
"""Users worked around the missing guardrailConfig support by passing the
AWS headers directly; an explicit header must keep winning over
guardrailConfig regardless of casing."""
headers = AmazonInvokeConfig().validate_environment(
headers={"x-amzn-bedrock-guardrailidentifier": "explicit-id"},
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
optional_params={
"guardrailConfig": {"guardrailIdentifier": "config-id", "guardrailVersion": "2"},
},
litellm_params={},
)
assert headers["x-amzn-bedrock-guardrailidentifier"] == "explicit-id"
assert "X-Amzn-Bedrock-GuardrailIdentifier" not in headers
assert headers["X-Amzn-Bedrock-GuardrailVersion"] == "2"
@pytest.mark.parametrize(
"bad_guardrail_config",
[
{"guardrailIdentifier": "gr-id", "trace": "verbose"},
{"guardrailIdentifier": ["gr-id"]},
"gr-id",
{},
{"trace": "enabled"},
],
)
def test_validate_environment_rejects_malformed_guardrail_config(bad_guardrail_config):
with pytest.raises(BedrockError) as excinfo:
AmazonInvokeConfig().validate_environment(
headers={},
model="anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
optional_params={"guardrailConfig": bad_guardrail_config},
litellm_params={},
)
assert excinfo.value.status_code == 400
assert "guardrailConfig" in str(excinfo.value)
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-3-sonnet-20240229-v1:0",
"amazon.titan-text-express-v1",
"mistral.mistral-7b-instruct-v0:2",
"meta.llama3-8b-instruct-v1:0",
],
)
def test_guardrail_config_flows_to_headers_not_request_body(model):
"""Mirrors the handler flow (validate_environment then transform_request):
guardrailConfig must end up in the signed headers and never leak into the
request body, where Bedrock rejects it as an extra input."""
config = AmazonInvokeConfig()
optional_params = {
"guardrailConfig": {
"guardrailIdentifier": "ff6ujrregl1q",
"guardrailVersion": "DRAFT",
"trace": "disabled",
},
"max_tokens": 10,
}
messages = [{"role": "user", "content": "hi"}]
headers = config.validate_environment(
headers={},
model=model,
messages=messages,
optional_params=optional_params,
litellm_params={},
)
request_body = config.transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params={},
headers=headers,
)
assert "guardrailConfig" not in json.dumps(request_body)
assert headers["X-Amzn-Bedrock-GuardrailIdentifier"] == "ff6ujrregl1q"
assert headers["X-Amzn-Bedrock-GuardrailVersion"] == "DRAFT"
assert headers["X-Amzn-Bedrock-Trace"] == "DISABLED"

View file

@ -611,6 +611,65 @@ def test_transform_request_helper_includes_anthropic_beta_and_tools():
assert fields["tools"][0]["type"] == "computer_20250124"
def test_parallel_tool_calls_config_kept_for_sonnet_5():
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
config = AmazonConverseConfig()
optional_params = config.map_openai_params(
model="anthropic.claude-sonnet-5",
non_default_params={"parallel_tool_calls": False},
optional_params={},
drop_params=False,
)
data = config._transform_request_helper(
model="anthropic.claude-sonnet-5",
system_content_blocks=[],
optional_params=optional_params,
messages=None,
)
assert data["additionalModelRequestFields"]["tool_choice"] == {
"disable_parallel_tool_use": True
}
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_parallel_tool_calls_config_dropped_for_ttl_only_model(
monkeypatch: pytest.MonkeyPatch,
):
model = "anthropic.claude-fable-5"
monkeypatch.setitem(
litellm.model_cost,
model,
{"cache_creation_input_token_cost_above_1hr": 2e-05},
)
config = AmazonConverseConfig()
optional_params = config.map_openai_params(
model=model,
non_default_params={"parallel_tool_calls": False},
optional_params={},
drop_params=False,
)
data = config._transform_request_helper(
model=model,
system_content_blocks=[],
optional_params=optional_params,
messages=None,
)
assert "tool_choice" not in data.get("additionalModelRequestFields", {})
def test_transform_response_with_computer_use_tool():
"""Test response transformation with computer use tool call."""
import httpx
@ -4130,6 +4189,42 @@ def test_parallel_tool_calls_newer_model_adds_disable_flag():
assert "parallel_tool_calls" not in request_data["additionalModelRequestFields"]
def test_parallel_tool_calls_flag_decoupled_from_ttl_pricing(monkeypatch):
"""
The disable_parallel_tool_use gate must read supports_parallel_tool_use_config,
not the 1h-TTL pricing field: a model carrying only the former still gets the flag.
"""
from litellm.llms.bedrock.common_utils import is_claude_4_5_on_bedrock
config = AmazonConverseConfig()
model = "anthropic.claude-parallel-tool-use-only"
monkeypatch.setitem(litellm.model_cost, model, {"supports_parallel_tool_use_config": True})
assert is_claude_4_5_on_bedrock(model) is False
messages = [{"role": "user", "content": "What's the weather in SF and NYC?"}]
optional_params = config.map_openai_params(
non_default_params={"parallel_tool_calls": False, "tools": _TOOL_PARAM},
optional_params={},
model=model,
drop_params=False,
)
request_data = config.transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={},
)
assert (
request_data["additionalModelRequestFields"]["tool_choice"][
"disable_parallel_tool_use"
]
is True
)
def test_parallel_tool_calls_older_model_drops_disable_flag():
"""Older Claude models (pre-4.5) must NOT receive disable_parallel_tool_use — Bedrock rejects it."""
config = AmazonConverseConfig()
@ -4554,6 +4649,154 @@ def test_cache_control_injection_tool_config_not_added_without_injection_point()
assert all("cachePoint" not in tool for tool in tools)
def test_cache_control_injection_tool_config_honors_ttl_for_supported_model():
"""
Regression test: cache_control_injection_points with location=tool_config
must honor the requested `control.ttl`, mirroring the message/system
cache_control behavior, instead of always emitting a bare
{"type": "default"} cachePoint with no ttl.
Forces the bundled local cost map so `is_claude_4_5_on_bedrock` (which
reads `cache_creation_input_token_cost_above_1hr` from litellm.model_cost)
sees this branch's pricing data rather than the network-fetched `main`
copy, which lacks it until merge.
"""
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
],
}
result = config._transform_request(
model="us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
assert tools[-1] == {"cachePoint": {"type": "default", "ttl": "1h"}}
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacking_own_pricing():
"""
Regression test: a regional pricing entry that omits
`cache_creation_input_token_cost_above_1hr` (e.g. `jp.anthropic.claude-opus-4-7`)
must not shadow the base model entry that carries it; the requested ttl
survives through the base-model fallback.
"""
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
assert "cache_creation_input_token_cost_above_1hr" not in litellm.model_cost["jp.anthropic.claude-opus-4-7"]
assert "cache_creation_input_token_cost_above_1hr" in litellm.model_cost["anthropic.claude-opus-4-7"]
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
],
}
result = config._transform_request(
model="jp.anthropic.claude-opus-4-7",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
assert tools[-1] == {"cachePoint": {"type": "default", "ttl": "1h"}}
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model():
"""
Models that don't support extended TTL caching (only Claude 4.5+ on
Bedrock does) must fall back to the default cachePoint with no ttl,
even if the caller requested one, matching message/system behavior.
"""
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
],
}
result = config._transform_request(
model="anthropic.claude-3-5-haiku-20241022-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
assert tools[-1] == {"cachePoint": {"type": "default"}}
def test_translate_response_format_json_schema_still_injects_tool():
"""
response_format with an explicit json_schema should still use the

View file

@ -492,7 +492,7 @@ def test_bedrock_invoke_messages_transform_converts_custom_tool_schema_type_to_o
assert result["tools"][0]["type"] == "custom"
def test_remove_ttl_from_cache_control_processes_tools():
def test_remove_ttl_from_cache_control_processes_tools(local_model_cost_map):
"""
Ensure _remove_ttl_from_cache_control also sanitizes cache_control on tools.
@ -538,7 +538,7 @@ def test_remove_ttl_from_cache_control_processes_tools():
assert "ttl" not in request["system"][0]["cache_control"]
def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5():
def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5(local_model_cost_map):
"""
For Claude 4.5+ models, ttl in ["5m", "1h"] should be preserved on tools,
just like it is for system and messages.
@ -564,7 +564,7 @@ def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5():
}
cfg._remove_ttl_from_cache_control(
request, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
request, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0"
)
# Both tools and system should preserve ttl for Claude 4.5

View file

@ -445,3 +445,31 @@ def test_explicit_invoke_route_does_not_match_async_invoke():
BedrockModelInfo._explicit_async_invoke_route(f"bedrock/{async_invoke_model}")
is True
)
def test_capability_lookups_fall_back_to_base_model_when_regional_entry_lacks_field(monkeypatch):
"""
Regression test: a regional model_cost entry without the capability field
must not shadow a base entry that has it (`get(model) or get(base)` used to
short-circuit on the truthy regional dict and drop the capability).
"""
import litellm
from litellm.llms.bedrock.common_utils import (
bedrock_converse_supports_parallel_tool_use_config,
is_claude_4_5_on_bedrock,
)
base = "anthropic.claude-fallback-test"
regional = f"eu.{base}"
monkeypatch.setitem(litellm.model_cost, regional, {"input_cost_per_token": 1e-06})
monkeypatch.setitem(
litellm.model_cost,
base,
{
"cache_creation_input_token_cost_above_1hr": 1e-05,
"supports_parallel_tool_use_config": True,
},
)
assert is_claude_4_5_on_bedrock(regional) is True
assert bedrock_converse_supports_parallel_tool_use_config(regional) is True

View file

@ -0,0 +1,207 @@
from unittest.mock import patch
from litellm.llms.tencent.chat.transformation import TencentChatConfig
def test_supported_openai_params_includes_thinking_and_reasoning_effort():
config = TencentChatConfig()
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
params = config.get_supported_openai_params(model="tencent/deepseek-v4-pro")
assert "thinking" in params
assert "reasoning_effort" in params
assert "stream" in params
assert "temperature" in params
def test_supported_openai_params_excludes_thinking_without_reasoning_support():
config = TencentChatConfig()
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=False,
):
params = config.get_supported_openai_params(model="tencent/non-reasoning-model")
assert "thinking" not in params
assert "reasoning_effort" not in params
assert "stream" in params
def test_map_openai_params_passes_thinking_dict_through():
config = TencentChatConfig()
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
result = config.map_openai_params(
non_default_params={"thinking": {"type": "enabled", "budget_tokens": 1024}},
optional_params={},
model="tencent/deepseek-v4-pro",
drop_params=False,
)
assert result["thinking"] == {"type": "enabled", "budget_tokens": 1024}
def test_map_openai_params_converts_reasoning_effort_to_thinking():
config = TencentChatConfig()
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
result = config.map_openai_params(
non_default_params={"reasoning_effort": "medium"},
optional_params={},
model="tencent/deepseek-v4-pro",
drop_params=False,
)
assert result["thinking"] == {"type": "enabled"}
def test_map_openai_params_drops_none_reasoning_effort():
config = TencentChatConfig()
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
result = config.map_openai_params(
non_default_params={"reasoning_effort": "none"},
optional_params={},
model="tencent/deepseek-v4-pro",
drop_params=False,
)
assert "thinking" not in result
assert "reasoning_effort" not in result
def test_map_openai_params_thinking_priority_over_reasoning_effort():
config = TencentChatConfig()
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
result = config.map_openai_params(
non_default_params={
"thinking": {"type": "enabled", "budget_tokens": 2048},
"reasoning_effort": "high",
},
optional_params={},
model="tencent/deepseek-v4-pro",
drop_params=False,
)
assert result["thinking"] == {"type": "enabled", "budget_tokens": 2048}
def test_map_openai_params_extracts_thinking_and_effort_from_optional_params():
config = TencentChatConfig()
result = config.map_openai_params(
non_default_params={},
optional_params={"thinking": {"type": "enabled"}, "reasoning_effort": "medium"},
model="tencent/deepseek-v4-pro",
drop_params=False,
)
assert "thinking" in result
assert "reasoning_effort" not in result
def test_get_complete_url_default():
config = TencentChatConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="tencent/deepseek-v4-pro",
optional_params={},
litellm_params={},
)
assert url == "https://tokenhub-intl.tencentcloudmaas.com/v1/chat/completions"
def test_get_complete_url_strips_trailing_slash():
config = TencentChatConfig()
url = config.get_complete_url(
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1/",
api_key=None,
model="tencent/deepseek-v4-pro",
optional_params={},
litellm_params={},
)
assert url == "https://tokenhub-intl.tencentcloudmaas.com/v1/chat/completions"
def test_get_complete_url_custom_base_preserves_v1():
config = TencentChatConfig()
url = config.get_complete_url(
api_base="https://tokenhub.tencentcloudmaas.com/v1",
api_key=None,
model="tencent/deepseek-v4-pro",
optional_params={},
litellm_params={},
)
assert url == "https://tokenhub.tencentcloudmaas.com/v1/chat/completions"
def test_get_complete_url_adds_v1_to_custom_base():
config = TencentChatConfig()
url = config.get_complete_url(
api_base="https://tokenhub.tencentcloudmaas.com",
api_key=None,
model="tencent/deepseek-v4-pro",
optional_params={},
litellm_params={},
)
assert url == "https://tokenhub.tencentcloudmaas.com/v1/chat/completions"
def test_get_complete_url_does_not_append_to_full_url():
config = TencentChatConfig()
url = config.get_complete_url(
api_base="https://tokenhub.tencentcloudmaas.com/v1/chat/completions",
api_key=None,
model="tencent/deepseek-v4-pro",
optional_params={},
litellm_params={},
)
assert url == "https://tokenhub.tencentcloudmaas.com/v1/chat/completions"
def test_provider_info_falls_back_to_default_base():
config = TencentChatConfig()
with patch("litellm.llms.tencent.chat.transformation.get_secret_str", return_value=None):
api_base, api_key = config._get_openai_compatible_provider_info(api_base=None, api_key="sk-arg")
assert api_base == "https://tokenhub-intl.tencentcloudmaas.com/v1"
assert api_key == "sk-arg"
def test_provider_info_reads_env_secrets():
config = TencentChatConfig()
secrets = {"TENCENT_API_BASE": "https://env.tencent/v1", "TENCENT_API_KEY": "sk-env"}
with patch(
"litellm.llms.tencent.chat.transformation.get_secret_str",
side_effect=lambda key: secrets.get(key),
):
api_base, api_key = config._get_openai_compatible_provider_info(api_base=None, api_key=None)
assert api_base == "https://env.tencent/v1"
assert api_key == "sk-env"

View file

@ -0,0 +1,173 @@
import litellm
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
from litellm.llms.tencent.messages.transformation import (
TencentAnthropicMessagesConfig,
)
from litellm.utils import ProviderConfigManager
def test_tencent_provider_uses_anthropic_messages_config():
config = ProviderConfigManager.get_provider_anthropic_messages_config(
model="deepseek-v4-pro",
provider=litellm.LlmProviders.TENCENT,
)
assert isinstance(config, TencentAnthropicMessagesConfig)
assert config.custom_llm_provider == "tencent"
def test_anthropic_provider_keeps_default_config_for_tencent_named_model():
config = ProviderConfigManager.get_provider_anthropic_messages_config(
model="deepseek-v4-pro",
provider=litellm.LlmProviders.ANTHROPIC,
)
assert isinstance(config, AnthropicMessagesConfig)
assert not isinstance(config, TencentAnthropicMessagesConfig)
def test_strips_billing_metadata():
config = TencentAnthropicMessagesConfig()
assert config.should_strip_billing_metadata() is True
def test_get_api_base_default():
config = TencentAnthropicMessagesConfig()
assert config.get_api_base() == "https://tokenhub-intl.tencentcloudmaas.com"
def test_get_api_base_from_arg():
config = TencentAnthropicMessagesConfig()
assert config.get_api_base(api_base="https://custom.example.com") == "https://custom.example.com"
def test_messages_url_default():
config = TencentAnthropicMessagesConfig()
assert (
config.get_complete_url(
api_base=None,
api_key=None,
model="deepseek-v4-pro",
optional_params={},
litellm_params={},
)
== "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
)
def test_messages_url_with_base_ending_in_v1():
config = TencentAnthropicMessagesConfig()
assert (
config.get_complete_url(
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1",
api_key=None,
model="deepseek-v4-pro",
optional_params={},
litellm_params={},
)
== "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
)
def test_messages_url_with_base_ending_in_v1_messages():
config = TencentAnthropicMessagesConfig()
url = config.get_complete_url(
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1/messages",
api_key=None,
model="deepseek-v4-pro",
optional_params={},
litellm_params={},
)
assert url == "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
def test_messages_url_with_base_ending_in_v1_chat_completions():
config = TencentAnthropicMessagesConfig()
assert (
config.get_complete_url(
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1/chat/completions",
api_key=None,
model="deepseek-v4-pro",
optional_params={},
litellm_params={},
)
== "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
)
def test_messages_url_with_custom_base_no_v1():
config = TencentAnthropicMessagesConfig()
assert (
config.get_complete_url(
api_base="https://tokenhub.tencentcloudmaas.com",
api_key=None,
model="deepseek-v4-pro",
optional_params={},
litellm_params={},
)
== "https://tokenhub.tencentcloudmaas.com/v1/messages"
)
def test_validate_environment_sets_headers():
config = TencentAnthropicMessagesConfig()
headers, api_base = config.validate_anthropic_messages_environment(
headers={},
model="deepseek-v4-pro",
messages=[],
optional_params={},
litellm_params={},
api_key="sk-tencent-key",
api_base="https://custom.test",
)
assert headers["x-api-key"] == "sk-tencent-key"
assert headers["anthropic-version"] == "2023-06-01"
assert headers["content-type"] == "application/json"
assert api_base == "https://custom.test"
def test_validate_environment_injects_anthropic_beta_headers():
config = TencentAnthropicMessagesConfig()
headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="deepseek-v4-pro",
messages=[],
optional_params={"speed": "fast"},
litellm_params={},
api_key="sk-tencent-key",
api_base=None,
)
assert "anthropic-beta" in headers
def test_validate_environment_preserves_existing_headers():
config = TencentAnthropicMessagesConfig()
headers, _ = config.validate_anthropic_messages_environment(
headers={"authorization": "Bearer existing", "anthropic-version": "2024-01-01"},
model="deepseek-v4-pro",
messages=[],
optional_params={},
litellm_params={},
api_key="sk-tencent-key",
api_base=None,
)
assert headers["authorization"] == "Bearer existing"
assert headers["anthropic-version"] == "2024-01-01"
assert "x-api-key" not in headers

View file

@ -0,0 +1,41 @@
import pytest
import litellm
from litellm.llms.tencent.cost_calculator import cost_per_token
from litellm.types.utils import Usage
@pytest.fixture
def local_model_cost_map(monkeypatch):
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
def test_cost_per_token_uses_tencent_model_pricing(local_model_cost_map):
usage = Usage(prompt_tokens=1000, completion_tokens=2000, total_tokens=3000)
prompt_cost, completion_cost = cost_per_token(model="tencent/deepseek-v4-pro", usage=usage)
assert prompt_cost == pytest.approx(1000 * 4.35e-07)
assert completion_cost == pytest.approx(2000 * 8.7e-07)
def test_top_level_dispatcher_routes_tencent_to_wrapper(local_model_cost_map):
from litellm.cost_calculator import cost_per_token as dispatch_cost_per_token
prompt_cost, completion_cost = dispatch_cost_per_token(
model="tencent/deepseek-v4-pro",
prompt_tokens=1000,
completion_tokens=1000,
custom_llm_provider="tencent",
)
assert prompt_cost == pytest.approx(1000 * 4.35e-07)
assert completion_cost == pytest.approx(1000 * 8.7e-07)

View file

@ -0,0 +1,184 @@
import asyncio
from typing import Dict, Optional
import pytest
from unittest.mock import patch
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
from litellm.types.mcp import MCPAuth, MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
HOLD_SECONDS = 0.1
class _ConcurrencyTracker:
"""Records how many call_tool invocations are simultaneously in flight."""
def __init__(self) -> None:
self.current_by_server: Dict[str, int] = {}
self.peak_by_server: Dict[str, int] = {}
self.global_current = 0
self.global_peak = 0
def enter(self, server_id: str) -> None:
self.current_by_server[server_id] = self.current_by_server.get(server_id, 0) + 1
self.peak_by_server[server_id] = max(self.peak_by_server.get(server_id, 0), self.current_by_server[server_id])
self.global_current += 1
self.global_peak = max(self.global_peak, self.global_current)
def exit(self, server_id: str) -> None:
self.current_by_server[server_id] -= 1
self.global_current -= 1
def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPServer:
return MCPServer(
server_id=server_id,
name=server_id,
server_name=server_id,
url="https://example.com",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
max_concurrent_requests=max_concurrent_requests,
)
def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker):
async def fake_create_mcp_client(server, **kwargs):
class _ProbeClient:
async def call_tool(self, params, host_progress_callback=None):
tracker.enter(server.server_id)
try:
await asyncio.sleep(HOLD_SECONDS)
return "ok"
finally:
tracker.exit(server.server_id)
return _ProbeClient()
return patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client)
async def _fire(manager: MCPServerManager, server: MCPServer, n: int) -> None:
await asyncio.gather(
*[
manager._call_regular_mcp_tool(
mcp_server=server,
original_tool_name="tool",
arguments={},
tasks=[],
mcp_auth_header=None,
mcp_server_auth_headers=None,
oauth2_headers=None,
raw_headers=None,
proxy_logging_obj=None,
)
for _ in range(n)
]
)
@pytest.mark.asyncio
async def test_max_concurrent_requests_caps_in_flight_tool_calls():
"""A configured cap of 2 must never let more than 2 calls hit one server at once."""
manager = MCPServerManager()
tracker = _ConcurrencyTracker()
server = _make_server("srv-limited", max_concurrent_requests=2)
with _patch_client_with_tracker(manager, tracker):
await _fire(manager, server, n=8)
assert tracker.peak_by_server["srv-limited"] == 2
@pytest.mark.asyncio
async def test_unset_limit_allows_unbounded_concurrency():
"""With no cap, all calls run concurrently (backward-compatible default)."""
manager = MCPServerManager()
tracker = _ConcurrencyTracker()
server = _make_server("srv-unbounded", max_concurrent_requests=None)
with _patch_client_with_tracker(manager, tracker):
await _fire(manager, server, n=6)
assert tracker.peak_by_server["srv-unbounded"] == 6
@pytest.mark.asyncio
async def test_non_positive_limit_is_treated_as_unlimited():
"""A cap of 0 must not deadlock; it means unlimited, not a zero-permit semaphore."""
manager = MCPServerManager()
tracker = _ConcurrencyTracker()
server = _make_server("srv-zero", max_concurrent_requests=0)
with _patch_client_with_tracker(manager, tracker):
await asyncio.wait_for(_fire(manager, server, n=5), timeout=5)
assert tracker.peak_by_server["srv-zero"] == 5
@pytest.mark.asyncio
async def test_limit_is_scoped_per_server():
"""Each server gets its own limiter; one server's cap must not throttle another."""
manager = MCPServerManager()
tracker = _ConcurrencyTracker()
server_a = _make_server("srv-a", max_concurrent_requests=1)
server_b = _make_server("srv-b", max_concurrent_requests=1)
with _patch_client_with_tracker(manager, tracker):
await asyncio.gather(
_fire(manager, server_a, n=3),
_fire(manager, server_b, n=3),
)
assert tracker.peak_by_server["srv-a"] == 1
assert tracker.peak_by_server["srv-b"] == 1
assert tracker.global_peak == 2
@pytest.mark.asyncio
async def test_openapi_backed_server_also_respects_the_cap():
"""OpenAPI (spec_path) servers dispatch through a different handler; the cap
must apply there too, not only on the regular MCP client path."""
manager = MCPServerManager()
tracker = _ConcurrencyTracker()
server = _make_server("srv-openapi", max_concurrent_requests=2)
server.spec_path = "/fake/openapi.json"
async def fake_openapi_handler(mcp_server, name, arguments):
tracker.enter(mcp_server.server_id)
try:
await asyncio.sleep(HOLD_SECONDS)
return "ok"
finally:
tracker.exit(mcp_server.server_id)
with (
patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server),
patch.object(manager, "_call_openapi_tool_handler", side_effect=fake_openapi_handler),
):
await asyncio.gather(
*[manager.call_tool(server_name="srv-openapi", name="tool", arguments={}) for _ in range(6)]
)
assert tracker.peak_by_server["srv-openapi"] == 2
def test_semaphore_is_reused_per_server_and_distinct_across_servers():
manager = MCPServerManager()
server_a = _make_server("srv-a", max_concurrent_requests=3)
server_b = _make_server("srv-b", max_concurrent_requests=3)
sem_a_first = manager._get_call_semaphore(server_a)
sem_a_second = manager._get_call_semaphore(server_a)
sem_b = manager._get_call_semaphore(server_b)
assert sem_a_first is sem_a_second
assert sem_a_first is not sem_b
def test_no_semaphore_created_when_limit_absent():
manager = MCPServerManager()
server = _make_server("srv-none", max_concurrent_requests=None)
assert manager._get_call_semaphore(server) is None

View file

@ -76,9 +76,7 @@ async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(
mock_client, passthrough_server.name, server=passthrough_server
)
await manager._fetch_tools_with_timeout(mock_client, passthrough_server.name)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == (
@ -113,9 +111,7 @@ async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(
mock_client, delegated_server.name, server=delegated_server
)
await manager._fetch_tools_with_timeout(mock_client, delegated_server.name)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == (
@ -126,7 +122,10 @@ async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
@pytest.mark.asyncio
async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior():
async def test_fetch_tools_from_client_credentials_oauth2_surfaces_upstream_401():
"""The auth_type carve-out was removed: a client_credentials (M2M) server now
surfaces an upstream 401 as MCPUpstreamAuthError too, instead of swallowing it
to an empty list, so single-server routes can return a 401 challenge."""
manager = MCPServerManager()
m2m_server = MCPServer(
server_id="oauth-m2m",
@ -150,12 +149,12 @@ async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
tools = await manager._fetch_tools_with_timeout(
mock_client, m2m_server.name, server=m2m_server
)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(mock_client, m2m_server.name)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
assert exc_info.value.status_code == 401
assert exc_info.value.server_name == "m2m_docs"
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
@pytest.mark.asyncio
@ -176,9 +175,7 @@ async def test_fetch_tools_from_passthrough_returns_tools_on_success():
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(return_value=[tool])
tools = await manager._fetch_tools_with_timeout(
mock_client, passthrough_server.name, server=passthrough_server
)
tools = await manager._fetch_tools_with_timeout(mock_client, passthrough_server.name)
assert tools == [tool]
@ -238,8 +235,11 @@ def test_to_http_exception_skips_challenge_for_non_401_status():
@pytest.mark.asyncio
async def test_fetch_tools_from_gateway_managed_swallows_errors():
"""Regression guard: non-pass-through servers keep returning [] on errors."""
async def test_fetch_tools_from_gateway_managed_surfaces_upstream_401():
"""An oauth2 server that is neither pass-through nor delegate now surfaces an
upstream 401 as MCPUpstreamAuthError as well; the auth_type carve-out that
swallowed it to [] was removed. A missing upstream WWW-Authenticate is carried
through as None (the single-server route fabricates one from the gateway URL)."""
manager = MCPServerManager()
oauth2_server = MCPServer(
server_id="o1",
@ -260,11 +260,13 @@ async def test_fetch_tools_from_gateway_managed_swallows_errors():
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
tools = await manager._fetch_tools_with_timeout(
mock_client, oauth2_server.name, server=oauth2_server
)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(mock_client, oauth2_server.name)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate is None
assert exc_info.value.server_name == "keycloak_whoami"
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
def _http_server(server_id: str, name: str, **kwargs) -> MCPServer:

View file

@ -5228,5 +5228,139 @@ class TestCreateMcpClientV2Graft:
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
def _upstream_status_error(status_code: int, challenge: str) -> httpx.HTTPStatusError:
request = httpx.Request("POST", "https://upstream.example/mcp")
response = httpx.Response(
status_code,
headers={"WWW-Authenticate": challenge},
request=request,
)
return httpx.HTTPStatusError(
"upstream rejected token", request=request, response=response
)
class TestMCPToolsListAuthSurfacing:
"""Regression: MCP tools/list 401 auth failures must surface as MCPUpstreamAuthError.
Previously a missing/expired per-user OAuth token, or an upstream 401 for any
non-carveout auth_type, was swallowed to an empty tool list, so a single-server
client saw a 200 with no tools instead of a 401 challenge. The listing helpers
now raise MCPUpstreamAuthError on a 401 regardless of auth_type; the single-server
routes turn it into a 401 + WWW-Authenticate while the aggregator absorbs it to an
empty list. Only a 401 challenges; a 403 (forbidden) degrades like any other error.
"""
@pytest.mark.asyncio
async def test_fetch_tools_with_timeout_surfaces_upstream_401(self):
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
manager = MCPServerManager()
challenge = 'Bearer resource_metadata="https://upstream.example/.well-known/oauth-protected-resource"'
client = MagicMock()
client.list_tools = AsyncMock(side_effect=_upstream_status_error(401, challenge))
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(client, "static-key-server")
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == challenge
assert exc_info.value.server_name == "static-key-server"
@pytest.mark.asyncio
async def test_fetch_tools_with_timeout_absorbs_upstream_403(self):
"""Only a 401 drives the re-auth challenge. A 403 (authenticated but
forbidden, e.g. insufficient scope) is not a re-auth signal, so even
with a WWW-Authenticate header it degrades to an empty list rather than
surfacing a challenge."""
manager = MCPServerManager()
challenge = 'Bearer error="insufficient_scope", scope="read:tools"'
client = MagicMock()
client.list_tools = AsyncMock(side_effect=_upstream_status_error(403, challenge))
assert await manager._fetch_tools_with_timeout(client, "forbidden-server") == []
@pytest.mark.asyncio
async def test_fetch_tools_with_timeout_returns_empty_on_non_auth_error(self):
manager = MCPServerManager()
client = MagicMock()
client.list_tools = AsyncMock(side_effect=RuntimeError("upstream 500"))
assert await manager._fetch_tools_with_timeout(client, "srv") == []
@pytest.mark.asyncio
async def test_get_tools_from_server_surfaces_unusable_user_token(self):
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
manager = MCPServerManager()
server = MCPServer(
server_id="oauth-srv", name="oauth-srv", transport=MCPTransport.http
)
challenge = 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/oauth-srv"'
manager._create_mcp_client = AsyncMock(
side_effect=HTTPException(
status_code=401,
detail="Unauthorized",
headers={"WWW-Authenticate": challenge},
)
)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._get_tools_from_server(server)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == challenge
assert exc_info.value.server_name == "oauth-srv"
@pytest.mark.asyncio
async def test_get_tools_from_server_absorbs_non_challenge_http_error(self):
manager = MCPServerManager()
server = MCPServer(
server_id="stdio-srv", name="stdio-srv", transport=MCPTransport.http
)
manager._create_mcp_client = AsyncMock(
side_effect=HTTPException(
status_code=403,
detail="MCP stdio command 'foo' is not in the allowlist",
)
)
assert await manager._get_tools_from_server(server) == []
@pytest.mark.asyncio
async def test_aggregate_list_tools_absorbs_unauthenticated_server(self):
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
manager = MCPServerManager()
good = MCPServer(server_id="good", name="good", transport=MCPTransport.http)
bad = MCPServer(server_id="bad", name="bad", transport=MCPTransport.http)
manager.get_allowed_mcp_servers = AsyncMock(return_value=["good", "bad"])
manager.get_mcp_server_by_id = MagicMock(
side_effect=lambda server_id: {"good": good, "bad": bad}.get(server_id)
)
good_tool = MCPTool(name="good-do_thing", description="do thing", inputSchema={})
async def fake_get_tools(server, **kwargs):
if server.server_id == "bad":
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate='Bearer realm="x"',
server_name="bad",
)
return [good_tool]
manager._get_tools_from_server = fake_get_tools
result = await manager.list_tools()
assert [t.name for t in result] == ["good-do_thing"]
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -843,6 +843,76 @@ class TestListToolsRestAPI:
assert exc_info.value.status_code == upstream_status
assert exc_info.value.headers == {"www-authenticate": challenge}
async def test_aggregate_list_absorbs_one_server_auth_failure(self, monkeypatch):
"""The multi-server aggregate listing degrades a server whose upstream
rejects auth to an empty contribution and still returns the healthy
server's tools with a 200, rather than surfacing a 401."""
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
class StubServer:
def __init__(self, name):
self.alias = name
self.server_name = name
self.name = name
self.allowed_tools = None
self.mcp_info = {"server_name": name}
self.available_on_public_internet = True
good = StubServer("good")
bad = StubServer("bad")
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return ["good", "bad"]
async def fake_get_tools(server, *args, **kwargs):
if server.server_name == "bad":
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate='Bearer realm="x"',
server_name="bad",
)
return ["good-tool"]
monkeypatch.setattr(
rest_endpoints,
"build_effective_auth_contexts",
fake_contexts,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: {"good": good, "bad": bad}.get(server_id),
raising=False,
)
monkeypatch.setattr(
rest_endpoints,
"_get_tools_for_single_server",
fake_get_tools,
raising=False,
)
request = _build_request(path="/mcp-rest/tools/list", method="GET")
result = await rest_endpoints.list_tool_rest_api(
request,
server_id=None,
user_api_key_dict=UserAPIKeyAuth(),
)
assert result["tools"] == ["good-tool"]
assert result["error"] is None
async def test_name_resolution_finds_server_by_uuid(self, monkeypatch):
"""When server_id is a name string, it should be resolved to its UUID
and used for the tools lookup when the UUID is in allowed_server_ids."""

View file

@ -58,15 +58,71 @@ class TestAnthropicEndpoints(unittest.TestCase):
self.assertEqual(result, expected_result)
# Assert safe_dumps was called for dictionary objects
mock_safe_dumps.assert_any_call(
{"type": "message_start", "message": {"id": "msg_123"}}
mock_safe_dumps.assert_any_call({"type": "message_start", "message": {"id": "msg_123"}})
mock_safe_dumps.assert_any_call({"type": "content_block_delta", "delta": {"text": "more data"}})
assert mock_safe_dumps.call_count == 2 # Called twice, once for each dict object
class TestBlockedResponseUsage:
"""Blocked responses report the blocked LLM response's real usage."""
def test_uses_original_response_usage(self):
from litellm.proxy.anthropic_endpoints.endpoints import _blocked_response_usage
# original_response is the AnthropicMessagesResponse the LLM produced
# before the guardrail blocked it; its usage is real.
original = {"usage": {"input_tokens": 31, "output_tokens": 9}}
assert _blocked_response_usage(original) == {
"input_tokens": 31,
"output_tokens": 9,
}
def test_zero_usage_when_no_original_response(self):
from litellm.proxy.anthropic_endpoints.endpoints import _blocked_response_usage
# Pre-call blocks never invoked the LLM -> nothing consumed.
assert _blocked_response_usage(None) == {
"input_tokens": 0,
"output_tokens": 0,
}
@pytest.mark.asyncio
async def test_blocked_endpoint_response_carries_original_usage(self):
"""The /v1/messages block handler reports the blocked response's real
usage, carried on ModifyResponseException.original_response."""
from unittest.mock import AsyncMock, MagicMock
import litellm.proxy.anthropic_endpoints.endpoints as ep
import litellm.proxy.proxy_server as proxy_server
from litellm.integrations.custom_guardrail import ModifyResponseException
exc = ModifyResponseException(
message="blocked by guardrail",
model="claude-3-5-sonnet-20240620",
request_data={"messages": [{"role": "user", "content": "hi"}]},
guardrail_name="rubrik",
original_response={"usage": {"input_tokens": 12, "output_tokens": 5}},
)
mock_safe_dumps.assert_any_call(
{"type": "content_block_delta", "delta": {"text": "more data"}}
)
assert (
mock_safe_dumps.call_count == 2
) # Called twice, once for each dict object
with (
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})),
patch.object(
ep.ProxyBaseLLMRequestProcessing,
"base_process_llm_request",
new=AsyncMock(side_effect=exc),
),
patch.object(proxy_server, "proxy_logging_obj") as mock_logging,
):
mock_logging.post_call_failure_hook = AsyncMock()
response = await ep.anthropic_response(
fastapi_response=MagicMock(),
request=MagicMock(),
user_api_key_dict=MagicMock(),
)
assert response["content"][0]["text"] == "blocked by guardrail"
assert response["usage"] == {"input_tokens": 12, "output_tokens": 5}
mock_logging.post_call_failure_hook.assert_awaited_once()
class TestEventLoggingBatchEndpoint:
@ -159,9 +215,7 @@ class TestStripTotalTokens(unittest.TestCase):
# SimpleNamespace mimics the .usage attribute access pattern; the
# helper's contract: if .usage is dict-shaped, strip total_tokens.
response = SimpleNamespace(
usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150}
)
response = SimpleNamespace(usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150})
_strip_total_tokens_from_anthropic_response(response)
assert "total_tokens" not in response.usage
assert response.usage == {"input_tokens": 100, "output_tokens": 50}

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