mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'litellm_internal_staging' into litellm_fix_mcp_streaming_tool_call_leak
This commit is contained in:
commit
1af14ae4bf
292 changed files with 12782 additions and 5012 deletions
32
.github/workflows/create-release.yml
vendored
32
.github/workflows/create-release.yml
vendored
|
|
@ -122,10 +122,28 @@ jobs:
|
|||
makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false";
|
||||
}
|
||||
|
||||
try {
|
||||
await github.rest.git.createRef({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
ref: `refs/tags/${tag}`,
|
||||
sha: commitHash,
|
||||
});
|
||||
} catch (error) {
|
||||
if (error.status !== 422) throw error;
|
||||
const existing = await github.rest.git.getRef({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
ref: `tags/${tag}`,
|
||||
});
|
||||
if (existing.data.object.sha !== commitHash) {
|
||||
throw new Error(`Tag ${tag} already exists at ${existing.data.object.sha}, expected ${commitHash}`);
|
||||
}
|
||||
}
|
||||
|
||||
const response = await github.rest.repos.createRelease({
|
||||
draft: true,
|
||||
generate_release_notes: true,
|
||||
target_commitish: commitHash,
|
||||
name: tag,
|
||||
owner: context.repo.owner,
|
||||
prerelease: isPrerelease,
|
||||
|
|
@ -138,11 +156,21 @@ jobs:
|
|||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
release_id: response.data.id,
|
||||
tag_name: tag,
|
||||
body: updatedBody,
|
||||
draft: false,
|
||||
make_latest: makeLatest,
|
||||
});
|
||||
|
||||
if (!isPrerelease) {
|
||||
await github.rest.repos.updateRelease({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
release_id: response.data.id,
|
||||
tag_name: tag,
|
||||
make_latest: makeLatest,
|
||||
});
|
||||
}
|
||||
|
||||
} catch (error) {
|
||||
core.setFailed(error.message);
|
||||
}
|
||||
|
|
|
|||
2
.gitignore
vendored
2
.gitignore
vendored
|
|
@ -130,3 +130,5 @@ crash.*.log
|
|||
|
||||
# pytest coverage data
|
||||
.coverage
|
||||
|
||||
ui/litellm-dashboard/out/
|
||||
|
|
|
|||
|
|
@ -17,6 +17,8 @@ Same thing for bug fixes. The tests should make it so that this specific bug can
|
|||
|
||||
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
|
||||
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`
|
||||
|
||||
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
|
||||
|
||||
Always use @.github/pull_request_template.md as a guide for your PR body
|
||||
|
|
|
|||
33
Makefile
33
Makefile
|
|
@ -4,7 +4,7 @@
|
|||
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc \
|
||||
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
|
||||
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
|
||||
info lint lint-dev format \
|
||||
info lint lint-dev lint-checks format \
|
||||
lint-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
|
||||
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
|
||||
install-dev install-proxy-dev install-test-deps install-hooks \
|
||||
|
|
@ -53,6 +53,11 @@ help:
|
|||
UV := uv
|
||||
UV_RUN := $(UV) run --no-sync
|
||||
|
||||
LINT_DEP_INSTALL ?= install-dev
|
||||
LINT_DEP_BASE ?= lint-fetch-base
|
||||
LINT_JOBS := $(shell sysctl -n hw.ncpu 2>/dev/null || nproc 2>/dev/null || echo 4)
|
||||
LINT_OUTPUT_SYNC := $(if $(filter output-sync,$(.FEATURES)),--output-sync=target,)
|
||||
|
||||
# Show info
|
||||
info:
|
||||
@echo "UV: $(UV)"
|
||||
|
|
@ -107,12 +112,12 @@ lint-fetch-base:
|
|||
# running proxy need.
|
||||
lint-install:
|
||||
$(UV) sync --inexact --frozen --group proxy-dev
|
||||
$(UV_RUN) prisma generate --schema litellm/proxy/schema.prisma
|
||||
$(UV_RUN) python scripts/prisma_generate_if_needed.py
|
||||
|
||||
# Diff-scoped format check, identical to test-linting.yml's "Check ruff format" step:
|
||||
# only the litellm Python files changed vs the base are checked, so a pre-existing
|
||||
# format issue elsewhere doesn't block an unrelated commit.
|
||||
lint-format-check-changed: install-dev lint-fetch-base
|
||||
lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
@files=$$(git diff --name-only origin/litellm_internal_staging...HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' || true); \
|
||||
if [ -z "$$files" ]; then \
|
||||
echo "No changed litellm Python files to format-check."; \
|
||||
|
|
@ -121,7 +126,7 @@ lint-format-check-changed: install-dev lint-fetch-base
|
|||
fi
|
||||
|
||||
# Linting targets
|
||||
lint-ruff: install-dev
|
||||
lint-ruff: $(LINT_DEP_INSTALL)
|
||||
cd litellm && $(UV_RUN) ruff check . && cd ..
|
||||
|
||||
# faster linter for developing ...
|
||||
|
|
@ -156,12 +161,12 @@ lint-ruff-FULL-dev: install-dev
|
|||
if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \
|
||||
else echo "No changed .py files to check."; fi
|
||||
|
||||
lint-basedpyright: install-dev lint-fetch-base
|
||||
lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
# Type-discipline budget (mutable collections / casts / type guards / kwargs /
|
||||
# unexplained suppressions), the test-linting.yml step `make lint` used to omit.
|
||||
lint-type-discipline: install-dev lint-fetch-base
|
||||
lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
# --update lowers each limit by what this branch fixed since its branch point, so
|
||||
|
|
@ -176,7 +181,7 @@ lint-ruff-budget: install-dev
|
|||
|
||||
# Strict gate, invoked the same way CI does in test-linting.yml so a local pass
|
||||
# means the CI check will pass too.
|
||||
lint-gate: install-dev lint-fetch-base
|
||||
lint-gate: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/ruff_strict_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
lint-ruff-budget-update: install-dev lint-fetch-base
|
||||
|
|
@ -188,10 +193,10 @@ lint-type-discipline-budget-update: install-dev lint-fetch-base
|
|||
# Ratchet all budgets in one shot (ruff strict + type-discipline + basedpyright)
|
||||
lint-budget-update: lint-ruff-budget-update lint-type-discipline-budget-update lint-basedpyright-budget-update
|
||||
|
||||
check-circular-imports: install-dev
|
||||
check-circular-imports: $(LINT_DEP_INSTALL)
|
||||
cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd ..
|
||||
|
||||
check-import-safety: install-dev
|
||||
check-import-safety: $(LINT_DEP_INSTALL)
|
||||
@$(UV_RUN) python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
|
||||
|
||||
# Combined linting, isomorphic to test-linting.yml's lint job so a local pass means a
|
||||
|
|
@ -199,9 +204,13 @@ check-import-safety: install-dev
|
|||
# runs the diff-scoped ruff format check, whole-tree ruff check, the strict-rule /
|
||||
# type-discipline / basedpyright budgets as a delta vs the base, then the circular-import
|
||||
# and import-safety checks. Steps that compare against the base resolve it the same way CI
|
||||
# does (merge-base with origin/litellm_internal_staging). lint-install is first so the
|
||||
# Prisma client exists before basedpyright runs.
|
||||
lint: lint-install lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright check-circular-imports check-import-safety
|
||||
# does (merge-base with origin/litellm_internal_staging). Setup (env sync, Prisma client,
|
||||
# base fetch) runs once up front; the checks themselves are independent, so a sub-make
|
||||
# fans them out with -j and the fast ones finish under basedpyright's shadow.
|
||||
lint: lint-install lint-fetch-base
|
||||
$(MAKE) -j $(LINT_JOBS) $(LINT_OUTPUT_SYNC) LINT_DEP_INSTALL= LINT_DEP_BASE= lint-checks
|
||||
|
||||
lint-checks: lint-format-check-changed lint-ruff lint-gate lint-type-discipline lint-basedpyright check-circular-imports check-import-safety
|
||||
|
||||
# Faster linting for local development (only checks changed code)
|
||||
lint-dev: lint-format-changed check-circular-imports check-import-safety
|
||||
|
|
|
|||
|
|
@ -141,6 +141,6 @@
|
|||
"limit": 1005
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"limit": 1298
|
||||
"limit": 1297
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -57,8 +57,6 @@ source ~/.nvm/nvm.sh
|
|||
nvm install v18.17.0
|
||||
nvm use v18.17.0
|
||||
|
||||
# copy _enterprise.json from this directory to /ui/litellm-dashboard, and rename it to ui_colors.json
|
||||
cp enterprise/enterprise_ui/enterprise_colors.json ui/litellm-dashboard/ui_colors.json
|
||||
|
||||
# cd in to /ui/litellm-dashboard
|
||||
cd ui/litellm-dashboard
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.45"
|
||||
version = "0.1.46"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.45"
|
||||
version = "0.1.46"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN "max_concurrent_requests" INTEGER;
|
||||
|
|
@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
max_concurrent_requests Int?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.74"
|
||||
version = "0.4.75"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.74"
|
||||
version = "0.4.75"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
|
|
@ -588,6 +588,7 @@ gemini_models: Set = set()
|
|||
xai_models: Set = set()
|
||||
zai_models: Set = set()
|
||||
deepseek_models: Set = set()
|
||||
tencent_models: Set = set()
|
||||
runwayml_models: Set = set()
|
||||
azure_ai_models: Set = set()
|
||||
jina_ai_models: Set = set()
|
||||
|
|
@ -801,6 +802,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
|
|||
fal_ai_models.add(key)
|
||||
elif value.get("litellm_provider") == "deepseek":
|
||||
deepseek_models.add(key)
|
||||
elif value.get("litellm_provider") == "tencent":
|
||||
tencent_models.add(key)
|
||||
elif value.get("litellm_provider") == "runwayml":
|
||||
runwayml_models.add(key)
|
||||
elif value.get("litellm_provider") == "meta_llama":
|
||||
|
|
@ -1093,6 +1096,7 @@ models_by_provider: dict = {
|
|||
"zai": zai_models,
|
||||
"fal_ai": fal_ai_models,
|
||||
"deepseek": deepseek_models,
|
||||
"tencent": tencent_models,
|
||||
"runwayml": runwayml_models,
|
||||
"mistral": mistral_chat_models,
|
||||
"azure_ai": azure_ai_models,
|
||||
|
|
@ -1804,6 +1808,9 @@ if TYPE_CHECKING:
|
|||
from .llms.deepseek.chat.transformation import (
|
||||
DeepSeekChatConfig as _DeepSeekChatConfig,
|
||||
)
|
||||
from .llms.tencent.chat.transformation import (
|
||||
TencentChatConfig as _TencentChatConfig,
|
||||
)
|
||||
from .llms.sap.chat.transformation import (
|
||||
GenAIHubOrchestrationConfig as _GenAIHubOrchestrationConfig,
|
||||
)
|
||||
|
|
@ -1846,6 +1853,7 @@ if TYPE_CHECKING:
|
|||
# Type stubs for lazy-loaded config classes (to help mypy understand types)
|
||||
VLLMConfig: Type[_VLLMConfig]
|
||||
DeepSeekChatConfig: Type[_DeepSeekChatConfig]
|
||||
TencentChatConfig: Type[_TencentChatConfig]
|
||||
GenAIHubOrchestrationConfig: Type[_GenAIHubOrchestrationConfig]
|
||||
GenAIHubEmbeddingConfig: Type[_GenAIHubEmbeddingConfig]
|
||||
AzureOpenAIO1Config: Type[_AzureOpenAIO1Config]
|
||||
|
|
|
|||
|
|
@ -284,6 +284,7 @@ LLM_CONFIG_NAMES = (
|
|||
"LiteLLMProxyChatConfig",
|
||||
"VLLMConfig",
|
||||
"DeepSeekChatConfig",
|
||||
"TencentChatConfig",
|
||||
"LMStudioChatConfig",
|
||||
"LmStudioEmbeddingConfig",
|
||||
"NscaleConfig",
|
||||
|
|
@ -1096,6 +1097,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
),
|
||||
"VLLMConfig": (".llms.vllm.completion.transformation", "VLLMConfig"),
|
||||
"DeepSeekChatConfig": (".llms.deepseek.chat.transformation", "DeepSeekChatConfig"),
|
||||
"TencentChatConfig": (".llms.tencent.chat.transformation", "TencentChatConfig"),
|
||||
"LMStudioChatConfig": (".llms.lm_studio.chat.transformation", "LMStudioChatConfig"),
|
||||
"LmStudioEmbeddingConfig": (
|
||||
".llms.lm_studio.embed.transformation",
|
||||
|
|
|
|||
|
|
@ -129,6 +129,33 @@ def _set_agent_id_on_logging_obj(
|
|||
litellm_logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
|
||||
_A2A_COST_PARAM_KEYS = ("cost_per_query", "input_cost_per_token", "output_cost_per_token")
|
||||
|
||||
|
||||
def _set_litellm_params_on_logging_obj(
|
||||
kwargs: dict[str, Any],
|
||||
litellm_params: dict[str, Any],
|
||||
) -> None:
|
||||
"""
|
||||
Merge the agent's pricing params into model_call_details["litellm_params"]
|
||||
so A2ACostCalculator can read them.
|
||||
|
||||
The non-streaming path reuses the proxy-built logging object, whose
|
||||
litellm_params already carries metadata / proxy_server_request / user-key
|
||||
context, so merge the pricing keys in rather than replacing the dict.
|
||||
"""
|
||||
logging_obj = kwargs.get("litellm_logging_obj")
|
||||
if logging_obj is None:
|
||||
return
|
||||
|
||||
cost_params = {key: litellm_params[key] for key in _A2A_COST_PARAM_KEYS if litellm_params.get(key) is not None}
|
||||
if not cost_params:
|
||||
return
|
||||
|
||||
existing = logging_obj.model_call_details.get("litellm_params") or {}
|
||||
logging_obj.model_call_details["litellm_params"] = {**existing, **cost_params}
|
||||
|
||||
|
||||
def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str:
|
||||
"""
|
||||
Extract agent info and set model/custom_llm_provider for cost tracking.
|
||||
|
|
@ -477,6 +504,9 @@ async def asend_message(
|
|||
completion_tokens=completion_tokens,
|
||||
)
|
||||
|
||||
# Merge agent pricing params into the logging obj so cost is calculated
|
||||
_set_litellm_params_on_logging_obj(kwargs=kwargs, litellm_params=litellm_params)
|
||||
|
||||
# Set agent_id on logging obj for SpendLogs tracking
|
||||
_set_agent_id_on_logging_obj(kwargs=kwargs, agent_id=agent_id)
|
||||
|
||||
|
|
|
|||
|
|
@ -121,8 +121,13 @@ class A2ARequestUtils:
|
|||
Returns:
|
||||
Tuple of (prompt_tokens, completion_tokens, total_tokens)
|
||||
"""
|
||||
# Count input tokens
|
||||
# Count input tokens. Dump the message to a dict first so extraction hits
|
||||
# the dict branch — request-side parts are a2a-sdk Part RootModels whose
|
||||
# kind/text live on part.root, which the object branch cannot read. This
|
||||
# mirrors how the response side already works (it operates on model_dump).
|
||||
input_message = A2ARequestUtils.get_input_message_from_request(request)
|
||||
if input_message is not None and hasattr(input_message, "model_dump"):
|
||||
input_message = input_message.model_dump(mode="json")
|
||||
input_text = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens = A2ARequestUtils.count_tokens(input_text)
|
||||
|
||||
|
|
|
|||
|
|
@ -508,6 +508,7 @@ LITELLM_CHAT_PROVIDERS = [
|
|||
"text-completion-codestral",
|
||||
"text-completion-inception",
|
||||
"deepseek",
|
||||
"tencent",
|
||||
"sambanova",
|
||||
"maritalk",
|
||||
"cloudflare",
|
||||
|
|
@ -729,6 +730,7 @@ openai_compatible_providers: List = [
|
|||
"volcengine",
|
||||
"codestral",
|
||||
"deepseek",
|
||||
"tencent",
|
||||
"deepinfra",
|
||||
"perplexity",
|
||||
"xinference",
|
||||
|
|
|
|||
|
|
@ -52,6 +52,9 @@ from litellm.llms.databricks.cost_calculator import (
|
|||
from litellm.llms.deepseek.cost_calculator import (
|
||||
cost_per_token as deepseek_cost_per_token,
|
||||
)
|
||||
from litellm.llms.tencent.cost_calculator import (
|
||||
cost_per_token as tencent_cost_per_token,
|
||||
)
|
||||
from litellm.llms.fireworks_ai.cost_calculator import (
|
||||
cost_per_token as fireworks_ai_cost_per_token,
|
||||
)
|
||||
|
|
@ -625,6 +628,8 @@ def cost_per_token(
|
|||
return gemini_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
|
||||
elif custom_llm_provider == "deepseek":
|
||||
return deepseek_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "tencent":
|
||||
return tencent_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "perplexity":
|
||||
return perplexity_cost_per_token(model=model, usage=usage_block)
|
||||
elif custom_llm_provider == "xai":
|
||||
|
|
|
|||
|
|
@ -1165,12 +1165,18 @@ class ModifyResponseException(Exception):
|
|||
request_data: Dict[str, Any],
|
||||
guardrail_name: Optional[str] = None,
|
||||
detection_info: Optional[Dict[str, Any]] = None,
|
||||
original_response: Optional[Any] = None,
|
||||
):
|
||||
self.message = message
|
||||
self.model = model
|
||||
self.request_data = request_data
|
||||
self.guardrail_name = guardrail_name
|
||||
self.detection_info = detection_info or {}
|
||||
# The LLM response that was blocked (post-call). Carries the real token
|
||||
# usage the upstream call consumed, so the synthetic block response can
|
||||
# report it instead of discarding it. None for pre-call blocks (the LLM
|
||||
# was never invoked).
|
||||
self.original_response = original_response
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timedelta
|
||||
|
|
@ -65,6 +66,26 @@ if TYPE_CHECKING:
|
|||
else:
|
||||
AsyncIOScheduler = Any
|
||||
|
||||
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT = 5.0
|
||||
|
||||
|
||||
def _get_budget_metrics_per_request_timeout() -> float:
|
||||
raw = os.getenv("PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT")
|
||||
if raw is None:
|
||||
return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
|
||||
try:
|
||||
parsed = float(raw)
|
||||
except ValueError:
|
||||
parsed = None
|
||||
if parsed is None or not math.isfinite(parsed) or parsed <= 0:
|
||||
verbose_logger.debug(
|
||||
"[Non-Blocking] Prometheus: invalid PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT=%r; using default %ss.",
|
||||
raw,
|
||||
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT,
|
||||
)
|
||||
return _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
|
||||
return parsed
|
||||
|
||||
|
||||
class PrometheusLogger(CustomLogger):
|
||||
# Class variables or attributes
|
||||
|
|
@ -1607,7 +1628,15 @@ class PrometheusLogger(CustomLogger):
|
|||
_user_spend = _metadata.get("user_api_key_user_spend", None)
|
||||
_user_max_budget = _metadata.get("user_api_key_user_max_budget", None)
|
||||
|
||||
results = await asyncio.gather(
|
||||
# Bound the per-request budget-metric emission so that slow Redis/DB
|
||||
# lookups under load cannot consume the whole LoggingWorker watchdog
|
||||
# (LOGGING_WORKER_MAX_TIME_PER_COROUTINE, default 20s) and get the entire
|
||||
# success-logging event cancelled. Budget gauges are also refreshed by the
|
||||
# periodic cron every PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES,
|
||||
# so dropping one slow per-request emission only loses sub-cron real-time
|
||||
# detail, not correctness.
|
||||
budget_metrics_timeout = _get_budget_metrics_per_request_timeout()
|
||||
gather_coro = asyncio.gather(
|
||||
self._set_api_key_budget_metrics_after_api_request(
|
||||
user_api_key=user_api_key,
|
||||
user_api_key_alias=user_api_key_alias,
|
||||
|
|
@ -1634,6 +1663,16 @@ class PrometheusLogger(CustomLogger):
|
|||
),
|
||||
return_exceptions=True,
|
||||
)
|
||||
try:
|
||||
results = await asyncio.wait_for(gather_coro, timeout=budget_metrics_timeout)
|
||||
except asyncio.TimeoutError:
|
||||
verbose_logger.debug(
|
||||
"[Non-Blocking] Prometheus: per-request budget metric emission "
|
||||
"exceeded %ss under load; skipping (values are refreshed by the "
|
||||
"periodic budget-metrics cron job).",
|
||||
budget_metrics_timeout,
|
||||
)
|
||||
return
|
||||
for i, r in enumerate(results):
|
||||
if isinstance(r, Exception):
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_strip_base64_files: bool = False,
|
||||
s3_use_key_prefix: bool = False,
|
||||
s3_use_virtual_hosted_style: bool = False,
|
||||
s3_server_side_encryption: Optional[str] = None,
|
||||
s3_callback_params_override: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -92,6 +93,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_strip_base64_files=s3_strip_base64_files,
|
||||
s3_use_key_prefix=s3_use_key_prefix,
|
||||
s3_use_virtual_hosted_style=s3_use_virtual_hosted_style,
|
||||
s3_server_side_encryption=s3_server_side_encryption,
|
||||
)
|
||||
verbose_logger.debug(f"s3 logger using endpoint url {s3_endpoint_url}")
|
||||
|
||||
|
|
@ -145,6 +147,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
s3_strip_base64_files: bool = False,
|
||||
s3_use_key_prefix: bool = False,
|
||||
s3_use_virtual_hosted_style: bool = False,
|
||||
s3_server_side_encryption: Optional[str] = None,
|
||||
params_source: Optional[dict] = None,
|
||||
):
|
||||
"""
|
||||
|
|
@ -194,6 +197,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
bool(params.get("s3_use_virtual_hosted_style", False)) or s3_use_virtual_hosted_style
|
||||
)
|
||||
|
||||
self.s3_server_side_encryption = params.get("s3_server_side_encryption") or s3_server_side_encryption
|
||||
|
||||
return
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
|
|
@ -273,6 +278,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
import requests
|
||||
|
|
@ -317,14 +323,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5 = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
# Prepare the request
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**(
|
||||
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
|
||||
if self.s3_server_side_encryption
|
||||
else {}
|
||||
),
|
||||
}
|
||||
req = requests.Request("PUT", url, data=json_string, headers=headers)
|
||||
prepped = req.prepare()
|
||||
|
|
@ -447,6 +462,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
def upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
|
||||
try:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
import requests
|
||||
|
|
@ -482,14 +498,23 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
|
||||
# Calculate SHA256 hash of the content
|
||||
content_hash = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
|
||||
content_md5 = base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
# Prepare the request
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Content-MD5": content_md5,
|
||||
"x-amz-content-sha256": content_hash,
|
||||
"Content-Language": "en",
|
||||
"Content-Disposition": f'inline; filename="{batch_logging_element.s3_object_download_filename}"',
|
||||
"Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
|
||||
**(
|
||||
{"x-amz-server-side-encryption": self.s3_server_side_encryption}
|
||||
if self.s3_server_side_encryption
|
||||
else {}
|
||||
),
|
||||
}
|
||||
req = requests.Request("PUT", url, data=json_string, headers=headers)
|
||||
prepped = req.prepare()
|
||||
|
|
|
|||
|
|
@ -652,6 +652,10 @@ def _get_openai_compatible_provider_info(
|
|||
api_base = api_base or get_secret("DEEPSEEK_API_BASE") or "https://api.deepseek.com/beta" # type: ignore
|
||||
|
||||
dynamic_api_key = api_key or get_secret_str("DEEPSEEK_API_KEY")
|
||||
elif custom_llm_provider == "tencent":
|
||||
api_base = api_base or get_secret("TENCENT_API_BASE") or "https://tokenhub-intl.tencentcloudmaas.com/v1"
|
||||
|
||||
dynamic_api_key = api_key or get_secret_str("TENCENT_API_KEY")
|
||||
elif custom_llm_provider == "fireworks_ai":
|
||||
# fireworks is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.fireworks.ai/inference/v1
|
||||
(
|
||||
|
|
|
|||
|
|
@ -106,6 +106,8 @@ def get_supported_openai_params(
|
|||
return litellm.VLLMConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "deepseek":
|
||||
return litellm.DeepSeekChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "tencent":
|
||||
return litellm.TencentChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere":
|
||||
return litellm.CohereChatConfig().get_supported_openai_params(model=model)
|
||||
elif custom_llm_provider == "maritalk":
|
||||
|
|
|
|||
|
|
@ -1,7 +1,23 @@
|
|||
from typing import Dict, Optional
|
||||
from typing import Any, Dict, Iterator, Optional
|
||||
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
_CLIENT_CALLBACK_METADATA_SLOTS: tuple[str, ...] = ("litellm_metadata", "metadata")
|
||||
|
||||
|
||||
def iter_client_callback_metadata_dicts(
|
||||
kwargs: dict[str, Any],
|
||||
) -> Iterator[tuple[str, dict[str, Any]]]:
|
||||
litellm_params = kwargs.get("litellm_params")
|
||||
if isinstance(litellm_params, dict):
|
||||
nested = litellm_params.get("metadata")
|
||||
if isinstance(nested, dict):
|
||||
yield "litellm_params.metadata", nested
|
||||
for key in _CLIENT_CALLBACK_METADATA_SLOTS:
|
||||
candidate = kwargs.get(key)
|
||||
if isinstance(candidate, dict):
|
||||
yield key, candidate
|
||||
|
||||
|
||||
def _is_env_reference(value: object) -> bool:
|
||||
return isinstance(value, str) and "os.environ/" in value
|
||||
|
|
@ -55,6 +71,7 @@ _supported_callback_params = [
|
|||
"dd_site",
|
||||
"dd_agent_host",
|
||||
"dd_agent_port",
|
||||
"turn_off_message_logging",
|
||||
]
|
||||
|
||||
_request_blocked_callback_params = {
|
||||
|
|
@ -87,19 +104,13 @@ def initialize_standard_callback_dynamic_params(
|
|||
validate_no_callback_env_reference(param, _param_value, source="request body")
|
||||
standard_callback_dynamic_params[param] = _param_value # type: ignore
|
||||
|
||||
# 2. Fallback: check "metadata" or "litellm_params" -> "metadata"
|
||||
metadata = (kwargs.get("metadata") or {}).copy()
|
||||
litellm_params = kwargs.get("litellm_params") or {}
|
||||
if isinstance(litellm_params, dict):
|
||||
metadata.update(litellm_params.get("metadata") or {})
|
||||
|
||||
if isinstance(metadata, dict):
|
||||
for slot_label, metadata in iter_client_callback_metadata_dicts(kwargs):
|
||||
for param in _supported_callback_params:
|
||||
if param in _request_blocked_callback_params:
|
||||
continue
|
||||
if param not in standard_callback_dynamic_params and param in metadata:
|
||||
_param_value = metadata.get(param)
|
||||
validate_no_callback_env_reference(param, _param_value, source="metadata")
|
||||
validate_no_callback_env_reference(param, _param_value, source=slot_label)
|
||||
standard_callback_dynamic_params[param] = _param_value # type: ignore
|
||||
|
||||
return standard_callback_dynamic_params
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ import httpx
|
|||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
|
||||
from ..common_utils import (
|
||||
A2AError,
|
||||
|
|
@ -312,6 +312,25 @@ class A2AConfig(BaseConfig):
|
|||
# Set ID from response
|
||||
model_response.id = response_json.get("id", str(uuid.uuid4()))
|
||||
|
||||
# A2A agents don't return token usage; estimate it so per-token pricing
|
||||
# produces real cost and callers don't receive usage of 0/0/0.
|
||||
try:
|
||||
from litellm.utils import token_counter
|
||||
|
||||
prompt_tokens = token_counter(model="gpt-3.5-turbo", messages=messages)
|
||||
completion_tokens = token_counter(model="gpt-3.5-turbo", text=text, count_response_tokens=True)
|
||||
setattr(
|
||||
model_response,
|
||||
"usage",
|
||||
Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
),
|
||||
)
|
||||
except Exception: # noqa: BLE001 - best-effort estimate; a tokenizer hiccup must not break the response
|
||||
pass
|
||||
|
||||
return model_response
|
||||
|
||||
def get_model_response_iterator(
|
||||
|
|
|
|||
|
|
@ -48,7 +48,10 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
|
|
@ -70,6 +73,170 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
super().__init__()
|
||||
self.adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
@staticmethod
|
||||
def _build_streaming_usage_response(
|
||||
responses_so_far: list[Any],
|
||||
request_data: Optional[dict],
|
||||
) -> Optional[ModelResponse]:
|
||||
chunks = tuple(response for response in responses_so_far if isinstance(response, (str, bytes)))
|
||||
if not chunks:
|
||||
return None
|
||||
try:
|
||||
return AnthropicPassthroughLoggingHandler._build_usage_only_response_from_chunks(
|
||||
all_chunks=chunks,
|
||||
model=str((request_data or {}).get("model") or ""),
|
||||
)
|
||||
except (AttributeError, TypeError, ValueError):
|
||||
return None
|
||||
|
||||
def build_block_sse_chunks(
|
||||
self,
|
||||
exc: "ModifyResponseException",
|
||||
stream_started: bool = False,
|
||||
responses_so_far: Optional[list[Any]] = None,
|
||||
) -> list[bytes]:
|
||||
"""
|
||||
Build an Anthropic SSE sequence delivering the guardrail block message
|
||||
and terminating the stream cleanly.
|
||||
|
||||
- ``stream_started`` False (buffered / pre-stream): nothing has been
|
||||
sent, so emit a complete standalone message (message_start ->
|
||||
content_block_* -> message_delta -> message_stop) via
|
||||
FakeAnthropicMessagesStreamIterator, the same converter the
|
||||
/v1/messages pre-stream block handler uses.
|
||||
- ``stream_started`` True (sampling / detect-only end-of-stream): real
|
||||
chunks were already sent, so *continue* the in-progress message --
|
||||
close the open content block, append the block message as a new text
|
||||
block, then end the message. Emitting a second ``message_start`` here
|
||||
would make Anthropic clients reject the stream.
|
||||
"""
|
||||
if stream_started:
|
||||
return self._block_continuation_chunks(exc, responses_so_far or [])
|
||||
return self._standalone_block_chunks(exc)
|
||||
|
||||
def _standalone_block_chunks(self, exc: "ModifyResponseException") -> list[bytes]:
|
||||
import uuid
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
|
||||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_response_usage,
|
||||
)
|
||||
from litellm.types.utils import AnthropicMessagesResponse
|
||||
|
||||
block_response = AnthropicMessagesResponse(
|
||||
id=f"msg_{uuid.uuid4()}",
|
||||
type="message",
|
||||
role="assistant",
|
||||
content=[{"type": "text", "text": exc.message}],
|
||||
model=exc.model,
|
||||
stop_reason="end_turn",
|
||||
usage=blocked_response_usage(getattr(exc, "original_response", None)),
|
||||
)
|
||||
return list(FakeAnthropicMessagesStreamIterator(response=block_response))
|
||||
|
||||
def _block_continuation_chunks(self, exc: "ModifyResponseException", responses_so_far: list[Any]) -> list[bytes]:
|
||||
"""Continue an already-started message: close the open content block,
|
||||
append the block message as a new text block, then end the message --
|
||||
without a second message_start."""
|
||||
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_response_usage,
|
||||
)
|
||||
|
||||
def _sse(event_type: str, payload: dict) -> bytes:
|
||||
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
|
||||
|
||||
output_tokens = blocked_response_usage(getattr(exc, "original_response", None))["output_tokens"]
|
||||
open_index, max_index = self._content_block_state(responses_so_far)
|
||||
new_index = (max_index + 1) if max_index is not None else 0
|
||||
chunks: list[bytes] = []
|
||||
if open_index is not None:
|
||||
chunks.append(_sse("content_block_stop", {"type": "content_block_stop", "index": open_index}))
|
||||
chunks += [
|
||||
_sse(
|
||||
"content_block_start",
|
||||
{
|
||||
"type": "content_block_start",
|
||||
"index": new_index,
|
||||
"content_block": {"type": "text", "text": ""},
|
||||
},
|
||||
),
|
||||
_sse(
|
||||
"content_block_delta",
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": new_index,
|
||||
"delta": {"type": "text_delta", "text": exc.message},
|
||||
},
|
||||
),
|
||||
_sse("content_block_stop", {"type": "content_block_stop", "index": new_index}),
|
||||
_sse(
|
||||
"message_delta",
|
||||
{
|
||||
"type": "message_delta",
|
||||
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
|
||||
"usage": {"output_tokens": output_tokens},
|
||||
},
|
||||
),
|
||||
_sse("message_stop", {"type": "message_stop"}),
|
||||
]
|
||||
return chunks
|
||||
|
||||
@staticmethod
|
||||
def _content_block_state(
|
||||
responses_so_far: list[Any],
|
||||
) -> tuple[Optional[int], Optional[int]]:
|
||||
"""From the SSE chunks already sent to the client, return (open
|
||||
content-block index or None, highest content-block index seen or None).
|
||||
|
||||
A single streamed item may bundle multiple SSE events (raw bytes) or be
|
||||
an already-parsed event dict, so every event across every item is
|
||||
considered -- matching how ``get_streaming_string_so_far`` reads the
|
||||
same stream."""
|
||||
open_indices: set[int] = set()
|
||||
max_index: Optional[int] = None
|
||||
for item in responses_so_far:
|
||||
for data in AnthropicMessagesHandler._iter_sse_events(item):
|
||||
event_type = data.get("type")
|
||||
index = data.get("index")
|
||||
if not isinstance(index, int):
|
||||
continue
|
||||
if event_type == "content_block_start":
|
||||
open_indices.add(index)
|
||||
max_index = index if max_index is None else max(max_index, index)
|
||||
elif event_type == "content_block_stop":
|
||||
open_indices.discard(index)
|
||||
open_index = max(open_indices) if open_indices else None
|
||||
return open_index, max_index
|
||||
|
||||
@staticmethod
|
||||
def _iter_sse_events(item: Any) -> list[dict]:
|
||||
"""Yield the event-data dicts in one stream chunk.
|
||||
|
||||
Handles both formats this stream can carry (see
|
||||
``get_streaming_string_so_far``): raw SSE ``bytes`` -- which may bundle
|
||||
several events separated by a blank line -- and an already-parsed event
|
||||
``dict``."""
|
||||
if isinstance(item, dict):
|
||||
return [item]
|
||||
if not isinstance(item, (bytes, bytearray)):
|
||||
return []
|
||||
events: list[dict] = []
|
||||
for block in item.decode("utf-8", errors="replace").split("\n\n"):
|
||||
for line in block.split("\n"):
|
||||
line = line.strip()
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
try:
|
||||
parsed = json.loads(line[len("data:") :].strip())
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(parsed, dict):
|
||||
events.append(parsed)
|
||||
return events
|
||||
|
||||
def _translate_to_openai(self, data: dict) -> ChatCompletionRequest:
|
||||
"""Translate Anthropic request to OpenAI chat completion format."""
|
||||
(
|
||||
|
|
@ -406,6 +573,8 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
Get the string so far, check the apply guardrail to the string so far, and return the list of responses so far.
|
||||
"""
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
has_ended = self._check_streaming_has_ended(responses_so_far)
|
||||
if has_ended:
|
||||
# build the model response from the responses_so_far
|
||||
|
|
@ -430,25 +599,35 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if tool_calls_list:
|
||||
guardrail_inputs["tool_calls"] = tool_calls_list
|
||||
|
||||
_guardrailed_inputs = (
|
||||
await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
|
||||
try:
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=guardrail_inputs,
|
||||
request_data=request_data if request_data is not None else {},
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = built_response or self._build_streaming_usage_response(
|
||||
responses_so_far, request_data
|
||||
)
|
||||
raise
|
||||
else:
|
||||
verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices")
|
||||
return responses_so_far
|
||||
|
||||
string_so_far = self.get_streaming_string_so_far(responses_so_far)
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
|
||||
inputs={"texts": [string_so_far]},
|
||||
request_data=request_data if request_data is not None else {},
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
try:
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [string_so_far]},
|
||||
request_data=request_data if request_data is not None else {},
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = self._build_streaming_usage_response(responses_so_far, request_data)
|
||||
raise
|
||||
return responses_so_far
|
||||
|
||||
def _prepare_request_data(
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from typing import Any, Dict
|
|||
from urllib.parse import quote
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
|
||||
|
|
@ -38,6 +39,30 @@ from litellm.secret_managers.main import get_secret_str
|
|||
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceLine(BaseModel):
|
||||
content: str | None = None
|
||||
|
||||
|
||||
class AzureDocumentIntelligencePage(BaseModel):
|
||||
pageNumber: int | None = None
|
||||
width: float | None = None
|
||||
height: float | None = None
|
||||
unit: str | None = None
|
||||
lines: tuple[AzureDocumentIntelligenceLine, ...] = ()
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceAnalyzeResult(BaseModel):
|
||||
content: str | None = None
|
||||
pages: tuple[AzureDocumentIntelligencePage, ...] = ()
|
||||
tables: list[dict[str, object]] | None = None
|
||||
keyValuePairs: list[dict[str, object]] | None = None
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceOperation(BaseModel):
|
||||
status: str | None = None
|
||||
analyzeResult: AzureDocumentIntelligenceAnalyzeResult | None = None
|
||||
|
||||
|
||||
class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
||||
"""
|
||||
Azure Document Intelligence OCR transformation configuration.
|
||||
|
|
@ -67,11 +92,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
(1-based, e.g. "1-3,5,7-9"). To keep the public request shape
|
||||
aligned with Mistral OCR, callers pass `pages` using Mistral
|
||||
semantics — a list of 0-based integers — or a pre-formatted
|
||||
Azure-style string. Other Mistral-specific params (e.g.
|
||||
Azure-style string. Azure DI also exposes a `features` query
|
||||
parameter enabling add-on capabilities (e.g. "keyValuePairs",
|
||||
"languages"), passed as a list of feature names or a
|
||||
comma-separated string. Other Mistral-specific params (e.g.
|
||||
`include_image_base64`) are not supported by Azure DI and are
|
||||
ignored during transformation.
|
||||
"""
|
||||
return ["pages"]
|
||||
return ["pages", "features"]
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
|
|
@ -85,16 +113,18 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
Translates Mistral-style `pages` (list[int], 0-based) into Azure's
|
||||
`pages` query string (1-based, e.g. "1,2,3" or "1-3,5"). A raw
|
||||
string that already matches Azure's format is passed through
|
||||
unchanged.
|
||||
unchanged. `features` (list[str] or comma-separated string) is
|
||||
normalized into Azure's comma-joined `features` query string.
|
||||
"""
|
||||
pages = non_default_params.get("pages")
|
||||
if pages is None:
|
||||
return optional_params
|
||||
|
||||
normalized = self._normalize_pages_param(pages)
|
||||
if normalized:
|
||||
optional_params["pages"] = normalized
|
||||
return optional_params
|
||||
features = non_default_params.get("features")
|
||||
normalized_pages = self._normalize_pages_param(pages) if pages is not None else ""
|
||||
normalized_features = self._normalize_features_param(features) if features is not None else ""
|
||||
return {
|
||||
**optional_params,
|
||||
**({"pages": normalized_pages} if normalized_pages else {}),
|
||||
**({"features": normalized_features} if normalized_features else {}),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _normalize_pages_param(pages: Any) -> str:
|
||||
|
|
@ -140,6 +170,39 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
|
||||
raise ValueError("`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'.")
|
||||
|
||||
@staticmethod
|
||||
def _normalize_features_param(features: object) -> str:
|
||||
"""
|
||||
Convert a caller-provided `features` value to Azure DI's query-string
|
||||
form (comma-joined feature names, e.g. "keyValuePairs,languages").
|
||||
|
||||
Accepted inputs:
|
||||
- list[str]: feature names like ["keyValuePairs", "languages"].
|
||||
- str: a single feature name or comma-separated names.
|
||||
"""
|
||||
invalid_features_error = ValueError(
|
||||
f"Invalid `features` for Azure Document Intelligence: {features!r}. "
|
||||
f"Expected a list of feature names or a comma-separated string like "
|
||||
f"'keyValuePairs' or 'keyValuePairs,languages'."
|
||||
)
|
||||
|
||||
if isinstance(features, str):
|
||||
raw_tokens = features.split(",")
|
||||
elif isinstance(features, list):
|
||||
if len(features) == 0:
|
||||
return ""
|
||||
raw_tokens = [feature for feature in features if isinstance(feature, str)]
|
||||
if len(raw_tokens) != len(features):
|
||||
raise invalid_features_error
|
||||
else:
|
||||
raise invalid_features_error
|
||||
|
||||
tokens = tuple(token.strip() for token in raw_tokens)
|
||||
feature_pattern = re.compile(r"^[A-Za-z][A-Za-z0-9]*$")
|
||||
if not all(feature_pattern.match(token) for token in tokens):
|
||||
raise invalid_features_error
|
||||
return ",".join(tokens)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Dict,
|
||||
|
|
@ -228,13 +291,15 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
|
||||
)
|
||||
|
||||
# Azure DI accepts `pages` as a query param (1-based, e.g. "1-3,5").
|
||||
# Azure DI accepts `pages` (1-based, e.g. "1-3,5") and `features`
|
||||
# (comma-joined names, e.g. "keyValuePairs") as query params.
|
||||
# `optional_params` has already been normalized in `map_ocr_params`.
|
||||
pages = optional_params.get("pages") if optional_params else None
|
||||
if pages:
|
||||
url += f"&pages={quote(str(pages), safe=',-')}"
|
||||
features = optional_params.get("features") if optional_params else None
|
||||
pages_query = f"&pages={quote(str(pages), safe=',-')}" if pages else ""
|
||||
features_query = f"&features={quote(str(features), safe=',')}" if features else ""
|
||||
|
||||
return url
|
||||
return f"{url}{pages_query}{features_query}"
|
||||
|
||||
def _extract_base64_from_data_uri(self, data_uri: str) -> str:
|
||||
"""
|
||||
|
|
@ -328,27 +393,15 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
|
||||
return OCRRequestData(data=data, files=None)
|
||||
|
||||
def _extract_page_markdown(self, page_data: Dict[str, Any]) -> str:
|
||||
"""
|
||||
Extract text from Azure DI page and format as markdown.
|
||||
|
||||
Azure DI provides text in 'lines' array. We concatenate them with newlines.
|
||||
|
||||
Args:
|
||||
page_data: Azure DI page object
|
||||
|
||||
Returns:
|
||||
Markdown-formatted text
|
||||
"""
|
||||
lines = page_data.get("lines", [])
|
||||
if not lines:
|
||||
return ""
|
||||
|
||||
# Extract text content from each line
|
||||
text_lines = [line.get("content", "") for line in lines]
|
||||
|
||||
# Join with newlines to preserve structure
|
||||
return "\n".join(text_lines)
|
||||
def _transform_azure_page(self, azure_page: AzureDocumentIntelligencePage) -> OCRPage:
|
||||
page_number = azure_page.pageNumber if azure_page.pageNumber is not None else 1
|
||||
markdown = "\n".join(line.content or "" for line in azure_page.lines)
|
||||
dimensions = self._convert_dimensions(
|
||||
width=azure_page.width if azure_page.width is not None else 8.5,
|
||||
height=azure_page.height if azure_page.height is not None else 11,
|
||||
unit=azure_page.unit if azure_page.unit is not None else "inch",
|
||||
)
|
||||
return OCRPage(index=page_number - 1, markdown=markdown, dimensions=dimensions)
|
||||
|
||||
def _convert_dimensions(self, width: float, height: float, unit: str) -> OCRPageDimensions:
|
||||
"""
|
||||
|
|
@ -526,6 +579,52 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
retry_after = self._get_retry_after(response=response)
|
||||
await asyncio.sleep(retry_after)
|
||||
|
||||
def _get_polling_target(self, raw_response: httpx.Response) -> tuple[str, Dict[str, str]]:
|
||||
operation_url = raw_response.headers.get("Operation-Location")
|
||||
if not operation_url:
|
||||
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
|
||||
|
||||
# Reject cross-origin polling URLs — the auth headers
|
||||
# below would otherwise leak to whatever URL the upstream
|
||||
# (or an attacker-controlled upstream) returns. VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_url, str(raw_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
|
||||
|
||||
poll_headers = {"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")}
|
||||
return operation_url, poll_headers
|
||||
|
||||
def _transform_completed_response(self, model: str, raw_response: httpx.Response) -> OCRResponse:
|
||||
"""
|
||||
Transform a completed Azure Document Intelligence analyze operation
|
||||
into the Mistral OCR response shape, preserving Azure-native
|
||||
`analyzeResult` fields (`content`, `tables`, `keyValuePairs`) as
|
||||
top-level response fields.
|
||||
"""
|
||||
operation = AzureDocumentIntelligenceOperation.model_validate(raw_response.json())
|
||||
|
||||
verbose_logger.debug(f"Azure Document Intelligence response status: {operation.status}")
|
||||
|
||||
if operation.status != "succeeded":
|
||||
raise ValueError(f"Azure Document Intelligence analysis failed with status: {operation.status}")
|
||||
|
||||
analyze_result = (
|
||||
operation.analyzeResult if operation.analyzeResult is not None else AzureDocumentIntelligenceAnalyzeResult()
|
||||
)
|
||||
mistral_pages = [self._transform_azure_page(azure_page) for azure_page in analyze_result.pages]
|
||||
usage_info = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
|
||||
|
||||
return OCRResponse(
|
||||
pages=mistral_pages,
|
||||
model=model,
|
||||
usage_info=usage_info,
|
||||
object="ocr",
|
||||
content=analyze_result.content,
|
||||
tables=analyze_result.tables,
|
||||
keyValuePairs=analyze_result.keyValuePairs,
|
||||
)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -552,11 +651,13 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
"unit": "inch",
|
||||
"lines": [{"content": "text", "boundingBox": [...]}]
|
||||
}
|
||||
]
|
||||
],
|
||||
"tables": [...],
|
||||
"keyValuePairs": [...]
|
||||
}
|
||||
}
|
||||
|
||||
Mistral OCR format:
|
||||
Mistral OCR format (with Azure-native fields preserved):
|
||||
{
|
||||
"pages": [
|
||||
{
|
||||
|
|
@ -567,7 +668,10 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
],
|
||||
"model": "azure_ai/doc-intelligence/prebuilt-layout",
|
||||
"usage_info": {"pages_processed": 1},
|
||||
"object": "ocr"
|
||||
"object": "ocr",
|
||||
"content": "Full document text...",
|
||||
"tables": [...],
|
||||
"keyValuePairs": [...]
|
||||
}
|
||||
|
||||
Args:
|
||||
|
|
@ -578,86 +682,17 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
Returns:
|
||||
OCRResponse in Mistral format
|
||||
"""
|
||||
try:
|
||||
# Check if we got 202 Accepted (async operation started)
|
||||
if raw_response.status_code == 202:
|
||||
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
|
||||
if raw_response.status_code != 202:
|
||||
return self._transform_completed_response(model=model, raw_response=raw_response)
|
||||
|
||||
# Get Operation-Location header
|
||||
operation_url = raw_response.headers.get("Operation-Location")
|
||||
if not operation_url:
|
||||
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
|
||||
|
||||
# Reject cross-origin polling URLs — the auth headers
|
||||
# below would otherwise leak to whatever URL the upstream
|
||||
# (or an attacker-controlled upstream) returns. VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_url, str(raw_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
|
||||
|
||||
# Get headers for polling (need auth)
|
||||
poll_headers = {
|
||||
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")
|
||||
}
|
||||
|
||||
# Get timeout from kwargs or use default
|
||||
timeout_secs = AZURE_OPERATION_POLLING_TIMEOUT
|
||||
|
||||
# Poll until operation completes
|
||||
raw_response = self._poll_operation_sync(
|
||||
operation_url=operation_url,
|
||||
headers=poll_headers,
|
||||
timeout_secs=timeout_secs,
|
||||
)
|
||||
|
||||
# Now parse the completed response
|
||||
response_json = raw_response.json()
|
||||
|
||||
verbose_logger.debug(f"Azure Document Intelligence response status: {response_json.get('status')}")
|
||||
|
||||
# Check if request succeeded
|
||||
status = response_json.get("status")
|
||||
if status != "succeeded":
|
||||
raise ValueError(f"Azure Document Intelligence analysis failed with status: {status}")
|
||||
|
||||
# Extract analyze result
|
||||
analyze_result = response_json.get("analyzeResult", {})
|
||||
azure_pages = analyze_result.get("pages", [])
|
||||
|
||||
# Transform pages to Mistral format
|
||||
mistral_pages = []
|
||||
for azure_page in azure_pages:
|
||||
page_number = azure_page.get("pageNumber", 1)
|
||||
index = page_number - 1 # Convert to 0-based index
|
||||
|
||||
# Extract markdown text
|
||||
markdown = self._extract_page_markdown(azure_page)
|
||||
|
||||
# Convert dimensions
|
||||
width = azure_page.get("width", 8.5)
|
||||
height = azure_page.get("height", 11)
|
||||
unit = azure_page.get("unit", "inch")
|
||||
dimensions = self._convert_dimensions(width=width, height=height, unit=unit)
|
||||
|
||||
# Build OCR page
|
||||
ocr_page = OCRPage(index=index, markdown=markdown, dimensions=dimensions)
|
||||
mistral_pages.append(ocr_page)
|
||||
|
||||
# Build usage info
|
||||
usage_info = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
|
||||
|
||||
# Return Mistral OCR response
|
||||
return OCRResponse(
|
||||
pages=mistral_pages,
|
||||
model=model,
|
||||
usage_info=usage_info,
|
||||
object="ocr",
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error parsing Azure Document Intelligence response: {e}")
|
||||
raise e
|
||||
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation...")
|
||||
operation_url, poll_headers = self._get_polling_target(raw_response)
|
||||
completed_response = self._poll_operation_sync(
|
||||
operation_url=operation_url,
|
||||
headers=poll_headers,
|
||||
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
|
||||
)
|
||||
return self._transform_completed_response(model=model, raw_response=completed_response)
|
||||
|
||||
async def async_transform_ocr_response(
|
||||
self,
|
||||
|
|
@ -680,81 +715,14 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
Returns:
|
||||
OCRResponse in Mistral format
|
||||
"""
|
||||
try:
|
||||
# Check if we got 202 Accepted (async operation started)
|
||||
if raw_response.status_code == 202:
|
||||
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
|
||||
if raw_response.status_code != 202:
|
||||
return self._transform_completed_response(model=model, raw_response=raw_response)
|
||||
|
||||
# Get Operation-Location header
|
||||
operation_url = raw_response.headers.get("Operation-Location")
|
||||
if not operation_url:
|
||||
raise ValueError("Azure Document Intelligence returned 202 but no Operation-Location header found")
|
||||
|
||||
# Reject cross-origin polling URLs (see sync path). VERIA-51.
|
||||
try:
|
||||
assert_same_origin(operation_url, str(raw_response.request.url))
|
||||
except SSRFError as ssrf_err:
|
||||
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
|
||||
|
||||
# Get headers for polling (need auth)
|
||||
poll_headers = {
|
||||
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")
|
||||
}
|
||||
|
||||
# Get timeout from kwargs or use default
|
||||
timeout_secs = AZURE_OPERATION_POLLING_TIMEOUT
|
||||
|
||||
# Poll until operation completes (async)
|
||||
raw_response = await self._poll_operation_async(
|
||||
operation_url=operation_url,
|
||||
headers=poll_headers,
|
||||
timeout_secs=timeout_secs,
|
||||
)
|
||||
|
||||
# Now parse the completed response
|
||||
response_json = raw_response.json()
|
||||
|
||||
verbose_logger.debug(f"Azure Document Intelligence response status: {response_json.get('status')}")
|
||||
|
||||
# Check if request succeeded
|
||||
status = response_json.get("status")
|
||||
if status != "succeeded":
|
||||
raise ValueError(f"Azure Document Intelligence analysis failed with status: {status}")
|
||||
|
||||
# Extract analyze result
|
||||
analyze_result = response_json.get("analyzeResult", {})
|
||||
azure_pages = analyze_result.get("pages", [])
|
||||
|
||||
# Transform pages to Mistral format
|
||||
mistral_pages = []
|
||||
for azure_page in azure_pages:
|
||||
page_number = azure_page.get("pageNumber", 1)
|
||||
index = page_number - 1 # Convert to 0-based index
|
||||
|
||||
# Extract markdown text
|
||||
markdown = self._extract_page_markdown(azure_page)
|
||||
|
||||
# Convert dimensions
|
||||
width = azure_page.get("width", 8.5)
|
||||
height = azure_page.get("height", 11)
|
||||
unit = azure_page.get("unit", "inch")
|
||||
dimensions = self._convert_dimensions(width=width, height=height, unit=unit)
|
||||
|
||||
# Build OCR page
|
||||
ocr_page = OCRPage(index=index, markdown=markdown, dimensions=dimensions)
|
||||
mistral_pages.append(ocr_page)
|
||||
|
||||
# Build usage info
|
||||
usage_info = OCRUsageInfo(pages_processed=len(mistral_pages), doc_size_bytes=None)
|
||||
|
||||
# Return Mistral OCR response
|
||||
return OCRResponse(
|
||||
pages=mistral_pages,
|
||||
model=model,
|
||||
usage_info=usage_info,
|
||||
object="ocr",
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error parsing Azure Document Intelligence response (async): {e}")
|
||||
raise e
|
||||
verbose_logger.debug("Azure DI returned 202 Accepted, polling operation (async)...")
|
||||
operation_url, poll_headers = self._get_polling_target(raw_response)
|
||||
completed_response = await self._poll_operation_async(
|
||||
operation_url=operation_url,
|
||||
headers=poll_headers,
|
||||
timeout_secs=AZURE_OPERATION_POLLING_TIMEOUT,
|
||||
)
|
||||
return self._transform_completed_response(model=model, raw_response=completed_response)
|
||||
|
|
|
|||
|
|
@ -2,7 +2,10 @@ from abc import ABC, abstractmethod
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
ModifyResponseException,
|
||||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -98,6 +101,30 @@ class BaseTranslation(ABC):
|
|||
"""
|
||||
return responses_so_far
|
||||
|
||||
def build_block_sse_chunks(
|
||||
self,
|
||||
exc: "ModifyResponseException",
|
||||
stream_started: bool = False,
|
||||
responses_so_far: Optional[list[Any]] = None,
|
||||
) -> Optional[list[bytes]]:
|
||||
"""
|
||||
Build the streaming chunks that deliver a guardrail block message and
|
||||
cleanly terminate the stream in this provider's wire format.
|
||||
|
||||
``stream_started`` is True when real chunks were already sent to the
|
||||
client: the result must *continue* the in-progress message (e.g. close
|
||||
the open content block and append the block message) rather than start
|
||||
a new one, which clients reject. ``responses_so_far`` provides the prior
|
||||
chunks needed to do so. When False, nothing has been sent and a
|
||||
standalone block message is emitted.
|
||||
|
||||
Returns None when the format has no safe terminator; the caller then
|
||||
re-raises ``exc`` so the proxy can surface a clean error instead.
|
||||
Override in provider subclasses that support synthesizing a block
|
||||
stream.
|
||||
"""
|
||||
return None
|
||||
|
||||
def get_structured_messages(self, data: dict) -> Optional[List["AllMessageValues"]]:
|
||||
"""
|
||||
Convert request data to OpenAI-spec structured messages.
|
||||
|
|
|
|||
|
|
@ -1,10 +1,100 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Any, List
|
||||
import json
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
def _anthropic_stream_chunk_events(item: Any) -> list[dict]:
|
||||
if isinstance(item, dict):
|
||||
return [item]
|
||||
if isinstance(item, bytes):
|
||||
chunk = item.decode("utf-8", errors="replace")
|
||||
elif isinstance(item, str):
|
||||
chunk = item
|
||||
else:
|
||||
return []
|
||||
|
||||
events: list[dict] = []
|
||||
for block in chunk.split("\n\n"):
|
||||
for line in block.splitlines():
|
||||
stripped = line.strip()
|
||||
if not stripped.startswith("data:"):
|
||||
continue
|
||||
payload = stripped[len("data:") :].strip()
|
||||
if not payload or payload == "[DONE]":
|
||||
continue
|
||||
try:
|
||||
parsed = json.loads(payload)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
if isinstance(parsed, dict):
|
||||
events.append(parsed)
|
||||
return events
|
||||
|
||||
|
||||
def _usage_from_anthropic_stream_chunks(original_response: list[Any]) -> Optional[AnthropicUsage]:
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
found_usage = False
|
||||
|
||||
for item in original_response:
|
||||
for event in _anthropic_stream_chunk_events(item):
|
||||
event_type = event.get("type")
|
||||
if event_type == "message_start":
|
||||
message = event.get("message") or {}
|
||||
usage_obj = message.get("usage") or {}
|
||||
elif event_type == "message_delta":
|
||||
usage_obj = event.get("usage") or {}
|
||||
else:
|
||||
usage_obj = {}
|
||||
if not isinstance(usage_obj, dict):
|
||||
continue
|
||||
if usage_obj.get("input_tokens") is not None:
|
||||
input_tokens = int(usage_obj.get("input_tokens") or 0)
|
||||
found_usage = True
|
||||
if usage_obj.get("output_tokens") is not None:
|
||||
output_tokens = int(usage_obj.get("output_tokens") or 0)
|
||||
found_usage = True
|
||||
|
||||
if not found_usage:
|
||||
return None
|
||||
return AnthropicUsage(input_tokens=input_tokens, output_tokens=output_tokens)
|
||||
|
||||
|
||||
def blocked_response_usage(original_response: Optional[Any]) -> AnthropicUsage:
|
||||
"""
|
||||
Token usage for a synthetic guardrail-blocked response.
|
||||
|
||||
A post-call block replaces the LLM's response with the violation message,
|
||||
but the upstream call already consumed tokens -- report that real usage
|
||||
(carried on ``ModifyResponseException.original_response``) rather than
|
||||
discarding it. Pre-call blocks never invoked the LLM (no original_response),
|
||||
so usage is zero.
|
||||
"""
|
||||
usage_obj: Any = None
|
||||
if isinstance(original_response, list):
|
||||
stream_usage = _usage_from_anthropic_stream_chunks(original_response)
|
||||
if stream_usage is not None:
|
||||
return stream_usage
|
||||
elif isinstance(original_response, dict):
|
||||
usage_obj = original_response.get("usage")
|
||||
elif original_response is not None:
|
||||
usage_obj = getattr(original_response, "usage", None)
|
||||
|
||||
def _tokens(key: str, fallback_key: str) -> int:
|
||||
if isinstance(usage_obj, dict):
|
||||
return int(usage_obj.get(key, usage_obj.get(fallback_key, 0)) or 0)
|
||||
return int(getattr(usage_obj, key, getattr(usage_obj, fallback_key, 0)) or 0)
|
||||
|
||||
return AnthropicUsage(
|
||||
input_tokens=_tokens("input_tokens", "prompt_tokens"),
|
||||
output_tokens=_tokens("output_tokens", "completion_tokens"),
|
||||
)
|
||||
|
||||
|
||||
def effective_skip_system_message_for_guardrail(guardrail_to_apply: Any) -> bool:
|
||||
per = getattr(guardrail_to_apply, "skip_system_message_in_guardrail", None)
|
||||
if per is not None:
|
||||
|
|
|
|||
|
|
@ -70,6 +70,9 @@ class OCRResponse(LiteLLMPydanticObjectBase):
|
|||
model: str
|
||||
document_annotation: Any | None = None
|
||||
usage_info: OCRUsageInfo | None = None
|
||||
content: str | None = None
|
||||
tables: list[dict[str, object]] | None = None
|
||||
keyValuePairs: list[dict[str, object]] | None = None
|
||||
object: str = "ocr"
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
|
|
|||
|
|
@ -76,6 +76,7 @@ from litellm.utils import (
|
|||
from ..common_utils import (
|
||||
BedrockError,
|
||||
BedrockModelInfo,
|
||||
bedrock_converse_supports_parallel_tool_use_config,
|
||||
get_anthropic_beta_from_headers,
|
||||
get_bedrock_tool_name,
|
||||
is_claude_4_5_on_bedrock,
|
||||
|
|
@ -1106,18 +1107,28 @@ class AmazonConverseConfig(BaseConfig):
|
|||
if cache_control is None:
|
||||
return None
|
||||
|
||||
cache_point = CachePointBlock(type="default")
|
||||
if isinstance(cache_control, dict) and "ttl" in cache_control:
|
||||
ttl = cache_control["ttl"]
|
||||
if ttl in ["5m", "1h"] and model is not None:
|
||||
if is_claude_4_5_on_bedrock(model):
|
||||
cache_point["ttl"] = ttl
|
||||
cache_point = self._build_cache_point_block(cache_control, model)
|
||||
|
||||
if block_type == "system":
|
||||
return SystemContentBlock(cachePoint=cache_point)
|
||||
else:
|
||||
return ContentBlock(cachePoint=cache_point)
|
||||
|
||||
@staticmethod
|
||||
def _build_cache_point_block(control: Optional[dict], model: Optional[str] = None) -> CachePointBlock:
|
||||
"""Build a Bedrock ``cachePoint`` block from an OpenAI-style ``cache_control``/``control`` dict.
|
||||
|
||||
``type`` is always ``"default"`` (the only value Bedrock's Converse API
|
||||
accepts). ``ttl`` is only honored for models that support extended TTL
|
||||
caching (Claude 4.5 family on Bedrock).
|
||||
"""
|
||||
cache_point = CachePointBlock(type="default")
|
||||
if isinstance(control, dict) and "ttl" in control:
|
||||
ttl = control["ttl"]
|
||||
if ttl in ["5m", "1h"] and model is not None and is_claude_4_5_on_bedrock(model):
|
||||
cache_point["ttl"] = ttl
|
||||
return cache_point
|
||||
|
||||
def _transform_system_message(
|
||||
self, messages: List[AllMessageValues], model: Optional[str] = None
|
||||
) -> Tuple[List[AllMessageValues], List[SystemContentBlock]]:
|
||||
|
|
@ -1241,7 +1252,7 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
# Handle parallel_tool_calls configuration
|
||||
parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None)
|
||||
if parallel_tool_use_config is not None and is_claude_4_5_on_bedrock(model):
|
||||
if parallel_tool_use_config is not None and bedrock_converse_supports_parallel_tool_use_config(model):
|
||||
for key, value in parallel_tool_use_config.items():
|
||||
if (
|
||||
key in additional_request_params
|
||||
|
|
@ -1526,7 +1537,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
if cache_injection_points and len(bedrock_tools) > 0:
|
||||
for point in cache_injection_points:
|
||||
if point.get("location") == "tool_config":
|
||||
bedrock_tools.append({"cachePoint": {"type": "default"}})
|
||||
cache_point = self._build_cache_point_block(point.get("control"), model)
|
||||
bedrock_tools.append(ToolBlock(cachePoint=cache_point))
|
||||
break
|
||||
|
||||
bedrock_tool_config: Optional[ToolConfigBlock] = None
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ from functools import partial
|
|||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast, get_args
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
|
@ -24,6 +25,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
HTTPHandler,
|
||||
_get_httpx_client,
|
||||
)
|
||||
from litellm.types.llms.bedrock import GuardrailConfigBlock
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
|
@ -37,6 +39,38 @@ else:
|
|||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
_GUARDRAIL_CONFIG_VALIDATOR: "TypeAdapter[GuardrailConfigBlock]" = TypeAdapter(GuardrailConfigBlock)
|
||||
|
||||
_GUARDRAIL_CONFIG_EXPECTED_FORMAT = (
|
||||
"{'guardrailIdentifier': str, 'guardrailVersion': str, 'trace': 'enabled'|'disabled'|'enabled_full'}"
|
||||
)
|
||||
|
||||
|
||||
def _bedrock_invoke_guardrail_headers(raw_guardrail_config: object) -> "dict[str, str]":
|
||||
try:
|
||||
guardrail_config = _GUARDRAIL_CONFIG_VALIDATOR.validate_python(raw_guardrail_config)
|
||||
except ValidationError as e:
|
||||
raise BedrockError(
|
||||
status_code=400,
|
||||
message="Invalid guardrailConfig={}. Expected format: {}. Error: {}".format(
|
||||
raw_guardrail_config, _GUARDRAIL_CONFIG_EXPECTED_FORMAT, e
|
||||
),
|
||||
)
|
||||
if "guardrailIdentifier" not in guardrail_config:
|
||||
raise BedrockError(
|
||||
status_code=400,
|
||||
message="guardrailConfig={} is missing 'guardrailIdentifier'. Expected format: {}".format(
|
||||
raw_guardrail_config, _GUARDRAIL_CONFIG_EXPECTED_FORMAT
|
||||
),
|
||||
)
|
||||
trace = guardrail_config.get("trace")
|
||||
candidate_headers = {
|
||||
"X-Amzn-Bedrock-GuardrailIdentifier": guardrail_config.get("guardrailIdentifier"),
|
||||
"X-Amzn-Bedrock-GuardrailVersion": guardrail_config.get("guardrailVersion"),
|
||||
"X-Amzn-Bedrock-Trace": trace.upper() if trace is not None else None,
|
||||
}
|
||||
return {name: value for name, value in candidate_headers.items() if value is not None}
|
||||
|
||||
|
||||
class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
||||
def __init__(self, **kwargs):
|
||||
|
|
@ -390,7 +424,16 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
return headers
|
||||
raw_guardrail_config = optional_params.pop("guardrailConfig", None)
|
||||
if raw_guardrail_config is None:
|
||||
return headers
|
||||
existing_header_names = frozenset(name.lower() for name in headers)
|
||||
guardrail_headers = {
|
||||
name: value
|
||||
for name, value in _bedrock_invoke_guardrail_headers(raw_guardrail_config).items()
|
||||
if name.lower() not in existing_header_names
|
||||
}
|
||||
return {**headers, **guardrail_headers}
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
|
|
|
|||
|
|
@ -685,39 +685,27 @@ def get_bedrock_base_model(model: str) -> str:
|
|||
return model
|
||||
|
||||
|
||||
def bedrock_converse_supports_parallel_tool_use_config(model: str) -> bool:
|
||||
return any(
|
||||
(litellm.model_cost.get(candidate) or {}).get("supports_parallel_tool_use_config") is True
|
||||
for candidate in (model, get_bedrock_base_model(model))
|
||||
)
|
||||
|
||||
|
||||
def is_claude_4_5_on_bedrock(model: str) -> bool:
|
||||
"""
|
||||
Check if the model is a Claude 4.5 model on Bedrock.
|
||||
Claude 4.5 models support prompt caching with '5m' and '1h' TTL on Bedrock.
|
||||
Check if the model supports Bedrock prompt caching with an extended '1h' TTL
|
||||
(in addition to the default 5m TTL).
|
||||
|
||||
Backed by the ``cache_creation_input_token_cost_above_1hr`` field in
|
||||
``model_prices_and_context_window.json`` instead of a hardcoded list of
|
||||
model-name patterns, so newly released models pick up support as soon as
|
||||
their pricing entry ships, with no code change required here.
|
||||
"""
|
||||
model_lower = model.lower()
|
||||
claude_4_5_patterns = [
|
||||
"sonnet-4.5",
|
||||
"sonnet_4.5",
|
||||
"sonnet-4-5",
|
||||
"sonnet_4_5",
|
||||
"haiku-4.5",
|
||||
"haiku_4.5",
|
||||
"haiku-4-5",
|
||||
"haiku_4_5",
|
||||
"opus-4.5",
|
||||
"opus_4.5",
|
||||
"opus-4-5",
|
||||
"opus_4_5",
|
||||
"sonnet-4.6",
|
||||
"sonnet_4.6",
|
||||
"sonnet-4-6",
|
||||
"sonnet_4_6",
|
||||
"opus-4.6",
|
||||
"opus_4.6",
|
||||
"opus-4-6",
|
||||
"opus_4_6",
|
||||
"opus-4.7",
|
||||
"opus_4.7",
|
||||
"opus-4-7",
|
||||
"opus_4_7",
|
||||
]
|
||||
return any(pattern in model_lower for pattern in claude_4_5_patterns)
|
||||
return any(
|
||||
(litellm.model_cost.get(candidate) or {}).get("cache_creation_input_token_cost_above_1hr") is not None
|
||||
for candidate in (model, get_bedrock_base_model(model))
|
||||
)
|
||||
|
||||
|
||||
_BEDROCK_MODEL_VERSION_SUFFIX_RE = re.compile(r"-v\d+(?::\d+)?$")
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ from litellm.types.llms.openai import (
|
|||
AllMessageValues,
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
ResponsesAPIStreamEvents,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
|
|
@ -586,7 +587,14 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
"""
|
||||
Check if the streaming has ended.
|
||||
"""
|
||||
return all(response.choices[0].finish_reason is not None for response in responses_so_far)
|
||||
if not responses_so_far:
|
||||
return False
|
||||
terminal_types = {
|
||||
ResponsesAPIStreamEvents.RESPONSE_COMPLETED.value,
|
||||
ResponsesAPIStreamEvents.RESPONSE_FAILED.value,
|
||||
ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value,
|
||||
}
|
||||
return responses_so_far[-1].get("type") in terminal_types
|
||||
|
||||
def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str:
|
||||
"""
|
||||
|
|
|
|||
0
litellm/llms/tencent/__init__.py
Normal file
0
litellm/llms/tencent/__init__.py
Normal file
0
litellm/llms/tencent/chat/__init__.py
Normal file
0
litellm/llms/tencent/chat/__init__.py
Normal file
68
litellm/llms/tencent/chat/transformation.py
Normal file
68
litellm/llms/tencent/chat/transformation.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
"""
|
||||
Translates from OpenAI's `/v1/chat/completions` to Tencent TokenHub's
|
||||
OpenAI-compatible endpoint.
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.utils import supports_reasoning
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
||||
class TencentChatConfig(OpenAIGPTConfig):
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
params = super().get_supported_openai_params(model)
|
||||
if supports_reasoning(model, custom_llm_provider="tencent"):
|
||||
params.extend(["thinking", "reasoning_effort"])
|
||||
return params
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
optional_params = super().map_openai_params(non_default_params, optional_params, model, drop_params)
|
||||
|
||||
thinking_value = optional_params.pop("thinking", None)
|
||||
reasoning_effort = optional_params.pop("reasoning_effort", None)
|
||||
|
||||
if thinking_value is not None:
|
||||
if isinstance(thinking_value, dict):
|
||||
optional_params["thinking"] = thinking_value
|
||||
elif reasoning_effort is not None and reasoning_effort != "none":
|
||||
optional_params["thinking"] = {"type": "enabled"}
|
||||
|
||||
return optional_params
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> tuple[Optional[str], Optional[str]]:
|
||||
api_base = api_base or get_secret_str("TENCENT_API_BASE") or "https://tokenhub-intl.tencentcloudmaas.com/v1"
|
||||
dynamic_api_key = api_key or get_secret_str("TENCENT_API_KEY")
|
||||
return api_base, dynamic_api_key
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
if not api_base:
|
||||
api_base = "https://tokenhub-intl.tencentcloudmaas.com/v1"
|
||||
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
if api_base.endswith("/chat/completions"):
|
||||
return api_base
|
||||
|
||||
if not api_base.endswith("/v1"):
|
||||
api_base = f"{api_base}/v1"
|
||||
|
||||
return f"{api_base}/chat/completions"
|
||||
6
litellm/llms/tencent/cost_calculator.py
Normal file
6
litellm/llms/tencent/cost_calculator.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
def cost_per_token(model: str, usage: Usage) -> tuple[float, float]:
|
||||
return generic_cost_per_token(model=model, usage=usage, custom_llm_provider="tencent")
|
||||
0
litellm/llms/tencent/messages/__init__.py
Normal file
0
litellm/llms/tencent/messages/__init__.py
Normal file
85
litellm/llms/tencent/messages/transformation.py
Normal file
85
litellm/llms/tencent/messages/transformation.py
Normal file
|
|
@ -0,0 +1,85 @@
|
|||
"""
|
||||
Tencent Anthropic-compatible messages transformation config.
|
||||
|
||||
Tencent TokenHub exposes an Anthropic-compatible Messages API endpoint
|
||||
alongside its standard OpenAI-compatible chat completions endpoint.
|
||||
"""
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class TencentAnthropicMessagesConfig(AnthropicMessagesConfig):
|
||||
"""
|
||||
Tencent TokenHub exposes an Anthropic-compatible Messages API.
|
||||
|
||||
Unlike the chat completions endpoint (which uses /v1), the Anthropic
|
||||
endpoint may use a different base URL. Configure via
|
||||
TENCENT_ANTHROPIC_API_BASE or TENCENT_API_BASE.
|
||||
"""
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "tencent"
|
||||
|
||||
def should_strip_billing_metadata(self) -> bool:
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def get_api_key(api_key: Optional[str] = None) -> Optional[str]:
|
||||
return api_key or get_secret_str("TENCENT_API_KEY") or litellm.api_key
|
||||
|
||||
@staticmethod
|
||||
def get_api_base(api_base: Optional[str] = None) -> str:
|
||||
return (
|
||||
api_base
|
||||
or get_secret_str("TENCENT_ANTHROPIC_API_BASE")
|
||||
or get_secret_str("TENCENT_API_BASE")
|
||||
or "https://tokenhub-intl.tencentcloudmaas.com"
|
||||
)
|
||||
|
||||
def validate_anthropic_messages_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: list[Any],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> tuple[dict, Optional[str]]:
|
||||
return super().validate_anthropic_messages_environment(
|
||||
headers=headers,
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
api_key=self.get_api_key(api_key=api_key),
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
base_url = self.get_api_base(api_base=api_base).rstrip("/")
|
||||
|
||||
if base_url.endswith("/v1/messages"):
|
||||
return base_url
|
||||
|
||||
if base_url.endswith("/v1/chat/completions"):
|
||||
base_url = base_url[: -len("/v1/chat/completions")]
|
||||
elif base_url.endswith("/v1"):
|
||||
base_url = base_url[: -len("/v1")]
|
||||
|
||||
return f"{base_url}/v1/messages"
|
||||
|
|
@ -1,5 +1,5 @@
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
|
@ -175,7 +175,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
|
|||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
|
|
@ -217,10 +217,10 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
|
|||
contents = [{"role": "user", "parts": [{"text": prompt}]}]
|
||||
|
||||
# Prepare generation config
|
||||
generation_config: Dict[str, Any] = {"responseModalities": ["IMAGE"]}
|
||||
generation_config: dict[str, Any] = {"responseModalities": ["IMAGE"]}
|
||||
|
||||
# Seed from user-supplied imageConfig dict; flat params are overlaid for backward compat.
|
||||
image_config: Dict[str, Any] = dict(optional_params.get("imageConfig") or {})
|
||||
image_config: dict[str, Any] = dict(optional_params.get("imageConfig") or {})
|
||||
|
||||
if "aspectRatio" in optional_params:
|
||||
image_config["aspectRatio"] = optional_params["aspectRatio"]
|
||||
|
|
@ -241,7 +241,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM):
|
|||
elif "n" in optional_params:
|
||||
generation_config["candidateCount"] = optional_params["n"]
|
||||
|
||||
request_body: Dict[str, Any] = {
|
||||
request_body: dict[str, Any] = {
|
||||
"contents": contents,
|
||||
"generationConfig": generation_config,
|
||||
}
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -93,6 +93,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
|
|||
has_user_credential: Optional[bool] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
approval_status: Optional[str] = Field(
|
||||
default="active",
|
||||
description="Approval status: 'pending_review', 'active', 'rejected'",
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1261,6 +1261,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
byok_api_key_help_url: Optional[str] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
# BYOM submission fields — set by the endpoint, not by the caller.
|
||||
# Any caller-provided values are silently overridden before persistence.
|
||||
approval_status: Optional[str] = Field(
|
||||
|
|
@ -1346,6 +1347,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
|
|||
byok_api_key_help_url: Optional[str] = None
|
||||
source_url: Optional[str] = None
|
||||
timeout: Optional[float] = None
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
|
|
|
|||
|
|
@ -12,6 +12,9 @@ from litellm.integrations.custom_guardrail import ModifyResponseException
|
|||
from litellm.llms.anthropic.experimental_pass_through.context_management import (
|
||||
AnthropicContextManagementError,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_response_usage as _blocked_response_usage,
|
||||
)
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import (
|
||||
|
|
@ -134,6 +137,10 @@ async def anthropic_response(
|
|||
|
||||
from litellm.types.utils import AnthropicMessagesResponse
|
||||
|
||||
# Report the blocked LLM response's real token usage (carried on the
|
||||
# exception) instead of discarding it; zero for pre-call blocks.
|
||||
_usage = _blocked_response_usage(e.original_response)
|
||||
|
||||
_anthropic_response = AnthropicMessagesResponse(
|
||||
id=f"msg_{str(uuid.uuid4())}",
|
||||
type="message",
|
||||
|
|
@ -141,7 +148,7 @@ async def anthropic_response(
|
|||
content=[{"type": "text", "text": e.message}],
|
||||
model=e.model,
|
||||
stop_reason="end_turn",
|
||||
usage={"input_tokens": 0, "output_tokens": 0},
|
||||
usage=_usage,
|
||||
)
|
||||
|
||||
if data.get("stream", None) is not None and data["stream"] is True:
|
||||
|
|
|
|||
|
|
@ -3017,7 +3017,9 @@ async def can_key_call_resolved_model(
|
|||
)
|
||||
|
||||
skip_key_model_check = valid_token.config or (
|
||||
isinstance(valid_token.models, list) and SpecialModelNames.all_team_models.value in valid_token.models
|
||||
isinstance(valid_token.models, list)
|
||||
and SpecialModelNames.all_team_models.value in valid_token.models
|
||||
and valid_token.team_id is not None
|
||||
)
|
||||
if not skip_key_model_check:
|
||||
await can_key_call_model(
|
||||
|
|
|
|||
|
|
@ -375,6 +375,16 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
|
|||
metadata = _coerce_metadata_to_dict(request_body.get(metadata_key))
|
||||
if metadata is not None:
|
||||
_check_banned_params(metadata, general_settings, llm_router, model)
|
||||
litellm_params = _coerce_metadata_to_dict(request_body.get("litellm_params"))
|
||||
if litellm_params is not None:
|
||||
litellm_params_metadata = _coerce_metadata_to_dict(litellm_params.get("metadata"))
|
||||
if litellm_params_metadata is not None:
|
||||
_check_banned_params(
|
||||
litellm_params_metadata,
|
||||
general_settings,
|
||||
llm_router,
|
||||
model,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2676,7 +2676,11 @@ async def _enforce_key_and_fallback_model_access(
|
|||
model_list = config.get("model_list", [])
|
||||
new_model_list = model_list
|
||||
verbose_proxy_logger.debug(f"\n new llm router model list {new_model_list}")
|
||||
elif isinstance(valid_token.models, list) and "all-team-models" in valid_token.models:
|
||||
elif (
|
||||
isinstance(valid_token.models, list)
|
||||
and "all-team-models" in valid_token.models
|
||||
and valid_token.team_id is not None
|
||||
):
|
||||
pass
|
||||
else:
|
||||
model = _get_model_from_request_context(
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint
|
|||
|
||||
import copy
|
||||
import json
|
||||
from typing import Any, AsyncGenerator, List, Optional, Union
|
||||
from typing import TYPE_CHECKING, Any, AsyncGenerator, List, Optional, Union
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
|
@ -23,6 +23,11 @@ from litellm.proxy._types import UserAPIKeyAuth
|
|||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.utils import CallTypes, CallTypesLiteral
|
||||
|
||||
if TYPE_CHECKING:
|
||||
# Imported lazily at runtime (inside the streaming hook) to avoid a
|
||||
# module-level cyclic import with litellm.integrations.custom_guardrail.
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
# Call types that use NDJSON streaming (A2A); guardrail HTTPException is emitted as in-stream error
|
||||
A2A_CALL_TYPES = (CallTypes.asend_message, CallTypes.send_message)
|
||||
|
||||
|
|
@ -197,6 +202,10 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
|
||||
# Local import avoids a module-level cyclic import with
|
||||
# litellm.integrations.custom_guardrail.
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
guardrail_to_apply: CustomGuardrail = data.pop("guardrail_to_apply", None)
|
||||
|
||||
if guardrail_to_apply is None:
|
||||
|
|
@ -238,18 +247,51 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
|
||||
|
||||
response = await endpoint_translation.process_output_response(
|
||||
response=response, # type: ignore
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
litellm_logging_obj=data.get("litellm_logging_obj"),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=data,
|
||||
)
|
||||
try:
|
||||
response = await endpoint_translation.process_output_response(
|
||||
response=response, # type: ignore
|
||||
guardrail_to_apply=guardrail_to_apply,
|
||||
litellm_logging_obj=data.get("litellm_logging_obj"),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=data,
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
# The guardrail blocked the response. Attach the original LLM
|
||||
# response so the endpoint handler can report its real token usage
|
||||
# instead of discarding it (the block replaces the content, but the
|
||||
# upstream call already consumed those tokens).
|
||||
if e.original_response is None:
|
||||
e.original_response = response
|
||||
raise
|
||||
# Add guardrail to applied guardrails header
|
||||
add_guardrail_to_applied_guardrails_header(request_data=data, guardrail_name=guardrail_to_apply.guardrail_name)
|
||||
|
||||
return response
|
||||
|
||||
async def _handle_streaming_block(
|
||||
self,
|
||||
exc: "ModifyResponseException",
|
||||
endpoint_translation: Any,
|
||||
stream_started: bool,
|
||||
responses_so_far: list[Any],
|
||||
) -> AsyncGenerator[Any, None]:
|
||||
"""
|
||||
Terminate a streamed response cleanly when a guardrail blocks it.
|
||||
|
||||
Format-agnostic routing: delegates to the provider translation handler's
|
||||
``build_block_sse_chunks`` (see ``BaseTranslation.build_block_sse_chunks``
|
||||
for the ``stream_started`` / ``responses_so_far`` contract). When the
|
||||
format has no safe terminator the handler returns None and we re-raise
|
||||
``exc`` so the proxy can surface a clean error.
|
||||
"""
|
||||
block_chunks = endpoint_translation.build_block_sse_chunks(
|
||||
exc, stream_started=stream_started, responses_so_far=responses_so_far
|
||||
)
|
||||
if block_chunks is None:
|
||||
raise exc
|
||||
for chunk in block_chunks:
|
||||
yield chunk
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -271,26 +313,53 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
global endpoint_guardrail_translation_mappings
|
||||
|
||||
# Local import avoids a module-level cyclic import with
|
||||
# litellm.integrations.custom_guardrail.
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
guardrail_to_apply: CustomGuardrail = request_data.pop("guardrail_to_apply", None)
|
||||
|
||||
# Get streaming configuration from guardrail or optional_params
|
||||
sampling_rate = 5
|
||||
end_of_stream_only = False # If True, only apply guardrail at end of stream
|
||||
# Get streaming configuration. Resolution order (later wins): default
|
||||
# < guardrail attribute < guardrail_config dict < this callback's
|
||||
# optional_params.
|
||||
def _streaming_flag(name: str, default: Any) -> Any:
|
||||
value = default
|
||||
if guardrail_to_apply is not None:
|
||||
value = getattr(guardrail_to_apply, name, value)
|
||||
config = getattr(guardrail_to_apply, "guardrail_config", {})
|
||||
if isinstance(config, dict):
|
||||
value = config.get(name, value)
|
||||
return self.optional_params.get(name, value)
|
||||
|
||||
if guardrail_to_apply is not None:
|
||||
# Check direct attributes on guardrail first
|
||||
sampling_rate = getattr(guardrail_to_apply, "streaming_sampling_rate", sampling_rate)
|
||||
end_of_stream_only = getattr(guardrail_to_apply, "streaming_end_of_stream_only", end_of_stream_only)
|
||||
sampling_rate = _streaming_flag("streaming_sampling_rate", 5)
|
||||
# Only apply the guardrail at end of stream (not per chunk).
|
||||
end_of_stream_only = _streaming_flag("streaming_end_of_stream_only", False)
|
||||
# Withhold every chunk until end-of-stream moderation passes, then
|
||||
# release the original chunks (clean) or only the block message
|
||||
# (blocked) -- moderating the whole response *before* any content
|
||||
# reaches the client. Only safe for allow/block guardrails: on
|
||||
# release the original chunks are replayed as-is, so a
|
||||
# content-rewriting guardrail (e.g. PII masking) would leak
|
||||
# unredacted content. Guarded below via mask_response_content.
|
||||
buffer_until_moderated = _streaming_flag("streaming_buffer_until_moderated", False)
|
||||
|
||||
# Also check guardrail_config dict if present
|
||||
guardrail_config = getattr(guardrail_to_apply, "guardrail_config", {})
|
||||
if isinstance(guardrail_config, dict):
|
||||
sampling_rate = guardrail_config.get("streaming_sampling_rate", sampling_rate)
|
||||
end_of_stream_only = guardrail_config.get("streaming_end_of_stream_only", end_of_stream_only)
|
||||
if (
|
||||
buffer_until_moderated
|
||||
and guardrail_to_apply is not None
|
||||
and getattr(guardrail_to_apply, "mask_response_content", False)
|
||||
):
|
||||
verbose_proxy_logger.warning(
|
||||
"UnifiedLLMGuardrails: streaming_buffer_until_moderated is disabled for %s "
|
||||
"because mask_response_content=True -- buffered replay would release "
|
||||
"unredacted original chunks instead of the moderated output.",
|
||||
guardrail_to_apply.guardrail_name,
|
||||
)
|
||||
buffer_until_moderated = False
|
||||
|
||||
# Also check optional_params as fallback
|
||||
sampling_rate = self.optional_params.get("streaming_sampling_rate", sampling_rate)
|
||||
end_of_stream_only = self.optional_params.get("streaming_end_of_stream_only", end_of_stream_only)
|
||||
# Buffering can only moderate the assembled response, so it always
|
||||
# defers to end-of-stream.
|
||||
if buffer_until_moderated:
|
||||
end_of_stream_only = True
|
||||
|
||||
if guardrail_to_apply is None:
|
||||
async for item in response:
|
||||
|
|
@ -315,6 +384,12 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
call_type = None
|
||||
chunk_counter = 0
|
||||
responses_so_far: List[Any] = []
|
||||
responses_yielded: list[Any] = []
|
||||
pending_end_of_stream_items: list[Any] = []
|
||||
# Whether any real response chunk has been forwarded to the client.
|
||||
# Drives how a block terminates the stream: continue the in-progress
|
||||
# message (True) vs emit a standalone block message (False, buffered).
|
||||
chunks_yielded = False
|
||||
|
||||
async for item in response:
|
||||
chunk_counter += 1
|
||||
|
|
@ -336,9 +411,22 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
yield remaining_item
|
||||
return
|
||||
|
||||
# If end_of_stream_only mode, yield chunks without processing
|
||||
# If end_of_stream_only mode, yield chunks without processing.
|
||||
# When buffering, withhold them instead -- they are released (or
|
||||
# replaced by the block message) only after end-of-stream
|
||||
# moderation runs below.
|
||||
if end_of_stream_only:
|
||||
yield item
|
||||
if not buffer_until_moderated:
|
||||
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
|
||||
stream_has_ended = hasattr(
|
||||
endpoint_translation, "_check_streaming_has_ended"
|
||||
) and endpoint_translation._check_streaming_has_ended(responses_so_far)
|
||||
if pending_end_of_stream_items or stream_has_ended:
|
||||
pending_end_of_stream_items.append(item)
|
||||
else:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
continue
|
||||
|
||||
# Process chunk based on sampling rate
|
||||
|
|
@ -368,6 +456,26 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
)
|
||||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = responses_so_far
|
||||
# Guardrail blocked the response mid-stream. Emit a clean
|
||||
# terminating SSE sequence delivering the block message
|
||||
# instead of letting the exception propagate into a bare
|
||||
# `data: {"error": ...}` blob (which truncates the stream).
|
||||
# Chunks have already been forwarded here, so the block
|
||||
# continues the in-progress message (stream_started=True).
|
||||
# The current chunk was appended to responses_so_far but not
|
||||
# yet yielded, so exclude it: the continuation must reflect
|
||||
# only what the client has actually received.
|
||||
async for block_chunk in self._handle_streaming_block(
|
||||
e,
|
||||
endpoint_translation,
|
||||
stream_started=chunks_yielded,
|
||||
responses_so_far=responses_yielded,
|
||||
):
|
||||
yield block_chunk
|
||||
return
|
||||
except HTTPException as e:
|
||||
# Response already started (we already yielded chunks); cannot send 400.
|
||||
# For A2A (NDJSON), yield an in-stream JSON-RPC error so the client sees it.
|
||||
|
|
@ -394,8 +502,12 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
yield error_chunk
|
||||
return
|
||||
raise
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(original_item)
|
||||
yield original_item
|
||||
else:
|
||||
chunks_yielded = True
|
||||
responses_yielded.append(item)
|
||||
yield item
|
||||
|
||||
# Stream has ended - do final processing with all collected chunks
|
||||
|
|
@ -408,6 +520,15 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
|
||||
endpoint_translation = endpoint_guardrail_translation_mappings[CallTypes(call_type)]()
|
||||
|
||||
# When buffering, snapshot the original chunks before moderation.
|
||||
# A shallow copy suffices: end-of-stream
|
||||
# process_output_streaming_response builds a separate assembled
|
||||
# response (it does not mutate the individual chunks in place), and
|
||||
# the chunks themselves are replayed verbatim -- so we only need to
|
||||
# preserve the list, not clone every chunk (deepcopy would double
|
||||
# peak memory for large responses).
|
||||
buffered_items = list(responses_so_far) if buffer_until_moderated else None
|
||||
|
||||
try:
|
||||
await endpoint_translation.process_output_streaming_response(
|
||||
responses_so_far=responses_so_far,
|
||||
|
|
@ -416,6 +537,28 @@ class UnifiedLLMGuardrails(CustomLogger):
|
|||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
)
|
||||
# Moderation passed: release the withheld original chunks.
|
||||
if buffered_items is not None:
|
||||
for buffered_item in buffered_items:
|
||||
yield buffered_item
|
||||
for pending_item in pending_end_of_stream_items:
|
||||
responses_yielded.append(pending_item)
|
||||
yield pending_item
|
||||
except ModifyResponseException as e:
|
||||
if e.original_response is None:
|
||||
e.original_response = responses_so_far
|
||||
# Block detected during end-of-stream processing. Emit a clean
|
||||
# terminating SSE sequence with the block message rather than
|
||||
# propagating into a bare error blob that truncates the stream.
|
||||
# The withheld original chunks are never released.
|
||||
async for block_chunk in self._handle_streaming_block(
|
||||
e,
|
||||
endpoint_translation,
|
||||
stream_started=bool(responses_yielded),
|
||||
responses_so_far=responses_yielded,
|
||||
):
|
||||
yield block_chunk
|
||||
return
|
||||
except HTTPException as e:
|
||||
if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES:
|
||||
request_id = _get_a2a_request_id(responses_so_far, request_data)
|
||||
|
|
|
|||
|
|
@ -15,6 +15,9 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
|
|||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
iter_client_callback_metadata_dicts,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.proxy._types import (
|
||||
|
|
@ -302,6 +305,24 @@ def _key_or_team_allows_client_pricing_override(
|
|||
)
|
||||
|
||||
|
||||
def _strip_client_message_redaction_opt_out(data: dict[str, Any]) -> None:
|
||||
stripped: list[str] = []
|
||||
if "turn_off_message_logging" in data and _is_false_like(data["turn_off_message_logging"]):
|
||||
stripped.append("turn_off_message_logging")
|
||||
data.pop("turn_off_message_logging", None)
|
||||
for slot_label, metadata in iter_client_callback_metadata_dicts(data):
|
||||
if "turn_off_message_logging" in metadata and _is_false_like(metadata["turn_off_message_logging"]):
|
||||
stripped.append(f"{slot_label}.turn_off_message_logging")
|
||||
metadata.pop("turn_off_message_logging", None)
|
||||
if stripped:
|
||||
verbose_proxy_logger.debug(
|
||||
"Stripped client-supplied message-redaction opt-out fields from request body: %s. "
|
||||
"Set `allow_client_message_redaction_opt_out: true` on the key or team metadata "
|
||||
"to keep these values.",
|
||||
", ".join(stripped),
|
||||
)
|
||||
|
||||
|
||||
def _strip_client_pricing_overrides(data: Dict[str, Any]) -> None:
|
||||
"""Drop pricing overrides from the request body and any metadata variant.
|
||||
|
||||
|
|
@ -1308,13 +1329,6 @@ async def add_litellm_data_to_request(
|
|||
_headers,
|
||||
allow_client_message_redaction_opt_out=_allow_client_message_redaction_opt_out,
|
||||
)
|
||||
if (
|
||||
not _allow_client_message_redaction_opt_out
|
||||
and litellm.turn_off_message_logging is True
|
||||
and "turn_off_message_logging" in data
|
||||
and _is_false_like(data["turn_off_message_logging"])
|
||||
):
|
||||
data.pop("turn_off_message_logging", None)
|
||||
verbose_proxy_logger.debug(f"Request Headers: {_headers}")
|
||||
verbose_proxy_logger.debug(f"Raw Headers: {_raw_headers}")
|
||||
|
||||
|
|
@ -1466,6 +1480,9 @@ async def add_litellm_data_to_request(
|
|||
if not _key_or_team_allows_client_pricing_override(user_api_key_dict):
|
||||
_strip_client_pricing_overrides(data)
|
||||
|
||||
if not _allow_client_message_redaction_opt_out and litellm.turn_off_message_logging is True:
|
||||
_strip_client_message_redaction_opt_out(data)
|
||||
|
||||
# Fill in the proxy_server_request body snapshot now that metadata has
|
||||
# been parsed. Consumers (standard_logging_payload, lago,
|
||||
# spend_tracking_utils, streaming_iterator) read `body` to audit the
|
||||
|
|
@ -1649,6 +1666,16 @@ async def add_litellm_data_to_request(
|
|||
tags_to_add=tags,
|
||||
)
|
||||
|
||||
if _metadata_variable_name != "metadata":
|
||||
_user_metadata = data.get("metadata")
|
||||
if isinstance(_user_metadata, dict):
|
||||
_user_tags = _user_metadata.get("tags")
|
||||
if isinstance(_user_tags, list) and _user_tags:
|
||||
data[_metadata_variable_name]["tags"] = LiteLLMProxyRequestSetup._merge_tags(
|
||||
request_tags=data[_metadata_variable_name].get("tags"),
|
||||
tags_to_add=_user_tags,
|
||||
)
|
||||
|
||||
# Team Callbacks controls
|
||||
callback_settings_obj = _get_dynamic_logging_metadata(
|
||||
user_api_key_dict=user_api_key_dict, proxy_config=proxy_config
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.proxy.management_endpoints.common_utils import (
|
|||
validate_finite_spend,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_check_permissions_caller_permission,
|
||||
generate_key_helper_fn,
|
||||
prepare_metadata_fields,
|
||||
)
|
||||
|
|
@ -440,6 +441,11 @@ async def new_user(
|
|||
detail=f"Only proxy admins can create administrative users (proxy_admin, proxy_admin_viewer). Attempted to create user with role: {data.user_role}. Your role: {user_api_key_dict.user_role}",
|
||||
)
|
||||
|
||||
_check_permissions_caller_permission(
|
||||
data=data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
data_json = data.json() # type: ignore
|
||||
data_json = _update_internal_new_user_params(data_json, data)
|
||||
_hash_password_in_dict(data_json)
|
||||
|
|
@ -1198,6 +1204,11 @@ async def _update_single_user_helper(
|
|||
if not user_request.user_id and not user_request.user_email:
|
||||
raise ValueError("Either user_id or user_email must be provided")
|
||||
|
||||
_check_permissions_caller_permission(
|
||||
data=user_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
data_json: dict = user_request.model_dump(exclude_unset=True)
|
||||
non_default_values = _update_internal_user_params(data_json=data_json, data=user_request)
|
||||
_hash_password_in_dict(non_default_values)
|
||||
|
|
|
|||
|
|
@ -530,26 +530,33 @@ def _check_allowed_routes_caller_permission(
|
|||
allowed_routes: Optional[list],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
*,
|
||||
allowed_routes_was_provided: bool = False,
|
||||
allow_safe_presets: bool = False,
|
||||
) -> None:
|
||||
"""
|
||||
Only proxy admins may set `allowed_routes` on a key.
|
||||
Require PROXY_ADMIN when `allowed_routes` is present in the request body,
|
||||
unless the caller went through the `key_type` preset flow.
|
||||
|
||||
`allowed_routes` overrides the standard role-based route gate in
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check, so the field is
|
||||
restricted to admins. Non-admins must instead use `key_type` to pick a
|
||||
preset bucket — that path goes through `handle_key_type` and re-enters
|
||||
this function with `allow_safe_presets=True`, which lets the derived
|
||||
`llm_api_routes` / `info_routes` values through. Raw-body call sites
|
||||
leave `allow_safe_presets=False` so non-admins can't write those values
|
||||
directly.
|
||||
Raw-body call sites pass
|
||||
`allowed_routes_was_provided="allowed_routes" in data.model_fields_set` so a
|
||||
caller that omits the field (model default flows through) is distinct from
|
||||
one that sends any explicit value.
|
||||
|
||||
Post-`handle_key_type` call sites pass `allow_safe_presets=True` with the
|
||||
values derived by `handle_key_type`; those values are not from the request
|
||||
body, so `allowed_routes_was_provided` stays False and the safe-preset
|
||||
carve-out below accepts any list of tokens in
|
||||
`_NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS`.
|
||||
"""
|
||||
# Empty list is the default on GenerateKeyRequest — treat as "not set".
|
||||
if not allowed_routes:
|
||||
if not allowed_routes_was_provided and not allowed_routes:
|
||||
return
|
||||
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value:
|
||||
return
|
||||
if allow_safe_presets and all(r in _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS for r in allowed_routes):
|
||||
if (
|
||||
allow_safe_presets
|
||||
and allowed_routes
|
||||
and all(r in _NON_ADMIN_SAFE_ALLOWED_ROUTES_PRESETS for r in allowed_routes)
|
||||
):
|
||||
return
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
|
|
@ -580,7 +587,7 @@ def _check_permissions_caller_permission(
|
|||
return
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail={"error": "Only proxy admins can set `permissions` on a key."},
|
||||
detail={"error": "Only proxy admins can set `permissions`."},
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -1566,6 +1573,7 @@ async def generate_key_fn(
|
|||
_check_allowed_routes_caller_permission(
|
||||
allowed_routes=data.allowed_routes,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
|
||||
)
|
||||
_check_passthrough_routes_caller_permission(
|
||||
data=data,
|
||||
|
|
@ -1736,6 +1744,7 @@ async def generate_service_account_key_fn(
|
|||
_check_allowed_routes_caller_permission(
|
||||
allowed_routes=data.allowed_routes,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
|
||||
)
|
||||
_check_passthrough_routes_caller_permission(
|
||||
data=data,
|
||||
|
|
@ -2060,6 +2069,11 @@ async def _process_single_key_update(
|
|||
# Validate max_budget
|
||||
_validate_max_budget(update_key_request.max_budget)
|
||||
|
||||
_check_permissions_caller_permission(
|
||||
data=update_key_request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Get and validate existing key
|
||||
if existing_key_row is None:
|
||||
existing_key_row = await _get_and_validate_existing_key(
|
||||
|
|
@ -2231,6 +2245,7 @@ async def _validate_update_key_data(
|
|||
_check_allowed_routes_caller_permission(
|
||||
allowed_routes=data.allowed_routes,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
|
||||
)
|
||||
_check_passthrough_routes_caller_permission(
|
||||
data=data,
|
||||
|
|
@ -4551,6 +4566,7 @@ async def regenerate_key_fn(
|
|||
_check_allowed_routes_caller_permission(
|
||||
allowed_routes=data.allowed_routes,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
allowed_routes_was_provided="allowed_routes" in data.model_fields_set,
|
||||
)
|
||||
_check_passthrough_routes_caller_permission(
|
||||
data=data,
|
||||
|
|
|
|||
|
|
@ -654,6 +654,7 @@ if MCP_AVAILABLE:
|
|||
allow_all_keys=payload.allow_all_keys,
|
||||
available_on_public_internet=payload.available_on_public_internet,
|
||||
timeout=payload.timeout,
|
||||
max_concurrent_requests=payload.max_concurrent_requests,
|
||||
)
|
||||
|
||||
def get_prisma_client_or_throw(message: str):
|
||||
|
|
|
|||
|
|
@ -8361,6 +8361,22 @@ async def model_info(
|
|||
)
|
||||
|
||||
|
||||
def _blocked_response_usage(original_response: Optional[Any]) -> "litellm.Usage":
|
||||
"""
|
||||
Token usage for a synthetic guardrail-blocked response.
|
||||
|
||||
A post-call block replaces the LLM's response with the violation message,
|
||||
but the upstream call already consumed tokens -- report that real usage
|
||||
(carried on ``ModifyResponseException.original_response``) rather than
|
||||
discarding it. Pre-call blocks never invoked the LLM (no original_response),
|
||||
so usage is zero.
|
||||
"""
|
||||
usage = getattr(original_response, "usage", None) if original_response is not None else None
|
||||
if isinstance(usage, litellm.Usage):
|
||||
return usage
|
||||
return litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/chat/completions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
|
|
@ -8467,6 +8483,9 @@ async def chat_completion(
|
|||
_chat_response.model = e.model # type: ignore
|
||||
_chat_response.choices[0].message.content = e.message # type: ignore
|
||||
_chat_response.choices[0].finish_reason = "content_filter" # type: ignore
|
||||
# Report the blocked LLM response's real usage (set before the stream
|
||||
# branch so both paths carry it); zero for pre-call blocks.
|
||||
_chat_response.usage = _blocked_response_usage(e.original_response) # type: ignore
|
||||
|
||||
if data.get("stream", None) is not None and data["stream"] is True:
|
||||
_iterator = litellm.utils.ModelResponseIterator(model_response=_chat_response, convert_to_delta=True)
|
||||
|
|
@ -8488,8 +8507,6 @@ async def chat_completion(
|
|||
media_type="text/event-stream",
|
||||
status_code=200, # Return 200 for passthrough mode
|
||||
)
|
||||
_usage = litellm.Usage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
|
||||
_chat_response.usage = _usage # type: ignore
|
||||
return _chat_response
|
||||
except RejectedRequestError as e:
|
||||
_data = e.request_data
|
||||
|
|
@ -8618,11 +8635,7 @@ async def completion(
|
|||
# Set text attribute dynamically for text completion format
|
||||
setattr(_text_response.choices[0], "text", e.message)
|
||||
_text_response.model = e.model # type: ignore[assignment]
|
||||
_usage = litellm.Usage(
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
total_tokens=0,
|
||||
)
|
||||
_usage = _blocked_response_usage(e.original_response)
|
||||
# Set usage attribute dynamically (ModelResponse accepts usage in __init__ but it's not in type definition)
|
||||
setattr(_text_response, "usage", _usage)
|
||||
_iterator = litellm.utils.ModelResponseIterator(model_response=_text_response, convert_to_delta=True)
|
||||
|
|
@ -8647,11 +8660,7 @@ async def completion(
|
|||
_response = litellm.TextCompletionResponse()
|
||||
_response.choices[0].text = e.message
|
||||
_response.model = e.model # type: ignore
|
||||
_usage = litellm.Usage(
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
total_tokens=0,
|
||||
)
|
||||
_usage = _blocked_response_usage(e.original_response)
|
||||
_response.usage = _usage # type: ignore
|
||||
return _response
|
||||
except RejectedRequestError as e:
|
||||
|
|
|
|||
|
|
@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
max_concurrent_requests Int?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ class CacheControlToolConfigInjectionPoint(TypedDict):
|
|||
"""Type for tool_config-level injection points (Bedrock)."""
|
||||
|
||||
location: Literal["tool_config"]
|
||||
control: Optional[ChatCompletionCachedContent]
|
||||
|
||||
|
||||
CacheControlInjectionPoint = Union[
|
||||
|
|
|
|||
|
|
@ -325,7 +325,7 @@ class ToolConfigBlock(TypedDict, total=False):
|
|||
class GuardrailConfigBlock(TypedDict, total=False):
|
||||
guardrailIdentifier: str
|
||||
guardrailVersion: str
|
||||
trace: Literal["enabled", "disabled"]
|
||||
trace: Literal["enabled", "disabled", "enabled_full"]
|
||||
|
||||
|
||||
class InferenceConfig(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -118,6 +118,9 @@ class MCPServer(BaseModel):
|
|||
# MCP_PER_USER_TOKEN_DEFAULT_TTL when expires_in is absent.
|
||||
token_storage_ttl_seconds: Optional[int] = None
|
||||
timeout: Optional[float] = None
|
||||
# Max concurrent outbound tool calls to this server; excess calls queue.
|
||||
# None or a value <= 0 means unlimited.
|
||||
max_concurrent_requests: Optional[int] = None
|
||||
# Resolved short-ID tool prefix when LITELLM_USE_SHORT_MCP_TOOL_PREFIX is
|
||||
# enabled. Set by ``MCPServerManager._assign_unique_short_prefix`` at
|
||||
# registration time so that natural-hash collisions between two
|
||||
|
|
|
|||
|
|
@ -3311,6 +3311,7 @@ class LlmProviders(str, Enum):
|
|||
CUSTOM = "custom"
|
||||
LITELLM_PROXY = "litellm_proxy"
|
||||
HOSTED_VLLM = "hosted_vllm"
|
||||
TENCENT = "tencent"
|
||||
LLAMAFILE = "llamafile"
|
||||
LM_STUDIO = "lm_studio"
|
||||
GALADRIEL = "galadriel"
|
||||
|
|
|
|||
|
|
@ -4204,6 +4204,13 @@ def get_optional_params(
|
|||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
)
|
||||
elif custom_llm_provider == "tencent":
|
||||
optional_params = litellm.TencentChatConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model=model,
|
||||
drop_params=(drop_params if drop_params is not None and isinstance(drop_params, bool) else False),
|
||||
)
|
||||
elif custom_llm_provider == "openrouter":
|
||||
optional_params = litellm.OpenrouterConfig().map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
|
|
@ -6017,6 +6024,11 @@ def validate_environment(
|
|||
keys_in_environment = True
|
||||
else:
|
||||
missing_keys.append("DEEPSEEK_API_KEY")
|
||||
elif custom_llm_provider == "tencent":
|
||||
if "TENCENT_API_KEY" in os.environ:
|
||||
keys_in_environment = True
|
||||
else:
|
||||
missing_keys.append("TENCENT_API_KEY")
|
||||
elif custom_llm_provider == "mistral":
|
||||
if "MISTRAL_API_KEY" in os.environ:
|
||||
keys_in_environment = True
|
||||
|
|
@ -7558,6 +7570,7 @@ class ProviderConfigManager:
|
|||
),
|
||||
# Simple provider mappings (no model parameter needed)
|
||||
LlmProviders.DEEPSEEK: (lambda: litellm.DeepSeekChatConfig(), False),
|
||||
LlmProviders.TENCENT: (lambda: litellm.TencentChatConfig(), False),
|
||||
LlmProviders.GROQ: (lambda: litellm.GroqChatConfig(), False),
|
||||
LlmProviders.BEDROCK_MANTLE: (
|
||||
lambda: litellm.BedrockMantleChatConfig(),
|
||||
|
|
@ -7996,6 +8009,12 @@ class ProviderConfigManager:
|
|||
)
|
||||
|
||||
return DeepSeekAnthropicMessagesConfig()
|
||||
elif litellm.LlmProviders.TENCENT == provider:
|
||||
from litellm.llms.tencent.messages.transformation import (
|
||||
TencentAnthropicMessagesConfig,
|
||||
)
|
||||
|
||||
return TencentAnthropicMessagesConfig()
|
||||
elif litellm.LlmProviders.GITHUB_COPILOT == provider:
|
||||
if "claude" in model_lower:
|
||||
from litellm.llms.github_copilot.messages.transformation import (
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -2295,6 +2295,24 @@
|
|||
"text_completion": true
|
||||
}
|
||||
},
|
||||
"tencent": {
|
||||
"display_name": "Tencent TokenHub (`tencent`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/tencent",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": false,
|
||||
"text_completion": false
|
||||
}
|
||||
},
|
||||
"text-completion-codestral": {
|
||||
"display_name": "Text Completion Codestral (`text-completion-codestral`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/codestral",
|
||||
|
|
|
|||
|
|
@ -62,8 +62,8 @@ proxy = [
|
|||
"azure-identity>=1.25.2,<2.0",
|
||||
"azure-storage-blob>=12.28.0,<13.0",
|
||||
"mcp>=1.26.0,<2.0",
|
||||
"litellm-proxy-extras==0.4.74",
|
||||
"litellm-enterprise==0.1.45",
|
||||
"litellm-proxy-extras==0.4.75",
|
||||
"litellm-enterprise==0.1.46",
|
||||
"RestrictedPython>=8.1,<9.0",
|
||||
"rich>=13.9.4,<14.0",
|
||||
"polars>=1.38.1,<2.0",
|
||||
|
|
|
|||
|
|
@ -222,7 +222,7 @@
|
|||
"limit": 38
|
||||
},
|
||||
"RET504": {
|
||||
"limit": 722
|
||||
"limit": 721
|
||||
},
|
||||
"RUF010": {
|
||||
"limit": 874
|
||||
|
|
@ -249,7 +249,7 @@
|
|||
"limit": 6
|
||||
},
|
||||
"RUF059": {
|
||||
"limit": 74
|
||||
"limit": 73
|
||||
},
|
||||
"RUF100": {
|
||||
"limit": 480
|
||||
|
|
@ -264,7 +264,7 @@
|
|||
"limit": 63
|
||||
},
|
||||
"SIM102": {
|
||||
"limit": 326
|
||||
"limit": 325
|
||||
},
|
||||
"SIM103": {
|
||||
"limit": 129
|
||||
|
|
@ -306,7 +306,7 @@
|
|||
"limit": 9
|
||||
},
|
||||
"TID251": {
|
||||
"limit": 2714
|
||||
"limit": 2710
|
||||
},
|
||||
"TRY002": {
|
||||
"limit": 548
|
||||
|
|
@ -324,7 +324,7 @@
|
|||
"limit": 883
|
||||
},
|
||||
"UP006": {
|
||||
"limit": 13041
|
||||
"limit": 12870
|
||||
},
|
||||
"UP007": {
|
||||
"limit": 2570
|
||||
|
|
@ -354,7 +354,7 @@
|
|||
"limit": 4
|
||||
},
|
||||
"UP035": {
|
||||
"limit": 2300
|
||||
"limit": 2295
|
||||
},
|
||||
"UP036": {
|
||||
"limit": 4
|
||||
|
|
|
|||
|
|
@ -338,6 +338,7 @@ model LiteLLM_MCPServerTable {
|
|||
byok_api_key_help_url String?
|
||||
source_url String?
|
||||
timeout Float?
|
||||
max_concurrent_requests Int?
|
||||
// BYOM submission lifecycle
|
||||
approval_status String? @default("active")
|
||||
submitted_by String?
|
||||
|
|
|
|||
|
|
@ -111,7 +111,7 @@ if [ -n "$spec_files" ]; then
|
|||
# and an up-to-date Prisma client; check-ui-api-types.yml installs those and runs
|
||||
# prisma generate before gen:api, so mirror that here or a stale client can mask
|
||||
# drift that CI will still flag.
|
||||
if ! uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma; then
|
||||
if ! uv run --no-sync python scripts/prisma_generate_if_needed.py; then
|
||||
echo "✗ Could not regenerate Prisma client (prisma generate failed)." >&2
|
||||
status=1
|
||||
elif ( cd ui/litellm-dashboard && LITELLM_PYTHON="uv run --no-sync python" npm run gen:api ); then
|
||||
|
|
|
|||
69
scripts/prisma_generate_if_needed.py
Normal file
69
scripts/prisma_generate_if_needed.py
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Run ``prisma generate`` only when its inputs changed since the last run.
|
||||
|
||||
The generated client is a pure function of ``litellm/proxy/schema.prisma`` and
|
||||
the installed prisma package version, so a stamp of those two written next to
|
||||
the venv is enough to prove the client is current. The stamp lives under
|
||||
``sys.prefix`` so recreating the venv discards it, and a missing generated
|
||||
client (a fresh or reinstalled prisma package) forces a regenerate even when
|
||||
the stamp matches. The prisma package itself is never imported here: once
|
||||
generated it re-exports the whole client on import, which costs more than the
|
||||
generate this script exists to skip.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import importlib.metadata
|
||||
import importlib.util
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent
|
||||
SCHEMA = REPO_ROOT / "litellm" / "proxy" / "schema.prisma"
|
||||
STAMP = Path(sys.prefix) / "litellm-prisma-schema.stamp"
|
||||
|
||||
|
||||
def stamp_value(schema_bytes: bytes, prisma_version: str) -> str:
|
||||
return f"{hashlib.sha256(schema_bytes).hexdigest()}:{prisma_version}"
|
||||
|
||||
|
||||
def should_skip(stamp: Path, expected: str, client_generated: bool) -> bool:
|
||||
if not client_generated:
|
||||
return False
|
||||
try:
|
||||
return stamp.read_text() == expected
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
def client_is_generated() -> bool:
|
||||
spec = importlib.util.find_spec("prisma")
|
||||
if spec is None or not spec.submodule_search_locations:
|
||||
return False
|
||||
return any(
|
||||
(Path(location) / "client.py").exists()
|
||||
for location in spec.submodule_search_locations
|
||||
)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
version = importlib.metadata.version("prisma")
|
||||
expected = stamp_value(SCHEMA.read_bytes(), version)
|
||||
if should_skip(STAMP, expected, client_is_generated()):
|
||||
print(
|
||||
f"Prisma client already generated for {SCHEMA.relative_to(REPO_ROOT)} "
|
||||
f"(prisma {version}); skipping prisma generate"
|
||||
)
|
||||
return 0
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-m", "prisma", "generate", "--schema", str(SCHEMA)],
|
||||
cwd=REPO_ROOT,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
return result.returncode
|
||||
STAMP.write_text(expected)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
|
@ -89,6 +89,18 @@ def base_counts(ref: str) -> dict:
|
|||
shutil.rmtree(parent, ignore_errors=True)
|
||||
|
||||
|
||||
def over_ceiling(head: dict, budget: dict) -> frozenset:
|
||||
"""Rules whose head count already exceeds their limit.
|
||||
|
||||
A rule can only breach when it is over its limit, so when none are the base
|
||||
comparison cannot change the verdict and the base worktree scan can be skipped.
|
||||
"""
|
||||
return frozenset(
|
||||
rule for rule, spec in budget.items()
|
||||
if head.get(rule, 0) > spec["limit"]
|
||||
)
|
||||
|
||||
|
||||
def evaluate(head: dict, base: dict, budget: dict) -> list:
|
||||
breaches = []
|
||||
for rule, spec in budget.items():
|
||||
|
|
@ -119,8 +131,12 @@ def introduced(violations: list, changed: dict) -> list:
|
|||
def cmd_check(base: str) -> None:
|
||||
budget = json.loads(BUDGET_PATH.read_text())
|
||||
head = head_violations()
|
||||
head_counts = count_by_rule(head)
|
||||
if not over_ceiling(head_counts, budget):
|
||||
print(f"OK: every strict rule is within its codebase ceiling (base {base})")
|
||||
return
|
||||
base_point = _run(["git", "merge-base", base, "HEAD"]).strip() or base
|
||||
breaches = evaluate(count_by_rule(head), base_counts(base_point), budget)
|
||||
breaches = evaluate(head_counts, base_counts(base_point), budget)
|
||||
if not breaches:
|
||||
print(f"OK: every strict rule is within its codebase ceiling (base {base})")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -13,9 +13,13 @@ bystander's count equals its base, so it is spared, while any PR that actually
|
|||
grows the rule past its limit still fails.
|
||||
|
||||
Head counts are read from stdin (the caller runs basedpyright once and pipes
|
||||
``--outputjson`` in); the base count is a second basedpyright pass over a
|
||||
detached worktree at the merge-base, run under the same environment so import
|
||||
resolution matches. ``--update`` ratchets each rule's ``limit`` down by the
|
||||
``--outputjson`` in). The base count only matters once some rule is over its
|
||||
limit, so when none is the base pass is skipped outright. When it is needed, it
|
||||
is a second basedpyright pass over a detached worktree at the merge-base, run
|
||||
under the same environment so import resolution matches, and its per-rule
|
||||
counts are cached under the repo's git common dir keyed by merge-base commit,
|
||||
``pyrightconfig.json``, and ``uv.lock``, so re-runs against the same branch
|
||||
point pay for it once. ``--update`` ratchets each rule's ``limit`` down by the
|
||||
number of errors this branch fixed relative to its branch point (the merge-base),
|
||||
so the headroom you were granted shrinks by exactly what you cleared and never
|
||||
grows.
|
||||
|
|
@ -28,20 +32,24 @@ carries an unambiguous ``rule`` field.
|
|||
|
||||
import argparse
|
||||
import contextlib
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from collections import Counter
|
||||
from collections.abc import Iterator, Mapping
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import NamedTuple
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent
|
||||
BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
|
||||
PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json"
|
||||
UV_LOCK = REPO_ROOT / "uv.lock"
|
||||
DEFAULT_BASE = "origin/litellm_internal_staging"
|
||||
CACHE_FILE_PREFIX = "basedpyright-base-"
|
||||
|
||||
# Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated.
|
||||
UNCODED = "<uncoded>"
|
||||
|
|
@ -129,6 +137,109 @@ def base_counts(ref: str) -> dict[str, int]:
|
|||
return count_basedpyright(proc.stdout, root=worktree)
|
||||
|
||||
|
||||
def over_ceiling(
|
||||
head: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]
|
||||
) -> frozenset[str]:
|
||||
"""Rules whose head count already exceeds their limit.
|
||||
|
||||
A rule can only breach when it is over its limit, so when none are the base
|
||||
comparison cannot change the verdict and the base worktree pass can be skipped.
|
||||
"""
|
||||
return frozenset(
|
||||
code
|
||||
for code, total in head.items()
|
||||
if total > (budget[code]["limit"] if code in budget else DEFAULT_LIMIT)
|
||||
)
|
||||
|
||||
|
||||
def environment_fingerprints() -> tuple[str, ...]:
|
||||
return tuple(
|
||||
hashlib.sha256(path.read_bytes()).hexdigest()
|
||||
for path in (PYRIGHT_CONFIG, UV_LOCK)
|
||||
if path.exists()
|
||||
)
|
||||
|
||||
|
||||
def cache_key(base_point: str, fingerprints: tuple[str, ...]) -> str:
|
||||
return hashlib.sha256("|".join((base_point, *fingerprints)).encode()).hexdigest()[
|
||||
:16
|
||||
]
|
||||
|
||||
|
||||
def cache_path(
|
||||
directory: Path, base_point: str, fingerprints: tuple[str, ...]
|
||||
) -> Path:
|
||||
return directory / f"{CACHE_FILE_PREFIX}{cache_key(base_point, fingerprints)}.json"
|
||||
|
||||
|
||||
def default_cache_dir() -> Path:
|
||||
common = Path(_run(["git", "rev-parse", "--git-common-dir"]).strip())
|
||||
resolved = common if common.is_absolute() else REPO_ROOT / common
|
||||
return resolved / "litellm-lint-cache"
|
||||
|
||||
|
||||
def load_cached_counts(path: Path) -> dict[str, int] | None:
|
||||
try:
|
||||
data = json.loads(path.read_text())
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return None
|
||||
counts = data.get("counts") if isinstance(data, dict) else None
|
||||
if not isinstance(counts, dict):
|
||||
return None
|
||||
if not all(
|
||||
isinstance(code, str) and isinstance(total, int) and not isinstance(total, bool)
|
||||
for code, total in counts.items()
|
||||
):
|
||||
return None
|
||||
return counts
|
||||
|
||||
|
||||
def scratch_path(path: Path) -> Path:
|
||||
"""In-flight scratch for the tmp+rename write. Dot-prefixed so the prune
|
||||
glob in `store_counts` can never match it (a concurrent run would otherwise
|
||||
unlink it between write and rename), and pid-suffixed so two concurrent
|
||||
writers of the same entry never share a scratch."""
|
||||
return path.with_name(f".{path.name}.{os.getpid()}.tmp")
|
||||
|
||||
|
||||
def store_counts(
|
||||
directory: Path, path: Path, base_point: str, counts: Mapping[str, int]
|
||||
) -> None:
|
||||
directory.mkdir(parents=True, exist_ok=True)
|
||||
for stale in directory.glob(f"{CACHE_FILE_PREFIX}*.json"):
|
||||
if stale != path:
|
||||
stale.unlink(missing_ok=True)
|
||||
scratch = scratch_path(path)
|
||||
scratch.write_text(
|
||||
json.dumps(
|
||||
{"base_point": base_point, "counts": dict(sorted(counts.items()))},
|
||||
indent=2,
|
||||
)
|
||||
+ "\n"
|
||||
)
|
||||
scratch.replace(path)
|
||||
|
||||
|
||||
def base_counts_cached(
|
||||
base_point: str,
|
||||
cache_dir: Path | None = None,
|
||||
compute: Callable[[str], dict[str, int]] = base_counts,
|
||||
) -> dict[str, int]:
|
||||
"""`base_counts` memoized on disk. The base tree at a given commit is
|
||||
immutable, so its counts are a pure function of the merge-base plus the
|
||||
environment fingerprints in the cache key; an empty result is never stored
|
||||
because it is the signature of a crashed pass, not a clean tree."""
|
||||
directory = default_cache_dir() if cache_dir is None else cache_dir
|
||||
path = cache_path(directory, base_point, environment_fingerprints())
|
||||
cached = load_cached_counts(path)
|
||||
if cached is not None:
|
||||
return cached
|
||||
counts = compute(base_point)
|
||||
if counts:
|
||||
store_counts(directory, path, base_point, counts)
|
||||
return counts
|
||||
|
||||
|
||||
def evaluate(
|
||||
head: Mapping[str, int],
|
||||
base: Mapping[str, int],
|
||||
|
|
@ -185,7 +296,7 @@ def cmd_update(current: Mapping[str, int], base_ref: str = DEFAULT_BASE) -> None
|
|||
"""
|
||||
budget = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {}
|
||||
base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref
|
||||
updated = ratcheted_budget(budget, current, base_counts(base_point))
|
||||
updated = ratcheted_budget(budget, current, base_counts_cached(base_point))
|
||||
BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n")
|
||||
cleared = sum(budget[code]["limit"] - updated[code]["limit"] for code in updated)
|
||||
print(
|
||||
|
|
@ -205,8 +316,13 @@ def cmd_check(base_ref: str) -> None:
|
|||
f"nothing; refusing to certify a vacuous run."
|
||||
)
|
||||
raise SystemExit(1)
|
||||
if not over_ceiling(head, budget):
|
||||
print(
|
||||
f"OK: every rule is within its basedpyright limit ({sum(head.values())} errors total)"
|
||||
)
|
||||
return
|
||||
base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref
|
||||
base = base_counts(base_point)
|
||||
base = base_counts_cached(base_point)
|
||||
if is_vacuous_run(base, budget):
|
||||
print(
|
||||
f"FAIL: basedpyright produced no errors for the base tree at "
|
||||
|
|
|
|||
|
|
@ -9,15 +9,18 @@ When contributing to this directory, please first discuss the change you wish to
|
|||
|
||||
## Setup
|
||||
|
||||
The suites run against a live proxy, so bring one up first. `docker-compose.yml` here starts that proxy with its Postgres and Redis, serving `gateway/litellm-config.yml`; add any model, pricing override, or guardrail your test needs to that file and read it back in the test rather than hardcoding values. `gateway/` holds proxy configuration only, so never put tests there
|
||||
The suites run against a live proxy, so bring one up first. `docker-compose.yml` here starts that proxy with a throwaway Postgres and Redis; `docker compose down -v` resets everything, so no state leaks between runs. The proxy config is inlined in the compose file under `configs`, prewired with example models (`gpt-5.5`, `claude-haiku-4-5`, `gemini-2.5-flash`, `openai-text-embedding-3-small`) whose keys come from your `.env`. If your test needs another model, a pricing override, or a guardrail declared up front, add it to that inline config and read it back in the test rather than hardcoding values
|
||||
|
||||
## Running the tests locally
|
||||
|
||||
1. Create a .env file and add provider keys:
|
||||
1. Create a `.env` file in this directory with the provider keys the example models use:
|
||||
|
||||
```bash
|
||||
OPENAI_API_KEY="sk-..."
|
||||
ANTHROPIC_API_KEY="sk-..."
|
||||
|
||||
GEMINI_API_KEY="..."
|
||||
```
|
||||
|
||||
2. Bring the stack up from this directory:
|
||||
|
||||
```bash
|
||||
|
|
@ -123,7 +126,17 @@ Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure cover
|
|||
|
||||
Before you push
|
||||
|
||||
- Run basedpyright over your changes; the harness is fully typed and new code must not add `Any` or widen the budgets
|
||||
- Bring the stack up with docker-compose from this directory and run your suite locally against it, so you exercise the same skip-vs-fail path CI does
|
||||
- Use the config at `tests/e2e/gateway/litellm-config.yml` if your feature needs a model, pricing override, guardrail, or other proxy setting declared up front; add the deployment there and read it back in the test rather than hardcoding values
|
||||
- Capture screenshots of the tests passing and attach them to the PR as proof of fix
|
||||
1. Run basedpyright over your changes; the harness is fully typed and new code must not add `Any` or widen the budgets
|
||||
|
||||
2. Add the models your test needs to the inline config in `docker-compose.yml`
|
||||
|
||||
3. Bring the stack up and run your suite against it:
|
||||
|
||||
```bash
|
||||
docker compose up -d
|
||||
uv run pytest tests/e2e/<your_suite>/ -v
|
||||
```
|
||||
|
||||
4. Capture screenshots of the test run and attach them to the PR as proof
|
||||
|
||||
5. If a test fails because it surfaced a real issue in the product, flag that explicitly in the PR rather than reworking the test until it passes
|
||||
|
|
|
|||
87
tests/e2e/docker-compose.yml
Normal file
87
tests/e2e/docker-compose.yml
Normal file
|
|
@ -0,0 +1,87 @@
|
|||
# local setup to run e2e tests
|
||||
configs:
|
||||
litellm_config:
|
||||
content: |
|
||||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
database_url: os.environ/DATABASE_URL
|
||||
store_prompts_in_spend_logs: true
|
||||
|
||||
litellm_settings:
|
||||
drop_params: true
|
||||
num_retries: 3
|
||||
request_timeout: 600
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
host: redis
|
||||
port: 6379
|
||||
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
num_retries: 3
|
||||
allowed_fails: 5
|
||||
cooldown_time: 30
|
||||
fallbacks:
|
||||
- gemini-2.5-flash: ["gpt-5.5", "claude-haiku-4-5"]
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-5.5
|
||||
litellm_params:
|
||||
model: openai/gpt-5.5
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: claude-haiku-4-5
|
||||
litellm_params:
|
||||
model: anthropic/claude-haiku-4-5
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
|
||||
- model_name: gemini-2.5-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
- model_name: openai-text-embedding-3-small
|
||||
litellm_params:
|
||||
model: openai/text-embedding-3-small
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
services:
|
||||
litellm:
|
||||
image: ghcr.io/berriai/litellm:main-latest
|
||||
depends_on:
|
||||
db:
|
||||
condition: service_healthy
|
||||
redis:
|
||||
condition: service_healthy
|
||||
env_file: .env
|
||||
environment:
|
||||
LITELLM_MASTER_KEY: sk-1234
|
||||
DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm
|
||||
ports:
|
||||
- "4000:4000"
|
||||
configs:
|
||||
- source: litellm_config
|
||||
target: /app/config.yaml
|
||||
command: ["--config", "/app/config.yaml", "--port", "4000"]
|
||||
|
||||
# throwaway db
|
||||
db:
|
||||
image: postgres:16
|
||||
environment:
|
||||
POSTGRES_USER: litellm
|
||||
POSTGRES_PASSWORD: litellm
|
||||
POSTGRES_DB: litellm
|
||||
healthcheck:
|
||||
test: ["CMD-SHELL", "pg_isready -U litellm"]
|
||||
interval: 3s
|
||||
timeout: 3s
|
||||
retries: 20
|
||||
|
||||
redis:
|
||||
image: redis:7
|
||||
healthcheck:
|
||||
test: ["CMD", "redis-cli", "ping"]
|
||||
interval: 3s
|
||||
timeout: 3s
|
||||
retries: 20
|
||||
|
|
@ -1,287 +0,0 @@
|
|||
# This default config file aims to support most popular model providers out of the box
|
||||
|
||||
#In general, the model name used by the client will be the same as the ones from the provider (For example, you will use "anthropic.claude-3-5-sonnet-20240620-v1:0" when you're calling LiteLLM just like you would when calling Amazon Bedrock directly)
|
||||
#In the case where there are model name conflicts, a prefix will be used (For example, the Azure and the openAI model names conflict, so when you are using Azure, you will use "azure/gpt-4o-realtime-preview-2024-10-01")
|
||||
|
||||
#Some model providers require additional user-specific configuration (such as Azure which requires you to specify your own api_base with your resource name, and your api_version).
|
||||
#In this case, the provider is commented out, and you should uncomment it and provide your specific info
|
||||
|
||||
#For more detailed information about each provider, refer to the docs: https://docs.litellm.ai/docs/providers
|
||||
|
||||
#If you are not interested in a particular provider, just remove it from your config.yaml, and redeploy, and it will no longer show up in your LiteLLM deployment
|
||||
|
||||
#If a particular provider is not working, double check your .env file, and make sure you have provided a valid api key for that provider, and then redeploy
|
||||
|
||||
#Full details on guardrails here: https://docs.litellm.ai/docs/proxy/guardrails/bedrock
|
||||
general_settings:
|
||||
store_prompts_in_spend_logs: true
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
proxy_batch_write_at: 60
|
||||
database_connection_pool_limit: 10
|
||||
# disable_error_logs: True
|
||||
forward_client_headers_to_llm_api: false
|
||||
maximum_spend_logs_retention_period: "60d" # GSE-13389: Cleanup logs older than 60 days
|
||||
maximum_spend_logs_cleanup_cron: "0 1 * * *" # 01:00 UTC daily = 18:00 PDT
|
||||
database_url: os.environ/DATABASE_URL
|
||||
control_plane_url: os.environ/CONTROL_PLANE_URL
|
||||
alerts: ["email"]
|
||||
|
||||
proxy_budget_rescheduler_min_time: 15
|
||||
proxy_budget_rescheduler_max_time: 20
|
||||
|
||||
# fallbacks: [{"gpt-4": ["anthropic.claude-3-5-sonnet-20240620-v1:0"]}] #Configure fallbacks for context window exeeded errors (In this example, we will fall back to Claude Sonnet if over 8000 tokens, which is gpt-4's limit)
|
||||
# default_fallbacks: ["anthropic.claude-3-haiku-20240307-v1:0"] #Configure fallbacks for any error for every model (the above fallback configurations override this one)
|
||||
# environment_variables:
|
||||
# STORE_MODEL_IN_DB: 'True'
|
||||
# LITELLM_LOG: "DEBUG"
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
# Spend counters inherit this as their Redis TTL, so an idle counter goes cold and
|
||||
# the next request reseeds it from the DB; kept short to exercise the cross-pod
|
||||
# reseed path in test_spend_counter_reseed_e2e. Response-cache writes pass their own
|
||||
# ttl and are unaffected.
|
||||
default_redis_ttl: 20
|
||||
request_timeout: 600
|
||||
num_retries: 3
|
||||
json_logs: true
|
||||
store_audit_logs: True
|
||||
cache: true
|
||||
cache_params:
|
||||
type: redis
|
||||
host: redis
|
||||
port: 6379
|
||||
password: os.environ/REDIS_PASSWORD
|
||||
namespace: litellm.caching
|
||||
ttl: 16600
|
||||
# max_budget: 1000000000.0 # (float) sets max budget in dollars across the entire proxy across all API keys. Note, the budget does not apply to the master key. That is the only exception.
|
||||
# budget_duration: 1mo # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
|
||||
# max_internal_user_budget: 1000000000.0 # (float) sets default budget in dollars for each internal user. (Doesn't apply to Admins. Doesn't apply to Teams. Doesn't apply to master key)
|
||||
# internal_user_budget_duration: "1mo" # (str) frequency of budget reset - You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
|
||||
# success_callback: ["s3_v2"]
|
||||
# failure_callback: ["s3_v2"]
|
||||
# service_callback: ["datadog"]
|
||||
callbacks: ["arize_phoenix", "datadog", "smtp_email", "prometheus", "otel"]
|
||||
require_auth_for_metrics_endpoint: false
|
||||
#type: redis-semantic
|
||||
#similarity_threshold: 0.8 # similarity threshold for semantic cache
|
||||
#redis_semantic_cache_embedding_model: text-embedding-ada-002 # only works with text-embedding-ada-002 for now... https://github.com/BerriAI/litellm/issues/4001
|
||||
|
||||
router_settings:
|
||||
routing_strategy: simple-shuffle
|
||||
num_retries: 3
|
||||
allowed_fails: 5
|
||||
cooldown_time: 30
|
||||
# When gemini deployments are exhausted (provider 429 / auth), cross over to
|
||||
# working models. Exercised by tests/e2e/router/test_rate_limiter.py.
|
||||
fallbacks:
|
||||
- gemini-2.5-flash: ["gpt-5.5", "claude-haiku-4-5"]
|
||||
|
||||
#ttl: Optional[float]
|
||||
#default_in_memory_ttl: Optional[float]
|
||||
#default_in_redis_ttl: Optional[float]
|
||||
|
||||
model_list:
|
||||
- model_name: gpt-5.5
|
||||
litellm_params:
|
||||
model: openai/gpt-5.5
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: claude-haiku-4-5
|
||||
litellm_params:
|
||||
model: anthropic/claude-haiku-4-5
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
|
||||
# Same underlying model via Vertex AI — distinct routing/auth path
|
||||
# # (service-account JSON), so it gets its own model_name.
|
||||
- model_name: gemini-2.5-flash-vertex
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-2.5-flash
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
|
||||
|
||||
- model_name: gemini-2.5-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
# load balancing to a different deployment, if gemini gets rate limited.
|
||||
- model_name: gemini-2.5-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
# Custom per-token pricing exercised by llm_translation/test_custom_pricing_e2e.py.
|
||||
# Rates deliberately exceed canonical gemini-2.5-flash (input 3e-7 / output 2.5e-6)
|
||||
# so an override that is ignored or under-applied reports spend at the base rate
|
||||
# and fails that test. The test reads these same rates back from this file.
|
||||
- model_name: custom-priced-flash
|
||||
litellm_params:
|
||||
model: gemini/gemini-2.5-flash
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
input_cost_per_token: 0.00005
|
||||
output_cost_per_token: 0.0001
|
||||
|
||||
# embedding models
|
||||
- model_name: openai-text-embedding-3-small
|
||||
litellm_params:
|
||||
model: openai/text-embedding-3-small
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: gemini-2-embedding
|
||||
litellm_params:
|
||||
model: gemini/gemini-2-embedding
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
|
||||
- model_name: openai-realtime
|
||||
litellm_params:
|
||||
model: openai/gpt-realtime
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
model_info:
|
||||
mode: realtime
|
||||
|
||||
- model_name: azure-realtime
|
||||
litellm_params:
|
||||
model: azure/gpt-realtime-2
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
api_base: os.environ/AZURE_API_BASE
|
||||
api_version: "2025-08-28"
|
||||
realtime_protocol: GA # Possible values: "GA"/ "v1", "beta"
|
||||
model_info:
|
||||
mode: realtime
|
||||
|
||||
- model_name: gemini-realtime
|
||||
litellm_params:
|
||||
model: gemini/gemini-3.1-flash-live-preview
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
model_info:
|
||||
mode: realtime
|
||||
|
||||
- model_name: vertex-realtime
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-live-2.5-flash-preview-native-audio-09-2025
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
|
||||
model_info:
|
||||
mode: realtime
|
||||
|
||||
- model_name: bedrock-realtime
|
||||
litellm_params:
|
||||
model: bedrock/amazon.nova-sonic-v1:0
|
||||
aws_region_name: us-east-1
|
||||
model_info:
|
||||
mode: realtime
|
||||
|
||||
- model_name: xai-realtime
|
||||
litellm_params:
|
||||
model: xai/grok-voice-latest
|
||||
api_key: os.environ/XAI_API_KEY
|
||||
model_info:
|
||||
mode: realtime
|
||||
- model_name: rust-ocr-mistral
|
||||
litellm_params:
|
||||
model: mistral/mistral-ocr-latest
|
||||
api_key: os.environ/MISTRAL_API_KEY
|
||||
|
||||
- model_name: rust-ocr-azure-ai
|
||||
litellm_params:
|
||||
model: azure_ai/mistral-document-ai-2505
|
||||
api_base: os.environ/AZURE_API_BASE
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
|
||||
- model_name: rust-ocr-azure-document-intelligence
|
||||
litellm_params:
|
||||
model: azure_ai/doc-intelligence/prebuilt-layout
|
||||
api_base: os.environ/AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT
|
||||
api_key: os.environ/AZURE_DOCUMENT_INTELLIGENCE_API_KEY
|
||||
|
||||
- model_name: rust-ocr-vertex-mistral
|
||||
litellm_params:
|
||||
model: vertex_ai/mistral-ocr-2505
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
|
||||
- model_name: rust-ocr-vertex-deepseek
|
||||
litellm_params:
|
||||
model: vertex_ai/deepseek-ocr-maas
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
|
||||
# batch models exercised by tests/e2e/batches/
|
||||
- model_name: openai-batch
|
||||
litellm_params:
|
||||
model: openai/gpt-4o-mini
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
model_info:
|
||||
mode: batch
|
||||
|
||||
- model_name: azure-batch
|
||||
litellm_params:
|
||||
model: azure/gpt-4.1-mini-batch
|
||||
api_base: os.environ/AZURE_API_BASE
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
api_version: "2024-07-01-preview"
|
||||
model_info:
|
||||
mode: batch
|
||||
|
||||
- model_name: vertex-batch
|
||||
litellm_params:
|
||||
model: vertex_ai/gemini-2.5-flash
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
|
||||
bucket_name: os.environ/GCS_BUCKET_NAME
|
||||
model_info:
|
||||
mode: batch
|
||||
|
||||
- model_name: bedrock-batch
|
||||
litellm_params:
|
||||
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
|
||||
s3_bucket_name: os.environ/AWS_BATCH_S3_BUCKET
|
||||
s3_region_name: us-west-2
|
||||
s3_access_key_id: os.environ/AWS_ACCESS_KEY_ID
|
||||
s3_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
|
||||
aws_batch_role_arn: os.environ/AWS_BATCH_ROLE_ARN
|
||||
model_info:
|
||||
mode: batch
|
||||
|
||||
files_settings:
|
||||
- custom_llm_provider: openai
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
- custom_llm_provider: azure
|
||||
api_base: os.environ/AZURE_API_BASE
|
||||
api_key: os.environ/AZURE_API_KEY
|
||||
api_version: "2024-07-01-preview"
|
||||
- custom_llm_provider: vertex_ai
|
||||
vertex_project: os.environ/VERTEXAI_PROJECT
|
||||
vertex_location: us-central1
|
||||
vertex_credentials: os.environ/VERTEXAI_CREDENTIALS
|
||||
bucket_name: os.environ/GCS_BUCKET_NAME
|
||||
|
||||
|
||||
mcp_servers:
|
||||
deepwiki_mcp:
|
||||
url: "https://mcp.deepwiki.com/mcp"
|
||||
auth_type: none
|
||||
description: "just a test"
|
||||
|
||||
atlassian:
|
||||
url: "https://mcp.atlassian.com/v1/mcp"
|
||||
auth_type: oauth2
|
||||
authorization_url: https://auth.atlassian.com/authorize
|
||||
|
||||
|
||||
guardrails:
|
||||
- guardrail_name: "presidio-pii"
|
||||
litellm_params:
|
||||
guardrail: presidio
|
||||
mode: pre_call
|
||||
presidio_analyzer_api_base: os.environ/PRESIDIO_ANALYZER_API_BASE
|
||||
presidio_anonymizer_api_base: os.environ/PRESIDIO_ANONYMIZER_API_BASE
|
||||
default_on: false
|
||||
pii_entities_config:
|
||||
EMAIL_ADDRESS: BLOCK
|
||||
CREDIT_CARD: BLOCK
|
||||
US_SSN: BLOCK
|
||||
PHONE_NUMBER: BLOCK
|
||||
|
|
@ -28,7 +28,7 @@ Status: `covered` / `partial` / `gap`.
|
|||
|----------|---------------|-----------|------------|-------------|--------|
|
||||
| Gemini (`/gemini/v1beta/models/{m}:generateContent` / `:streamGenerateContent`) | live | live | live | live | **covered** |
|
||||
| Anthropic (`/anthropic/v1/messages`) | live | live | live | live | **covered** |
|
||||
| Vertex AI (`/vertex_ai/...`) | - | - | - | - | gap (gcloud auth) |
|
||||
| Vertex AI (`/vertex_ai/v1/projects/{p}/locations/{loc}/.../models/{m}:generateContent`) | live | - | - | live | **partial** |
|
||||
| OpenAI / Bedrock / Cohere / Mistral / VLLM | - | - | - | - | gap |
|
||||
|
||||
Each covered cell asserts: `call_type == "pass_through_endpoint"`, `spend > 0`,
|
||||
|
|
@ -60,11 +60,21 @@ most likely to silently break and the one a mock can't prove works.
|
|||
| `test_anthropic_passthrough_nonstreaming_logs_cost` | anthropic native, non-stream, cost |
|
||||
| `test_anthropic_passthrough_streaming_logs_cost` | anthropic native, stream, cost |
|
||||
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
|
||||
| `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost |
|
||||
|
||||
Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is
|
||||
added at runtime instead of declared in the gateway config: the test POSTs `/model/new`
|
||||
with `use_in_pass_through`, so the proxy registers that deployment's service account for
|
||||
the `/vertex_ai` route, then deletes it on teardown. The passthrough call sends only its
|
||||
litellm virtual key (`x-litellm-api-key`), no upstream bearer, and the proxy mints the
|
||||
Vertex token itself. Credentials (`VERTEXAI_PROJECT` / `VERTEXAI_CREDENTIALS`) are read
|
||||
from the same env the proxy uses, so the test never mints a token.
|
||||
|
||||
## Gaps
|
||||
|
||||
- Vertex / OpenAI / Bedrock / Cohere passthrough (same shape; add once the
|
||||
provider credential is configured; Vertex is closest - route exists, auth stale).
|
||||
- Vertex streaming / tool-call passthrough (non-streaming + cost now covered).
|
||||
- OpenAI / Bedrock / Cohere passthrough (same shape; add once the provider
|
||||
credential is configured).
|
||||
- Non-passthrough tool calls over `/chat/completions` end to end with cost.
|
||||
- Image / audio / rerank / responses / realtime translation + cost.
|
||||
- Streaming cost-injection (`include_cost_in_streaming_usage`); passthrough on
|
||||
|
|
|
|||
|
|
@ -48,6 +48,16 @@ class AnthropicHeaders(Headers):
|
|||
tags: str | None = None
|
||||
|
||||
|
||||
class VertexHeaders(Headers):
|
||||
# Only the litellm virtual key; the /vertex_ai passthrough mints the Vertex token
|
||||
# from the proxy's own service account (the deployment marked use_in_pass_through),
|
||||
# so no upstream Authorization bearer is sent from the client.
|
||||
x_litellm_api_key: str = Field(serialization_alias="x-litellm-api-key")
|
||||
content_type: str = Field(
|
||||
default="application/json", serialization_alias="Content-Type"
|
||||
)
|
||||
|
||||
|
||||
class AltSseParams(BaseModel):
|
||||
alt: str = "sse"
|
||||
|
||||
|
|
@ -132,6 +142,23 @@ class PassthroughClient:
|
|||
stream=True,
|
||||
)
|
||||
|
||||
# ---- Vertex AI native passthrough (/vertex_ai/v1/projects/...) -------
|
||||
|
||||
def vertex_generate(
|
||||
self, key: str, project: str, location: str, model: str, text: str
|
||||
) -> StreamingResponse:
|
||||
path = (
|
||||
f"/vertex_ai/v1/projects/{project}/locations/{location}"
|
||||
f"/publishers/google/models/{model}:generateContent"
|
||||
)
|
||||
return self.gateway.transport.send(
|
||||
path,
|
||||
headers=VertexHeaders(x_litellm_api_key=key),
|
||||
json=GeminiGenerateBody(
|
||||
contents=[GeminiContent(parts=[GeminiPart(text=text)])]
|
||||
),
|
||||
)
|
||||
|
||||
# ---- Anthropic native passthrough (/anthropic/v1/messages) ----------
|
||||
|
||||
def anthropic_message(
|
||||
|
|
|
|||
166
tests/e2e/llm_translation/test_vertex_passthrough_e2e.py
Normal file
166
tests/e2e/llm_translation/test_vertex_passthrough_e2e.py
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
"""Live e2e: a native Vertex AI generateContent call over the proxy's /vertex_ai
|
||||
passthrough is forwarded to Vertex and still logged as a costed SpendLogs row.
|
||||
|
||||
Ports the de-flake of the SDK-based spend test (#31689). That test configured the
|
||||
vertexai SDK with an api_endpoint override pointing at the proxy, but the SDK
|
||||
intermittently ignored the override and billed the public Vertex endpoint directly,
|
||||
so the request never reached LiteLLM and no spend was recorded; the bypass, not
|
||||
logging lag, was the flake. Driving raw HTTP through the shared transport always
|
||||
reaches the proxy, which the harness already guarantees, so the only residual
|
||||
nondeterminism is the ~60s async spend flush the poll absorbs.
|
||||
|
||||
The vertex deployment is added at runtime through the management endpoint rather than
|
||||
declared in the gateway config: the test POSTs /model/new with use_in_pass_through so
|
||||
the proxy registers that deployment's service account for the /vertex_ai route, then
|
||||
deletes it on teardown. The credential is the one the proxy already holds (read from
|
||||
the same VERTEXAI_CREDENTIALS the deployment uses), so the passthrough call sends only
|
||||
its litellm virtual key in x-litellm-api-key and no upstream bearer, and the proxy
|
||||
mints the Vertex token itself. The test never mints a token.
|
||||
|
||||
Asserts both sides of the promise: the forward succeeds (2xx with a candidate) and
|
||||
the costed row lands (call_type pass_through_endpoint, vertex_ai provider, a gemini
|
||||
model, spend > 0), correlated by the x-litellm-call-id header.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_http import NoBody, require_successful_call, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import SpendLogRow
|
||||
from passthrough_client import PassthroughClient
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
VERTEX_MODEL = "gemini-2.5-flash"
|
||||
# The added deployment's region and the passthrough URL's region are the same constant,
|
||||
# so they always agree; the proxy registers passthrough credentials per project+region.
|
||||
VERTEX_LOCATION = os.environ.get("VERTEXAI_LOCATION", "us-central1")
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def vertex_project() -> str:
|
||||
"""The Vertex project to bill, read from the same VERTEXAI_PROJECT the proxy uses.
|
||||
Skip when unset, since that is an environment gap rather than a behavior failure."""
|
||||
project = os.environ.get("VERTEXAI_PROJECT")
|
||||
if not project:
|
||||
pytest.skip("set VERTEXAI_PROJECT (the project the vertex deployment bills)")
|
||||
return project
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def vertex_credentials() -> str:
|
||||
"""The service-account JSON the added deployment authenticates with, the same
|
||||
VERTEXAI_CREDENTIALS the proxy holds. Skip when unset."""
|
||||
credentials = os.environ.get("VERTEXAI_CREDENTIALS")
|
||||
if not credentials:
|
||||
pytest.skip("set VERTEXAI_CREDENTIALS (the vertex service-account JSON)")
|
||||
return credentials
|
||||
|
||||
|
||||
class _VertexDeploymentParams(BaseModel):
|
||||
model: str
|
||||
vertex_project: str
|
||||
vertex_location: str
|
||||
vertex_credentials: str
|
||||
use_in_pass_through: bool
|
||||
|
||||
|
||||
class _ModelInfoId(BaseModel):
|
||||
id: str
|
||||
|
||||
|
||||
class _ModelNewBody(BaseModel):
|
||||
model_name: str
|
||||
litellm_params: _VertexDeploymentParams
|
||||
model_info: _ModelInfoId
|
||||
|
||||
|
||||
class _ModelNewResponse(BaseModel):
|
||||
model_id: str
|
||||
|
||||
|
||||
class _ModelDeleteBody(BaseModel):
|
||||
id: str
|
||||
|
||||
|
||||
def _add_vertex_passthrough_model(
|
||||
client: PassthroughClient, model_name: str, project: str, credentials: str
|
||||
) -> str:
|
||||
return unwrap(
|
||||
client.gateway.transport.post(
|
||||
"/model/new",
|
||||
headers=client.gateway.transport.master,
|
||||
json=_ModelNewBody(
|
||||
model_name=model_name,
|
||||
litellm_params=_VertexDeploymentParams(
|
||||
model=f"vertex_ai/{VERTEX_MODEL}",
|
||||
vertex_project=project,
|
||||
vertex_location=VERTEX_LOCATION,
|
||||
vertex_credentials=credentials,
|
||||
use_in_pass_through=True,
|
||||
),
|
||||
model_info=_ModelInfoId(id=model_name),
|
||||
),
|
||||
response_type=_ModelNewResponse,
|
||||
)
|
||||
).model_id
|
||||
|
||||
|
||||
def _delete_model(client: PassthroughClient, model_id: str) -> None:
|
||||
_ = client.gateway.transport.post(
|
||||
"/model/delete",
|
||||
headers=client.gateway.transport.master,
|
||||
json=_ModelDeleteBody(id=model_id),
|
||||
response_type=NoBody,
|
||||
)
|
||||
|
||||
|
||||
def _costed_row(client: PassthroughClient, call_id: str | None) -> SpendLogRow:
|
||||
"""The passthrough call's SpendLogs row, polled until it carries a cost.
|
||||
|
||||
A 2xx passthrough call that produced no costed row is a hard failure, not a skip:
|
||||
a billed Vertex call that LiteLLM did not track is the exact regression #31689
|
||||
guards against."""
|
||||
assert call_id, "vertex passthrough response had no x-litellm-call-id header"
|
||||
rows = client.gateway.poll_logs_for_request_id(
|
||||
call_id,
|
||||
predicate=lambda rs: (rs[0].spend or 0) > 0,
|
||||
)
|
||||
assert rows, f"no SpendLogs row for vertex passthrough call_id {call_id}"
|
||||
row = rows[0]
|
||||
assert row.call_type == "pass_through_endpoint", f"unexpected call_type: {row}"
|
||||
assert (row.spend or 0) > 0, f"vertex passthrough call was not costed: {row}"
|
||||
assert row.status == "success", f"unexpected status: {row}"
|
||||
return row
|
||||
|
||||
|
||||
class TestVertexPassthroughSpendTracking:
|
||||
def test_vertex_passthrough_via_managed_model_logs_cost(
|
||||
self,
|
||||
client: PassthroughClient,
|
||||
scoped_key: str,
|
||||
resources: ResourceManager,
|
||||
vertex_project: str,
|
||||
vertex_credentials: str,
|
||||
) -> None:
|
||||
model_name = f"e2e-vertex-pt-{unique_marker()}"
|
||||
model_id = _add_vertex_passthrough_model(client, model_name, vertex_project, vertex_credentials)
|
||||
resources.defer(lambda: _delete_model(client, model_id))
|
||||
|
||||
result = client.vertex_generate(
|
||||
key=scoped_key,
|
||||
project=vertex_project,
|
||||
location=VERTEX_LOCATION,
|
||||
model=VERTEX_MODEL,
|
||||
text=f"reply with one word {unique_marker()}",
|
||||
)
|
||||
require_successful_call(result)
|
||||
assert '"candidates"' in result.body, f"vertex passthrough returned no candidates: {result.body[:300]}"
|
||||
|
||||
row = _costed_row(client, result.call_id)
|
||||
assert row.custom_llm_provider == "vertex_ai", f"passthrough spend logged under the wrong provider: {row}"
|
||||
assert "gemini" in (row.model or ""), f"unexpected model in spend log: {row}"
|
||||
|
|
@ -56,69 +56,56 @@ async def test_global_redaction_on():
|
|||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("turn_off_message_logging", [True, False])
|
||||
@pytest.mark.parametrize(
|
||||
"dynamic_turn_off, expect_redacted",
|
||||
[(True, True), (False, False)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_global_redaction_ignores_dynamic_param(turn_off_message_logging):
|
||||
"""
|
||||
Request-body `turn_off_message_logging` is no longer honored as a dynamic
|
||||
callback param — global setting (or admin-configured key/team config) wins.
|
||||
With global redaction ON, the caller cannot disable redaction via the
|
||||
request body.
|
||||
"""
|
||||
async def test_dynamic_turn_off_message_logging_overrides_global_on(dynamic_turn_off, expect_redacted):
|
||||
litellm.turn_off_message_logging = True
|
||||
test_custom_logger = TestCustomLogger()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
response = await litellm.acompletion(
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
turn_off_message_logging=turn_off_message_logging,
|
||||
turn_off_message_logging=dynamic_turn_off,
|
||||
mock_response="hello",
|
||||
)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
standard_logging_payload = test_custom_logger.logged_standard_logging_payload
|
||||
assert standard_logging_payload is not None
|
||||
print(
|
||||
"logged standard logging payload",
|
||||
json.dumps(standard_logging_payload, indent=2),
|
||||
)
|
||||
|
||||
response = standard_logging_payload["response"]
|
||||
assert response["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
||||
assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm"
|
||||
expected_response_content = "redacted-by-litellm" if expect_redacted else "hello"
|
||||
expected_message_content = "redacted-by-litellm" if expect_redacted else "hi"
|
||||
assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content
|
||||
assert standard_logging_payload["messages"][0]["content"] == expected_message_content
|
||||
|
||||
|
||||
@pytest.mark.parametrize("turn_off_message_logging", [True, False])
|
||||
@pytest.mark.parametrize(
|
||||
"dynamic_turn_off, expect_redacted",
|
||||
[(True, True), (False, False)],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_global_redaction_off_ignores_dynamic_param(turn_off_message_logging):
|
||||
"""
|
||||
Request-body `turn_off_message_logging` is no longer honored as a dynamic
|
||||
callback param — global setting (or admin-configured key/team config) wins.
|
||||
With global redaction OFF, the caller cannot enable redaction via the
|
||||
request body.
|
||||
"""
|
||||
async def test_dynamic_turn_off_message_logging_overrides_global_off(dynamic_turn_off, expect_redacted):
|
||||
litellm.turn_off_message_logging = False
|
||||
test_custom_logger = TestCustomLogger()
|
||||
litellm.callbacks = [test_custom_logger]
|
||||
response = await litellm.acompletion(
|
||||
await litellm.acompletion(
|
||||
model="gpt-5-mini",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
turn_off_message_logging=turn_off_message_logging,
|
||||
turn_off_message_logging=dynamic_turn_off,
|
||||
mock_response="hello",
|
||||
)
|
||||
|
||||
await asyncio.sleep(1)
|
||||
standard_logging_payload = test_custom_logger.logged_standard_logging_payload
|
||||
assert standard_logging_payload is not None
|
||||
print(
|
||||
"logged standard logging payload",
|
||||
json.dumps(standard_logging_payload, indent=2),
|
||||
)
|
||||
assert (
|
||||
standard_logging_payload["response"]["choices"][0]["message"]["content"]
|
||||
== "hello"
|
||||
)
|
||||
assert standard_logging_payload["messages"][0]["content"] == "hi"
|
||||
|
||||
expected_response_content = "redacted-by-litellm" if expect_redacted else "hello"
|
||||
expected_message_content = "redacted-by-litellm" if expect_redacted else "hi"
|
||||
assert standard_logging_payload["response"]["choices"][0]["message"]["content"] == expected_response_content
|
||||
assert standard_logging_payload["messages"][0]["content"] == expected_message_content
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -61,8 +61,8 @@ class TestAzureDocumentIntelligencePagesParam:
|
|||
def cfg(self) -> AzureDocumentIntelligenceOCRConfig:
|
||||
return AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
def test_get_supported_ocr_params_includes_pages(self, cfg):
|
||||
assert cfg.get_supported_ocr_params("prebuilt-layout") == ["pages"]
|
||||
def test_get_supported_ocr_params_includes_pages_and_features(self, cfg):
|
||||
assert cfg.get_supported_ocr_params("prebuilt-layout") == ["pages", "features"]
|
||||
|
||||
def test_map_ocr_params_mistral_zero_based_int_list(self, cfg):
|
||||
mapped = cfg.map_ocr_params({"pages": [0, 1, 2]}, {}, "prebuilt-layout")
|
||||
|
|
|
|||
0
tests/test_litellm/a2a_protocol/__init__.py
Normal file
0
tests/test_litellm/a2a_protocol/__init__.py
Normal file
|
|
@ -103,9 +103,7 @@ async def _mock_stream_messages(a2a_client: Any, request: Any) -> AsyncIterator[
|
|||
kind="message",
|
||||
)
|
||||
for _ in range(2):
|
||||
yield SendStreamingMessageResponse(
|
||||
root=SendStreamingMessageSuccessResponse(id=request.id, result=msg)
|
||||
)
|
||||
yield SendStreamingMessageResponse(root=SendStreamingMessageSuccessResponse(id=request.id, result=msg))
|
||||
|
||||
|
||||
class CostLogger(CustomLogger):
|
||||
|
|
@ -119,9 +117,7 @@ class CostLogger(CustomLogger):
|
|||
slp = kwargs.get("standard_logging_object")
|
||||
if slp:
|
||||
self.response_cost = (
|
||||
slp.get("response_cost")
|
||||
if isinstance(slp, dict)
|
||||
else getattr(slp, "response_cost", None)
|
||||
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -160,6 +156,43 @@ async def test_asend_message_uses_cost_per_query():
|
|||
assert cost_logger.response_cost == 0.05
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_uses_cost_per_query_from_litellm_params_dict():
|
||||
"""
|
||||
Proxy passes agent pricing as the litellm_params dict param (not top-level
|
||||
kwargs). Regression for cost_per_query landing at $0 on the native path.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message
|
||||
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
cost_logger = CostLogger()
|
||||
litellm.callbacks = [cost_logger]
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
mock_request = _make_send_message_request("test-123")
|
||||
|
||||
with patch(
|
||||
"litellm.a2a_protocol.main._execute_a2a_send_with_retry",
|
||||
new=_mock_execute_a2a_send,
|
||||
):
|
||||
await asend_message(
|
||||
a2a_client=mock_client,
|
||||
request=mock_request,
|
||||
litellm_params={
|
||||
"cost_per_query": 0.5,
|
||||
"input_cost_per_token": 0.099999,
|
||||
"output_cost_per_token": 0.1,
|
||||
},
|
||||
)
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert cost_logger.response_cost == 0.5
|
||||
|
||||
|
||||
class TokenAndCostLogger(CustomLogger):
|
||||
"""Custom logger to capture both token counts and cost."""
|
||||
|
||||
|
|
@ -173,19 +206,13 @@ class TokenAndCostLogger(CustomLogger):
|
|||
slp = kwargs.get("standard_logging_object")
|
||||
if slp:
|
||||
self.response_cost = (
|
||||
slp.get("response_cost")
|
||||
if isinstance(slp, dict)
|
||||
else getattr(slp, "response_cost", None)
|
||||
slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
||||
)
|
||||
self.prompt_tokens = (
|
||||
slp.get("prompt_tokens")
|
||||
if isinstance(slp, dict)
|
||||
else getattr(slp, "prompt_tokens", None)
|
||||
slp.get("prompt_tokens") if isinstance(slp, dict) else getattr(slp, "prompt_tokens", None)
|
||||
)
|
||||
self.completion_tokens = (
|
||||
slp.get("completion_tokens")
|
||||
if isinstance(slp, dict)
|
||||
else getattr(slp, "completion_tokens", None)
|
||||
slp.get("completion_tokens") if isinstance(slp, dict) else getattr(slp, "completion_tokens", None)
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -207,9 +234,7 @@ async def test_asend_message_uses_input_output_cost_per_token():
|
|||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
mock_request = _make_send_message_request(
|
||||
"test-123", user_text="Hello, what can you do?"
|
||||
)
|
||||
mock_request = _make_send_message_request("test-123", user_text="Hello, what can you do?")
|
||||
|
||||
# Define specific cost per token values
|
||||
input_cost_per_token = 0.00001 # $0.01 per 1000 tokens
|
||||
|
|
@ -246,15 +271,11 @@ async def test_asend_message_uses_input_output_cost_per_token():
|
|||
assert response_cost is not None, "response_cost should be captured"
|
||||
|
||||
# Calculate expected cost
|
||||
expected_cost = (prompt_tokens * input_cost_per_token) + (
|
||||
completion_tokens * output_cost_per_token
|
||||
)
|
||||
expected_cost = (prompt_tokens * input_cost_per_token) + (completion_tokens * output_cost_per_token)
|
||||
print(f"expected_cost: {expected_cost}")
|
||||
|
||||
# Verify exact cost calculation
|
||||
assert (
|
||||
response_cost == expected_cost
|
||||
), f"response_cost {response_cost} should equal expected {expected_cost}"
|
||||
assert response_cost == expected_cost, f"response_cost {response_cost} should equal expected {expected_cost}"
|
||||
|
||||
|
||||
class AgentIdLogger(CustomLogger):
|
||||
|
|
@ -305,9 +326,9 @@ async def test_asend_message_passes_agent_id_to_callback():
|
|||
await asyncio.sleep(0.1)
|
||||
|
||||
# Verify agent_id was passed to callback
|
||||
assert (
|
||||
agent_id_logger.agent_id == test_agent_id
|
||||
), f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'"
|
||||
assert agent_id_logger.agent_id == test_agent_id, (
|
||||
f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'"
|
||||
)
|
||||
|
||||
|
||||
class MetadataLogger(CustomLogger):
|
||||
|
|
@ -418,9 +439,7 @@ async def test_asend_message_streaming_triggers_callbacks():
|
|||
assert len(chunks) == 2
|
||||
|
||||
# Verify callbacks WERE triggered after stream completed
|
||||
assert (
|
||||
callback_logger.kwargs is not None
|
||||
), "Streaming should trigger callbacks after completion"
|
||||
assert (
|
||||
callback_logger.agent_id == test_agent_id
|
||||
), f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
|
||||
assert callback_logger.kwargs is not None, "Streaming should trigger callbacks after completion"
|
||||
assert callback_logger.agent_id == test_agent_id, (
|
||||
f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
|
||||
)
|
||||
|
|
|
|||
41
tests/test_litellm/a2a_protocol/test_utils.py
Normal file
41
tests/test_litellm/a2a_protocol/test_utils.py
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
"""Tests for litellm/a2a_protocol/utils.py token/usage extraction."""
|
||||
|
||||
import pytest
|
||||
|
||||
pytest.importorskip("a2a.compat.v0_3.types")
|
||||
|
||||
from a2a.compat.v0_3.types import MessageSendParams, SendMessageRequest
|
||||
|
||||
from litellm.a2a_protocol.utils import A2ARequestUtils
|
||||
|
||||
|
||||
def _request(user_text: str) -> SendMessageRequest:
|
||||
return SendMessageRequest(
|
||||
id="r1",
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"messageId": "m1",
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": user_text}],
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_calculate_usage_counts_input_tokens_from_request_object():
|
||||
"""Regression: request-side Part is a RootModel; input tokens must be counted."""
|
||||
request = _request("count these input tokens please")
|
||||
response_dict = {
|
||||
"result": {
|
||||
"kind": "message",
|
||||
"parts": [{"kind": "text", "text": "ok"}],
|
||||
}
|
||||
}
|
||||
|
||||
prompt_tokens, completion_tokens, total_tokens = A2ARequestUtils.calculate_usage_from_request_response(
|
||||
request=request, response_dict=response_dict
|
||||
)
|
||||
|
||||
assert prompt_tokens > 0
|
||||
assert completion_tokens > 0
|
||||
assert total_tokens == prompt_tokens + completion_tokens
|
||||
|
|
@ -0,0 +1,168 @@
|
|||
"""
|
||||
Unit tests for the per-request budget-metric emission timeout in
|
||||
PrometheusLogger._increment_remaining_budget_metrics.
|
||||
|
||||
A slow Redis/DB lookup in one of the budget branches must not let the gather run
|
||||
unbounded; it is wrapped in asyncio.wait_for so the success-logging coroutine
|
||||
cannot exceed the LoggingWorker watchdog and get the whole event cancelled.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from prometheus_client import REGISTRY
|
||||
|
||||
from litellm.integrations.prometheus import (
|
||||
PrometheusLogger,
|
||||
_DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT,
|
||||
_get_budget_metrics_per_request_timeout,
|
||||
)
|
||||
|
||||
TIMEOUT_ENV = "PROMETHEUS_BUDGET_METRICS_PER_REQUEST_TIMEOUT"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def cleanup_prometheus_registry():
|
||||
collectors = list(REGISTRY._collector_to_names.keys())
|
||||
for collector in collectors:
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
yield
|
||||
collectors = list(REGISTRY._collector_to_names.keys())
|
||||
for collector in collectors:
|
||||
try:
|
||||
REGISTRY.unregister(collector)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def prometheus_logger():
|
||||
return PrometheusLogger()
|
||||
|
||||
|
||||
def _call_increment(logger: PrometheusLogger):
|
||||
return logger._increment_remaining_budget_metrics(
|
||||
user_api_team="team-1",
|
||||
user_api_team_alias="team-alias",
|
||||
user_api_key="key-1",
|
||||
user_api_key_alias="key-alias",
|
||||
litellm_params={"metadata": {}},
|
||||
response_cost=0.01,
|
||||
user_id="user-1",
|
||||
user_api_key_org_id="org-1",
|
||||
)
|
||||
|
||||
|
||||
def _skip_logged(debug_mock) -> bool:
|
||||
return any("skipping" in str(call.args[0]) for call in debug_mock.call_args_list if call.args)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_budget_metric_emission_skips_on_timeout(prometheus_logger, monkeypatch):
|
||||
"""A branch slower than the timeout is skipped without propagating, and the
|
||||
skip is logged instead of cancelling the success-logging event."""
|
||||
monkeypatch.setenv(TIMEOUT_ENV, "0.05")
|
||||
|
||||
async def _slow_branch(**kwargs):
|
||||
await asyncio.sleep(30)
|
||||
|
||||
prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch
|
||||
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
|
||||
|
||||
with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger:
|
||||
await _call_increment(prometheus_logger)
|
||||
|
||||
assert _skip_logged(mock_logger.debug)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_budget_metric_emission_completes_within_timeout(prometheus_logger, monkeypatch):
|
||||
"""With a generous timeout every branch is awaited and no skip is logged."""
|
||||
monkeypatch.setenv(TIMEOUT_ENV, "5.0")
|
||||
|
||||
prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
|
||||
|
||||
with patch("litellm.integrations.prometheus.verbose_logger") as mock_logger:
|
||||
await _call_increment(prometheus_logger)
|
||||
|
||||
assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1
|
||||
assert prometheus_logger._set_team_budget_metrics_after_api_request.await_count == 1
|
||||
assert prometheus_logger._set_user_budget_metrics_after_api_request.await_count == 1
|
||||
assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1
|
||||
assert not _skip_logged(mock_logger.debug)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_timeout_env_falls_back_to_default(prometheus_logger, monkeypatch):
|
||||
"""A malformed timeout env value must not raise (which would recreate the
|
||||
failure mode); it falls back to the default and every branch still runs."""
|
||||
monkeypatch.setenv(TIMEOUT_ENV, "not-a-number")
|
||||
|
||||
prometheus_logger._set_api_key_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
|
||||
|
||||
await _call_increment(prometheus_logger)
|
||||
|
||||
assert prometheus_logger._set_api_key_budget_metrics_after_api_request.await_count == 1
|
||||
assert prometheus_logger._set_org_budget_metrics_after_api_request.await_count == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", ["not-a-number", "0", "-1", "nan", "inf", "-inf"])
|
||||
def test_unusable_timeout_env_falls_back_to_default(value, monkeypatch):
|
||||
"""Values that parse but disable or unbound the timeout (0, negative, nan,
|
||||
inf) must fall back to the default instead of being used; otherwise they
|
||||
either skip every emission or recreate the unbounded-wait failure mode."""
|
||||
monkeypatch.setenv(TIMEOUT_ENV, value)
|
||||
|
||||
assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value,expected", [("0.05", 0.05), ("5.0", 5.0), ("30", 30.0)])
|
||||
def test_valid_timeout_env_is_used(value, expected, monkeypatch):
|
||||
"""A finite positive value is parsed and returned unchanged."""
|
||||
monkeypatch.setenv(TIMEOUT_ENV, value)
|
||||
|
||||
assert _get_budget_metrics_per_request_timeout() == expected
|
||||
|
||||
|
||||
def test_missing_timeout_env_uses_default(monkeypatch):
|
||||
"""With the env unset the default is returned."""
|
||||
monkeypatch.delenv(TIMEOUT_ENV, raising=False)
|
||||
|
||||
assert _get_budget_metrics_per_request_timeout() == _DEFAULT_BUDGET_METRICS_PER_REQUEST_TIMEOUT
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_outer_cancellation_still_propagates(prometheus_logger, monkeypatch):
|
||||
"""Only asyncio.TimeoutError is swallowed; an outer cancellation (cooperative
|
||||
shutdown / watchdog) injected while awaiting must still propagate."""
|
||||
monkeypatch.setenv(TIMEOUT_ENV, "30")
|
||||
|
||||
started = asyncio.Event()
|
||||
|
||||
async def _slow_branch(**kwargs):
|
||||
started.set()
|
||||
await asyncio.sleep(30)
|
||||
|
||||
prometheus_logger._set_api_key_budget_metrics_after_api_request = _slow_branch
|
||||
prometheus_logger._set_team_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_user_budget_metrics_after_api_request = AsyncMock()
|
||||
prometheus_logger._set_org_budget_metrics_after_api_request = AsyncMock()
|
||||
|
||||
task = asyncio.create_task(_call_increment(prometheus_logger))
|
||||
await started.wait()
|
||||
task.cancel()
|
||||
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
|
@ -1194,3 +1194,156 @@ def test_s3_callback_params_override_empty_dict_is_opt_in():
|
|||
assert logger.s3_bucket_name is None
|
||||
finally:
|
||||
litellm.s3_callback_params = original
|
||||
|
||||
|
||||
def _expected_content_md5(payload: dict) -> str:
|
||||
import base64
|
||||
import hashlib
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
json_string = safe_dumps(payload)
|
||||
return base64.b64encode(
|
||||
hashlib.md5(json_string.encode("utf-8"), usedforsecurity=False).digest()
|
||||
).decode()
|
||||
|
||||
|
||||
def _require_non_security_md5(monkeypatch):
|
||||
import hashlib
|
||||
|
||||
original_md5 = hashlib.md5
|
||||
|
||||
def fips_md5(data=b"", *, usedforsecurity=True):
|
||||
if usedforsecurity:
|
||||
raise ValueError("MD5 blocked for security use")
|
||||
return original_md5(data, usedforsecurity=usedforsecurity)
|
||||
|
||||
monkeypatch.setattr(hashlib, "md5", fips_md5)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_upload_sets_content_md5_header(monkeypatch):
|
||||
"""
|
||||
Object Lock buckets reject PUTs without a Content-MD5 header (AWS spec).
|
||||
The async upload must send a base64 md5 of the exact signed body.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||||
|
||||
logger = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1",
|
||||
)
|
||||
|
||||
payload = {"test": "content-md5"}
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/test-md5.json",
|
||||
payload=payload,
|
||||
s3_object_download_filename="test-md5.json",
|
||||
)
|
||||
_require_non_security_md5(monkeypatch)
|
||||
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.raise_for_status = MagicMock()
|
||||
logger.async_httpx_client = AsyncMock()
|
||||
logger.async_httpx_client.put.return_value = response
|
||||
|
||||
await logger.async_upload_data_to_s3(test_element)
|
||||
|
||||
headers = logger.async_httpx_client.put.call_args.kwargs["headers"]
|
||||
assert headers["Content-MD5"] == _expected_content_md5(payload)
|
||||
assert "x-amz-server-side-encryption" not in headers
|
||||
|
||||
|
||||
def test_sync_upload_sets_content_md5_header(monkeypatch):
|
||||
"""The sync upload path must also send Content-MD5 for Object Lock buckets."""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||||
|
||||
logger = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1",
|
||||
)
|
||||
|
||||
payload = {"test": "sync-content-md5"}
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/test-sync-md5.json",
|
||||
payload=payload,
|
||||
s3_object_download_filename="test-sync-md5.json",
|
||||
)
|
||||
_require_non_security_md5(monkeypatch)
|
||||
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.raise_for_status = MagicMock()
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.put.return_value = response
|
||||
|
||||
with patch(
|
||||
"litellm.integrations.s3_v2._get_httpx_client",
|
||||
return_value=mock_sync_client,
|
||||
):
|
||||
logger.upload_data_to_s3(test_element)
|
||||
|
||||
headers = mock_sync_client.put.call_args.kwargs["headers"]
|
||||
assert headers["Content-MD5"] == _expected_content_md5(payload)
|
||||
assert "x-amz-server-side-encryption" not in headers
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_upload_sets_server_side_encryption_header_when_configured():
|
||||
"""
|
||||
When s3_server_side_encryption is set (e.g. buckets with a KMS default
|
||||
encryption policy), the PUT must carry x-amz-server-side-encryption.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||||
|
||||
logger = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1",
|
||||
s3_server_side_encryption="aws:kms",
|
||||
)
|
||||
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/test-sse.json",
|
||||
payload={"test": "sse"},
|
||||
s3_object_download_filename="test-sse.json",
|
||||
)
|
||||
|
||||
response = MagicMock()
|
||||
response.status_code = 200
|
||||
response.raise_for_status = MagicMock()
|
||||
logger.async_httpx_client = AsyncMock()
|
||||
logger.async_httpx_client.put.return_value = response
|
||||
|
||||
await logger.async_upload_data_to_s3(test_element)
|
||||
|
||||
headers = logger.async_httpx_client.put.call_args.kwargs["headers"]
|
||||
assert headers["x-amz-server-side-encryption"] == "aws:kms"
|
||||
|
||||
|
||||
def test_s3_server_side_encryption_read_from_callback_params():
|
||||
"""s3_server_side_encryption can be configured via s3_callback_params."""
|
||||
import litellm
|
||||
|
||||
original = litellm.s3_callback_params
|
||||
litellm.s3_callback_params = {
|
||||
"s3_bucket_name": "from-global",
|
||||
"s3_server_side_encryption": "aws:kms",
|
||||
}
|
||||
try:
|
||||
logger = S3Logger()
|
||||
assert logger.s3_server_side_encryption == "aws:kms"
|
||||
finally:
|
||||
litellm.s3_callback_params = original
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -2744,102 +2745,132 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5():
|
|||
TTL ordering constraint (tools -> system -> messages).
|
||||
|
||||
Ref: https://github.com/BerriAI/litellm/issues/XXXXX
|
||||
|
||||
Forces the bundled local cost map so ttl eligibility (driven by
|
||||
`cache_creation_input_token_cost_above_1hr` in litellm.model_cost) reads
|
||||
this branch's pricing data rather than the network-fetched `main` copy,
|
||||
which lacks the fix until merge.
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import (
|
||||
add_cache_point_tool_block,
|
||||
)
|
||||
|
||||
tool_with_1h = {
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {"type": "object"}},
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h"},
|
||||
}
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
tool_with_1h = {
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {"type": "object"}},
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h"},
|
||||
}
|
||||
|
||||
# Claude 4.5 model: ttl should be preserved
|
||||
result = add_cache_point_tool_block(
|
||||
tool_with_1h, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
|
||||
)
|
||||
assert result is not None
|
||||
assert result["cachePoint"]["type"] == "default"
|
||||
assert result["cachePoint"]["ttl"] == "1h"
|
||||
# Claude 4.5 model: ttl should be preserved
|
||||
result = add_cache_point_tool_block(
|
||||
tool_with_1h, model="jp.anthropic.claude-opus-4-7"
|
||||
)
|
||||
assert result is not None
|
||||
assert result["cachePoint"]["type"] == "default"
|
||||
assert result["cachePoint"]["ttl"] == "1h"
|
||||
|
||||
# Claude 4.5 model with 5m ttl: also preserved
|
||||
tool_with_5m = {
|
||||
"cache_control": {"type": "ephemeral", "ttl": "5m"},
|
||||
}
|
||||
result_5m = add_cache_point_tool_block(
|
||||
tool_with_5m, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
|
||||
)
|
||||
assert result_5m is not None
|
||||
assert result_5m["cachePoint"]["ttl"] == "5m"
|
||||
# Claude 4.5 model with 5m ttl: also preserved
|
||||
tool_with_5m = {
|
||||
"cache_control": {"type": "ephemeral", "ttl": "5m"},
|
||||
}
|
||||
result_5m = add_cache_point_tool_block(
|
||||
tool_with_5m, model="jp.anthropic.claude-opus-4-7"
|
||||
)
|
||||
assert result_5m is not None
|
||||
assert result_5m["cachePoint"]["ttl"] == "5m"
|
||||
|
||||
# Older model: ttl should be stripped
|
||||
result_old = add_cache_point_tool_block(
|
||||
tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||
)
|
||||
assert result_old is not None
|
||||
assert result_old["cachePoint"]["type"] == "default"
|
||||
assert "ttl" not in result_old["cachePoint"]
|
||||
# Older model: ttl should be stripped
|
||||
result_old = add_cache_point_tool_block(
|
||||
tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||
)
|
||||
assert result_old is not None
|
||||
assert result_old["cachePoint"]["type"] == "default"
|
||||
assert "ttl" not in result_old["cachePoint"]
|
||||
|
||||
# No model provided: ttl should be stripped (safe default)
|
||||
result_no_model = add_cache_point_tool_block(tool_with_1h, model=None)
|
||||
assert result_no_model is not None
|
||||
assert "ttl" not in result_no_model["cachePoint"]
|
||||
# No model provided: ttl should be stripped (safe default)
|
||||
result_no_model = add_cache_point_tool_block(tool_with_1h, model=None)
|
||||
assert result_no_model is not None
|
||||
assert "ttl" not in result_no_model["cachePoint"]
|
||||
|
||||
# No cache_control: returns None (unchanged behavior)
|
||||
tool_no_cache = {
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {"type": "object"}},
|
||||
}
|
||||
assert add_cache_point_tool_block(tool_no_cache) is None
|
||||
# No cache_control: returns None (unchanged behavior)
|
||||
tool_no_cache = {
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "parameters": {"type": "object"}},
|
||||
}
|
||||
assert add_cache_point_tool_block(tool_no_cache) is None
|
||||
|
||||
# cache_control without ttl: returns default cachePoint (unchanged behavior)
|
||||
tool_no_ttl = {"cache_control": {"type": "ephemeral"}}
|
||||
result_no_ttl = add_cache_point_tool_block(
|
||||
tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
|
||||
)
|
||||
assert result_no_ttl is not None
|
||||
assert result_no_ttl["cachePoint"]["type"] == "default"
|
||||
assert "ttl" not in result_no_ttl["cachePoint"]
|
||||
# cache_control without ttl: returns default cachePoint (unchanged behavior)
|
||||
tool_no_ttl = {"cache_control": {"type": "ephemeral"}}
|
||||
result_no_ttl = add_cache_point_tool_block(
|
||||
tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0"
|
||||
)
|
||||
assert result_no_ttl is not None
|
||||
assert result_no_ttl["cachePoint"]["type"] == "default"
|
||||
assert "ttl" not in result_no_ttl["cachePoint"]
|
||||
finally:
|
||||
litellm.model_cost = old_cost
|
||||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
|
||||
|
||||
def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
|
||||
"""
|
||||
End-to-end: _bedrock_tools_pt should produce cachePoint blocks with ttl
|
||||
for Claude 4.5+ models when tools have cache_control with ttl.
|
||||
|
||||
Forces the bundled local cost map so ttl eligibility (driven by
|
||||
`cache_creation_input_token_cost_above_1hr` in litellm.model_cost) reads
|
||||
this branch's pricing data rather than the network-fetched `main` copy,
|
||||
which lacks the fix until merge.
|
||||
"""
|
||||
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
|
||||
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
tools = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
},
|
||||
},
|
||||
},
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h"},
|
||||
}
|
||||
]
|
||||
"cache_control": {"type": "ephemeral", "ttl": "1h"},
|
||||
}
|
||||
]
|
||||
|
||||
# Claude 4.5: cachePoint should have ttl
|
||||
result = _bedrock_tools_pt(
|
||||
tools, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
|
||||
)
|
||||
cache_blocks = [b for b in result if "cachePoint" in b]
|
||||
assert len(cache_blocks) == 1
|
||||
assert cache_blocks[0]["cachePoint"]["ttl"] == "1h"
|
||||
# Claude 4.5: cachePoint should have ttl
|
||||
result = _bedrock_tools_pt(tools, model="jp.anthropic.claude-opus-4-7")
|
||||
cache_blocks = [b for b in result if "cachePoint" in b]
|
||||
assert len(cache_blocks) == 1
|
||||
assert cache_blocks[0]["cachePoint"]["ttl"] == "1h"
|
||||
|
||||
# Older model: cachePoint should not have ttl
|
||||
result_old = _bedrock_tools_pt(
|
||||
tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||
)
|
||||
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
|
||||
assert len(cache_blocks_old) == 1
|
||||
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
|
||||
# Older model: cachePoint should not have ttl
|
||||
result_old = _bedrock_tools_pt(
|
||||
tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
|
||||
)
|
||||
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
|
||||
assert len(cache_blocks_old) == 1
|
||||
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
|
||||
finally:
|
||||
litellm.model_cost = old_cost
|
||||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
|
||||
|
||||
def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document():
|
||||
|
|
|
|||
|
|
@ -7,9 +7,53 @@ sys.path.insert(0, os.path.abspath("../../.."))
|
|||
|
||||
from litellm.litellm_core_utils.initialize_dynamic_callback_params import (
|
||||
initialize_standard_callback_dynamic_params,
|
||||
iter_client_callback_metadata_dicts,
|
||||
)
|
||||
|
||||
|
||||
def test_iter_client_callback_metadata_dicts_covers_all_read_paths():
|
||||
md = {"m": 1}
|
||||
lm = {"lm": 1}
|
||||
lp_md = {"lp": 1}
|
||||
slots = dict(
|
||||
iter_client_callback_metadata_dicts(
|
||||
{
|
||||
"metadata": md,
|
||||
"litellm_metadata": lm,
|
||||
"litellm_params": {"metadata": lp_md},
|
||||
}
|
||||
)
|
||||
)
|
||||
assert slots == {
|
||||
"metadata": md,
|
||||
"litellm_metadata": lm,
|
||||
"litellm_params.metadata": lp_md,
|
||||
}
|
||||
|
||||
|
||||
def test_iter_client_callback_metadata_dicts_skips_non_dict_slots():
|
||||
slots = list(
|
||||
iter_client_callback_metadata_dicts(
|
||||
{
|
||||
"metadata": "not-a-dict",
|
||||
"litellm_metadata": None,
|
||||
"litellm_params": {"metadata": []},
|
||||
}
|
||||
)
|
||||
)
|
||||
assert slots == []
|
||||
|
||||
|
||||
def test_extractor_reads_turn_off_message_logging_from_every_slot():
|
||||
for kwargs in (
|
||||
{"metadata": {"turn_off_message_logging": True}},
|
||||
{"litellm_metadata": {"turn_off_message_logging": True}},
|
||||
{"litellm_params": {"metadata": {"turn_off_message_logging": True}}},
|
||||
):
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
assert params.get("turn_off_message_logging") is True, kwargs
|
||||
|
||||
|
||||
def test_resolves_plain_values_at_top_level():
|
||||
kwargs = {
|
||||
"langfuse_public_key": "pk-test",
|
||||
|
|
@ -36,6 +80,33 @@ def test_resolves_plain_values_from_metadata():
|
|||
assert params.get("langfuse_host") == "https://test.langfuse.com"
|
||||
|
||||
|
||||
def test_litellm_params_metadata_overrides_metadata():
|
||||
kwargs = {
|
||||
"metadata": {
|
||||
"langfuse_public_key": "pk-meta",
|
||||
},
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"langfuse_public_key": "pk-litellm-params",
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
|
||||
assert params.get("langfuse_public_key") == "pk-litellm-params"
|
||||
|
||||
|
||||
def test_top_level_kwargs_overrides_metadata_slots():
|
||||
kwargs = {
|
||||
"langfuse_public_key": "from-top-level",
|
||||
"metadata": {"langfuse_public_key": "from-metadata"},
|
||||
"litellm_params": {"metadata": {"langfuse_public_key": "from-litellm-params"}},
|
||||
}
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
assert params.get("langfuse_public_key") == "from-top-level"
|
||||
|
||||
|
||||
def test_env_reference_at_top_level_raises_with_guidance():
|
||||
kwargs = {"langfuse_public_key": "os.environ/LANGFUSE_PUBLIC_KEY"}
|
||||
|
||||
|
|
@ -100,11 +171,17 @@ def test_non_string_values_are_not_flagged():
|
|||
assert params.get("langsmith_sampling_rate") == 0.5
|
||||
|
||||
|
||||
def test_turn_off_message_logging_not_extracted_from_request():
|
||||
"""turn_off_message_logging is admin-only — must not be settable via request."""
|
||||
kwargs = {"turn_off_message_logging": True}
|
||||
@pytest.mark.parametrize(
|
||||
"kwargs,expected",
|
||||
[
|
||||
({"turn_off_message_logging": False}, False),
|
||||
({"turn_off_message_logging": "False"}, "False"),
|
||||
({"metadata": {"turn_off_message_logging": True}}, True),
|
||||
],
|
||||
)
|
||||
def test_turn_off_message_logging_extracted_from_kwargs(kwargs, expected):
|
||||
params = initialize_standard_callback_dynamic_params(kwargs)
|
||||
assert params.get("turn_off_message_logging") is None
|
||||
assert params.get("turn_off_message_logging") == expected
|
||||
|
||||
|
||||
def test_empty_kwargs_returns_empty_params():
|
||||
|
|
|
|||
|
|
@ -0,0 +1,42 @@
|
|||
"""Tests for litellm/llms/a2a/chat/transformation.py response transform."""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.a2a.chat.transformation import A2AConfig
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
def _raw_response(text: str) -> MagicMock:
|
||||
raw = MagicMock()
|
||||
raw.status_code = 200
|
||||
raw.headers = {}
|
||||
raw.json.return_value = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "resp-1",
|
||||
"result": {
|
||||
"kind": "message",
|
||||
"parts": [{"kind": "text", "text": text}],
|
||||
},
|
||||
}
|
||||
return raw
|
||||
|
||||
|
||||
def test_transform_response_sets_usage():
|
||||
"""Regression: A2AConfig.transform_response must populate usage so per-token
|
||||
pricing computes real cost and callers don't get usage 0/0/0."""
|
||||
result = A2AConfig().transform_response(
|
||||
model="a2a/test-agent",
|
||||
raw_response=_raw_response("hello from the agent"),
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data={},
|
||||
messages=[{"role": "user", "content": "hi there agent"}],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
assert result.usage is not None
|
||||
assert result.usage.prompt_tokens > 0
|
||||
assert result.usage.completion_tokens > 0
|
||||
assert result.usage.total_tokens == (result.usage.prompt_tokens + result.usage.completion_tokens)
|
||||
|
|
@ -1,3 +1,6 @@
|
|||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
|
|
@ -31,3 +34,217 @@ def test_should_reject_dot_segment_azure_document_intelligence_model_id():
|
|||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
|
||||
AZURE_TABLES = [
|
||||
{
|
||||
"rowCount": 2,
|
||||
"columnCount": 2,
|
||||
"cells": [
|
||||
{"kind": "columnHeader", "rowIndex": 0, "columnIndex": 0, "content": "Item"},
|
||||
{"kind": "columnHeader", "rowIndex": 0, "columnIndex": 1, "content": "Price"},
|
||||
{"rowIndex": 1, "columnIndex": 0, "content": "Widget"},
|
||||
{"rowIndex": 1, "columnIndex": 1, "content": "$100.00"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"rowCount": 1,
|
||||
"columnCount": 1,
|
||||
"cells": [{"rowIndex": 0, "columnIndex": 0, "content": "Totals"}],
|
||||
},
|
||||
]
|
||||
|
||||
AZURE_KEY_VALUE_PAIRS = [
|
||||
{"key": {"content": "Invoice No"}, "value": {"content": "INV-12345"}, "confidence": 0.98},
|
||||
{"key": {"content": "Total"}, "value": {"content": "$100.00"}, "confidence": 0.95},
|
||||
]
|
||||
|
||||
AZURE_ANALYZE_SUCCEEDED = {
|
||||
"status": "succeeded",
|
||||
"createdDateTime": "2026-07-02T00:00:00Z",
|
||||
"lastUpdatedDateTime": "2026-07-02T00:00:05Z",
|
||||
"analyzeResult": {
|
||||
"apiVersion": "2024-11-30",
|
||||
"modelId": "prebuilt-layout",
|
||||
"content": "Invoice\nInvoice No: INV-12345\nTotal: $100.00",
|
||||
"pages": [
|
||||
{
|
||||
"pageNumber": 1,
|
||||
"width": 8.5,
|
||||
"height": 11,
|
||||
"unit": "inch",
|
||||
"lines": [
|
||||
{"content": "Invoice"},
|
||||
{"content": "Invoice No: INV-12345"},
|
||||
{"content": "Total: $100.00"},
|
||||
],
|
||||
}
|
||||
],
|
||||
"tables": AZURE_TABLES,
|
||||
"keyValuePairs": AZURE_KEY_VALUE_PAIRS,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _completed_response(payload: dict) -> httpx.Response:
|
||||
return httpx.Response(
|
||||
status_code=200,
|
||||
json=payload,
|
||||
request=httpx.Request("GET", "https://example.cognitiveservices.azure.com/analyzeResults/xyz"),
|
||||
)
|
||||
|
||||
|
||||
def _assert_native_fields_preserved(serialized: dict) -> None:
|
||||
assert serialized["content"] == "Invoice\nInvoice No: INV-12345\nTotal: $100.00"
|
||||
assert serialized["tables"] == AZURE_TABLES
|
||||
assert serialized["keyValuePairs"] == AZURE_KEY_VALUE_PAIRS
|
||||
assert serialized["object"] == "ocr"
|
||||
assert serialized["usage_info"]["pages_processed"] == 1
|
||||
assert serialized["pages"][0]["index"] == 0
|
||||
assert serialized["pages"][0]["markdown"] == "Invoice\nInvoice No: INV-12345\nTotal: $100.00"
|
||||
assert serialized["pages"][0]["dimensions"] == {"width": 816, "height": 1056, "dpi": 96}
|
||||
|
||||
|
||||
def test_transform_ocr_response_preserves_azure_native_fields():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
result = config.transform_ocr_response(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
raw_response=_completed_response(AZURE_ANALYZE_SUCCEEDED),
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
_assert_native_fields_preserved(result.model_dump())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_transform_ocr_response_preserves_azure_native_fields():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
result = await config.async_transform_ocr_response(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
raw_response=_completed_response(AZURE_ANALYZE_SUCCEEDED),
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
_assert_native_fields_preserved(result.model_dump())
|
||||
|
||||
|
||||
def test_transform_ocr_response_tolerates_missing_native_fields():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
payload = {
|
||||
"status": "succeeded",
|
||||
"analyzeResult": {
|
||||
"pages": [
|
||||
{
|
||||
"pageNumber": 1,
|
||||
"width": 8.5,
|
||||
"height": 11,
|
||||
"unit": "inch",
|
||||
"lines": [{"content": "hello"}],
|
||||
}
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
result = config.transform_ocr_response(
|
||||
model="azure_ai/doc-intelligence/prebuilt-read",
|
||||
raw_response=_completed_response(payload),
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
serialized = result.model_dump()
|
||||
assert serialized["pages"][0]["markdown"] == "hello"
|
||||
assert serialized["content"] is None
|
||||
assert serialized["tables"] is None
|
||||
assert serialized["keyValuePairs"] is None
|
||||
|
||||
|
||||
def test_transform_ocr_response_non_succeeded_status_raises():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
with pytest.raises(ValueError, match="failed with status: failed"):
|
||||
config.transform_ocr_response(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
raw_response=_completed_response({"status": "failed"}),
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
def test_get_supported_ocr_params_includes_features():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
assert config.get_supported_ocr_params("prebuilt-layout") == ["pages", "features"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"features,expected",
|
||||
[
|
||||
(["keyValuePairs"], "keyValuePairs"),
|
||||
(["keyValuePairs", "languages"], "keyValuePairs,languages"),
|
||||
("keyValuePairs", "keyValuePairs"),
|
||||
("keyValuePairs,languages", "keyValuePairs,languages"),
|
||||
("keyValuePairs, languages", "keyValuePairs,languages"),
|
||||
],
|
||||
)
|
||||
def test_map_ocr_params_features(features, expected):
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
mapped = config.map_ocr_params({"features": features}, {}, "prebuilt-layout")
|
||||
|
||||
assert mapped == {"features": expected}
|
||||
|
||||
|
||||
def test_map_ocr_params_empty_features_list_omitted():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
assert config.map_ocr_params({"features": []}, {}, "prebuilt-layout") == {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"features",
|
||||
[
|
||||
"keyValuePairs&pages=9",
|
||||
"key value pairs",
|
||||
"",
|
||||
[1, 2],
|
||||
[["keyValuePairs"]],
|
||||
{"feature": "keyValuePairs"},
|
||||
5,
|
||||
],
|
||||
)
|
||||
def test_map_ocr_params_invalid_features_raises(features):
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
with pytest.raises(ValueError, match="Invalid `features`"):
|
||||
config.map_ocr_params({"features": features}, {}, "prebuilt-layout")
|
||||
|
||||
|
||||
def test_get_complete_url_appends_features_query():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://example.cognitiveservices.azure.com",
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
optional_params={"features": "keyValuePairs"},
|
||||
)
|
||||
|
||||
assert "&features=keyValuePairs" in url
|
||||
|
||||
|
||||
def test_get_complete_url_combines_pages_and_features():
|
||||
config = AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
optional_params = config.map_ocr_params(
|
||||
{"pages": [0, 1, 2], "features": ["keyValuePairs", "languages"]},
|
||||
{},
|
||||
"prebuilt-layout",
|
||||
)
|
||||
url = config.get_complete_url(
|
||||
api_base="https://example.cognitiveservices.azure.com",
|
||||
model="prebuilt-layout",
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
assert "&pages=1,2,3" in url
|
||||
assert "&features=keyValuePairs,languages" in url
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transfor
|
|||
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
|
||||
AmazonInvokeConfig,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -39,3 +40,145 @@ def test_transform_request_drops_stream_chunk_size(config, model):
|
|||
)
|
||||
|
||||
assert "stream_chunk_size" not in json.dumps(request_body)
|
||||
|
||||
|
||||
def test_validate_environment_maps_guardrail_config_to_invoke_headers():
|
||||
"""The InvokeModel API takes the guardrail identifier/version/trace as
|
||||
X-Amzn-Bedrock-* request headers, unlike Converse which takes them in the
|
||||
body. https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_InvokeModel.html"""
|
||||
optional_params = {
|
||||
"guardrailConfig": {
|
||||
"guardrailIdentifier": "ff6ujrregl1q",
|
||||
"guardrailVersion": "DRAFT",
|
||||
"trace": "enabled",
|
||||
},
|
||||
"max_tokens": 10,
|
||||
}
|
||||
|
||||
headers = AmazonInvokeConfig().validate_environment(
|
||||
headers={},
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert headers["X-Amzn-Bedrock-GuardrailIdentifier"] == "ff6ujrregl1q"
|
||||
assert headers["X-Amzn-Bedrock-GuardrailVersion"] == "DRAFT"
|
||||
assert headers["X-Amzn-Bedrock-Trace"] == "ENABLED"
|
||||
assert "guardrailConfig" not in optional_params
|
||||
|
||||
|
||||
def test_validate_environment_without_guardrail_config_leaves_headers_untouched():
|
||||
headers = AmazonInvokeConfig().validate_environment(
|
||||
headers={"foo": "bar"},
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"max_tokens": 10},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert headers == {"foo": "bar"}
|
||||
|
||||
|
||||
def test_validate_environment_skips_absent_guardrail_fields():
|
||||
headers = AmazonInvokeConfig().validate_environment(
|
||||
headers={},
|
||||
model="amazon.titan-text-express-v1",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"guardrailConfig": {"guardrailIdentifier": "gr-id", "guardrailVersion": "1"}},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert headers == {
|
||||
"X-Amzn-Bedrock-GuardrailIdentifier": "gr-id",
|
||||
"X-Amzn-Bedrock-GuardrailVersion": "1",
|
||||
}
|
||||
|
||||
|
||||
def test_validate_environment_does_not_clobber_explicit_guardrail_headers():
|
||||
"""Users worked around the missing guardrailConfig support by passing the
|
||||
AWS headers directly; an explicit header must keep winning over
|
||||
guardrailConfig regardless of casing."""
|
||||
headers = AmazonInvokeConfig().validate_environment(
|
||||
headers={"x-amzn-bedrock-guardrailidentifier": "explicit-id"},
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={
|
||||
"guardrailConfig": {"guardrailIdentifier": "config-id", "guardrailVersion": "2"},
|
||||
},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert headers["x-amzn-bedrock-guardrailidentifier"] == "explicit-id"
|
||||
assert "X-Amzn-Bedrock-GuardrailIdentifier" not in headers
|
||||
assert headers["X-Amzn-Bedrock-GuardrailVersion"] == "2"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_guardrail_config",
|
||||
[
|
||||
{"guardrailIdentifier": "gr-id", "trace": "verbose"},
|
||||
{"guardrailIdentifier": ["gr-id"]},
|
||||
"gr-id",
|
||||
{},
|
||||
{"trace": "enabled"},
|
||||
],
|
||||
)
|
||||
def test_validate_environment_rejects_malformed_guardrail_config(bad_guardrail_config):
|
||||
with pytest.raises(BedrockError) as excinfo:
|
||||
AmazonInvokeConfig().validate_environment(
|
||||
headers={},
|
||||
model="anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
optional_params={"guardrailConfig": bad_guardrail_config},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert excinfo.value.status_code == 400
|
||||
assert "guardrailConfig" in str(excinfo.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
"amazon.titan-text-express-v1",
|
||||
"mistral.mistral-7b-instruct-v0:2",
|
||||
"meta.llama3-8b-instruct-v1:0",
|
||||
],
|
||||
)
|
||||
def test_guardrail_config_flows_to_headers_not_request_body(model):
|
||||
"""Mirrors the handler flow (validate_environment then transform_request):
|
||||
guardrailConfig must end up in the signed headers and never leak into the
|
||||
request body, where Bedrock rejects it as an extra input."""
|
||||
config = AmazonInvokeConfig()
|
||||
optional_params = {
|
||||
"guardrailConfig": {
|
||||
"guardrailIdentifier": "ff6ujrregl1q",
|
||||
"guardrailVersion": "DRAFT",
|
||||
"trace": "disabled",
|
||||
},
|
||||
"max_tokens": 10,
|
||||
}
|
||||
messages = [{"role": "user", "content": "hi"}]
|
||||
|
||||
headers = config.validate_environment(
|
||||
headers={},
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
request_body = config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
assert "guardrailConfig" not in json.dumps(request_body)
|
||||
assert headers["X-Amzn-Bedrock-GuardrailIdentifier"] == "ff6ujrregl1q"
|
||||
assert headers["X-Amzn-Bedrock-GuardrailVersion"] == "DRAFT"
|
||||
assert headers["X-Amzn-Bedrock-Trace"] == "DISABLED"
|
||||
|
|
|
|||
|
|
@ -611,6 +611,65 @@ def test_transform_request_helper_includes_anthropic_beta_and_tools():
|
|||
assert fields["tools"][0]["type"] == "computer_20250124"
|
||||
|
||||
|
||||
def test_parallel_tool_calls_config_kept_for_sonnet_5():
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
model="anthropic.claude-sonnet-5",
|
||||
non_default_params={"parallel_tool_calls": False},
|
||||
optional_params={},
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
data = config._transform_request_helper(
|
||||
model="anthropic.claude-sonnet-5",
|
||||
system_content_blocks=[],
|
||||
optional_params=optional_params,
|
||||
messages=None,
|
||||
)
|
||||
|
||||
assert data["additionalModelRequestFields"]["tool_choice"] == {
|
||||
"disable_parallel_tool_use": True
|
||||
}
|
||||
finally:
|
||||
litellm.model_cost = old_cost
|
||||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
|
||||
|
||||
def test_parallel_tool_calls_config_dropped_for_ttl_only_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
model = "anthropic.claude-fable-5"
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
model,
|
||||
{"cache_creation_input_token_cost_above_1hr": 2e-05},
|
||||
)
|
||||
config = AmazonConverseConfig()
|
||||
optional_params = config.map_openai_params(
|
||||
model=model,
|
||||
non_default_params={"parallel_tool_calls": False},
|
||||
optional_params={},
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
data = config._transform_request_helper(
|
||||
model=model,
|
||||
system_content_blocks=[],
|
||||
optional_params=optional_params,
|
||||
messages=None,
|
||||
)
|
||||
|
||||
assert "tool_choice" not in data.get("additionalModelRequestFields", {})
|
||||
|
||||
|
||||
def test_transform_response_with_computer_use_tool():
|
||||
"""Test response transformation with computer use tool call."""
|
||||
import httpx
|
||||
|
|
@ -4130,6 +4189,42 @@ def test_parallel_tool_calls_newer_model_adds_disable_flag():
|
|||
assert "parallel_tool_calls" not in request_data["additionalModelRequestFields"]
|
||||
|
||||
|
||||
def test_parallel_tool_calls_flag_decoupled_from_ttl_pricing(monkeypatch):
|
||||
"""
|
||||
The disable_parallel_tool_use gate must read supports_parallel_tool_use_config,
|
||||
not the 1h-TTL pricing field: a model carrying only the former still gets the flag.
|
||||
"""
|
||||
from litellm.llms.bedrock.common_utils import is_claude_4_5_on_bedrock
|
||||
|
||||
config = AmazonConverseConfig()
|
||||
model = "anthropic.claude-parallel-tool-use-only"
|
||||
monkeypatch.setitem(litellm.model_cost, model, {"supports_parallel_tool_use_config": True})
|
||||
assert is_claude_4_5_on_bedrock(model) is False
|
||||
messages = [{"role": "user", "content": "What's the weather in SF and NYC?"}]
|
||||
|
||||
optional_params = config.map_openai_params(
|
||||
non_default_params={"parallel_tool_calls": False, "tools": _TOOL_PARAM},
|
||||
optional_params={},
|
||||
model=model,
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
request_data = config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert (
|
||||
request_data["additionalModelRequestFields"]["tool_choice"][
|
||||
"disable_parallel_tool_use"
|
||||
]
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_parallel_tool_calls_older_model_drops_disable_flag():
|
||||
"""Older Claude models (pre-4.5) must NOT receive disable_parallel_tool_use — Bedrock rejects it."""
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -4554,6 +4649,154 @@ def test_cache_control_injection_tool_config_not_added_without_injection_point()
|
|||
assert all("cachePoint" not in tool for tool in tools)
|
||||
|
||||
|
||||
def test_cache_control_injection_tool_config_honors_ttl_for_supported_model():
|
||||
"""
|
||||
Regression test: cache_control_injection_points with location=tool_config
|
||||
must honor the requested `control.ttl`, mirroring the message/system
|
||||
cache_control behavior, instead of always emitting a bare
|
||||
{"type": "default"} cachePoint with no ttl.
|
||||
|
||||
Forces the bundled local cost map so `is_claude_4_5_on_bedrock` (which
|
||||
reads `cache_creation_input_token_cost_above_1hr` from litellm.model_cost)
|
||||
sees this branch's pricing data rather than the network-fetched `main`
|
||||
copy, which lacks it until merge.
|
||||
"""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
messages = [
|
||||
{"role": "user", "content": "What is the weather?"},
|
||||
]
|
||||
optional_params = {
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
"cache_control_injection_points": [
|
||||
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
],
|
||||
}
|
||||
result = config._transform_request(
|
||||
model="us.anthropic.claude-sonnet-4-5-20250929-v1:0",
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
tools = result["toolConfig"]["tools"]
|
||||
assert tools[-1] == {"cachePoint": {"type": "default", "ttl": "1h"}}
|
||||
finally:
|
||||
litellm.model_cost = old_cost
|
||||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
|
||||
|
||||
def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacking_own_pricing():
|
||||
"""
|
||||
Regression test: a regional pricing entry that omits
|
||||
`cache_creation_input_token_cost_above_1hr` (e.g. `jp.anthropic.claude-opus-4-7`)
|
||||
must not shadow the base model entry that carries it; the requested ttl
|
||||
survives through the base-model fallback.
|
||||
"""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
assert "cache_creation_input_token_cost_above_1hr" not in litellm.model_cost["jp.anthropic.claude-opus-4-7"]
|
||||
assert "cache_creation_input_token_cost_above_1hr" in litellm.model_cost["anthropic.claude-opus-4-7"]
|
||||
config = AmazonConverseConfig()
|
||||
messages = [
|
||||
{"role": "user", "content": "What is the weather?"},
|
||||
]
|
||||
optional_params = {
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
"cache_control_injection_points": [
|
||||
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
],
|
||||
}
|
||||
result = config._transform_request(
|
||||
model="jp.anthropic.claude-opus-4-7",
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
tools = result["toolConfig"]["tools"]
|
||||
assert tools[-1] == {"cachePoint": {"type": "default", "ttl": "1h"}}
|
||||
finally:
|
||||
litellm.model_cost = old_cost
|
||||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
|
||||
|
||||
def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model():
|
||||
"""
|
||||
Models that don't support extended TTL caching (only Claude 4.5+ on
|
||||
Bedrock does) must fall back to the default cachePoint with no ttl,
|
||||
even if the caller requested one, matching message/system behavior.
|
||||
"""
|
||||
config = AmazonConverseConfig()
|
||||
messages = [
|
||||
{"role": "user", "content": "What is the weather?"},
|
||||
]
|
||||
optional_params = {
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get weather for a location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
"cache_control_injection_points": [
|
||||
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
|
||||
],
|
||||
}
|
||||
result = config._transform_request(
|
||||
model="anthropic.claude-3-5-haiku-20241022-v1:0",
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
)
|
||||
tools = result["toolConfig"]["tools"]
|
||||
assert tools[-1] == {"cachePoint": {"type": "default"}}
|
||||
|
||||
|
||||
def test_translate_response_format_json_schema_still_injects_tool():
|
||||
"""
|
||||
response_format with an explicit json_schema should still use the
|
||||
|
|
|
|||
|
|
@ -492,7 +492,7 @@ def test_bedrock_invoke_messages_transform_converts_custom_tool_schema_type_to_o
|
|||
assert result["tools"][0]["type"] == "custom"
|
||||
|
||||
|
||||
def test_remove_ttl_from_cache_control_processes_tools():
|
||||
def test_remove_ttl_from_cache_control_processes_tools(local_model_cost_map):
|
||||
"""
|
||||
Ensure _remove_ttl_from_cache_control also sanitizes cache_control on tools.
|
||||
|
||||
|
|
@ -538,7 +538,7 @@ def test_remove_ttl_from_cache_control_processes_tools():
|
|||
assert "ttl" not in request["system"][0]["cache_control"]
|
||||
|
||||
|
||||
def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5():
|
||||
def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5(local_model_cost_map):
|
||||
"""
|
||||
For Claude 4.5+ models, ttl in ["5m", "1h"] should be preserved on tools,
|
||||
just like it is for system and messages.
|
||||
|
|
@ -564,7 +564,7 @@ def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5():
|
|||
}
|
||||
|
||||
cfg._remove_ttl_from_cache_control(
|
||||
request, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
|
||||
request, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0"
|
||||
)
|
||||
|
||||
# Both tools and system should preserve ttl for Claude 4.5
|
||||
|
|
|
|||
|
|
@ -445,3 +445,31 @@ def test_explicit_invoke_route_does_not_match_async_invoke():
|
|||
BedrockModelInfo._explicit_async_invoke_route(f"bedrock/{async_invoke_model}")
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
def test_capability_lookups_fall_back_to_base_model_when_regional_entry_lacks_field(monkeypatch):
|
||||
"""
|
||||
Regression test: a regional model_cost entry without the capability field
|
||||
must not shadow a base entry that has it (`get(model) or get(base)` used to
|
||||
short-circuit on the truthy regional dict and drop the capability).
|
||||
"""
|
||||
import litellm
|
||||
from litellm.llms.bedrock.common_utils import (
|
||||
bedrock_converse_supports_parallel_tool_use_config,
|
||||
is_claude_4_5_on_bedrock,
|
||||
)
|
||||
|
||||
base = "anthropic.claude-fallback-test"
|
||||
regional = f"eu.{base}"
|
||||
monkeypatch.setitem(litellm.model_cost, regional, {"input_cost_per_token": 1e-06})
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
base,
|
||||
{
|
||||
"cache_creation_input_token_cost_above_1hr": 1e-05,
|
||||
"supports_parallel_tool_use_config": True,
|
||||
},
|
||||
)
|
||||
|
||||
assert is_claude_4_5_on_bedrock(regional) is True
|
||||
assert bedrock_converse_supports_parallel_tool_use_config(regional) is True
|
||||
|
|
|
|||
0
tests/test_litellm/llms/tencent/__init__.py
Normal file
0
tests/test_litellm/llms/tencent/__init__.py
Normal file
0
tests/test_litellm/llms/tencent/chat/__init__.py
Normal file
0
tests/test_litellm/llms/tencent/chat/__init__.py
Normal file
|
|
@ -0,0 +1,207 @@
|
|||
from unittest.mock import patch
|
||||
|
||||
from litellm.llms.tencent.chat.transformation import TencentChatConfig
|
||||
|
||||
|
||||
def test_supported_openai_params_includes_thinking_and_reasoning_effort():
|
||||
config = TencentChatConfig()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.supports_reasoning",
|
||||
return_value=True,
|
||||
):
|
||||
params = config.get_supported_openai_params(model="tencent/deepseek-v4-pro")
|
||||
|
||||
assert "thinking" in params
|
||||
assert "reasoning_effort" in params
|
||||
assert "stream" in params
|
||||
assert "temperature" in params
|
||||
|
||||
|
||||
def test_supported_openai_params_excludes_thinking_without_reasoning_support():
|
||||
config = TencentChatConfig()
|
||||
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.supports_reasoning",
|
||||
return_value=False,
|
||||
):
|
||||
params = config.get_supported_openai_params(model="tencent/non-reasoning-model")
|
||||
|
||||
assert "thinking" not in params
|
||||
assert "reasoning_effort" not in params
|
||||
assert "stream" in params
|
||||
|
||||
|
||||
def test_map_openai_params_passes_thinking_dict_through():
|
||||
config = TencentChatConfig()
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.supports_reasoning",
|
||||
return_value=True,
|
||||
):
|
||||
result = config.map_openai_params(
|
||||
non_default_params={"thinking": {"type": "enabled", "budget_tokens": 1024}},
|
||||
optional_params={},
|
||||
model="tencent/deepseek-v4-pro",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["thinking"] == {"type": "enabled", "budget_tokens": 1024}
|
||||
|
||||
|
||||
def test_map_openai_params_converts_reasoning_effort_to_thinking():
|
||||
config = TencentChatConfig()
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.supports_reasoning",
|
||||
return_value=True,
|
||||
):
|
||||
result = config.map_openai_params(
|
||||
non_default_params={"reasoning_effort": "medium"},
|
||||
optional_params={},
|
||||
model="tencent/deepseek-v4-pro",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["thinking"] == {"type": "enabled"}
|
||||
|
||||
|
||||
def test_map_openai_params_drops_none_reasoning_effort():
|
||||
config = TencentChatConfig()
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.supports_reasoning",
|
||||
return_value=True,
|
||||
):
|
||||
result = config.map_openai_params(
|
||||
non_default_params={"reasoning_effort": "none"},
|
||||
optional_params={},
|
||||
model="tencent/deepseek-v4-pro",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "thinking" not in result
|
||||
assert "reasoning_effort" not in result
|
||||
|
||||
|
||||
def test_map_openai_params_thinking_priority_over_reasoning_effort():
|
||||
config = TencentChatConfig()
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.supports_reasoning",
|
||||
return_value=True,
|
||||
):
|
||||
result = config.map_openai_params(
|
||||
non_default_params={
|
||||
"thinking": {"type": "enabled", "budget_tokens": 2048},
|
||||
"reasoning_effort": "high",
|
||||
},
|
||||
optional_params={},
|
||||
model="tencent/deepseek-v4-pro",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert result["thinking"] == {"type": "enabled", "budget_tokens": 2048}
|
||||
|
||||
|
||||
def test_map_openai_params_extracts_thinking_and_effort_from_optional_params():
|
||||
config = TencentChatConfig()
|
||||
result = config.map_openai_params(
|
||||
non_default_params={},
|
||||
optional_params={"thinking": {"type": "enabled"}, "reasoning_effort": "medium"},
|
||||
model="tencent/deepseek-v4-pro",
|
||||
drop_params=False,
|
||||
)
|
||||
|
||||
assert "thinking" in result
|
||||
assert "reasoning_effort" not in result
|
||||
|
||||
|
||||
def test_get_complete_url_default():
|
||||
config = TencentChatConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="tencent/deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert url == "https://tokenhub-intl.tencentcloudmaas.com/v1/chat/completions"
|
||||
|
||||
|
||||
def test_get_complete_url_strips_trailing_slash():
|
||||
config = TencentChatConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1/",
|
||||
api_key=None,
|
||||
model="tencent/deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert url == "https://tokenhub-intl.tencentcloudmaas.com/v1/chat/completions"
|
||||
|
||||
|
||||
def test_get_complete_url_custom_base_preserves_v1():
|
||||
config = TencentChatConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://tokenhub.tencentcloudmaas.com/v1",
|
||||
api_key=None,
|
||||
model="tencent/deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert url == "https://tokenhub.tencentcloudmaas.com/v1/chat/completions"
|
||||
|
||||
|
||||
def test_get_complete_url_adds_v1_to_custom_base():
|
||||
config = TencentChatConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://tokenhub.tencentcloudmaas.com",
|
||||
api_key=None,
|
||||
model="tencent/deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert url == "https://tokenhub.tencentcloudmaas.com/v1/chat/completions"
|
||||
|
||||
|
||||
def test_get_complete_url_does_not_append_to_full_url():
|
||||
config = TencentChatConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://tokenhub.tencentcloudmaas.com/v1/chat/completions",
|
||||
api_key=None,
|
||||
model="tencent/deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert url == "https://tokenhub.tencentcloudmaas.com/v1/chat/completions"
|
||||
|
||||
|
||||
def test_provider_info_falls_back_to_default_base():
|
||||
config = TencentChatConfig()
|
||||
|
||||
with patch("litellm.llms.tencent.chat.transformation.get_secret_str", return_value=None):
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(api_base=None, api_key="sk-arg")
|
||||
|
||||
assert api_base == "https://tokenhub-intl.tencentcloudmaas.com/v1"
|
||||
assert api_key == "sk-arg"
|
||||
|
||||
|
||||
def test_provider_info_reads_env_secrets():
|
||||
config = TencentChatConfig()
|
||||
|
||||
secrets = {"TENCENT_API_BASE": "https://env.tencent/v1", "TENCENT_API_KEY": "sk-env"}
|
||||
with patch(
|
||||
"litellm.llms.tencent.chat.transformation.get_secret_str",
|
||||
side_effect=lambda key: secrets.get(key),
|
||||
):
|
||||
api_base, api_key = config._get_openai_compatible_provider_info(api_base=None, api_key=None)
|
||||
|
||||
assert api_base == "https://env.tencent/v1"
|
||||
assert api_key == "sk-env"
|
||||
0
tests/test_litellm/llms/tencent/messages/__init__.py
Normal file
0
tests/test_litellm/llms/tencent/messages/__init__.py
Normal file
|
|
@ -0,0 +1,173 @@
|
|||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
|
||||
AnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.llms.tencent.messages.transformation import (
|
||||
TencentAnthropicMessagesConfig,
|
||||
)
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
|
||||
def test_tencent_provider_uses_anthropic_messages_config():
|
||||
config = ProviderConfigManager.get_provider_anthropic_messages_config(
|
||||
model="deepseek-v4-pro",
|
||||
provider=litellm.LlmProviders.TENCENT,
|
||||
)
|
||||
|
||||
assert isinstance(config, TencentAnthropicMessagesConfig)
|
||||
assert config.custom_llm_provider == "tencent"
|
||||
|
||||
|
||||
def test_anthropic_provider_keeps_default_config_for_tencent_named_model():
|
||||
config = ProviderConfigManager.get_provider_anthropic_messages_config(
|
||||
model="deepseek-v4-pro",
|
||||
provider=litellm.LlmProviders.ANTHROPIC,
|
||||
)
|
||||
|
||||
assert isinstance(config, AnthropicMessagesConfig)
|
||||
assert not isinstance(config, TencentAnthropicMessagesConfig)
|
||||
|
||||
|
||||
def test_strips_billing_metadata():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert config.should_strip_billing_metadata() is True
|
||||
|
||||
|
||||
def test_get_api_base_default():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert config.get_api_base() == "https://tokenhub-intl.tencentcloudmaas.com"
|
||||
|
||||
|
||||
def test_get_api_base_from_arg():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert config.get_api_base(api_base="https://custom.example.com") == "https://custom.example.com"
|
||||
|
||||
|
||||
def test_messages_url_default():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key=None,
|
||||
model="deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
== "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
|
||||
)
|
||||
|
||||
|
||||
def test_messages_url_with_base_ending_in_v1():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1",
|
||||
api_key=None,
|
||||
model="deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
== "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
|
||||
)
|
||||
|
||||
|
||||
def test_messages_url_with_base_ending_in_v1_messages():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
url = config.get_complete_url(
|
||||
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1/messages",
|
||||
api_key=None,
|
||||
model="deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert url == "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
|
||||
|
||||
|
||||
def test_messages_url_with_base_ending_in_v1_chat_completions():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base="https://tokenhub-intl.tencentcloudmaas.com/v1/chat/completions",
|
||||
api_key=None,
|
||||
model="deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
== "https://tokenhub-intl.tencentcloudmaas.com/v1/messages"
|
||||
)
|
||||
|
||||
|
||||
def test_messages_url_with_custom_base_no_v1():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
assert (
|
||||
config.get_complete_url(
|
||||
api_base="https://tokenhub.tencentcloudmaas.com",
|
||||
api_key=None,
|
||||
model="deepseek-v4-pro",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
== "https://tokenhub.tencentcloudmaas.com/v1/messages"
|
||||
)
|
||||
|
||||
|
||||
def test_validate_environment_sets_headers():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
headers, api_base = config.validate_anthropic_messages_environment(
|
||||
headers={},
|
||||
model="deepseek-v4-pro",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key="sk-tencent-key",
|
||||
api_base="https://custom.test",
|
||||
)
|
||||
|
||||
assert headers["x-api-key"] == "sk-tencent-key"
|
||||
assert headers["anthropic-version"] == "2023-06-01"
|
||||
assert headers["content-type"] == "application/json"
|
||||
assert api_base == "https://custom.test"
|
||||
|
||||
|
||||
def test_validate_environment_injects_anthropic_beta_headers():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
headers, _ = config.validate_anthropic_messages_environment(
|
||||
headers={},
|
||||
model="deepseek-v4-pro",
|
||||
messages=[],
|
||||
optional_params={"speed": "fast"},
|
||||
litellm_params={},
|
||||
api_key="sk-tencent-key",
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert "anthropic-beta" in headers
|
||||
|
||||
|
||||
def test_validate_environment_preserves_existing_headers():
|
||||
config = TencentAnthropicMessagesConfig()
|
||||
|
||||
headers, _ = config.validate_anthropic_messages_environment(
|
||||
headers={"authorization": "Bearer existing", "anthropic-version": "2024-01-01"},
|
||||
model="deepseek-v4-pro",
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
api_key="sk-tencent-key",
|
||||
api_base=None,
|
||||
)
|
||||
|
||||
assert headers["authorization"] == "Bearer existing"
|
||||
assert headers["anthropic-version"] == "2024-01-01"
|
||||
assert "x-api-key" not in headers
|
||||
41
tests/test_litellm/llms/tencent/test_cost_calculator.py
Normal file
41
tests/test_litellm/llms/tencent/test_cost_calculator.py
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.tencent.cost_calculator import cost_per_token
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def local_model_cost_map(monkeypatch):
|
||||
original_model_cost = litellm.model_cost
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.get_model_info.cache_clear()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
litellm.model_cost = original_model_cost
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
def test_cost_per_token_uses_tencent_model_pricing(local_model_cost_map):
|
||||
usage = Usage(prompt_tokens=1000, completion_tokens=2000, total_tokens=3000)
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(model="tencent/deepseek-v4-pro", usage=usage)
|
||||
|
||||
assert prompt_cost == pytest.approx(1000 * 4.35e-07)
|
||||
assert completion_cost == pytest.approx(2000 * 8.7e-07)
|
||||
|
||||
|
||||
def test_top_level_dispatcher_routes_tencent_to_wrapper(local_model_cost_map):
|
||||
from litellm.cost_calculator import cost_per_token as dispatch_cost_per_token
|
||||
|
||||
prompt_cost, completion_cost = dispatch_cost_per_token(
|
||||
model="tencent/deepseek-v4-pro",
|
||||
prompt_tokens=1000,
|
||||
completion_tokens=1000,
|
||||
custom_llm_provider="tencent",
|
||||
)
|
||||
|
||||
assert prompt_cost == pytest.approx(1000 * 4.35e-07)
|
||||
assert completion_cost == pytest.approx(1000 * 8.7e-07)
|
||||
|
|
@ -0,0 +1,184 @@
|
|||
import asyncio
|
||||
from typing import Dict, Optional
|
||||
|
||||
import pytest
|
||||
from unittest.mock import patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.types.mcp import MCPAuth, MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
HOLD_SECONDS = 0.1
|
||||
|
||||
|
||||
class _ConcurrencyTracker:
|
||||
"""Records how many call_tool invocations are simultaneously in flight."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.current_by_server: Dict[str, int] = {}
|
||||
self.peak_by_server: Dict[str, int] = {}
|
||||
self.global_current = 0
|
||||
self.global_peak = 0
|
||||
|
||||
def enter(self, server_id: str) -> None:
|
||||
self.current_by_server[server_id] = self.current_by_server.get(server_id, 0) + 1
|
||||
self.peak_by_server[server_id] = max(self.peak_by_server.get(server_id, 0), self.current_by_server[server_id])
|
||||
self.global_current += 1
|
||||
self.global_peak = max(self.global_peak, self.global_current)
|
||||
|
||||
def exit(self, server_id: str) -> None:
|
||||
self.current_by_server[server_id] -= 1
|
||||
self.global_current -= 1
|
||||
|
||||
|
||||
def _make_server(server_id: str, max_concurrent_requests: Optional[int]) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=server_id,
|
||||
name=server_id,
|
||||
server_name=server_id,
|
||||
url="https://example.com",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
max_concurrent_requests=max_concurrent_requests,
|
||||
)
|
||||
|
||||
|
||||
def _patch_client_with_tracker(manager: MCPServerManager, tracker: _ConcurrencyTracker):
|
||||
async def fake_create_mcp_client(server, **kwargs):
|
||||
class _ProbeClient:
|
||||
async def call_tool(self, params, host_progress_callback=None):
|
||||
tracker.enter(server.server_id)
|
||||
try:
|
||||
await asyncio.sleep(HOLD_SECONDS)
|
||||
return "ok"
|
||||
finally:
|
||||
tracker.exit(server.server_id)
|
||||
|
||||
return _ProbeClient()
|
||||
|
||||
return patch.object(manager, "_create_mcp_client", side_effect=fake_create_mcp_client)
|
||||
|
||||
|
||||
async def _fire(manager: MCPServerManager, server: MCPServer, n: int) -> None:
|
||||
await asyncio.gather(
|
||||
*[
|
||||
manager._call_regular_mcp_tool(
|
||||
mcp_server=server,
|
||||
original_tool_name="tool",
|
||||
arguments={},
|
||||
tasks=[],
|
||||
mcp_auth_header=None,
|
||||
mcp_server_auth_headers=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
proxy_logging_obj=None,
|
||||
)
|
||||
for _ in range(n)
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_max_concurrent_requests_caps_in_flight_tool_calls():
|
||||
"""A configured cap of 2 must never let more than 2 calls hit one server at once."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-limited", max_concurrent_requests=2)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await _fire(manager, server, n=8)
|
||||
|
||||
assert tracker.peak_by_server["srv-limited"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unset_limit_allows_unbounded_concurrency():
|
||||
"""With no cap, all calls run concurrently (backward-compatible default)."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-unbounded", max_concurrent_requests=None)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await _fire(manager, server, n=6)
|
||||
|
||||
assert tracker.peak_by_server["srv-unbounded"] == 6
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_positive_limit_is_treated_as_unlimited():
|
||||
"""A cap of 0 must not deadlock; it means unlimited, not a zero-permit semaphore."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-zero", max_concurrent_requests=0)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await asyncio.wait_for(_fire(manager, server, n=5), timeout=5)
|
||||
|
||||
assert tracker.peak_by_server["srv-zero"] == 5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_limit_is_scoped_per_server():
|
||||
"""Each server gets its own limiter; one server's cap must not throttle another."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server_a = _make_server("srv-a", max_concurrent_requests=1)
|
||||
server_b = _make_server("srv-b", max_concurrent_requests=1)
|
||||
|
||||
with _patch_client_with_tracker(manager, tracker):
|
||||
await asyncio.gather(
|
||||
_fire(manager, server_a, n=3),
|
||||
_fire(manager, server_b, n=3),
|
||||
)
|
||||
|
||||
assert tracker.peak_by_server["srv-a"] == 1
|
||||
assert tracker.peak_by_server["srv-b"] == 1
|
||||
assert tracker.global_peak == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_openapi_backed_server_also_respects_the_cap():
|
||||
"""OpenAPI (spec_path) servers dispatch through a different handler; the cap
|
||||
must apply there too, not only on the regular MCP client path."""
|
||||
manager = MCPServerManager()
|
||||
tracker = _ConcurrencyTracker()
|
||||
server = _make_server("srv-openapi", max_concurrent_requests=2)
|
||||
server.spec_path = "/fake/openapi.json"
|
||||
|
||||
async def fake_openapi_handler(mcp_server, name, arguments):
|
||||
tracker.enter(mcp_server.server_id)
|
||||
try:
|
||||
await asyncio.sleep(HOLD_SECONDS)
|
||||
return "ok"
|
||||
finally:
|
||||
tracker.exit(mcp_server.server_id)
|
||||
|
||||
with (
|
||||
patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server),
|
||||
patch.object(manager, "_call_openapi_tool_handler", side_effect=fake_openapi_handler),
|
||||
):
|
||||
await asyncio.gather(
|
||||
*[manager.call_tool(server_name="srv-openapi", name="tool", arguments={}) for _ in range(6)]
|
||||
)
|
||||
|
||||
assert tracker.peak_by_server["srv-openapi"] == 2
|
||||
|
||||
|
||||
def test_semaphore_is_reused_per_server_and_distinct_across_servers():
|
||||
manager = MCPServerManager()
|
||||
server_a = _make_server("srv-a", max_concurrent_requests=3)
|
||||
server_b = _make_server("srv-b", max_concurrent_requests=3)
|
||||
|
||||
sem_a_first = manager._get_call_semaphore(server_a)
|
||||
sem_a_second = manager._get_call_semaphore(server_a)
|
||||
sem_b = manager._get_call_semaphore(server_b)
|
||||
|
||||
assert sem_a_first is sem_a_second
|
||||
assert sem_a_first is not sem_b
|
||||
|
||||
|
||||
def test_no_semaphore_created_when_limit_absent():
|
||||
manager = MCPServerManager()
|
||||
server = _make_server("srv-none", max_concurrent_requests=None)
|
||||
|
||||
assert manager._get_call_semaphore(server) is None
|
||||
|
|
@ -76,9 +76,7 @@ async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
|
|||
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
|
||||
|
||||
with pytest.raises(MCPUpstreamAuthError) as exc_info:
|
||||
await manager._fetch_tools_with_timeout(
|
||||
mock_client, passthrough_server.name, server=passthrough_server
|
||||
)
|
||||
await manager._fetch_tools_with_timeout(mock_client, passthrough_server.name)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.www_authenticate == (
|
||||
|
|
@ -113,9 +111,7 @@ async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
|
|||
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
|
||||
|
||||
with pytest.raises(MCPUpstreamAuthError) as exc_info:
|
||||
await manager._fetch_tools_with_timeout(
|
||||
mock_client, delegated_server.name, server=delegated_server
|
||||
)
|
||||
await manager._fetch_tools_with_timeout(mock_client, delegated_server.name)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.www_authenticate == (
|
||||
|
|
@ -126,7 +122,10 @@ async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior():
|
||||
async def test_fetch_tools_from_client_credentials_oauth2_surfaces_upstream_401():
|
||||
"""The auth_type carve-out was removed: a client_credentials (M2M) server now
|
||||
surfaces an upstream 401 as MCPUpstreamAuthError too, instead of swallowing it
|
||||
to an empty list, so single-server routes can return a 401 challenge."""
|
||||
manager = MCPServerManager()
|
||||
m2m_server = MCPServer(
|
||||
server_id="oauth-m2m",
|
||||
|
|
@ -150,12 +149,12 @@ async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior
|
|||
mock_client = MagicMock()
|
||||
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
|
||||
|
||||
tools = await manager._fetch_tools_with_timeout(
|
||||
mock_client, m2m_server.name, server=m2m_server
|
||||
)
|
||||
with pytest.raises(MCPUpstreamAuthError) as exc_info:
|
||||
await manager._fetch_tools_with_timeout(mock_client, m2m_server.name)
|
||||
|
||||
assert tools == []
|
||||
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.server_name == "m2m_docs"
|
||||
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -176,9 +175,7 @@ async def test_fetch_tools_from_passthrough_returns_tools_on_success():
|
|||
mock_client = MagicMock()
|
||||
mock_client.list_tools = AsyncMock(return_value=[tool])
|
||||
|
||||
tools = await manager._fetch_tools_with_timeout(
|
||||
mock_client, passthrough_server.name, server=passthrough_server
|
||||
)
|
||||
tools = await manager._fetch_tools_with_timeout(mock_client, passthrough_server.name)
|
||||
assert tools == [tool]
|
||||
|
||||
|
||||
|
|
@ -238,8 +235,11 @@ def test_to_http_exception_skips_challenge_for_non_401_status():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_tools_from_gateway_managed_swallows_errors():
|
||||
"""Regression guard: non-pass-through servers keep returning [] on errors."""
|
||||
async def test_fetch_tools_from_gateway_managed_surfaces_upstream_401():
|
||||
"""An oauth2 server that is neither pass-through nor delegate now surfaces an
|
||||
upstream 401 as MCPUpstreamAuthError as well; the auth_type carve-out that
|
||||
swallowed it to [] was removed. A missing upstream WWW-Authenticate is carried
|
||||
through as None (the single-server route fabricates one from the gateway URL)."""
|
||||
manager = MCPServerManager()
|
||||
oauth2_server = MCPServer(
|
||||
server_id="o1",
|
||||
|
|
@ -260,11 +260,13 @@ async def test_fetch_tools_from_gateway_managed_swallows_errors():
|
|||
mock_client = MagicMock()
|
||||
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
|
||||
|
||||
tools = await manager._fetch_tools_with_timeout(
|
||||
mock_client, oauth2_server.name, server=oauth2_server
|
||||
)
|
||||
assert tools == []
|
||||
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
|
||||
with pytest.raises(MCPUpstreamAuthError) as exc_info:
|
||||
await manager._fetch_tools_with_timeout(mock_client, oauth2_server.name)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.www_authenticate is None
|
||||
assert exc_info.value.server_name == "keycloak_whoami"
|
||||
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
|
||||
|
||||
|
||||
def _http_server(server_id: str, name: str, **kwargs) -> MCPServer:
|
||||
|
|
|
|||
|
|
@ -5228,5 +5228,139 @@ class TestCreateMcpClientV2Graft:
|
|||
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
|
||||
|
||||
|
||||
def _upstream_status_error(status_code: int, challenge: str) -> httpx.HTTPStatusError:
|
||||
request = httpx.Request("POST", "https://upstream.example/mcp")
|
||||
response = httpx.Response(
|
||||
status_code,
|
||||
headers={"WWW-Authenticate": challenge},
|
||||
request=request,
|
||||
)
|
||||
return httpx.HTTPStatusError(
|
||||
"upstream rejected token", request=request, response=response
|
||||
)
|
||||
|
||||
|
||||
class TestMCPToolsListAuthSurfacing:
|
||||
"""Regression: MCP tools/list 401 auth failures must surface as MCPUpstreamAuthError.
|
||||
|
||||
Previously a missing/expired per-user OAuth token, or an upstream 401 for any
|
||||
non-carveout auth_type, was swallowed to an empty tool list, so a single-server
|
||||
client saw a 200 with no tools instead of a 401 challenge. The listing helpers
|
||||
now raise MCPUpstreamAuthError on a 401 regardless of auth_type; the single-server
|
||||
routes turn it into a 401 + WWW-Authenticate while the aggregator absorbs it to an
|
||||
empty list. Only a 401 challenges; a 403 (forbidden) degrades like any other error.
|
||||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_tools_with_timeout_surfaces_upstream_401(self):
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
challenge = 'Bearer resource_metadata="https://upstream.example/.well-known/oauth-protected-resource"'
|
||||
client = MagicMock()
|
||||
client.list_tools = AsyncMock(side_effect=_upstream_status_error(401, challenge))
|
||||
|
||||
with pytest.raises(MCPUpstreamAuthError) as exc_info:
|
||||
await manager._fetch_tools_with_timeout(client, "static-key-server")
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.www_authenticate == challenge
|
||||
assert exc_info.value.server_name == "static-key-server"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_tools_with_timeout_absorbs_upstream_403(self):
|
||||
"""Only a 401 drives the re-auth challenge. A 403 (authenticated but
|
||||
forbidden, e.g. insufficient scope) is not a re-auth signal, so even
|
||||
with a WWW-Authenticate header it degrades to an empty list rather than
|
||||
surfacing a challenge."""
|
||||
manager = MCPServerManager()
|
||||
challenge = 'Bearer error="insufficient_scope", scope="read:tools"'
|
||||
client = MagicMock()
|
||||
client.list_tools = AsyncMock(side_effect=_upstream_status_error(403, challenge))
|
||||
|
||||
assert await manager._fetch_tools_with_timeout(client, "forbidden-server") == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_tools_with_timeout_returns_empty_on_non_auth_error(self):
|
||||
manager = MCPServerManager()
|
||||
client = MagicMock()
|
||||
client.list_tools = AsyncMock(side_effect=RuntimeError("upstream 500"))
|
||||
|
||||
assert await manager._fetch_tools_with_timeout(client, "srv") == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_server_surfaces_unusable_user_token(self):
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="oauth-srv", name="oauth-srv", transport=MCPTransport.http
|
||||
)
|
||||
challenge = 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/oauth-srv"'
|
||||
manager._create_mcp_client = AsyncMock(
|
||||
side_effect=HTTPException(
|
||||
status_code=401,
|
||||
detail="Unauthorized",
|
||||
headers={"WWW-Authenticate": challenge},
|
||||
)
|
||||
)
|
||||
|
||||
with pytest.raises(MCPUpstreamAuthError) as exc_info:
|
||||
await manager._get_tools_from_server(server)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
assert exc_info.value.www_authenticate == challenge
|
||||
assert exc_info.value.server_name == "oauth-srv"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tools_from_server_absorbs_non_challenge_http_error(self):
|
||||
manager = MCPServerManager()
|
||||
server = MCPServer(
|
||||
server_id="stdio-srv", name="stdio-srv", transport=MCPTransport.http
|
||||
)
|
||||
manager._create_mcp_client = AsyncMock(
|
||||
side_effect=HTTPException(
|
||||
status_code=403,
|
||||
detail="MCP stdio command 'foo' is not in the allowlist",
|
||||
)
|
||||
)
|
||||
|
||||
assert await manager._get_tools_from_server(server) == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aggregate_list_tools_absorbs_unauthenticated_server(self):
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
|
||||
manager = MCPServerManager()
|
||||
good = MCPServer(server_id="good", name="good", transport=MCPTransport.http)
|
||||
bad = MCPServer(server_id="bad", name="bad", transport=MCPTransport.http)
|
||||
manager.get_allowed_mcp_servers = AsyncMock(return_value=["good", "bad"])
|
||||
manager.get_mcp_server_by_id = MagicMock(
|
||||
side_effect=lambda server_id: {"good": good, "bad": bad}.get(server_id)
|
||||
)
|
||||
good_tool = MCPTool(name="good-do_thing", description="do thing", inputSchema={})
|
||||
|
||||
async def fake_get_tools(server, **kwargs):
|
||||
if server.server_id == "bad":
|
||||
raise MCPUpstreamAuthError(
|
||||
status_code=401,
|
||||
www_authenticate='Bearer realm="x"',
|
||||
server_name="bad",
|
||||
)
|
||||
return [good_tool]
|
||||
|
||||
manager._get_tools_from_server = fake_get_tools
|
||||
|
||||
result = await manager.list_tools()
|
||||
|
||||
assert [t.name for t in result] == ["good-do_thing"]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
|
|
@ -843,6 +843,76 @@ class TestListToolsRestAPI:
|
|||
assert exc_info.value.status_code == upstream_status
|
||||
assert exc_info.value.headers == {"www-authenticate": challenge}
|
||||
|
||||
async def test_aggregate_list_absorbs_one_server_auth_failure(self, monkeypatch):
|
||||
"""The multi-server aggregate listing degrades a server whose upstream
|
||||
rejects auth to an empty contribution and still returns the healthy
|
||||
server's tools with a 200, rather than surfacing a 401."""
|
||||
from litellm.proxy._experimental.mcp_server.exceptions import (
|
||||
MCPUpstreamAuthError,
|
||||
)
|
||||
|
||||
class StubServer:
|
||||
def __init__(self, name):
|
||||
self.alias = name
|
||||
self.server_name = name
|
||||
self.name = name
|
||||
self.allowed_tools = None
|
||||
self.mcp_info = {"server_name": name}
|
||||
self.available_on_public_internet = True
|
||||
|
||||
good = StubServer("good")
|
||||
bad = StubServer("bad")
|
||||
|
||||
async def fake_contexts(user_api_key_auth):
|
||||
return [user_api_key_auth]
|
||||
|
||||
async def fake_get_allowed_mcp_servers(*args, **kwargs):
|
||||
return ["good", "bad"]
|
||||
|
||||
async def fake_get_tools(server, *args, **kwargs):
|
||||
if server.server_name == "bad":
|
||||
raise MCPUpstreamAuthError(
|
||||
status_code=401,
|
||||
www_authenticate='Bearer realm="x"',
|
||||
server_name="bad",
|
||||
)
|
||||
return ["good-tool"]
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"build_effective_auth_contexts",
|
||||
fake_contexts,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_allowed_mcp_servers",
|
||||
fake_get_allowed_mcp_servers,
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
"get_mcp_server_by_id",
|
||||
lambda server_id: {"good": good, "bad": bad}.get(server_id),
|
||||
raising=False,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints,
|
||||
"_get_tools_for_single_server",
|
||||
fake_get_tools,
|
||||
raising=False,
|
||||
)
|
||||
|
||||
request = _build_request(path="/mcp-rest/tools/list", method="GET")
|
||||
result = await rest_endpoints.list_tool_rest_api(
|
||||
request,
|
||||
server_id=None,
|
||||
user_api_key_dict=UserAPIKeyAuth(),
|
||||
)
|
||||
|
||||
assert result["tools"] == ["good-tool"]
|
||||
assert result["error"] is None
|
||||
|
||||
async def test_name_resolution_finds_server_by_uuid(self, monkeypatch):
|
||||
"""When server_id is a name string, it should be resolved to its UUID
|
||||
and used for the tools lookup when the UUID is in allowed_server_ids."""
|
||||
|
|
|
|||
|
|
@ -58,15 +58,71 @@ class TestAnthropicEndpoints(unittest.TestCase):
|
|||
self.assertEqual(result, expected_result)
|
||||
|
||||
# Assert safe_dumps was called for dictionary objects
|
||||
mock_safe_dumps.assert_any_call(
|
||||
{"type": "message_start", "message": {"id": "msg_123"}}
|
||||
mock_safe_dumps.assert_any_call({"type": "message_start", "message": {"id": "msg_123"}})
|
||||
mock_safe_dumps.assert_any_call({"type": "content_block_delta", "delta": {"text": "more data"}})
|
||||
assert mock_safe_dumps.call_count == 2 # Called twice, once for each dict object
|
||||
|
||||
|
||||
class TestBlockedResponseUsage:
|
||||
"""Blocked responses report the blocked LLM response's real usage."""
|
||||
|
||||
def test_uses_original_response_usage(self):
|
||||
from litellm.proxy.anthropic_endpoints.endpoints import _blocked_response_usage
|
||||
|
||||
# original_response is the AnthropicMessagesResponse the LLM produced
|
||||
# before the guardrail blocked it; its usage is real.
|
||||
original = {"usage": {"input_tokens": 31, "output_tokens": 9}}
|
||||
assert _blocked_response_usage(original) == {
|
||||
"input_tokens": 31,
|
||||
"output_tokens": 9,
|
||||
}
|
||||
|
||||
def test_zero_usage_when_no_original_response(self):
|
||||
from litellm.proxy.anthropic_endpoints.endpoints import _blocked_response_usage
|
||||
|
||||
# Pre-call blocks never invoked the LLM -> nothing consumed.
|
||||
assert _blocked_response_usage(None) == {
|
||||
"input_tokens": 0,
|
||||
"output_tokens": 0,
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_blocked_endpoint_response_carries_original_usage(self):
|
||||
"""The /v1/messages block handler reports the blocked response's real
|
||||
usage, carried on ModifyResponseException.original_response."""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import litellm.proxy.anthropic_endpoints.endpoints as ep
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
||||
exc = ModifyResponseException(
|
||||
message="blocked by guardrail",
|
||||
model="claude-3-5-sonnet-20240620",
|
||||
request_data={"messages": [{"role": "user", "content": "hi"}]},
|
||||
guardrail_name="rubrik",
|
||||
original_response={"usage": {"input_tokens": 12, "output_tokens": 5}},
|
||||
)
|
||||
mock_safe_dumps.assert_any_call(
|
||||
{"type": "content_block_delta", "delta": {"text": "more data"}}
|
||||
)
|
||||
assert (
|
||||
mock_safe_dumps.call_count == 2
|
||||
) # Called twice, once for each dict object
|
||||
|
||||
with (
|
||||
patch.object(ep, "_read_request_body", new=AsyncMock(return_value={})),
|
||||
patch.object(
|
||||
ep.ProxyBaseLLMRequestProcessing,
|
||||
"base_process_llm_request",
|
||||
new=AsyncMock(side_effect=exc),
|
||||
),
|
||||
patch.object(proxy_server, "proxy_logging_obj") as mock_logging,
|
||||
):
|
||||
mock_logging.post_call_failure_hook = AsyncMock()
|
||||
response = await ep.anthropic_response(
|
||||
fastapi_response=MagicMock(),
|
||||
request=MagicMock(),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert response["content"][0]["text"] == "blocked by guardrail"
|
||||
assert response["usage"] == {"input_tokens": 12, "output_tokens": 5}
|
||||
mock_logging.post_call_failure_hook.assert_awaited_once()
|
||||
|
||||
|
||||
class TestEventLoggingBatchEndpoint:
|
||||
|
|
@ -159,9 +215,7 @@ class TestStripTotalTokens(unittest.TestCase):
|
|||
|
||||
# SimpleNamespace mimics the .usage attribute access pattern; the
|
||||
# helper's contract: if .usage is dict-shaped, strip total_tokens.
|
||||
response = SimpleNamespace(
|
||||
usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150}
|
||||
)
|
||||
response = SimpleNamespace(usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150})
|
||||
_strip_total_tokens_from_anthropic_response(response)
|
||||
assert "total_tokens" not in response.usage
|
||||
assert response.usage == {"input_tokens": 100, "output_tokens": 50}
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue