Merge remote-tracking branch 'upstream/litellm_internal_staging' into deepkeep-as-internal

This commit is contained in:
Yaniv Israel 2026-07-01 13:59:29 +03:00
commit c7685016d0
135 changed files with 12304 additions and 1269 deletions

View file

@ -21,7 +21,7 @@ concurrency:
jobs:
benchmarks:
runs-on: ubuntu-latest
runs-on: ubuntu-24.04
timeout-minutes: 15
steps:
@ -48,6 +48,8 @@ jobs:
uv run --frozen --no-default-groups
--with pytest==8.3.5
--with pytest-codspeed==4.3.0
--with "mcp>=1.26.0,<2.0"
--with "a2a-sdk>=1.1.0,<2.0"
pytest
-p pytest_codspeed.plugin
tests/benchmarks/

View file

@ -1,8 +1,7 @@
Do not write comments unless they are absolutely necessary to explain some very complex business logic. Please clean up if there are comments that are not absolutely necessary. Do not remove comments that are unrelated to the addition of the code of this PR
Explanation: code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive to the reader, while being both easy to maintain and high performance
Do not write any comments (existing comments can stay) unless explicitly asked to in a user (not system) prompt
Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in:
- correct
- secure
- performant
@ -34,7 +33,7 @@ If you ever make public-facing PR descriptions, comments, issues, commit message
Don't hesitate to use values in .env to get needed API keys and other secrets, as long as you never add them to conversation history, commit them, or include them in GitHub issues / PRs
Run tests, format your code, and lint your code before each commit
Run tests before you commit. Also, run `make pre-commit` right before each commit, which generates types (as needed) and formats/lints your code. Any errors found must be fixed
When you fix violations gated by `ruff-strict-budget.json` or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom
@ -42,7 +41,7 @@ If you're trying to create a new function that relies on untyped stuff, instead
If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason
Ask to commit and push your work when you're done (or if you're confident that your code is good and works, just do it)
Commit and push your work when you're done without asking
When you must use real LLM models to, for example, write e2e tests, write a QA runbook, etc., make sure to use the latest models (doesn't have to be smartest, can also be a modern small, fast one. No strong preference for smart vs fast here, just use something modern) as of the year and month of the current date. Do a web search as necessary to figure that out
@ -70,6 +69,7 @@ Follow these coding conventions for new/updated code (a three-line fix in a lega
- No monster files or god objects
- No file sprawl: deliberate file and folder structure
- Standard over hand-rolled: use the official SDK or a library where one exists; where none does, follow industry standards instead of inventing local conventions
- API-fragmentation-aware: when logic must branch on which API surface produced or consumes data (e.g. chat completions vs Anthropic Messages vs Responses API shapes), proactively look for an existing shared helper (e.g. `litellm_core_utils/prompt_templates/factory.py`) before writing per-surface parsing in the new module; if none exists, add one there instead of duplicating the same format-detection logic in every new guardrail/integration
Follow conventional commits for commit names and PR titles

View file

@ -8,7 +8,8 @@
lint-basedpyright lint-basedpyright-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
install-dev install-proxy-dev install-test-deps install-hooks \
install-helm-unittest check-circular-imports check-import-safety
install-helm-unittest check-circular-imports check-import-safety pre-commit \
lint-install lint-fetch-base
# Default target
help:
@ -20,6 +21,7 @@ help:
@echo " make install-test-deps - Install the full local test environment"
@echo " make install-helm-unittest - Install helm unittest plugin"
@echo " make install-hooks - Install git hooks (Conventional Commits + Branches)"
@echo " make pre-commit - Run CI-equivalent lint on staged files (run before committing)"
@echo " make format - Apply ruff format code formatting"
@echo " make format-check - Check ruff format code formatting (matches CI)"
@echo " make lint - Run all linting (Ruff, basedpyright, format check, circular imports, import safety)"
@ -56,8 +58,11 @@ info:
@echo "UV: $(UV)"
# Installation targets
# --inexact: sync the locked deps without pruning anything already installed, so running
# a lint/format target doesn't tear the proxy extras (prisma, websockets, ...) out from
# under a dev's venv (CI installs its own env per job, so it is unaffected by this).
install-dev:
$(UV) sync --frozen
$(UV) sync --inexact --frozen
install-proxy-dev:
$(UV) sync --frozen --group proxy-dev --extra proxy
@ -90,6 +95,31 @@ format: install-dev
format-check: install-dev
cd litellm && $(UV_RUN) ruff format --check --exclude '/enterprise/' . && cd ..
# Single fetch of the PR base so the delta-based gates below share one network round
# trip instead of each re-fetching when chained from `lint`.
lint-fetch-base:
git fetch origin litellm_internal_staging 2>/dev/null || true
# Mirror test-linting.yml's lint job environment: the proxy-dev group plus a generated
# Prisma client, so basedpyright resolves the same modules CI does (without the generated
# client the DB wrappers typed against it degrade to Unknown, drifting the budget from
# CI's). --inexact tops up the venv instead of pruning the proxy extras gen:api and the
# running proxy need.
lint-install:
$(UV) sync --inexact --frozen --group proxy-dev
$(UV_RUN) prisma generate --schema litellm/proxy/schema.prisma
# 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
@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."; \
else \
echo "$$files" | xargs $(UV_RUN) ruff format --check --exclude '/enterprise/'; \
fi
# Linting targets
lint-ruff: install-dev
cd litellm && $(UV_RUN) ruff check . && cd ..
@ -126,9 +156,14 @@ 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-basedpyright: install-dev lint-fetch-base
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py
# 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
$(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging
lint-basedpyright-budget-update: install-dev
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update
@ -142,8 +177,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
git fetch origin litellm_internal_staging
lint-gate: install-dev lint-fetch-base
$(UV_RUN) python scripts/ruff_strict_gate.py --base origin/litellm_internal_staging
lint-ruff-budget-update: install-dev
@ -158,12 +192,25 @@ check-circular-imports: install-dev
check-import-safety: install-dev
@$(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 (matches test-linting.yml workflow)
lint: format-check lint-ruff lint-basedpyright check-circular-imports check-import-safety lint-ruff-budget
# Combined linting, isomorphic to test-linting.yml's lint job so a local pass means a
# green CI lint: it installs the same env (proxy-dev + generated Prisma client) and then
# 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
# Faster linting for local development (only checks changed code)
lint-dev: lint-format-changed check-circular-imports check-import-safety
# Run the gating CI checks against your staged files right before committing. Mirrors
# test-linting.yml (Python), test-litellm-ui-build.yml's frontend-lint (dashboard), and
# check-ui-api-types.yml (API-type drift), skipping any whose files you didn't stage.
# Not auto-installed as a git hook so it never slows an unrelated human commit.
pre-commit:
./scripts/pre_commit_lint.sh
# Testing targets
test: install-test-deps
$(UV_RUN) pytest tests/

View file

@ -239,6 +239,7 @@ class BaseEmailLogger(CustomLogger):
max_budget_info=max_budget_info,
base_url=email_params.base_url,
email_support_contact=email_params.support_contact,
email_footer=email_params.signature,
)
await self.send_email(
from_email=self.DEFAULT_LITELLM_EMAIL,
@ -311,6 +312,7 @@ class BaseEmailLogger(CustomLogger):
max_budget_info=max_budget_info,
base_url=email_params.base_url,
email_support_contact=email_params.support_contact,
email_footer=email_params.signature,
)
# Send email to all recipients
@ -381,6 +383,7 @@ class BaseEmailLogger(CustomLogger):
alert_threshold=alert_threshold_str,
base_url=email_params.base_url,
email_support_contact=email_params.support_contact,
email_footer=email_params.signature,
)
await self.send_email(
from_email=self.DEFAULT_LITELLM_EMAIL,
@ -405,6 +408,7 @@ class BaseEmailLogger(CustomLogger):
alert_threshold=alert_threshold_str,
base_url=email_params.base_url,
email_support_contact=email_params.support_contact,
email_footer=email_params.signature,
)
await self.send_email(
from_email=self.DEFAULT_LITELLM_EMAIL,

View file

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

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN "mcp_tool_search_enabled" BOOLEAN;

View file

@ -279,6 +279,7 @@ model LiteLLM_ObjectPermissionTable {
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
mcp_tool_search_enabled Boolean?
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -23,7 +23,11 @@ from litellm._redis_credential_provider import (
GCPIAMCredentialProvider,
_generate_gcp_iam_access_token,
)
from litellm.constants import REDIS_CONNECTION_POOL_TIMEOUT, REDIS_SOCKET_TIMEOUT
from litellm.constants import (
REDIS_CLUSTER_HEALTH_CHECK_INTERVAL,
REDIS_CONNECTION_POOL_TIMEOUT,
REDIS_SOCKET_TIMEOUT,
)
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from ._logging import verbose_logger
@ -102,6 +106,8 @@ def _get_redis_cluster_kwargs(client=None):
"max_connections",
"socket_timeout",
"socket_connect_timeout",
"health_check_interval",
"socket_keepalive",
}
return available_args
@ -579,6 +585,13 @@ def get_redis_async_client(
new_startup_nodes.append(ClusterNode(**item))
cluster_kwargs.pop("startup_nodes", None)
# Default to a periodic health check + TCP keepalive so a connection silently dropped
# by a cluster restart (e.g. ElastiCache Serverless maintenance) is revalidated and
# reconnected before reuse instead of stalling in re-initialization; an explicit value
# from config still wins.
cluster_kwargs.setdefault("health_check_interval", REDIS_CLUSTER_HEALTH_CHECK_INTERVAL)
cluster_kwargs.setdefault("socket_keepalive", True)
# Create async RedisCluster with IAM token as password if available
cluster_client = async_redis.RedisCluster(
startup_nodes=new_startup_nodes,

View file

@ -332,6 +332,10 @@ REDIS_CONNECTION_POOL_TIMEOUT = int(os.getenv("REDIS_CONNECTION_POOL_TIMEOUT", 5
REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD = int(os.getenv("REDIS_CIRCUIT_BREAKER_FAILURE_THRESHOLD", 5))
REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT = int(os.getenv("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", 60))
REDIS_CIRCUIT_BREAKER_ENABLED = os.getenv("REDIS_CIRCUIT_BREAKER_ENABLED", "true").lower() == "true"
# Seconds of idle before a Redis cluster connection is validated with a PING and
# reconnected if dead, so a connection silently dropped by a cluster restart
# (e.g. ElastiCache Serverless maintenance) is not reused while broken
REDIS_CLUSTER_HEALTH_CHECK_INTERVAL = 25
# Default Redis major version to assume when version cannot be determined
# Using 7 as it's the modern version that supports LPOP with count parameter
DEFAULT_REDIS_MAJOR_VERSION = int(os.getenv("DEFAULT_REDIS_MAJOR_VERSION", 7))
@ -1123,6 +1127,7 @@ BEDROCK_CONVERSE_MODELS = [
"anthropic.claude-haiku-4-5-20251001-v1:0",
"anthropic.claude-sonnet-4-5-20250929-v1:0",
"anthropic.claude-fable-5",
"anthropic.claude-sonnet-5",
"anthropic.claude-opus-4-8",
"anthropic.claude-opus-4-7",
"anthropic.claude-opus-4-6-v1:0",

View file

@ -1,9 +1,12 @@
"""
This hook is used to inject cache control directives into the messages of a chat completion.
This hook is used to inject cache control directives into messages.
Users can define
- `cache_control_injection_points` in the completion params and litellm will inject the cache control directives into the messages at the specified injection points.
Supported for both `v1/chat/completions` (via the prompt-management hook) and
`v1/messages` (via `apply_to_anthropic_messages_request`).
"""
import copy
@ -225,6 +228,98 @@ class AnthropicCacheControlHook(CustomPromptManagement):
message_content[-1]["cache_control"] = control # type: ignore
return message
@staticmethod
def apply_to_anthropic_messages_request(
messages: List[Dict],
system: str | list | None,
injection_points: List[CacheControlInjectionPoint],
) -> Tuple[List[Dict], str | list | None, List[CacheControlInjectionPoint]]:
"""Apply cache control injection for the Anthropic-native v1/messages endpoint.
Returns (messages, system, remaining_non_message_points).
"""
if not injection_points:
return messages, system, []
processed_messages: List[Dict] = copy.deepcopy(messages)
processed_system = copy.deepcopy(system) if system is not None else None
message_points: List[CacheControlMessageInjectionPoint] = []
system_points: List[CacheControlMessageInjectionPoint] = []
remaining_points: List[CacheControlInjectionPoint] = []
for point in injection_points:
if point.get("location") == "message":
msg_point = cast(CacheControlMessageInjectionPoint, point)
if msg_point.get("role") == "system":
system_points.append(msg_point)
else:
message_points.append(msg_point)
else:
remaining_points.append(point)
reserved_blocks = 1 if any(p.get("location") == "tool_config" for p in remaining_points) else 0
max_blocks = MAX_CACHE_CONTROL_BLOCKS - reserved_blocks
used_blocks = sum(
AnthropicCacheControlHook._count_cache_control_blocks(cast(AllMessageValues, msg))
for msg in processed_messages
)
if isinstance(processed_system, list):
used_blocks += sum(
1 for b in processed_system if isinstance(b, dict) and b.get("cache_control") is not None
)
if system_points and processed_system is not None and used_blocks < max_blocks:
system_already_has_cc = isinstance(processed_system, list) and any(
isinstance(b, dict) and b.get("cache_control") is not None for b in processed_system
)
if not system_already_has_cc:
control = system_points[0].get("control") or ChatCompletionCachedContent(type="ephemeral")
if isinstance(processed_system, str):
processed_system = [{"type": "text", "text": processed_system, "cache_control": control}]
used_blocks += 1
elif len(processed_system) > 0 and isinstance(processed_system[-1], dict):
processed_system[-1] = {**processed_system[-1], "cache_control": control}
used_blocks += 1
for i, msg in enumerate(processed_messages):
content = msg.get("content")
if isinstance(content, str):
processed_messages[i] = {**msg, "content": [{"type": "text", "text": content}]}
processed_messages = AnthropicCacheControlHook._apply_message_injections(
points=message_points,
messages=cast(List[AllMessageValues], processed_messages),
max_blocks=max_blocks - used_blocks,
)
return processed_messages, processed_system, remaining_points
@staticmethod
def maybe_inject_cache_control(
messages: List[Dict],
system: str | list | None,
kwargs: Dict[str, Any],
) -> Tuple[List[Dict], str | list | None]:
"""Extract cache_control_injection_points from kwargs and apply if present.
Pops the key from kwargs; if remaining (non-message) points exist they
are written back so downstream transforms can handle them.
"""
injection_points = kwargs.pop("cache_control_injection_points", None)
if not injection_points:
return messages, system
messages, system, remaining = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
if remaining:
kwargs["cache_control_injection_points"] = remaining
return messages, system
@property
def integration_name(self) -> str:
"""Return the integration name for this hook."""

View file

@ -40,9 +40,11 @@ from litellm.types.utils import (
LITELLM_CODE_EXECUTION_TOOL_NAME = "litellm_code_execution"
_INTERCEPTION_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
_SESSION_SCOPED_KEY = "_code_interpreter_interception_session_scoped"
_CONVERTED_STREAM_KEY = "_code_interpreter_interception_converted_stream"
_LITELLM_METADATA_KEY = "litellm_metadata"
_CACHE_TTL_SECONDS = 15 * 60
_SESSION_SCOPED_PER_IDENTITY_CAP = 10
class CodeExecutionToolCall(TypedDict, total=False):
@ -107,6 +109,20 @@ class ChatCompletionFunctionToolChoice(TypedDict):
CodeExecutionFunctionToolChoice = ResponsesFunctionToolChoice | ChatCompletionFunctionToolChoice
def _extract_session_id(kwargs: dict[str, Any]) -> str | None:
for meta_key in ("metadata", "litellm_metadata"):
meta = kwargs.get(meta_key)
if isinstance(meta, dict):
sid = meta.get("session_id")
if sid and isinstance(sid, str):
return sid
return None
def _extract_identity(kwargs: dict[str, Any]) -> str:
return kwargs.get("user_api_key_hash") or ""
def _resolve_sandbox_tool(sandbox_tool_name: str | None) -> dict[str, Any] | None:
try:
from litellm.sandbox.sandbox_tools import resolve_sandbox_tool
@ -140,7 +156,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
self.enabled_providers = enabled_providers
self.sandbox_tool_name = sandbox_tool_name
self.sandbox_config = sandbox_config
self._container_cache: dict[str, tuple[Any, dict[str, Any] | None, float]] = {}
self._container_cache: dict[str, tuple[Any, dict[str, Any] | None, float, str | None]] = {}
@classmethod
def from_config_yaml(cls, config: CodeInterpreterInterceptionConfig) -> "CodeInterpreterInterceptionLogger":
@ -191,7 +207,13 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
return None
kwargs[_INTERCEPTION_ACTIVE_KEY] = True
kwargs[_SANDBOX_KEY] = uuid.uuid4().hex
session_id = _extract_session_id(kwargs)
if session_id:
identity = _extract_identity(kwargs)
kwargs[_SANDBOX_KEY] = f"{identity}:{session_id}" if identity else session_id
kwargs[_SESSION_SCOPED_KEY] = True
else:
kwargs[_SANDBOX_KEY] = uuid.uuid4().hex
if kwargs.get("stream"):
kwargs["stream"] = False
kwargs[_CONVERTED_STREAM_KEY] = True
@ -217,6 +239,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
if not is_interception_internal_key(key)
and not key.startswith("_agentic_loop")
and key != "max_agentic_loops"
and key != _SESSION_SCOPED_KEY
}
if filtered_metadata:
kwargs[_LITELLM_METADATA_KEY] = filtered_metadata
@ -227,7 +250,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
def _write_interception_metadata(kwargs: dict[str, Any]) -> None:
metadata = kwargs.get(_LITELLM_METADATA_KEY)
metadata = dict(metadata) if isinstance(metadata, dict) else {}
for key in (_INTERCEPTION_ACTIVE_KEY, _SANDBOX_KEY, _CONVERTED_STREAM_KEY):
for key in (_INTERCEPTION_ACTIVE_KEY, _SANDBOX_KEY, _SESSION_SCOPED_KEY, _CONVERTED_STREAM_KEY):
if key in kwargs:
metadata[key] = kwargs[key]
kwargs[_LITELLM_METADATA_KEY] = metadata
@ -347,7 +370,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
await self._prune_expired_cache()
tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
sandbox_key = kwargs.get(_SANDBOX_KEY)
container, params = await self._get_or_create_container(cache_key=sandbox_key)
is_session = bool(kwargs.get(_SESSION_SCOPED_KEY))
identity = _extract_identity(kwargs) if is_session else None
container, params = await self._get_or_create_container(cache_key=sandbox_key, identity=identity)
try:
container_id = cast(str | None, getattr(container, "id", None))
@ -404,6 +429,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
metadata={
"tool_type": "code_interpreter",
"sandbox_key": sandbox_key or "",
"is_session_scoped": bool(kwargs.get(_SESSION_SCOPED_KEY)),
"code_interpreter_calls": code_interpreter_calls,
},
)
@ -419,7 +445,9 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
await self._prune_expired_cache()
tool_calls = cast(list[CodeExecutionToolCall], tools.get("tool_calls", []))
sandbox_key = cast(str | None, kwargs.get(_SANDBOX_KEY))
container, params = await self._get_or_create_container(cache_key=sandbox_key)
is_session = bool(kwargs.get(_SESSION_SCOPED_KEY))
identity = _extract_identity(cast(dict[str, Any], kwargs)) if is_session else None
container, params = await self._get_or_create_container(cache_key=sandbox_key, identity=identity)
try:
container_id = cast(str | None, getattr(container, "id", None))
@ -455,6 +483,7 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
metadata={
"tool_type": "code_interpreter",
"sandbox_key": sandbox_key or "",
"is_session_scoped": bool(kwargs.get(_SESSION_SCOPED_KEY)),
"code_interpreter_calls": code_interpreter_calls,
"response_format": "openai",
},
@ -489,6 +518,8 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
async def async_agentic_loop_cleanup_hook(self, plan: AgenticLoopPlan, kwargs: dict) -> None:
metadata = plan.metadata or {} if plan else {}
if metadata.get("is_session_scoped"):
return
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
@staticmethod
@ -520,7 +551,8 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
async def async_post_agentic_loop_response_hook(self, response: Any, plan: AgenticLoopPlan, kwargs: dict) -> Any:
metadata = plan.metadata or {} if plan else {}
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
if not metadata.get("is_session_scoped"):
await self._delete_container_for_cache_key(metadata.get("sandbox_key"))
calls = metadata.get("code_interpreter_calls")
if not calls:
@ -565,17 +597,32 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
return f"[execution error] {message}"
return getattr(result, "stdout", "") or ""
async def _get_or_create_container(self, cache_key: str | None) -> tuple[Any, dict[str, Any] | None]:
async def _get_or_create_container(
self,
cache_key: str | None,
identity: str | None = None,
) -> tuple[Any, dict[str, Any] | None]:
if cache_key:
cached = self._container_cache.get(cache_key)
if cached is not None:
self._container_cache[cache_key] = (cached[0], cached[1], time.time(), cached[3])
return cached[0], cached[1]
container, params = await self._create_container()
if cache_key:
self._container_cache[cache_key] = (container, params, time.time())
if identity is not None:
await self._evict_lru_session_if_over_cap(identity)
self._container_cache[cache_key] = (container, params, time.time(), identity)
return container, params
async def _evict_lru_session_if_over_cap(self, identity: str) -> None:
identity_entries = [(k, v) for k, v in self._container_cache.items() if v[3] == identity]
if len(identity_entries) < _SESSION_SCOPED_PER_IDENTITY_CAP:
return
lru_key, lru_entry = min(identity_entries, key=lambda item: item[1][2])
self._container_cache.pop(lru_key, None)
await self._delete_container(container=lru_entry[0], params=lru_entry[1])
async def _create_container(self) -> tuple[Any, dict[str, Any] | None]:
if self.sandbox_config is not None:
return await self.sandbox_config.acreate_sandbox(), None
@ -739,12 +786,8 @@ class CodeInterpreterInterceptionLogger(CustomLogger):
now = time.time()
expired = [
(cache_key, container, params)
for cache_key, (
container,
params,
created_at,
) in self._container_cache.items()
if now - created_at > _CACHE_TTL_SECONDS
for cache_key, (container, params, last_accessed, *_) in self._container_cache.items()
if now - last_accessed > _CACHE_TTL_SECONDS
]
for cache_key, container, params in expired:
self._container_cache.pop(cache_key, None)

View file

@ -81,8 +81,7 @@ SOFT_BUDGET_ALERT_EMAIL_TEMPLATE = """
If you have any questions, please send an email to {email_support_contact} <br /> <br />
Best, <br />
The LiteLLM team <br />
{email_footer}
"""
TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE = """
@ -105,8 +104,7 @@ TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE = """
If you have any questions, please send an email to {email_support_contact} <br /> <br />
Best, <br />
The LiteLLM team <br />
{email_footer}
"""
MAX_BUDGET_ALERT_EMAIL_TEMPLATE = """
@ -129,6 +127,5 @@ MAX_BUDGET_ALERT_EMAIL_TEMPLATE = """
If you have any questions, please send an email to {email_support_contact} <br /> <br />
Best, <br />
The LiteLLM team <br />
{email_footer}
"""

View file

@ -32,11 +32,13 @@ from litellm.integrations.otel.model.payloads import (
LLMCallSpanData,
LLMRequestParams,
LLMUsage,
MCPListToolsSpanData,
MCPToolCallSpanData,
ProxyRequestSpanData,
ServerInfo,
ServiceSpanData,
SpanError,
is_mcp_list_tools,
is_mcp_tool_call,
)
from litellm.integrations.otel.model.semconv import (
@ -106,6 +108,7 @@ __all__ = [
"LLMCallSpanData",
"LLMRequestParams",
"LLMUsage",
"MCPListToolsSpanData",
"MCPToolCallSpanData",
"ProxyRequestSpanData",
"RequestContext",
@ -113,6 +116,7 @@ __all__ = [
"ServerInfo",
"ServiceSpanData",
"SpanError",
"is_mcp_list_tools",
"is_mcp_tool_call",
"promoted_baggage",
]

View file

@ -4,7 +4,7 @@ from collections import OrderedDict
from typing import Callable, Sequence
from opentelemetry.context import Context
from opentelemetry.trace import Span, Tracer
from opentelemetry.trace import Link, Span, Tracer
from opentelemetry.trace.status import Status, StatusCode
from litellm.integrations.otel.model.config import OpenTelemetryV2Config
@ -13,6 +13,7 @@ from litellm.integrations.otel.mappers.base import AttributeMapper, SpanData
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPListToolsSpanData,
MCPToolCallSpanData,
ServiceSpanData,
)
@ -23,6 +24,7 @@ from litellm.integrations.otel.model.spans import (
SpanRole,
guardrail_span_name,
llm_call_span_name,
mcp_list_tools_span_name,
mcp_tool_call_span_name,
service_span_name,
)
@ -33,6 +35,7 @@ from litellm.integrations.otel.model.spans import (
_NAME_BUILDERS: dict[SpanRole, Callable[..., str]] = {
SpanRole.LLM_CALL: llm_call_span_name,
SpanRole.MCP_TOOL_CALL: mcp_tool_call_span_name,
SpanRole.MCP_LIST_TOOLS: mcp_list_tools_span_name,
SpanRole.GUARDRAIL: guardrail_span_name,
# DB_CALL and SERVICE are both built from ServiceSpanData; they differ only in
# span kind (CLIENT vs INTERNAL) and attribute vocabulary, not in naming.
@ -74,18 +77,21 @@ class SpanEmitter:
start_time_ns: int | None = None,
*,
tracer: Tracer | None = None,
links: Sequence[Link] | None = None,
) -> Span:
"""Start a span for ``role`` without dedup or attribute mapping.
For callers that own and manage their own span lifecycle. ``tracer``
overrides the bound tracer for this span only, used for per-request
multi-tenant credential routing.
multi-tenant credential routing. ``links`` records related-but-not-parent
spans (e.g. the transport span of an MCP message, per MCP semconv).
"""
return (tracer or self._tracer).start_span(
name,
context=parent_context,
kind=to_otel_span_kind(SPAN_REGISTRY[role].kind),
start_time=start_time_ns,
links=list(links) if links else None,
)
def _seen(self, dedup_key: str | None, role: SpanRole) -> bool:
@ -116,16 +122,23 @@ class SpanEmitter:
start_time_ns: int | None = None,
end_time_ns: int | None = None,
tracer: Tracer | None = None,
links: Sequence[Link] | None = None,
) -> Span | None:
"""Emit one complete span: dedup, start, map attributes, status, end.
Return the span, or ``None`` if it was deduplicated away. ``tracer``
overrides the bound tracer for this span, used for per-request routing.
``links`` records related-but-not-parent spans (the transport span of an
MCP message).
"""
# LLM-call and MCP tool-call spans carry a dedup key (their request's
# call id), so a sync+async double-firing coalesces. ``isinstance`` narrows
# the type for mypy and keeps the engine free of duck-typed attribute reads.
dedup_key = data.identity.call_id if isinstance(data, (LLMCallSpanData, MCPToolCallSpanData)) else None
dedup_key = (
data.identity.call_id
if isinstance(data, (LLMCallSpanData, MCPToolCallSpanData, MCPListToolsSpanData))
else None
)
if self._seen(dedup_key, role):
return None
span = self.start_span(
@ -134,6 +147,7 @@ class SpanEmitter:
parent_context=parent_context,
start_time_ns=start_time_ns,
tracer=tracer,
links=links,
)
self.finish_span(role, span, data, end_time_ns=end_time_ns)
return span
@ -166,6 +180,7 @@ class SpanEmitter:
(
LLMCallSpanData,
MCPToolCallSpanData,
MCPListToolsSpanData,
ServiceSpanData,
GuardrailSpanData,
),

View file

@ -5,7 +5,7 @@ from contextlib import contextmanager
from datetime import datetime
from typing import TYPE_CHECKING, Any, Callable, Iterator, Mapping, Sequence, cast
from opentelemetry.context import attach, get_current
from opentelemetry.context import Context, attach, get_current
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.trace import Span, Tracer, get_current_span, use_span
@ -17,6 +17,7 @@ from litellm.integrations.otel.model.config import OpenTelemetryV2Config
from litellm.integrations.otel.plumbing.context import (
is_recordable_span,
request_root_span,
resolve_mcp_span_context,
resolve_parent_context,
resolve_request_span_context,
set_request_baggage,
@ -32,9 +33,11 @@ from litellm.integrations.otel.model.metadata import (
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPListToolsSpanData,
MCPToolCallSpanData,
ServiceSpanData,
SpanError,
is_mcp_list_tools,
is_mcp_tool_call,
)
from litellm.integrations.otel.plumbing.metrics import (
@ -218,6 +221,8 @@ class OpenTelemetryV2(CustomLogger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if self._emit_mcp_tool_call(kwargs, start_time, end_time):
return
if self._emit_mcp_list_tools(kwargs, start_time, end_time):
return
self._close_llm_call(kwargs, start_time, end_time)
self._record_metrics(kwargs, response_obj, start_time, end_time)
@ -242,8 +247,24 @@ class OpenTelemetryV2(CustomLogger):
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if self._emit_mcp_tool_call(kwargs, start_time, end_time):
return
if self._emit_mcp_list_tools(kwargs, start_time, end_time):
return
self._close_llm_call(kwargs, start_time, end_time)
def _seed_identity_baggage(self, identity: RequestIdentity, model: str | None, context: Context) -> Context:
"""Seed authenticated request-identity Baggage onto ``context`` so the Baggage
processor stamps team/key/metadata onto the span. Identity is read from the
parsed payload, never the client's ``params._meta`` carrier, so it can't be
spoofed."""
bag = promoted_baggage(
identity,
model,
promoted_keys=tuple(self.config.baggage_promoted_keys),
metadata_keys=tuple(self.config.baggage_metadata_keys),
team_metadata_keys=tuple(self.config.baggage_team_metadata_keys),
)
return set_request_baggage(bag, context=context) if bag else context
def _emit_mcp_tool_call(
self,
kwargs: Mapping[str, Any],
@ -254,10 +275,12 @@ class OpenTelemetryV2(CustomLogger):
MCP tool calls reach the success/failure callbacks like any other request
(with ``call_type`` ``call_mcp_tool``), but they are not LLM calls and have
no ``pre_call`` carrier — so they get their own CLIENT span here, parented
to the request's server span. Returns whether it handled the event, so the
caller skips the LLM-call path. The whole span is emitted at once (there is
no boundary to open it at), deduped on the call id by the emitter.
no ``pre_call`` carrier — so they get their own CLIENT span here. Per the MCP
semconv it parents to the trace context the client propagated in
``params._meta`` (or starts a new root) and links the transport span, rather
than nesting under the HTTP/session span. Returns whether it handled the
event, so the caller skips the LLM-call path. The whole span is emitted at
once (there is no boundary to open it at), deduped on the call id.
"""
raw_payload = kwargs.get("standard_logging_object")
if not raw_payload or not is_mcp_tool_call(cast(Mapping[str, object], raw_payload)):
@ -271,12 +294,51 @@ class OpenTelemetryV2(CustomLogger):
# as a phantom LLM span.
if data.identity.call_id:
self._open_llm_calls.pop(data.identity.call_id, None)
parent_context, links = resolve_mcp_span_context()
parent_context = self._seed_identity_baggage(data.identity, None, parent_context)
self._emitter.emit(
SpanRole.MCP_TOOL_CALL,
data,
parent_context=resolve_request_span_context(),
parent_context=parent_context,
start_time_ns=to_ns(start_time),
end_time_ns=to_ns(end_time),
links=links,
)
return True
def _emit_mcp_list_tools(
self,
kwargs: Mapping[str, object],
start_time: datetime | float | None,
end_time: datetime | float | None,
) -> bool:
"""Emit an MCP ``tools/list`` span when the closed request was a discovery call.
Like a tool call, listing reaches the success/failure callbacks (here with
``call_type`` ``list_mcp_tools``) with no ``pre_call`` carrier, so it gets its
own CLIENT span. Per the MCP semconv it parents to the ``params._meta`` trace
context (or starts a new root) and links the transport span, rather than
nesting under the HTTP/session span. Returns whether it handled the event so
the caller skips the LLM-call path.
"""
raw_payload = kwargs.get("standard_logging_object")
if not raw_payload or not is_mcp_list_tools(cast(Mapping[str, object], raw_payload)):
return False
payload = cast("StandardLoggingPayload", raw_payload)
data = MCPListToolsSpanData.from_standard_logging_payload(
payload, capture_content=self.config.capture_span_content
)
if data.identity.call_id:
self._open_llm_calls.pop(data.identity.call_id, None)
parent_context, links = resolve_mcp_span_context()
parent_context = self._seed_identity_baggage(data.identity, None, parent_context)
self._emitter.emit(
SpanRole.MCP_LIST_TOOLS,
data,
parent_context=parent_context,
start_time_ns=to_ns(start_time),
end_time_ns=to_ns(end_time),
links=links,
)
return True
@ -319,16 +381,7 @@ class OpenTelemetryV2(CustomLogger):
# root span — parent to it (ambient fallback on the SDK path). Seed identity
# Baggage so the span — and the SDK path, which has none — is labeled
# consistently.
parent_ctx = resolve_request_span_context()
bag = promoted_baggage(
data.identity,
data.request_model,
promoted_keys=tuple(self.config.baggage_promoted_keys),
metadata_keys=tuple(self.config.baggage_metadata_keys),
team_metadata_keys=tuple(self.config.baggage_team_metadata_keys),
)
if bag:
parent_ctx = set_request_baggage(bag, context=parent_ctx)
parent_ctx = self._seed_identity_baggage(data.identity, data.request_model, resolve_request_span_context())
return self._emitter.emit(
SpanRole.LLM_CALL,
data,

View file

@ -7,6 +7,7 @@ from typing_extensions import Protocol, runtime_checkable
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPListToolsSpanData,
MCPToolCallSpanData,
ServiceSpanData,
)
@ -20,7 +21,7 @@ AttributeMap = dict[str, AttrValue]
# The closed set of span-data types the engine routes through the mapper chain.
# Server spans (PROXY_REQUEST + management routes) belong to the mounted FastAPI
# instrumentor, not the mapper chain.
SpanData = LLMCallSpanData | MCPToolCallSpanData | GuardrailSpanData | ServiceSpanData
SpanData = LLMCallSpanData | MCPToolCallSpanData | MCPListToolsSpanData | GuardrailSpanData | ServiceSpanData
@runtime_checkable

View file

@ -19,6 +19,7 @@ from litellm.integrations.otel.mappers.utils import (
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPListToolsSpanData,
MCPToolCallSpanData,
ServiceSpanData,
ToolDefinition,
@ -100,6 +101,15 @@ class GenAIMapper:
f"{LiteLLM.COST_PREFIX}total": lambda d: d.response_cost,
}
# A tools/list discovery span: the method and session only. Per semconv it must
# NOT carry gen_ai.operation.name (execute_tool) or gen_ai.tool.name — those are
# for tool calls, and listing executes no tool.
_MCP_LIST_ATTRS: dict[str, Callable[[MCPListToolsSpanData], AttrValue | None]] = {
MCP.METHOD_NAME: lambda d: d.method,
MCP.SESSION_ID: lambda d: d.session_id,
LiteLLM.CALL_ID: lambda d: d.identity.call_id or None,
}
_GUARDRAIL_ATTRS: dict[str, Callable[[GuardrailSpanData], AttrValue | None]] = {
LiteLLM.GUARDRAIL_NAME: lambda d: d.guardrail_name,
LiteLLM.GUARDRAIL_MODE: lambda d: d.mode,
@ -130,6 +140,8 @@ class GenAIMapper:
return self._llm_call(data)
case MCPToolCallSpanData():
return collect(self._MCP_ATTRS, data)
case MCPListToolsSpanData():
return collect(self._MCP_LIST_ATTRS, data)
case GuardrailSpanData():
return self._guardrail(data)
case ServiceSpanData():

View file

@ -37,12 +37,14 @@ __all__ = [
"LLMCost",
"LLMRequestParams",
"LLMUsage",
"MCPListToolsSpanData",
"MCPToolCallSpanData",
"ProxyRequestSpanData",
"ServerInfo",
"ServiceSpanData",
"SpanError",
"ToolDefinition",
"is_mcp_list_tools",
"is_mcp_tool_call",
]
@ -415,6 +417,42 @@ def is_mcp_tool_call(payload: Mapping[str, object]) -> bool:
return bool(_mcp_tool_call_metadata(payload)) or (payload.get("call_type") == "call_mcp_tool")
@dataclass(frozen=True)
class MCPListToolsSpanData:
"""One MCP ``tools/list`` discovery call, parsed from a closed request's payload.
The proxy is an MCP *client* enumerating an upstream server's tools, so this is
a CLIENT span. It carries neither ``gen_ai.operation.name`` nor ``gen_ai.tool.name``:
the GenAI semconv sets ``execute_tool`` (and the tool name) only for tool *calls*,
and listing executes no tool.
"""
method: str
session_id: str | None
error: SpanError | None
identity: RequestIdentity
@classmethod
def from_standard_logging_payload(
cls, payload: StandardLoggingPayload, capture_content: bool = False
) -> MCPListToolsSpanData:
# The list-tools logging path does not thread an MCP session id into the
# payload (only the tool-call path stamps ``mcp_tool_call_metadata``), so
# there is none to read here; ``mcp.session.id`` is simply omitted.
return cls(
method=MCPMethod.TOOLS_LIST.value,
session_id=None,
error=_parse_error(payload),
identity=RequestContext.from_standard_logging_payload(payload).identity,
)
def is_mcp_list_tools(payload: Mapping[str, object]) -> bool:
"""Whether a closed request's payload is an MCP ``tools/list`` discovery call
rather than a tool call or an LLM call — true when the call type says so."""
return payload.get("call_type") == "list_mcp_tools"
# --- service event_metadata sanitization ------------------------------------ #
# Substrings (case-insensitive) of keys that must never reach a span: secrets,

View file

@ -18,6 +18,13 @@ before the LLM call even starts), so a guardrail is a sibling of the LLM call,
not a child of it. The emitter parents every span to the ambient OTel context
(the active server span), which matches this.
MCP spans (``MCP_TOOL_CALL``, ``MCP_LIST_TOOLS``) are intentionally NOT in this
tree. Per the OTel GenAI MCP semconv, MCP and the HTTP transport are independent
contexts, so an MCP span parents to the trace context the client propagated in
``params._meta`` (or starts its own root when none is propagated) and records the
``PROXY_REQUEST`` transport span as a span *link*, never a parent. The registry
encodes this as ``parent=None, links=PROXY_REQUEST``.
Not every service call becomes a span — :func:`span_role_for_service` decides:
- ``DB_CALL`` (CLIENT) — outbound datastores (redis, postgres,
@ -46,6 +53,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPListToolsSpanData,
MCPToolCallSpanData,
ProxyRequestSpanData,
ServiceSpanData,
@ -56,6 +64,7 @@ class SpanRole(str, Enum):
PROXY_REQUEST = "proxy_request"
LLM_CALL = "llm_call"
MCP_TOOL_CALL = "mcp_tool_call"
MCP_LIST_TOOLS = "mcp_list_tools"
GUARDRAIL = "guardrail"
DB_CALL = "db_call"
SERVICE = "service"
@ -74,14 +83,24 @@ class SpanSpec:
role: SpanRole
kind: LiteLLMSpanKind
parent: SpanRole | None
links: SpanRole | None = None
SPAN_REGISTRY: dict[SpanRole, SpanSpec] = {
SpanRole.PROXY_REQUEST: SpanSpec(SpanRole.PROXY_REQUEST, LiteLLMSpanKind.SERVER, parent=None),
SpanRole.LLM_CALL: SpanSpec(SpanRole.LLM_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
# The proxy is an MCP client to the upstream server it dispatches the tool
# call to, so this is a CLIENT span, sibling of the LLM call under the request.
SpanRole.MCP_TOOL_CALL: SpanSpec(SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
# MCP and the HTTP transport are independent contexts (OTel GenAI MCP semconv),
# so an MCP span does not nest under the transport span. The proxy is an MCP
# client to the upstream server, so it's a CLIENT span; it parents to the trace
# context the client propagated in ``params._meta`` (or starts its own root when
# none is propagated) and records the PROXY_REQUEST transport span as a span
# *link*, never a parent — hence ``parent=None, links=PROXY_REQUEST``.
SpanRole.MCP_TOOL_CALL: SpanSpec(
SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=None, links=SpanRole.PROXY_REQUEST
),
SpanRole.MCP_LIST_TOOLS: SpanSpec(
SpanRole.MCP_LIST_TOOLS, LiteLLMSpanKind.CLIENT, parent=None, links=SpanRole.PROXY_REQUEST
),
SpanRole.GUARDRAIL: SpanSpec(SpanRole.GUARDRAIL, LiteLLMSpanKind.INTERNAL, parent=SpanRole.PROXY_REQUEST),
SpanRole.DB_CALL: SpanSpec(SpanRole.DB_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST),
SpanRole.SERVICE: SpanSpec(SpanRole.SERVICE, LiteLLMSpanKind.INTERNAL, parent=SpanRole.PROXY_REQUEST),
@ -163,6 +182,12 @@ def mcp_tool_call_span_name(data: "MCPToolCallSpanData") -> str:
return f"{data.method} {data.tool_name}".strip()
def mcp_list_tools_span_name(data: "MCPListToolsSpanData") -> str:
"""``"{mcp.method.name}"`` i.e. ``"tools/list"`` — no low-cardinality target, so
the method name alone names the span (MCP semconv)."""
return data.method
def proxy_request_span_name(data: "ProxyRequestSpanData") -> str:
"""``"{method} {route}"`` (HTTP semconv)."""
return f"{data.http_method} {data.route}".strip()
@ -179,7 +204,8 @@ def service_span_name(data: "ServiceSpanData") -> str:
def root_roles() -> list[SpanRole]:
"""Roles that start a new trace (no in-process parent)."""
"""Roles with no in-process parent. They start a new trace unless they adopt a
remote parent (e.g. an MCP span joining the client's propagated context)."""
return [role for role, spec in SPAN_REGISTRY.items() if spec.parent is None]
@ -196,6 +222,8 @@ def validate_registry(
raise ValueError(f"SPAN_REGISTRY[{role}] has mismatched role {spec.role}")
if spec.parent is not None and spec.parent not in reg:
raise ValueError(f"span role {role} declares unknown parent {spec.parent}")
if spec.links is not None and spec.links not in reg:
raise ValueError(f"span role {role} declares unknown link target {spec.links}")
missing = [role for role in SpanRole if role not in reg]
if missing:
raise ValueError(f"SPAN_REGISTRY is missing roles: {missing}")

View file

@ -1,11 +1,11 @@
"""Trace-context + Baggage helpers."""
from contextvars import ContextVar
from contextvars import ContextVar, Token
from typing import Mapping
from opentelemetry import baggage
from opentelemetry.context import Context, get_current
from opentelemetry.trace import Span, get_current_span, set_span_in_context
from opentelemetry.trace import Link, Span, get_current_span, set_span_in_context
from opentelemetry.trace.propagation.tracecontext import (
TraceContextTextMapPropagator,
)
@ -47,6 +47,31 @@ def request_root_span() -> "Span | None":
return span if is_recordable_span(span) else None
# The W3C trace-context carrier (``traceparent``/``tracestate``/``baggage``) the
# MCP client propagated in the current request's ``params._meta``. The MCP gateway
# sets it per message so the MCP span can parent to the client's span rather than
# to the transport. A ``ContextVar`` because, like the root-span anchor, it must
# ride the request task and be readable by the inline success-logging callback.
_mcp_message_trace_carrier: "ContextVar[Mapping[str, str] | None]" = ContextVar(
"litellm_otel_mcp_message_trace_carrier", default=None
)
def set_mcp_message_trace_carrier(
carrier: "Mapping[str, str] | None",
) -> "Token[Mapping[str, str] | None]":
"""Stash the current MCP message's propagated trace-context carrier.
Returns the reset token; the caller must reset it once the message is handled
so the carrier never leaks to the next message on the same session task.
"""
return _mcp_message_trace_carrier.set(carrier)
def reset_mcp_message_trace_carrier(token: "Token[Mapping[str, str] | None]") -> None:
_mcp_message_trace_carrier.reset(token)
def set_request_baggage(values: Mapping[str, str], context: Context | None = None) -> Context:
"""Return a context with ``values`` written into Baggage."""
ctx = context
@ -104,6 +129,38 @@ def resolve_request_span_context() -> Context:
return get_current()
def resolve_mcp_span_context(
carrier: "Mapping[str, str] | None" = None,
) -> "tuple[Context, tuple[Link, ...]]":
"""Parent context + links for an MCP message span, per the OTel GenAI MCP semconv.
MCP and the underlying transport (HTTP) are independent lifecycles — one
streamable-HTTP session multiplexes many messages, so nesting the message span
under the HTTP/session span is wrong (it renders the message at the session's
start, skewed by however long the session has been open). Instead:
* parent to the trace context the client propagated in the request's
``params._meta`` (a *remote* parent), and
* record the transport/session span as a *link*, never the parent.
Only trace context (``traceparent``/``tracestate``) is extracted, never the
client's W3C Baggage: ``params._meta`` is caller-controlled, and the otel
baggage processor stamps allowlisted baggage keys (``litellm.team.id``,
``litellm.metadata.*``, ...) onto the span as attributes, so honoring remote
baggage would let a client spoof a span's identity attribution.
With no propagated context the returned context carries no span, so the span
starts its own root trace (still linked to the transport). The base context is
explicitly empty so an absent ``traceparent`` can never fall through to the
ambient (stale session) span.
"""
source = carrier if carrier is not None else _mcp_message_trace_carrier.get()
parent = _PROPAGATOR.extract(dict(source or {}), context=Context())
transport = request_root_span()
links = (Link(transport.get_span_context()),) if transport is not None else ()
return parent, links
def is_recordable_span(obj: object) -> bool:
"""True if ``obj`` is a live span with a valid context (safe to parent under)."""
if not isinstance(obj, Span):

View file

@ -8,7 +8,23 @@ identity unconditionally.
"""
from contextlib import contextmanager
from typing import Any, Iterator
from functools import cache
from typing import Any, Callable, Iterator, Optional
@cache
def _otel_runtime() -> "Optional[tuple[Callable[[str], Any], Callable[..., None]]]":
"""Resolve the SDK-backed hooks once and cache the outcome, absence included.
CPython never caches a failed import, so without this memoization every call
site re-attempts the import on each request; when the OTel SDK is not installed
that re-scans ``sys.path`` and contends on the import lock on the hot path.
"""
try:
from litellm.integrations.otel import logger
except Exception:
return None
return (logger.phase_span, logger.seed_request_identity)
@contextmanager
@ -18,21 +34,17 @@ def phase_span(name: str) -> "Iterator[Any]":
Yields ``None`` (a plain no-op) when the OTel SDK is unavailable or V2 is not
the active logger.
"""
try:
from litellm.integrations.otel.logger import phase_span as _phase_span
except Exception:
runtime = _otel_runtime()
if runtime is None:
yield None
return
with _phase_span(name) as span:
with runtime[0](name) as span:
yield span
def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None:
"""Seed request-identity Baggage at the auth boundary (no-op without V2)."""
try:
from litellm.integrations.otel.logger import (
seed_request_identity as _seed_request_identity,
)
except Exception:
runtime = _otel_runtime()
if runtime is None:
return
_seed_request_identity(user_api_key_dict, model=model)
runtime[1](user_api_key_dict, model=model)

View file

@ -3804,6 +3804,10 @@ def _get_combined_custom_metadata_from_standard_logging_payload(
) -> Dict[str, Any]:
"""
Combine the metadata sources that can supply custom Prometheus labels.
Includes top-level scalar fields from the standard logging metadata (e.g.
user_api_key_project_alias, user_api_key_team_alias) so they are accessible
via custom_prometheus_metadata_labels configuration.
"""
if not isinstance(standard_logging_payload, dict):
return {}
@ -3817,6 +3821,7 @@ def _get_combined_custom_metadata_from_standard_logging_payload(
spend_logs_metadata = standard_logging_metadata.get("spend_logs_metadata")
return {
**{k: v for k, v in standard_logging_metadata.items() if not isinstance(v, dict)},
**(requester_metadata if isinstance(requester_metadata, dict) else {}),
**(user_api_key_auth_metadata if isinstance(user_api_key_auth_metadata, dict) else {}),
**(spend_logs_metadata if isinstance(spend_logs_metadata, dict) else {}),

View file

@ -31,6 +31,7 @@ from litellm.types.integrations.websearch_interception import (
WebSearchInterceptionConfig,
)
from litellm.types.integrations.custom_logger import (
CHAT_COMPLETION_AGENTIC_SURFACE,
AgenticLoopPlan,
AgenticLoopRequestPatch,
)
@ -440,12 +441,16 @@ class WebSearchInterceptionLogger(CustomLogger):
custom_llm_provider: str,
kwargs: Dict,
) -> Tuple[bool, Dict]:
"""
Check if WebSearch tool interception is needed for Anthropic Messages API.
This is the legacy method for Anthropic-style responses.
For chat completions, use async_should_run_chat_completion_agentic_loop instead.
"""
if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE:
return await self.async_should_run_chat_completion_agentic_loop(
response=response,
model=model,
messages=messages,
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
)
verbose_logger.debug(f"WebSearchInterception: Hook called! provider={custom_llm_provider}, stream={stream}")
verbose_logger.debug(f"WebSearchInterception: Response type: {type(response)}")
@ -629,6 +634,18 @@ class WebSearchInterceptionLogger(CustomLogger):
stream: bool,
kwargs: Dict,
) -> AgenticLoopPlan:
if kwargs.get("_agentic_loop_api_surface") == CHAT_COMPLETION_AGENTIC_SURFACE:
return await self.async_build_chat_completion_agentic_loop_plan(
tools=tools,
model=model,
messages=messages,
response=response,
optional_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
stream=stream,
kwargs=kwargs,
)
tool_calls = tools["tool_calls"]
thinking_blocks = tools.get("thinking_blocks", [])
request_patch, structured_results = await self._build_anthropic_request_patch(
@ -1088,6 +1105,7 @@ class WebSearchInterceptionLogger(CustomLogger):
raise ValueError("WebSearchInterception: missing follow-up messages")
params = dict(optional_params)
params.update(request_patch.optional_params)
params.pop("tool_choice", None)
return await litellm.acompletion(
model=request_patch.model or model,
messages=request_patch.messages,
@ -1203,6 +1221,7 @@ class WebSearchInterceptionLogger(CustomLogger):
if k
not in {
"tools",
"tool_choice",
"extra_body",
"model_alias_map",
"stream_response",

View file

@ -137,8 +137,8 @@ async def _execute_chat_completion_agentic_plan(
optional_params_for_followup = {**optional_params, **patch.optional_params}
if patch.tools is not None:
optional_params_for_followup["tools"] = patch.tools
if "tool_choice" not in patch.optional_params:
optional_params_for_followup.pop("tool_choice", None)
if "tool_choice" not in patch.optional_params:
optional_params_for_followup.pop("tool_choice", None)
kwargs_for_followup = _filter_followup_kwargs(kwargs)
kwargs_for_followup.update(
@ -206,10 +206,11 @@ async def maybe_run_chat_completion_agentic_loop(
for callback in callbacks:
if not isinstance(callback, CustomLogger):
continue
if not _gate_overridden(callback):
continue
gate_kwargs = {
hook_kwargs = {
**kwargs,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
"custom_llm_provider": custom_llm_provider,
@ -222,7 +223,7 @@ async def maybe_run_chat_completion_agentic_loop(
tools=tools,
stream=stream,
custom_llm_provider=custom_llm_provider,
kwargs=gate_kwargs,
kwargs=hook_kwargs,
)
except Exception as e:
verbose_logger.exception(
@ -243,11 +244,6 @@ async def maybe_run_chat_completion_agentic_loop(
)
try:
plan_kwargs = {
**kwargs,
"_agentic_loop_api_surface": CHAT_COMPLETION_AGENTIC_SURFACE,
"custom_llm_provider": custom_llm_provider,
}
if not _build_plan_overridden(callback):
return await callback.async_run_agentic_loop(
tools=tool_calls,
@ -258,7 +254,7 @@ async def maybe_run_chat_completion_agentic_loop(
anthropic_messages_optional_request_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=plan_kwargs,
kwargs=hook_kwargs,
)
plan = await callback.async_build_agentic_loop_plan(
@ -270,7 +266,7 @@ async def maybe_run_chat_completion_agentic_loop(
anthropic_messages_optional_request_params=optional_params,
logging_obj=logging_obj,
stream=stream,
kwargs=plan_kwargs,
kwargs=hook_kwargs,
)
if plan.response_override is not None:

View file

@ -6,7 +6,7 @@ import mimetypes
import re
import xml.etree.ElementTree as ET
from enum import Enum
from typing import Any, Dict, List, Optional, Set, Tuple, Union, cast, overload
from typing import Any, Dict, List, Optional, Set, Tuple, TypedDict, Union, cast, overload
from jinja2.sandbox import ImmutableSandboxedEnvironment
@ -5322,3 +5322,146 @@ def get_attribute_or_key(tool_or_function, attribute, default=None):
if hasattr(tool_or_function, attribute):
return getattr(tool_or_function, attribute)
return tool_or_function.get(attribute, default)
class NormalizedToolCall(TypedDict):
id: Optional[str]
name: Optional[str]
arguments: dict[str, Any]
def _parse_tool_call_arguments(raw: Any, tool_name: Optional[str], context: str) -> dict[str, Any]:
# Anthropic's tool_use blocks already carry a parsed dict in "input";
# chat completions and the Responses API carry a JSON string that may be
# truncated by the model, so route those through the repair-aware parser.
if isinstance(raw, dict):
return raw
if not isinstance(raw, str):
return {}
from litellm.litellm_core_utils.prompt_templates.common_utils import (
parse_tool_call_arguments,
)
try:
parsed = parse_tool_call_arguments(raw, tool_name=tool_name, context=context)
except ValueError as e:
verbose_logger.warning("Failed to parse tool call arguments: %s", e)
return {}
return parsed if isinstance(parsed, dict) else {}
def _tool_calls_from_chat_completion_response(response: Any) -> list[NormalizedToolCall]:
choices = get_attribute_or_key(response, "choices", None)
if not (isinstance(choices, list) and choices):
return []
message = get_attribute_or_key(choices[0], "message", None)
tool_calls = get_attribute_or_key(message, "tool_calls", None) if message else None
if not isinstance(tool_calls, list):
return []
result: list[NormalizedToolCall] = []
for tc in tool_calls:
fn = get_attribute_or_key(tc, "function", None)
if fn is None:
continue
name = get_attribute_or_key(fn, "name")
result.append(
NormalizedToolCall(
id=get_attribute_or_key(tc, "id"),
name=name,
arguments=_parse_tool_call_arguments(
get_attribute_or_key(fn, "arguments", "{}"),
tool_name=name,
context="chat completions",
),
)
)
return result
def _tool_calls_from_responses_api_response(response: Any) -> list[NormalizedToolCall]:
output = get_attribute_or_key(response, "output", None)
if not isinstance(output, list):
return []
result: list[NormalizedToolCall] = []
for item in output:
if get_attribute_or_key(item, "type") != "function_call":
continue
name = get_attribute_or_key(item, "name")
result.append(
NormalizedToolCall(
id=get_attribute_or_key(item, "call_id") or get_attribute_or_key(item, "id"),
name=name,
arguments=_parse_tool_call_arguments(
get_attribute_or_key(item, "arguments", "{}"),
tool_name=name,
context="responses API",
),
)
)
return result
def _tool_calls_from_anthropic_messages_response(response: Any) -> list[NormalizedToolCall]:
content = get_attribute_or_key(response, "content", None)
if not isinstance(content, list):
return []
result: list[NormalizedToolCall] = []
for block in content:
if get_attribute_or_key(block, "type") != "tool_use":
continue
raw_input = get_attribute_or_key(block, "input", {})
result.append(
NormalizedToolCall(
id=get_attribute_or_key(block, "id"),
name=get_attribute_or_key(block, "name"),
arguments=raw_input if isinstance(raw_input, dict) else {},
)
)
return result
def get_tool_calls_from_response(response: Any) -> list[NormalizedToolCall]:
"""
Extract tool/function calls from a response object into a normalized
``{"id", "name", "arguments"}`` shape, regardless of which API surface
produced it: chat completions (``choices[].message.tool_calls``),
the Responses API (``output`` items of type ``function_call``), or the
Anthropic Messages API (``content`` blocks of type ``tool_use``).
Callers that only care about a specific tool should filter the result by
``name`` themselves -- this returns every tool call found.
"""
for extractor in (
_tool_calls_from_chat_completion_response,
_tool_calls_from_responses_api_response,
_tool_calls_from_anthropic_messages_response,
):
tool_calls = extractor(response)
if tool_calls:
return tool_calls
return []
def has_tool_with_name(tools: Any, tool_name: str) -> bool:
"""
Check whether a tools list (as sent to an LLM) includes a tool with the
given name, regardless of shape: OpenAI-style function tools
(``{"type": "function", "function": {"name": ...}}``) or Anthropic's
native tool shape (a top-level ``"name"``, e.g.
``{"name": ..., "input_schema": ...}``). Anthropic's documented client
tool format doesn't require a ``"type"`` key at all -- ``"custom"`` is
only one of several possible values -- so any non-OpenAI-shaped tool is
matched on its top-level ``"name"``.
"""
if not isinstance(tools, list):
return False
for tool in tools:
if not isinstance(tool, dict):
continue
function = tool.get("function")
if tool.get("type") == "function" and isinstance(function, dict):
if function.get("name") == tool_name:
return True
elif tool.get("name") == tool_name:
return True
return False

View file

@ -5,6 +5,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Protocol, Union, ca
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.types.llms.openai import (
OpenAIRealtimeEvents,
@ -315,8 +316,10 @@ class RealTimeStreaming:
self.logging_obj.model_call_details["realtime_tools"] = self.session_tools
self.logging_obj.model_call_details["realtime_tool_calls"] = self.tool_calls
## ASYNC LOGGING
# Create an event loop for the new thread
asyncio.create_task(self.logging_obj.async_success_handler(self.messages))
# Route through the bounded logging worker (per-coroutine timeout +
# concurrency cap) instead of a bare create_task, so a slow callback
# can't leave suspended tasks pinning each call's response in memory.
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(self.logging_obj.async_success_handler(self.messages))
## SYNC LOGGING
executor.submit(self.logging_obj.success_handler(self.messages))

View file

@ -407,6 +407,37 @@ def token_counter(
return num_tokens
def _count_function_call_tokens(
key: str,
value: Any,
message: Mapping[str, Any],
count_function: TokenCounterFunction,
) -> int:
"""
Count tokens contributed by an assistant message's tool/function call payload.
Handles both the modern `tool_calls` list and the legacy OpenAI
`function_call` dict. Only the `arguments` string is counted (matching the
existing tool_calls behavior); names are accounted for elsewhere via the
tool/function definitions and `tool_choice`.
"""
if key == "tool_calls":
if not isinstance(value, List):
raise ValueError(f"Unsupported type {type(value)} for key tool_calls in message {message}")
total = 0
for tool_call in value:
if "function" not in tool_call:
raise ValueError(f"Unsupported tool call {tool_call} must contain a function key")
function_arguments = tool_call["function"].get("arguments", "")
total += count_function(str(function_arguments))
return total
if key == "function_call":
if not isinstance(value, Mapping):
raise ValueError(f"Unsupported type {type(value)} for key function_call in message {message}")
return count_function(str(value.get("arguments", "")))
raise ValueError(f"Unexpected key {key!r}; expected 'tool_calls' or 'function_call'")
def _count_messages(
params: _MessageCountParams,
messages: List[AllMessageValues],
@ -430,16 +461,8 @@ def _count_messages(
for key, value in message.items():
if value is None:
pass
elif key == "tool_calls":
if isinstance(value, List):
for tool_call in value:
if "function" in tool_call:
function_arguments = tool_call["function"].get("arguments", [])
num_tokens += params.count_function(str(function_arguments))
else:
raise ValueError(f"Unsupported tool call {tool_call} must contain a function key")
else:
raise ValueError(f"Unsupported type {type(value)} for key tool_calls in message {message}")
elif key in ("tool_calls", "function_call"):
num_tokens += _count_function_call_tokens(key, value, message, params.count_function)
elif isinstance(value, str):
num_tokens += params.count_function(value)
if key == "name":

View file

@ -61,6 +61,19 @@ def _should_route_to_responses_api(custom_llm_provider: Optional[str]) -> bool:
return custom_llm_provider in _RESPONSES_API_PROVIDERS
def _deployment_passes_through_anthropic_messages(model_info: object) -> bool:
"""Whether the deployment opted into forwarding /v1/messages untranslated.
The opt-in is ``model_info.supported_endpoints`` containing ``"/v1/messages"``,
declared per deployment in config.yaml and plumbed here as ``kwargs["model_info"]``
by the router.
"""
if not isinstance(model_info, dict):
return False
supported_endpoints = model_info.get("supported_endpoints")
return isinstance(supported_endpoints, (list, tuple)) and "/v1/messages" in supported_endpoints
####### ENVIRONMENT VARIABLES ###################
# Initialize any necessary instances or variables here
base_llm_http_handler = BaseLLMHTTPHandler()
@ -186,7 +199,7 @@ async def anthropic_messages(
metadata: Optional[Dict] = None,
stop_sequences: Optional[List[str]] = None,
stream: Optional[bool] = False,
system: Optional[str] = None,
system: Optional[Union[str, list]] = None,
temperature: Optional[float] = None,
thinking: Optional[Dict] = None,
tool_choice: Optional[Dict] = None,
@ -217,6 +230,12 @@ async def anthropic_messages(
# ids like ``functions.Bash:0`` that violate Anthropic's id pattern.
messages = sanitize_tool_use_ids_in_anthropic_messages(messages)
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
original_stream = stream or kwargs.get("_websearch_interception_converted_stream", False)
# Execute pre-request hooks to allow CustomLoggers to modify request.
@ -362,7 +381,7 @@ def anthropic_messages_handler(
metadata: Optional[Dict] = None,
stop_sequences: Optional[List[str]] = None,
stream: Optional[bool] = False,
system: Optional[str] = None,
system: Optional[Union[str, list]] = None,
temperature: Optional[float] = None,
thinking: Optional[Dict] = None,
tool_choice: Optional[Dict] = None,
@ -399,6 +418,12 @@ def anthropic_messages_handler(
messages = strip_empty_text_blocks_from_anthropic_messages(messages)
messages = sanitize_tool_use_ids_in_anthropic_messages(messages)
from litellm.integrations.anthropic_cache_control_hook import (
AnthropicCacheControlHook,
)
messages, system = AnthropicCacheControlHook.maybe_inject_cache_control(messages, system, kwargs)
metadata = validate_anthropic_api_metadata(metadata)
local_vars = locals()
@ -456,6 +481,14 @@ def anthropic_messages_handler(
model=model,
provider=litellm.LlmProviders(custom_llm_provider),
)
if anthropic_messages_provider_config is None and _deployment_passes_through_anthropic_messages(
kwargs.get("model_info")
):
from litellm.llms.openai_like.messages.transformation import (
OpenAILikeAnthropicMessagesConfig,
)
anthropic_messages_provider_config = OpenAILikeAnthropicMessagesConfig()
if anthropic_messages_provider_config is None:
# Route to Responses API for OpenAI / Azure, chat/completions for everything else.
_shared_kwargs = dict(

View file

@ -103,6 +103,17 @@ class BaseAnthropicMessagesConfig(ABC):
"""
return headers, None
def should_filter_anthropic_beta_headers(self) -> bool:
"""
Whether ``anthropic-beta`` header values should be filtered down to the
ones the routed provider supports before the upstream request.
Cross-provider translation paths (bedrock, vertex_ai, ...) need this so
unsupported betas are dropped. Configs that forward natively to an
Anthropic-compatible endpoint return False to pass betas through verbatim.
"""
return True
def get_async_streaming_response_iterator(
self,
model: str,

View file

@ -1986,7 +1986,8 @@ class BaseLLMHTTPHandler:
api_base=api_base,
)
headers = update_headers_with_filtered_beta(headers=headers, provider=custom_llm_provider)
if anthropic_messages_provider_config.should_filter_anthropic_beta_headers():
headers = update_headers_with_filtered_beta(headers=headers, provider=custom_llm_provider)
logging_obj.update_from_kwargs(
kwargs=kwargs,
@ -2111,7 +2112,7 @@ class BaseLLMHTTPHandler:
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
)
return initial_response
else:
@ -2121,6 +2122,10 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
)
# Inject api_key into kwargs so follow-up calls in agentic hooks can
# authenticate. api_key is a named param here (not in kwargs), so
# _prepare_followup_kwargs would miss it otherwise.
kwargs_for_agentic = {**kwargs, "api_key": api_key} if api_key else kwargs
# Call agentic completion hooks (non-streaming path only)
final_response = await self._call_agentic_completion_hooks(
response=initial_response,
@ -2131,7 +2136,7 @@ class BaseLLMHTTPHandler:
logging_obj=logging_obj,
stream=False,
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
kwargs=kwargs_for_agentic,
)
return self._maybe_wrap_in_fake_stream(

View file

@ -0,0 +1,69 @@
from typing import Any, Optional
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
DEFAULT_ANTHROPIC_API_VERSION = "2023-06-01"
class OpenAILikeAnthropicMessagesConfig(AnthropicMessagesConfig):
"""
Forwards Anthropic /v1/messages requests to an OpenAI-compatible server that
also natively exposes the Anthropic Messages API, with no translation.
Opted into per deployment via ``model_info.supported_endpoints`` containing
``"/v1/messages"``. The inbound Anthropic payload (system, cache_control,
thinking, tools, ...) is forwarded essentially unchanged to
``{api_base}/v1/messages``, so Anthropic-only features that the
Anthropic->OpenAI translation would otherwise drop are preserved. Response
parsing and streaming are inherited from the native Anthropic config.
"""
def validate_anthropic_messages_environment(
self,
headers: dict[str, str],
model: str,
messages: list[Any],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> tuple[dict[str, str], Optional[str]]:
present = {key.lower() for key in headers}
needs_auth = bool(api_key) and "authorization" not in present and "x-api-key" not in present
defaults: dict[str, str] = {
**({"authorization": f"Bearer {api_key}"} if needs_auth else {}),
**({"anthropic-version": DEFAULT_ANTHROPIC_API_VERSION} if "anthropic-version" not in present else {}),
**({"content-type": "application/json"} if "content-type" not in present else {}),
}
combined = {**headers, **defaults}
normalized = {
("anthropic-beta" if key.lower() == "anthropic-beta" else key): value for key, value in combined.items()
}
merged = self._update_headers_with_anthropic_beta(
headers=normalized,
optional_params=optional_params,
)
return merged, api_base
def should_filter_anthropic_beta_headers(self) -> bool:
return False
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:
raise ValueError("api_base is required to forward Anthropic /v1/messages to a native endpoint")
base = api_base.rstrip("/")
if base.endswith("/v1/messages"):
return base
if base.endswith("/v1"):
base = base[: -len("/v1")]
return f"{base}/v1/messages"

View file

@ -1671,6 +1671,204 @@
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"global.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"au.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"jp.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"anthropic.claude-sonnet-4-6": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
@ -2511,6 +2709,36 @@
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true
},
"azure_ai/claude-sonnet-4-6": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
@ -10245,6 +10473,40 @@
"supports_vision": true,
"supports_web_search": true
},
"claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true,
"provider_specific_entry": {
"us": 1.1
},
"supports_output_config": true
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -34944,6 +35206,36 @@
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true
},
"vertex_ai/claude-sonnet-4-6": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
@ -42381,6 +42673,36 @@
"search_context_size_high": 0.035
}
},
"vertex_ai/claude-sonnet-5@default": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true
},
"vertex_ai/claude-sonnet-4-6@default": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,

View file

@ -21,6 +21,7 @@ class LiteLLM_EndUserTable(LiteLLMPydanticObjectBase):
spend: float = 0.0
allowed_model_region: Optional[Literal["eu", "us"]] = None
default_model: Optional[str] = None
budget_id: Optional[str] = None
litellm_budget_table: Optional[LiteLLM_BudgetTable] = None
object_permission_id: Optional[str] = None
object_permission: Optional[LiteLLM_ObjectPermissionTable] = None

View file

@ -24,3 +24,4 @@ class LiteLLM_ObjectPermissionTable(LiteLLMPydanticObjectBase):
mcp_toolsets: Optional[List[str]] = None
blocked_tools: Optional[List[str]] = []
search_tools: Optional[List[str]] = []
mcp_tool_search_enabled: Optional[bool] = None

View file

@ -41,6 +41,7 @@ litellm/proxy/_experimental/mcp_server/
sampling_handler.py # MCP sampling to LiteLLM completion flow
elicitation_handler.py # MCP elicitation relay flow
semantic_tool_filter.py # semantic filtering of available MCP tools
tool_search.py # opt-in virtual tools (mcp_tool_search + mcp_tool_call) for large catalogs
guardrail_translation/
handler.py # MCP guardrail result translation
sse_transport.py # SSE transport implementation
@ -79,6 +80,11 @@ module materially harder to understand.
encryption need focused tests for both allowed and rejected paths.
- Avoid adding comments to new code unless they explain non-obvious security or
protocol behavior. Prefer clear names and small functions.
- The virtual tool path (`tool_search.py`, gated by `mcp_tool_search_enabled`)
must mirror the normal tool flow: IP filtering, server allowlist, per-key tool
permissions, no-accessible-server rejection, per-request auth headers, server
scope, error to `isError` conversion, and spend logging. Reuse `_list_mcp_tools`
and `execute_mcp_tool` rather than reimplementing any of these checks.
## Tests

View file

@ -569,6 +569,21 @@ if MCP_AVAILABLE:
include_disabled_tools and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN
)
if apply_tool_filters and getattr(
getattr(user_api_key_dict, "object_permission", None),
"mcp_tool_search_enabled",
False,
):
from litellm.proxy._experimental.mcp_server.tool_search import (
get_virtual_tool_definitions,
)
return {
"tools": get_virtual_tool_definitions(),
"error": None,
"message": "Successfully retrieved tools",
}
# Extract auth headers from request
headers = request.headers
raw_headers_from_request = dict(headers)
@ -727,6 +742,74 @@ if MCP_AVAILABLE:
try:
data = await request.json()
tool_name = data.get("name")
tool_arguments = data.get("arguments") or {}
from litellm.proxy._experimental.mcp_server.tool_search import (
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
handle_mcp_tool_call,
handle_mcp_tool_search,
)
if tool_name in (MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME):
if not getattr(
getattr(user_api_key_dict, "object_permission", None),
"mcp_tool_search_enabled",
False,
):
raise HTTPException(
status_code=403,
detail={
"error": "forbidden",
"message": f"{tool_name} requires mcp_tool_search_enabled on the key",
},
)
rest_client_ip = IPAddressUtils.get_mcp_client_ip(request)
(
virtual_mcp_auth_header,
virtual_mcp_server_auth_headers,
virtual_raw_headers,
) = _extract_mcp_headers_from_request(request, MCPRequestHandler)
virtual_oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(request.headers)
if tool_name == MCP_TOOL_SEARCH_TOOL_NAME:
return await handle_mcp_tool_search(
query=tool_arguments.get("query", ""),
top_k=coerce_top_k(tool_arguments.get("top_k", 5)),
user_api_key_dict=user_api_key_dict,
client_ip=rest_client_ip,
mcp_auth_header=virtual_mcp_auth_header,
mcp_server_auth_headers=virtual_mcp_server_auth_headers,
oauth2_headers=virtual_oauth2_headers,
raw_headers=virtual_raw_headers,
)
else: # MCP_TOOL_CALL_TOOL_NAME
# Run the same pre-call pipeline as the normal call path so the
# tool execution is spend-logged and guardrail-checked.
(
_,
virtual_logging_obj,
) = await ProxyBaseLLMRequestProcessing(data=data).common_processing_pre_call_logic(
request=request,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
route_type=CallTypes.call_mcp_tool.value,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
)
return await handle_mcp_tool_call(
tool_name=tool_arguments.get("tool_name", ""),
arguments=tool_arguments.get("arguments") or {},
user_api_key_dict=user_api_key_dict,
client_ip=rest_client_ip,
mcp_auth_header=virtual_mcp_auth_header,
mcp_server_auth_headers=virtual_mcp_server_auth_headers,
oauth2_headers=virtual_oauth2_headers,
raw_headers=virtual_raw_headers,
litellm_logging_obj=virtual_logging_obj,
)
# Validate required parameters early
server_id = data.get("server_id")
if not server_id:
@ -738,7 +821,6 @@ if MCP_AVAILABLE:
},
)
tool_name = data.get("name")
if not tool_name:
raise HTTPException(
status_code=400,
@ -748,8 +830,6 @@ if MCP_AVAILABLE:
},
)
tool_arguments = data.get("arguments") or {}
proxy_base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
(
data,

View file

@ -10,8 +10,8 @@ import contextvars
import hashlib
import json
import time
import types
import traceback
import types
import uuid
from datetime import datetime
from typing import (
@ -37,13 +37,17 @@ from starlette.types import Message, Receive, Scope, Send
from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
get_request_base_url,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.mcp_context import (
_mcp_active_toolset_id,
_mcp_gateway_initialize_instructions,
@ -59,10 +63,6 @@ from litellm.proxy._experimental.mcp_server.utils import (
get_server_prefix,
iter_known_server_prefixes,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import (
ProxyException,
SpecialMCPServerNames,
@ -122,9 +122,12 @@ def _write_byok_cred_cache(user_id: str, server_id: str, credential: Optional[st
# TODO: Make this a util function for litellm client usage
MCP_AVAILABLE: bool = True
try:
import weakref
from mcp import ReadResourceResult, Resource
from mcp.server import Server
from mcp.server.lowlevel.helper_types import ReadResourceContents
from mcp.server.session import ServerSession as _McpServerSession
from mcp.types import (
BlobResourceContents,
GetPromptResult,
@ -132,8 +135,6 @@ try:
TextResourceContents,
Tool,
)
from mcp.server.session import ServerSession as _McpServerSession
import weakref
# Robust auth lookup keyed by session_object.
_session_obj_auth_storage: "weakref.WeakKeyDictionary[Any, MCPAuthenticatedUser]" = weakref.WeakKeyDictionary()
@ -229,6 +230,56 @@ def _jsonrpc_text_has_top_level_method(text: str) -> bool:
return False
def _mcp_meta_trace_carrier(req_ctx: object) -> Optional[dict[str, str]]:
"""The W3C trace context (``traceparent``/``tracestate``) the MCP client
propagated in the request's ``params._meta`` (SEP-414), or ``None``.
Per the OTel MCP semconv the MCP span parents to this propagated context rather
than to the HTTP/session transport (which is recorded as a link instead), so a
streamable-HTTP session that multiplexes many messages does not glue every
message under the session's first request. The client's W3C Baggage is
deliberately excluded: it is caller-controlled, and the otel baggage processor
stamps allowlisted baggage keys (``litellm.team.id``, ``litellm.metadata.*``,
...) onto the span, so honoring remote baggage would let a client spoof a
span's identity attribution.
"""
meta = getattr(req_ctx, "meta", None)
extra = getattr(meta, "model_extra", None)
if not isinstance(extra, dict):
return None
carrier = {key: extra[key] for key in ("traceparent", "tracestate") if isinstance(extra.get(key), str)}
return carrier or None
def _otel_set_mcp_trace_carrier(carrier: Optional[dict[str, str]]) -> object:
"""Stash ``carrier`` for the otel_v2 MCP span and return a reset token, or
``None`` when otel_v2 is unavailable. Lazily imported so opentelemetry stays an
optional dependency."""
try:
from litellm.integrations.otel.plumbing.context import (
set_mcp_message_trace_carrier,
)
return set_mcp_message_trace_carrier(carrier)
except ImportError:
return None
def _otel_reset_mcp_trace_carrier(token: object) -> None:
"""Clear the per-message trace carrier so it never leaks to the next message on
the same session task. Paired with ``_otel_set_mcp_trace_carrier``."""
if token is None:
return
try:
from litellm.integrations.otel.plumbing.context import (
reset_mcp_message_trace_carrier,
)
reset_mcp_message_trace_carrier(token)
except ImportError:
return
def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
"""Map a ``ProxyException`` to an ``HTTPException`` that preserves its real
status code and headers.
@ -253,14 +304,14 @@ def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException:
if MCP_AVAILABLE:
from mcp.server import Server
from mcp.server.lowlevel.server import NotificationOptions
from mcp.server.models import InitializationOptions
# Import auth context variables and middleware
from mcp.server.auth.middleware.auth_context import (
AuthContextMiddleware,
auth_context_var,
)
from mcp.server.lowlevel.server import NotificationOptions
from mcp.server.models import InitializationOptions
try:
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
@ -595,8 +646,10 @@ if MCP_AVAILABLE:
_session_reset_token = None
if req_ctx:
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
_trace_token = None
try:
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
# Get user authentication from context variable
(
user_api_key_auth,
@ -612,6 +665,19 @@ if MCP_AVAILABLE:
verbose_logger.debug(
f"MCP list_tools - MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
if getattr(
getattr(user_api_key_auth, "object_permission", None),
"mcp_tool_search_enabled",
False,
):
from mcp.types import Tool
from litellm.proxy._experimental.mcp_server.tool_search import (
get_virtual_tool_definitions,
)
return [Tool(**d) for d in get_virtual_tool_definitions()]
# Get mcp_servers from context variable
verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools")
tools = await _list_mcp_tools(
@ -632,9 +698,154 @@ if MCP_AVAILABLE:
# This prevents the HTTP stream from failing and allows the client to get a response
return []
finally:
_otel_reset_mcp_trace_carrier(_trace_token)
if _session_reset_token is not None:
active_mcp_session_var.reset(_session_reset_token)
def _capture_host_progress_callback(host_server) -> Optional[Callable]:
"""Return a progress-forwarding callback bound to the host MCP session.
Returns ``None`` when the host did not supply a progress token.
"""
try:
host_ctx = host_server.request_context
except Exception as e:
verbose_logger.warning(f"Could not capture host progress context: {e}")
return None
if not (host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta):
return None
host_token = getattr(host_ctx.meta, "progressToken", None)
if not (host_token and hasattr(host_ctx, "session") and host_ctx.session):
return None
host_session = host_ctx.session
async def forward_progress(progress: float, total: Optional[float]):
"""Forward progress notifications from external MCP to Host"""
try:
await host_session.send_progress_notification(
progress_token=host_token,
progress=progress,
total=total,
)
verbose_logger.debug(f"Forwarded progress {progress}/{total} to Host")
except Exception as e:
verbose_logger.error(f"Failed to forward progress to Host: {e}")
verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...")
return forward_progress
async def _build_virtual_call_logging_obj(
name: str,
arguments: dict[str, Any],
user_api_key_auth: UserAPIKeyAuth,
) -> Optional[LiteLLMLoggingObj]:
"""Run the pre-call pipeline (guardrails + logging setup) for a virtual
mcp_tool_call so the SSE path spend-logs like the REST path."""
from fastapi import Request
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
from litellm.proxy.proxy_server import (
general_settings,
proxy_config,
proxy_logging_obj,
)
request = Request(
scope={
"type": "http",
"method": "POST",
"path": "/mcp/tools/call",
"headers": [(b"content-type", b"application/json")],
}
)
_, virtual_logging_obj = await ProxyBaseLLMRequestProcessing(
data={"name": name, "arguments": arguments}
).common_processing_pre_call_logic(
request=request,
user_api_key_dict=user_api_key_auth,
proxy_config=proxy_config,
route_type=CallTypes.call_mcp_tool.value,
proxy_logging_obj=proxy_logging_obj,
general_settings=general_settings,
)
return virtual_logging_obj
async def _dispatch_virtual_mcp_tool(
name: str,
arguments: Optional[dict[str, Any]],
user_api_key_auth: Optional[UserAPIKeyAuth],
client_ip: Optional[str],
mcp_servers: Optional[list[str]] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None,
oauth2_headers: Optional[dict[str, str]] = None,
raw_headers: Optional[dict[str, str]] = None,
) -> Optional[CallToolResult]:
"""Handle the mcp_tool_search / mcp_tool_call virtual tools.
Returns a CallToolResult when ``name`` is a virtual tool, else ``None`` so
the caller falls through to normal tool routing.
"""
from litellm.proxy._experimental.mcp_server.tool_search import (
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
handle_mcp_tool_call,
handle_mcp_tool_search,
)
if name not in (MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME):
return None
if not getattr(
getattr(user_api_key_auth, "object_permission", None),
"mcp_tool_search_enabled",
False,
):
return CallToolResult(
content=[
TextContent(
type="text",
text=f"Tool {name} requires mcp_tool_search_enabled on the key",
)
],
isError=True,
)
args = arguments or {}
if name == MCP_TOOL_SEARCH_TOOL_NAME:
return await handle_mcp_tool_search(
query=args.get("query", ""),
top_k=coerce_top_k(args.get("top_k", 5)),
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
assert user_api_key_auth is not None # guaranteed by the flag check above
virtual_logging_obj = await _build_virtual_call_logging_obj(
name=name, arguments=args, user_api_key_auth=user_api_key_auth
)
return await handle_mcp_tool_call(
tool_name=args.get("tool_name", ""),
arguments=args.get("arguments") or {},
user_api_key_dict=user_api_key_auth,
client_ip=client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=virtual_logging_obj,
)
@server.call_tool()
async def mcp_server_tool_call(name: str, arguments: Dict[str, Any] | None) -> CallToolResult:
"""
@ -648,18 +859,21 @@ if MCP_AVAILABLE:
HTTPException: If tool not found or arguments missing
"""
from fastapi import Request
from mcp.server.lowlevel.server import request_ctx
from mcp.types import CallToolResult
from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import proxy_config
from mcp.types import CallToolResult
from mcp.server.lowlevel.server import request_ctx
req_ctx = request_ctx.get(None)
_session_reset_token = None
if req_ctx:
_session_reset_token = active_mcp_session_var.set(req_ctx.session)
_trace_token = None
try:
_trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx))
# Validate arguments
(
user_api_key_auth,
@ -675,31 +889,25 @@ if MCP_AVAILABLE:
)
verbose_logger.debug(f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}")
host_progress_callback = None
try:
host_ctx = server.request_context
if host_ctx and hasattr(host_ctx, "meta") and host_ctx.meta:
host_token = getattr(host_ctx.meta, "progressToken", None)
if host_token and hasattr(host_ctx, "session") and host_ctx.session:
host_session = host_ctx.session
async def forward_progress(progress: float, total: Optional[float]):
"""Forward progress notifications from external MCP to Host"""
try:
await host_session.send_progress_notification(
progress_token=host_token,
progress=progress,
total=total,
)
verbose_logger.debug(f"Forwarded progress {progress}/{total} to Host")
except Exception as e:
verbose_logger.error(f"Failed to forward progress to Host: {e}")
host_progress_callback = forward_progress
verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...")
except Exception as e:
verbose_logger.warning(f"Could not capture host progress context: {e}")
try:
# Inside this try so virtual-tool errors convert to isError
# CallToolResult instead of raising out of the protocol handler.
virtual_tool_result = await _dispatch_virtual_mcp_tool(
name=name,
arguments=arguments,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
mcp_servers=mcp_servers,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
if virtual_tool_result is not None:
return virtual_tool_result
host_progress_callback = _capture_host_progress_callback(server)
# Create a body date for logging
body_data = {"name": name, "arguments": arguments}
# Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A)
@ -778,6 +986,7 @@ if MCP_AVAILABLE:
return response
finally:
_otel_reset_mcp_trace_carrier(_trace_token)
if _session_reset_token is not None:
active_mcp_session_var.reset(_session_reset_token)
@ -1472,6 +1681,7 @@ if MCP_AVAILABLE:
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: Optional[str] = None,
litellm_trace_id: Optional[str] = None,
client_ip: Optional[str] = None,
) -> List[MCPTool]:
"""
Helper method to fetch tools from MCP servers based on server filtering criteria.
@ -1559,6 +1769,7 @@ if MCP_AVAILABLE:
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
# Pre-fetch OAuth credentials only when at least one server uses OAuth2,
@ -1643,12 +1854,13 @@ if MCP_AVAILABLE:
)
return filtered_tools
except MCPUpstreamAuthError:
# Surface upstream 401/403 to the outer handler so the
# client receives a proper WWW-Authenticate challenge
# instead of a silently empty tool list. Without this
# re-raise the broad ``except Exception`` below would
# swallow the auth error.
raise
# Absorb so one unauthenticated server does not empty every other server's
# tools. Surfacing the upstream 401 to the client as a re-auth challenge is
# intentionally not done here: raising from this list handler cannot produce a
# 401 + WWW-Authenticate (the MCP session manager serializes it as a JSON-RPC
# error), so that belongs in a request-scope preemptive check, tracked separately.
verbose_logger.debug(f"MCP list_tools: omitting {server.name}; it needs upstream auth")
return []
except Exception as e:
verbose_logger.exception(f"Error getting tools from server {server.name}: {str(e)}")
return []
@ -1967,6 +2179,7 @@ if MCP_AVAILABLE:
raw_headers: Optional[Dict[str, str]] = None,
log_list_tools_to_spendlogs: bool = False,
list_tools_log_source: Optional[str] = None,
client_ip: Optional[str] = None,
) -> List[MCPTool]:
"""
List all available MCP tools.
@ -1976,6 +2189,7 @@ if MCP_AVAILABLE:
mcp_auth_header: Optional auth header for MCP server (deprecated)
mcp_servers: Optional list of server names/aliases to filter by
mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value}
client_ip: Client IP for IP-based server access control
Returns:
List[MCPTool]: Combined list of tools from all accessible servers
@ -1999,6 +2213,7 @@ if MCP_AVAILABLE:
raw_headers=raw_headers,
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
list_tools_log_source=list_tools_log_source,
client_ip=client_ip,
)
verbose_logger.debug(f"Successfully fetched {len(managed_tools)} tools from managed MCP servers")
except Exception as e:

View file

@ -0,0 +1,157 @@
from __future__ import annotations
import json
from datetime import datetime
from typing import TYPE_CHECKING, Any, Optional
if TYPE_CHECKING:
from mcp.types import CallToolResult
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
MCP_TOOL_SEARCH_TOOL_NAME: str = "mcp_tool_search"
MCP_TOOL_CALL_TOOL_NAME: str = "mcp_tool_call"
def coerce_top_k(value: Any, default: int = 5) -> int:
try:
return int(value)
except (TypeError, ValueError):
return default
def search_tools(query: str, tools: list[dict[str, Any]], top_k: int = 5) -> list[dict[str, Any]]:
if not query:
return []
tokens = query.lower().split()
def _score(tool: dict[str, Any]) -> int:
haystack = (tool.get("name", "") + " " + tool.get("description", "")).lower()
return sum(1 for t in tokens if t in haystack)
scored = ((s, tool) for tool in tools if (s := _score(tool)) > 0)
return [tool for _, tool in sorted(scored, key=lambda x: x[0], reverse=True)[:top_k]]
def get_virtual_tool_definitions() -> list[dict[str, Any]]:
return [
{
"name": MCP_TOOL_SEARCH_TOOL_NAME,
"description": "Search for MCP tools by keyword. Returns top matching tools with names, descriptions, and input schemas.",
"inputSchema": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "Keywords to search for in tool names and descriptions.",
},
"top_k": {
"type": "integer",
"description": "Maximum number of results to return.",
"default": 5,
},
},
"required": ["query"],
},
},
{
"name": MCP_TOOL_CALL_TOOL_NAME,
"description": "Call an MCP tool by name with the given arguments.",
"inputSchema": {
"type": "object",
"properties": {
"tool_name": {
"type": "string",
"description": "The exact name of the MCP tool to call.",
},
"arguments": {
"type": "object",
"description": "Arguments to pass to the tool.",
},
},
"required": ["tool_name"],
},
},
]
async def handle_mcp_tool_search(
query: str,
top_k: int,
user_api_key_dict: UserAPIKeyAuth,
client_ip: Optional[str] = None,
mcp_servers: Optional[list[str]] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None,
oauth2_headers: Optional[dict[str, str]] = None,
raw_headers: Optional[dict[str, str]] = None,
) -> CallToolResult:
from mcp.types import CallToolResult, TextContent
from litellm.proxy._experimental.mcp_server.server import _list_mcp_tools
mcp_tools = await _list_mcp_tools(
user_api_key_auth=user_api_key_dict,
mcp_servers=mcp_servers,
client_ip=client_ip,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
)
tools = [
{
"name": t.name,
"description": t.description or "",
"inputSchema": t.inputSchema,
}
for t in mcp_tools
]
results = search_tools(query, tools, top_k)
return CallToolResult(content=[TextContent(type="text", text=json.dumps(results))], isError=False)
async def handle_mcp_tool_call(
tool_name: str,
arguments: dict[str, Any],
user_api_key_dict: UserAPIKeyAuth,
client_ip: Optional[str] = None,
mcp_servers: Optional[list[str]] = None,
mcp_auth_header: Optional[str] = None,
mcp_server_auth_headers: Optional[dict[str, dict[str, str]]] = None,
oauth2_headers: Optional[dict[str, str]] = None,
raw_headers: Optional[dict[str, str]] = None,
litellm_logging_obj: Optional[LiteLLMLoggingObj] = None,
) -> CallToolResult:
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers,
execute_mcp_tool,
)
allowed_mcp_servers = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_dict,
mcp_servers=mcp_servers,
client_ip=client_ip,
)
# Reject before dispatch when the key has no accessible servers; otherwise an
# unprefixed local tool name would fall through to the local registry in
# execute_mcp_tool, which has no server permission check.
if not allowed_mcp_servers:
from fastapi import HTTPException
raise HTTPException(status_code=403, detail="User not allowed to call this tool.")
return await execute_mcp_tool(
name=tool_name,
arguments=arguments,
allowed_mcp_servers=allowed_mcp_servers,
start_time=datetime.now(),
user_api_key_auth=user_api_key_dict,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
)

View file

@ -189,6 +189,9 @@ class LitellmTableNames(str, enum.Enum):
TOOL_TABLE_NAME = "LiteLLM_ToolTable"
CACHE_CONFIG_TABLE_NAME = "LiteLLM_CacheConfig"
CONFIG_OVERRIDES_TABLE_NAME = "LiteLLM_ConfigOverrides"
CONFIG_TABLE_NAME = "LiteLLM_Config"
SSO_CONFIG_TABLE_NAME = "LiteLLM_SSOConfig"
UI_SETTINGS_TABLE_NAME = "LiteLLM_UISettings"
class Litellm_EntityType(enum.Enum):
@ -1003,6 +1006,7 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
agent_access_groups: Optional[List[str]] = None
models: Optional[List[str]] = None
search_tools: Optional[List[str]] = None
mcp_tool_search_enabled: Optional[bool] = None
from litellm.types.object_permission import ( # noqa: E402

View file

@ -278,6 +278,12 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
"s3_endpoint_url",
"sagemaker_base_url",
"deployment_url",
# NVIDIA Riva fields consumed by the audio-transcription handler
# via ``optional_params``. Banned for the same reason as the
# provider-specific entries above: a caller-supplied value retargets
# the request away from the admin's pinned configuration.
"nvcf_function_id",
"use_ssl",
# SDK-only field; also rejected outright in is_request_body_safe.
"model_list",
# Observability credentials, hosts, and project identifiers: derived

View file

@ -175,16 +175,9 @@ class DBSpendUpdateWriter:
if team_id is not None and team_id != "":
payload["team_id"] = team_id
# One deepcopy shared by all 6 daily spend helpers (was 5, fixes agent bug)
payload_copy = copy.deepcopy(payload)
# Deepcopy request_tags for _update_tag_db
request_tags = copy.deepcopy(payload.get("request_tags"))
# Keep _insert_spend_log_to_db awaited inline (not a task, preserve current behavior)
if disable_spend_logs is False:
await self._insert_spend_log_to_db(
payload=copy.deepcopy(payload),
payload=payload,
prisma_client=prisma_client,
)
else:
@ -204,8 +197,7 @@ class DBSpendUpdateWriter:
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
payload_copy=payload_copy,
request_tags=request_tags,
payload=payload,
)
)
@ -336,14 +328,18 @@ class DBSpendUpdateWriter:
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
payload_copy: SpendLogsPayload,
request_tags: Optional[Any],
payload: SpendLogsPayload,
):
"""
Runs all 11 spend-update helpers sequentially inside a single asyncio task.
Each helper is wrapped in try/except so one failure doesn't prevent the others.
The deepcopy runs here, off the awaited request path, so the daily spend
helpers get a payload isolated from the spend-log queue entry and the caller.
"""
payload_copy = copy.deepcopy(payload)
request_tags = payload_copy.get("request_tags")
try:
await self._update_user_db(
response_cost=response_cost,

View file

@ -66,6 +66,29 @@ class PrismaDBExceptionHandler:
return True
return False
@staticmethod
def is_prisma_data_error(e: Exception) -> bool:
"""True iff ``e`` is a base prisma ``DataError``: the database processed
the statement and refused the data itself (e.g. ``invalid byte sequence
for encoding "UTF8": 0x00``), as opposed to a connectivity failure.
Matched by exact type, not ``isinstance``: the specific data-layer
subclasses (``UniqueViolationError``, ``TableNotFoundError``,
``MissingRequiredValueError`` ...) all derive from ``DataError`` but
carry their own semantics, and a systemic one like a missing table must
not be mistaken for a single poison row and bisected away. A raw
Postgres execution error with no prisma P-code surfaces as the base
``DataError``.
prisma also wraps the P1001 "can't reach database server" outage as a
base ``DataError``, so a caller that must not treat an outage as a
per-row data rejection has to additionally consult
``is_database_service_unavailable_error`` before acting on a True here.
"""
import prisma
return type(e) is prisma.errors.DataError
@staticmethod
def is_database_transport_error(e: Exception) -> bool:
"""

View file

@ -182,10 +182,25 @@ model_list:
litellm_params:
model: openai/gpt-5.5
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-4o-mini
litellm_params:
model: openai/gpt-4o-mini
api_key: os.environ/OPENAI_API_KEY
general_settings:
master_key: sk-1234
sandbox_tools:
- sandbox_tool_name: e2b_sandbox
litellm_params:
sandbox_provider: e2b
api_key: os.environ/E2B_API_KEY
litellm_settings:
drop_params: True
telemetry: False
code_interpreter_interception_params:
enabled: true
sandbox_tool_name: e2b_sandbox
callbacks:
- code_interpreter_interception

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, Optional
from litellm.types.guardrails import SupportedGuardrailIntegrations
@ -8,9 +8,23 @@ if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def _get_config_value(litellm_params: Any, optional_params: Any, attribute_name: str) -> Optional[Any]:
if optional_params is not None:
value = (
optional_params.get(attribute_name)
if isinstance(optional_params, dict)
else getattr(optional_params, attribute_name, None)
)
if value is not None:
return value
return getattr(litellm_params, attribute_name, None)
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
optional_params = getattr(litellm_params, "optional_params", None)
_generic_guardrail_api_callback = GenericGuardrailAPI(
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
@ -22,6 +36,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -33,6 +33,7 @@ from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
GUARDRAIL_NAME = "generic_guardrail_api"
@ -178,6 +179,8 @@ class GenericGuardrailAPI(CustomGuardrail):
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
fail_on_error: Optional[bool] = True,
extra_headers: Optional[list] = None,
streaming_end_of_stream_only: Optional[bool] = None,
streaming_sampling_rate: Optional[int] = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
@ -209,6 +212,15 @@ class GenericGuardrailAPI(CustomGuardrail):
self.fail_on_error: bool = True if fail_on_error is None else fail_on_error
# Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook
# via getattr(guardrail_to_apply, "streaming_*", default).
self.streaming_end_of_stream_only: bool = (
False if streaming_end_of_stream_only is None else streaming_end_of_stream_only
)
if streaming_sampling_rate is not None and streaming_sampling_rate < 1:
raise ValueError(f"streaming_sampling_rate must be >= 1 (got {streaming_sampling_rate})")
self.streaming_sampling_rate: int = 5 if streaming_sampling_rate is None else streaming_sampling_rate
# Set supported event hooks
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = [
@ -470,3 +482,11 @@ class GenericGuardrailAPI(CustomGuardrail):
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
except Exception as e:
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj, is_unreachable=False)
@staticmethod
def get_config_model() -> Optional[type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIConfigModel,
)
return GenericGuardrailAPIConfigModel

View file

@ -1,6 +1,10 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Literal
import json
import re
import time
import uuid
from typing import TYPE_CHECKING, Any, Literal, Optional
import httpx
from fastapi import HTTPException
@ -12,12 +16,18 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.litellm_core_utils.prompt_templates.factory import (
get_attribute_or_key,
get_tool_calls_from_response,
has_tool_with_name,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType]
httpxSpecialProvider,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.guardrails import GuardrailEventHooks, Mode
from litellm.types.integrations.custom_logger import AgenticLoopPlan, AgenticLoopRequestPatch
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
@ -25,6 +35,9 @@ if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
BYPASS_HEADER = "x-headroom-bypass"
HEADROOM_RETRIEVE_TOOL_NAME = "headroom_retrieve"
_HASH_PATTERN = re.compile(r"hash=([a-f0-9]{24})")
_HASH_CACHE_TTL_SECONDS = 15 * 60
def _is_str_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
@ -35,6 +48,163 @@ def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isin
return isinstance(value, list)
def extract_hashes_from_messages(messages: list[dict[str, object]]) -> list[str]:
hashes: list[str] = []
for msg in messages:
content = msg.get("content")
if isinstance(content, str):
hashes.extend(_HASH_PATTERN.findall(content))
elif isinstance(content, list):
for block in content:
if isinstance(block, dict):
text = block.get("text")
if isinstance(text, str):
hashes.extend(_HASH_PATTERN.findall(text))
return hashes
def _build_headroom_retrieve_tool() -> dict[str, object]:
return {
"type": "function",
"function": {
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"description": (
"Retrieve original content that was compressed by Headroom. "
"Call this when you encounter a compression marker containing a hash."
),
"parameters": {
"type": "object",
"properties": {
"hash": {
"type": "string",
"description": "The 24-character hex hash from the compression marker.",
},
"query": {
"type": "string",
"description": "Optional search query for BM25-ranked retrieval.",
},
},
"required": ["hash"],
},
},
}
def _resolve_call_id(logging_obj: object, request_state: dict[str, object]) -> Optional[str]:
"""Resolve the litellm_call_id shared by a request's pre-call hook and its
agentic-loop hooks, so CCR hash validation can be scoped per call instead
of trusting any hash-shaped string that shows up in message text."""
logging_call_id = getattr(logging_obj, "litellm_call_id", None)
if isinstance(logging_call_id, str) and logging_call_id:
return logging_call_id
kwargs_call_id = request_state.get("litellm_call_id")
return kwargs_call_id if isinstance(kwargs_call_id, str) else None
def has_headroom_retrieve_tool(tools: object) -> bool:
return has_tool_with_name(tools, HEADROOM_RETRIEVE_TOOL_NAME)
def _extract_headroom_tool_calls(response: object) -> list[dict[str, object]]:
return [
{"id": tc["id"], "type": "function", "name": tc["name"], "arguments": tc["arguments"]}
for tc in get_tool_calls_from_response(response)
if tc["name"] == HEADROOM_RETRIEVE_TOOL_NAME
]
def _build_assistant_message_from_response(response: object) -> dict[str, object]:
choices = getattr(response, "choices", None)
if not isinstance(choices, list) or not choices:
return {"role": "assistant", "content": None, "tool_calls": []}
message = getattr(choices[0], "message", None)
if message is None:
return {"role": "assistant", "content": None, "tool_calls": []}
content = getattr(message, "content", None)
tool_calls = getattr(message, "tool_calls", None)
raw_tool_calls: list[dict[str, object]] = []
if isinstance(tool_calls, list):
for tc in tool_calls:
fn = getattr(tc, "function", None)
raw_tool_calls.append(
{
"id": getattr(tc, "id", None),
"type": "function",
"function": {
"name": getattr(fn, "name", None) if fn else None,
"arguments": getattr(fn, "arguments", "{}") if fn else "{}",
},
}
)
return {"role": "assistant", "content": content, "tool_calls": raw_tool_calls}
def _is_responses_api_response(response: object) -> bool:
# Real response objects can be plain dicts at runtime (e.g. TypedDict-based
# response types), so getattr alone would silently miss the key -- use the
# same dict-or-object accessor as the tool-call extractors.
return isinstance(get_attribute_or_key(response, "output", None), list)
def _is_anthropic_messages_response(response: object) -> bool:
return isinstance(get_attribute_or_key(response, "content", None), list)
def _build_anthropic_followup_messages(
retrieved: list[tuple[dict[str, object], str]],
) -> list[dict[str, object]]:
"""Build Anthropic Messages API follow-up messages for a tool round-trip.
Anthropic requires the tool_use block to be echoed back in an assistant
message, paired with a tool_result block in a user message keyed by the
same tool_use_id -- it does not accept chat-style tool-role messages.
"""
assistant_message: dict[str, object] = {
"role": "assistant",
"content": [
{
"type": "tool_use",
"id": tool_call.get("id"),
"name": tool_call.get("name"),
"input": tool_call.get("arguments", {}),
}
for tool_call, _ in retrieved
],
}
user_message: dict[str, object] = {
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": tool_call.get("id"), "content": content}
for tool_call, content in retrieved
],
}
return [assistant_message, user_message]
def _build_responses_followup_items(
retrieved: list[tuple[dict[str, object], str]],
) -> list[dict[str, object]]:
"""Build Responses API input items for a tool round-trip.
The Responses API does not accept chat-style assistant/tool messages as
follow-up input; it requires the model's function_call to be echoed back
paired with a function_call_output keyed by the same call_id.
"""
items: list[dict[str, object]] = []
for tool_call, content in retrieved:
call_id = tool_call.get("id")
items.append(
{
"type": "function_call",
"call_id": call_id,
"name": tool_call.get("name"),
"arguments": json.dumps(tool_call.get("arguments", {})),
}
)
items.append({"type": "function_call_output", "call_id": call_id, "output": content})
return items
class HeadroomGuardrail(CustomGuardrail):
def __init__(
self,
@ -56,6 +226,7 @@ class HeadroomGuardrail(CustomGuardrail):
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
self._issued_hashes_by_call_id: dict[str, tuple[frozenset[str], float]] = {}
super().__init__( # pyright: ignore[reportUnknownMemberType]
guardrail_name=guardrail_name,
event_hook=event_hook,
@ -72,6 +243,20 @@ class HeadroomGuardrail(CustomGuardrail):
value = headers.get(BYPASS_HEADER)
return str(value).lower() == "true"
def _request_headers(self) -> dict[str, str]:
headers: dict[str, str] = {"Content-Type": "application/json"}
if self.headroom_api_key:
headers["Authorization"] = f"Bearer {self.headroom_api_key}"
return headers
def _prune_expired_hashes(self) -> None:
now = time.monotonic()
self._issued_hashes_by_call_id = {
call_id: (hashes, expiry)
for call_id, (hashes, expiry) in self._issued_hashes_by_call_id.items()
if expiry > now
}
async def _call_compress(
self,
messages: list[dict[str, object]],
@ -81,15 +266,11 @@ class HeadroomGuardrail(CustomGuardrail):
if model:
payload["model"] = model
request_headers: dict[str, str] = {"Content-Type": "application/json"}
if self.headroom_api_key:
request_headers["Authorization"] = f"Bearer {self.headroom_api_key}"
try:
raw_response: HttpxResponse | None = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
url=f"{self.headroom_api_base}/v1/compress",
json=payload,
headers=request_headers,
headers=self._request_headers(),
)
except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError) as e:
raise HTTPException(
@ -118,7 +299,7 @@ class HeadroomGuardrail(CustomGuardrail):
try:
body: object = response.json()
except Exception:
except ValueError:
raise HTTPException(
status_code=502,
detail={
@ -163,6 +344,44 @@ class HeadroomGuardrail(CustomGuardrail):
)
return filtered
async def _call_retrieve(self, hash_value: str, query: str | None = None) -> str:
params: dict[str, str] = {}
if query:
params["query"] = query
try:
raw_response: HttpxResponse | None = await self.async_handler.get( # pyright: ignore[reportUnknownMemberType]
url=f"{self.headroom_api_base}/v1/retrieve/{hash_value}",
params=params,
headers=self._request_headers(),
)
except (httpx.ConnectError, httpx.TimeoutException, httpx.TransportError) as e:
verbose_proxy_logger.warning("Headroom: retrieve failed for hash=%s: %s", hash_value, e)
return f"[Headroom: retrieval failed for hash={hash_value}]"
if raw_response is None or raw_response.status_code == 404:
return f"[Headroom: hash={hash_value} not found or expired]"
if raw_response.status_code != 200:
verbose_proxy_logger.warning(
"Headroom: retrieve returned %s for hash=%s",
raw_response.status_code,
hash_value,
)
return f"[Headroom: retrieval error {raw_response.status_code} for hash={hash_value}]"
try:
body: object = raw_response.json()
except ValueError:
return raw_response.text
if _is_str_object_dict(body):
original_content = body.get("original_content")
if isinstance(original_content, str):
return original_content
return str(body)
@log_guardrail_information
async def apply_guardrail(
self,
@ -192,7 +411,127 @@ class HeadroomGuardrail(CustomGuardrail):
model=model if isinstance(model, str) else None,
)
return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType]
hashes = extract_hashes_from_messages(compressed)
if not hashes:
return {**inputs, "structured_messages": compressed} # pyright: ignore[reportReturnType]
self._prune_expired_hashes()
call_id = _resolve_call_id(logging_obj, request_data)
if not call_id:
call_id = str(uuid.uuid4())
request_data["litellm_call_id"] = call_id
self._issued_hashes_by_call_id[call_id] = (frozenset(hashes), time.monotonic() + _HASH_CACHE_TTL_SECONDS)
existing_tools = inputs.get("tools")
retrieve_tool = _build_headroom_retrieve_tool()
if isinstance(existing_tools, list) and not has_headroom_retrieve_tool(existing_tools):
merged_tools: list[object] = list(existing_tools) + [retrieve_tool]
elif existing_tools is None:
merged_tools = [retrieve_tool]
else:
merged_tools = list(existing_tools) if isinstance(existing_tools, list) else [retrieve_tool]
return {**inputs, "structured_messages": compressed, "tools": merged_tools} # pyright: ignore[reportReturnType]
async def async_should_run_agentic_loop(
self,
response: Any,
model: str,
messages: list[dict],
tools: Optional[list[dict]],
stream: bool,
custom_llm_provider: str,
kwargs: dict,
) -> tuple[bool, dict]:
if not has_headroom_retrieve_tool(tools):
return False, {}
tool_calls = _extract_headroom_tool_calls(response)
if not tool_calls:
return False, {}
return True, {"tool_calls": tool_calls}
async def async_build_agentic_loop_plan(
self,
tools: dict,
model: str,
messages: list[dict],
response: Any,
anthropic_messages_provider_config: Any,
anthropic_messages_optional_request_params: dict,
logging_obj: Any,
stream: bool,
kwargs: dict,
) -> AgenticLoopPlan:
tool_calls: list[dict[str, object]] = tools.get("tool_calls", []) # type: ignore[assignment]
self._prune_expired_hashes()
call_id = _resolve_call_id(logging_obj, kwargs)
valid_hashes = self._issued_hashes_by_call_id.get(call_id, (frozenset(), 0.0))[0] if call_id else frozenset()
retrieved: list[tuple[dict[str, object], str]] = []
for tc in tool_calls:
arguments = tc.get("arguments", {})
hash_value = arguments.get("hash", "") if isinstance(arguments, dict) else ""
query = arguments.get("query") if isinstance(arguments, dict) else None
# A hash is only honored if it was issued by *this request's own*
# Headroom /v1/compress call, scoped by litellm_call_id. Scoping by
# message text alone is forgeable -- an attacker can plant a
# hash-shaped string in their own prompt, and a hash issued for one
# request would validate for any other request that echoes it back.
if str(hash_value) not in valid_hashes:
verbose_proxy_logger.warning(
"Headroom CCR: rejecting hash=%s not produced by current request compression",
hash_value,
)
content = f"[Headroom: hash={hash_value} was not produced by the current request]"
else:
content = await self._call_retrieve(
hash_value=str(hash_value),
query=str(query) if query else None,
)
verbose_proxy_logger.debug("Headroom CCR: retrieved hash=%s (%d chars)", hash_value, len(content))
retrieved.append((tc, content))
if _is_responses_api_response(response):
follow_up_messages = list(messages) + _build_responses_followup_items(retrieved)
elif _is_anthropic_messages_response(response):
follow_up_messages = list(messages) + _build_anthropic_followup_messages(retrieved)
else:
assistant_message = _build_assistant_message_from_response(response)
tool_results = [
{"role": "tool", "tool_call_id": tc.get("id"), "content": content} for tc, content in retrieved
]
follow_up_messages = list(messages) + [assistant_message] + tool_results
max_tokens: Optional[int] = anthropic_messages_optional_request_params.get("max_tokens") or kwargs.get(
"max_tokens"
)
optional_params_without_max_tokens = {
k: v for k, v in anthropic_messages_optional_request_params.items() if k != "max_tokens"
}
full_model_name = model
if logging_obj is not None:
agentic_params = getattr(logging_obj, "model_call_details", {}).get("agentic_loop_params", {})
candidate = agentic_params.get("model", model)
if isinstance(candidate, str) and candidate:
full_model_name = candidate
return AgenticLoopPlan(
run_agentic_loop=True,
request_patch=AgenticLoopRequestPatch(
model=full_model_name,
messages=follow_up_messages,
max_tokens=max_tokens,
optional_params=optional_params_without_max_tokens,
kwargs={
k: v for k, v in kwargs.items() if not k.startswith("_headroom") and k != "litellm_logging_obj"
},
),
metadata={"tool_type": "headroom_ccr"},
)
@staticmethod
def get_config_model() -> type[GuardrailConfigModel[object]] | None:

View file

@ -153,6 +153,7 @@ _UNTRUSTED_ROOT_CONTROL_FIELDS = (
"_code_interpreter_interception_active",
"_code_interpreter_interception_converted_stream",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"max_agentic_loops",
)

View file

@ -15,6 +15,7 @@ from typing import List, Optional
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import BaseModel
import litellm
from litellm._logging import verbose_proxy_logger
@ -32,10 +33,26 @@ from litellm.repositories.table_repositories import EndUserRepository
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.customer_endpoints import (
BlockUsersResponse,
CustomerResponse,
DeleteCustomersResponse,
UnblockUsersResponse,
)
router = APIRouter()
def _to_customer_response(record: BaseModel) -> CustomerResponse:
"""Validate a raw end-user DB row into the typed customer response.
object_permission reverse relations and the budget's audit fields are
dropped here by the response model's field set, so callers need no manual
cleanup.
"""
return CustomerResponse.model_validate(record.model_dump())
@router.post(
"/end_user/block",
tags=["Customer Management"],
@ -46,6 +63,7 @@ router = APIRouter()
"/customer/block",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=BlockUsersResponse,
)
async def block_user(data: BlockUsers):
"""
@ -100,6 +118,7 @@ async def block_user(data: BlockUsers):
"/customer/unblock",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=UnblockUsersResponse,
)
async def unblock_user(data: BlockUsers):
"""
@ -213,11 +232,12 @@ async def _handle_customer_object_permission_update(
"/customer/new",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=CustomerResponse,
)
async def new_end_user(
data: NewCustomerRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
) -> CustomerResponse:
"""
Allow creating a new Customer
@ -370,20 +390,7 @@ async def new_end_user(
include={"litellm_budget_table": True, "object_permission": True},
)
# Convert to dict and clean up recursive fields
response_dict = end_user_record.model_dump()
if response_dict.get("object_permission"):
# Remove reverse relations from object_permission
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
response_dict["object_permission"].pop(field, None)
return response_dict
return _to_customer_response(end_user_record)
except Exception as e:
verbose_proxy_logger.exception(
"litellm.proxy.management_endpoints.customer_endpoints.new_end_user(): Exception occured - {}".format(
@ -404,7 +411,7 @@ async def new_end_user(
"/customer/info",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=LiteLLM_EndUserTable,
response_model=CustomerResponse,
)
@router.get(
"/end_user/info",
@ -414,7 +421,7 @@ async def new_end_user(
)
async def end_user_info(
end_user_id: str = fastapi.Query(description="End User ID in the request parameters"),
):
) -> CustomerResponse:
"""
Get information about an end-user. An `end_user` is a customer (external user) of the proxy.
@ -449,20 +456,7 @@ async def end_user_info(
param="end_user_id",
)
# Convert to dict and clean up recursive fields
response_dict = user_info.model_dump(exclude_none=True)
if response_dict.get("object_permission"):
# Remove reverse relations from object_permission
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
response_dict["object_permission"].pop(field, None)
return response_dict
return _to_customer_response(user_info)
except Exception as e:
verbose_proxy_logger.exception(
@ -477,6 +471,7 @@ async def end_user_info(
"/customer/update",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=CustomerResponse,
)
@router.post(
"/end_user/update",
@ -487,7 +482,7 @@ async def end_user_info(
async def update_end_user(
data: UpdateCustomerRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
) -> CustomerResponse:
"""
Example curl
@ -641,20 +636,7 @@ async def update_end_user(
raise ValueError(f"Failed updating customer data. User ID does not exist passed user_id={data.user_id}")
verbose_proxy_logger.debug(f"received response from updating prisma client. response={response}")
# Convert to dict and clean up recursive fields
response_dict = response.model_dump()
if response_dict.get("object_permission"):
# Remove reverse relations from object_permission
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
response_dict["object_permission"].pop(field, None)
return response_dict
return _to_customer_response(response)
else:
raise ValueError(f"user_id is required, passed user_id = {data.user_id}")
@ -671,6 +653,7 @@ async def update_end_user(
"/customer/delete",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=DeleteCustomersResponse,
)
@router.post(
"/end_user/delete",
@ -681,7 +664,7 @@ async def update_end_user(
async def delete_end_user(
data: DeleteCustomerRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
) -> DeleteCustomersResponse:
"""
Delete multiple end-users.
@ -728,10 +711,10 @@ async def delete_end_user(
where={"user_id": {"in": data.user_ids}}
)
verbose_proxy_logger.debug(f"received response from updating prisma client. response={response}")
return {
"deleted_customers": response,
"message": "Successfully deleted customers with ids: " + str(data.user_ids),
}
return DeleteCustomersResponse(
deleted_customers=response,
message="Successfully deleted customers with ids: " + str(data.user_ids),
)
else:
raise ValueError(f"user_id is required, passed user_id = {data.user_ids}")
@ -747,7 +730,7 @@ async def delete_end_user(
"/customer/list",
tags=["Customer Management"],
dependencies=[Depends(user_api_key_auth)],
response_model=List[LiteLLM_EndUserTable],
response_model=List[CustomerResponse],
)
@router.get(
"/end_user/list",
@ -758,7 +741,7 @@ async def delete_end_user(
async def list_end_user(
http_request: Request,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
) -> List[CustomerResponse]:
"""
[Admin-only] List all available customers
@ -791,21 +774,7 @@ async def list_end_user(
include={"litellm_budget_table": True, "object_permission": True}
)
returned_response: List[LiteLLM_EndUserTable] = []
for item in response:
item_dict = item.model_dump()
# Remove reverse relations from object_permission
if item_dict.get("object_permission"):
for field in [
"teams",
"verification_tokens",
"organizations",
"users",
"end_users",
]:
item_dict["object_permission"].pop(field, None)
returned_response.append(LiteLLM_EndUserTable(**item_dict))
return returned_response
return [_to_customer_response(item) for item in response]
except Exception as e:
verbose_proxy_logger.exception(

View file

@ -965,6 +965,23 @@ async def _common_key_generation_helper(
is_proxy_admin=_is_proxy_admin_caller,
)
# Merge default_key_generate_params.object_permission in *after* the team-scope
# checks above, so an admin-configured default (e.g. vector_stores, search_tools)
# is never mistaken for a caller-requested permission and rejected by those
# non-admin/no-team checks. Only fields the caller left unset are filled in.
_default_object_permission = (
litellm.default_key_generate_params.get("object_permission")
if litellm.default_key_generate_params is not None
else None
)
if isinstance(_default_object_permission, dict):
_caller_object_permission = data_json.get("object_permission")
if _caller_object_permission is None:
data_json["object_permission"] = dict(_default_object_permission)
elif isinstance(_caller_object_permission, dict):
for _op_field, _op_default_value in _default_object_permission.items():
_caller_object_permission.setdefault(_op_field, _op_default_value)
data_json = await _set_object_permission(
data_json=data_json,
prisma_client=prisma_client,

View file

@ -425,7 +425,10 @@ from litellm.proxy.management_endpoints.user_agent_analytics_endpoints import (
from litellm.proxy.management_endpoints.workflow_management_endpoints import (
router as workflow_management_router,
)
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
create_object_audit_log,
)
from litellm.proxy.memory.memory_endpoints import router as memory_router
from litellm.proxy.plugin_routes import (
router as plugin_router,
@ -470,6 +473,7 @@ from litellm.proxy.response_api_endpoints.endpoints import router as response_ro
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.search_endpoints.endpoints import router as search_router
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_management_router,
)
@ -2263,111 +2267,133 @@ async def increment_spend_counters(
budget_reservation["finalized"] = True
return
if token is not None:
# token arrives pre-hashed from metadata["user_api_key"] (auth flow
cost: float = response_cost
async def _key_scope(key_token: str) -> None:
# key_token arrives pre-hashed from metadata["user_api_key"] (auth flow
# hashes raw "sk-..." keys before they reach the callback). The
# startswith("sk-") check is a safety net matching update_cache —
# if a raw key somehow arrives, hash it; otherwise use as-is to
# avoid double-hashing (budget checks read valid_token.token which
# is single-hashed).
hashed_token = hash_token(token=token) if isinstance(token, str) and token.startswith("sk-") else token
hashed_token = (
hash_token(token=key_token) if isinstance(key_token, str) and key_token.startswith("sk-") else key_token
)
key_counter_key = f"spend:key:{hashed_token}"
if key_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=key_counter_key,
source_cache_key=hashed_token,
increment=response_cost,
increment=cost,
)
# Increment per-window budget counters for multi-budget keys
key_obj = await user_api_key_cache.async_get_cache(key=hashed_token)
if key_obj is not None:
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
)
if isinstance(key_budget_limits, str):
key_budget_limits = json.loads(key_budget_limits)
if isinstance(key_budget_limits, list):
for window in key_budget_limits:
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
if key_window_counter not in reserved_counter_keys:
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
)
if key_obj is None:
return
key_budget_limits = getattr(key_obj, "budget_limits", None) or (
key_obj.get("budget_limits") if isinstance(key_obj, dict) else None
)
if isinstance(key_budget_limits, str):
key_budget_limits = json.loads(key_budget_limits)
if not isinstance(key_budget_limits, list):
return
for window in key_budget_limits:
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
key_window_counter = f"spend:key:{hashed_token}:window:{duration}"
if key_window_counter not in reserved_counter_keys:
await _init_and_increment_window_spend_counter(
counter_key=key_window_counter,
entity_type="Key",
entity_id=hashed_token,
window_start=get_budget_window_start(window),
increment=cost,
)
await _init_and_increment_window_spend_counter(
counter_key=key_window_counter,
entity_type="Key",
entity_id=hashed_token,
window_start=get_budget_window_start(window),
increment=response_cost,
)
if team_id is not None:
team_counter_key = f"spend:team:{team_id}"
async def _team_scope(scope_team_id: str) -> None:
team_counter_key = f"spend:team:{scope_team_id}"
if team_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_counter_key,
source_cache_key=f"team_id:{team_id}",
increment=response_cost,
source_cache_key=f"team_id:{scope_team_id}",
increment=cost,
)
# Increment per-window budget counters for multi-budget teams
team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{team_id}")
if team_obj is not None:
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{scope_team_id}")
if team_obj is None:
return
team_budget_limits = getattr(team_obj, "budget_limits", None) or (
team_obj.get("budget_limits") if isinstance(team_obj, dict) else None
)
if isinstance(team_budget_limits, str):
team_budget_limits = json.loads(team_budget_limits)
if not isinstance(team_budget_limits, list):
return
for window in team_budget_limits:
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
team_window_counter = f"spend:team:{scope_team_id}:window:{duration}"
if team_window_counter not in reserved_counter_keys:
await _init_and_increment_window_spend_counter(
counter_key=team_window_counter,
entity_type="Team",
entity_id=scope_team_id,
window_start=get_budget_window_start(window),
increment=cost,
)
async def _team_member_scope(scope_user_id: str, scope_team_id: str) -> None:
team_member_counter_key = f"spend:team_member:{scope_user_id}:{scope_team_id}"
if team_member_counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=team_member_counter_key,
source_cache_key=f"team_membership:{scope_user_id}:{scope_team_id}",
increment=cost,
)
async def _user_scope(scope_user_id: str) -> None:
user_counter_key = f"spend:user:{scope_user_id}"
if user_counter_key in reserved_counter_keys:
return
await _init_and_increment_spend_counter(
counter_key=user_counter_key,
source_cache_key=scope_user_id,
increment=cost,
)
scope_coros = tuple(
coro
for coro in (
_key_scope(token) if token is not None else None,
_team_scope(team_id) if team_id is not None else None,
_team_member_scope(user_id, team_id) if user_id is not None and team_id is not None else None,
_user_scope(user_id) if user_id is not None else None,
_increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
tags=tags,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
if isinstance(team_budget_limits, str):
team_budget_limits = json.loads(team_budget_limits)
if isinstance(team_budget_limits, list):
for window in team_budget_limits:
duration = window["budget_duration"] if isinstance(window, dict) else window.budget_duration
team_window_counter = f"spend:team:{team_id}:window:{duration}"
if team_window_counter not in reserved_counter_keys:
from litellm.proxy.spend_tracking.budget_reservation import (
get_budget_window_start,
)
await _init_and_increment_window_spend_counter(
counter_key=team_window_counter,
entity_type="Team",
entity_id=team_id,
window_start=get_budget_window_start(window),
increment=response_cost,
)
if user_id is not None and team_id is not None:
team_member_counter_key = f"spend:team_member:{user_id}:{team_id}"
if team_member_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=team_member_counter_key,
source_cache_key=f"team_membership:{user_id}:{team_id}",
increment=response_cost,
if end_user_id is not None or tags is not None
else None,
_increment_org_spend_counter(
org_id=org_id,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
if user_id is not None:
user_counter_key = f"spend:user:{user_id}"
if user_counter_key not in reserved_counter_keys:
await _init_and_increment_spend_counter(
counter_key=user_counter_key,
source_cache_key=user_id,
increment=response_cost,
)
await _increment_end_user_and_tag_spend_counters(
end_user_id=end_user_id,
tags=tags,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
if org_id is not None
else None,
)
if coro is not None
)
await _increment_org_spend_counter(
org_id=org_id,
response_cost=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
# return_exceptions so a failing scope does not leave its siblings running
# as orphaned tasks that race the caller's reservation-counter invalidation;
# all scopes settle, then the first error propagates as before.
scope_results = await asyncio.gather(*scope_coros, return_exceptions=True)
scope_errors = [r for r in scope_results if isinstance(r, BaseException)]
if scope_errors:
raise scope_errors[0]
if budget_reservation is not None:
budget_reservation["finalized"] = True
@ -6288,6 +6314,10 @@ class ProxyConfig:
"litellm.proxy.proxy_server.py::ProxyConfig:_init_mcp_servers_in_db - {}".format(str(e))
)
async def init_mcp_servers_from_db(self) -> None:
if self._should_load_db_object(object_type="mcp"):
await self._init_mcp_servers_in_db()
async def _init_agents_in_db(self, prisma_client: PrismaClient):
from litellm.proxy.agent_endpoints.agent_registry import (
global_agent_registry as AGENT_REGISTRY,
@ -7535,6 +7565,9 @@ class ProxyStartupEvent:
)
await proxy_config.get_credentials(prisma_client=prisma_client)
if store_model_in_db is not True:
await proxy_config.init_mcp_servers_from_db()
await cls._initialize_slack_alerting_jobs(
scheduler=scheduler,
general_settings=general_settings,
@ -13936,6 +13969,7 @@ async def update_config(
# effect of auto-enabling slack alerting.
if config_info.general_settings is not None:
existing = await _read_section("general_settings")
before_general_settings = copy.deepcopy(existing)
updates = config_info.general_settings.dict(exclude_none=True)
for k, v in updates.items():
if k == "alert_to_webhook_url":
@ -13945,6 +13979,11 @@ async def update_config(
existing["alerting"].append("slack")
existing[k] = v
await _upsert_section("general_settings", existing)
asyncio.create_task(
create_config_audit_log(
"general_settings", "updated", before_general_settings, existing, user_api_key_dict
)
)
# environment_variables: idempotently encrypt the request values
# (plaintext on first write, OR ciphertext the UI read back via
@ -13953,10 +13992,16 @@ async def update_config(
# their stored ciphertext byte-for-byte.
if config_info.environment_variables is not None:
existing = await _read_section("environment_variables")
before_environment_variables = copy.deepcopy(existing)
existing.update(
proxy_config._encrypt_env_variables_for_db(environment_variables=config_info.environment_variables)
)
await _upsert_section("environment_variables", existing)
asyncio.create_task(
create_config_audit_log(
"environment_variables", "updated", before_environment_variables, existing, user_api_key_dict
)
)
# litellm_settings: merge existing + request, request wins (matching
# router_settings semantics — the caller's value for any given key is
@ -13968,6 +14013,7 @@ async def update_config(
# entries that delete_callback (lowercase lookup) cannot find.
if config_info.litellm_settings is not None:
existing = await _read_section("litellm_settings")
before_litellm_settings = copy.deepcopy(existing)
updated_litellm_settings = dict(config_info.litellm_settings)
incoming_cb = updated_litellm_settings.get("success_callback")
@ -13989,12 +14035,24 @@ async def update_config(
merged["success_callback"] = list(set(incoming_cb))
await _upsert_section("litellm_settings", merged)
asyncio.create_task(
create_config_audit_log(
"litellm_settings", "updated", before_litellm_settings, merged, user_api_key_dict
)
)
# router_settings: merge existing + request, request wins.
if config_info.router_settings is not None:
existing = await _read_section("router_settings")
before_router_settings = copy.deepcopy(existing)
updates = config_info.router_settings.dict(exclude_none=True)
await _upsert_section("router_settings", {**existing, **updates})
new_router_settings = {**existing, **updates}
await _upsert_section("router_settings", new_router_settings)
asyncio.create_task(
create_config_audit_log(
"router_settings", "updated", before_router_settings, new_router_settings, user_api_key_dict
)
)
await proxy_config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)
@ -14126,6 +14184,8 @@ async def update_config_general_settings(
else:
general_settings = dict(db_general_settings.param_value)
before_general_settings = copy.deepcopy(general_settings)
## update db
field_value = data.field_value
@ -14145,6 +14205,11 @@ async def update_config_general_settings(
},
)
await invalidate_config_param("general_settings")
asyncio.create_task(
create_config_audit_log(
"general_settings", "updated", before_general_settings, general_settings, user_api_key_dict
)
)
if data.field_name == "plugins":
register_plugins_from_config(general_settings)
@ -14205,6 +14270,47 @@ def _redact_general_setting_value(field_name: str, value: JsonValue, is_full_adm
return value
def _dump_redacted_config(value: Optional[JsonValue], *, redact_all_values: bool = False) -> Optional[str]:
# `default=str` matches the sibling audit-log serializers in
# team_endpoints.py and the LiteLLM_AuditLogs validator, so a YAML-loaded
# value with a non-JSON-native leaf (datetime, custom object) cannot turn
# an audit write into a 500.
if value is None:
return None
if redact_all_values and isinstance(value, dict):
return json.dumps({key: "REDACTED" for key in value}, default=str)
return json.dumps(_redact_secret_values_in_obj(value), default=str)
async def create_config_audit_log(
param_name: str,
action: AUDIT_ACTIONS,
before_value: Optional[JsonValue],
after_value: Optional[JsonValue],
user_api_key_dict: UserAPIKeyAuth,
table_name: LitellmTableNames = LitellmTableNames.CONFIG_TABLE_NAME,
) -> None:
"""Record a system-wide settings change in LiteLLM_AuditLog.
Secret leaves are redacted before the row is written. environment_variables
hold arbitrary credentials under non-secret-looking uppercase keys (e.g.
DATABASE_URL), so every value in that section is redacted rather than
relying on key-name matching; other sections reuse the same matcher
/config/field/info applies for non-admins.
"""
redact_all_values = param_name == "environment_variables"
await create_object_audit_log(
object_id=param_name,
action=action,
table_name=table_name,
before_value=_dump_redacted_config(before_value, redact_all_values=redact_all_values),
after_value=_dump_redacted_config(after_value, redact_all_values=redact_all_values),
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
litellm_proxy_admin_name=LITELLM_PROXY_ADMIN_NAME,
)
@router.get(
"/config/field/info",
tags=["config.yaml"],
@ -14488,6 +14594,8 @@ async def delete_config_general_settings(
else:
general_settings = dict(db_general_settings.param_value)
before_general_settings = copy.deepcopy(general_settings)
## update db
general_settings.pop(data.field_name, None)
@ -14503,6 +14611,11 @@ async def delete_config_general_settings(
},
)
await invalidate_config_param("general_settings")
asyncio.create_task(
create_config_audit_log(
"general_settings", "deleted", before_general_settings, general_settings, user_api_key_dict
)
)
return response
@ -14560,6 +14673,8 @@ async def delete_callback(
detail={"error": f"Callback '{callback_name}' not found in active configuration"},
)
before_success_callbacks = list(success_callbacks)
# Remove callback from success_callback list
success_callbacks.remove(callback_name)
config.setdefault("litellm_settings", {})["success_callback"] = success_callbacks
@ -14567,6 +14682,16 @@ async def delete_callback(
# Save the updated configuration
await proxy_config.save_config(new_config=config)
asyncio.create_task(
create_config_audit_log(
"litellm_settings",
"deleted",
{"success_callback": before_success_callbacks},
{"success_callback": success_callbacks},
user_api_key_dict,
)
)
# Restart the proxy to apply changes
await proxy_config.add_deployment(prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj)

View file

@ -279,6 +279,7 @@ model LiteLLM_ObjectPermissionTable {
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
mcp_tool_search_enabled Boolean?
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -1,4 +1,5 @@
#### CRUD ENDPOINTS for UI Settings #####
import asyncio
import json
from typing import Any, Dict, List, Optional, Set, Tuple, Type, Union
from urllib.parse import urlparse
@ -322,8 +323,12 @@ async def get_allowed_ips():
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
)
async def add_allowed_ip(ip_address: IPAddress):
async def add_allowed_ip(
ip_address: IPAddress,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.proxy_server import (
create_config_audit_log,
general_settings,
prisma_client,
proxy_config,
@ -355,11 +360,22 @@ async def add_allowed_ip(ip_address: IPAddress):
if "allowed_ips" not in config["general_settings"]:
config["general_settings"]["allowed_ips"] = []
before_allowed_ips = list(config["general_settings"]["allowed_ips"])
if ip_address.ip not in config["general_settings"]["allowed_ips"]:
config["general_settings"]["allowed_ips"].append(ip_address.ip)
await proxy_config.save_config(new_config=config)
asyncio.create_task(
create_config_audit_log(
param_name="general_settings",
action="updated",
before_value={"allowed_ips": before_allowed_ips},
after_value={"allowed_ips": config["general_settings"]["allowed_ips"]},
user_api_key_dict=user_api_key_dict,
)
)
return {
"message": f"IP {ip_address.ip} address added successfully",
"status": "success",
@ -371,8 +387,15 @@ async def add_allowed_ip(ip_address: IPAddress):
tags=["Budget & Spend Tracking"],
dependencies=[Depends(user_api_key_auth)],
)
async def delete_allowed_ip(ip_address: IPAddress):
from litellm.proxy.proxy_server import general_settings, proxy_config
async def delete_allowed_ip(
ip_address: IPAddress,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.proxy_server import (
create_config_audit_log,
general_settings,
proxy_config,
)
_allowed_ips: List = general_settings.get("allowed_ips", [])
if ip_address.ip in _allowed_ips:
@ -390,11 +413,22 @@ async def delete_allowed_ip(ip_address: IPAddress):
if "allowed_ips" not in config["general_settings"]:
config["general_settings"]["allowed_ips"] = []
before_allowed_ips = list(config["general_settings"]["allowed_ips"])
if ip_address.ip in config["general_settings"]["allowed_ips"]:
config["general_settings"]["allowed_ips"].remove(ip_address.ip)
await proxy_config.save_config(new_config=config)
asyncio.create_task(
create_config_audit_log(
param_name="general_settings",
action="deleted",
before_value={"allowed_ips": before_allowed_ips},
after_value={"allowed_ips": config["general_settings"]["allowed_ips"]},
user_api_key_dict=user_api_key_dict,
)
)
return {"message": f"IP {ip_address.ip} deleted successfully", "status": "success"}
@ -553,6 +587,7 @@ async def _update_litellm_setting(
settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings],
settings_key: str,
success_message: str,
user_api_key_dict: UserAPIKeyAuth,
):
"""
Common utility function to update `litellm_settings` in both memory and config.
@ -561,8 +596,13 @@ async def _update_litellm_setting(
settings: The settings object to update
settings_key: The key in litellm_settings to update
success_message: Message to return on success
user_api_key_dict: The acting admin, recorded as the audit-log actor.
"""
from litellm.proxy.proxy_server import proxy_config, store_model_in_db
from litellm.proxy.proxy_server import (
create_config_audit_log,
proxy_config,
store_model_in_db,
)
if store_model_in_db is not True:
raise HTTPException(
@ -576,6 +616,7 @@ async def _update_litellm_setting(
# because get_config() may overwrite litellm.<key> with stale DB values
# via LITELLM_SETTINGS_SAFE_DB_OVERRIDES.
config = await proxy_config.get_config()
before_value = config.get("litellm_settings", {}).get(settings_key)
# Update the in-memory settings (after get_config to avoid stale override)
setattr(litellm, settings_key, in_memory_var)
@ -589,6 +630,20 @@ async def _update_litellm_setting(
# Save the updated config
await proxy_config.save_config(new_config=config)
# Fire-and-forget so an audit-log failure (transient DB blip, etc.)
# never surfaces as a 500 after save_config has already committed,
# matching the create_object_audit_log pattern used elsewhere
# (e.g. model_management_endpoints).
asyncio.create_task(
create_config_audit_log(
param_name=settings_key,
action="updated",
before_value=before_value,
after_value=in_memory_var,
user_api_key_dict=user_api_key_dict,
)
)
return {
"message": success_message,
"status": "success",
@ -619,6 +674,7 @@ async def update_internal_user_settings(
settings=settings,
settings_key="default_internal_user_params",
success_message="Internal user settings updated successfully",
user_api_key_dict=user_api_key_dict,
)
@ -627,7 +683,10 @@ async def update_internal_user_settings(
tags=["SSO Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_default_team_settings(settings: DefaultTeamSSOParams):
async def update_default_team_settings(
settings: DefaultTeamSSOParams,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Update the default team parameters for SSO users.
These settings will be applied to new teams created from SSO.
@ -636,6 +695,7 @@ async def update_default_team_settings(settings: DefaultTeamSSOParams):
settings=settings,
settings_key="default_team_params",
success_message="Default team settings updated successfully",
user_api_key_dict=user_api_key_dict,
)
@ -746,7 +806,10 @@ async def get_sso_settings():
tags=["SSO Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_sso_settings(sso_config: SSOConfig):
async def update_sso_settings(
sso_config: SSOConfig,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Update SSO configuration by saving to the dedicated SSO table.
"""
@ -754,6 +817,7 @@ async def update_sso_settings(sso_config: SSOConfig):
import os
from litellm.proxy.proxy_server import (
create_config_audit_log,
prisma_client,
proxy_config,
store_model_in_db,
@ -786,6 +850,20 @@ async def update_sso_settings(sso_config: SSOConfig):
"proxy_base_url": "PROXY_BASE_URL",
}
# Read the existing SSO row first so the audit log captures a real
# before/after diff. Stored values are encrypted; decrypt them so the
# before-snapshot has the same shape as after_value, and rely on
# create_config_audit_log's secret-name redaction to mask the
# *_client_secret fields before the audit row is written.
existing_sso_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"})
before_sso_data: Optional[Dict[str, Any]] = None
if existing_sso_record and existing_sso_record.sso_settings:
stored = existing_sso_record.sso_settings
if isinstance(stored, str):
stored = json.loads(stored)
if isinstance(stored, dict):
before_sso_data = proxy_config._decrypt_db_variables(stored)
# Load existing config
config = await proxy_config.get_config()
@ -824,6 +902,17 @@ async def update_sso_settings(sso_config: SSOConfig):
},
)
asyncio.create_task(
create_config_audit_log(
param_name="sso_config",
action="updated",
before_value=before_sso_data,
after_value=sso_data,
user_api_key_dict=user_api_key_dict,
table_name=LitellmTableNames.SSO_CONFIG_TABLE_NAME,
)
)
# Remove SSO-related env vars from config.environment_variables
try:
env_var_entry = await ConfigRepository(prisma_client).table.find_unique(
@ -917,14 +1006,21 @@ def _validate_public_image_url(value: Optional[str], field_name: str) -> None:
tags=["UI Theme Settings"],
dependencies=[Depends(user_api_key_auth)],
)
async def update_ui_theme_settings(theme_config: UIThemeConfig):
async def update_ui_theme_settings(
theme_config: UIThemeConfig,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
Update UI theme configuration.
Updates logo settings for the admin UI.
"""
import os
from litellm.proxy.proxy_server import proxy_config, store_model_in_db
from litellm.proxy.proxy_server import (
create_config_audit_log,
proxy_config,
store_model_in_db,
)
_validate_public_image_url(theme_config.logo_url, "logo_url")
_validate_public_image_url(theme_config.favicon_url, "favicon_url")
@ -937,6 +1033,7 @@ async def update_ui_theme_settings(theme_config: UIThemeConfig):
# Load existing config
config = await proxy_config.get_config()
before_theme = config.get("litellm_settings", {}).get("ui_theme_config")
# Update config with UI theme settings
if "general_settings" not in config:
@ -1003,6 +1100,16 @@ async def update_ui_theme_settings(theme_config: UIThemeConfig):
# Save the updated config
await proxy_config.save_config(new_config=stored_config)
asyncio.create_task(
create_config_audit_log(
param_name="ui_theme_config",
action="updated",
before_value=before_theme,
after_value=theme_data,
user_api_key_dict=user_api_key_dict,
)
)
return {
"message": "UI theme settings updated successfully.",
"status": "success",
@ -1057,6 +1164,7 @@ async def update_mcp_semantic_filter_settings(
settings=settings,
settings_key="mcp_semantic_tool_filter",
success_message="MCP Semantic Filter settings updated successfully. Changes will be applied across all pods within 10 seconds.",
user_api_key_dict=user_api_key_dict,
)
try:
from litellm.proxy.proxy_server import prisma_client, proxy_config
@ -1174,7 +1282,11 @@ async def update_ui_settings(
Update UI-specific configuration flags.
Only proxy admins are allowed to modify these settings.
"""
from litellm.proxy.proxy_server import prisma_client, store_model_in_db
from litellm.proxy.proxy_server import (
create_config_audit_log,
prisma_client,
store_model_in_db,
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(status_code=403, detail="Only proxy admins can update UI settings.")
@ -1256,6 +1368,17 @@ async def update_ui_settings(
sanitized = {k: v for k, v in ui_settings.items() if k in ALLOWED_UI_SETTINGS_FIELDS}
await user_api_key_cache.async_set_cache(key=UI_SETTINGS_CACHE_KEY, value=sanitized, ttl=UI_SETTINGS_CACHE_TTL)
asyncio.create_task(
create_config_audit_log(
param_name="ui_settings",
action="updated",
before_value=existing,
after_value=ui_settings,
user_api_key_dict=user_api_key_dict,
table_name=LitellmTableNames.UI_SETTINGS_TABLE_NAME,
)
)
return {
"message": "UI settings updated successfully",
"status": "success",

View file

@ -23,7 +23,9 @@ from typing import (
Dict,
List,
Literal,
Mapping,
Optional,
Sequence,
Tuple,
Union,
cast,
@ -5194,8 +5196,10 @@ class ProxyUpdateSpend:
for j in range(0, len(logs_to_process), BATCH_SIZE):
batch = logs_to_process[j : j + BATCH_SIZE]
batch_with_dates = [prisma_client.jsonify_object({**entry}) for entry in batch]
await SpendLogsRepository(prisma_client).table.create_many(
data=batch_with_dates, skip_duplicates=True
await _create_spend_logs_with_poison_isolation(
SpendLogsRepository(prisma_client),
batch_with_dates,
MAX_SPEND_LOG_ISOLATION_ATTEMPTS_PER_BATCH,
)
verbose_proxy_logger.debug(f"Flushed {len(batch)} logs to the DB.")
# Explicitly clear batch memory
@ -5462,6 +5466,65 @@ async def _monitor_spend_logs_queue(
await asyncio.sleep(current_interval)
MAX_SPEND_LOG_ISOLATION_ATTEMPTS_PER_BATCH = 256
async def _create_spend_logs_with_poison_isolation(
repo: SpendLogsRepository,
rows: Sequence[Mapping[str, object]],
attempts_left: int,
) -> int:
"""Write spend-log rows, isolating any row Postgres rejects on its data.
``create_many`` writes the whole batch in a single statement, so one row
carrying bytes Postgres refuses (a residual NUL byte is the canonical case)
fails the entire insert and drops every good row alongside it. On a genuine
data-layer rejection the batch is bisected so the good rows still persist
and only the offending row is dropped and logged. Transport failures,
including the "can't reach database server" outage that prisma mislabels as
a ``DataError``, are re-raised unchanged so the caller's connection-retry
path still runs.
``attempts_left`` is a hard ceiling on the number of ``create_many`` calls
the isolation may issue for this batch, so an authenticated caller flooding
poisoned rows cannot amplify one failed bulk insert into unbounded failed
inserts and log lines. It is checked before any insert (so an exhausted
budget never even attempts a write), decremented once per ``create_many``
call, and threaded through the recursion so the whole bisection shares one
allowance; total inserts are therefore bounded by the initial value
regardless of how many rows are poisoned. When it runs out the still-failing
remainder is dropped wholesale (the pre-existing drop-the-batch behavior)
under one log line. Returns the budget left after this subtree.
"""
if attempts_left <= 0:
spend_log_error(
"Spend tracking - dropping %d spend log rows without per-row isolation; "
"isolation attempt budget exhausted for this flush",
len(rows),
)
return 0
try:
await repo.table.create_many(data=rows, skip_duplicates=True)
return attempts_left - 1
except Exception as e:
if not PrismaDBExceptionHandler.is_prisma_data_error(e):
raise
if PrismaDBExceptionHandler.is_database_service_unavailable_error(e):
raise
if len(rows) == 1:
request_id = rows[0].get("request_id")
spend_log_error(
"Spend tracking - dropping spend log row Postgres rejected. request_id=%s error=%s",
request_id,
str(e),
exc=e,
)
return attempts_left - 1
mid = len(rows) // 2
remaining = await _create_spend_logs_with_poison_isolation(repo, rows[:mid], attempts_left - 1)
return await _create_spend_logs_with_poison_isolation(repo, rows[mid:], remaining)
def _raise_failed_update_spend_exception(e: Exception, start_time: float, proxy_logging_obj: ProxyLogging):
"""
Raise an exception for failed update spend logs

View file

@ -7,7 +7,7 @@ Use this to route requests between Teams
"""
import re
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union
from typing import TYPE_CHECKING, Any, Literal, Optional, Union
from litellm._logging import verbose_logger
from litellm.types.router import RouterErrors
@ -21,8 +21,8 @@ else:
def _is_valid_deployment_tag_regex(
tag_regexes: List[str],
header_strings: List[str],
tag_regexes: list[str],
header_strings: list[str],
) -> Optional[str]:
"""
Test compiled regex patterns against "Header-Name: value" strings.
@ -43,7 +43,7 @@ def _is_valid_deployment_tag_regex(
return None
def is_valid_deployment_tag(deployment_tags: List[str], request_tags: List[str], match_any: bool = True) -> bool:
def is_valid_deployment_tag(deployment_tags: list[str], request_tags: list[str], match_any: bool = True) -> bool:
"""
Check if a tag is valid, the matching can be either any or all based on `match_any` flag
"""
@ -71,10 +71,10 @@ def is_valid_deployment_tag(deployment_tags: List[str], request_tags: List[str],
def _match_deployment(
deployment: Any,
request_tags: Optional[List[str]],
header_strings: List[str],
request_tags: Optional[list[str]],
header_strings: list[str],
match_any: bool,
) -> Optional[Dict[str, str]]:
) -> Optional[dict[str, str]]:
"""
Determine whether *deployment* matches the current request.
@ -87,8 +87,8 @@ def _match_deployment(
ran and failed, so the regex cannot override strict-tag policy.
"""
litellm_params = deployment.get("litellm_params", {})
deployment_tags: Optional[List[str]] = litellm_params.get("tags")
deployment_tag_regex: Optional[List[str]] = litellm_params.get("tag_regex")
deployment_tags: Optional[list[str]] = litellm_params.get("tags")
deployment_tag_regex: Optional[list[str]] = litellm_params.get("tag_regex")
# 1. Exact tag match (existing behaviour).
if deployment_tags and request_tags:
@ -114,11 +114,46 @@ def _match_deployment(
return None
def _split_tags(tags: list[str]) -> tuple[list[str], list[str]]:
positive = [t for t in tags if not t.startswith("!")]
excluded = [tag[1:] for tag in tags if tag.startswith("!") and len(tag) > 1]
return positive, excluded
def _exclude_deployments(
deployments: Union[list[Any], dict[Any, Any]],
excluded_set: frozenset[str],
) -> list[Any]:
if not excluded_set:
return list(deployments)
return [d for d in deployments if not excluded_set.intersection(d.get("litellm_params", {}).get("tags") or [])]
def _require_candidates(
candidates: list[Any],
model: str,
request_tags: Any,
) -> list[Any]:
if not candidates:
raise ValueError(
f"{RouterErrors.no_deployments_with_tag_routing.value}. Passed model={model} and tags={request_tags}"
)
return candidates
def _ban_only_base_pool(
deployments: Union[list[Any], dict[Any, Any]],
) -> list[Any]:
# Mirrors untagged-request semantics so callers can't use !tags to escape the default pool.
defaults = [d for d in deployments if "default" in (d.get("litellm_params", {}).get("tags") or [])]
return defaults if defaults else list(deployments)
async def get_deployments_for_tag(
llm_router_instance: LitellmRouter,
model: str, # used to raise the correct error
healthy_deployments: Union[List[Any], Dict[Any, Any]],
request_kwargs: Optional[Dict[Any, Any]] = None,
healthy_deployments: Union[list[Any], dict[Any, Any]],
request_kwargs: Optional[dict[Any, Any]] = None,
metadata_variable_name: Literal["metadata", "litellm_metadata"] = "metadata",
):
"""
@ -136,13 +171,8 @@ async def get_deployments_for_tag(
)
return healthy_deployments
if healthy_deployments is None:
verbose_logger.debug("get_deployments_for_tag: healthy_deployments is None returning healthy_deployments")
return healthy_deployments
# Tag filtering applies only when there is at least one deployment to evaluate.
if isinstance(healthy_deployments, list) and len(healthy_deployments) == 0:
verbose_logger.debug("get_deployments_for_tag: empty candidate set; skipping tag filter")
if not healthy_deployments:
verbose_logger.debug("get_deployments_for_tag: empty or None healthy_deployments; skipping tag filter")
return healthy_deployments
verbose_logger.debug("request metadata: %s", request_kwargs.get(metadata_variable_name))
@ -154,30 +184,36 @@ async def get_deployments_for_tag(
# Build header strings for regex matching from what the proxy already stores.
# Currently we match against User-Agent; format matches "^User-Agent: claude-code/..."
user_agent = metadata.get("user_agent", "")
header_strings: List[str] = [f"User-Agent: {user_agent}"] if user_agent else []
header_strings: list[str] = [f"User-Agent: {user_agent}"] if user_agent else []
new_healthy_deployments: List[Any] = []
default_deployments: List[Any] = []
positive_tags, excluded_patterns = _split_tags(request_tags or [])
excluded_set = frozenset(excluded_patterns)
candidates = _exclude_deployments(healthy_deployments, excluded_set)
has_regex_deployments = any(d.get("litellm_params", {}).get("tag_regex") for d in candidates)
has_tag_filter = bool(positive_tags) or (bool(header_strings) and has_regex_deployments)
ban_only = bool(excluded_set) and not has_tag_filter
if ban_only:
pool = _exclude_deployments(_ban_only_base_pool(healthy_deployments), excluded_set)
return _require_candidates(pool, model, request_tags)
new_healthy_deployments: list[Any] = []
default_deployments: list[Any] = []
# Only activate header-based regex filtering when at least one deployment in
# the candidate set has tag_regex configured. This preserves existing
# behaviour for operators who use plain tags: a request that carries a
# User-Agent (all proxy requests do) but targets deployments with no
# tag_regex will continue to use the original tag-only code path.
has_regex_deployments = any(d.get("litellm_params", {}).get("tag_regex") for d in healthy_deployments)
has_tag_filter = bool(request_tags) or (bool(header_strings) and has_regex_deployments)
if has_tag_filter:
verbose_logger.debug(
"get_deployments_for_tag routing: request_tags=%s user_agent=%s",
request_tags,
user_agent,
)
for deployment in healthy_deployments:
for deployment in candidates:
deployment_tags = deployment.get("litellm_params", {}).get("tags")
match_result = _match_deployment(
deployment=deployment,
request_tags=request_tags,
request_tags=positive_tags,
header_strings=header_strings,
match_any=match_any,
)
@ -189,10 +225,6 @@ async def get_deployments_for_tag(
match_result["matched_via"],
match_result["matched_value"],
)
# Record provenance in metadata so it flows to SpendLogs.
# Written only for the first match — load balancer selects one
# deployment from new_healthy_deployments, so overwriting on
# subsequent matches would produce misleading observability data.
if "tag_routing" not in metadata:
metadata["tag_routing"] = {
"matched_deployment": deployment.get("model_name"),
@ -208,7 +240,8 @@ async def get_deployments_for_tag(
if len(new_healthy_deployments) == 0 and len(default_deployments) == 0:
raise ValueError(
f"{RouterErrors.no_deployments_with_tag_routing.value}. Passed model={model} and tags={request_tags}"
f"{RouterErrors.no_deployments_with_tag_routing.value}."
f" Passed model={model} and tags={request_tags}"
)
return new_healthy_deployments if len(new_healthy_deployments) > 0 else default_deployments
@ -231,9 +264,9 @@ async def get_deployments_for_tag(
def _get_tags_from_request_kwargs(
request_kwargs: Optional[Dict[Any, Any]] = None,
request_kwargs: Optional[dict[Any, Any]] = None,
metadata_variable_name: Literal["metadata", "litellm_metadata"] = "metadata",
) -> List[str]:
) -> list[str]:
"""
Helper to get tags from request kwargs

View file

@ -52,6 +52,13 @@ def _admin_config_fields_to_clear_on_base_override() -> List[str]:
"oci_tenancy",
"oci_key",
"oci_key_file",
# NVIDIA Riva fields — consumed by
# ``litellm/llms/nvidia_riva/audio_transcription/handler.py`` via
# optional_params and not declared on CredentialLiteLLMParams.
# Admin-pinned values must not flow through on a caller-redirected
# ``api_base`` for the same reason as the OCI entries above.
"nvcf_function_id",
"use_ssl",
]
return typed_fields + kwargs_only_fields

View file

@ -58,6 +58,7 @@ PROVIDERS: List[Dict] = [
"test_model": "claude-haiku-4-5-20251001",
"models": [
"claude-fable-5",
"claude-sonnet-5",
"claude-opus-4-8",
"claude-opus-4-7",
"claude-opus-4-6",

View file

@ -24,3 +24,4 @@ class ObjectPermissionDict(TypedDict, total=False):
agent_access_groups: Optional[list[str]]
models: Optional[list[str]]
search_tools: Optional[list[str]]
mcp_tool_search_enabled: Optional[bool]

View file

@ -1,7 +1,7 @@
from typing import Any, Dict, List, Literal, Optional, Union
from pydantic import BaseModel, ConfigDict, Field
from typing_extensions import TYPE_CHECKING, TypedDict
from typing_extensions import TypedDict
from litellm.types.llms.openai import (
AllMessageValues,
@ -60,6 +60,30 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
streaming_end_of_stream_only: Optional[bool] = Field(
default=None,
description=(
"If False (default when unset), the guardrail runs on sampled chunks during "
"the stream at the cadence set by streaming_sampling_rate, and an in-flight "
"BLOCKED stops further chunks from streaming. If True, the guardrail runs "
"once at end of stream over the assembled response; lower cost and latency, "
"but flagged content has already streamed to the client before the terminal "
"block. Defaults are applied in GenericGuardrailAPI.__init__ when None so "
"unset optional_params does not shadow top-level litellm_params."
),
)
streaming_sampling_rate: Optional[int] = Field(
default=None,
ge=1,
description=(
"When streaming_end_of_stream_only is False, the guardrail runs every Nth "
"streamed chunk. Ignored when streaming_end_of_stream_only is True. "
"Must be >= 1 when set. Defaults to 5 in GenericGuardrailAPI.__init__ "
"when None so unset optional_params does not shadow top-level litellm_params."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],

View file

@ -0,0 +1,30 @@
from typing import List, Optional
from pydantic import BaseModel, Field
from litellm.models.budget import LiteLLM_BudgetTableFull
from litellm.models.end_user import LiteLLM_EndUserTable
class CustomerResponse(LiteLLM_EndUserTable):
"""Customer object returned by the /customer read+write endpoints.
Nests the full budget response model so server-managed budget fields
(budget_reset_at, created_at) survive response_model filtering, rather than
the narrow write-allowlist shape LiteLLM_EndUserTable carries for internal use.
"""
litellm_budget_table: Optional[LiteLLM_BudgetTableFull] = None # pyright: ignore
class BlockUsersResponse(BaseModel):
blocked_users: List[LiteLLM_EndUserTable]
class UnblockUsersResponse(BaseModel):
blocked_users: List[str] = Field(description="User IDs that remain blocked after this unblock call")
class DeleteCustomersResponse(BaseModel):
deleted_customers: int
message: str

View file

@ -3062,6 +3062,7 @@ agentic_loop_internal_litellm_params = [
"max_agentic_loops",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_code_interpreter_interception_converted_stream",
]

View file

@ -1671,6 +1671,204 @@
"supports_max_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"global.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"au.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"jp.anthropic.claude-sonnet-5": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.65e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_native_structured_output": true,
"supports_max_reasoning_effort": true,
"supports_output_config": true,
"bedrock_output_config_effort_ceiling": "xhigh"
},
"anthropic.claude-sonnet-4-6": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
@ -2511,6 +2709,36 @@
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "azure_ai",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true
},
"azure_ai/claude-sonnet-4-6": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
@ -10245,6 +10473,40 @@
"supports_vision": true,
"supports_web_search": true
},
"claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "anthropic",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true,
"provider_specific_entry": {
"us": 1.1
},
"supports_output_config": true
},
"claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -35121,6 +35383,36 @@
"supports_tool_choice": true,
"supports_vision": true
},
"vertex_ai/claude-sonnet-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true
},
"vertex_ai/claude-sonnet-4-6": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
@ -42616,6 +42908,36 @@
"search_context_size_high": 0.035
}
},
"vertex_ai/claude-sonnet-5@default": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"supports_adaptive_thinking": true,
"supports_assistant_prefill": false,
"supports_computer_use": true,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_sampling_params": false,
"supports_tool_choice": true,
"supports_vision": true,
"supports_xhigh_reasoning_effort": true,
"supports_max_reasoning_effort": true
},
"vertex_ai/claude-sonnet-4-6@default": {
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.91.0"
version = "1.92.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@ -63,7 +63,7 @@ proxy = [
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.26.0,<2.0",
"litellm-proxy-extras==0.4.74",
"litellm-enterprise==0.1.44",
"litellm-enterprise==0.1.45",
"RestrictedPython>=8.1,<9.0",
"rich>=13.9.4,<14.0",
"polars>=1.38.1,<2.0",
@ -274,7 +274,7 @@ source-exclude = [
profile = "black"
[tool.commitizen]
version = "1.91.0"
version = "1.92.0"
version_files = [
"pyproject.toml:^version",
]

59
qa_sticky_session.sh Executable file
View file

@ -0,0 +1,59 @@
#!/usr/bin/env bash
# QA: code interpreter sandbox stickiness via metadata.session_id
# bash qa_sticky_session.sh
# LITELLM_BASE_URL=http://localhost:4000 LITELLM_KEY=sk-1234 bash qa_sticky_session.sh
set -euo pipefail
BASE="${LITELLM_BASE_URL:-http://localhost:4000}"
KEY="${LITELLM_KEY:-sk-1234}"
MODEL="${LITELLM_MODEL:-gpt-4o-mini}"
# proxy running at http://localhost:4000 (master key: sk-1234)
SESSION_A="qa-session-$(date +%s)-A"
SESSION_B="qa-session-$(date +%s)-B"
content() {
echo "$1" | python3 -c "import sys,json; d=json.load(sys.stdin); print(d.get('choices',[{}])[0].get('message',{}).get('content','<error>'))"
}
call() {
local session="${1:-}" code="$2" meta=""
[[ -n "$session" ]] && meta=", \"metadata\": {\"session_id\": \"$session\"}"
curl -s -X POST "$BASE/chat/completions" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $KEY" \
-d "{\"model\":\"$MODEL\"$meta,\"tools\":[{\"type\":\"code_interpreter\"}],\"messages\":[{\"role\":\"user\",\"content\":\"Run this Python code and tell me the result: $code\"}]}"
}
assert_match() {
local label="$1" body="$2" pattern="$3"
if echo "$body" | grep -qiE "$pattern"; then
echo "PASS $label"
else
echo "FAIL $label (expected /$pattern/)"
echo " $(content "$body")"
exit 1
fi
}
echo "=== Sticky Session Sandbox QA ==="
echo "base: $BASE session A: $SESSION_A session B: $SESSION_B"
echo
R=$(call "$SESSION_A" "x = 42; print(x)")
assert_match "same session_id reuses sandbox (set x=42)" "$R" "42"
R=$(call "$SESSION_A" "print(x)")
assert_match "same session_id keeps state (x still 42)" "$R" "42"
R=$(call "$SESSION_B" "print(x)")
assert_match "different session_id is isolated" "$R" "not defined|NameError|undefined|error"
R=$(call "" "y = 99; print(y)")
assert_match "no session_id runs code" "$R" "99"
R=$(call "" "print(y)")
assert_match "no session_id gets fresh sandbox each request" "$R" "not defined|NameError|undefined|error"
echo
echo "All checks passed."

View file

@ -279,6 +279,7 @@ model LiteLLM_ObjectPermissionTable {
blocked_tools String[] @default([]) // Tool names blocked for any key/team/user with this permission
mcp_toolsets String[] @default([]) // Toolset IDs granted to this key/team/user
search_tools String[] @default([]) // search_tool_name values this key/team/user may call
mcp_tool_search_enabled Boolean?
teams LiteLLM_TeamTable[]
projects LiteLLM_ProjectTable[]
verification_tokens LiteLLM_VerificationToken[]

View file

@ -34,5 +34,8 @@ cat <<EOF
These hooks enforce Conventional Commits and Conventional Branches.
Bypass with --no-verify when you need to (e.g. for emergency hotfixes).
The CI-equivalent lint is deliberately not installed as an auto-firing hook
(it can take minutes); run it on demand with 'make pre-commit' before committing.
To uninstall: git config --unset core.hooksPath
EOF

112
scripts/pre_commit_lint.sh Executable file
View file

@ -0,0 +1,112 @@
#!/usr/bin/env bash
#
# pre_commit_lint.sh — shift CI lint left. Run it (via `make pre-commit`) right
# before `git commit`; it inspects your staged files and runs only the matching
# gating CI checks, so a clean run means a green CI lint:
# - litellm/ Python staged -> `make lint` (test-linting.yml's lint job)
# - dashboard staged -> prettier + eslint + lint budgets (test-litellm-ui-build.yml's frontend-lint)
# - proxy/types staged -> regenerate dashboard API types and fail on drift (check-ui-api-types.yml)
#
# Each block is skipped when no matching files are staged, so unrelated commits stay
# fast. This is intentionally not auto-installed as a git hook (see scripts/install_git_hooks.sh):
# the dashboard and basedpyright passes can take minutes, so it's run on demand rather
# than firing on every human commit. It is hook-compatible if you want that anyway:
# `ln -s ../../scripts/pre_commit_lint.sh .git/hooks/pre-commit`.
set -eu
repo_root=$(git rev-parse --show-toplevel)
cd "$repo_root"
staged=$(git diff --cached --name-only --diff-filter=ACMR)
staged_match() { printf '%s\n' "$staged" | grep -E "$1" || true; }
# CI's lint job (test-linting.yml) only inspects litellm/, so a tests-only or
# scripts-only commit can't turn it red; scope the trigger there to skip the slow
# make lint when it couldn't catch anything.
litellm_py_files=$(staged_match '^litellm/.*\.py$')
# ruff format (and CI's format step) skip enterprise; the rest of make lint covers it.
fmt_files=$(printf '%s\n' "$litellm_py_files" | grep -v '^litellm/enterprise/' || true)
# check-ui-api-types.yml triggers on any file under litellm/proxy or litellm/types
# (Prisma schema and configs included, not just Python) plus the generator and its
# lockfiles, so match that whole trigger set rather than a Python subset.
spec_files=$(staged_match '^(litellm/(proxy|types)/.*|ui/litellm-dashboard/(scripts/gen-api-types\.mjs|package\.json|package-lock\.json|src/lib/http/schema\.d\.ts))$')
# CI's frontend-lint runs prettier over a wider extension set than eslint; keep that
# split so this flags exactly what the job would.
ui_prettier_files=$(staged_match '^ui/litellm-dashboard/.*\.(js|jsx|ts|tsx|mjs|cjs|json|css|scss|md|mdx|yml|yaml|html)$')
ui_eslint_files=$(staged_match '^ui/litellm-dashboard/.*\.(js|jsx|ts|tsx|mjs|cjs)$')
lint_dashboard() {
(
rc=0
prettier_rel=()
eslint_rel=()
while IFS= read -r f; do
[ -n "$f" ] && prettier_rel+=("${f#ui/litellm-dashboard/}")
done <<EOF
$ui_prettier_files
EOF
while IFS= read -r f; do
[ -n "$f" ] && eslint_rel+=("${f#ui/litellm-dashboard/}")
done <<EOF
$ui_eslint_files
EOF
cd ui/litellm-dashboard
if [ ${#prettier_rel[@]} -gt 0 ]; then
npx prettier --check "${prettier_rel[@]}" || rc=1
fi
if [ ${#eslint_rel[@]} -gt 0 ]; then
npx eslint --no-warn-ignored --pass-on-unpruned-suppressions "${eslint_rel[@]}" || rc=1
fi
# Whole-folder lint budgets, exactly as the frontend-lint job runs them: the
# counts and the committed metrics file are not diff-scoped, so a local pass
# here means the budget step will pass in CI too.
report=$(mktemp)
npx eslint . -f json -o "$report" || true
node scripts/check-lint-budgets.mjs "$report" eslint-budgets.json --check eslint-metrics.json || rc=1
rm -f "$report"
exit $rc
)
}
status=0
if [ -n "$litellm_py_files" ]; then
echo "pre-commit: linting Python (make lint)"
make lint || { echo "✗ Python lint failed. Fix the reds above, then re-run make pre-commit." >&2; status=1; }
# `make lint` format-checks files in origin/base...HEAD, which at pre-commit time
# predates the staged change, so format-check the staged litellm files directly to
# cover a brand-new commit before it lands.
if [ -n "$fmt_files" ]; then
echo "pre-commit: ruff format --check (staged litellm files)"
printf '%s\n' "$fmt_files" | xargs uv run --no-sync ruff format --check --exclude '/enterprise/' \
|| { echo "✗ Unformatted staged files. Fix with: make format, then re-stage." >&2; status=1; }
fi
fi
if [ -n "$ui_prettier_files" ] || [ -n "$ui_eslint_files" ]; then
echo "pre-commit: linting dashboard (prettier + eslint + lint budgets)"
lint_dashboard || { echo "✗ Dashboard lint failed. See above; format with: (cd ui/litellm-dashboard && npm run format)." >&2; status=1; }
fi
if [ -n "$spec_files" ]; then
echo "pre-commit: checking dashboard API types are in sync (npm run gen:api)"
# gen-api-types.mjs imports litellm.proxy.proxy_server, which needs the proxy deps
# 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
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
if ! git diff --quiet -- ui/litellm-dashboard/src/lib/http/schema.d.ts; then
echo "✗ Dashboard API types are stale; regenerated src/lib/http/schema.d.ts. Stage it and re-run make pre-commit." >&2
status=1
fi
else
echo "✗ Could not regenerate API types (npm run gen:api failed)." >&2
status=1
fi
fi
exit $status

View file

@ -0,0 +1,76 @@
"""
Performance benchmarks for the A2A (agent-to-agent) message-translation hot path.
Both directions are covered: the client direction (litellm.completion talking to
an upstream A2A agent) converts OpenAI messages into a prompt and extracts text
from the A2A response, and the proxy server-ingress direction converts an inbound
A2A message into OpenAI messages before bridging to a completion. All are pure-CPU
per-request transforms.
"""
import pytest
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
A2ACompletionBridgeTransformation,
)
from litellm.llms.a2a.common_utils import (
convert_messages_to_prompt,
extract_text_from_a2a_response,
)
MESSAGES = [
{"role": "system", "content": "You are a helpful research assistant."},
{"role": "user", "content": "What is the capital of France?"},
{"role": "assistant", "content": "The capital of France is Paris."},
{"role": "user", "content": "And what is its population?"},
]
MESSAGE_RESPONSE = {
"result": {
"kind": "message",
"parts": [
{"kind": "text", "text": "The population of Paris is about 2.1 million."},
{"kind": "text", "text": "The metro area has over 12 million people."},
],
}
}
TASK_RESPONSE = {
"result": {
"kind": "task",
"artifacts": [{"parts": [{"kind": "text", "text": "Paris has a population of about 2.1 million."}]}],
}
}
A2A_INBOUND_MESSAGE = {
"role": "user",
"parts": [
{"kind": "text", "text": "Summarize the latest quarterly report."},
{"kind": "text", "text": "Focus on revenue and margins."},
],
"messageId": "msg-1",
}
@pytest.mark.benchmark
def test_convert_messages_to_a2a_prompt():
"""Benchmark converting OpenAI messages into an A2A prompt string."""
convert_messages_to_prompt(messages=MESSAGES)
@pytest.mark.benchmark
def test_extract_text_from_a2a_message_response():
"""Benchmark extracting text from a direct-message A2A response."""
extract_text_from_a2a_response(response_dict=MESSAGE_RESPONSE)
@pytest.mark.benchmark
def test_extract_text_from_a2a_task_response():
"""Benchmark extracting text from a task-with-artifacts A2A response."""
extract_text_from_a2a_response(response_dict=TASK_RESPONSE)
@pytest.mark.benchmark
def test_a2a_inbound_message_to_openai_messages():
"""Benchmark the proxy converting an inbound A2A message into OpenAI messages."""
A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(A2A_INBOUND_MESSAGE)

View file

@ -0,0 +1,113 @@
"""
Performance benchmarks for the LLM inference (chat completion) hot path.
The end-to-end cases use ``mock_response`` so the full SDK overhead is exercised
-- provider resolution, request/response transformation, ``ModelResponse``
construction, token counting and cost calculation -- without any network I/O. The
``convert_to_model_response_object`` case isolates the provider-response to
``ModelResponse`` translation, the single deterministic core every non-streaming
completion runs.
"""
import pytest
import litellm
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
convert_to_model_response_object,
)
from litellm.types.utils import ModelResponse
SIMPLE_MESSAGES = [{"role": "user", "content": "Hello, how are you?"}]
MULTI_TURN_MESSAGES = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "What is the capital of France?"},
{
"role": "assistant",
"content": "The capital of France is Paris. It is known as the City of Light.",
},
{"role": "user", "content": "Tell me more about Paris."},
]
TOOL_DEFINITIONS = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather in a given location",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "The city and state, e.g. San Francisco, CA",
},
"unit": {"type": "string", "enum": ["celsius", "fahrenheit"]},
},
"required": ["location"],
},
},
}
]
MOCK_RESPONSE = "The capital of France is Paris, the country's largest city and cultural centre."
PROVIDER_RESPONSE = {
"id": "chatcmpl-abc123",
"object": "chat.completion",
"created": 1700000000,
"model": "gpt-4o",
"choices": [
{
"index": 0,
"finish_reason": "stop",
"message": {"role": "assistant", "content": MOCK_RESPONSE},
}
],
"usage": {"prompt_tokens": 12, "completion_tokens": 16, "total_tokens": 28},
}
@pytest.mark.benchmark
def test_completion_simple_message():
"""Benchmark a single-message completion through the full SDK path."""
litellm.completion(model="gpt-4o", messages=SIMPLE_MESSAGES, mock_response=MOCK_RESPONSE)
@pytest.mark.benchmark
def test_completion_multi_turn():
"""Benchmark a multi-turn completion through the full SDK path."""
litellm.completion(model="gpt-4o", messages=MULTI_TURN_MESSAGES, mock_response=MOCK_RESPONSE)
@pytest.mark.benchmark
def test_completion_with_tools():
"""Benchmark a completion that has to process tool schemas."""
litellm.completion(
model="gpt-4o",
messages=SIMPLE_MESSAGES,
tools=TOOL_DEFINITIONS,
mock_response=MOCK_RESPONSE,
)
@pytest.mark.benchmark
def test_completion_streaming():
"""Benchmark consuming a full streamed completion (CustomStreamWrapper)."""
stream = litellm.completion(
model="gpt-4o",
messages=SIMPLE_MESSAGES,
mock_response=MOCK_RESPONSE,
stream=True,
)
for _ in stream:
pass
@pytest.mark.benchmark
def test_response_to_model_response_object():
"""Benchmark the provider-response to ModelResponse translation core."""
convert_to_model_response_object(
response_object=PROVIDER_RESPONSE,
model_response_object=ModelResponse(),
)

View file

@ -0,0 +1,84 @@
"""
Performance benchmarks for the MCP tool hot path.
Two layers are covered: the client-side translation between MCP and OpenAI
function-calling formats, and the server-side tool-name prefixing that the proxy
runs on every list-tools (prefix each tool) and call-tool (strip prefix to route)
request. Both are pure-CPU and deterministic.
"""
import pytest
from mcp.types import Tool as MCPTool
from litellm.experimental_mcp_client.tools import (
transform_mcp_tool_to_openai_tool,
transform_openai_tool_call_request_to_mcp_tool_call_request,
)
from litellm.proxy._experimental.mcp_server.utils import (
add_server_prefix_to_name,
split_server_prefix_from_name,
)
def _make_tool(index: int) -> MCPTool:
return MCPTool(
name=f"tool_{index}",
description=f"Test tool number {index} that performs an operation",
inputSchema={
"type": "object",
"properties": {
"query": {"type": "string", "description": "The search query"},
"limit": {"type": "integer", "description": "Max results"},
},
"required": ["query"],
},
)
SINGLE_TOOL = _make_tool(0)
TOOL_LIST = tuple(_make_tool(i) for i in range(20))
TOOL_NAMES = tuple(t.name for t in TOOL_LIST)
SERVER_NAME = "github_mcp"
PREFIXED_TOOL_NAME = add_server_prefix_to_name("tool_0", SERVER_NAME)
OPENAI_TOOL_CALL = {
"id": "call_abc123",
"type": "function",
"function": {
"name": "tool_0",
"arguments": '{"query": "weather in San Francisco", "limit": 5}',
},
}
@pytest.mark.benchmark
def test_transform_single_mcp_tool_to_openai():
"""Benchmark translating one MCP tool into OpenAI tool format."""
transform_mcp_tool_to_openai_tool(mcp_tool=SINGLE_TOOL)
@pytest.mark.benchmark
def test_transform_mcp_tool_list_to_openai():
"""Benchmark translating a full list-tools response into OpenAI format."""
for tool in TOOL_LIST:
transform_mcp_tool_to_openai_tool(mcp_tool=tool)
@pytest.mark.benchmark
def test_transform_openai_tool_call_to_mcp():
"""Benchmark translating an OpenAI tool call into an MCP call request."""
transform_openai_tool_call_request_to_mcp_tool_call_request(openai_tool=OPENAI_TOOL_CALL)
@pytest.mark.benchmark
def test_mcp_server_prefix_tool_list():
"""Benchmark the proxy prefixing every tool name on a list-tools response."""
for name in TOOL_NAMES:
add_server_prefix_to_name(name, SERVER_NAME)
@pytest.mark.benchmark
def test_mcp_server_strip_prefix_on_call():
"""Benchmark the proxy stripping the server prefix to route a tool call."""
split_server_prefix_from_name(PREFIXED_TOOL_NAME)

View file

@ -181,6 +181,13 @@ ANTHROPIC_DIRECT_MODELS: Tuple[ModelEntry, ...] = (
required_env=_ANTHROPIC_REQ,
caps=_CAPS_XHIGH_MAX,
),
ModelEntry(
alias="claude-sonnet-5",
model="anthropic/claude-sonnet-5",
mode="adaptive",
required_env=_ANTHROPIC_REQ,
caps=_CAPS_XHIGH_MAX,
),
ModelEntry(
alias="claude-sonnet-4-6",
model="anthropic/claude-sonnet-4-6",

View file

@ -201,8 +201,8 @@ async def test_reasoning_effort_grid(
def test_grid_cell_count() -> None:
assert len(_PARAMS) == 29 * 11, (
f"expected 319 cells (29 provider x model combos x 11 efforts), "
assert len(_PARAMS) == 30 * 11, (
f"expected 330 cells (30 provider x model combos x 11 efforts), "
f"got {len(_PARAMS)}"
)

View file

@ -20,6 +20,8 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
convert_to_anthropic_tool_invoke,
convert_url_to_base64,
create_anthropic_image_param,
get_tool_calls_from_response,
has_tool_with_name,
llama_2_chat_pt,
prompt_factory,
)
@ -2385,3 +2387,100 @@ def test_anthropic_messages_pt_list_content_with_thinking_preserves_order():
# Verify signatures preserved in correct positions
assert content[0]["signature"] == "sig_1"
assert content[3]["signature"] == "sig_2"
def test_get_tool_calls_from_response_chat_completions():
response = MagicMock()
response.output = None
response.content = None
tool_call = MagicMock()
tool_call.id = "call_abc"
tool_call.function.name = "my_tool"
tool_call.function.arguments = '{"x": 1}'
response.choices = [MagicMock(message=MagicMock(tool_calls=[tool_call]))]
result = get_tool_calls_from_response(response)
assert result == [{"id": "call_abc", "name": "my_tool", "arguments": {"x": 1}}]
def test_get_tool_calls_from_response_responses_api():
response = MagicMock()
response.choices = None
response.content = None
response.output = [
{
"type": "function_call",
"id": "fc_1",
"call_id": "call_1",
"name": "my_tool",
"arguments": '{"x": 2}',
}
]
result = get_tool_calls_from_response(response)
assert result == [{"id": "call_1", "name": "my_tool", "arguments": {"x": 2}}]
def test_get_tool_calls_from_response_anthropic_messages():
response = MagicMock()
response.choices = None
response.output = None
response.content = [
{"type": "tool_use", "id": "toolu_1", "name": "my_tool", "input": {"x": 3}},
]
result = get_tool_calls_from_response(response)
assert result == [{"id": "toolu_1", "name": "my_tool", "arguments": {"x": 3}}]
def test_get_tool_calls_from_response_anthropic_messages_plain_dict():
# AnthropicMessagesResponse is a TypedDict -- real responses are plain
# dicts at runtime, not objects with attribute access. A MagicMock-only
# test would pass even if the extractor used bare getattr() and silently
# returned nothing for a real response.
response = {
"content": [
{"type": "tool_use", "id": "toolu_1", "name": "my_tool", "input": {"x": 3}},
]
}
result = get_tool_calls_from_response(response)
assert result == [{"id": "toolu_1", "name": "my_tool", "arguments": {"x": 3}}]
def test_get_tool_calls_from_response_no_tool_calls():
response = MagicMock()
response.choices = None
response.output = None
response.content = None
assert get_tool_calls_from_response(response) == []
def test_has_tool_with_name_openai_function_shape():
tools = [{"type": "function", "function": {"name": "my_tool"}}]
assert has_tool_with_name(tools, "my_tool")
assert not has_tool_with_name(tools, "other_tool")
def test_has_tool_with_name_anthropic_custom_shape():
tools = [{"type": "custom", "name": "my_tool", "input_schema": {}}]
assert has_tool_with_name(tools, "my_tool")
assert not has_tool_with_name(tools, "other_tool")
def test_has_tool_with_name_anthropic_shape_without_type_field():
# Anthropic's documented client tool format is just name + input_schema;
# "type" isn't required at all (type: "custom" is only one possible value).
tools = [{"name": "my_tool", "input_schema": {}}]
assert has_tool_with_name(tools, "my_tool")
assert not has_tool_with_name(tools, "other_tool")
def test_has_tool_with_name_not_a_list():
assert not has_tool_with_name(None, "my_tool")
assert not has_tool_with_name("not a list", "my_tool")

View file

@ -11,6 +11,7 @@ import json
import os
import pytest
import asyncio
import requests
# Path to your service account JSON file
SERVICE_ACCOUNT_FILE = "path/to/your/service-account.json"
@ -57,98 +58,114 @@ def load_vertex_ai_credentials():
os.environ["GOOGLE_APPLICATION_CREDENTIALS"] = os.path.abspath(temp_file.name)
async def call_spend_logs_endpoint():
"""
Call this
curl -X GET "http://0.0.0.0:4000/spend/logs" -H "Authorization: Bearer sk-1234"
"""
import datetime
import requests
todays_date = datetime.datetime.now().strftime("%Y-%m-%d")
url = f"http://0.0.0.0:4000/global/spend/logs?api_key=best-api-key-ever"
headers = {"Authorization": f"Bearer sk-1234"}
response = requests.get(url, headers=headers)
print("response from call_spend_logs_endpoint", response)
if response.status_code != 200:
print(f"spend logs endpoint returned {response.status_code}: {response.text}")
return None
json_response = response.json()
# get spend for today
"""
json response looks like this
[{'date': '2024-08-30', 'spend': 0.00016600000000000002, 'api_key': 'best-api-key-ever'}]
"""
print("json_response", json_response)
todays_date = datetime.datetime.now().strftime("%Y-%m-%d")
for spend_log in json_response:
if spend_log["date"] == todays_date:
return spend_log["spend"]
LITE_LLM_ENDPOINT = "http://localhost:4000"
SPEND_LOG_API_KEY = "best-api-key-ever"
def _is_vertex_quota_error(exc: Exception) -> bool:
message = str(exc)
return (
"429" in message
or "Too Many Requests" in message
or "RESOURCE_EXHAUSTED" in message
def get_tracked_spend() -> float:
"""
Total spend recorded under the pass-through key in the global spend view.
Sums every day the endpoint returns instead of matching the runner's local
"today" so a UTC date rollover mid-test can't hide a freshly billed call, and
treats an unreachable endpoint as "nothing recorded yet" (0.0).
"""
url = f"{LITE_LLM_ENDPOINT}/global/spend/logs?api_key={SPEND_LOG_API_KEY}"
response = requests.get(url, headers={"Authorization": "Bearer sk-1234"})
if response.status_code != 200:
print(f"global spend logs endpoint returned {response.status_code}: {response.text}")
return 0.0
rows = response.json()
print("global spend logs rows", rows)
return sum(float(row.get("spend") or 0.0) for row in rows)
VERTEX_PROJECT = "litellm-ci-cd"
VERTEX_MODEL = "gemini-3.1-flash-lite"
VERTEX_GENERATE_CONTENT_URL = (
f"{LITE_LLM_ENDPOINT}/vertex_ai/v1/projects/{VERTEX_PROJECT}"
f"/locations/global/publishers/google/models/{VERTEX_MODEL}:generateContent"
)
def _vertex_access_token() -> str:
import google.auth
import google.auth.transport.requests
credentials, _ = google.auth.default(
scopes=["https://www.googleapis.com/auth/cloud-platform"]
)
credentials.refresh(google.auth.transport.requests.Request())
return credentials.token
def _spend_log_for_request(call_id: str) -> dict | None:
response = requests.get(
f"{LITE_LLM_ENDPOINT}/spend/logs?request_id={call_id}",
headers={"Authorization": "Bearer sk-1234"},
timeout=30,
)
if response.status_code != 200:
return None
rows = response.json()
return rows[0] if rows else None
def _is_vertex_quota_error(response: requests.Response) -> bool:
return response.status_code == 429 or "RESOURCE_EXHAUSTED" in response.text
@pytest.mark.asyncio()
async def test_basic_vertex_ai_pass_through_with_spendlog():
spend_before = await call_spend_logs_endpoint() or 0.0
load_vertex_ai_credentials()
access_token = _vertex_access_token()
vertexai.init(
project="litellm-ci-cd",
location="global",
api_endpoint=f"{LITE_LLM_ENDPOINT}/vertex_ai",
api_transport="rest",
)
# Drive the pass-through over HTTP instead of the vertexai SDK: the SDK intermittently
# routes generateContent to the public Vertex endpoint rather than the proxy override,
# so the call never reaches LiteLLM and no spend is logged. A direct request always
# hits the proxy. Spend logging then runs on a best-effort background worker that can
# drop a single event, so retry a few billed calls and assert that one specific call's
# spend log lands. Failing every attempt still fails hard, which is the signal we want
# if cost tracking is broken.
max_attempts = 3
poll_seconds = 60
poll_interval = 5
model = GenerativeModel(model_name="gemini-3.1-flash-lite")
try:
response = model.generate_content("hi")
except Exception as exc:
if _is_vertex_quota_error(exc):
for attempt in range(1, max_attempts + 1):
response = requests.post(
VERTEX_GENERATE_CONTENT_URL,
headers={
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
},
json={"contents": [{"role": "user", "parts": [{"text": "hi"}]}]},
timeout=60,
)
if _is_vertex_quota_error(response):
pytest.skip("Vertex AI quota exhausted")
raise
assert (
response.status_code == 200
), f"vertex pass-through call failed: {response.status_code} {response.text}"
print("response", response)
call_id = response.headers.get("x-litellm-call-id")
assert call_id, "proxy response missing x-litellm-call-id header"
# Spend logging is async/batched and can lag under CI load, so poll instead of
# sleeping a fixed amount. A transient empty read is skipped, not counted as 0.0
# spend, which would spuriously fail the assertion on an otherwise-billed call.
max_wait = 240 # total seconds to wait
poll_interval = 10 # seconds between checks
elapsed = 0
spend_after = spend_before
while elapsed < max_wait:
await asyncio.sleep(poll_interval)
elapsed += poll_interval
latest_spend = await call_spend_logs_endpoint()
if latest_spend is None:
print(f"spend logs unavailable (elapsed={elapsed}s), retrying")
continue
spend_after = latest_spend
print(f"spend_after (elapsed={elapsed}s)", spend_after)
if spend_after > spend_before:
break
for _ in range(poll_seconds // poll_interval):
await asyncio.sleep(poll_interval)
row = _spend_log_for_request(call_id)
if row is not None and float(row.get("spend") or 0) > 0:
assert "gemini" in row["model"], f"unexpected model in spend log: {row}"
assert (
row["custom_llm_provider"] == "vertex_ai"
), f"unexpected provider in spend log: {row}"
return
assert (
spend_after > spend_before
), "Spend should be greater than before after {}s. spend_before: {}, spend_after: {}".format(
elapsed, spend_before, spend_after
print(f"attempt {attempt}: spend log for call {call_id} not found yet, re-billing")
pytest.fail(
f"Vertex pass-through spend never recorded after {max_attempts} billed calls"
)
@ -156,7 +173,7 @@ async def test_basic_vertex_ai_pass_through_with_spendlog():
@pytest.mark.skip(reason="skip flaky test - vertex pass through streaming is flaky")
async def test_basic_vertex_ai_pass_through_streaming_with_spendlog():
spend_before = await call_spend_logs_endpoint() or 0.0
spend_before = get_tracked_spend()
print("spend_before", spend_before)
load_vertex_ai_credentials()
@ -176,7 +193,7 @@ async def test_basic_vertex_ai_pass_through_streaming_with_spendlog():
print("response", response)
await asyncio.sleep(20)
spend_after = await call_spend_logs_endpoint()
spend_after = get_tracked_spend()
print("spend_after", spend_after)
assert (
spend_after > spend_before

View file

@ -1112,3 +1112,133 @@ async def test_no_map_preserves_old_single_threshold(
# Old path cache key has no threshold percentage
cache_key = mock_cache.async_set_cache.call_args[1]["key"]
assert cache_key == "email_budget_alerts:max_budget_alert:test_user"
CUSTOM_SIGNATURE = "<div>Best,<br/>The Acme Platform Team</div>"
@pytest.mark.asyncio
async def test_send_soft_budget_alert_email_uses_custom_signature(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Soft budget alert honors EMAIL_SIGNATURE for premium users."""
event = WebhookEvent(
user_id="test_user",
user_email="test@example.com",
event_group=Litellm_EntityType.USER,
event="soft_budget_crossed",
event_message="Soft Budget Crossed",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
)
with mock.patch.dict(
os.environ,
{"PROXY_BASE_URL": "http://test.com", "EMAIL_SIGNATURE": CUSTOM_SIGNATURE},
), patch("litellm.proxy.proxy_server.premium_user", True):
await base_email_logger.send_soft_budget_alert_email(event)
html_body = mock_send_email.call_args[1]["html_body"]
assert CUSTOM_SIGNATURE in html_body
assert "The LiteLLM team" not in html_body
@pytest.mark.asyncio
async def test_send_team_soft_budget_alert_email_uses_custom_signature(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Team soft budget alert honors EMAIL_SIGNATURE for premium users."""
event = WebhookEvent(
user_id="test_user",
event_group=Litellm_EntityType.TEAM,
event="soft_budget_crossed",
event_message="Team Soft Budget Crossed",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
team_alias="Acme",
alert_emails=["teamlead@example.com"],
)
with mock.patch.dict(
os.environ,
{"PROXY_BASE_URL": "http://test.com", "EMAIL_SIGNATURE": CUSTOM_SIGNATURE},
), patch("litellm.proxy.proxy_server.premium_user", True):
await base_email_logger.send_team_soft_budget_alert_email(event)
html_body = mock_send_email.call_args[1]["html_body"]
assert CUSTOM_SIGNATURE in html_body
assert "The LiteLLM team" not in html_body
@pytest.mark.asyncio
async def test_send_max_budget_alert_email_single_recipient_uses_custom_signature(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Max budget alert (single-recipient path) honors EMAIL_SIGNATURE."""
event = WebhookEvent(
user_id="test_user",
user_email="test@example.com",
event_group=Litellm_EntityType.USER,
event="max_budget_alert",
event_message="Max Budget Alert",
spend=165.0,
max_budget=200.0,
)
with mock.patch.dict(
os.environ,
{"PROXY_BASE_URL": "http://test.com", "EMAIL_SIGNATURE": CUSTOM_SIGNATURE},
), patch("litellm.proxy.proxy_server.premium_user", True):
await base_email_logger.send_max_budget_alert_email(event)
html_body = mock_send_email.call_args[1]["html_body"]
assert CUSTOM_SIGNATURE in html_body
assert "The LiteLLM team" not in html_body
@pytest.mark.asyncio
async def test_send_max_budget_alert_email_multi_recipient_uses_custom_signature(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Max budget alert (multi-threshold/recipient path) honors EMAIL_SIGNATURE."""
event = WebhookEvent(
user_id="test_user",
user_email="owner@example.com",
event_group=Litellm_EntityType.USER,
event="max_budget_alert",
event_message="Max Budget Alert",
spend=165.0,
max_budget=200.0,
)
with mock.patch.dict(
os.environ,
{"PROXY_BASE_URL": "http://test.com", "EMAIL_SIGNATURE": CUSTOM_SIGNATURE},
), patch("litellm.proxy.proxy_server.premium_user", True):
await base_email_logger.send_max_budget_alert_email(
event, threshold_pct=75, recipient_emails=["a@example.com", "b@example.com"]
)
html_body = mock_send_email.call_args[1]["html_body"]
assert CUSTOM_SIGNATURE in html_body
assert "The LiteLLM team" not in html_body
@pytest.mark.asyncio
async def test_send_soft_budget_alert_email_default_footer_when_no_signature(
base_email_logger, mock_send_email, mock_lookup_user_email
):
"""Without EMAIL_SIGNATURE, budget alert falls back to the default EMAIL_FOOTER."""
event = WebhookEvent(
user_id="test_user",
user_email="test@example.com",
event_group=Litellm_EntityType.USER,
event="soft_budget_crossed",
event_message="Soft Budget Crossed",
spend=105.0,
max_budget=200.0,
soft_budget=100.0,
)
with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}):
await base_email_logger.send_soft_budget_alert_email(event)
html_body = mock_send_email.call_args[1]["html_body"]
assert EMAIL_FOOTER in html_body

View file

@ -14,6 +14,7 @@ from litellm.integrations.code_interpreter_interception.handler import (
LITELLM_CODE_EXECUTION_TOOL_NAME,
_INTERCEPTION_ACTIVE_KEY as _ACTIVE_KEY,
_SANDBOX_KEY,
_SESSION_SCOPED_KEY,
)
from litellm.types.integrations.custom_logger import (
CHAT_COMPLETION_AGENTIC_SURFACE,
@ -138,11 +139,7 @@ async def test_build_plan_runs_code_and_feeds_output_back():
assert sandbox.run_calls[0]["code"] == "print(40 + 2)"
messages = _iter_messages(plan)
outputs = [
m
for m in messages
if isinstance(m, dict) and m.get("type") == "function_call_output"
]
outputs = [m for m in messages if isinstance(m, dict) and m.get("type") == "function_call_output"]
assert outputs, "expected a function_call_output item appended"
output_item = next(m for m in outputs if m.get("call_id") == "c1")
assert "42" in str(output_item["output"])
@ -160,9 +157,7 @@ async def test_pre_call_converts_code_interpreter_tool():
assert result is not None
tools = result["tools"]
assert not any(
t.get("type") == "code_interpreter" for t in tools
), "code_interpreter tool must be removed"
assert not any(t.get("type") == "code_interpreter" for t in tools), "code_interpreter tool must be removed"
names = [t.get("name") or (t.get("function") or {}).get("name") for t in tools]
assert LITELLM_CODE_EXECUTION_TOOL_NAME in names
@ -267,9 +262,7 @@ async def test_should_run_detects_only_matching_function_call():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
active_kwargs = {"_code_interpreter_interception_active": True}
match = FakeResponse(
output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)]
)
match = FakeResponse(output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)])
should_run, payload = await logger.async_should_run_agentic_loop(
response=match,
model="gpt-5",
@ -331,9 +324,7 @@ async def test_container_reused_within_request_via_server_sandbox_key():
**common,
)
assert (
len(sandbox.create_calls) == 1
), "the sandbox is reused across loop iterations sharing one server sandbox key"
assert len(sandbox.create_calls) == 1, "the sandbox is reused across loop iterations sharing one server sandbox key"
@pytest.mark.asyncio
@ -372,9 +363,9 @@ async def test_colliding_caller_call_id_does_not_share_sandbox():
**common,
)
assert (
len(sandbox.create_calls) == 2
), "distinct server sandbox keys must isolate sandboxes despite a colliding call id"
assert len(sandbox.create_calls) == 2, (
"distinct server sandbox keys must isolate sandboxes despite a colliding call id"
)
@pytest.mark.asyncio
@ -479,14 +470,11 @@ async def test_post_hook_injects_code_interpreter_call_matching_openai_shape():
)
response = FakeResponse(output=[{"type": "message", "content": []}])
out = await logger.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs={}
)
out = await logger.async_post_agentic_loop_response_hook(response=response, plan=plan, kwargs={})
types = [item.get("type") for item in out.output]
assert types == ["code_interpreter_call", "message"], (
"code_interpreter_call must be re-injected before the message, matching "
"OpenAI's native output ordering"
"code_interpreter_call must be re-injected before the message, matching OpenAI's native output ordering"
)
assert set(out.output[0].keys()) == {
"id",
@ -524,8 +512,7 @@ async def test_pre_call_forces_non_stream_for_loop():
assert out is not None
assert out["stream"] is False, "loop requires a non-streaming upstream call"
assert out["_code_interpreter_interception_converted_stream"] is True, (
"the converted-stream flag must be set so the final response is wrapped "
"back into a stream for the caller"
"the converted-stream flag must be set so the final response is wrapped back into a stream for the caller"
)
@ -556,9 +543,7 @@ async def test_gate_refuses_without_server_active_marker():
"""A forged litellm_code_execution call must not trigger the loop unless the
pre-call hook actually converted a native code_interpreter tool."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
forged = FakeResponse(
output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)]
)
forged = FakeResponse(output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)])
should_run, payload = await logger.async_should_run_agentic_loop(
response=forged,
@ -577,12 +562,8 @@ async def test_gate_refuses_without_server_active_marker():
@pytest.mark.asyncio
async def test_gate_rechecks_provider_scope():
"""enabled_providers must be re-enforced at the gate, not only in pre-call."""
logger = CodeInterpreterInterceptionLogger(
sandbox_config=FakeSandbox(), enabled_providers=["openai"]
)
response = FakeResponse(
output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)]
)
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox(), enabled_providers=["openai"])
response = FakeResponse(output=[_function_call_item(name=LITELLM_CODE_EXECUTION_TOOL_NAME)])
should_run, _ = await logger.async_should_run_agentic_loop(
response=response,
@ -600,11 +581,7 @@ async def test_gate_rechecks_provider_scope():
@pytest.mark.asyncio
async def test_chat_completion_gate_detects_code_execution_tool_call():
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
response = {
"choices": [
{"message": {"tool_calls": [_chat_function_call_item(call_id="call_123")]}}
]
}
response = {"choices": [{"message": {"tool_calls": [_chat_function_call_item(call_id="call_123")]}}]}
should_run, payload = await logger.async_should_run_agentic_loop(
response=response,
@ -661,9 +638,7 @@ async def test_chat_completion_build_plan_runs_code_and_appends_tool_message():
},
model="gpt-5",
messages=[{"role": "user", "content": "x"}],
response={
"choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]
},
response={"choices": [{"message": {"tool_calls": [_chat_function_call_item()]}}]},
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={
"tools": [native_chat_tool],
@ -738,8 +713,7 @@ async def test_pre_call_strips_client_forged_marker_on_initial_request():
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert _ACTIVE_KEY not in kwargs, (
"no native code_interpreter tool was present, so a client-supplied "
"active marker must be cleared"
"no native code_interpreter tool was present, so a client-supplied active marker must be cleared"
)
assert kwargs["litellm_metadata"] == {"safe_user_value": "kept"}
@ -774,8 +748,7 @@ async def test_pre_call_strips_forged_loop_controls_then_mints_own_markers():
assert metadata[_ACTIVE_KEY] is True
assert metadata[_SANDBOX_KEY] == result[_SANDBOX_KEY]
assert metadata[_SANDBOX_KEY] != "client-forged", (
"the surviving sandbox key must be the server-minted one, not the forged "
"value the client supplied"
"the surviving sandbox key must be the server-minted one, not the forged value the client supplied"
)
@ -793,8 +766,7 @@ async def test_pre_call_preserves_marker_on_server_followup():
await logger.async_pre_call_deployment_hook(kwargs, CallTypes.aresponses)
assert kwargs.get(_ACTIVE_KEY) is True, (
"the server-set marker must survive followup requests so multi-round "
"code execution keeps working"
"the server-set marker must survive followup requests so multi-round code execution keeps working"
)
@ -805,9 +777,7 @@ async def test_sandbox_deleted_after_loop_completes():
plan = await _build_plan(logger, sandbox, call_id="k1")
assert sandbox.create_calls, "sandbox must be created during the loop"
assert (
not sandbox.delete_calls
), "sandbox must outlive the loop until the final hook"
assert not sandbox.delete_calls, "sandbox must outlive the loop until the final hook"
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
@ -816,8 +786,7 @@ async def test_sandbox_deleted_after_loop_completes():
)
assert len(sandbox.delete_calls) == 1, (
"the sandbox must be deleted once the final response is assembled, "
"otherwise it keeps running and billing"
"the sandbox must be deleted once the final response is assembled, otherwise it keeps running and billing"
)
assert "sbxkey1" not in logger._container_cache
@ -829,16 +798,11 @@ async def test_post_hook_delete_is_idempotent_across_loop_levels():
plan = await _build_plan(logger, sandbox, call_id="k1")
response = FakeResponse(output=[{"type": "message", "content": []}])
await logger.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs={}
)
await logger.async_post_agentic_loop_response_hook(
response=response, plan=plan, kwargs={}
)
await logger.async_post_agentic_loop_response_hook(response=response, plan=plan, kwargs={})
await logger.async_post_agentic_loop_response_hook(response=response, plan=plan, kwargs={})
assert len(sandbox.delete_calls) == 1, (
"deleting an already-removed container must be a no-op so unwinding "
"loop levels do not double-delete"
"deleting an already-removed container must be a no-op so unwinding loop levels do not double-delete"
)
@ -860,8 +824,7 @@ async def test_build_plan_deletes_sandbox_when_execution_raises():
assert len(sandbox.create_calls) == 1, "the sandbox must have been created"
assert len(sandbox.delete_calls) == 1, (
"a build failure must delete the cached sandbox so it does not keep "
"running and billing"
"a build failure must delete the cached sandbox so it does not keep running and billing"
)
assert "sbxkey1" not in logger._container_cache
@ -875,8 +838,7 @@ async def test_cleanup_hook_deletes_sandbox():
await logger.async_agentic_loop_cleanup_hook(plan=plan, kwargs={})
assert len(sandbox.delete_calls) == 1, (
"the cleanup hook must delete the sandbox so a rerun failure cannot "
"leak a running container"
"the cleanup hook must delete the sandbox so a rerun failure cannot leak a running container"
)
assert "sbxkey1" not in logger._container_cache
@ -895,8 +857,7 @@ async def test_cleanup_hook_is_idempotent_with_post_hook():
await logger.async_agentic_loop_cleanup_hook(plan=plan, kwargs={})
assert len(sandbox.delete_calls) == 1, (
"cleanup running in finally after the success-path post hook already "
"deleted the sandbox must not double-delete"
"cleanup running in finally after the success-path post hook already deleted the sandbox must not double-delete"
)
@ -923,9 +884,7 @@ async def test_responses_plan_cleans_up_sandbox_when_followup_raises():
plan = AgenticLoopPlan(
run_agentic_loop=True,
request_patch=AgenticLoopRequestPatch(
model="gpt-5", messages=[{"role": "user", "content": "x"}]
),
request_patch=AgenticLoopRequestPatch(model="gpt-5", messages=[{"role": "user", "content": "x"}]),
metadata={"sandbox_key": "sbxkey1"},
)
@ -995,9 +954,7 @@ async def test_run_code_does_not_re_resolve_registry(monkeypatch):
sandbox_tools.clear_sandbox_tools()
stdout = await logger._run_tool_call(
container=container, params=params, arguments='{"code":"print(1)"}'
)
stdout = await logger._run_tool_call(container=container, params=params, arguments='{"code":"print(1)"}')
finally:
sandbox_tools.clear_sandbox_tools()
@ -1013,9 +970,7 @@ async def test_run_tool_call_surfaces_execution_error():
class ErroringSandbox(FakeSandbox):
async def arun_code(self, *, container, code, **kwargs):
self.run_calls.append({"container": container, "code": code})
return CodeExecutionResult(
stdout="", error={"name": "ValueError", "value": "boom"}
)
return CodeExecutionResult(stdout="", error={"name": "ValueError", "value": "boom"})
sandbox = ErroringSandbox()
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
@ -1036,9 +991,7 @@ async def test_run_tool_call_reports_unparseable_arguments():
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
container = await logger._create_container()
stdout = await logger._run_tool_call(
container=container[0], params=None, arguments="not-json"
)
stdout = await logger._run_tool_call(container=container[0], params=None, arguments="not-json")
assert stdout == "[invalid tool arguments: could not parse code]"
assert not sandbox.run_calls, "code must not run when arguments cannot be parsed"
@ -1048,9 +1001,7 @@ async def test_run_tool_call_reports_unparseable_arguments():
async def test_pre_call_skips_provider_outside_scope():
"""enabled_providers must filter the pre-call conversion so a request to an
out-of-scope provider is left untouched."""
logger = CodeInterpreterInterceptionLogger(
sandbox_config=FakeSandbox(), enabled_providers=["openai"]
)
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox(), enabled_providers=["openai"])
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "anthropic",
@ -1119,6 +1070,7 @@ async def test_prune_expired_cache_deletes_underlying_container():
container,
params,
time.time() - handler_mod._CACHE_TTL_SECONDS - 1,
None,
)
await logger._prune_expired_cache()
@ -1217,3 +1169,258 @@ async def test_extract_tool_calls_reads_object_attributes():
assert len(calls) == 1
assert calls[0]["call_id"] == "c9"
assert calls[0]["arguments"] == '{"code":"print(1)"}'
# ---------------------------------------------------------------------------
# Sticky session tests
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_pre_call_uses_session_id_from_metadata_as_sandbox_key():
"""When session_id is in request metadata, it becomes the sandbox key so the
container is shared across requests in the same session."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
session_id = "conv-abc-123"
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
"metadata": {"session_id": session_id},
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is not None
assert result[_SANDBOX_KEY] == session_id
assert result[_SESSION_SCOPED_KEY] is True
assert result["litellm_metadata"][_SANDBOX_KEY] == session_id
assert result["litellm_metadata"][_SESSION_SCOPED_KEY] is True
@pytest.mark.asyncio
async def test_pre_call_uses_session_id_from_litellm_metadata():
"""session_id in litellm_metadata also works as the sticky key."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
session_id = "sess-xyz-789"
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
"litellm_metadata": {"session_id": session_id},
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is not None
assert result[_SANDBOX_KEY] == session_id
assert result[_SESSION_SCOPED_KEY] is True
@pytest.mark.asyncio
async def test_pre_call_without_session_id_still_mints_random_key():
"""Requests without a session_id still get a server-minted random sandbox key."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
kwargs = {
"tools": [{"type": "code_interpreter", "container": {"type": "auto"}}],
"custom_llm_provider": "openai",
}
result = await logger.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion)
assert result is not None
assert _SESSION_SCOPED_KEY not in result or result[_SESSION_SCOPED_KEY] is False
assert len(result[_SANDBOX_KEY]) >= 16
@pytest.mark.asyncio
async def test_session_scoped_sandbox_survives_agentic_loop_cleanup():
"""A session-scoped sandbox must NOT be deleted by the cleanup or post hooks;
it needs to persist across requests within the same session."""
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
session_id = "conv-persist-me"
plan = await logger.async_build_agentic_loop_plan(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"x = 10"}',
}
]
},
model="gpt-4o-mini",
messages=[{"role": "user", "content": "set x"}],
response=FakeResponse(output=[_function_call_item()]),
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
logging_obj=FakeLogging(litellm_call_id="k1"),
stream=False,
kwargs={
"litellm_call_id": "k1",
_SANDBOX_KEY: session_id,
_SESSION_SCOPED_KEY: True,
},
)
assert plan.metadata["is_session_scoped"] is True
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
plan=plan,
kwargs={},
)
await logger.async_agentic_loop_cleanup_hook(plan=plan, kwargs={})
assert not sandbox.delete_calls, (
"session-scoped sandbox must not be deleted after a single agentic loop; "
"it must persist for the next request in the session"
)
assert session_id in logger._container_cache, "session-scoped container must remain in cache after loop ends"
@pytest.mark.asyncio
async def test_session_scoped_sandbox_reused_across_sequential_requests():
"""Two sequential requests with the same session_id must share one container,
confirming state (e.g. assigned variables) can persist across HTTP requests."""
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
session_id = "conv-reuse-me"
common_plan_args = dict(
tools={
"tool_calls": [
{
"call_id": "c1",
"name": LITELLM_CODE_EXECUTION_TOOL_NAME,
"arguments": '{"code":"print(1)"}',
}
]
},
model="gpt-4o-mini",
messages=[{"role": "user", "content": "x"}],
response=FakeResponse(output=[_function_call_item()]),
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={"tools": []},
stream=False,
)
session_kwargs = {_SANDBOX_KEY: session_id, _SESSION_SCOPED_KEY: True}
plan1 = await logger.async_build_agentic_loop_plan(
logging_obj=FakeLogging(litellm_call_id="req1"),
kwargs={"litellm_call_id": "req1", **session_kwargs},
**common_plan_args,
)
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
plan=plan1,
kwargs={},
)
plan2 = await logger.async_build_agentic_loop_plan(
logging_obj=FakeLogging(litellm_call_id="req2"),
kwargs={"litellm_call_id": "req2", **session_kwargs},
**common_plan_args,
)
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
plan=plan2,
kwargs={},
)
assert len(sandbox.create_calls) == 1, (
"a single container must serve both requests in the same session; "
"two creates means state cannot persist between requests"
)
assert len(sandbox.delete_calls) == 0, "the session container must still be alive after both requests complete"
@pytest.mark.asyncio
async def test_non_session_sandbox_still_deleted_after_loop():
"""Without a session_id, the existing per-request ephemeral behavior is unchanged."""
sandbox = FakeSandbox(stdout="42")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
plan = await _build_plan(logger, sandbox, call_id="k1")
await logger.async_post_agentic_loop_response_hook(
response=FakeResponse(output=[{"type": "message", "content": []}]),
plan=plan,
kwargs={},
)
assert len(sandbox.delete_calls) == 1, "non-session sandbox must still be cleaned up after each request"
@pytest.mark.asyncio
async def test_sandbox_key_scoped_to_api_key_hash_isolates_users():
"""Two callers supplying the same session_id but different API key hashes must
each get their own sandbox; sharing across tenants would let one read or mutate
the other's interpreter state."""
logger = CodeInterpreterInterceptionLogger(sandbox_config=FakeSandbox())
session_id = "same-session-id"
result_a = await logger.async_pre_call_deployment_hook(
{
"tools": [{"type": "code_interpreter"}],
"custom_llm_provider": "openai",
"metadata": {"session_id": session_id},
"user_api_key_hash": "hash-for-tenant-a",
},
CallTypes.acompletion,
)
result_b = await logger.async_pre_call_deployment_hook(
{
"tools": [{"type": "code_interpreter"}],
"custom_llm_provider": "openai",
"metadata": {"session_id": session_id},
"user_api_key_hash": "hash-for-tenant-b",
},
CallTypes.acompletion,
)
assert result_a is not None and result_b is not None
assert result_a[_SANDBOX_KEY] != result_b[_SANDBOX_KEY], (
"same session_id from different API keys must yield different sandbox keys; "
"otherwise tenant A can read tenant B's sandbox state"
)
assert "hash-for-tenant-a" in result_a[_SANDBOX_KEY]
assert "hash-for-tenant-b" in result_b[_SANDBOX_KEY]
@pytest.mark.asyncio
async def test_per_identity_cap_evicts_lru_session():
"""When a single identity holds the cap limit of session sandboxes and opens a
new one, the least-recently-used session is evicted so the allocation stays
bounded. Without this, rotating session IDs is an unbounded sandbox leak."""
from litellm.integrations.code_interpreter_interception.handler import _SESSION_SCOPED_PER_IDENTITY_CAP
sandbox = FakeSandbox(stdout="ok")
logger = CodeInterpreterInterceptionLogger(sandbox_config=sandbox)
identity = "hash-for-identity-x"
for i in range(_SESSION_SCOPED_PER_IDENTITY_CAP):
await logger._get_or_create_container(
cache_key=f"{identity}:session-{i}",
identity=identity,
)
logger._container_cache[f"{identity}:session-{i}"] = (
logger._container_cache[f"{identity}:session-{i}"][0],
logger._container_cache[f"{identity}:session-{i}"][1],
float(i),
identity,
)
assert len(logger._container_cache) == _SESSION_SCOPED_PER_IDENTITY_CAP
await logger._get_or_create_container(
cache_key=f"{identity}:session-new",
identity=identity,
)
assert len(logger._container_cache) == _SESSION_SCOPED_PER_IDENTITY_CAP, (
"adding a new session beyond the cap must evict one entry so total stays bounded"
)
assert f"{identity}:session-0" not in logger._container_cache, (
"the entry with the oldest last_accessed timestamp must be evicted first (LRU)"
)
assert len(sandbox.delete_calls) == 1, "evicted sandbox must be deleted, not just removed from cache"

View file

@ -27,9 +27,11 @@ from litellm.integrations.otel import ( # noqa: E402
OpenTelemetryV2Config,
)
from litellm.integrations.otel.plumbing import providers # noqa: E402
from litellm.integrations.otel.plumbing.context import (
from litellm.integrations.otel.plumbing.context import ( # noqa: E402
reset_mcp_message_trace_carrier,
set_mcp_message_trace_carrier,
set_request_root_span,
) # noqa: E402
)
from litellm.integrations.otel.logger import OpenTelemetryV2 # noqa: E402
from litellm.integrations.otel.model.spans import ( # noqa: E402
LITELLM_PROXY_REQUEST_SPAN_NAME,
@ -53,8 +55,10 @@ def _reset_request_root_span():
from litellm.integrations.otel.plumbing import context as _otel_context
_otel_context._request_root_span.set(None)
_otel_context._mcp_message_trace_carrier.set(None)
yield
_otel_context._request_root_span.set(None)
_otel_context._mcp_message_trace_carrier.set(None)
def _payload(**overrides):
@ -387,6 +391,190 @@ def test_mcp_tool_call_metadata_read_from_nested_metadata_not_top_level():
assert LiteLLM.MCP_SERVER_NAME not in span.attributes
def _mcp_list_payload(**overrides):
payload = {
"call_type": "list_mcp_tools",
"status": "success",
"litellm_call_id": "mcp_list_1",
"metadata": {
"user_api_key_team_id": "t1",
"spend_logs_metadata": {"mcp_operation": "list_tools"},
},
"hidden_params": {},
}
payload.update(overrides)
return payload
def test_mcp_list_tools_emits_client_span():
"""An MCP ``tools/list`` discovery call becomes a CLIENT span named ``tools/list``,
carrying only the MCP method and the call id. Per the GenAI MCP semconv the list
span omits ``gen_ai.operation.name`` and ``gen_ai.tool.name`` (tool-call-only) and
``mcp.session.id`` (the list path threads no session id), so a naive reuse of the
tool-call mapper would wrongly stamp them, and the pre-fix code emitted no span at
all for a ``list_mcp_tools`` payload."""
logger, exporter = _logger()
kwargs = {"standard_logging_object": _mcp_list_payload()}
asyncio.run(logger.async_log_success_event(kwargs, None, None, None))
(span,) = exporter.get_finished_spans()
assert span.name == "tools/list"
assert span.kind is SpanKind.CLIENT
assert span.attributes["mcp.method.name"] == "tools/list"
assert span.attributes[LiteLLM.CALL_ID] == "mcp_list_1"
assert span.status.status_code is StatusCode.UNSET
# Bug-killers: no span pre-fix (empty exporter -> the unpack above raises), and a
# tool-call-shaped fix would leak execute_tool / tool name / session id here.
assert GenAI.OPERATION_NAME not in span.attributes
assert "gen_ai.tool.name" not in span.attributes
assert "mcp.session.id" not in span.attributes
_MCP_SPAN_CASES = [
(_mcp_payload, "tools/call get_weather"),
(_mcp_list_payload, "tools/list"),
]
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_roots_and_links_transport_without_propagated_context(
make_payload, span_name
):
"""MCP and the HTTP transport are independent lifecycles (one streamable-HTTP
session multiplexes many messages), so per the MCP semconv the message span
must NOT nest under the session/transport span — that is what made it render
skewed at the session's start. With no propagated ``params._meta`` context it
starts its own root trace and records the transport span as a *link*, never
the parent."""
logger, exporter = _logger()
transport = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
set_request_root_span(transport)
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": make_payload()}, None, None, None
)
)
transport.end()
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
assert span.parent is None
assert span.context.trace_id != transport.get_span_context().trace_id
assert [link.context.span_id for link in span.links] == [
transport.get_span_context().span_id
]
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_parents_to_propagated_meta_trace_context(make_payload, span_name):
"""When the client propagates W3C trace context in the request's
``params._meta`` (SEP-414), the MCP span parents to it (one distributed trace)
and still links the transport span — never falling through to the
ambient/session span."""
logger, exporter = _logger()
transport = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
set_request_root_span(transport)
token = set_mcp_message_trace_carrier(
{"traceparent": "00-11111111111111111111111111111111-2222222222222222-01"}
)
try:
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": make_payload()}, None, None, None
)
)
finally:
reset_mcp_message_trace_carrier(token)
transport.end()
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
assert span.context.trace_id == 0x11111111111111111111111111111111
assert span.parent is not None
assert span.parent.span_id == 0x2222222222222222
assert [link.context.span_id for link in span.links] == [
transport.get_span_context().span_id
]
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_ignores_client_supplied_baggage(make_payload, span_name):
"""The MCP span must NOT honor W3C Baggage from the client's ``params._meta``.
``params._meta`` is caller-controlled and the baggage processor stamps
allowlisted baggage keys onto every span, so extracting remote baggage would
let a client spoof a span's identity (e.g. ``litellm.team.id``). The propagator
extracts trace context only, so the spoofed keys never reach the span while the
legitimate traceparent parenting still works."""
logger, exporter = _logger()
transport = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
set_request_root_span(transport)
token = set_mcp_message_trace_carrier(
{
"traceparent": "00-11111111111111111111111111111111-2222222222222222-01",
"baggage": "litellm.team.id=spoofed-team,litellm.metadata.user_api_key_user_id=attacker",
}
)
try:
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": make_payload()}, None, None, None
)
)
finally:
reset_mcp_message_trace_carrier(token)
transport.end()
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
# Trace context still honored: proves the carrier was processed, not dropped wholesale.
assert span.parent is not None and span.parent.span_id == 0x2222222222222222
# Identity is the authenticated payload's team, never the client's spoofed value.
assert span.attributes[LiteLLM.TEAM_ID] == "t1"
assert "litellm.metadata.user_api_key_user_id" not in span.attributes
@pytest.mark.parametrize("make_payload, span_name", _MCP_SPAN_CASES)
def test_mcp_span_carries_authenticated_identity(make_payload, span_name):
"""An MCP span is labeled with the authenticated request's identity (team/key),
seeded from the parsed payload like the LLM-call span. Without this seeding the
span — parented to an empty remote context — would carry no team/key attribute at
all, so it couldn't be attributed or filtered by team in the traces backend."""
logger, exporter = _logger()
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": make_payload()}, None, None, None
)
)
span = next(s for s in exporter.get_finished_spans() if s.name == span_name)
assert span.attributes[LiteLLM.TEAM_ID] == "t1"
def test_mcp_span_malformed_traceparent_starts_root():
"""A malformed traceparent in ``params._meta`` must not crash or parent to a
bogus span: the propagator ignores it, so the span starts its own root trace and
still links the transport span."""
logger, exporter = _logger()
transport = logger._emitter.start_span(
SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME
)
set_request_root_span(transport)
token = set_mcp_message_trace_carrier({"traceparent": "not-a-valid-traceparent"})
try:
asyncio.run(
logger.async_log_success_event(
{"standard_logging_object": _mcp_list_payload()}, None, None, None
)
)
finally:
reset_mcp_message_trace_carrier(token)
transport.end()
span = next(s for s in exporter.get_finished_spans() if s.name == "tools/list")
assert span.parent is None
assert [link.context.span_id for link in span.links] == [
transport.get_span_context().span_id
]
def test_pre_call_idempotent_keeps_first_span():
"""A retried call may re-enter ``pre_call`` with the same call id; the first
span (with the true start time) is kept, not replaced."""

View file

@ -94,19 +94,32 @@ def test_registry_parent_integrity_no_orphans():
def test_registry_hierarchy_shape():
assert set(root_roles()) == {SpanRole.PROXY_REQUEST}
# MCP roles have no in-process parent: per the MCP semconv they root (or adopt
# the client's propagated _meta context), so they sit alongside PROXY_REQUEST.
assert set(root_roles()) == {
SpanRole.PROXY_REQUEST,
SpanRole.MCP_TOOL_CALL,
SpanRole.MCP_LIST_TOOLS,
}
# Guardrails parent to the request span, not the LLM call: a pre-call
# guardrail runs before the LLM call exists, so it's a sibling of it.
assert set(child_roles(SpanRole.PROXY_REQUEST)) == {
SpanRole.LLM_CALL,
SpanRole.MCP_TOOL_CALL,
SpanRole.GUARDRAIL,
SpanRole.DB_CALL,
SpanRole.SERVICE,
}
assert SPAN_REGISTRY[SpanRole.LLM_CALL].kind is LiteLLMSpanKind.CLIENT
# The proxy is an MCP client to the upstream tool server: CLIENT span.
# The proxy is an MCP client to the upstream tool server: CLIENT span. Listing
# tools is the same client relationship, so it's a CLIENT span too.
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].kind is LiteLLMSpanKind.CLIENT
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].kind is LiteLLMSpanKind.CLIENT
# MCP spans don't nest under the transport: they link the PROXY_REQUEST span
# instead of parenting to it (OTel GenAI MCP semconv).
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].parent is None
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].parent is None
assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].links is SpanRole.PROXY_REQUEST
assert SPAN_REGISTRY[SpanRole.MCP_LIST_TOOLS].links is SpanRole.PROXY_REQUEST
assert SPAN_REGISTRY[SpanRole.PROXY_REQUEST].kind is LiteLLMSpanKind.SERVER
assert SPAN_REGISTRY[SpanRole.GUARDRAIL].parent is SpanRole.PROXY_REQUEST
# An outbound datastore call is a CLIENT span; an internal service is INTERNAL.

View file

@ -0,0 +1,64 @@
"""Regression tests for the SDK-free OTel runtime shim.
The proxy auth hot path calls ``phase_span`` and ``seed_request_identity`` on
every request. These wrappers resolve the SDK-backed implementations with a
lazy import. CPython never caches a failed import, so before memoization an
absent OTel SDK made every request re-scan ``sys.path`` and contend on the
import lock. These tests pin the import to a single resolution.
"""
import builtins
import litellm.integrations.otel.runtime as runtime
def test_logger_not_reimported_after_first_resolution(monkeypatch):
runtime._otel_runtime.cache_clear()
counts = {"n": 0}
real_import = builtins.__import__
def counting_import(name, globals=None, locals=None, fromlist=(), level=0):
if name == "litellm.integrations.otel" and fromlist and "logger" in fromlist:
counts["n"] += 1
return real_import(name, globals, locals, fromlist, level)
monkeypatch.setattr(builtins, "__import__", counting_import)
with runtime.phase_span("auth /v1/chat/completions"):
pass
after_first = counts["n"]
for _ in range(49):
with runtime.phase_span("auth /v1/chat/completions"):
pass
assert counts["n"] == after_first, (
f"otel.logger re-imported {counts['n'] - after_first} times after the first "
"resolution; it must be memoized so it does not re-scan sys.path per request"
)
runtime._otel_runtime.cache_clear()
def test_resolution_is_memoized():
runtime._otel_runtime.cache_clear()
for _ in range(25):
with runtime.phase_span("p"):
pass
info = runtime._otel_runtime.cache_info()
assert info.misses == 1
assert info.hits >= 24
runtime._otel_runtime.cache_clear()
def test_wrappers_no_op_when_runtime_absent(monkeypatch):
monkeypatch.setattr(runtime, "_otel_runtime", lambda: None)
with runtime.phase_span("auth") as span:
assert span is None
assert runtime.seed_request_identity({"token": "sk-x"}, model="gpt-4o") is None

View file

@ -1,3 +1,4 @@
import copy
import datetime
import json
import os
@ -9,9 +10,7 @@ from unittest.mock import ANY, MagicMock, Mock, patch
import httpx
import pytest
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system-path
sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system-path
import litellm
from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
@ -93,13 +92,9 @@ async def test_anthropic_cache_control_hook_system_message():
# Verify that cache control was applied (Bedrock transforms it to a separate item)
cache_control_count = sum(
1
for item in request_body["system"]
if isinstance(item, dict) and "cachePoint" in item
1 for item in request_body["system"] if isinstance(item, dict) and "cachePoint" in item
)
assert (
cache_control_count == 1
), f"Expected exactly 1 cache control point, found {cache_control_count}"
assert cache_control_count == 1, f"Expected exactly 1 cache control point, found {cache_control_count}"
@pytest.mark.asyncio
@ -171,9 +166,7 @@ async def test_anthropic_cache_control_hook_user_message():
print("request_body: ", json.dumps(request_body, indent=4))
# Verify the request body
assert request_body["messages"][1]["content"][1]["cachePoint"] == {
"type": "default"
}
assert request_body["messages"][1]["content"][1]["cachePoint"] == {"type": "default"}
@pytest.mark.asyncio
@ -262,14 +255,10 @@ async def test_anthropic_cache_control_hook_negative_indices():
# Verify the last message (input index -1 -> request index 2) has cache control
last_message_content = request_body["messages"][2]["content"]
assert isinstance(
last_message_content, list
), "Last message content should be a list"
assert any(
"cachePoint" in item
for item in last_message_content
if isinstance(item, dict)
), "CachePoint missing in last message"
assert isinstance(last_message_content, list), "Last message content should be a list"
assert any("cachePoint" in item for item in last_message_content if isinstance(item, dict)), (
"CachePoint missing in last message"
)
# Note: Based on debug output, the hook correctly applies cache control to both messages,
# but the Bedrock API transformation appears to only preserve cache control for user messages,
@ -278,30 +267,20 @@ async def test_anthropic_cache_control_hook_negative_indices():
# The second-to-last message (assistant) gets cache_control from the hook but loses it
# during API transformation. This test documents this behavior.
second_last_message_content = request_body["messages"][1]["content"]
assert isinstance(
second_last_message_content, list
), "Second-to-last message content should be a list"
assert isinstance(second_last_message_content, list), "Second-to-last message content should be a list"
# Check if assistant message cache control is preserved (currently it's not)
assistant_has_cache_control = any(
"cachePoint" in item
for item in second_last_message_content
if isinstance(item, dict)
)
print(
f"Assistant message has cache control in final request: {assistant_has_cache_control}"
"cachePoint" in item for item in second_last_message_content if isinstance(item, dict)
)
print(f"Assistant message has cache control in final request: {assistant_has_cache_control}")
# Verify the first user message (request index 0) was NOT modified
first_user_message_content = request_body["messages"][0]["content"]
assert isinstance(
first_user_message_content, list
), "First user message content should be a list"
assert not any(
"cachePoint" in item
for item in first_user_message_content
if isinstance(item, dict)
), "CachePoint unexpectedly found in first user message"
assert isinstance(first_user_message_content, list), "First user message content should be a list"
assert not any("cachePoint" in item for item in first_user_message_content if isinstance(item, dict)), (
"CachePoint unexpectedly found in first user message"
)
@pytest.mark.asyncio
@ -342,9 +321,7 @@ async def test_anthropic_cache_control_hook_out_of_bounds_logging():
client = AsyncHTTPHandler()
# Mock the verbose_logger to capture warning calls
with patch(
"litellm.integrations.anthropic_cache_control_hook.verbose_logger"
) as mock_logger:
with patch("litellm.integrations.anthropic_cache_control_hook.verbose_logger") as mock_logger:
with patch.object(client, "post", return_value=mock_response) as mock_post:
messages = [
{"role": "user", "content": "Message 1"},
@ -354,9 +331,7 @@ async def test_anthropic_cache_control_hook_out_of_bounds_logging():
await litellm.acompletion(
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
cache_control_injection_points=[
{"location": "message", "index": 10}
], # Out of bounds index
cache_control_injection_points=[{"location": "message", "index": 10}], # Out of bounds index
client=client,
)
@ -365,10 +340,7 @@ async def test_anthropic_cache_control_hook_out_of_bounds_logging():
warning_call = mock_logger.warning.call_args[0][0]
# Check that the warning message contains the expected information
assert (
"AnthropicCacheControlHook: Provided index 10 is out of bounds"
in warning_call
)
assert "AnthropicCacheControlHook: Provided index 10 is out of bounds" in warning_call
assert "message list of length 2" in warning_call
assert "Targeted index was 10" in warning_call
assert "Skipping cache control injection for this point" in warning_call
@ -411,9 +383,7 @@ async def test_anthropic_cache_control_hook_negative_out_of_bounds_logging():
client = AsyncHTTPHandler()
# Mock the verbose_logger to capture warning calls
with patch(
"litellm.integrations.anthropic_cache_control_hook.verbose_logger"
) as mock_logger:
with patch("litellm.integrations.anthropic_cache_control_hook.verbose_logger") as mock_logger:
with patch.object(client, "post", return_value=mock_response) as mock_post:
messages = [
{"role": "user", "content": "Single message"},
@ -436,14 +406,9 @@ async def test_anthropic_cache_control_hook_negative_out_of_bounds_logging():
warning_call = mock_logger.warning.call_args[0][0]
# Check that the warning message contains the original negative index
assert (
"AnthropicCacheControlHook: Provided index -5 is out of bounds"
in warning_call
)
assert "AnthropicCacheControlHook: Provided index -5 is out of bounds" in warning_call
assert "message list of length 1" in warning_call
assert (
"Targeted index was -4" in warning_call
) # -5 + 1 = -4 (converted index)
assert "Targeted index was -4" in warning_call # -5 + 1 = -4 (converted index)
assert "Skipping cache control injection for this point" in warning_call
@ -531,15 +496,11 @@ async def test_anthropic_cache_control_hook_multiple_user_messages():
# Count cache control points - should have 2 since both injection points were applied
cache_control_count = sum(
1
for item in combined_message_content
if isinstance(item, dict) and "cachePoint" in item
1 for item in combined_message_content if isinstance(item, dict) and "cachePoint" in item
)
assert cache_control_count == 2
print(
f"Found {cache_control_count} cache control points in the combined message"
)
print(f"Found {cache_control_count} cache control points in the combined message")
@pytest.mark.asyncio
@ -588,9 +549,7 @@ async def test_anthropic_cache_control_hook_out_of_bounds(bad_index):
await litellm.acompletion(
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
cache_control_injection_points=[
{"location": "message", "index": bad_index}
],
cache_control_injection_points=[{"location": "message", "index": bad_index}],
client=client,
)
@ -601,19 +560,13 @@ async def test_anthropic_cache_control_hook_out_of_bounds(bad_index):
for msg in request_body["messages"]:
content = msg.get("content", [])
if isinstance(content, list):
assert not any(
"cachePoint" in item
for item in content
if isinstance(item, dict)
)
assert not any("cachePoint" in item for item in content if isinstance(item, dict))
@pytest.mark.asyncio
@pytest.mark.parametrize(
"message_list",
[
[{"role": "user", "content": "Single message"}]
], # Single message only - empty list will fail at API level
[[{"role": "user", "content": "Single message"}]], # Single message only - empty list will fail at API level
)
async def test_anthropic_cache_control_hook_single_message(message_list):
"""
@ -662,9 +615,7 @@ async def test_anthropic_cache_control_hook_single_message(message_list):
# For the single message, verify cache control was applied
content = request_body["messages"][0]["content"]
assert isinstance(content, list)
assert any(
"cachePoint" in item for item in content if isinstance(item, dict)
)
assert any("cachePoint" in item for item in content if isinstance(item, dict))
@pytest.mark.asyncio
@ -693,9 +644,7 @@ async def test_anthropic_cache_control_hook_empty_message_list():
await litellm.acompletion(
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=[],
cache_control_injection_points=[
{"location": "message", "index": -1}
],
cache_control_injection_points=[{"location": "message", "index": -1}],
client=client,
)
@ -755,11 +704,7 @@ async def test_anthropic_cache_control_hook_no_op():
for msg in request_body["messages"]:
content = msg.get("content", [])
if isinstance(content, list):
assert not any(
"cachePoint" in item
for item in content
if isinstance(item, dict)
)
assert not any("cachePoint" in item for item in content if isinstance(item, dict))
@pytest.mark.asyncio
@ -827,14 +772,10 @@ async def test_anthropic_cache_control_hook_multiple_content_items_last_only():
message_content = request_body["messages"][0]["content"]
assert isinstance(message_content, list)
cache_control_count = sum(
1
for item in message_content
if isinstance(item, dict) and "cachePoint" in item
cache_control_count = sum(1 for item in message_content if isinstance(item, dict) and "cachePoint" in item)
assert cache_control_count == 1, (
f"Expected exactly 1 cache control point, found {cache_control_count}. This test verifies the fix for issue 15696 where cache_control was incorrectly applied to ALL content items."
)
assert (
cache_control_count == 1
), f"Expected exactly 1 cache control point, found {cache_control_count}. This test verifies the fix for issue 15696 where cache_control was incorrectly applied to ALL content items."
@pytest.mark.asyncio
@ -891,30 +832,22 @@ async def test_anthropic_cache_control_hook_document_analysis_multiple_pages():
],
}
],
cache_control_injection_points=[
{"location": "message", "role": "user"}
],
cache_control_injection_points=[{"location": "message", "role": "user"}],
client=client,
)
mock_post.assert_called_once()
request_body = json.loads(mock_post.call_args.kwargs["data"])
print(
"Document analysis request_body: ", json.dumps(request_body, indent=4)
)
print("Document analysis request_body: ", json.dumps(request_body, indent=4))
message_content = request_body["messages"][0]["content"]
assert isinstance(message_content, list)
cache_control_count = sum(
1
for item in message_content
if isinstance(item, dict) and "cachePoint" in item
cache_control_count = sum(1 for item in message_content if isinstance(item, dict) and "cachePoint" in item)
assert cache_control_count == 1, (
f"Expected exactly 1 cache control point (last item only), found {cache_control_count}. Before fix, this would be 6 (one for each content item)."
)
assert (
cache_control_count == 1
), f"Expected exactly 1 cache control point (last item only), found {cache_control_count}. Before fix, this would be 6 (one for each content item)."
def test_gemini_cache_control_injection_points_detected():
@ -1076,13 +1009,8 @@ async def test_anthropic_cache_control_hook_string_negative_index():
# The last user message should have cache control applied
last_message = request_body["messages"][-1]
last_message_content = last_message["content"]
assert isinstance(
last_message_content, list
), f"Expected list content, got {type(last_message_content)}"
has_cache_point = any(
isinstance(item, dict) and "cachePoint" in item
for item in last_message_content
)
assert isinstance(last_message_content, list), f"Expected list content, got {type(last_message_content)}"
has_cache_point = any(isinstance(item, dict) and "cachePoint" in item for item in last_message_content)
assert has_cache_point, (
f"Expected cachePoint in last message content, got: {last_message_content}. "
"String index '-1' was not parsed correctly (str.isdigit() returns False for negative strings)."
@ -1146,17 +1074,13 @@ def test_cache_control_hook_caps_at_four_blocks_with_client_cache_control():
_, processed, _ = hook.get_chat_completion_prompt(
model="bedrock/us.anthropic.claude-opus-4-6-v1:0",
messages=messages,
non_default_params={
"cache_control_injection_points": _build_injection_points()
},
non_default_params={"cache_control_injection_points": _build_injection_points()},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
)
assert (
_count_cache_control(processed) == 4
), "Hook must cap cache_control at Anthropic's limit of 4 blocks"
assert _count_cache_control(processed) == 4, "Hook must cap cache_control at Anthropic's limit of 4 blocks"
# Client TTL on system blocks must be preserved (not overwritten by config).
for i in range(4):
@ -1170,11 +1094,7 @@ def test_cache_control_hook_caps_at_four_blocks_with_client_cache_control():
assert user_message.get("cache_control") is None
user_content = user_message.get("content")
if isinstance(user_content, list):
assert all(
block.get("cache_control") is None
for block in user_content
if isinstance(block, dict)
)
assert all(block.get("cache_control") is None for block in user_content if isinstance(block, dict))
def test_cache_control_hook_caps_at_four_blocks_without_client_cache_control():
@ -1184,17 +1104,13 @@ def test_cache_control_hook_caps_at_four_blocks_without_client_cache_control():
"""
hook = AnthropicCacheControlHook()
messages: List[AllMessageValues] = [
{"role": "system", "content": f"System {i}"} for i in range(4)
]
messages: List[AllMessageValues] = [{"role": "system", "content": f"System {i}"} for i in range(4)]
messages.append({"role": "user", "content": "hello"})
_, processed, _ = hook.get_chat_completion_prompt(
model="bedrock/us.anthropic.claude-opus-4-6-v1:0",
messages=messages,
non_default_params={
"cache_control_injection_points": _build_injection_points()
},
non_default_params={"cache_control_injection_points": _build_injection_points()},
prompt_id=None,
prompt_variables=None,
dynamic_callback_params={},
@ -1303,18 +1219,12 @@ async def test_cache_control_hook_bedrock_payload_caps_cachepoints_at_four():
request_body = json.loads(mock_post.call_args.kwargs["data"])
cache_points = sum(
1
for block in request_body.get("system", [])
if isinstance(block, dict) and "cachePoint" in block
1 for block in request_body.get("system", []) if isinstance(block, dict) and "cachePoint" in block
)
for msg in request_body.get("messages", []):
content = msg.get("content", [])
if isinstance(content, list):
cache_points += sum(
1
for block in content
if isinstance(block, dict) and "cachePoint" in block
)
cache_points += sum(1 for block in content if isinstance(block, dict) and "cachePoint" in block)
assert cache_points <= 4, (
f"Bedrock payload exceeded Anthropic's 4 cache_control block limit: "
@ -1331,9 +1241,7 @@ def test_cache_control_hook_reserves_slot_for_tool_config_point():
"""
hook = AnthropicCacheControlHook()
messages: List[AllMessageValues] = [
{"role": "system", "content": f"System {i}"} for i in range(4)
]
messages: List[AllMessageValues] = [{"role": "system", "content": f"System {i}"} for i in range(4)]
messages.append({"role": "user", "content": "hello"})
_, processed, non_default_params = hook.get_chat_completion_prompt(
@ -1356,9 +1264,7 @@ def test_cache_control_hook_reserves_slot_for_tool_config_point():
assert _count_cache_control(processed) == 3
# The tool_config point is passed through for the provider transform.
assert non_default_params["cache_control_injection_points"] == [
{"location": "tool_config"}
]
assert non_default_params["cache_control_injection_points"] == [{"location": "tool_config"}]
@pytest.mark.asyncio
@ -1384,9 +1290,7 @@ async def test_cache_control_hook_bedrock_payload_caps_with_tool_config_point():
client = AsyncHTTPHandler()
with patch.object(client, "post", return_value=mock_response) as mock_post:
messages = [
{"role": "system", "content": f"System block {i}"} for i in range(4)
]
messages = [{"role": "system", "content": f"System block {i}"} for i in range(4)]
messages.append({"role": "user", "content": "What is the weather?"})
await litellm.acompletion(
@ -1421,18 +1325,12 @@ async def test_cache_control_hook_bedrock_payload_caps_with_tool_config_point():
request_body = json.loads(mock_post.call_args.kwargs["data"])
cache_points = sum(
1
for block in request_body.get("system", [])
if isinstance(block, dict) and "cachePoint" in block
1 for block in request_body.get("system", []) if isinstance(block, dict) and "cachePoint" in block
)
for msg in request_body.get("messages", []):
content = msg.get("content", [])
if isinstance(content, list):
cache_points += sum(
1
for block in content
if isinstance(block, dict) and "cachePoint" in block
)
cache_points += sum(1 for block in content if isinstance(block, dict) and "cachePoint" in block)
for tool in request_body.get("toolConfig", {}).get("tools", []):
if isinstance(tool, dict) and "cachePoint" in tool:
cache_points += 1
@ -1441,3 +1339,197 @@ async def test_cache_control_hook_bedrock_payload_caps_with_tool_config_point():
f"Bedrock payload exceeded Anthropic's 4 cache_control block limit "
f"when mixing message and tool_config injection: found {cache_points}"
)
class TestApplyToAnthropicMessagesRequest:
"""Tests for apply_to_anthropic_messages_request (v1/messages cache control)."""
def test_system_string_injection(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
system = "You are helpful"
injection_points = [{"location": "message", "role": "system"}]
result_msgs, result_sys, remaining = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
assert result_sys == [{"type": "text", "text": "You are helpful", "cache_control": {"type": "ephemeral"}}]
assert result_msgs == messages
assert remaining == []
def test_system_list_injection(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
system = [
{"type": "text", "text": "Part 1"},
{"type": "text", "text": "Part 2"},
]
injection_points = [{"location": "message", "role": "system"}]
_, result_sys, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
assert result_sys[0] == {"type": "text", "text": "Part 1"}
assert result_sys[1] == {"type": "text", "text": "Part 2", "cache_control": {"type": "ephemeral"}}
def test_user_message_injection_by_role(self):
messages = [
{"role": "user", "content": [{"type": "text", "text": "First"}]},
{"role": "assistant", "content": [{"type": "text", "text": "Response"}]},
{"role": "user", "content": [{"type": "text", "text": "Second"}]},
]
injection_points = [{"location": "message", "role": "user"}]
result_msgs, _, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=None,
injection_points=injection_points,
)
assert result_msgs[0]["content"][-1].get("cache_control") == {"type": "ephemeral"}
assert result_msgs[2]["content"][-1].get("cache_control") == {"type": "ephemeral"}
assert result_msgs[1]["content"][-1].get("cache_control") is None
def test_message_injection_by_index(self):
messages = [
{"role": "user", "content": [{"type": "text", "text": "First"}]},
{"role": "assistant", "content": [{"type": "text", "text": "Response"}]},
{"role": "user", "content": [{"type": "text", "text": "Second"}]},
]
injection_points = [{"location": "message", "index": -1}]
result_msgs, _, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=None,
injection_points=injection_points,
)
assert result_msgs[2]["content"][-1].get("cache_control") == {"type": "ephemeral"}
assert result_msgs[0]["content"][-1].get("cache_control") is None
assert result_msgs[1]["content"][-1].get("cache_control") is None
def test_mixed_system_and_message_injection(self):
messages = [
{"role": "user", "content": [{"type": "text", "text": "Hello"}]},
{"role": "assistant", "content": [{"type": "text", "text": "Hi"}]},
{"role": "user", "content": [{"type": "text", "text": "Question"}]},
]
system = "System prompt"
injection_points = [
{"location": "message", "role": "system"},
{"location": "message", "index": -1},
]
result_msgs, result_sys, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
assert result_sys[0]["cache_control"] == {"type": "ephemeral"}
assert result_msgs[2]["content"][-1].get("cache_control") == {"type": "ephemeral"}
def test_respects_max_4_blocks(self):
messages = [{"role": "user", "content": [{"type": "text", "text": f"Msg {i}"}]} for i in range(6)]
system = "System"
injection_points = [
{"location": "message", "role": "system"},
{"location": "message", "role": "user"},
]
result_msgs, result_sys, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
sys_blocks = sum(1 for b in (result_sys or []) if isinstance(b, dict) and b.get("cache_control") is not None)
total_blocks = sys_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(m) for m in result_msgs)
assert total_blocks <= 4
def test_tool_config_points_forwarded_as_remaining(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
injection_points = [
{"location": "message", "role": "user"},
{"location": "tool_config"},
]
_, _, remaining = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=None,
injection_points=injection_points,
)
assert remaining == [{"location": "tool_config"}]
def test_no_injection_points_returns_unchanged(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
system = "System"
result_msgs, result_sys, remaining = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=[],
)
assert result_msgs == messages
assert result_sys == system
assert remaining == []
def test_does_not_mutate_input(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
system = [{"type": "text", "text": "System"}]
injection_points = [{"location": "message", "role": "system"}]
original_system = copy.deepcopy(system)
original_messages = copy.deepcopy(messages)
AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
assert messages == original_messages
assert system == original_system
def test_system_none_with_system_point_skipped(self):
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
injection_points = [{"location": "message", "role": "system"}]
result_msgs, result_sys, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=None,
injection_points=injection_points,
)
assert result_sys is None
def test_existing_cache_control_counted_toward_limit(self):
messages = [
{"role": "user", "content": [{"type": "text", "text": "A", "cache_control": {"type": "ephemeral"}}]},
{"role": "assistant", "content": [{"type": "text", "text": "B", "cache_control": {"type": "ephemeral"}}]},
{"role": "user", "content": [{"type": "text", "text": "C", "cache_control": {"type": "ephemeral"}}]},
{"role": "user", "content": [{"type": "text", "text": "D"}]},
{"role": "user", "content": [{"type": "text", "text": "E"}]},
]
system = "System"
injection_points = [
{"location": "message", "role": "system"},
{"location": "message", "index": 3},
{"location": "message", "index": 4},
]
result_msgs, result_sys, _ = AnthropicCacheControlHook.apply_to_anthropic_messages_request(
messages=messages,
system=system,
injection_points=injection_points,
)
sys_blocks = sum(1 for b in (result_sys or []) if isinstance(b, dict) and b.get("cache_control") is not None)
total_blocks = sys_blocks + sum(AnthropicCacheControlHook._count_cache_control_blocks(m) for m in result_msgs)
assert total_blocks <= 4

View file

@ -5,7 +5,10 @@ Verifies that metadata from x-litellm-spend-logs-metadata header is available
in Prometheus custom labels via combined_metadata.
"""
from litellm.integrations.prometheus import get_custom_labels_from_metadata
from litellm.integrations.prometheus import (
_get_combined_custom_metadata_from_standard_logging_payload,
get_custom_labels_from_metadata,
)
def test_get_custom_labels_includes_spend_logs_metadata(monkeypatch):
@ -109,3 +112,96 @@ def test_combined_metadata_with_none_spend_logs(monkeypatch):
result = get_custom_labels_from_metadata(combined_metadata)
assert result == {"metadata_foo": "bar"}
def test_combined_metadata_includes_top_level_fields():
"""
Regression test for LIT-3741: user_api_key_project_alias (and other
top-level metadata fields) must be included in the combined metadata
so they can be referenced via custom_prometheus_metadata_labels.
"""
standard_logging_payload = {
"metadata": {
"user_api_key_hash": "sk-abc123",
"user_api_key_alias": "hotel-key",
"user_api_key_team_id": "team-1",
"user_api_key_team_alias": "hotel-team",
"user_api_key_project_id": "proj-1",
"user_api_key_project_alias": "hotel-recommendations",
"user_api_key_user_id": "user-1",
"user_api_key_user_email": "user@example.com",
"user_api_key_end_user_id": None,
"user_api_key_org_id": None,
"user_api_key_org_alias": None,
"user_api_key_request_route": "/v1/chat/completions",
"requester_metadata": {"custom_field": "custom_value"},
"user_api_key_auth_metadata": {"auth_field": "auth_value"},
"spend_logs_metadata": None,
}
}
combined = _get_combined_custom_metadata_from_standard_logging_payload(
standard_logging_payload
)
assert combined["user_api_key_project_alias"] == "hotel-recommendations"
assert combined["user_api_key_project_id"] == "proj-1"
assert combined["user_api_key_team_alias"] == "hotel-team"
assert combined["user_api_key_request_route"] == "/v1/chat/completions"
assert combined["custom_field"] == "custom_value"
assert combined["auth_field"] == "auth_value"
def test_project_alias_accessible_via_custom_prometheus_labels(monkeypatch):
"""
Regression test for LIT-3741: configuring
custom_prometheus_metadata_labels with "metadata.user_api_key_project_alias"
should produce a label with the project's alias value.
"""
monkeypatch.setattr(
"litellm.custom_prometheus_metadata_labels",
["metadata.user_api_key_project_alias"],
)
standard_logging_payload = {
"metadata": {
"user_api_key_project_alias": "hotel-recommendations",
"requester_metadata": None,
"user_api_key_auth_metadata": None,
"spend_logs_metadata": None,
}
}
combined = _get_combined_custom_metadata_from_standard_logging_payload(
standard_logging_payload
)
result = get_custom_labels_from_metadata(combined)
assert result == {"metadata_user_api_key_project_alias": "hotel-recommendations"}
def test_project_alias_accessible_without_prefix(monkeypatch):
"""
user_api_key_project_alias should also be accessible without
the "metadata." prefix in custom_prometheus_metadata_labels config.
"""
monkeypatch.setattr(
"litellm.custom_prometheus_metadata_labels",
["user_api_key_project_alias"],
)
standard_logging_payload = {
"metadata": {
"user_api_key_project_alias": "hotel-recommendations",
"requester_metadata": None,
"user_api_key_auth_metadata": None,
"spend_logs_metadata": None,
}
}
combined = _get_combined_custom_metadata_from_standard_logging_payload(
standard_logging_payload
)
result = get_custom_labels_from_metadata(combined)
assert result == {"user_api_key_project_alias": "hotel-recommendations"}

View file

@ -6,7 +6,7 @@ litellm.acompletion() for transparent server-side web search execution.
"""
import os
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import MagicMock
import pytest
@ -34,12 +34,14 @@ def mock_search_response():
@pytest.fixture
def websearch_logger():
"""Create a WebSearchInterceptionLogger instance"""
return WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI, LlmProviders.MINIMAX]
)
return WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI, LlmProviders.MINIMAX])
@pytest.mark.asyncio
@pytest.mark.skipif(
os.environ.get("OPENAI_API_KEY") is None,
reason="OPENAI_API_KEY not set",
)
async def test_websearch_chat_completion_with_openai():
"""Test websearch interception with OpenAI chat completions API.
@ -48,112 +50,56 @@ async def test_websearch_chat_completion_with_openai():
2. Server executes web search automatically
3. Server makes follow-up request with search results
4. User gets final answer without tool_calls
Uses mocked acompletion so no real API key is needed.
"""
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Choices,
Function,
Message,
)
# First call returns a tool_call response; second call returns a final answer.
tool_call_response = ModelResponse(
id="chatcmpl-tool",
choices=[
Choices(
finish_reason="tool_calls",
index=0,
message=Message(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
id="call_001",
type="function",
function=Function(
name="litellm_web_search",
arguments='{"query": "weather San Francisco"}',
),
)
],
),
)
],
model="gpt-4o-mini",
object="chat.completion",
created=1234567890,
)
final_response = ModelResponse(
id="chatcmpl-final",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
role="assistant",
content="The weather in San Francisco today is 65°F and partly cloudy.",
tool_calls=None,
),
)
],
model="gpt-4o-mini",
object="chat.completion",
created=1234567891,
)
call_count = 0
async def mock_acompletion(*args, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
return tool_call_response
return final_response
# Configure WebSearch interception
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI]
)
websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
litellm.callbacks = [websearch_logger]
try:
with patch("litellm.acompletion", side_effect=mock_acompletion), \
patch("litellm.integrations.websearch_interception.handler.litellm.acompletion",
side_effect=mock_acompletion):
response = await litellm.acompletion(
model="gpt-4o-mini",
messages=[
{
"role": "user",
"content": "What's the weather in San Francisco today?",
}
],
tools=[
{
"type": "function",
"function": {
"name": "litellm_web_search",
"description": "Search the web for information",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "Search query",
}
},
"required": ["query"],
response = await litellm.acompletion(
model="gpt-4o-mini", # Use cheaper model for testing
messages=[
{
"role": "user",
"content": "What's the weather in San Francisco today?",
}
],
tools=[
{
"type": "function",
"function": {
"name": "litellm_web_search",
"description": "Search the web for information",
"parameters": {
"type": "object",
"properties": {
"query": {
"type": "string",
"description": "Search query",
}
},
"required": ["query"],
},
}
],
)
},
}
],
)
# Verify response structure
assert isinstance(response, ModelResponse)
assert response.choices[0].message.content is not None
assert len(response.choices[0].message.content) > 0
# If agentic loop worked, we should NOT have tool_calls in final response
# (they should have been executed and replaced with final answer)
if hasattr(response.choices[0].message, "tool_calls"):
# If tool_calls exist, it means agentic loop didn't run
# This could happen if search tool is not configured
pytest.skip("Agentic loop did not execute - search tool may not be configured")
# Verify we got a meaningful response
assert response.choices[0].finish_reason in ["stop", "end_turn"]
finally:
# Restore original callbacks
@ -170,9 +116,7 @@ async def test_websearch_chat_completion_hook_detection():
Message,
)
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI]
)
websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
# Mock response with litellm_web_search tool call
mock_response = ModelResponse(
@ -203,21 +147,19 @@ async def test_websearch_chat_completion_hook_detection():
)
# Test should_run_chat_completion_agentic_loop
should_run, tools_dict = (
await websearch_logger.async_should_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=[
{
"type": "function",
"function": {"name": "litellm_web_search"},
}
],
stream=False,
custom_llm_provider="openai",
kwargs={},
)
should_run, tools_dict = await websearch_logger.async_should_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "What's the weather?"}],
tools=[
{
"type": "function",
"function": {"name": "litellm_web_search"},
}
],
stream=False,
custom_llm_provider="openai",
kwargs={},
)
# Verify hook detected the tool call
@ -233,9 +175,7 @@ async def test_websearch_not_triggered_without_tool():
"""Test that websearch hook is NOT triggered when no web search tool in request."""
from litellm.types.utils import Choices, Message
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.OPENAI]
)
websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
mock_response = ModelResponse(
id="test-123",
@ -256,21 +196,19 @@ async def test_websearch_not_triggered_without_tool():
)
# Test without web search tool
should_run, tools_dict = (
await websearch_logger.async_should_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
tools=[
{
"type": "function",
"function": {"name": "some_other_tool"},
}
],
stream=False,
custom_llm_provider="openai",
kwargs={},
)
should_run, tools_dict = await websearch_logger.async_should_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
tools=[
{
"type": "function",
"function": {"name": "some_other_tool"},
}
],
stream=False,
custom_llm_provider="openai",
kwargs={},
)
# Verify hook did NOT trigger
@ -289,9 +227,7 @@ async def test_websearch_not_triggered_for_disabled_provider():
)
# Only enable bedrock
websearch_logger = WebSearchInterceptionLogger(
enabled_providers=[LlmProviders.BEDROCK]
)
websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.BEDROCK])
mock_response = ModelResponse(
id="test-123",
@ -321,21 +257,19 @@ async def test_websearch_not_triggered_for_disabled_provider():
)
# Test with OpenAI provider (not enabled)
should_run, tools_dict = (
await websearch_logger.async_should_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "test"}],
tools=[
{
"type": "function",
"function": {"name": "litellm_web_search"},
}
],
stream=False,
custom_llm_provider="openai", # Not in enabled_providers
kwargs={},
)
should_run, tools_dict = await websearch_logger.async_should_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "test"}],
tools=[
{
"type": "function",
"function": {"name": "litellm_web_search"},
}
],
stream=False,
custom_llm_provider="openai", # Not in enabled_providers
kwargs={},
)
# Verify hook did NOT trigger
@ -388,9 +322,11 @@ async def test_websearch_json_serialization_fix():
@pytest.mark.asyncio
@pytest.mark.skipif(
os.environ.get("OPENAI_API_KEY") is None or os.environ.get("PERPLEXITY_API_KEY") is None,
reason="OPENAI_API_KEY or PERPLEXITY_API_KEY not set",
)
async def test_websearch_streaming_conversion():
if not os.environ.get("OPENAI_API_KEY") or not os.environ.get("PERPLEXITY_API_KEY"):
pytest.skip("OPENAI_API_KEY or PERPLEXITY_API_KEY not set")
"""Test that streaming requests are converted to non-streaming for web search.
When stream=True is passed with web search tools, the handler should:
@ -440,6 +376,174 @@ async def test_websearch_streaming_conversion():
litellm.callbacks = []
@pytest.mark.asyncio
async def test_maybe_run_chat_completion_agentic_loop_calls_chat_completion_hook():
"""Regression test: maybe_run_chat_completion_agentic_loop must call
async_should_run_chat_completion_agentic_loop, not async_should_run_agentic_loop.
Before the fix, the function used the wrong gate check and wrong hook,
causing WebSearchInterceptionLogger to never intercept chat completion requests
even when the LLM returned a litellm_web_search tool call.
"""
from litellm.litellm_core_utils.chat_completion_agentic_loop import (
maybe_run_chat_completion_agentic_loop,
)
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Choices,
Function,
Message,
)
mock_response = ModelResponse(
id="test-regression-123",
choices=[
Choices(
finish_reason="tool_calls",
index=0,
message=Message(
role="assistant",
content=None,
tool_calls=[
ChatCompletionMessageToolCall(
id="call_abc",
type="function",
function=Function(
name="litellm_web_search",
arguments='{"query": "latest news"}',
),
)
],
),
)
],
model="gpt-4o",
object="chat.completion",
created=1234567890,
)
sentinel = ModelResponse(
id="sentinel-final",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(role="assistant", content="Here is the news."),
)
],
model="gpt-4o",
object="chat.completion",
created=1234567890,
)
websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
chat_completion_hook_called = False
async def fake_should_run_chat_completion(response, model, messages, tools, stream, custom_llm_provider, kwargs):
nonlocal chat_completion_hook_called
chat_completion_hook_called = True
return True, {
"tool_calls": [{"id": "call_abc", "name": "litellm_web_search", "input": {"query": "latest news"}}],
"tool_type": "websearch",
"provider": "openai",
"response_format": "openai",
}
async def fake_build_plan(tools, model, messages, response, optional_params, logging_obj, stream, kwargs):
from litellm.types.integrations.custom_logger import AgenticLoopPlan
return AgenticLoopPlan(run_agentic_loop=False, response_override=sentinel)
websearch_logger.async_should_run_chat_completion_agentic_loop = fake_should_run_chat_completion
websearch_logger.async_build_chat_completion_agentic_loop_plan = fake_build_plan
import litellm as _litellm
original_callbacks = _litellm.callbacks[:]
_litellm.callbacks = [websearch_logger]
mock_logging_obj = MagicMock()
mock_logging_obj.dynamic_success_callbacks = None
try:
result = await maybe_run_chat_completion_agentic_loop(
response=mock_response,
model="gpt-4o",
messages=[{"role": "user", "content": "Latest news?"}],
optional_params={
"tools": [
{
"type": "function",
"function": {"name": "litellm_web_search"},
}
]
},
kwargs={},
logging_obj=mock_logging_obj,
custom_llm_provider="openai",
stream=False,
)
finally:
_litellm.callbacks = original_callbacks
assert chat_completion_hook_called, (
"async_should_run_chat_completion_agentic_loop was never called; "
"maybe_run_chat_completion_agentic_loop used the wrong hook"
)
assert result is sentinel, "Expected agentic loop to return sentinel final response"
@pytest.mark.asyncio
async def test_execute_chat_completion_agentic_loop_strips_tool_choice():
"""Regression: _execute_chat_completion_agentic_loop must not forward tool_choice
from the original request into the follow-up synthesis call.
When the original request forces tool_choice to litellm_web_search, merging
optional_params into the follow-up params without explicit removal causes the
model to call the search tool again instead of synthesizing an answer.
"""
from unittest.mock import patch
websearch_logger = WebSearchInterceptionLogger(enabled_providers=[LlmProviders.OPENAI])
captured_kwargs: dict = {}
async def fake_acompletion(**kwargs):
captured_kwargs.update(kwargs)
return ModelResponse(id="followup", model="gpt-4o", object="chat.completion")
async def fake_search(query):
return ("Bitcoin price is $60,000", None)
with patch.object(websearch_logger, "_execute_search", side_effect=fake_search):
with patch("litellm.acompletion", side_effect=fake_acompletion):
await websearch_logger._execute_chat_completion_agentic_loop(
model="gpt-4o",
messages=[{"role": "user", "content": "What is Bitcoin price?"}],
tool_calls=[
{
"id": "call_1",
"name": "litellm_web_search",
"input": {"query": "bitcoin price"},
}
],
optional_params={
"tools": [{"type": "function", "function": {"name": "litellm_web_search"}}],
"tool_choice": {"type": "function", "function": {"name": "litellm_web_search"}},
"max_tokens": 512,
},
logging_obj=MagicMock(),
stream=False,
kwargs={},
)
assert "tool_choice" not in captured_kwargs, (
"tool_choice must not appear in follow-up acompletion kwargs; "
"it would force the model to call the search tool again instead of synthesizing"
)
if __name__ == "__main__":
# Run with: pytest test_websearch_chat_completion.py -v -s
pytest.main([__file__, "-v", "-s"])

View file

@ -181,8 +181,7 @@ async def test_internal_control_fields_never_leak_into_provider_body(restore_cal
# The loop must have actually fired (sanity: two provider calls).
assert create.await_count == 2, (
"expected the agentic loop to issue a follow-up provider call; "
f"got {create.await_count} call(s)"
f"expected the agentic loop to issue a follow-up provider call; got {create.await_count} call(s)"
)
for idx, call in enumerate(create.await_args_list):
@ -194,8 +193,7 @@ async def test_internal_control_fields_never_leak_into_provider_body(restore_cal
f"top-level request body: {sorted(body.keys())}"
)
assert field not in extra_body, (
f"provider call #{idx}: internal field {field!r} leaked into "
f"extra_body: {sorted(extra_body.keys())}"
f"provider call #{idx}: internal field {field!r} leaked into extra_body: {sorted(extra_body.keys())}"
)
# The native code_interpreter tool must have been swapped for the
# function tool, never sent raw to OpenAI as a chat-completions request.
@ -254,9 +252,7 @@ class _GateOnlyLogger(CustomLogger):
) -> AgenticLoopPlan:
return self._plan
async def async_agentic_loop_cleanup_hook(
self, plan: AgenticLoopPlan, kwargs: Dict[str, Any]
) -> None:
async def async_agentic_loop_cleanup_hook(self, plan: AgenticLoopPlan, kwargs: Dict[str, Any]) -> None:
self.cleanup_calls += 1
@ -343,9 +339,7 @@ async def test_dispatcher_runs_followup_with_incremented_depth_and_patched_messa
assert call_kwargs["max_agentic_loops"] >= 1
assert "_agentic_loop_fingerprints" in call_kwargs
# Interception markers are mirrored into litellm_metadata for the follow-up.
assert (
call_kwargs["litellm_metadata"]["_code_interpreter_interception_active"] is True
)
assert call_kwargs["litellm_metadata"]["_code_interpreter_interception_active"] is True
# The transient surface marker is NOT forwarded to the follow-up call.
assert "_agentic_loop_api_surface" not in call_kwargs
# Cleanup hook always runs.
@ -390,9 +384,7 @@ async def test_dispatcher_raises_on_repeated_tool_call_fingerprint(restore_callb
# The dispatcher fingerprints the whole value the gate returns as its second
# tuple element, so the seeded fingerprint must mirror that dict exactly.
gate_tool_calls = {
"tool_calls": [{"id": "call_abc", "name": "litellm_code_execution"}]
}
gate_tool_calls = {"tool_calls": [{"id": "call_abc", "name": "litellm_code_execution"}]}
fingerprint = json.dumps(gate_tool_calls, sort_keys=True, default=str)
logger = _GateOnlyLogger(

View file

@ -2945,3 +2945,26 @@ def test_non_bidi_setup_left_untouched_for_followup_capable_providers():
assert streaming._maybe_inject_guardrail_auto_response_disable(msg) == msg
finally:
litellm.callbacks = []
@pytest.mark.asyncio
async def test_log_messages_routes_async_logging_through_bounded_worker():
"""Realtime success logging must go through GLOBAL_LOGGING_WORKER (bounded
queue + per-coroutine timeout), not a bare asyncio.create_task. A bare task
has no timeout/concurrency cap, so when a logging callback is slow every
realtime turn leaves a suspended task pinning its response in memory -> an
unbounded leak. Regression for that fix."""
logging_obj = MagicMock()
streaming = RealTimeStreaming(MagicMock(), MagicMock(), logging_obj)
streaming.messages = [{"type": "session.created"}]
with (
patch("litellm.litellm_core_utils.realtime_streaming.GLOBAL_LOGGING_WORKER") as mock_worker,
patch("litellm.litellm_core_utils.realtime_streaming.asyncio.create_task") as mock_create_task,
patch("litellm.litellm_core_utils.realtime_streaming.executor.submit"),
):
await streaming.log_messages()
mock_worker.ensure_initialized_and_enqueue.assert_called_once()
# the bare create_task path must no longer be used for success logging
mock_create_task.assert_not_called()

View file

@ -97,6 +97,50 @@ def test_token_counter_normal_plus_function_calling():
# test_token_counter_normal_plus_function_calling()
def test_token_counter_legacy_function_call_counts_arguments():
"""
Regression for VERIA-492 (Token-counter function_call bypass).
The legacy OpenAI assistant `function_call` field carries arbitrary text in
`arguments`. Before the fix, `_count_messages` had no branch for
`function_call` and fell through to the unsupported-key `continue`, so an
assistant turn could smuggle unlimited text past `token_counter` and the
proxy `/utils/token_counter` endpoint (and downstream pre-call budget /
`get_modified_max_tokens` math). After the fix it must be counted the
same as the equivalent `tool_calls` payload.
"""
long_arg = "A" * 4000
fc_messages = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": None,
"function_call": {"name": "search", "arguments": long_arg},
},
]
tc_messages = [
{"role": "user", "content": "hi"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "search", "arguments": long_arg},
}
],
},
]
fc_tokens = token_counter(model="gpt-3.5-turbo", messages=fc_messages)
tc_tokens = token_counter(model="gpt-3.5-turbo", messages=tc_messages)
assert fc_tokens == tc_tokens, (
f"function_call arguments must count like tool_calls arguments; "
f"got function_call={fc_tokens}, tool_calls={tc_tokens}"
)
assert fc_tokens > 500, f"4000-char arguments payload must contribute real tokens, got {fc_tokens}"
@pytest.mark.parametrize(
"message_count_pair",
MESSAGES_TEXT,

View file

@ -715,3 +715,109 @@ async def test_async_wrapper_sets_presanitized_and_sanitizes_once():
assert spy.call_count == 1
assert captured["presanitized"] is True
assert [b["type"] for b in captured["messages"][0]["content"]] == ["tool_use"]
def _gate_stubs(monkeypatch):
"""Patch the gate's downstream dispatch targets so config selection can be
observed without making a network call.
Returns ``(captured, translation_calls)`` where ``captured["config"]`` is the
provider config handed to the native passthrough path and ``translation_calls``
counts hits on the Anthropic->OpenAI translation handlers.
"""
from litellm.llms.anthropic.experimental_pass_through.messages import handler
captured = {}
translation_calls = {"count": 0}
def fake_native(**kwargs):
captured["config"] = kwargs.get("anthropic_messages_provider_config")
return "native-passthrough"
def fake_translation(**kwargs):
translation_calls["count"] += 1
return "translated"
monkeypatch.setattr(handler.base_llm_http_handler, "anthropic_messages_handler", fake_native)
monkeypatch.setattr(
handler.LiteLLMMessagesToResponsesAPIHandler,
"anthropic_messages_handler",
staticmethod(fake_translation),
)
monkeypatch.setattr(
handler.LiteLLMMessagesToCompletionTransformationHandler,
"anthropic_messages_handler",
staticmethod(fake_translation),
)
return captured, translation_calls
def test_gate_passthrough_when_supported_endpoints_opts_in(monkeypatch):
"""provider=openai + model_info.supported_endpoints containing /v1/messages
must route to the native passthrough config, NOT the translation handlers."""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
)
from litellm.llms.openai_like.messages.transformation import (
OpenAILikeAnthropicMessagesConfig,
)
captured, translation_calls = _gate_stubs(monkeypatch)
result = anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "Hello"}],
model="openai/some-model",
api_key="sk-test",
api_base="https://host/v1",
model_info={"supported_endpoints": ["/v1/chat/completions", "/v1/messages"]},
)
assert result == "native-passthrough"
assert isinstance(captured["config"], OpenAILikeAnthropicMessagesConfig)
assert translation_calls["count"] == 0
def test_gate_translates_when_supported_endpoints_absent(monkeypatch):
"""Default behavior is unchanged: without the /v1/messages opt-in, an openai
deployment is translated (Responses API), never passed through natively."""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
)
captured, translation_calls = _gate_stubs(monkeypatch)
result = anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "Hello"}],
model="openai/some-model",
api_key="sk-test",
api_base="https://host/v1",
)
assert result == "translated"
assert translation_calls["count"] == 1
assert "config" not in captured
def test_gate_passthrough_skipped_when_only_chat_completions_supported(monkeypatch):
"""A deployment that lists only /v1/chat/completions is still translated;
the opt-in is specifically the /v1/messages entry."""
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
anthropic_messages_handler,
)
captured, translation_calls = _gate_stubs(monkeypatch)
result = anthropic_messages_handler(
max_tokens=100,
messages=[{"role": "user", "content": "Hello"}],
model="openai/some-model",
api_key="sk-test",
api_base="https://host/v1",
model_info={"supported_endpoints": ["/v1/chat/completions"]},
)
assert result == "translated"
assert translation_calls["count"] == 1
assert "config" not in captured

View file

@ -1221,6 +1221,73 @@ def test_async_compact_handler_sends_json_when_not_signed():
assert "data" not in kwargs
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks():
"""
Regression: async_anthropic_messages_handler must inject api_key into the
kwargs dict forwarded to _call_agentic_completion_hooks.
Without this, follow-up calls made by agentic hooks (e.g. websearch
interception's second LLM call after executing searches) have no api_key
and fail with "x-api-key header is required".
"""
handler = BaseLLMHTTPHandler()
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
return_value=({"x-api-key": "sk-test"}, "https://api.anthropic.com")
)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude-haiku", "messages": [], "max_tokens": 16}
)
mock_config.sign_request = Mock(return_value=({}, None))
fake_raw_response = {"id": "msg_1", "type": "message", "role": "assistant", "content": [], "stop_reason": "end_turn"}
mock_config.transform_anthropic_messages_response = Mock(return_value=fake_raw_response)
mock_logging_obj = Mock()
mock_logging_obj.update_environment_variables = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.stream = False
mock_logging_obj.dynamic_success_callbacks = None
captured_kwargs: dict = {}
sentinel_response = object()
async def fake_agentic_hooks(**call_kwargs):
captured_kwargs.update(call_kwargs)
return sentinel_response
mock_httpx_response = Mock()
mock_httpx_response.status_code = 200
with (
patch.object(handler, "_async_post_anthropic_messages_with_http_error_retry", new=AsyncMock(return_value=mock_httpx_response)),
patch.object(handler, "_call_agentic_completion_hooks", side_effect=fake_agentic_hooks),
patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client"),
patch("litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", return_value=None),
):
result = await handler.async_anthropic_messages_handler(
model="claude-haiku",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=mock_config,
anthropic_messages_optional_request_params={"stream": False},
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(api_key="sk-real-anthropic-key"),
logging_obj=mock_logging_obj,
api_key="sk-real-anthropic-key",
stream=False,
)
assert result is sentinel_response
assert "kwargs" in captured_kwargs, "_call_agentic_completion_hooks not called"
forwarded = captured_kwargs["kwargs"]
assert forwarded.get("api_key") == "sk-real-anthropic-key", (
"api_key must be injected into kwargs passed to _call_agentic_completion_hooks "
"so follow-up calls in agentic hooks (e.g. websearch) can authenticate"
)
class _FakeWSExceptions:
class WebSocketException(Exception):
pass

View file

@ -0,0 +1,301 @@
import pytest
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.openai_like.messages.transformation import (
OpenAILikeAnthropicMessagesConfig,
)
from litellm.types.router import GenericLiteLLMParams
@pytest.fixture
def config() -> OpenAILikeAnthropicMessagesConfig:
return OpenAILikeAnthropicMessagesConfig()
@pytest.mark.parametrize(
"api_base, expected",
[
("https://host/v1", "https://host/v1/messages"),
("https://host/v1/", "https://host/v1/messages"),
("https://host", "https://host/v1/messages"),
("https://host/v1/messages", "https://host/v1/messages"),
("https://api.deepseek.com/anthropic", "https://api.deepseek.com/anthropic/v1/messages"),
("https://api.deepseek.com/anthropic/v1", "https://api.deepseek.com/anthropic/v1/messages"),
],
)
def test_get_complete_url_handles_api_base_variants(config, api_base, expected):
url = config.get_complete_url(
api_base=api_base,
api_key="sk-test",
model="some-model",
optional_params={},
litellm_params={},
)
assert url == expected
def test_get_complete_url_requires_api_base(config):
with pytest.raises(ValueError, match="api_base is required"):
config.get_complete_url(
api_base=None,
api_key="sk-test",
model="some-model",
optional_params={},
litellm_params={},
)
def test_request_stays_in_anthropic_shape(config):
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": "Summarize this",
"cache_control": {"type": "ephemeral"},
}
],
}
]
optional_params = {
"max_tokens": 256,
"system": "You are a careful assistant",
"thinking": {"type": "enabled", "budget_tokens": 1024},
"temperature": 0.3,
"tools": [{"name": "lookup", "input_schema": {"type": "object"}}],
"stream": False,
}
payload = config.transform_anthropic_messages_request(
model="some-model",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert payload["model"] == "some-model"
assert payload["messages"] == messages
assert payload["messages"][0]["content"][0]["cache_control"] == {"type": "ephemeral"}
assert payload["system"] == "You are a careful assistant"
assert payload["thinking"] == {"type": "enabled", "budget_tokens": 1024}
assert payload["max_tokens"] == 256
assert payload["tools"] == optional_params["tools"]
openai_only_keys = {
"max_completion_tokens",
"stop",
"n",
"logprobs",
"response_format",
"frequency_penalty",
}
assert openai_only_keys.isdisjoint(payload.keys())
def test_request_requires_max_tokens(config):
with pytest.raises(AnthropicError, match="max_tokens is required"):
config.transform_anthropic_messages_request(
model="some-model",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_optional_request_params={"system": "s"},
litellm_params=GenericLiteLLMParams(),
headers={},
)
def test_validate_environment_sets_bearer_and_anthropic_defaults(config):
headers, api_base = config.validate_anthropic_messages_environment(
headers={},
model="some-model",
messages=[],
optional_params={},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
assert headers["authorization"] == "Bearer sk-test"
assert headers["anthropic-version"] == "2023-06-01"
assert headers["content-type"] == "application/json"
assert api_base == "https://host/v1"
def test_validate_environment_does_not_overwrite_caller_headers(config):
headers, _ = config.validate_anthropic_messages_environment(
headers={
"authorization": "Bearer caller-token",
"anthropic-version": "2024-10-22",
"content-type": "application/json",
},
model="some-model",
messages=[],
optional_params={},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
assert headers["authorization"] == "Bearer caller-token"
assert headers["anthropic-version"] == "2024-10-22"
def test_validate_environment_preserves_standard_cased_caller_headers(config):
headers, _ = config.validate_anthropic_messages_environment(
headers={
"Authorization": "Bearer caller-token",
"Anthropic-Version": "2024-10-22",
"Content-Type": "application/json",
},
model="some-model",
messages=[],
optional_params={},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
lowercased = {key.lower() for key in headers}
assert len(lowercased) == len(headers)
assert headers["Authorization"] == "Bearer caller-token"
assert headers["Anthropic-Version"] == "2024-10-22"
assert headers["Content-Type"] == "application/json"
def test_validate_environment_honors_x_api_key_when_present(config):
headers, _ = config.validate_anthropic_messages_environment(
headers={"X-Api-Key": "caller-key"},
model="some-model",
messages=[],
optional_params={},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
assert "authorization" not in {key.lower() for key in headers}
assert headers["X-Api-Key"] == "caller-key"
def test_validate_environment_injects_anthropic_beta_for_context_management(config):
headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="some-model",
messages=[],
optional_params={
"context_management": {"edits": [{"type": "clear_tool_uses_20250919"}]},
},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
assert "context-management-2025-06-27" in headers["anthropic-beta"].split(",")
def test_validate_environment_injects_anthropic_beta_for_fast_mode(config):
headers, _ = config.validate_anthropic_messages_environment(
headers={},
model="some-model",
messages=[],
optional_params={"speed": "fast"},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
assert "fast-mode-2026-02-01" in headers["anthropic-beta"].split(",")
def test_validate_environment_merges_existing_anthropic_beta(config):
headers, _ = config.validate_anthropic_messages_environment(
headers={"anthropic-beta": "caller-flag"},
model="some-model",
messages=[],
optional_params={"speed": "fast"},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
beta_values = set(headers["anthropic-beta"].split(","))
assert "caller-flag" in beta_values
assert "fast-mode-2026-02-01" in beta_values
def test_request_strips_advisor_blocks_when_advisor_tool_absent(config):
messages = [
{"role": "user", "content": "hello"},
{
"role": "assistant",
"content": [
{"type": "text", "text": "thinking out loud"},
{"type": "server_tool_use", "id": "advisor_1", "name": "advisor", "input": {}},
{"type": "advisor_tool_result", "tool_use_id": "advisor_1", "content": "stale"},
],
},
]
payload = config.transform_anthropic_messages_request(
model="some-model",
messages=messages,
anthropic_messages_optional_request_params={"max_tokens": 64},
litellm_params=GenericLiteLLMParams(),
headers={},
)
flattened_types = [
block.get("type")
for message in payload["messages"]
if isinstance(message.get("content"), list)
for block in message["content"]
if isinstance(block, dict)
]
assert "advisor_tool_result" not in flattened_types
assert "server_tool_use" not in flattened_types
def test_request_maps_reasoning_effort_to_thinking(config):
payload = config.transform_anthropic_messages_request(
model="claude-sonnet-4-20250514",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_optional_request_params={
"max_tokens": 1024,
"reasoning_effort": "medium",
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "reasoning_effort" not in payload
assert isinstance(payload.get("thinking"), dict)
assert payload["thinking"].get("type") == "enabled"
def test_passthrough_disables_anthropic_beta_filtering(config):
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
assert config.should_filter_anthropic_beta_headers() is False
assert AnthropicMessagesConfig().should_filter_anthropic_beta_headers() is True
def test_anthropic_beta_survives_provider_filter_on_passthrough_path(config):
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
headers, _ = config.validate_anthropic_messages_environment(
headers={"Anthropic-Beta": "caller-flag"},
model="some-model",
messages=[],
optional_params={"speed": "fast"},
litellm_params={},
api_key="sk-test",
api_base="https://host/v1",
)
# The deployment routes as provider "openai", which has no beta mapping, so an
# unconditional filter would drop every anthropic-beta value. The handler must
# skip filtering for this config so the native upstream still receives them.
if config.should_filter_anthropic_beta_headers():
headers = update_headers_with_filtered_beta(headers=dict(headers), provider="openai")
survived = set(headers.get("anthropic-beta", "").split(","))
assert {"caller-flag", "fast-mode-2026-02-01"} <= survived
stripped = update_headers_with_filtered_beta(headers=dict(headers), provider="openai")
assert "anthropic-beta" not in stripped

View file

@ -265,3 +265,131 @@ async def test_fetch_tools_from_gateway_managed_swallows_errors():
)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
def _http_server(server_id: str, name: str, **kwargs) -> MCPServer:
return MCPServer(
server_id=server_id,
name=name,
url=f"https://{name}/mcp",
transport=MCPTransport.http,
**kwargs,
)
@pytest.mark.asyncio
async def test_aggregate_list_tools_absorbs_one_unauthenticated_server():
"""Regression: across the aggregate (/mcp), a delegate/passthrough server that raises
MCPUpstreamAuthError must not empty every other server's tools. Re-raising it on the
aggregate path (introduced with the passthrough feature) zeroed the whole list because the
fan-out gather propagated it."""
from unittest.mock import patch
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._types import UserAPIKeyAuth
delegate = _http_server(
"s1", "delegate_docs", auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True
)
working = _http_server("s2", "working_docs", auth_type=MCPAuth.none)
good_tool = MCPTool(name="working_docs-read", description="d", inputSchema={"type": "object"})
async def fake_get_tools(server, **kwargs):
if server.server_id == delegate.server_id:
raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name)
return [good_tool]
with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate, working])), patch.object(
mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={})
), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object(
mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None)
), patch.object(
mcp_server, "filter_tools_by_key_team_permissions", AsyncMock(side_effect=lambda tools, **k: tools)
), patch.object(
mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools)
):
tools = await mcp_server._get_tools_from_mcp_servers(
user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"),
mcp_auth_header=None,
mcp_servers=None,
)
assert [t.name for t in tools] == ["working_docs-read"]
@pytest.mark.asyncio
async def test_single_server_route_also_absorbs_upstream_auth_error():
"""A single-server route (/<server>/mcp) absorbs an upstream-auth error just like the aggregate:
the failing server is omitted (empty list) rather than re-raised. Surfacing it to the client as a
401 + WWW-Authenticate challenge cannot be done from this list handler — the MCP session manager
serializes a raise into a JSON-RPC error, not an HTTP 401 — so re-auth surfacing is handled by a
request-scope preemptive check, tracked separately."""
from unittest.mock import patch
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_gateway_server_name
from litellm.proxy._types import UserAPIKeyAuth
delegate = _http_server(
"s1", "delegate_docs", auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True
)
async def fake_get_tools(server, **kwargs):
raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name)
# /<server>/mcp sets the path-derived single-server scope; absorption must hold even then.
token = _mcp_gateway_server_name.set("delegate_docs")
try:
with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object(
mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={})
), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object(
mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None)
), patch.object(
mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools)
):
tools = await mcp_server._get_tools_from_mcp_servers(
user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"),
mcp_auth_header=None,
mcp_servers=["delegate_docs"],
)
assert tools == []
finally:
_mcp_gateway_server_name.reset(token)
@pytest.mark.asyncio
async def test_aggregate_with_single_accessible_server_still_absorbs():
"""Regression for the route-misclassification: an aggregate request (/mcp, mcp_servers=None)
from a key that can access exactly one server must still absorb that server's
MCPUpstreamAuthError, not surface it. Keying the surface decision off the allowed count rather
than the request filter would re-raise here and leave the aggregate broken for one-server
permission sets."""
from unittest.mock import patch
from litellm.proxy._experimental.mcp_server import server as mcp_server
from litellm.proxy._types import UserAPIKeyAuth
delegate = _http_server(
"s1", "delegate_docs", auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True
)
async def fake_get_tools(server, **kwargs):
raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name)
with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object(
mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={})
), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object(
mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None)
), patch.object(
mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools)
):
# Aggregate route: no explicit server filter, even though only one server is accessible.
tools = await mcp_server._get_tools_from_mcp_servers(
user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"),
mcp_auth_header=None,
mcp_servers=None,
)
assert tools == []

View file

@ -6447,3 +6447,54 @@ class TestStreamableHttpAuthErrorMapping:
m.get("type") == "http.response.start" and m.get("status") == 500
for m in sent
)
class TestMCPMetaTraceCarrier:
"""`_mcp_meta_trace_carrier` extracts the W3C trace context the MCP client
propagated in the request's params._meta (SEP-414) so the otel_v2 MCP span can
parent to the client's span. Exercises the real MCP SDK `RequestParams.Meta`
shape (extra='allow' preserves the unprefixed keys), not just an injected
carrier."""
def test_extracts_trace_context_and_excludes_baggage_and_other_meta(self):
"""Only traceparent/tracestate are carried. The client's W3C ``baggage`` is
deliberately dropped even though it rides in params._meta: it is
caller-controlled, and the otel baggage processor stamps allowlisted baggage
keys onto the span, so honoring it would let a client spoof a span's identity
(e.g. ``litellm.team.id``). Dropping it at the source is the regression guard."""
from types import SimpleNamespace
from mcp.types import RequestParams
from litellm.proxy._experimental.mcp_server.server import (
_mcp_meta_trace_carrier,
)
meta = RequestParams.Meta.model_validate(
{
"traceparent": "00-11111111111111111111111111111111-2222222222222222-01",
"tracestate": "rojo=1",
"baggage": "litellm.team.id=spoofed-team,litellm.metadata.user_api_key_user_id=attacker",
"progressToken": "p1",
}
)
carrier = _mcp_meta_trace_carrier(SimpleNamespace(meta=meta))
assert carrier == {
"traceparent": "00-11111111111111111111111111111111-2222222222222222-01",
"tracestate": "rojo=1",
}
assert "baggage" not in carrier
def test_none_when_no_trace_context(self):
from types import SimpleNamespace
from mcp.types import RequestParams
from litellm.proxy._experimental.mcp_server.server import (
_mcp_meta_trace_carrier,
)
assert _mcp_meta_trace_carrier(None) is None
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=None)) is None
only_progress = RequestParams.Meta.model_validate({"progressToken": "p1"})
assert _mcp_meta_trace_carrier(SimpleNamespace(meta=only_progress)) is None

View file

@ -0,0 +1,837 @@
"""
Tests for MCP tool search feature.
Covers:
- search_tools() pure function
- get_virtual_tool_definitions() shape
- list_tool_rest_api returns only virtual tools when mcp_tool_search_enabled=True
- call_tool_rest_api intercepts mcp_tool_search calls
- call_tool_rest_api intercepts mcp_tool_call calls
"""
import json
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy._experimental.mcp_server.tool_search import (
MCP_TOOL_CALL_TOOL_NAME,
MCP_TOOL_SEARCH_TOOL_NAME,
coerce_top_k,
get_virtual_tool_definitions,
search_tools,
)
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
def _make_tools(specs: list[tuple[str, str]]) -> list[dict[str, Any]]:
return [
{
"name": name,
"description": desc,
"inputSchema": {"type": "object", "properties": {}},
}
for name, desc in specs
]
def _make_perm(**kwargs: Any) -> LiteLLM_ObjectPermissionTable:
return LiteLLM_ObjectPermissionTable(object_permission_id="test", **kwargs)
SAMPLE_TOOLS = _make_tools(
[
("github-create_issue", "Create a new issue in a GitHub repository"),
("github-list_repos", "List all repositories for a GitHub user"),
("slack-send_message", "Send a message to a Slack channel"),
("slack-list_channels", "List all Slack channels in a workspace"),
("notion-create_page", "Create a new page in Notion"),
]
)
class TestCoerceTopK:
def test_int_passthrough(self) -> None:
assert coerce_top_k(3) == 3
def test_numeric_string_coerced(self) -> None:
assert coerce_top_k("7") == 7
def test_float_truncated(self) -> None:
assert coerce_top_k(3.9) == 3
def test_non_numeric_string_returns_default(self) -> None:
assert coerce_top_k("abc") == 5
def test_none_returns_default(self) -> None:
assert coerce_top_k(None) == 5
def test_custom_default(self) -> None:
assert coerce_top_k("nope", default=10) == 10
class TestSearchTools:
def test_returns_matching_tools(self) -> None:
results = search_tools("github issue", SAMPLE_TOOLS)
names = [t["name"] for t in results]
assert "github-create_issue" in names
def test_ranks_by_relevance(self) -> None:
results = search_tools("github", SAMPLE_TOOLS)
names = [t["name"] for t in results]
github_positions = [i for i, n in enumerate(names) if n.startswith("github")]
other_positions = [i for i, n in enumerate(names) if not n.startswith("github")]
assert all(g < o for g in github_positions for o in other_positions)
def test_top_k_limits_results(self) -> None:
results = search_tools("a", SAMPLE_TOOLS, top_k=2)
assert len(results) <= 2
def test_empty_query_returns_empty(self) -> None:
assert search_tools("", SAMPLE_TOOLS) == []
def test_no_match_returns_empty(self) -> None:
assert search_tools("xyzzy_nonexistent_zzz", SAMPLE_TOOLS) == []
def test_matches_description_not_just_name(self) -> None:
results = search_tools("channel", SAMPLE_TOOLS)
names = [t["name"] for t in results]
assert "slack-list_channels" in names
def test_case_insensitive(self) -> None:
lower = [t["name"] for t in search_tools("github", SAMPLE_TOOLS)]
upper = [t["name"] for t in search_tools("GITHUB", SAMPLE_TOOLS)]
assert lower == upper
def test_result_tools_have_full_schema(self) -> None:
for tool in search_tools("github", SAMPLE_TOOLS):
assert "name" in tool
assert "description" in tool
assert "inputSchema" in tool
class TestGetVirtualToolDefinitions:
def test_returns_two_tools(self) -> None:
assert len(get_virtual_tool_definitions()) == 2
def test_has_mcp_tool_search(self) -> None:
names = [t["name"] for t in get_virtual_tool_definitions()]
assert MCP_TOOL_SEARCH_TOOL_NAME in names
def test_has_mcp_tool_call(self) -> None:
names = [t["name"] for t in get_virtual_tool_definitions()]
assert MCP_TOOL_CALL_TOOL_NAME in names
def test_mcp_tool_search_schema_has_query(self) -> None:
tools = get_virtual_tool_definitions()
search_tool = next(t for t in tools if t["name"] == MCP_TOOL_SEARCH_TOOL_NAME)
props = search_tool["inputSchema"]["properties"]
assert "query" in props
assert search_tool["inputSchema"]["required"] == ["query"]
def test_mcp_tool_call_schema_has_tool_name_and_arguments(self) -> None:
tools = get_virtual_tool_definitions()
call_tool = next(t for t in tools if t["name"] == MCP_TOOL_CALL_TOOL_NAME)
props = call_tool["inputSchema"]["properties"]
assert "tool_name" in props
assert "arguments" in props
assert "tool_name" in call_tool["inputSchema"]["required"]
def test_all_tools_have_description(self) -> None:
for tool in get_virtual_tool_definitions():
assert tool.get("description"), f"{tool['name']} missing description"
def test_definitions_construct_mcp_protocol_tool(self) -> None:
"""The MCP protocol list_tools handler builds mcp.types.Tool(**d) from
each definition, so the dict keys must stay valid Tool fields."""
from mcp.types import Tool
built = [Tool(**d) for d in get_virtual_tool_definitions()]
assert {t.name for t in built} == {
MCP_TOOL_SEARCH_TOOL_NAME,
MCP_TOOL_CALL_TOOL_NAME,
}
class TestListToolRestApiWithToolSearch:
@pytest.mark.asyncio
async def test_returns_only_virtual_tools_when_flag_enabled(self) -> None:
from litellm.proxy._experimental.mcp_server.rest_endpoints import router
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(
mcp_tool_search_enabled=True,
mcp_servers=["github", "slack"],
),
)
mock_request = MagicMock()
mock_request.headers = {}
list_fn = next(
r.endpoint
for r in router.routes
if hasattr(r, "path") and r.path.endswith("/tools/list") and hasattr(r, "methods") and "GET" in r.methods
)
result = await list_fn(
request=mock_request,
server_id=None,
include_disabled_tools=False,
user_api_key_dict=user_api_key_dict,
)
assert result["error"] is None
tool_names = [t["name"] for t in result["tools"]]
assert set(tool_names) == {MCP_TOOL_SEARCH_TOOL_NAME, MCP_TOOL_CALL_TOOL_NAME}
@pytest.mark.asyncio
async def test_returns_full_catalog_when_flag_disabled(self) -> None:
from litellm.proxy._experimental.mcp_server.rest_endpoints import router
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(
mcp_tool_search_enabled=False,
mcp_servers=["github"],
),
)
mock_request = MagicMock()
mock_request.headers = {}
fake_tools = [
{
"name": "github-create_issue",
"description": "Create issue",
"inputSchema": {"type": "object"},
}
]
list_fn = next(
r.endpoint
for r in router.routes
if hasattr(r, "path") and r.path.endswith("/tools/list") and hasattr(r, "methods") and "GET" in r.methods
)
with (
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.build_effective_auth_contexts",
new_callable=AsyncMock,
return_value=[user_api_key_dict],
),
patch("litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager") as mock_manager,
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server",
new_callable=AsyncMock,
return_value=fake_tools,
),
patch("litellm.proxy._experimental.mcp_server.rest_endpoints.IPAddressUtils"),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._prefetch_user_oauth_creds",
new_callable=AsyncMock,
return_value={},
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids",
return_value=[],
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_server_auth_header",
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_user_oauth_extra_headers",
new_callable=AsyncMock,
return_value=None,
),
):
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["github"], 0))
mock_manager.get_mcp_server_by_id = MagicMock(return_value=MagicMock(name="github", server_id="github"))
result = await list_fn(
request=mock_request,
server_id=None,
include_disabled_tools=False,
user_api_key_dict=user_api_key_dict,
)
tool_names = [t["name"] for t in result["tools"]]
assert MCP_TOOL_SEARCH_TOOL_NAME not in tool_names
assert "github-create_issue" in tool_names
@pytest.mark.asyncio
async def test_admin_include_disabled_tools_bypasses_virtual_catalog(self) -> None:
"""Regression: an admin listing with include_disabled_tools must see the
real catalog (to configure allowlists) even when mcp_tool_search_enabled is
set, instead of the two virtual tools."""
from litellm.proxy._experimental.mcp_server.rest_endpoints import router
user_api_key_dict = UserAPIKeyAuth(
api_key="admin_key",
user_role=LitellmUserRoles.PROXY_ADMIN,
object_permission=_make_perm(
mcp_tool_search_enabled=True,
mcp_servers=["github"],
),
)
mock_request = MagicMock()
mock_request.headers = {}
fake_tools = [
{
"name": "github-create_issue",
"description": "Create issue",
"inputSchema": {"type": "object"},
}
]
list_fn = next(
r.endpoint
for r in router.routes
if hasattr(r, "path") and r.path.endswith("/tools/list") and hasattr(r, "methods") and "GET" in r.methods
)
with (
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.build_effective_auth_contexts",
new_callable=AsyncMock,
return_value=[user_api_key_dict],
),
patch("litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager") as mock_manager,
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_tools_for_single_server",
new_callable=AsyncMock,
return_value=fake_tools,
),
patch("litellm.proxy._experimental.mcp_server.rest_endpoints.IPAddressUtils"),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._prefetch_user_oauth_creds",
new_callable=AsyncMock,
return_value={},
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_oauth2_server_ids",
return_value=[],
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_server_auth_header",
return_value=None,
),
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints._get_user_oauth_extra_headers",
new_callable=AsyncMock,
return_value=None,
),
):
mock_manager.get_allowed_mcp_servers = AsyncMock(return_value=["github"])
mock_manager.filter_server_ids_by_ip_with_info = MagicMock(return_value=(["github"], 0))
mock_manager.get_mcp_server_by_id = MagicMock(return_value=MagicMock(name="github", server_id="github"))
result = await list_fn(
request=mock_request,
server_id=None,
include_disabled_tools=True,
user_api_key_dict=user_api_key_dict,
)
tool_names = [t["name"] for t in result["tools"]]
assert MCP_TOOL_SEARCH_TOOL_NAME not in tool_names
assert "github-create_issue" in tool_names
class TestCallToolRestApiVirtualTools:
def _make_request(self, body: dict[str, Any]) -> MagicMock:
mock_request = MagicMock()
mock_request.json = AsyncMock(return_value=body)
mock_request.headers = {}
mock_request.url = MagicMock()
mock_request.url.path = "/mcp-rest/tools/call"
return mock_request
def _get_call_fn(self) -> Any:
from litellm.proxy._experimental.mcp_server.rest_endpoints import router
return next(
r.endpoint
for r in router.routes
if hasattr(r, "path") and r.path.endswith("/tools/call") and hasattr(r, "methods") and "POST" in r.methods
)
@pytest.mark.asyncio
async def test_mcp_tool_search_call_returns_tool_defs(self) -> None:
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(
mcp_tool_search_enabled=True,
mcp_servers=["github"],
),
)
request = self._make_request({"name": MCP_TOOL_SEARCH_TOOL_NAME, "arguments": {"query": "create issue"}})
mock_tool = MagicMock()
mock_tool.name = "github-create_issue"
mock_tool.description = "Create a GitHub issue"
mock_tool.inputSchema = {"type": "object", "properties": {}}
with patch(
"litellm.proxy._experimental.mcp_server.server._list_mcp_tools",
new_callable=AsyncMock,
return_value=[mock_tool],
):
result = await self._get_call_fn()(
request=request,
user_api_key_dict=user_api_key_dict,
)
assert result.content
assert result.content[0].type == "text"
returned_tools = json.loads(result.content[0].text)
assert isinstance(returned_tools, list)
assert any(t["name"] == "github-create_issue" for t in returned_tools)
@pytest.mark.asyncio
async def test_mcp_tool_call_executes_discovered_tool(self) -> None:
from mcp.types import CallToolResult, TextContent
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(
mcp_tool_search_enabled=True,
mcp_servers=["github"],
),
)
request = self._make_request(
{
"name": MCP_TOOL_CALL_TOOL_NAME,
"arguments": {
"tool_name": "github-create_issue",
"arguments": {"title": "bug", "repo": "myrepo"},
},
}
)
fake_result = CallToolResult(
content=[TextContent(type="text", text="Issue created")],
isError=False,
)
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
),
patch(
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
new_callable=AsyncMock,
return_value=fake_result,
) as mock_execute,
):
result = await self._get_call_fn()(
request=request,
user_api_key_dict=user_api_key_dict,
)
mock_execute.assert_awaited_once()
assert mock_execute.await_args.kwargs["name"] == "github-create_issue"
assert result.isError is False
assert result.content[0].text == "Issue created"
@pytest.mark.asyncio
async def test_mcp_tool_call_forwards_client_ip_for_ip_filtering(self) -> None:
"""Regression: the virtual call path must resolve allowed servers with the
request's client IP so IP-restricted servers (available_on_public_internet:
false) cannot be reached from a public IP."""
from mcp.types import CallToolResult, TextContent
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(mcp_tool_search_enabled=True, mcp_servers=["github"]),
)
request = self._make_request(
{
"name": MCP_TOOL_CALL_TOOL_NAME,
"arguments": {"tool_name": "github-create_issue", "arguments": {}},
}
)
fake_result = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
with (
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="203.0.113.7",
),
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
new_callable=AsyncMock,
return_value=fake_result,
),
):
await self._get_call_fn()(request=request, user_api_key_dict=user_api_key_dict)
mock_allowed.assert_awaited_once()
assert mock_allowed.await_args.kwargs["client_ip"] == "203.0.113.7"
@pytest.mark.asyncio
async def test_mcp_tool_search_forwards_client_ip_for_ip_filtering(self) -> None:
"""Search must list tools through the IP-filtered catalog, not the raw one."""
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(mcp_tool_search_enabled=True, mcp_servers=["github"]),
)
request = self._make_request({"name": MCP_TOOL_SEARCH_TOOL_NAME, "arguments": {"query": "issue"}})
with (
patch(
"litellm.proxy._experimental.mcp_server.rest_endpoints.IPAddressUtils.get_mcp_client_ip",
return_value="203.0.113.7",
),
patch(
"litellm.proxy._experimental.mcp_server.server._list_mcp_tools",
new_callable=AsyncMock,
return_value=[],
) as mock_list,
):
await self._get_call_fn()(request=request, user_api_key_dict=user_api_key_dict)
mock_list.assert_awaited_once()
assert mock_list.await_args.kwargs["client_ip"] == "203.0.113.7"
@pytest.mark.asyncio
async def test_mcp_tool_search_requires_flag_enabled(self) -> None:
from fastapi import HTTPException
user_api_key_dict = UserAPIKeyAuth(
api_key="test_key",
object_permission=_make_perm(mcp_tool_search_enabled=False),
)
request = self._make_request({"name": MCP_TOOL_SEARCH_TOOL_NAME, "arguments": {"query": "create issue"}})
with pytest.raises(HTTPException) as exc_info:
await self._get_call_fn()(
request=request,
user_api_key_dict=user_api_key_dict,
)
assert exc_info.value.status_code in (400, 403, 404)
class TestDispatchVirtualMcpTool:
"""Covers the SSE/protocol-path interception helper in server.py."""
@pytest.mark.asyncio
async def test_returns_none_for_non_virtual_tool(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_dispatch_virtual_mcp_tool,
)
result = await _dispatch_virtual_mcp_tool(
name="github-create_issue",
arguments={},
user_api_key_auth=UserAPIKeyAuth(api_key="k"),
client_ip=None,
)
assert result is None
@pytest.mark.asyncio
async def test_rejects_when_flag_disabled(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_dispatch_virtual_mcp_tool,
)
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=False))
result = await _dispatch_virtual_mcp_tool(
name=MCP_TOOL_SEARCH_TOOL_NAME,
arguments={"query": "x"},
user_api_key_auth=uak,
client_ip=None,
)
assert result is not None
assert result.isError is True
@pytest.mark.asyncio
async def test_routes_search_with_client_ip(self) -> None:
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_search",
new_callable=AsyncMock,
return_value="SEARCH_RESULT",
) as mock_search:
result = await srv._dispatch_virtual_mcp_tool(
name=MCP_TOOL_SEARCH_TOOL_NAME,
arguments={"query": "q", "top_k": 3},
user_api_key_auth=uak,
client_ip="203.0.113.9",
)
assert result == "SEARCH_RESULT"
assert mock_search.await_args.kwargs["client_ip"] == "203.0.113.9"
assert mock_search.await_args.kwargs["query"] == "q"
assert mock_search.await_args.kwargs["top_k"] == 3
@pytest.mark.asyncio
async def test_routes_call_with_client_ip(self) -> None:
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_call",
new_callable=AsyncMock,
return_value="CALL_RESULT",
) as mock_call:
result = await srv._dispatch_virtual_mcp_tool(
name=MCP_TOOL_CALL_TOOL_NAME,
arguments={"tool_name": "math-add", "arguments": {"a": 1, "b": 2}},
user_api_key_auth=uak,
client_ip="203.0.113.9",
mcp_auth_header="bearer-xyz",
mcp_server_auth_headers={"github": {"Authorization": "Bearer gh"}},
oauth2_headers={"Authorization": "Bearer oauth"},
raw_headers={"x-mcp-auth": "tok"},
)
assert result == "CALL_RESULT"
kw = mock_call.await_args.kwargs
assert kw["tool_name"] == "math-add"
assert kw["client_ip"] == "203.0.113.9"
assert kw["mcp_auth_header"] == "bearer-xyz"
assert kw["mcp_server_auth_headers"] == {"github": {"Authorization": "Bearer gh"}}
assert kw["oauth2_headers"] == {"Authorization": "Bearer oauth"}
assert kw["raw_headers"] == {"x-mcp-auth": "tok"}
@pytest.mark.asyncio
async def test_call_builds_and_forwards_logging_obj(self) -> None:
"""Regression: the SSE dispatch must run the pre-call pipeline and forward
the resulting logging object to handle_mcp_tool_call, otherwise mcp_tool_call
over /mcp/ skips spend logging and guardrails (unlike the REST path)."""
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
sentinel_logging_obj = object()
with (
patch.object(
srv,
"_build_virtual_call_logging_obj",
new_callable=AsyncMock,
return_value=sentinel_logging_obj,
) as mock_build,
patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_call",
new_callable=AsyncMock,
return_value="CALL_RESULT",
) as mock_call,
):
await srv._dispatch_virtual_mcp_tool(
name=MCP_TOOL_CALL_TOOL_NAME,
arguments={"tool_name": "math-add", "arguments": {"a": 1}},
user_api_key_auth=uak,
client_ip=None,
)
assert mock_build.await_count == 1
assert mock_call.await_args.kwargs["litellm_logging_obj"] is sentinel_logging_obj
@pytest.mark.asyncio
async def test_search_coerces_non_int_top_k(self) -> None:
"""Regression: a non-integer top_k from an MCP client must not raise; it
falls back to the default instead of ValueError propagating out."""
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with patch(
"litellm.proxy._experimental.mcp_server.tool_search.handle_mcp_tool_search",
new_callable=AsyncMock,
return_value="SEARCH_RESULT",
) as mock_search:
await srv._dispatch_virtual_mcp_tool(
name=MCP_TOOL_SEARCH_TOOL_NAME,
arguments={"query": "issue", "top_k": "not-a-number"},
user_api_key_auth=uak,
client_ip=None,
)
assert mock_search.await_args.kwargs["top_k"] == 5
@pytest.mark.asyncio
async def test_call_handler_forwards_auth_headers_to_execute(self) -> None:
"""Regression: per-request auth headers must reach execute_mcp_tool so
upstream MCP servers needing pass-through auth can be called."""
from mcp.types import CallToolResult, TextContent
from litellm.proxy._experimental.mcp_server.tool_search import (
handle_mcp_tool_call,
)
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
fake = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False)
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[MagicMock()],
) as mock_allowed,
patch(
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
new_callable=AsyncMock,
return_value=fake,
) as mock_exec,
):
sentinel_logging_obj = object()
await handle_mcp_tool_call(
tool_name="github-create_issue",
arguments={},
user_api_key_dict=uak,
mcp_servers=["github"],
mcp_auth_header="bearer-xyz",
mcp_server_auth_headers={"github": {"Authorization": "Bearer gh"}},
oauth2_headers={"Authorization": "Bearer oauth"},
raw_headers={"x-mcp-auth": "tok"},
litellm_logging_obj=sentinel_logging_obj,
)
kw = mock_exec.await_args.kwargs
assert kw["mcp_auth_header"] == "bearer-xyz"
assert kw["mcp_server_auth_headers"] == {"github": {"Authorization": "Bearer gh"}}
assert kw["oauth2_headers"] == {"Authorization": "Bearer oauth"}
assert kw["raw_headers"] == {"x-mcp-auth": "tok"}
# Spend logging: the logging object must reach execute_mcp_tool
assert kw["litellm_logging_obj"] is sentinel_logging_obj
# Scoped session: the requested mcp_servers scope must reach server resolution
assert mock_allowed.await_args.kwargs["mcp_servers"] == ["github"]
@pytest.mark.asyncio
async def test_call_rejected_when_no_accessible_servers(self) -> None:
"""Regression: a key with no accessible MCP servers must not reach
execute_mcp_tool, where an unprefixed local tool name would otherwise
run via the local registry without a server permission check."""
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.tool_search import (
handle_mcp_tool_call,
)
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with (
patch(
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new_callable=AsyncMock,
return_value=[],
),
patch(
"litellm.proxy._experimental.mcp_server.server.execute_mcp_tool",
new_callable=AsyncMock,
) as mock_exec,
):
with pytest.raises(HTTPException) as exc_info:
await handle_mcp_tool_call(
tool_name="local_secret_tool",
arguments={},
user_api_key_dict=uak,
)
assert exc_info.value.status_code == 403
mock_exec.assert_not_awaited()
class TestCaptureHostProgressCallback:
"""Covers the host progress-forwarding helper extracted from the tool call path."""
def test_returns_none_when_request_context_unavailable(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
class _NoCtx:
@property
def request_context(self): # type: ignore[no-untyped-def]
raise RuntimeError("no context")
assert _capture_host_progress_callback(_NoCtx()) is None
def test_returns_none_when_no_progress_token(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = None
assert _capture_host_progress_callback(host) is None
def test_returns_callable_when_token_present(self) -> None:
from litellm.proxy._experimental.mcp_server.server import (
_capture_host_progress_callback,
)
host = MagicMock()
host.request_context.meta.progressToken = "tok12345"
host.request_context.session = MagicMock()
assert callable(_capture_host_progress_callback(host))
class TestHandleListToolsVirtual:
"""Covers the protocol list_tools early-return when the flag is enabled."""
@pytest.mark.asyncio
async def test_returns_virtual_tools_when_flag_enabled(self) -> None:
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with patch(
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new_callable=AsyncMock,
return_value=(uak, None, None, None, None, None, None),
):
tools = await srv.handle_list_tools()
assert {t.name for t in tools} == {
MCP_TOOL_SEARCH_TOOL_NAME,
MCP_TOOL_CALL_TOOL_NAME,
}
class TestMcpServerToolCallErrorHandling:
"""The protocol tool-call handler must convert virtual-tool errors to an
isError CallToolResult instead of letting them raise out of the handler."""
@pytest.mark.asyncio
async def test_virtual_tool_error_returns_iserror_not_raised(self) -> None:
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server import server as srv
uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True))
with (
patch(
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new_callable=AsyncMock,
return_value=(uak, None, None, None, None, None, None),
),
patch(
"litellm.proxy._experimental.mcp_server.server._dispatch_virtual_mcp_tool",
new_callable=AsyncMock,
side_effect=HTTPException(status_code=403, detail="User not allowed to call this tool"),
),
):
result = await srv.mcp_server_tool_call(
name=MCP_TOOL_CALL_TOOL_NAME,
arguments={"tool_name": "other-server-tool", "arguments": {}},
)
assert result.isError is True
assert "User not allowed to call this tool" in result.content[0].text

View file

@ -1520,6 +1520,42 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride:
assert "vertex_credentials" not in out
assert "vertex_project" not in out
def test_clears_nvcf_function_id_on_base_override(self):
from litellm.router_utils.clientside_credential_handler import (
get_dynamic_litellm_params,
)
admin_params = {
"model": "nvidia_riva/parakeet",
"api_base": "grpc.nvcf.nvidia.com:443",
"api_key": "nvapi-admin",
"nvcf_function_id": "admin-pinned-function",
}
out = get_dynamic_litellm_params(
litellm_params=dict(admin_params),
request_kwargs={"api_base": "self-hosted.example.com:50051"},
)
assert out["api_base"] == "self-hosted.example.com:50051"
assert "nvcf_function_id" not in out
def test_clears_use_ssl_on_base_override(self):
from litellm.router_utils.clientside_credential_handler import (
get_dynamic_litellm_params,
)
admin_params = {
"model": "nvidia_riva/parakeet",
"api_base": "grpc.nvcf.nvidia.com:443",
"api_key": "nvapi-admin",
"use_ssl": True,
}
out = get_dynamic_litellm_params(
litellm_params=dict(admin_params),
request_kwargs={"api_base": "self-hosted.example.com:50051"},
)
assert out["api_base"] == "self-hosted.example.com:50051"
assert "use_ssl" not in out
def test_caller_resupplied_value_overrides_admin_value_on_base_override(self):
# When the caller redirects ``api_base`` and *also* supplies their
# own value for one of the admin fields (e.g. ``organization``),
@ -1712,6 +1748,127 @@ class TestIsRequestBodySafeBlocksBedrockProjectOverride:
)
class TestIsRequestBodySafeBlocksNVCFFunctionOverride:
"""``nvcf_function_id`` is rejected as a request-body param unless the
admin opted in proxy-wide or per-deployment."""
def test_nvcf_function_id_in_request_body_is_rejected(self):
with pytest.raises(ValueError, match="nvcf_function_id"):
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"nvcf_function_id": "caller-supplied",
},
general_settings={},
llm_router=None,
model="nvidia_riva/parakeet",
)
def test_nvcf_function_id_with_api_key_still_rejected(self):
with pytest.raises(ValueError, match="nvcf_function_id"):
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"api_key": "sk-anything",
"nvcf_function_id": "caller-supplied",
},
general_settings={},
llm_router=None,
model="nvidia_riva/parakeet",
)
def test_admin_opt_in_proxy_wide_allows_nvcf_function_id(self):
assert (
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"nvcf_function_id": "byok-function-id",
},
general_settings={"allow_client_side_credentials": True},
llm_router=None,
model="nvidia_riva/parakeet",
)
is True
)
def test_admin_opt_in_per_deployment_allows_nvcf_function_id(self, monkeypatch):
"""The error message lists per-deployment ``configurable_clientside_auth_params``
as a second opt-in. Cover that path too so it can't silently regress."""
from litellm.proxy.auth import auth_utils
monkeypatch.setattr(
auth_utils,
"_allow_model_level_clientside_configurable_parameters",
lambda model, param, request_body_value, llm_router: param == "nvcf_function_id",
)
assert (
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"nvcf_function_id": "byok-function-id",
},
general_settings={},
llm_router=None,
model="nvidia_riva/parakeet",
)
is True
)
class TestIsRequestBodySafeBlocksRivaUseSsl:
"""``use_ssl`` is rejected as a request-body param unless the admin
opted in proxy-wide or per-deployment."""
def test_use_ssl_in_request_body_is_rejected(self):
with pytest.raises(ValueError, match="use_ssl"):
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"use_ssl": False,
},
general_settings={},
llm_router=None,
model="nvidia_riva/parakeet",
)
def test_admin_opt_in_proxy_wide_allows_use_ssl(self):
assert (
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"use_ssl": True,
},
general_settings={"allow_client_side_credentials": True},
llm_router=None,
model="nvidia_riva/parakeet",
)
is True
)
def test_admin_opt_in_per_deployment_allows_use_ssl(self, monkeypatch):
from litellm.proxy.auth import auth_utils
monkeypatch.setattr(
auth_utils,
"_allow_model_level_clientside_configurable_parameters",
lambda model, param, request_body_value, llm_router: param == "use_ssl",
)
assert (
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"use_ssl": True,
},
general_settings={},
llm_router=None,
model="nvidia_riva/parakeet",
)
is True
)
# ── is_request_body_safe nested-config recursion (VERIA-6) ────────────────────
@ -1748,6 +1905,22 @@ class TestIsRequestBodySafeNestedConfig:
model="milvus-store",
)
def test_nested_nvcf_function_id_in_metadata_blocked(self):
"""Smuggling ``nvcf_function_id`` via ``metadata`` / ``extra_body``
is the same shape as the VERIA-6 ``api_base`` bypass — must be
rejected by the recursive walk so the NVCF override gate cannot
be sidestepped with nesting."""
with pytest.raises(ValueError, match="nvcf_function_id"):
is_request_body_safe(
request_body={
"model": "nvidia_riva/parakeet",
"litellm_metadata": {"nvcf_function_id": "attacker-via-metadata"},
},
general_settings={},
llm_router=None,
model="nvidia_riva/parakeet",
)
def test_nested_langfuse_host_in_embedding_config_blocked(self):
"""The recursion uses the *full* banned-param list, not a special
subset — so any flag that's banned at the root is also banned

View file

@ -1,4 +1,5 @@
import asyncio
import copy
import json
import os
import sys
@ -1419,8 +1420,7 @@ async def test_batch_database_updates_isolation_on_failure():
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
litellm_proxy_budget_name="budget",
payload_copy={"key": "value"},
request_tags=None,
payload={"key": "value"},
)
# _update_key_db raised, but all others should still have been called
@ -1704,3 +1704,117 @@ async def test_commit_spend_updates_iterates_in_sorted_order(
)
assert captured_where_values == expected_order
@pytest.mark.asyncio
async def test_update_database_does_not_deepcopy_on_request_path():
"""
Regression for LIT-4088: copy.deepcopy must not run while the caller awaits
update_database(). The deepcopy used to isolate the daily-spend helpers is
relocated into the _batch_database_updates background task, and the spend-log
insert receives the payload directly (all consumers are read-only).
Asserts:
- zero copy.deepcopy calls happen on the awaited request path
- the batch background task still hands the daily helpers an isolated copy
(mutating the original after the task ran does not bleed into it)
- the spend-log insert receives the payload on the request path with the
correct content
"""
db_writer = DBSpendUpdateWriter()
captured_batch_payloads = []
captured_spend_log = {}
async def capture_batch_payload(**kwargs):
captured_batch_payloads.append(kwargs.get("payload"))
async def capture_spend_log(**kwargs):
payload = kwargs.get("payload")
captured_spend_log["ref"] = payload
captured_spend_log["model_at_call"] = payload["model"]
db_writer._insert_spend_log_to_db = AsyncMock(side_effect=capture_spend_log)
db_writer._update_user_db = AsyncMock()
db_writer._update_key_db = AsyncMock()
db_writer._update_team_db = AsyncMock()
db_writer._update_org_db = AsyncMock()
db_writer._update_tag_db = AsyncMock()
db_writer._update_agent_db = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_user_transaction = AsyncMock(
side_effect=capture_batch_payload
)
db_writer.add_spend_log_transaction_to_daily_end_user_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_agent_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_team_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_org_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_tag_transaction = AsyncMock()
fake_payload = {
"startTime": "2024-01-01T00:00:00",
"endTime": "2024-01-01T00:01:00",
"model": "gpt-4",
"custom_llm_provider": "openai",
"request_tags": '["prod-tag"]',
"spend": 0.0,
"nested": {"a": 1},
}
deepcopy_calls = []
real_deepcopy = copy.deepcopy
def counting_deepcopy(obj, *args, **kwargs):
deepcopy_calls.append(obj)
return real_deepcopy(obj, *args, **kwargs)
with (
patch("litellm.proxy.proxy_server.disable_spend_logs", False),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()),
patch("litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"),
patch(
"litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload",
return_value=fake_payload,
),
patch(
"litellm.proxy.db.db_spend_update_writer.copy.deepcopy",
counting_deepcopy,
),
):
await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id="test-end-user",
team_id="test-team",
org_id="test-org",
kwargs={"model": "gpt-4", "custom_llm_provider": "openai"},
completion_response=MagicMock(),
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.1,
)
# Request path is clean: nothing was deepcopied while the caller awaited.
assert len(deepcopy_calls) == 0
# The spend-log insert ran inline on the request path with the real payload.
assert captured_spend_log["ref"] is fake_payload
assert captured_spend_log["model_at_call"] == "gpt-4"
assert fake_payload["spend"] == 0.1
# Now let the batch background task run; the deepcopy happens here.
await asyncio.sleep(0)
assert len(deepcopy_calls) >= 1
assert len(captured_batch_payloads) == 1
batch_payload = captured_batch_payloads[0]
assert batch_payload is not fake_payload
assert batch_payload["model"] == "gpt-4"
assert batch_payload["spend"] == 0.1
# Mutating the original after the batch task captured its snapshot must not
# leak into the daily helper's isolated copy.
fake_payload["model"] = "MUTATED"
fake_payload["nested"]["a"] = 999
assert batch_payload["model"] == "gpt-4"
assert batch_payload["nested"]["a"] == 1

View file

@ -148,6 +148,37 @@ def test_is_database_service_unavailable_error_prisma_p1001_masquerades_as_datae
)
def test_is_prisma_data_error_only_true_for_dataerror():
"""The spend-log poison-row isolation gates on this: only a prisma
``DataError`` (the DB refused the data, e.g. a NUL byte) may be bisected
into a per-row drop. A connectivity failure or any non-prisma exception
must not be treated as a data rejection, so the whole batch surfaces."""
import httpx
data_error = DataError(data={"user_facing_error": {"message": "invalid byte sequence for encoding UTF8: 0x00"}})
assert PrismaDBExceptionHandler.is_prisma_data_error(data_error) is True
for non_data in (
httpx.ConnectError("conn refused"),
PrismaError("can't reach database server"),
UniqueViolationError(data={"user_facing_error": {"meta": {"table": "t"}}}),
RuntimeError("boom"),
):
assert PrismaDBExceptionHandler.is_prisma_data_error(non_data) is False
def test_is_prisma_data_error_true_for_connection_masquerade_dataerror():
"""The P1001 outage prisma mislabels as a ``DataError`` is still a
``DataError`` by type, so this returns True; the spend-log helper relies on
``is_database_service_unavailable_error`` (not this check) to keep that
outage on the retry path instead of dropping rows."""
p1001_as_dataerror = DataError(
data={"user_facing_error": {"message": "Can't reach database server at `127.0.0.1`:`5499`"}}
)
assert PrismaDBExceptionHandler.is_prisma_data_error(p1001_as_dataerror) is True
assert PrismaDBExceptionHandler.is_database_service_unavailable_error(p1001_as_dataerror) is True
def test_is_database_service_unavailable_error_cached_plan_escapes_as_503():
"""Composes with the cached-plan retry: when that recovery fails and the
Postgres "cached plan must not change result type" error escapes (raised by

View file

@ -609,7 +609,6 @@ class TestImageSupport:
request_data=mock_request_data_input,
input_type="request",
)
result_texts = guardrailed_inputs.get("texts", [])
result_images = guardrailed_inputs.get("images", None)
# Verify API was called with images
@ -943,7 +942,7 @@ class TestMultimodalSupport:
guardrail.async_handler, "post", return_value=mock_response
) as mock_post:
# This should not raise SerializationIterator error
result = await guardrail.apply_guardrail(
await guardrail.apply_guardrail(
inputs={
"texts": ["What's in this image?"],
"images": ["https://example.com/image.jpg"],
@ -1006,7 +1005,7 @@ class TestMultimodalSupport:
with patch.object(
guardrail.async_handler, "post", return_value=mock_response
) as mock_post:
result = await guardrail.apply_guardrail(
await guardrail.apply_guardrail(
inputs={
"texts": ["Hello", "World"],
"structured_messages": messages_with_iterable,
@ -1023,6 +1022,717 @@ class TestMultimodalSupport:
assert isinstance(json_payload["structured_messages"], list)
def _make_stream_chunk(content: str, finish_reason=None):
"""Build a real ModelResponseStream so the handler's isinstance checks pass."""
from litellm.types.utils import Delta, ModelResponseStream
return ModelResponseStream(
model="gpt-4",
choices=[
litellm.StreamingChoices(
index=0,
delta=Delta(role="assistant", content=content),
finish_reason=finish_reason,
)
],
)
def _make_assembled_model_response(content: str) -> ModelResponse:
return ModelResponse(
id="mock-response",
model="gpt-4",
choices=[
litellm.Choices(
index=0,
message=litellm.Message(role="assistant", content=content),
finish_reason="stop",
)
],
)
def _mock_guardrail_post_response(action: str = "NONE", texts=None, blocked_reason=None):
mock_response = MagicMock()
payload = {"action": action}
if texts is not None:
payload["texts"] = texts
if blocked_reason is not None:
payload["blocked_reason"] = blocked_reason
mock_response.json.return_value = payload
mock_response.raise_for_status = MagicMock()
return mock_response
def _make_responses_stream_events(text: str):
"""Minimal /v1/responses SSE event sequence ending in response.completed."""
return (
{"type": "response.created", "response": {"id": "resp_test"}},
{
"type": "response.output_item.added",
"item": {"type": "message", "id": "msg_test"},
},
{
"type": "response.content_part.added",
"part": {"type": "output_text", "text": ""},
},
{"type": "response.output_text.delta", "delta": text},
{
"type": "response.output_text.done",
"text": text,
},
{
"type": "response.completed",
"response": {
"id": "resp_test",
"output": [
{
"type": "message",
"id": "msg_test",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": text}],
}
],
"status": "completed",
},
},
)
class TestGenericGuardrailAPIStreamingConfig:
"""Streaming knobs on GenericGuardrailAPI and initialize_guardrail plumbing."""
def test_streaming_defaults(self):
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
)
assert guardrail.streaming_end_of_stream_only is False
assert guardrail.streaming_sampling_rate == 5
def test_streaming_overrides(self):
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_end_of_stream_only=True,
streaming_sampling_rate=2,
)
assert guardrail.streaming_end_of_stream_only is True
assert guardrail.streaming_sampling_rate == 2
@pytest.mark.parametrize("invalid_rate", [0, -1, -5])
def test_streaming_sampling_rate_rejects_non_positive(self, invalid_rate):
with pytest.raises(ValueError, match="streaming_sampling_rate must be >= 1"):
GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_sampling_rate=invalid_rate,
)
def test_optional_params_streaming_sampling_rate_ge_one(self):
from pydantic import ValidationError
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIOptionalParams,
)
with pytest.raises(ValidationError):
GenericGuardrailAPIOptionalParams(streaming_sampling_rate=0)
def test_get_config_model(self):
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIConfigModel,
)
assert GenericGuardrailAPI.get_config_model() is GenericGuardrailAPIConfigModel
def test_initialize_guardrail_forwards_streaming_flags(self):
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
initialize_guardrail,
)
from litellm.types.guardrails import LitellmParams
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="post_call",
api_base="https://api.test.guardrail.com",
default_on=False,
)
# LitellmParams uses extra="allow" on the base; set streaming knobs dynamically
litellm_params.streaming_end_of_stream_only = False # type: ignore[attr-defined]
litellm_params.streaming_sampling_rate = 3 # type: ignore[attr-defined]
guardrail_config = {"guardrail_name": "test-generic-streaming"}
with patch(
"litellm.logging_callback_manager.add_litellm_callback"
):
guardrail = initialize_guardrail(litellm_params, guardrail_config)
assert guardrail.streaming_end_of_stream_only is False
assert guardrail.streaming_sampling_rate == 3
def test_initialize_guardrail_optional_params_defaults_do_not_shadow_top_level(
self,
):
"""Top-level streaming knobs win when optional_params only carries siblings."""
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
initialize_guardrail,
)
from litellm.types.guardrails import LitellmParams
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIOptionalParams,
)
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="post_call",
api_base="https://api.test.guardrail.com",
default_on=False,
)
litellm_params.streaming_end_of_stream_only = True # type: ignore[attr-defined]
litellm_params.streaming_sampling_rate = 2 # type: ignore[attr-defined]
# Sibling optional_params only; streaming fields stay at Pydantic default None.
litellm_params.optional_params = GenericGuardrailAPIOptionalParams( # type: ignore[attr-defined]
additional_provider_specific_params={"tenant": "acme"},
)
guardrail_config = {"guardrail_name": "test-generic-streaming-mixed"}
with patch(
"litellm.logging_callback_manager.add_litellm_callback"
):
guardrail = initialize_guardrail(litellm_params, guardrail_config)
assert guardrail.streaming_end_of_stream_only is True
assert guardrail.streaming_sampling_rate == 2
def test_initialize_guardrail_explicit_optional_params_streaming_wins(self):
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
initialize_guardrail,
)
from litellm.types.guardrails import LitellmParams
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIOptionalParams,
)
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="post_call",
api_base="https://api.test.guardrail.com",
default_on=False,
)
litellm_params.streaming_end_of_stream_only = False # type: ignore[attr-defined]
litellm_params.streaming_sampling_rate = 9 # type: ignore[attr-defined]
litellm_params.optional_params = GenericGuardrailAPIOptionalParams( # type: ignore[attr-defined]
streaming_end_of_stream_only=True,
streaming_sampling_rate=1,
)
guardrail_config = {"guardrail_name": "test-generic-streaming-nested-wins"}
with patch(
"litellm.logging_callback_manager.add_litellm_callback"
):
guardrail = initialize_guardrail(litellm_params, guardrail_config)
assert guardrail.streaming_end_of_stream_only is True
assert guardrail.streaming_sampling_rate == 1
def test_initialize_guardrail_dict_optional_params_streaming_wins(self):
"""Guardrail API/UI delivers optional_params as a plain dict, not a model."""
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
initialize_guardrail,
)
from litellm.types.guardrails import LitellmParams
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="post_call",
api_base="https://api.test.guardrail.com",
default_on=False,
)
litellm_params.streaming_end_of_stream_only = False # type: ignore[attr-defined]
litellm_params.streaming_sampling_rate = 9 # type: ignore[attr-defined]
# Plain dict mirrors how configs arrive from the guardrail API/UI.
litellm_params.optional_params = { # type: ignore[attr-defined]
"streaming_end_of_stream_only": True,
"streaming_sampling_rate": 1,
}
guardrail_config = {"guardrail_name": "test-generic-streaming-dict-optional"}
with patch(
"litellm.logging_callback_manager.add_litellm_callback"
):
guardrail = initialize_guardrail(litellm_params, guardrail_config)
assert guardrail.streaming_end_of_stream_only is True
assert guardrail.streaming_sampling_rate == 1
def test_initialize_guardrail_dict_optional_params_sibling_only_falls_through(
self,
):
"""Dict optional_params without streaming keys must not shadow top-level knobs."""
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
initialize_guardrail,
)
from litellm.types.guardrails import LitellmParams
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="post_call",
api_base="https://api.test.guardrail.com",
default_on=False,
)
litellm_params.streaming_end_of_stream_only = True # type: ignore[attr-defined]
litellm_params.streaming_sampling_rate = 2 # type: ignore[attr-defined]
litellm_params.optional_params = { # type: ignore[attr-defined]
"additional_provider_specific_params": {"tenant": "acme"},
}
guardrail_config = {"guardrail_name": "test-generic-streaming-dict-sibling"}
with patch(
"litellm.logging_callback_manager.add_litellm_callback"
):
guardrail = initialize_guardrail(litellm_params, guardrail_config)
assert guardrail.streaming_end_of_stream_only is True
assert guardrail.streaming_sampling_rate == 2
class TestGenericGuardrailAPIStreamingViaUnified:
"""Streaming output checks routed through UnifiedLLMGuardrails."""
@pytest.mark.asyncio
async def test_streaming_safe_content_yields_all_chunks(self):
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["Hello", " ", "world", "!", " Goodbye"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["Hello world! Goodbye"]
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("Hello world! Goodbye"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
chunks_received = 0
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
chunks_received += 1
assert chunks_received == 5
assert mock_post.await_count >= 1
@pytest.mark.asyncio
async def test_streaming_blocked_content_raises(self):
from litellm.exceptions import GuardrailRaisedException
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_sampling_rate=1,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["Hello", " ishaan", " here"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="BLOCKED", blocked_reason="Ishaan is not allowed"
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("Hello ishaan here"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
with pytest.raises(GuardrailRaisedException) as exc_info:
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert "Ishaan is not allowed" in str(exc_info.value)
@pytest.mark.asyncio
async def test_streaming_default_uses_sampled_cadence(self):
"""Default samples every 5th chunk + final pass: 10 chunks → calls at 5, 10, and final = 3."""
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["A", "B", "C", "D", "E", "F", "G", "H", "I", "J"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["ABCDEFGHIJ"]
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABCDEFGHIJ"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert mock_post.await_count == 3, (
f"Expected 3 guardrail calls (2 sampled at chunks 5 / 10 + 1 final), "
f"got {mock_post.await_count}"
)
for call in mock_post.await_args_list:
assert call.kwargs["json"]["input_type"] == "response"
@pytest.mark.asyncio
async def test_streaming_end_of_stream_only_calls_guardrail_once(self):
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_end_of_stream_only=True,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["A", "B", "C", "D", "E", "F", "G", "H", "I", "J"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["ABCDEFGHIJ"]
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABCDEFGHIJ"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert mock_post.await_count == 1, (
f"Expected exactly one guardrail call at end of stream, "
f"got {mock_post.await_count}"
)
@pytest.mark.asyncio
async def test_streaming_sampling_rate_override(self):
"""sampling_rate=2 on 6 chunks → in-stream at 2,4,6 plus final = 4 calls."""
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_end_of_stream_only=False,
streaming_sampling_rate=2,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["A", "B", "C", "D", "E", "F"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(action="NONE", texts=["ABCDEF"])
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABCDEF"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert mock_post.await_count == 4, (
f"Expected 4 guardrail calls (3 sampled + 1 final aggregate), "
f"got {mock_post.await_count}"
)
@pytest.mark.asyncio
async def test_streaming_fail_open_on_unreachable_continues_stream(self):
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
unreachable_fallback="fail_open",
streaming_end_of_stream_only=True,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
for i, content in enumerate(["A", "B", "C"]):
yield _make_stream_chunk(
content, finish_reason="stop" if i == 2 else None
)
mock_post = AsyncMock(side_effect=httpx.ConnectError("connection refused"))
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABC"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
chunks_received = 0
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
chunks_received += 1
assert chunks_received == 3
@pytest.mark.asyncio
async def test_responses_api_streaming_end_of_stream_only_calls_guardrail_once(self):
"""/v1/responses path through unified hook; end-of-stream-only = one call."""
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_end_of_stream_only=True,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_responses_stream():
for event in _make_responses_stream_events("Hello world"):
yield event
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["Hello world"]
)
)
with patch.object(guardrail.async_handler, "post", mock_post):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/v1/responses"
)
request_data = {
"input": "hi",
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
events_received = 0
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_responses_stream(),
request_data=request_data,
):
events_received += 1
assert events_received == 6
assert mock_post.await_count == 1, (
f"Expected exactly one guardrail call at end of /v1/responses stream, "
f"got {mock_post.await_count}"
)
assert mock_post.await_args.kwargs["json"]["input_type"] == "response"
@pytest.mark.asyncio
async def test_responses_api_streaming_blocked_raises(self):
"""Mid-stream BLOCKED on /v1/responses surfaces GuardrailRaisedException."""
from litellm.exceptions import GuardrailRaisedException
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_sampling_rate=1,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_responses_stream():
for event in _make_responses_stream_events("blocked content"):
yield event
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="BLOCKED", blocked_reason="Responses content not allowed"
)
)
with patch.object(guardrail.async_handler, "post", mock_post):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/v1/responses"
)
request_data = {
"input": "hi",
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
with pytest.raises(GuardrailRaisedException) as exc_info:
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_responses_stream(),
request_data=request_data,
):
pass
assert "Responses content not allowed" in str(exc_info.value)
class TestToolSupport:
"""Test tool handling in guardrail requests"""

View file

@ -8,15 +8,25 @@ Tests cover:
- response-type input is passed through unchanged
- /v1/compress HTTP error raises HTTPException
- /v1/compress returning malformed JSON raises HTTPException
- CCR: headroom_retrieve tool injected when compressed messages contain hashes
- CCR: async_should_run_agentic_loop returns True when response has headroom_retrieve tool calls
- CCR: async_build_agentic_loop_plan calls retrieve endpoint and builds follow-up messages
"""
import json
import time
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException
from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import HeadroomGuardrail
from litellm.proxy.guardrails.guardrail_hooks.headroom.headroom import (
HeadroomGuardrail,
extract_hashes_from_messages,
has_headroom_retrieve_tool,
HEADROOM_RETRIEVE_TOOL_NAME,
)
from litellm.types.utils import GenericGuardrailAPIInputs
FAKE_API_BASE = "https://headroom.example.com"
@ -30,6 +40,13 @@ COMPRESSED_MESSAGES = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "A" * 500},
]
COMPRESSED_MESSAGES_WITH_HASH = [
{"role": "system", "content": "You are a helpful assistant."},
{
"role": "user",
"content": "Summary. Retrieve more: hash=b573993006976af767214fac",
},
]
def _make_guardrail(**kwargs) -> HeadroomGuardrail:
@ -57,6 +74,36 @@ def _make_compress_response(messages: list, status: int = 200) -> MagicMock:
return mock
def _make_retrieve_response(original_content: str, status: int = 200) -> MagicMock:
mock = MagicMock()
mock.status_code = status
mock.json.return_value = {"original_content": original_content}
mock.text = original_content
return mock
def _make_openai_response_with_tool_call(tool_name: str, arguments: dict, tool_id: str = "call_abc123") -> MagicMock:
fn = MagicMock()
fn.name = tool_name
fn.arguments = json.dumps(arguments)
tc = MagicMock()
tc.id = tool_id
tc.type = "function"
tc.function = fn
message = MagicMock()
message.content = None
message.tool_calls = [tc]
choice = MagicMock()
choice.message = message
response = MagicMock()
response.choices = [choice]
return response
@pytest.fixture
def guardrail() -> HeadroomGuardrail:
return _make_guardrail()
@ -87,6 +134,567 @@ async def test_apply_guardrail_compresses_and_returns_structured_messages(
assert result.get("structured_messages") == COMPRESSED_MESSAGES
@pytest.mark.asyncio
async def test_apply_guardrail_injects_retrieve_tool_when_hashes_present(
guardrail: HeadroomGuardrail,
):
inputs = GenericGuardrailAPIInputs(
texts=["A" * 5000],
structured_messages=ORIGINAL_MESSAGES,
)
mock_response = _make_compress_response(COMPRESSED_MESSAGES_WITH_HASH)
with patch.object(
guardrail.async_handler,
"post",
new_callable=AsyncMock,
return_value=mock_response,
):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"model": "gpt-4o"},
input_type="request",
)
tools = result.get("tools")
assert tools is not None
assert has_headroom_retrieve_tool(tools)
@pytest.mark.asyncio
async def test_apply_guardrail_no_tool_injected_when_no_hashes(
guardrail: HeadroomGuardrail,
):
inputs = GenericGuardrailAPIInputs(
texts=["A" * 5000],
structured_messages=ORIGINAL_MESSAGES,
)
mock_response = _make_compress_response(COMPRESSED_MESSAGES)
with patch.object(
guardrail.async_handler,
"post",
new_callable=AsyncMock,
return_value=mock_response,
):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"model": "gpt-4o"},
input_type="request",
)
tools = result.get("tools")
assert not has_headroom_retrieve_tool(tools or [])
@pytest.mark.asyncio
async def test_apply_guardrail_preserves_existing_tools_when_injecting(
guardrail: HeadroomGuardrail,
):
existing_tool = {"type": "function", "function": {"name": "my_tool", "parameters": {}}}
inputs = GenericGuardrailAPIInputs(
texts=["A" * 5000],
structured_messages=ORIGINAL_MESSAGES,
tools=[existing_tool],
)
mock_response = _make_compress_response(COMPRESSED_MESSAGES_WITH_HASH)
with patch.object(
guardrail.async_handler,
"post",
new_callable=AsyncMock,
return_value=mock_response,
):
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"model": "gpt-4o"},
input_type="request",
)
tools = result.get("tools")
assert tools is not None
assert isinstance(tools, list)
assert any(isinstance(t, dict) and t.get("function", {}).get("name") == "my_tool" for t in tools)
assert has_headroom_retrieve_tool(tools)
@pytest.mark.asyncio
async def test_async_should_run_agentic_loop_returns_true_for_retrieve_call(
guardrail: HeadroomGuardrail,
):
retrieve_tool_def = [{"type": "function", "function": {"name": HEADROOM_RETRIEVE_TOOL_NAME}}]
response = _make_openai_response_with_tool_call(
tool_name=HEADROOM_RETRIEVE_TOOL_NAME,
arguments={"hash": "b573993006976af767214fac"},
)
should_run, ctx = await guardrail.async_should_run_agentic_loop(
response=response,
model="gpt-4o",
messages=[],
tools=retrieve_tool_def,
stream=False,
custom_llm_provider="openai",
kwargs={},
)
assert should_run is True
assert len(ctx["tool_calls"]) == 1
assert ctx["tool_calls"][0]["arguments"]["hash"] == "b573993006976af767214fac"
@pytest.mark.asyncio
async def test_async_should_run_agentic_loop_returns_false_without_retrieve_tool(
guardrail: HeadroomGuardrail,
):
other_tools = [{"type": "function", "function": {"name": "other_tool"}}]
response = _make_openai_response_with_tool_call(
tool_name="other_tool",
arguments={},
)
should_run, _ = await guardrail.async_should_run_agentic_loop(
response=response,
model="gpt-4o",
messages=[],
tools=other_tools,
stream=False,
custom_llm_provider="openai",
kwargs={},
)
assert should_run is False
@pytest.mark.asyncio
async def test_async_should_run_agentic_loop_returns_false_when_no_retrieve_calls(
guardrail: HeadroomGuardrail,
):
retrieve_tool_def = [{"type": "function", "function": {"name": HEADROOM_RETRIEVE_TOOL_NAME}}]
response = _make_openai_response_with_tool_call(
tool_name="some_other_function",
arguments={},
)
should_run, _ = await guardrail.async_should_run_agentic_loop(
response=response,
model="gpt-4o",
messages=[],
tools=retrieve_tool_def,
stream=False,
custom_llm_provider="openai",
kwargs={},
)
assert should_run is False
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_calls_retrieve_and_builds_messages(
guardrail: HeadroomGuardrail,
):
original_content = "This is the full compressed content."
mock_retrieve = _make_retrieve_response(original_content)
tool_calls = [
{
"id": "call_abc123",
"type": "function",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": {"hash": "b573993006976af767214fac"},
}
]
response = _make_openai_response_with_tool_call(
tool_name=HEADROOM_RETRIEVE_TOOL_NAME,
arguments={"hash": "b573993006976af767214fac"},
tool_id="call_abc123",
)
messages = [{"role": "user", "content": "What does it say? hash=b573993006976af767214fac"}]
guardrail._issued_hashes_by_call_id["call-1"] = (
frozenset({"b573993006976af767214fac"}),
time.monotonic() + 999,
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
return_value=mock_retrieve,
) as mock_get:
plan = await guardrail.async_build_agentic_loop_plan(
tools={"tool_calls": tool_calls},
model="gpt-4o",
messages=messages,
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=None,
stream=False,
kwargs={"litellm_call_id": "call-1"},
)
assert plan.run_agentic_loop is True
assert plan.request_patch is not None
follow_up = plan.request_patch.messages
assert follow_up is not None
tool_result_message = next((m for m in follow_up if m.get("role") == "tool"), None)
assert tool_result_message is not None
assert tool_result_message["content"] == original_content
assert tool_result_message["tool_call_id"] == "call_abc123"
mock_get.assert_called_once()
call_url = mock_get.call_args.kwargs.get("url") or mock_get.call_args.args[0]
assert "b573993006976af767214fac" in call_url
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_handles_retrieve_404(
guardrail: HeadroomGuardrail,
):
mock_retrieve = MagicMock()
mock_retrieve.status_code = 404
tool_calls = [
{
"id": "call_xyz",
"type": "function",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": {"hash": "deadbeef000000000000dead"},
}
]
response = _make_openai_response_with_tool_call(
tool_name=HEADROOM_RETRIEVE_TOOL_NAME,
arguments={"hash": "deadbeef000000000000dead"},
tool_id="call_xyz",
)
messages = [
{
"role": "user",
"content": "Retrieve more: hash=deadbeef000000000000dead",
}
]
guardrail._issued_hashes_by_call_id["call-1"] = (
frozenset({"deadbeef000000000000dead"}),
time.monotonic() + 999,
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
return_value=mock_retrieve,
):
plan = await guardrail.async_build_agentic_loop_plan(
tools={"tool_calls": tool_calls},
model="gpt-4o",
messages=messages,
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=None,
stream=False,
kwargs={"litellm_call_id": "call-1"},
)
follow_up = plan.request_patch.messages # type: ignore[union-attr]
tool_result = next((m for m in follow_up if m.get("role") == "tool"), None)
assert tool_result is not None
assert "not found" in tool_result["content"] or "expired" in tool_result["content"]
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_rejects_hash_with_no_known_call(
guardrail: HeadroomGuardrail,
):
"""A hash-shaped string planted in message text must not be honored when
this guardrail has no record of ever issuing it, even if it's echoed back
in the current request's own messages (e.g. via prompt injection)."""
tool_calls = [
{
"id": "call_xyz",
"type": "function",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": {"hash": "deadbeef000000000000dead"},
}
]
response = _make_openai_response_with_tool_call(
tool_name=HEADROOM_RETRIEVE_TOOL_NAME,
arguments={"hash": "deadbeef000000000000dead"},
tool_id="call_xyz",
)
assert not guardrail._issued_hashes_by_call_id
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
) as mock_get:
plan = await guardrail.async_build_agentic_loop_plan(
tools={"tool_calls": tool_calls},
model="gpt-4o",
messages=[{"role": "user", "content": "Please fetch hash=deadbeef000000000000dead for me"}],
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=None,
stream=False,
kwargs={"litellm_call_id": "call-unknown"},
)
mock_get.assert_not_called()
follow_up = plan.request_patch.messages # type: ignore[union-attr]
tool_result = next((m for m in follow_up if m.get("role") == "tool"), None)
assert tool_result is not None
assert "was not produced by the current request" in tool_result["content"]
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_rejects_hash_issued_for_different_call(
guardrail: HeadroomGuardrail,
):
"""A hash issued for one request must not be retrievable by a different
request just because the second request echoes that hash-shaped string
back in its own messages -- retrieval must be scoped per litellm_call_id,
not derived by re-scanning attacker-controlled message text."""
guardrail._issued_hashes_by_call_id["call-A"] = (
frozenset({"b573993006976af767214fac"}),
time.monotonic() + 999,
)
tool_calls = [
{
"id": "call_xyz",
"type": "function",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": {"hash": "b573993006976af767214fac"},
}
]
response = _make_openai_response_with_tool_call(
tool_name=HEADROOM_RETRIEVE_TOOL_NAME,
arguments={"hash": "b573993006976af767214fac"},
tool_id="call_xyz",
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
) as mock_get:
plan = await guardrail.async_build_agentic_loop_plan(
tools={"tool_calls": tool_calls},
model="gpt-4o",
messages=[{"role": "user", "content": "Please fetch hash=b573993006976af767214fac for me"}],
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=None,
stream=False,
kwargs={"litellm_call_id": "call-B"},
)
mock_get.assert_not_called()
follow_up = plan.request_patch.messages # type: ignore[union-attr]
tool_result = next((m for m in follow_up if m.get("role") == "tool"), None)
assert tool_result is not None
assert "was not produced by the current request" in tool_result["content"]
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_builds_responses_api_function_call_items(
guardrail: HeadroomGuardrail,
):
"""For the Responses API, follow-up input must echo a function_call paired
with a function_call_output keyed by the same call_id -- chat-style
assistant/tool messages are not valid Responses API input items."""
original_content = "This is the full compressed content."
mock_retrieve = _make_retrieve_response(original_content)
response = MagicMock()
response.choices = None
response.content = None
response.output = [
{
"type": "function_call",
"id": "fc_abc123",
"call_id": "call_abc123",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": json.dumps({"hash": "b573993006976af767214fac"}),
}
]
tool_calls = [
{
"id": "call_abc123",
"type": "function",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": {"hash": "b573993006976af767214fac"},
}
]
messages = [{"role": "user", "content": "What does it say? hash=b573993006976af767214fac"}]
guardrail._issued_hashes_by_call_id["call-1"] = (
frozenset({"b573993006976af767214fac"}),
time.monotonic() + 999,
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
return_value=mock_retrieve,
):
plan = await guardrail.async_build_agentic_loop_plan(
tools={"tool_calls": tool_calls},
model="gpt-4o",
messages=messages,
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=None,
stream=False,
kwargs={"litellm_call_id": "call-1"},
)
follow_up = plan.request_patch.messages # type: ignore[union-attr]
assert all("role" not in item for item in follow_up if item not in messages)
function_call_item = next((i for i in follow_up if i.get("type") == "function_call"), None)
assert function_call_item is not None
assert function_call_item["call_id"] == "call_abc123"
assert function_call_item["name"] == HEADROOM_RETRIEVE_TOOL_NAME
output_item = next((i for i in follow_up if i.get("type") == "function_call_output"), None)
assert output_item is not None
assert output_item["call_id"] == "call_abc123"
assert output_item["output"] == original_content
@pytest.mark.asyncio
async def test_async_build_agentic_loop_plan_builds_anthropic_tool_result_messages(
guardrail: HeadroomGuardrail,
):
"""For the Anthropic Messages API, follow-up must echo a tool_use content
block in an assistant message paired with a tool_result content block in a
user message keyed by the same tool_use_id -- chat-style tool-role
messages are not valid Anthropic input.
AnthropicMessagesResponse is a TypedDict, so real responses are plain
dicts at runtime; a MagicMock response here would pass even if branch
selection used bare getattr() and silently fell through to the
chat-completions replay shape for every real Anthropic response.
"""
original_content = "This is the full compressed content."
mock_retrieve = _make_retrieve_response(original_content)
response = {
"content": [
{
"type": "tool_use",
"id": "toolu_abc123",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"input": {"hash": "b573993006976af767214fac"},
}
]
}
tool_calls = [
{
"id": "toolu_abc123",
"type": "function",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": {"hash": "b573993006976af767214fac"},
}
]
messages = [{"role": "user", "content": "What does it say? hash=b573993006976af767214fac"}]
guardrail._issued_hashes_by_call_id["call-1"] = (
frozenset({"b573993006976af767214fac"}),
time.monotonic() + 999,
)
with patch.object(
guardrail.async_handler,
"get",
new_callable=AsyncMock,
return_value=mock_retrieve,
):
plan = await guardrail.async_build_agentic_loop_plan(
tools={"tool_calls": tool_calls},
model="claude-sonnet-4-5",
messages=messages,
response=response,
anthropic_messages_provider_config=None,
anthropic_messages_optional_request_params={},
logging_obj=None,
stream=False,
kwargs={"litellm_call_id": "call-1"},
)
follow_up = plan.request_patch.messages # type: ignore[union-attr]
assert all(m.get("role") != "tool" for m in follow_up)
assistant_message = next((m for m in follow_up if m.get("role") == "assistant"), None)
assert assistant_message is not None
tool_use_block = next((b for b in assistant_message["content"] if b.get("type") == "tool_use"), None)
assert tool_use_block is not None
assert tool_use_block["id"] == "toolu_abc123"
user_message = follow_up[-1]
assert user_message["role"] == "user"
tool_result_block = next((b for b in user_message["content"] if b.get("type") == "tool_result"), None)
assert tool_result_block is not None
assert tool_result_block["tool_use_id"] == "toolu_abc123"
assert tool_result_block["content"] == original_content
def test_extract_hashes_from_messages_finds_hashes():
messages = [
{"role": "user", "content": "Retrieve more: hash=b573993006976af767214fac"},
{"role": "assistant", "content": "Also: hash=aabbccdd001122334455aabb"},
]
hashes = extract_hashes_from_messages(messages)
assert "b573993006976af767214fac" in hashes
assert "aabbccdd001122334455aabb" in hashes
def test_extract_hashes_from_messages_ignores_short_hashes():
messages = [{"role": "user", "content": "hash=tooshort"}]
hashes = extract_hashes_from_messages(messages)
assert not hashes
def test_extract_hashes_from_list_content_blocks():
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "hash=b573993006976af767214fac found here"},
],
}
]
hashes = extract_hashes_from_messages(messages)
assert "b573993006976af767214fac" in hashes
def test_has_headroom_retrieve_tool_recognizes_anthropic_native_shape():
"""By the time an Anthropic Messages API response reaches the agentic-loop
gate, the OpenAI-shaped tool this guardrail injects (type: "function")
has already been transformed into Anthropic's native tool shape
(type: "custom", top-level "name", no nested "function" object)."""
anthropic_native_tools = [
{
"type": "custom",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"input_schema": {"type": "object", "properties": {"hash": {"type": "string"}}},
}
]
assert has_headroom_retrieve_tool(anthropic_native_tools)
assert not has_headroom_retrieve_tool([{"type": "custom", "name": "some_other_tool"}])
@pytest.mark.asyncio
async def test_apply_guardrail_bypass_header_skips_compression(
guardrail: HeadroomGuardrail,
@ -97,9 +705,7 @@ async def test_apply_guardrail_bypass_header_skips_compression(
)
request_data = {"proxy_server_request": {"headers": {"x-headroom-bypass": "true"}}}
with patch.object(
guardrail.async_handler, "post", new_callable=AsyncMock
) as mock_post:
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data=request_data,
@ -119,9 +725,7 @@ async def test_apply_guardrail_response_type_passthrough(
structured_messages=ORIGINAL_MESSAGES,
)
with patch.object(
guardrail.async_handler, "post", new_callable=AsyncMock
) as mock_post:
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
@ -138,9 +742,7 @@ async def test_apply_guardrail_empty_structured_messages_passthrough(
):
inputs = GenericGuardrailAPIInputs(texts=["hello"])
with patch.object(
guardrail.async_handler, "post", new_callable=AsyncMock
) as mock_post:
with patch.object(guardrail.async_handler, "post", new_callable=AsyncMock) as mock_post:
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={},
@ -277,9 +879,7 @@ def test_bypass_header_case_insensitive():
guardrail = _make_guardrail()
for header_value in ("true", "True", "TRUE"):
data = {
"proxy_server_request": {"headers": {"x-headroom-bypass": header_value}}
}
data = {"proxy_server_request": {"headers": {"x-headroom-bypass": header_value}}}
assert guardrail._should_bypass(data) is True
data = {"proxy_server_request": {"headers": {"x-headroom-bypass": "false"}}}
@ -344,3 +944,118 @@ async def test_apply_guardrail_sends_model_from_request_data_when_no_config_mode
call_kwargs = mock_post.call_args
sent_payload = call_kwargs.kwargs.get("json") or call_kwargs.args[1]
assert sent_payload.get("model") == "gpt-4o"
@pytest.mark.asyncio
async def test_async_should_run_agentic_loop_detects_anthropic_content_block_format(
guardrail: HeadroomGuardrail,
):
# Anthropic's native tool format (type: "custom", top-level "name") --
# by the time a Messages API response reaches this gate, the OpenAI-shaped
# tool this guardrail injects has already been transformed into this shape.
retrieve_tool_def = [
{
"type": "custom",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"input_schema": {"type": "object", "properties": {"hash": {"type": "string"}}},
}
]
response = MagicMock()
response.choices = None
response.content = [
{
"type": "tool_use",
"id": "toolu_abc",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"input": {"hash": "b573993006976af767214fac"},
}
]
should_run, ctx = await guardrail.async_should_run_agentic_loop(
response=response,
model="claude-sonnet-4-6",
messages=[],
tools=retrieve_tool_def,
stream=False,
custom_llm_provider="anthropic",
kwargs={},
)
assert should_run is True
assert len(ctx["tool_calls"]) == 1
assert ctx["tool_calls"][0]["arguments"]["hash"] == "b573993006976af767214fac"
@pytest.mark.asyncio
async def test_async_should_run_agentic_loop_detects_anthropic_response_as_plain_dict(
guardrail: HeadroomGuardrail,
):
"""AnthropicMessagesResponse is a TypedDict -- real Messages API responses
are plain dicts at runtime, not objects with attribute access. A
MagicMock-only test would pass even if detection used bare getattr() and
silently treated every real response as having no tool calls."""
retrieve_tool_def = [
{
"type": "custom",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"input_schema": {"type": "object", "properties": {"hash": {"type": "string"}}},
}
]
response = {
"content": [
{
"type": "tool_use",
"id": "toolu_abc",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"input": {"hash": "b573993006976af767214fac"},
}
]
}
should_run, ctx = await guardrail.async_should_run_agentic_loop(
response=response,
model="claude-sonnet-4-6",
messages=[],
tools=retrieve_tool_def,
stream=False,
custom_llm_provider="anthropic",
kwargs={},
)
assert should_run is True
assert len(ctx["tool_calls"]) == 1
assert ctx["tool_calls"][0]["arguments"]["hash"] == "b573993006976af767214fac"
@pytest.mark.asyncio
async def test_async_should_run_agentic_loop_detects_responses_api_output_format(
guardrail: HeadroomGuardrail,
):
retrieve_tool_def = [{"type": "function", "function": {"name": HEADROOM_RETRIEVE_TOOL_NAME}}]
response = MagicMock()
response.choices = None
response.content = None
response.output = [
{
"type": "function_call",
"id": "fc_abc123",
"name": HEADROOM_RETRIEVE_TOOL_NAME,
"arguments": json.dumps({"hash": "b573993006976af767214fac"}),
}
]
should_run, ctx = await guardrail.async_should_run_agentic_loop(
response=response,
model="gpt-4o",
messages=[],
tools=retrieve_tool_def,
stream=False,
custom_llm_provider="openai",
kwargs={},
)
assert should_run is True
assert len(ctx["tool_calls"]) == 1
assert ctx["tool_calls"][0]["arguments"]["hash"] == "b573993006976af767214fac"

View file

@ -258,6 +258,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp
)
import litellm
from litellm.proxy._types import UserAPIKeyAuth
settings = DefaultInternalUserParams(
user_role=LitellmUserRoles.INTERNAL_USER,
@ -266,6 +267,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp
settings=settings,
settings_key="default_internal_user_params",
success_message="ok",
user_api_key_dict=UserAPIKeyAuth(user_id="test-admin"),
)
# Verify the in-memory variable was actually updated

View file

@ -134,8 +134,10 @@ async def test_update_customer_creates_budget_with_proper_relations(
)
# Mock end user update
mock_updated_user = MagicMock()
mock_updated_user.model_dump.return_value = {"user_id": "test-user", "blocked": False}
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=MagicMock()
return_value=mock_updated_user
)
# Create update request with budget creation fields (not just budget_id)
@ -190,8 +192,10 @@ async def test_update_customer_creates_budget_with_required_fields(
)
# Mock end user update
mock_updated_user = MagicMock()
mock_updated_user.model_dump.return_value = {"user_id": "test-user", "blocked": False}
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=MagicMock()
return_value=mock_updated_user
)
# Create update request with budget creation fields
@ -253,8 +257,10 @@ async def test_update_customer_budget_creation_with_fallback_admin(
)
# Mock end user update
mock_updated_user = MagicMock()
mock_updated_user.model_dump.return_value = {"user_id": "test-user", "blocked": False}
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=MagicMock()
return_value=mock_updated_user
)
# Create update request with budget creation fields
@ -309,6 +315,7 @@ async def test_update_customer_with_budget_id_and_creation_fields(
# Mock end user update
mock_updated_user = MagicMock()
mock_updated_user.model_dump.return_value = {"user_id": "test-user", "blocked": False}
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=mock_updated_user
)

View file

@ -1,18 +1,28 @@
from typing import List
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import FastAPI, HTTPException, Request, status
from fastapi.responses import JSONResponse
from fastapi.routing import APIRoute
from fastapi.testclient import TestClient
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_EndUserTable,
LitellmUserRoles,
ProxyException,
)
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.management_endpoints.customer_endpoints import router
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.customer_endpoints import (
BlockUsersResponse,
CustomerResponse,
DeleteCustomersResponse,
UnblockUsersResponse,
)
app = FastAPI()
@ -22,9 +32,7 @@ async def openai_exception_handler(request: Request, exc: ProxyException):
headers = exc.headers
error_dict = exc.to_dict()
return JSONResponse(
status_code=(
int(exc.code) if exc.code else status.HTTP_500_INTERNAL_SERVER_ERROR
),
status_code=(int(exc.code) if exc.code else status.HTTP_500_INTERNAL_SERVER_ERROR),
content={"error": error_dict},
headers=headers,
)
@ -54,30 +62,20 @@ def mock_user_api_key_auth():
def test_update_customer_success(mock_prisma_client, mock_user_api_key_auth):
# Mock the database responses
mock_end_user = LiteLLM_EndUserTable(
user_id="test-user-1", alias="Test User", blocked=False
)
updated_mock_end_user = LiteLLM_EndUserTable(
user_id="test-user-1", alias="Updated Test User", blocked=False
)
mock_end_user = LiteLLM_EndUserTable(user_id="test-user-1", alias="Test User", blocked=False)
updated_mock_end_user = LiteLLM_EndUserTable(user_id="test-user-1", alias="Updated Test User", blocked=False)
# Mock the find_first response
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=mock_end_user
)
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=mock_end_user)
# Mock the update response
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(
return_value=updated_mock_end_user
)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=updated_mock_end_user)
# Test data
test_data = {"user_id": "test-user-1", "alias": "Updated Test User"}
# Make the request
response = client.post(
"/customer/update", json=test_data, headers={"Authorization": "Bearer test-key"}
)
response = client.post("/customer/update", json=test_data, headers={"Authorization": "Bearer test-key"})
# Assert response
assert response.status_code == 200
@ -106,10 +104,7 @@ def test_update_customer_not_found(mock_prisma_client, mock_user_api_key_auth):
assert response.status_code == 404
response_json = response.json()
assert "error" in response_json
assert (
response_json["error"]["message"]
== "End User Id=non-existent-user does not exist in db"
)
assert response_json["error"]["message"] == "End User Id=non-existent-user does not exist in db"
assert response_json["error"]["type"] == "not_found"
assert response_json["error"]["param"] == "user_id"
assert response_json["error"]["code"] == "404"
@ -132,10 +127,7 @@ def test_info_customer_not_found(mock_prisma_client, mock_user_api_key_auth):
assert response.status_code == 404
response_json = response.json()
assert "error" in response_json
assert (
response_json["error"]["message"]
== "End User Id=non-existent-user does not exist in db"
)
assert response_json["error"]["message"] == "End User Id=non-existent-user does not exist in db"
assert response_json["error"]["type"] == "not_found"
assert response_json["error"]["param"] == "end_user_id"
assert response_json["error"]["code"] == "404"
@ -220,11 +212,6 @@ def test_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
assert error["code"] == "404"
# Test /customer/new - duplicate user error
from unittest.mock import MagicMock
mock_end_user = LiteLLM_EndUserTable(
user_id="existing-user", alias="Existing User", blocked=False
)
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(
side_effect=Exception("Unique constraint failed on the fields: (`user_id`)")
)
@ -238,9 +225,7 @@ def test_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
assert error["code"] == "400"
def test_customer_endpoints_error_schema_consistency(
mock_prisma_client, mock_user_api_key_auth
):
def test_customer_endpoints_error_schema_consistency(mock_prisma_client, mock_user_api_key_auth):
"""
Test the exact scenarios from the curl examples provided.
@ -307,9 +292,7 @@ def test_customer_endpoints_error_schema_consistency(
assert "Customer already exists" in error2["message"]
# Verify both errors have the same schema structure
assert set(error1.keys()) == set(
error2.keys()
), "Both errors should have the same top-level keys"
assert set(error1.keys()) == set(error2.keys()), "Both errors should have the same top-level keys"
# Both should have string values for all fields
for key in ["message", "type", "code"]:
@ -317,6 +300,153 @@ def test_customer_endpoints_error_schema_consistency(
assert isinstance(error2[key], str), f"error2[{key}] should be a string"
EXPECTED_RESPONSE_MODELS = {
"/customer/block": BlockUsersResponse,
"/customer/unblock": UnblockUsersResponse,
"/customer/new": CustomerResponse,
"/customer/update": CustomerResponse,
"/customer/delete": DeleteCustomersResponse,
"/customer/info": CustomerResponse,
"/customer/list": List[CustomerResponse],
"/customer/daily/activity": SpendAnalyticsPaginatedResponse,
}
@pytest.mark.parametrize("path, expected_model", EXPECTED_RESPONSE_MODELS.items())
def test_customer_routes_declare_response_model(path, expected_model):
"""
Every public /customer/* operation must declare a typed response_model so
the generated OpenAPI schema documents the response body. Regression for the
OpenAPI response-type coverage goal: drop a response_model and this fails.
"""
route = next(r for r in router.routes if isinstance(r, APIRoute) and r.path == path)
assert route.response_model == expected_model
def test_customer_new_documented_in_openapi_schema():
"""
The response_model must surface in the OpenAPI schema as a concrete ref, not
an empty/default response. This is what the coverage metric measures.
"""
schema = app.openapi()["paths"]["/customer/new"]["post"]
json_schema = schema["responses"]["200"]["content"]["application/json"]["schema"]
assert json_schema["$ref"].endswith("/CustomerResponse")
def test_update_customer_response_preserves_budget_id(mock_prisma_client, mock_user_api_key_auth):
"""
Regression for the response_model field-stripping concern: budget_id is a real
column on the end-user table that /customer/update echoes. response_model=
LiteLLM_EndUserTable must NOT drop it, so budget_id stays in LiteLLM_EndUserTable.
"""
existing = LiteLLM_EndUserTable(user_id="cust-1", blocked=False)
updated = LiteLLM_EndUserTable(user_id="cust-1", blocked=False, budget_id="budget-123")
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=existing)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=updated)
response = client.post(
"/customer/update",
json={"user_id": "cust-1", "budget_id": "budget-123"},
headers={"Authorization": "Bearer test-key"},
)
assert response.status_code == 200
assert response.json()["budget_id"] == "budget-123"
def test_update_customer_response_keeps_nested_budget_server_fields(mock_prisma_client, mock_user_api_key_auth):
"""
Faithfulness regression: /customer/update embeds the full budget row. The
response_model must keep the server-managed budget fields the endpoint used
to return (budget_reset_at, created_at) instead of the narrow write-allowlist
shape. The intentionally-internal audit fields (created_by/updated_by) stay out.
"""
existing = LiteLLM_EndUserTable(user_id="cust-1", blocked=False)
raw_row = MagicMock()
raw_row.model_dump.return_value = {
"user_id": "cust-1",
"blocked": False,
"alias": "renamed",
"spend": 0.0,
"allowed_model_region": None,
"default_model": None,
"budget_id": "b-1",
"object_permission_id": None,
"object_permission": None,
"litellm_budget_table": {
"budget_id": "b-1",
"max_budget": 10.0,
"budget_duration": "30d",
"budget_reset_at": "2024-02-01T00:00:00",
"created_at": "2024-01-01T00:00:00",
"created_by": "admin",
"updated_at": "2024-01-02T00:00:00",
"updated_by": "admin",
},
}
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=existing)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=raw_row)
response = client.post(
"/customer/update",
json={"user_id": "cust-1", "alias": "renamed"},
headers={"Authorization": "Bearer test-key"},
)
assert response.status_code == 200
budget = response.json()["litellm_budget_table"]
assert budget["budget_reset_at"] == "2024-02-01T00:00:00"
assert budget["created_at"] == "2024-01-01T00:00:00"
assert "created_by" not in budget
assert "updated_by" not in budget
def test_block_customer_success_serializes_through_response_model(mock_prisma_client, mock_user_api_key_auth):
"""
/customer/block returns {"blocked_users": [<end user rows>]}. With
response_model=BlockUsersResponse, a shape mismatch would raise a 500
ResponseValidationError, so a clean 200 proves the model matches runtime output.
"""
blocked_row = LiteLLM_EndUserTable(user_id="blocked-1", blocked=True)
mock_prisma_client.db.litellm_endusertable.upsert = AsyncMock(return_value=blocked_row)
response = client.post(
"/customer/block",
json={"user_ids": ["blocked-1"]},
headers={"Authorization": "Bearer test-key"},
)
assert response.status_code == 200
body = response.json()
assert body["blocked_users"][0]["user_id"] == "blocked-1"
assert body["blocked_users"][0]["blocked"] is True
def test_delete_customer_success_serializes_through_response_model(mock_prisma_client, mock_user_api_key_auth):
"""
/customer/delete returns {"deleted_customers": <int>, "message": <str>}.
response_model=DeleteCustomersResponse enforces that exact shape.
"""
existing = [
LiteLLM_EndUserTable(user_id="u1", blocked=False),
LiteLLM_EndUserTable(user_id="u2", blocked=False),
]
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=existing)
mock_prisma_client.db.litellm_endusertable.delete_many = AsyncMock(return_value=2)
response = client.post(
"/customer/delete",
json={"user_ids": ["u1", "u2"]},
headers={"Authorization": "Bearer test-key"},
)
assert response.status_code == 200
assert response.json() == {
"deleted_customers": 2,
"message": "Successfully deleted customers with ids: ['u1', 'u2']",
}
@pytest.mark.asyncio
async def test_get_customer_daily_activity_admin_param_passing(monkeypatch):
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
@ -331,9 +461,7 @@ async def test_get_customer_daily_activity_admin_param_passing(monkeypatch):
mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse")
get_daily_activity_mock = AsyncMock(return_value=mocked_response)
monkeypatch.setattr(
customer_endpoints, "get_daily_activity", get_daily_activity_mock
)
monkeypatch.setattr(customer_endpoints, "get_daily_activity", get_daily_activity_mock)
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin1")
result = await get_customer_daily_activity(
@ -380,16 +508,12 @@ async def test_get_customer_daily_activity_with_end_user_aliases(monkeypatch):
mock_end_user2.user_id = "end-user-2"
mock_end_user2.alias = "Customer Two"
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(
return_value=[mock_end_user1, mock_end_user2]
)
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[mock_end_user1, mock_end_user2])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse")
get_daily_activity_mock = AsyncMock(return_value=mocked_response)
monkeypatch.setattr(
customer_endpoints, "get_daily_activity", get_daily_activity_mock
)
monkeypatch.setattr(customer_endpoints, "get_daily_activity", get_daily_activity_mock)
auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin1")
await get_customer_daily_activity(
@ -436,9 +560,7 @@ async def test_get_customer_daily_activity_non_admin_is_rejected(monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
get_daily_activity_mock = AsyncMock()
monkeypatch.setattr(
customer_endpoints, "get_daily_activity", get_daily_activity_mock
)
monkeypatch.setattr(customer_endpoints, "get_daily_activity", get_daily_activity_mock)
non_admin_key = UserAPIKeyAuth(
user_id="regular-user-abc",
@ -482,9 +604,7 @@ async def test_get_customer_daily_activity_service_account_key_is_rejected(monke
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
get_daily_activity_mock = AsyncMock()
monkeypatch.setattr(
customer_endpoints, "get_daily_activity", get_daily_activity_mock
)
monkeypatch.setattr(customer_endpoints, "get_daily_activity", get_daily_activity_mock)
service_account_key = UserAPIKeyAuth(
user_id=None,
@ -507,3 +627,158 @@ async def test_get_customer_daily_activity_service_account_key_is_rejected(monke
assert exc_info.value.status_code == 401
assert "Admin-only endpoint" in str(exc_info.value.detail)
get_daily_activity_mock.assert_not_called()
# ---------------------------------------------------------------------------
# Characterization (golden-master) tests.
#
# These lock the EXACT JSON body every customer-object endpoint returns today,
# so a type-safety refactor of the handlers is only allowed to land if it
# reproduces these byte for byte. The input below is what a Prisma row's
# .model_dump() yields (full nested budget incl. audit fields + object_permission
# incl. reverse relations); the expected output is what the live endpoint emits.
# ---------------------------------------------------------------------------
_FULL_DB_ROW = {
"user_id": "c1",
"blocked": False,
"alias": "Acme",
"spend": 1.5,
"allowed_model_region": None,
"default_model": None,
"budget_id": "b1",
"object_permission_id": "p1",
"litellm_budget_table": {
"budget_id": "b1",
"max_budget": 10.0,
"soft_budget": None,
"max_parallel_requests": None,
"tpm_limit": None,
"rpm_limit": None,
"model_max_budget": None,
"budget_duration": "30d",
"allowed_models": [],
"budget_reset_at": "2024-02-01T00:00:00",
"created_at": "2024-01-01T00:00:00",
"created_by": "admin",
"updated_at": "2024-01-02T00:00:00",
"updated_by": "admin",
},
"object_permission": {
"object_permission_id": "p1",
"mcp_servers": ["s1"],
"mcp_access_groups": [],
"mcp_tool_permissions": None,
"vector_stores": [],
"agents": [],
"agent_access_groups": [],
"models": [],
"mcp_toolsets": None,
"blocked_tools": [],
"search_tools": [],
"teams": [{"team_id": "t1"}],
"users": [{"user_id": "x"}],
"end_users": [],
"organizations": [],
"verification_tokens": [],
},
}
_EXPECTED_CUSTOMER = {
"user_id": "c1",
"blocked": False,
"alias": "Acme",
"spend": 1.5,
"allowed_model_region": None,
"default_model": None,
"budget_id": "b1",
"litellm_budget_table": {
"budget_id": "b1",
"soft_budget": None,
"max_budget": 10.0,
"max_parallel_requests": None,
"tpm_limit": None,
"rpm_limit": None,
"model_max_budget": None,
"budget_duration": "30d",
"allowed_models": [],
"budget_reset_at": "2024-02-01T00:00:00",
"created_at": "2024-01-01T00:00:00",
},
"object_permission_id": "p1",
"object_permission": {
"object_permission_id": "p1",
"mcp_servers": ["s1"],
"mcp_access_groups": [],
"mcp_tool_permissions": None,
"vector_stores": [],
"agents": [],
"agent_access_groups": [],
"models": [],
"mcp_toolsets": None,
"blocked_tools": [],
"search_tools": [],
"mcp_tool_search_enabled": None,
},
}
def _row(dump: dict) -> MagicMock:
row = MagicMock()
row.model_dump.return_value = dump
return row
def test_char_info_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.get("/customer/info?end_user_id=c1", headers={"Authorization": "Bearer k"})
assert response.status_code == 200
assert response.json() == _EXPECTED_CUSTOMER
def test_char_list_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(return_value=[_row(_FULL_DB_ROW)])
response = client.get("/customer/list", headers={"Authorization": "Bearer k"})
assert response.status_code == 200
assert response.json() == [_EXPECTED_CUSTOMER]
def test_char_new_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.create = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.post("/customer/new", json={"user_id": "c1"}, headers={"Authorization": "Bearer k"})
assert response.status_code == 200
assert response.json() == _EXPECTED_CUSTOMER
def test_char_update_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_first = AsyncMock(
return_value=_row({"user_id": "c1", "blocked": False})
)
mock_prisma_client.db.litellm_endusertable.update = AsyncMock(return_value=_row(_FULL_DB_ROW))
response = client.post(
"/customer/update",
json={"user_id": "c1", "alias": "Acme"},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200
assert response.json() == _EXPECTED_CUSTOMER
def test_char_delete_body(mock_prisma_client, mock_user_api_key_auth):
mock_prisma_client.db.litellm_endusertable.find_many = AsyncMock(
return_value=[
LiteLLM_EndUserTable(user_id="c1", blocked=False),
LiteLLM_EndUserTable(user_id="c2", blocked=False),
]
)
mock_prisma_client.db.litellm_endusertable.delete_many = AsyncMock(return_value=2)
response = client.post(
"/customer/delete",
json={"user_ids": ["c1", "c2"]},
headers={"Authorization": "Bearer k"},
)
assert response.status_code == 200
assert response.json() == {
"deleted_customers": 2,
"message": "Successfully deleted customers with ids: ['c1', 'c2']",
}

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