Merge branch 'litellm_/guardrail-automation-testing-3ecd3d' into litellm_presidio_ui_user_story_e2e

This commit is contained in:
Yuneng Jiang 2026-09-07 10:08:30 -07:00
commit b61ef7c232
No known key found for this signature in database
48 changed files with 2452 additions and 462 deletions

View file

@ -1,129 +0,0 @@
name: Report LiteLLM Rust release wheel
on: # zizmor: ignore[dangerous-triggers] reporter executes no PR code and consumes no PR artifacts or outputs
workflow_run:
workflows:
- LiteLLM Rust
types:
- completed
permissions: {}
concurrency:
group: ${{ github.workflow }}-${{ github.event.workflow_run.pull_requests[0].number || github.event.workflow_run.id }}
cancel-in-progress: false
jobs:
report-release-wheel:
name: report release wheel
if: >-
github.event.workflow_run.event == 'pull_request' &&
github.event.workflow_run.path == '.github/workflows/test-rust.yml' &&
github.event.workflow_run.head_repository.full_name == github.repository &&
github.event.workflow_run.pull_requests[0].number != null
runs-on: ubuntu-latest
timeout-minutes: 5
permissions:
pull-requests: write
steps:
- name: Link release wheel report on PR
uses: actions/github-script@f28e40c7f34bde8b3046d885e986cb6290c5673b # v7.1.0
env:
COMMENT_MARKER: "<!-- litellm-release-wheel-size -->"
with:
script: |
const marker = process.env.COMMENT_MARKER;
const workflowRun = context.payload.workflow_run;
const allowedConclusions = new Set([
"action_required",
"cancelled",
"failure",
"neutral",
"skipped",
"stale",
"startup_failure",
"success",
"timed_out",
]);
if (
!allowedConclusions.has(workflowRun.conclusion) ||
workflowRun.event !== "pull_request" ||
workflowRun.path !== ".github/workflows/test-rust.yml" ||
workflowRun.head_repository?.full_name !==
`${context.repo.owner}/${context.repo.repo}` ||
workflowRun.pull_requests?.length !== 1
) {
throw new Error("unexpected source workflow");
}
const pullRequest = workflowRun.pull_requests[0];
const pullRequestNumber = pullRequest.number;
const headSha = workflowRun.head_sha;
const runId = workflowRun.id;
if (
!Number.isSafeInteger(pullRequestNumber) ||
pullRequestNumber <= 0 ||
!Number.isSafeInteger(runId) ||
runId <= 0 ||
!/^[0-9a-f]{40}$/.test(headSha) ||
pullRequest.head?.sha !== headSha
) {
throw new Error("invalid source workflow metadata");
}
const runUrl =
`${context.serverUrl}/${context.repo.owner}/${context.repo.repo}` +
`/actions/runs/${runId}`;
const result =
workflowRun.conclusion === "success"
? "successfully"
: `with \`${workflowRun.conclusion}\``;
const body = [
marker,
"## LiteLLM Rust workflow",
"",
`Workflow completed ${result} for \`${headSha}\``,
"",
`[View workflow run](${runUrl})`,
].join("\n");
const comments = await github.paginate(github.rest.issues.listComments, {
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: pullRequestNumber,
per_page: 100,
});
const existing = comments.find(
(comment) =>
comment.user?.login === "github-actions[bot]" &&
comment.body?.startsWith(marker),
);
const currentPullRequest = (
await github.rest.pulls.get({
owner: context.repo.owner,
repo: context.repo.repo,
pull_number: pullRequestNumber,
})
).data;
if (
currentPullRequest.state !== "open" ||
currentPullRequest.head.repo?.full_name !==
`${context.repo.owner}/${context.repo.repo}` ||
currentPullRequest.head.sha !== headSha
) {
core.info("source workflow no longer matches the current pull request head");
return;
}
if (existing) {
await github.rest.issues.updateComment({
owner: context.repo.owner,
repo: context.repo.repo,
comment_id: existing.id,
body,
});
} else {
await github.rest.issues.createComment({
owner: context.repo.owner,
repo: context.repo.repo,
issue_number: pullRequestNumber,
body,
});
}

View file

@ -7,6 +7,7 @@ on:
- ".cargo/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/actions/setup-uv-with-retries/**"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
@ -22,6 +23,7 @@ on:
- ".cargo/**"
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/actions/setup-uv-with-retries/**"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/verify_linux_native_wheel.py"
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
@ -34,102 +36,89 @@ concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
env:
CARGO_TERM_COLOR: always
jobs:
rust-checks:
name: rustfmt, clippy, test
rust-lint:
runs-on: ubuntu-latest
timeout-minutes: 10
defaults:
run:
working-directory: litellm-rust
env:
CARGO_TERM_COLOR: always
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Rust
run: rustup toolchain install
- run: rustup toolchain install --no-self-update
- name: Cache Cargo registry and target
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
- run: cargo fmt --check
- uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cargo/registry
~/.cargo/git
litellm-rust/target
key: ${{ runner.os }}-cargo-${{ hashFiles('rust-toolchain.toml', 'litellm-rust/Cargo.lock') }}
key: ${{ runner.os }}-cargo-${{ github.job }}-${{ hashFiles('rust-toolchain.toml', '.cargo/**', 'litellm-rust/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-
${{ runner.os }}-cargo-${{ github.job }}-
- name: Check Rust formatting
run: cargo fmt --check
- run: cargo clippy --workspace --all-targets --locked -- -D warnings
- name: Run Clippy
run: cargo clippy --workspace --all-targets --locked -- -D warnings
- run: cargo clippy -p litellm-core --all-targets --features bedrock-auth --locked -- -D warnings
- name: Run Clippy with Bedrock auth
run: cargo clippy -p litellm-core --all-targets --features bedrock-auth --locked -- -D warnings
- run: cargo clippy -p litellm-ai-gateway --all-targets --all-features --locked -- -D warnings
- name: Run Clippy with all gateway features
run: cargo clippy -p litellm-ai-gateway --all-targets --all-features --locked -- -D warnings
- name: Run Rust tests
run: cargo test --workspace --locked
- name: Run core tests with Bedrock auth
run: cargo test -p litellm-core --features bedrock-auth --locked
# Not --all-features: python-config links libpython, which this job does not install.
- name: Run gateway tests with the server feature
run: cargo test -p litellm-ai-gateway --features server --locked
release-wheel:
name: release wheel
rust-test:
runs-on: ubuntu-latest
timeout-minutes: 20
permissions:
contents: read
env:
CARGO_TERM_COLOR: always
timeout-minutes: 30
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
- uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
- uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Set up Rust
run: rustup toolchain install
- run: rustup toolchain install --no-self-update
- name: Build release wheel
run: uv build --wheel --out-dir dist
- uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |
~/.cargo/registry
~/.cargo/git
litellm-rust/target
key: ${{ runner.os }}-cargo-${{ github.job }}-${{ hashFiles('rust-toolchain.toml', '.cargo/**', 'litellm-rust/Cargo.lock') }}
restore-keys: |
${{ runner.os }}-cargo-${{ github.job }}-
- name: Build panic contract wheel
run: >-
- run: cargo test --workspace --locked
working-directory: litellm-rust
- run: cargo test -p litellm-core --features bedrock-auth --locked
working-directory: litellm-rust
- run: cargo test -p litellm-ai-gateway --features server --locked
working-directory: litellm-rust
- run: uv build --wheel --out-dir dist
- run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
- run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
- run: >-
uv build --wheel --out-dir panic-dist
--config-setting "maturin.build-args=--features panic-test,extension-module"
- name: Smoke-test native panic unwinding
run: python .github/scripts/smoke_test_native_wheel.py panic-dist/*.whl
- name: Verify stripped native extension
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl
- name: Test native route wheel
run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
- run: python .github/scripts/smoke_test_native_wheel.py panic-dist/*.whl

View file

@ -52,7 +52,7 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a
Python max line length is 120, not 88
When you fix violations gated by `ruff-strict-budget.json`, `type-discipline-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered limits so the ceilings ratchet down instead of leaving stale headroom. It measures the working tree, so it must contain exactly the fixes you're committing
Never edit or commit `ruff-strict-budget.json`, `type-discipline-budget.json`, `basedpyright-code-budget.json`, or `test-quality-budget.json` on a PR branch, and don't run `make lint-budget-update` there. A scheduled Devin automation lowers the limits on `litellm_internal_staging` in its own PR by exactly what landed since the last ratchet, so concurrent PRs don't fight over the same `"limit"` lines. If your branch already carries a budget edit, drop it before opening the PR
`make check` (f.k.a. `make pre-commit`, which still works identically as an alias) saves its complete output to a log file in .git (overwriting previous logs) and prints that path as its first and last output lines. To inspect a run, read or grep that log instead of re-running the multi-minute checks just to see a different slice

View file

@ -84,7 +84,7 @@
"limit": 56
},
"reportPrivateUsage": {
"limit": 1808
"limit": 1804
},
"reportRedeclaration": {
"limit": 8
@ -135,7 +135,7 @@
"limit": 21
},
"reportUnusedFunction": {
"limit": 138
"limit": 136
},
"reportUnusedImport": {
"limit": 542

View file

@ -25,6 +25,27 @@ class BatchCostUsageResult:
failed_requests: int
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
def batch_cost_is_final(batch: Batch) -> bool:
"""Whether this retrieve of the batch is the one to account its cost from.
A batch still in flight has nothing to price, and a "completed" batch can report
no output_file_id for a moment before the output populates; pricing either records
$0 under the batch's single spend row and pins it there. Final means a completed
batch whose output file has arrived or whose counts prove no line succeeded, or
any other terminal status (failed, cancelled, expired).
"""
if batch.status not in _TERMINAL_BATCH_STATUSES:
return False
if batch.status not in _COMPLETED_BATCH_STATUSES or batch.output_file_id is not None:
return True
request_counts: Final = batch.request_counts
return request_counts is not None and request_counts.total > 0 and request_counts.completed == 0
async def calculate_batch_cost_and_usage(
file_content_dictionary: list[dict],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic"],

View file

@ -36,7 +36,7 @@ from litellm._logging import (
verbose_logger,
)
from litellm._uuid import uuid
from litellm.batches.batch_utils import _handle_completed_batch
from litellm.batches.batch_utils import _handle_completed_batch, batch_cost_is_final
from litellm.caching.caching import DualCache, InMemoryCache
from litellm.caching.caching_handler import LLMCachingHandler
from litellm.constants import (
@ -2899,13 +2899,6 @@ class Logging(LiteLLMLoggingBaseClass):
): # polling job will query these frequently, don't spam db logs
return
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
)
# check if file id is a unified file id
is_base64_unified_file_id: Final = _is_base64_encoded_unified_file_id(result.id)
batch_cost: Final = kwargs.get("batch_cost", None)
batch_usage = kwargs.get("batch_usage", None)
batch_models = kwargs.get("batch_models", None)
@ -2913,9 +2906,7 @@ class Logging(LiteLLMLoggingBaseClass):
batch_failed_requests: Final = kwargs.get("batch_failed_requests", None)
has_explicit_batch_data: Final = all(x is not None for x in (batch_cost, batch_usage, batch_models))
should_compute_batch_data: Final = (
not is_base64_unified_file_id or not has_explicit_batch_data and result.status == "completed"
)
should_compute_batch_data: Final = not has_explicit_batch_data and batch_cost_is_final(result)
if has_explicit_batch_data:
result._hidden_params["response_cost"] = batch_cost
result._hidden_params["batch_models"] = batch_models

View file

@ -17,6 +17,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.llms.azure.common_utils import BaseAzureLLM
from litellm.llms.azure_ai.common_utils import is_foundry_model_inference_base
from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error
from litellm.llms.openai.openai import OpenAIConfig
from litellm.llms.xai.chat.transformation import XAIChatConfig
@ -42,12 +43,37 @@ NON_OPENAI_SPEC_MESSAGE_FIELDS: Final = (
)
class AzureAIGPT5Config(OpenAIGPT5Config):
@classmethod
def _model_map_lookup_name(cls, model: str) -> str:
"""Normalise a Foundry routing name to its cost-map key, when the map has one.
A Foundry deployment and its OpenAI-hosted namesake are different products with
different capabilities, so ``azure_ai/<model>`` is the entry to read whenever the map
carries it. Most gpt-5-family names have no ``azure_ai/`` row, though, and prefixing
those anyway costs them every flag: ``get_llm_provider`` re-resolves an ``azure_ai/``
name to the azure provider when a global AZURE_AI_API_BASE points at an
openai.azure.com host, ``azure/<model>`` is not a key either, so the lookup lands
nowhere and every effort answer degrades to False. A missing key defers to the base
resolver instead.
"""
prefixed: Final = model if model.startswith("azure_ai/") else f"azure_ai/{model}"
return prefixed if prefixed in litellm.model_cost else super()._model_map_lookup_name(model)
azureAIGPT5Config: Final = AzureAIGPT5Config()
class AzureAIStudioConfig(OpenAIConfig):
def get_supported_openai_params(self, model: str) -> list:
model_supports_tool_choice = True # azure ai supports this by default
if not supports_tool_choice(model=f"azure_ai/{model}"):
model_supports_tool_choice = False
supported_params = super().get_supported_openai_params(model)
supported_params = (
azureAIGPT5Config.get_supported_openai_params(model)
if azureAIGPT5Config.is_model_gpt_5_model(model)
else super().get_supported_openai_params(model)
)
if not model_supports_tool_choice:
filtered_supported_params: Final = []
for param in supported_params:
@ -61,6 +87,27 @@ class AzureAIStudioConfig(OpenAIConfig):
return supported_params
def map_openai_params(
self,
non_default_params: dict[str, object], # mutable-ok: OpenAIConfig.map_openai_params signature
optional_params: dict[str, object], # mutable-ok: OpenAIConfig.map_openai_params signature
model: str,
drop_params: bool,
) -> dict[str, object]: # mutable-ok: OpenAIConfig.map_openai_params signature
if not azureAIGPT5Config.is_model_gpt_5_model(model):
return super().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=drop_params,
)
return azureAIGPT5Config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=drop_params,
)
def _supports_stop_reason(self, model: str) -> bool:
"""
Check if the model supports stop tokens.

View file

@ -0,0 +1,210 @@
"""
Support for Mistral Voxtral text-to-speech via ``/v1/audio/speech``.
API reference: https://docs.mistral.ai/api/#tag/audio/operation/audio_speech_v1_audio_speech_post
"""
import base64
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
import httpx
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.text_to_speech.transformation import (
BaseTextToSpeechConfig,
TextToSpeechRequestData,
)
from litellm.secret_managers.main import get_secret_str
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import HttpxBinaryResponseContent
class MistralTextToSpeechException(BaseLLMException):
pass
class MistralTextToSpeechConfig(BaseTextToSpeechConfig):
TTS_BASE_URL: Final[str] = "https://api.mistral.ai/v1"
AUDIO_CONTENT_TYPES: Final[MappingProxyType[str, str]] = MappingProxyType(
{
"mp3": "audio/mpeg",
"wav": "audio/wav",
"pcm": "audio/pcm",
"flac": "audio/flac",
"opus": "audio/ogg",
}
)
DROPPED_RESPONSE_HEADERS: Final[frozenset[str]] = frozenset(
{"content-encoding", "transfer-encoding", "content-length", "content-type"}
)
OPENAI_VOICE_ALIASES: Final[MappingProxyType[str, str]] = MappingProxyType(
{
"alloy": "en_paul_neutral",
"echo": "gb_oliver_neutral",
"fable": "en_paul_cheerful",
"onyx": "en_paul_confident",
"nova": "gb_jane_sarcasm",
"shimmer": "gb_jane_sarcasm",
}
)
def get_supported_openai_params(self, model: str) -> list: # mutable-ok: base class contract returns a plain list
return ["voice", "response_format"] # mutable-ok: base class contract returns a plain list
def _map_openai_voice(self, voice_id: str) -> str:
return self.OPENAI_VOICE_ALIASES.get(voice_id.lower(), voice_id)
def _resolve_voice_id(self, voice: object) -> str | None:
if isinstance(voice, str) and voice.strip():
return self._map_openai_voice(voice.strip())
if isinstance(voice, Mapping):
candidates: Final = (voice.get(key) for key in ("voice_id", "id", "name"))
resolved: Final = next(
(candidate.strip() for candidate in candidates if isinstance(candidate, str) and candidate.strip()),
None,
)
return self._map_openai_voice(resolved) if resolved else None
return None
def map_openai_params(
self,
model: str,
optional_params: Mapping[str, object],
voice: object = None,
drop_params: bool = False,
kwargs: Mapping[str, object] | None = None,
) -> tuple[str | None, dict]: # mutable-ok: base class contract returns a plain dict
response_format: Final = optional_params.get("response_format")
ref_audio: Final = kwargs.get("ref_audio") if kwargs else None
voice_id_kwarg: Final = kwargs.get("voice_id") if kwargs else None
mapped_voice: Final = self._resolve_voice_id(voice) or self._resolve_voice_id(voice_id_kwarg)
mapped_params: Final = { # mutable-ok: base class contract returns a plain dict
key: value
for key, value in (("response_format", response_format), ("ref_audio", ref_audio))
if isinstance(value, str)
}
return mapped_voice, mapped_params
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
api_key: str | None = None,
api_base: str | None = None,
) -> dict: # mutable-ok: base class contract returns a plain dict
resolved_key: Final = api_key or get_secret_str("MISTRAL_API_KEY")
if resolved_key is None:
raise MistralTextToSpeechException(
status_code=401,
message="Mistral API key is required. Set MISTRAL_API_KEY or pass api_key.",
)
return { # mutable-ok: base class contract returns a plain dict
**headers,
"Authorization": f"Bearer {resolved_key}",
"Content-Type": "application/json",
}
def get_complete_url(
self,
model: str,
api_base: str | None,
litellm_params: Mapping[str, object],
) -> str:
configured_base: Final = (api_base or self.TTS_BASE_URL).rstrip("/")
versioned_base: Final = configured_base if configured_base.endswith("/v1") else f"{configured_base}/v1"
return f"{versioned_base}/audio/speech"
def transform_text_to_speech_request(
self,
model: str,
input: str,
voice: str | None,
optional_params: Mapping[str, object],
litellm_params: Mapping[str, object],
headers: Mapping[str, str],
) -> TextToSpeechRequestData:
response_format: Final = optional_params.get("response_format")
ref_audio: Final = optional_params.get("ref_audio")
request_data: Final[TextToSpeechRequestData] = {
"dict_body": {
"model": model,
"input": input,
**({"voice_id": voice} if voice else {}),
**({"response_format": response_format} if isinstance(response_format, str) else {}),
**({"ref_audio": ref_audio} if isinstance(ref_audio, str) else {}),
},
"headers": {"Content-Type": "application/json"},
}
return request_data
def _requested_content_type(self, request: httpx.Request) -> str:
request_body: Final = json.loads(request.content or b"{}")
requested_format: Final = request_body.get("response_format")
if not isinstance(requested_format, str):
return "audio/mpeg"
return self.AUDIO_CONTENT_TYPES.get(requested_format, "audio/mpeg")
def transform_text_to_speech_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
) -> "HttpxBinaryResponseContent":
from litellm.types.llms.openai import HttpxBinaryResponseContent
try:
response_json: Final = raw_response.json()
except (json.JSONDecodeError, ValueError):
raise MistralTextToSpeechException(
status_code=raw_response.status_code,
message=f"Non-JSON response from Mistral speech API: {raw_response.text[:500]}",
headers=raw_response.headers,
)
audio_b64: Final = response_json.get("audio_data")
if not isinstance(audio_b64, str) or not audio_b64:
raise MistralTextToSpeechException(
status_code=500,
message=f"No audio_data in Mistral speech response. Response keys: {tuple(response_json.keys())}",
headers=raw_response.headers,
)
try:
audio_bytes: Final = base64.b64decode(audio_b64, validate=True)
except ValueError:
raise MistralTextToSpeechException(
status_code=500,
message="Invalid base64 audio_data in Mistral speech response.",
headers=raw_response.headers,
)
retained_headers: Final = tuple(
(key, value)
for key, value in raw_response.headers.items()
if key.lower() not in self.DROPPED_RESPONSE_HEADERS
)
response_headers: Final = retained_headers + (
("content-length", str(len(audio_bytes))),
("content-type", self._requested_content_type(raw_response.request)),
)
binary_response: Final = httpx.Response(
status_code=200,
headers=response_headers,
content=audio_bytes,
request=raw_response.request,
)
return HttpxBinaryResponseContent(binary_response)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict | httpx.Headers, # mutable-ok: BaseLLMException takes a plain dict or httpx.Headers
) -> BaseLLMException:
return MistralTextToSpeechException(
message=error_message,
status_code=status_code,
headers=headers,
)

View file

@ -8389,6 +8389,34 @@ def speech(
client=client,
_is_async=aspeech or False,
)
elif custom_llm_provider == "mistral":
from litellm.llms.mistral.audio_speech.transformation import (
MistralTextToSpeechConfig,
)
mistral_tts_config: Final = text_to_speech_provider_config or MistralTextToSpeechConfig()
if api_base is not None:
litellm_params_dict["api_base"] = api_base
if api_key is not None:
litellm_params_dict["api_key"] = api_key
mistral_voice: Final[str | None] = voice if isinstance(voice, str) else None
response = base_llm_http_handler.text_to_speech_handler(
model=model,
input=input,
voice=mistral_voice,
text_to_speech_provider_config=mistral_tts_config,
text_to_speech_optional_params=optional_params,
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params_dict,
logging_obj=logging_obj,
timeout=timeout,
extra_headers=extra_headers,
client=client,
_is_async=aspeech or False,
)
elif custom_llm_provider == "aws_polly":
from litellm.llms.aws_polly.text_to_speech.transformation import (
AWSPollyTextToSpeechConfig,

View file

@ -3485,6 +3485,55 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"azure_ai/gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"source": "https://ai.azure.com/catalog/models/gpt-6-astra",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_cache_breakpoint": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure_ai/gpt-5.5": {
"deprecation_date": "2027-10-26",
"cache_read_input_token_cost": 5e-07,
@ -7189,7 +7238,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
@ -7455,7 +7504,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
@ -34898,9 +34947,9 @@
"supports_audio_input": true
},
"mistral/voxtral-mini-tts-latest": {
"input_cost_per_character": 1.6e-05,
"litellm_provider": "mistral",
"mode": "audio_speech",
"output_cost_per_character": 1.6e-05,
"source": "https://docs.mistral.ai/models/model-cards/voxtral-tts-26-03",
"supported_endpoints": [
"/v1/audio/speech"
@ -56467,9 +56516,9 @@
"supports_audio_input": true
},
"mistral/voxtral-mini-tts-2603": {
"input_cost_per_character": 1.6e-05,
"litellm_provider": "mistral",
"mode": "audio_speech",
"output_cost_per_character": 1.6e-05,
"source": "https://docs.mistral.ai/models/model-cards/voxtral-tts-26-03",
"supported_endpoints": [
"/v1/audio/speech"

View file

@ -14,6 +14,7 @@ import time
import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
import litellm
@ -84,6 +85,25 @@ else:
RESPONSES_SESSION_CALL_TYPES: Final = frozenset({CallTypes.responses.value, CallTypes.aresponses.value})
def _is_batch_cost_row(payload: SpendLogsPayload) -> bool:
return payload.get("call_type") == CallTypes.aretrieve_batch.value and payload.get("status") == "success"
_BATCH_COST_CLAIM_FIELDS: Final = frozenset({"request_id", "call_type", "spend", "startTime", "endTime", "status"})
def _batch_cost_row_to_write(payload: SpendLogsPayload, disable_spend_logs: bool) -> Mapping[str, object]:
"""Reduce a batch's cost row to what tells the retrieves apart when logging is off.
A proxy run with spend logs disabled still needs one row per batch to charge it once,
so the row is written either way, but it carries no request of its own: no metadata,
no requester IP, no key, model, or token counts (LIT-7048).
"""
if disable_spend_logs is False:
return payload
return MappingProxyType({field: value for field, value in payload.items() if field in _BATCH_COST_CLAIM_FIELDS})
class _SpendBatch(Protocol):
litellm_usertable: BatchTable
litellm_verificationtoken: BatchTable
@ -215,7 +235,12 @@ class DBSpendUpdateWriter:
start_time: datetime | None,
end_time: datetime | None,
response_cost: float | None,
) -> None:
) -> bool:
"""Record the request's spend, answering whether its cost still needs charging.
False only for a batch retrieve whose cost row another retrieve already wrote,
so the caller leaves the key, team, and user counters alone (LIT-7048).
"""
from litellm.proxy.proxy_server import (
disable_spend_logs,
litellm_proxy_budget_name,
@ -232,7 +257,7 @@ class DBSpendUpdateWriter:
team_id,
)
if ProxyUpdateSpend.disable_spend_updates() is True:
return
return True
if token is not None and isinstance(token, str) and token.startswith("sk-"):
hashed_token = hash_token(token=token)
else:
@ -262,11 +287,12 @@ class DBSpendUpdateWriter:
if team_id is not None and team_id != "":
payload["team_id"] = team_id
if not await self._record_spend_log(
payload=payload, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs
):
return False
if disable_spend_logs is False:
await self._insert_spend_log_to_db(
payload=payload,
prisma_client=prisma_client,
)
await self._enqueue_tool_usage_transaction(
payload=payload,
completion_response=completion_response,
@ -306,6 +332,7 @@ class DBSpendUpdateWriter:
)
verbose_proxy_logger.debug("Runs spend update on all tables")
return True
except Exception:
spend_log_error(
"Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue "
@ -318,7 +345,102 @@ class DBSpendUpdateWriter:
org_id,
end_user_id,
)
return
return True
async def _record_spend_log(
self, payload: SpendLogsPayload, prisma_client: "PrismaClient | None", disable_spend_logs: bool
) -> bool:
if prisma_client is not None and _is_batch_cost_row(payload):
return await self._claim_batch_cost_spend_log(
payload=payload, prisma_client=prisma_client, disable_spend_logs=disable_spend_logs
)
if disable_spend_logs is False:
await self._insert_spend_log_to_db(payload=payload, prisma_client=prisma_client)
return True
async def _claim_batch_cost_spend_log(
self, payload: SpendLogsPayload, prisma_client: "PrismaClient", disable_spend_logs: bool
) -> bool:
"""Write the batch's cost row now, or learn that another retrieve already did.
Every retrieve of one batch shares this row, so the insert that lands first owns
the charge and every later one finds the row and charges nothing (LIT-7048). Only
a row that recorded a charge counts: a failed retrieve, a request whose client
picked the batch id as its call id, and the $0 row an older proxy left behind
while the batch was still running all leave the charge to be made.
"""
from litellm.repositories.table_repositories import SpendLogsRepository
request_id: Final = payload["request_id"]
row: Final = _batch_cost_row_to_write(payload, disable_spend_logs)
spend_logs: Final = SpendLogsRepository(prisma_client).table
try:
claimed: Final = await spend_logs.create_many(
data=[prisma_client.jsonify_object(row)], # mutable-ok: prisma create_many takes a list
skip_duplicates=True,
)
if claimed == 1:
return True
existing: Final = await spend_logs.find_unique(
where={"request_id": request_id} # mutable-ok: prisma where clause
)
except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; an unreachable DB queues the row like any other spend log
verbose_proxy_logger.warning(
"Could not claim spend row %s for a batch's cost, queueing it: %s", request_id, e
)
await self._insert_spend_log_to_db(payload=prisma_client.jsonify_object(row), prisma_client=prisma_client)
return True
if existing is None or existing.call_type != CallTypes.aretrieve_batch.value or existing.status != "success":
verbose_proxy_logger.warning(
"Spend row %s belongs to a %s request, so this batch's cost is charged without a row of its own",
request_id,
getattr(existing, "call_type", None),
)
return True
if existing.spend > 0:
verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id)
return False
return await self._take_over_uncharged_batch_cost_row(payload=payload, prisma_client=prisma_client, row=row)
async def _take_over_uncharged_batch_cost_row(
self, payload: SpendLogsPayload, prisma_client: "PrismaClient", row: Mapping[str, object]
) -> bool:
"""Take the batch's cost row over from the poll that left it charging nothing.
A pre-upgrade proxy wrote that row every time it polled the batch while it was still
running, so the charge is still to be made and the row still has to end up carrying
it. The row stops matching the moment it carries a charge, so it is one retrieve that
takes it over and charges, and every later one reads the charge and charges nothing.
"""
from litellm.repositories.table_repositories import SpendLogsRepository
request_id: Final = payload["request_id"]
if payload["spend"] <= 0:
verbose_proxy_logger.debug(
"Cost tracking skipped: this batch costs nothing and spend row %s says so", request_id
)
return False
try:
taken_over: Final = await SpendLogsRepository(prisma_client).table.update_many(
data=prisma_client.jsonify_object(
MappingProxyType({field: value for field, value in row.items() if field != "request_id"})
),
where={ # mutable-ok: prisma where clause
"request_id": request_id,
"call_type": CallTypes.aretrieve_batch.value,
"status": "success",
"spend": 0.0,
},
)
except Exception as e: # noqa: BLE001 # prisma raises its own hierarchy; the next retrieve takes the row over
verbose_proxy_logger.warning(
"Could not take over spend row %s, leaving this batch's cost to the next retrieve: %s", request_id, e
)
return False
if taken_over == 0:
verbose_proxy_logger.debug("Cost tracking skipped: spend row %s already charged this batch", request_id)
return False
return True
async def _enqueue_tool_usage_transaction(
self,

View file

@ -6,6 +6,7 @@ from typing import TYPE_CHECKING, Any, Final, cast
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.batches.batch_utils import batch_cost_is_final
from litellm.constants import BACKGROUND_INTERACTION_COST_POLLING_ENABLED
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import (
@ -37,6 +38,7 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import (
from litellm.proxy.utils import ProxyUpdateSpend
from litellm.types.utils import (
CallTypes,
LiteLLMBatch,
StandardLoggingPayload,
StandardLoggingPayloadErrorInformation,
)
@ -248,6 +250,18 @@ class _ProxyDBLogger(CustomLogger):
)
_write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata)
budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata)
if (
isinstance(completion_response, LiteLLMBatch)
and kwargs.get("call_type") == CallTypes.aretrieve_batch.value
and not batch_cost_is_final(completion_response)
):
verbose_proxy_logger.debug(
"Cost tracking deferred for batch %s still in status %s",
completion_response.id,
completion_response.status,
)
await _release_budget_reservation(budget_reservation=budget_reservation)
return
user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None))
team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None))
org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None))
@ -285,7 +299,7 @@ class _ProxyDBLogger(CustomLogger):
call_type=call_type,
):
## UPDATE DATABASE
await _update_database_and_spend_counters(
charged: Final = await _update_database_and_spend_counters(
proxy_logging_obj=proxy_logging_obj,
increment_spend_counters=increment_spend_counters,
user_api_key=user_api_key,
@ -302,6 +316,8 @@ class _ProxyDBLogger(CustomLogger):
request_tags=tags,
model_access_groups=model_access_groups,
)
if not charged:
return
# update cache (fire-and-forget for backward compat:
# cached object fields, soft budget alerts, etc.)
@ -578,9 +594,9 @@ async def _update_database_and_spend_counters(
budget_reservation: dict | None,
request_tags: list[str] | None = None,
model_access_groups: Sequence[str] | None = None,
) -> None:
) -> bool:
try:
await proxy_logging_obj.db_spend_update_writer.update_database(
charged: Final = await proxy_logging_obj.db_spend_update_writer.update_database(
token=user_api_key,
response_cost=response_cost,
user_id=user_id,
@ -605,6 +621,9 @@ async def _update_database_and_spend_counters(
"Failed to invalidate budget reservation counters after release failed"
)
raise
if not charged:
await _release_budget_reservation(budget_reservation=budget_reservation)
return False
try:
await increment_spend_counters(
@ -630,6 +649,7 @@ async def _update_database_and_spend_counters(
finally:
budget_reservation["finalized"] = True
raise
return True
async def _release_budget_reservation(budget_reservation: dict | None) -> None:

View file

@ -2882,6 +2882,28 @@ def _add_guardrails_from_policies_in_metadata(
)
def add_guardrails_from_auth_metadata(
user_api_key_dict: UserAPIKeyAuth,
data: dict, # mutable-ok: writes guardrails into the live request dict, same contract as the helpers it wraps
metadata_variable_name: str,
) -> None:
"""Resolve key, team, and project guardrails, direct and via policies, onto the request metadata."""
_add_guardrails_from_key_or_team_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=user_api_key_dict.project_metadata,
data=data,
metadata_variable_name=metadata_variable_name,
)
_add_guardrails_from_policies_in_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=user_api_key_dict.project_metadata,
data=data,
metadata_variable_name=metadata_variable_name,
)
async def move_guardrails_to_metadata(
data: dict,
_metadata_variable_name: str,
@ -2914,22 +2936,8 @@ async def move_guardrails_to_metadata(
data.pop("policies", None)
return
# Check key/team/project-level guardrails
_add_guardrails_from_key_or_team_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=project_metadata,
data=data,
metadata_variable_name=_metadata_variable_name,
)
#########################################################################################
# Add guardrails from policies attached to key/team/project metadata
#########################################################################################
_add_guardrails_from_policies_in_metadata(
key_metadata=user_api_key_dict.metadata,
team_metadata=user_api_key_dict.team_metadata,
project_metadata=project_metadata,
add_guardrails_from_auth_metadata(
user_api_key_dict=user_api_key_dict,
data=data,
metadata_variable_name=_metadata_variable_name,
)

View file

@ -15,6 +15,7 @@ from typing import (
runtime_checkable,
)
from litellm.batches.batch_utils import batch_cost_is_final
from litellm.proxy._types import ProxyException
from litellm.repositories.table_repositories import (
ManagedFileRepository,
@ -1357,12 +1358,7 @@ def _completed_batch_safe_to_retire(response: "LiteLLMBatch") -> bool:
enumerated the batch and none succeeded. A zero or unknown total means counts
are unreported, so stay eligible and let the next poller pass revisit it. (#37713)
"""
if response.output_file_id is not None:
return True
request_counts = response.request_counts
if request_counts is None:
return False
return request_counts.total > 0 and request_counts.completed == 0
return batch_cost_is_final(response)
async def update_batch_in_database(

View file

@ -152,7 +152,7 @@ from litellm.proxy.hooks.parallel_request_limiter_v3 import (
from litellm.proxy.hooks.sensitive_data_routing import (
_PROXY_SensitiveDataRoutingHandler,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_guardrails_from_auth_metadata
from litellm.proxy.management_helpers.key_settings_audit import with_settings_updated_at
from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor
from litellm.repositories.budget_repository import BudgetRepository
@ -924,7 +924,13 @@ class ProxyLogging:
"incoming_bearer_token": kwargs.get("incoming_bearer_token"),
"metadata": {"headers": kwargs.get("headers") or {}},
}
user_api_key_auth: Final = kwargs.get("user_api_key_auth")
if isinstance(user_api_key_auth, UserAPIKeyAuth):
add_guardrails_from_auth_metadata(
user_api_key_dict=user_api_key_auth,
data=synthetic_data,
metadata_variable_name="metadata",
)
return synthetic_data
def _convert_llm_result_to_mcp_response(self, llm_result, request_obj) -> MCPPreCallResponseObject | None:

View file

@ -403,6 +403,10 @@ _NO_SESSION_KWARGS: Final[Mapping[str, Mapping[str, object]]] = MappingProxyType
_SESSION_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _as_retry_skipped_deployment_ids(value: object) -> tuple[str, ...]:
return tuple(item for item in value if isinstance(item, str)) if isinstance(value, tuple) else ()
def _with_router_resolved_session_model(session: object, model_name: str) -> Mapping[str, Mapping[str, object]]:
"""
Realtime client-secret requests carry the model inside ``session`` as well, and the caller's copy of it still
@ -4449,7 +4453,7 @@ class Router:
self.fail_calls[model_name] += 1
raise e
async def aspeech(self, model: str, input: str, voice: str, **kwargs):
async def aspeech(self, model: str, input: str, voice: str | None = None, **kwargs):
"""
Example Usage:
@ -4501,7 +4505,7 @@ class Router:
)
raise e
async def _aspeech(self, model: str, input: str, voice: str, **kwargs):
async def _aspeech(self, model: str, input: str, voice: str | None = None, **kwargs):
model_name: Final = model
try:
verbose_router_logger.debug("Inside _aspeech()- model: %s; kwargs: %s", model, kwargs)
@ -4525,7 +4529,7 @@ class Router:
**{
**data,
"input": input,
"voice": voice,
"voice": data.get("voice") if voice is None else voice,
"client": model_client,
**kwargs,
}
@ -7458,6 +7462,21 @@ class Router:
Context_Policy_Fallbacks={content_policy_fallbacks}",
)
@staticmethod
def _deployment_ids_to_skip_on_retry(exception: Exception, already_skipped: object) -> tuple[str, ...]:
failed_deployment_id: Final[str | None] = getattr(exception, "failed_deployment_id", None)
status_code: Final = getattr(exception, "status_code", None)
if not failed_deployment_id or not isinstance(status_code, int):
return ()
if litellm._should_retry(status_code): # pyright: ignore[reportPrivateUsage] # as in should_retry_this_error
return ()
already_skipped_ids: Final = _as_retry_skipped_deployment_ids(already_skipped)
skipped: Final = tuple(sorted(frozenset((*already_skipped_ids, failed_deployment_id))))
verbose_router_logger.debug(
"Retry skips deployments that already answered %s to this request: %s", status_code, skipped
)
return skipped
@tracer.wrap()
async def async_function_with_retries(self, *args, **kwargs):
verbose_router_logger.debug("Inside async function with retries.")
@ -7553,6 +7572,12 @@ class Router:
## LOGGING
if num_retries > 0:
kwargs = self.log_retry(kwargs=kwargs, e=original_exception)
first_skipped_ids: Final = self._deployment_ids_to_skip_on_retry(
exception=original_exception,
already_skipped=kwargs.get("_retry_skipped_deployment_ids"),
)
if first_skipped_ids:
kwargs["_retry_skipped_deployment_ids"] = first_skipped_ids # rebind-ok: the next attempt reads it
else:
raise
@ -7622,6 +7647,12 @@ class Router:
except Exception:
raise e
skipped_ids = self._deployment_ids_to_skip_on_retry(
exception=e,
already_skipped=kwargs.get("_retry_skipped_deployment_ids"),
)
if skipped_ids:
kwargs["_retry_skipped_deployment_ids"] = skipped_ids # rebind-ok: the next attempt reads it
_timeout = self._time_to_sleep_before_retry(
e=e,
remaining_retries=remaining_retries,
@ -12454,7 +12485,7 @@ class Router:
## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
_target_order: Final = (request_kwargs or {}).pop("_target_order", None)
healthy_deployments = litellm.utils._get_order_filtered_deployments(
healthy_deployments = litellm.utils.get_order_filtered_deployments(
cast(list[dict], healthy_deployments), target_order=_target_order
)
@ -12462,11 +12493,24 @@ class Router:
## this request via weighted-failover. Always honored, regardless of the
## router-level flag, so a stale exclusion key on kwargs cannot escape.
_excluded_deployment_ids: Final = (request_kwargs or {}).pop("_excluded_deployment_ids", None)
healthy_deployments = litellm.utils._get_excluded_filtered_deployments(
healthy_deployments = litellm.utils.get_excluded_filtered_deployments(
cast(list[dict], healthy_deployments),
excluded_deployment_ids=_excluded_deployment_ids,
)
## RETRY SKIP ## -> drop deployments that already refused this request with a
## non-retryable status, unless that leaves nothing, so the caller still gets
## the provider's own error instead of a no-deployments error.
_retry_skipped_deployment_ids: Final = _as_retry_skipped_deployment_ids(
request_kwargs.pop("_retry_skipped_deployment_ids", None) if request_kwargs else None
)
healthy_deployments = (
litellm.utils.get_excluded_filtered_deployments(
healthy_deployments, excluded_deployment_ids=_retry_skipped_deployment_ids
)
or healthy_deployments
)
if len(healthy_deployments) == 0:
exception: Final = await async_raise_no_deployment_exception(
litellm_router_instance=self,
@ -13359,7 +13403,7 @@ class Router:
## ORDER FILTERING ## -> if user set 'order' in deployments, return deployments with lowest order (e.g. order=1 > order=2)
_target_order: Final = (request_kwargs or {}).pop("_target_order", None)
healthy_deployments = litellm.utils._get_order_filtered_deployments(
healthy_deployments = litellm.utils.get_order_filtered_deployments(
healthy_deployments, target_order=_target_order
)
@ -13367,11 +13411,22 @@ class Router:
## this request via weighted-failover. See async counterpart in
## async_get_healthy_deployments for details.
_excluded_deployment_ids: Final = (request_kwargs or {}).pop("_excluded_deployment_ids", None)
healthy_deployments = litellm.utils._get_excluded_filtered_deployments(
healthy_deployments = litellm.utils.get_excluded_filtered_deployments(
healthy_deployments,
excluded_deployment_ids=_excluded_deployment_ids,
)
## RETRY SKIP ## -> see async counterpart in async_get_healthy_deployments.
_retry_skipped_deployment_ids: Final = _as_retry_skipped_deployment_ids(
request_kwargs.pop("_retry_skipped_deployment_ids", None) if request_kwargs else None
)
healthy_deployments = (
litellm.utils.get_excluded_filtered_deployments(
healthy_deployments, excluded_deployment_ids=_retry_skipped_deployment_ids
)
or healthy_deployments
)
if len(healthy_deployments) == 0:
model_ids = self.get_model_ids(model_name=model)
_cooldown_time = self.cooldown_cache.get_min_cooldown(

View file

@ -2835,8 +2835,9 @@ class ComplexityRouter(CustomLogger):
where the prompt never arrives as messages.
Probed on a COPY of request_kwargs because the owner pops routing bookkeeping off the
dict it is handed (`_target_order`, `_excluded_deployment_ids`), and this is a
speculative question about a model that may never be picked.
dict it is handed (`_target_order`, `_excluded_deployment_ids`,
`_retry_skipped_deployment_ids`), and this is a speculative question about a model
that may never be picked.
Every way the owner says "nothing here can serve this" is a negative verdict: no healthy
deployment for the group at all (BadRequestError, which ContextWindowExceededError

View file

@ -10,8 +10,8 @@ opt-in. none is opt-out everywhere except the azure gpt-5 family, whose config r
UnsupportedParamsError without an explicit true.
xhigh is gated on the request path by the openai and azure gpt-5 configs. max is not gated there at
all: every entry carrying supports_max_reasoning_effort is Claude-family, and
anthropic/chat/transformation.py gates max on the output_config path while its reasoning_effort
all: outside the gpt-6-astra rows every entry carrying supports_max_reasoning_effort is Claude-family,
and anthropic/chat/transformation.py gates max on the output_config path while its reasoning_effort
path maps any level to a thinking budget. Making max opt-in is a deliberate trade, then, since an
explicit flag is the only signal that the tier is a real one rather than litellm rounding the level
to a budget, and a missing flag costs advisory metadata rather than a rejected request.

View file

@ -4889,7 +4889,7 @@ def _get_deployment_order(deployment: dict | Any) -> int | None:
return order
def _get_order_filtered_deployments(healthy_deployments: list[dict], target_order: int | None = None) -> list:
def get_order_filtered_deployments(healthy_deployments: list[dict], target_order: int | None = None) -> list:
if target_order is not None:
return [d for d in healthy_deployments if _get_deployment_order(d) == target_order]
@ -4908,7 +4908,7 @@ def _get_order_filtered_deployments(healthy_deployments: list[dict], target_orde
return healthy_deployments
def _get_excluded_filtered_deployments(
def get_excluded_filtered_deployments(
healthy_deployments: list[dict],
excluded_deployment_ids: Iterable[str] | None = None,
) -> list:
@ -4919,10 +4919,12 @@ def _get_excluded_filtered_deployments(
across the remaining deployments in the same model group after one of them
has failed.
If the filter would leave no deployments, an empty list is returned so the
caller raises its usual no-deployments error and the weighted-failover
helper falls through to the cross-group fallback path. Returning the
original unfiltered list here would re-include the just-failed deployment.
If the filter would leave no deployments, an empty list is returned and the
caller decides what that means. Weighted failover lets it raise the usual
no-deployments error and fall through to the cross-group fallback path; the
retry skip in `async_get_healthy_deployments` deliberately falls back to the
unfiltered list, so a request every deployment refused still comes back with
the provider's own error rather than a no-deployments one.
"""
if not excluded_deployment_ids:
return healthy_deployments
@ -9451,6 +9453,12 @@ class ProviderConfigManager:
)
return MinimaxTextToSpeechConfig()
elif litellm.LlmProviders.MISTRAL == provider:
from litellm.llms.mistral.audio_speech.transformation import (
MistralTextToSpeechConfig,
)
return MistralTextToSpeechConfig()
elif litellm.LlmProviders.AWS_POLLY == provider:
from litellm.llms.aws_polly.text_to_speech.transformation import (
AWSPollyTextToSpeechConfig,

View file

@ -3485,6 +3485,55 @@
"supports_response_schema": true,
"supports_tool_choice": true
},
"azure_ai/gpt-6-astra": {
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
"cache_read_input_token_cost": 1e-06,
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
"input_cost_per_token": 1e-05,
"input_cost_per_token_above_272k_tokens": 2e-05,
"litellm_provider": "azure_ai",
"max_input_tokens": 922000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 5e-05,
"output_cost_per_token_above_272k_tokens": 7.5e-05,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
"search_context_size_medium": 0.01
},
"source": "https://ai.azure.com/catalog/models/gpt-6-astra",
"supported_endpoints": [
"/v1/chat/completions",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_cache_breakpoint": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": true,
"supports_xhigh_reasoning_effort": true
},
"azure_ai/gpt-5.5": {
"deprecation_date": "2027-10-26",
"cache_read_input_token_cost": 5e-07,
@ -7189,7 +7238,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
@ -7455,7 +7504,7 @@
],
"supports_computer_use": true,
"supports_function_calling": true,
"supports_max_reasoning_effort": true,
"supports_max_reasoning_effort": false,
"supports_minimal_reasoning_effort": false,
"supports_native_streaming": true,
"supports_none_reasoning_effort": true,
@ -34898,9 +34947,9 @@
"supports_audio_input": true
},
"mistral/voxtral-mini-tts-latest": {
"input_cost_per_character": 1.6e-05,
"litellm_provider": "mistral",
"mode": "audio_speech",
"output_cost_per_character": 1.6e-05,
"source": "https://docs.mistral.ai/models/model-cards/voxtral-tts-26-03",
"supported_endpoints": [
"/v1/audio/speech"
@ -56467,9 +56516,9 @@
"supports_audio_input": true
},
"mistral/voxtral-mini-tts-2603": {
"input_cost_per_character": 1.6e-05,
"litellm_provider": "mistral",
"mode": "audio_speech",
"output_cost_per_character": 1.6e-05,
"source": "https://docs.mistral.ai/models/model-cards/voxtral-tts-26-03",
"supported_endpoints": [
"/v1/audio/speech"

View file

@ -10,17 +10,12 @@ base.
Every rule is seeded at exactly its count on the day the gate landed, so the
suite's existing debt is grandfathered and any net-new violation trips the gate
immediately. ``--update`` ratchets a limit down by the violations this branch
fixed relative to its branch point (the merge-base), so the ceilings only ever
fall. Base counts are measured with the *current* checker, so a rule introduced
on this branch is counted at the base too and ratchets like every other one.
Only ever falling is not the same as always falling, so the gate enforces the
second half: a branch that clears violations and leaves the ceiling above its
new count fails, naming the rules and telling the author to run
``make lint-budget-update``. Without that, a removed violation could come back
later under a ceiling nobody lowered. Drift already in the base is never
blamed, so this fires only on the branch that did the clearing.
immediately. ``--update`` ratchets a limit down by the violations fixed relative
to ``--base``, so the ceilings only ever fall. Base counts are measured with the
*current* checker, so a rule introduced on this branch is counted at the base too
and ratchets like every other one. The ratchet runs as a scheduled automation
against litellm_internal_staging, not on PR branches, so concurrent PRs never
race to edit the same limit.
The deliberate difference from its sibling: this gate has no headroom anywhere.
Type discipline seeded LIT010/LIT011 at 1.5x to leave room for an in-flight
@ -144,21 +139,6 @@ def over_ceiling(head: Mapping[str, int], budget: Mapping[str, Mapping[str, int]
)
def unratcheted(
head: Mapping[str, int],
base: Mapping[str, int],
budget: Mapping[str, Mapping[str, int]],
) -> tuple[Breach, ...]:
"""Rules this branch cleared without lowering the ceiling behind them. Requires
both `head < base`, so drift already in the base is never blamed on this change,
and `head < limit`, so a ceiling already at the count is left alone."""
return tuple(sorted(
Breach(rule, head.get(rule, 0), spec["limit"], head.get(rule, 0) - base.get(rule, 0))
for rule, spec in budget.items()
if head.get(rule, 0) < base.get(rule, 0) and head.get(rule, 0) < spec["limit"]
))
def evaluate(
head: Mapping[str, int],
base: Mapping[str, int],
@ -198,38 +178,15 @@ def introduced(
return tuple(v for v in violations if v.line in changed.get(v.file, frozenset()))
def touches_measured_tree(base_point: str) -> bool:
"""Whether this branch changed anything that can move a count. A branch that
touches neither the test tree nor the checker cannot have cleared a violation,
so the base scan is skipped and the gate stays cheap on the common change."""
changed: Final = _run(
["git", "diff", "--name-only", base_point, "--", TARGET, str(CHECKER.relative_to(REPO_ROOT))]
)
return bool(changed.strip())
def cmd_check(base: str) -> None:
budget: Final = json.loads(BUDGET_PATH.read_text())
head: Final = head_violations()
head_counts: Final = count_by_rule(head)
base_point: Final = resolve_base_point(base)
if not over_ceiling(head_counts, budget) and not touches_measured_tree(base_point):
if not over_ceiling(head_counts, budget):
print(f"OK: every TQ rule is within its test-suite ceiling (base {base})")
return
base_point: Final = resolve_base_point(base)
base_at_point: Final = base_counts(base_point)
stale: Final = unratcheted(head_counts, base_at_point, budget)
if stale:
print(f"FAIL: TQ-rule limits were left above the count this branch reached (base {base}):")
for breach in stale:
print(
f" {breach.rule}: this branch cleared {-breach.added} down to {breach.total}, "
f"but the limit is still {breach.cap}"
)
print(
"Run `make lint-budget-update` and commit the lowered limits, so the "
"violations you cleared cannot come back under a ceiling nobody moved."
)
raise SystemExit(1)
breaches: Final = evaluate(head_counts, base_at_point, budget)
if not breaches:
print(f"OK: every TQ rule is within its test-suite ceiling (base {base})")

View file

@ -7,7 +7,7 @@
- {id: reliability.retry.timeout.succeeds_within_retries, module: reliability, tier: P0, behavior: retry, variant: timeout, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:44", rationale: "Timeout retried per policy"}
- {id: reliability.retry.429.succeeds_within_retries, module: reliability, tier: P0, behavior: retry, variant: "429", assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:46", rationale: "429 retried per RateLimitErrorRetries policy"}
- {id: reliability.retry.auth.succeeds_within_retries, module: reliability, tier: P1, behavior: retry, variant: auth, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:42", rationale: "Transient auth glitch retry"}
- {id: reliability.retry.context_window.succeeds_within_retries, module: reliability, tier: P1, behavior: retry, variant: context_window, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:51", rationale: "Multi-attempt on context error"}
- {id: reliability.retry.context_window.succeeds_within_retries, module: reliability, tier: P1, behavior: retry, variant: context_window, assertions: [succeeds_within_retries], exercised_on: [chat_completions, messages], source: "get_retry_from_policy.py:51", fail_before_fix: proven, rationale: "A context-window 400 under BadRequestErrorRetries retries onto a sibling deployment in the same model group, instead of coming straight back as the 400 the deployment that just refused it returned"}
- {id: reliability.cooldown.5xx.trips_then_recovers, module: reliability, tier: P0, behavior: cooldown, variant: "5xx", assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:40", rationale: "Deployment cools after repeated 5xx, recovers after cooldown_time"}
- {id: reliability.cooldown.429.trips_then_recovers, module: reliability, tier: P0, behavior: cooldown, variant: "429", assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:69", rationale: "Cools on 429, avoids hammering exhausted provider"}
- {id: reliability.cooldown.auth.trips_then_recovers, module: reliability, tier: P1, behavior: cooldown, variant: auth, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages], source: "cooldown_handlers.py:74", rationale: "Cools on 401 auth error"}

View file

@ -291,6 +291,7 @@ class RouterSettingsOverride(BaseModel):
context_window_fallbacks: list[dict[str, list[str]]] | None = None
content_policy_fallbacks: list[dict[str, list[str]]] | None = None
num_retries: int | None = None
model_group_retry_policy: dict[str, dict[str, int]] | None = None
enable_tag_filtering: bool | None = None

View file

@ -73,10 +73,25 @@ def create_always_timing_out_deployment(proxy: ProxyClient, name: str) -> str:
)
def create_always_picked_small_context_deployment(proxy: ProxyClient, name: str) -> str:
"""The always-picked half of a retry pair on the smallest-context model OpenAI
still serves: it holds all of the model group's shuffle weight, so an oversized
prompt opens on it and earns a real context-window refusal, which never benches
a deployment, so only the retry itself can steer the request off it."""
return proxy.register_model(
ModelNewBody(
model_name=name,
litellm_params=LiteLLMParamsBody(model=SMALL_CONTEXT_MODEL, api_key=REAL_KEY, weight=1),
model_info=ModelInfoBody(),
)
)
def create_zero_weight_backup_deployment(proxy: ProxyClient, name: str) -> str:
"""The other half of a retry pair: healthy, but weight 0, so the weighted shuffle
never opens on it. It is reachable only once its sibling is benched and the
weighted pick falls through to a uniform one over what is left."""
never opens on it. It is reachable only once its sibling is out of the running,
benched by a cooldown or skipped by the retry, and the weighted pick falls through
to a uniform one over what is left."""
return proxy.register_model(
ModelNewBody(
model_name=name,

View file

@ -1,13 +1,17 @@
"""Live e2e: a request that fails on its first deployment is retried inside its own
model group and still comes back a completion.
The model group is a pair: an always-timing-out deployment that holds all of the
group's shuffle weight, and a healthy backup at weight 0. The weighted pick always
opens on the timing-out one, its first Timeout benches it (an
`allowed_fails_policy` of `TimeoutErrorAllowedFails: 0`), and the retry falls
through to the only deployment left. So the customer sees a completion and the
proxy reports that it took a retry to get there, with no random first pick in the
middle of it.
Each model group is a pair: a deployment that always refuses and holds all of the
group's shuffle weight, plus a healthy backup at weight 0. The weighted pick always
opens on the refusing one, so the customer sees a completion only if the retry
lands on the backup, and the proxy reports that it took a retry to get there, with
no random first pick in the middle of it.
The timeout pair relies on cooldown: the first Timeout benches the timing-out
deployment (an `allowed_fails_policy` of `TimeoutErrorAllowedFails: 0`) and the
retry falls through to the only deployment left. The context-window pair cannot:
a 400 never benches a deployment, so the retry policy's `BadRequestErrorRetries`
has to steer the retry off the deployment that just refused the prompt.
"""
from __future__ import annotations
@ -16,20 +20,48 @@ import pytest
from complexity_router_client import ComplexityRouterClient
from e2e_config import unique_marker
from e2e_http import StreamingResponse
from lifecycle import ResourceManager
from models import RouterSettingsOverride
from reliability_support import (
chat_override,
completion_tokens_of,
content_of,
create_always_picked_small_context_deployment,
create_always_timing_out_deployment,
create_zero_weight_backup_deployment,
finish_reason_of,
oversized_prompt,
)
pytestmark = pytest.mark.e2e
def assert_retry_landed_on_backup(resp: StreamingResponse) -> None:
assert resp.status_code == 200, (
f"the retry should have landed on the healthy backup, got {resp.status_code}: {resp.body[:300]}"
)
attempted = resp.headers.get("x-litellm-attempted-retries")
assert attempted is not None, "response is missing the x-litellm-attempted-retries header"
assert int(attempted) >= 1, (
f"x-litellm-attempted-retries is {attempted!r}; a 200 with no retry means the request never "
"opened on the refusing deployment, so this proves nothing about retries"
)
content = content_of(resp)
finish_reason = finish_reason_of(resp)
completion_tokens = completion_tokens_of(resp) or 0
assert isinstance(content, str), (
f"the retry should have returned a completion body, got content {content!r} (body={resp.body[:300]})"
)
assert content or (finish_reason == "length" and completion_tokens > 0), (
f"the retry returned empty content with finish_reason={finish_reason!r}, "
f"completion_tokens={completion_tokens}; empty content is only acceptable when the budget "
f"was spent on non-visible reasoning (body={resp.body[:300]})"
)
class TestReliabilityRetries:
@pytest.mark.covers("reliability.retry.timeout.succeeds_within_retries")
def test_timeout_on_first_deployment_succeeds_on_retry(
@ -49,25 +81,27 @@ class TestReliabilityRetries:
override=RouterSettingsOverride(num_retries=2),
)
assert resp.status_code == 200, (
f"the retry should have landed on the healthy backup, got {resp.status_code}: {resp.body[:300]}"
assert_retry_landed_on_backup(resp)
@pytest.mark.covers("reliability.retry.context_window.succeeds_within_retries")
def test_context_window_refusal_on_first_deployment_succeeds_on_retry(
self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str
) -> None:
group = f"reliability-retry-{unique_marker()}"
small_context = create_always_picked_small_context_deployment(client.proxy, group)
resources.defer(lambda: client.proxy.delete_model(small_context))
backup = create_zero_weight_backup_deployment(client.proxy, group)
resources.defer(lambda: client.proxy.delete_model(backup))
resp = chat_override(
client.proxy,
scoped_key,
group,
oversized_prompt(unique_marker()),
override=RouterSettingsOverride(
num_retries=2,
model_group_retry_policy={group: {"BadRequestErrorRetries": 2}},
),
)
attempted = resp.headers.get("x-litellm-attempted-retries")
assert attempted is not None, "response is missing the x-litellm-attempted-retries header"
assert int(attempted) >= 1, (
f"x-litellm-attempted-retries is {attempted!r}; a 200 with no retry means the request never "
"opened on the timing-out deployment, so this proves nothing about retries"
)
content = content_of(resp)
finish_reason = finish_reason_of(resp)
completion_tokens = completion_tokens_of(resp) or 0
assert isinstance(content, str), (
f"the retry should have returned a completion body, got content {content!r} (body={resp.body[:300]})"
)
assert content or (finish_reason == "length" and completion_tokens > 0), (
f"the retry returned empty content with finish_reason={finish_reason!r}, "
f"completion_tokens={completion_tokens}; empty content is only acceptable when the budget "
f"was spent on non-visible reasoning (body={resp.body[:300]})"
)
assert_retry_landed_on_backup(resp)

View file

@ -21,11 +21,12 @@ from types import MappingProxyType
import httpx
import pytest
import respx
from openai.types.batch import BatchRequestCounts
import litellm
import litellm.batches.batch_utils as bu
from litellm.types.utils import Usage
from litellm.types.utils import LiteLLMBatch, Usage
# --------------------------------------------------------------------------- #
# Builders for batch OUTPUT file rows.
@ -1718,3 +1719,57 @@ def test_unparsable_bedrock_batch_usage_warns(caplog):
assert usage.total_tokens == 0
assert "does not understand" in caplog.text
assert "inputTextTokenCount" in caplog.text
# --------------------------------------------------------------------------- #
# batch_cost_is_final
# --------------------------------------------------------------------------- #
def _retrieved_batch(
status: str, output_file_id: str | None = None, counts: BatchRequestCounts | None = None
) -> LiteLLMBatch:
return LiteLLMBatch(
id="batch_abc",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status="validating",
output_file_id=output_file_id,
request_counts=counts,
).model_copy(update={"status": status})
class TestBatchCostIsFinal:
"""Every retrieve of one batch writes the same spend row, so the first retrieve
that prices it decides the row for good. A poll before the output exists must
therefore not count as final: pricing it recorded $0 and pinned it (LIT-7048)."""
@pytest.mark.parametrize("status", ["validating", "in_progress", "finalizing", "cancelling"])
def test_in_flight_batch_is_not_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status)) is False
@pytest.mark.parametrize("status", ["completed", "complete"])
def test_completed_with_output_is_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status, output_file_id="file-out")) is True
def test_completed_without_output_and_unknown_counts_is_not_final(self):
assert bu.batch_cost_is_final(_retrieved_batch("completed")) is False
def test_completed_without_output_and_zero_counts_is_not_final(self):
counts = BatchRequestCounts(total=0, completed=0, failed=0)
assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False
def test_completed_without_output_but_successful_lines_is_not_final(self):
counts = BatchRequestCounts(total=2, completed=2, failed=0)
assert bu.batch_cost_is_final(_retrieved_batch("completed", counts=counts)) is False
@pytest.mark.parametrize("status", ["completed", "complete"])
def test_completed_without_output_and_every_line_failed_is_final(self, status):
counts = BatchRequestCounts(total=2, completed=0, failed=2)
assert bu.batch_cost_is_final(_retrieved_batch(status, counts=counts)) is True
@pytest.mark.parametrize("status", ["failed", "expired", "cancelled"])
def test_other_terminal_statuses_are_final(self, status):
assert bu.batch_cost_is_final(_retrieved_batch(status)) is True

View file

@ -2008,7 +2008,14 @@ def test_generic_cost_per_token_azure_gpt56(_local_model_cost_map,
assert round(completion_cost, 10) == round(output_cost * completion_tokens, 10)
@pytest.mark.parametrize("model,zone_multiplier", [("azure/gpt-6-astra", 1.0), ("azure/us/gpt-6-astra", 1.1)])
@pytest.mark.parametrize(
"model,custom_llm_provider,zone_multiplier",
[
("azure/gpt-6-astra", "azure", 1.0),
("azure/us/gpt-6-astra", "azure", 1.1),
("azure_ai/gpt-6-astra", "azure_ai", 1.0),
],
)
@pytest.mark.parametrize(
"prompt_tokens,input_side_multiplier,output_multiplier",
[(100000, 1.0, 1.0), (300000, 2.0, 1.5)],
@ -2016,6 +2023,7 @@ def test_generic_cost_per_token_azure_gpt56(_local_model_cost_map,
def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
_local_model_cost_map,
model,
custom_llm_provider,
zone_multiplier,
prompt_tokens,
input_side_multiplier,
@ -2023,7 +2031,8 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
):
"""Microsoft Foundry sells gpt-6-astra at the OpenAI rates: $10 input, $1 cache read, $12.50 cache write,
$50 output per 1M tokens on Standard Global, with the input side doubling and output 1.5x above 272K
prompt tokens. Standard US Data Zone carries the usual 10% uplift on every rate.
prompt tokens. Standard US Data Zone carries the usual 10% uplift on every rate. A Foundry
deployment reached through the azure_ai route bills the same Standard Global sheet.
"""
cached_tokens = 50000
cache_write_tokens = 40000
@ -2041,7 +2050,7 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
prompt_cost, completion_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="azure",
custom_llm_provider=custom_llm_provider,
)
input_side = zone_multiplier * input_side_multiplier
@ -2051,6 +2060,18 @@ def test_generic_cost_per_token_azure_gpt_6_astra_foundry_price_sheet(
assert completion_cost == pytest.approx(zone_multiplier * output_multiplier * completion_tokens * 5e-5)
def test_generic_cost_per_token_azure_ai_gpt_6_astra_flex_bills_the_standard_rate(_local_model_cost_map):
usage = Usage(prompt_tokens=1000, completion_tokens=100, total_tokens=1100)
standard = generic_cost_per_token(model="azure_ai/gpt-6-astra", usage=usage, custom_llm_provider="azure_ai")
flex = generic_cost_per_token(
model="azure_ai/gpt-6-astra", usage=usage, custom_llm_provider="azure_ai", service_tier="flex"
)
assert flex == standard
assert standard == pytest.approx((1000 * 1e-05, 100 * 5e-05))
@pytest.mark.parametrize(
"model,expected_none,expected_xhigh,expected_minimal",
[

View file

@ -632,6 +632,86 @@ class TestRetrieveBatchCostPassesModelIdentity:
assert captured["model_info"]["input_cost_per_token"] == 0.0
class TestRetrieveBatchPricesOnlyFinalBatches:
"""Regression (LIT-7048): retrieving a provider-id batch priced it on every poll.
Every retrieve of one batch logs under the same spend row, so pricing a poll
that landed before the output existed wrote that row at $0 and pinned it there.
Only a final batch gets priced; an in-flight poll carries no cost at all.
"""
@staticmethod
def _logging_obj() -> LitellmLogging:
obj = LitellmLogging(
model="gpt-5.6-luna",
messages=[{"role": "user", "content": "Hey"}],
stream=False,
call_type="aretrieve_batch",
start_time=time.time(),
litellm_call_id="batch-call-2",
function_id="f",
)
obj.custom_llm_provider = "openai"
return obj
@staticmethod
def _batch(status: str, output_file_id: str | None):
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="batch_6a9c99e185588190877d391f8b9d7f8a",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status="validating",
output_file_id=output_file_id,
).model_copy(update={"status": status})
@pytest.mark.asyncio
@pytest.mark.parametrize(
("status", "output_file_id"),
[("validating", None), ("in_progress", None), ("finalizing", None), ("completed", None), ("complete", None)],
)
async def test_non_final_batch_is_not_priced(self, monkeypatch, status, output_file_id) -> None:
from litellm.litellm_core_utils import litellm_logging as logging_module
handle_completed_batch = AsyncMock()
monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch)
batch = self._batch(status, output_file_id)
await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None)
handle_completed_batch.assert_not_awaited()
assert "response_cost" not in batch._hidden_params
@pytest.mark.asyncio
async def test_completed_batch_with_output_is_priced(self, monkeypatch) -> None:
from litellm.batches.batch_utils import BatchCostUsageResult
from litellm.litellm_core_utils import litellm_logging as logging_module
from litellm.types.utils import Usage
handle_completed_batch = AsyncMock(
return_value=BatchCostUsageResult(
cost=8e-06,
usage=Usage(prompt_tokens=26, completion_tokens=9, total_tokens=35),
models=["gpt-5.6-luna"],
successful_requests=2,
failed_requests=0,
)
)
monkeypatch.setattr(logging_module, "_handle_completed_batch", handle_completed_batch)
batch = self._batch("completed", "file-out")
await self._logging_obj()._async_success_handler_body(result=batch, start_time=None, end_time=None)
handle_completed_batch.assert_awaited_once()
assert batch._hidden_params["response_cost"] == 8e-06
assert batch.usage is not None
assert batch.usage.total_tokens == 35
class TestAnthropicPassthroughCustomPricing:
"""Verify the Anthropic pass-through handler forwards custom pricing."""

View file

@ -82,3 +82,22 @@ class TestTheNormalizedTierIsTheTierSent:
self, local_model_cost_map, model, provider, effort, expected
):
assert _reasoning_effort_sent(model, provider, effort) == expected
@pytest.mark.parametrize(
"model, provider",
[
("gpt-6-astra", "azure_ai"),
("azure_ai/gpt-6-astra", "azure_ai"),
("gpt-6-astra", "azure"),
("us/gpt-6-astra", "azure"),
],
)
def test_an_azure_hosted_astra_deployment_drops_to_the_tier_it_accepts(
self, local_model_cost_map, model, provider
):
"""The deployment answers ``max`` with a 400 naming ``none`` through ``xhigh``, so the rows
say so and the adapter sends the tier below instead of the rejected one."""
assert _reasoning_effort_sent(model, provider, "max") == "xhigh"
def test_the_openai_hosted_twin_still_sends_max(self, local_model_cost_map):
assert _reasoning_effort_sent("gpt-6-astra", "openai", "max") == "max"

View file

@ -3,6 +3,8 @@ from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
from litellm.llms.azure_ai.azure_model_router.transformation import (
AzureModelRouterConfig,
)
@ -138,6 +140,46 @@ def test_azure_ai_validate_environment_with_azure_ad_token():
assert headers["Content-Type"] == "application/json"
@pytest.fixture
def _local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url))
def test_foundry_gpt_6_astra_keeps_sampling_params_when_reasoning_effort_is_none(_local_model_cost_map):
optional_params = AzureAIStudioConfig().map_openai_params(
non_default_params={"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9},
optional_params={},
model="gpt-6-astra",
drop_params=False,
)
assert optional_params == {"reasoning_effort": "none", "temperature": 0.2, "top_p": 0.9}
def test_a_gpt_5_name_without_a_foundry_row_keeps_reading_its_own_entry(
monkeypatch: pytest.MonkeyPatch, _local_model_cost_map
):
"""Most gpt-5-family names have no azure_ai/ row. Reading an azure_ai/ key for those finds
nothing, and an openai.azure.com base sends the name down the azure provider, which has no key
for it either, so every effort answer would silently fall back to false and take temperature,
top_p and logprobs down with it."""
monkeypatch.setenv("AZURE_AI_API_BASE", "https://example-resource.openai.azure.com")
monkeypatch.setenv("AZURE_AI_API_KEY", "placeholder")
optional_params = litellm.utils.get_optional_params(
model="gpt-5.1-chat-latest",
custom_llm_provider="azure_ai",
temperature=0.2,
top_p=0.9,
logprobs=True,
)
assert optional_params["temperature"] == 0.2
assert optional_params["top_p"] == 0.9
assert optional_params["logprobs"] is True
def test_azure_ai_grok_stop_parameter_handling():
"""
Test that Grok models properly handle stop parameter filtering in Azure AI Studio.

View file

@ -0,0 +1,197 @@
import base64
from typing import Final
from unittest.mock import MagicMock
import httpx
import pytest
import litellm
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
from litellm.llms.mistral.audio_speech.transformation import (
MistralTextToSpeechConfig,
MistralTextToSpeechException,
)
from litellm.utils import ProviderConfigManager
SPEECH_URL: Final = "https://api.mistral.ai/v1/audio/speech"
def test_mistral_text_to_speech_config_installed():
config: Final = ProviderConfigManager.get_provider_text_to_speech_config(
model="voxtral-mini-tts-2603",
provider=litellm.LlmProviders.MISTRAL,
)
assert isinstance(config, BaseTextToSpeechConfig)
assert isinstance(config, MistralTextToSpeechConfig)
def test_map_openai_params_drops_speed_and_instructions():
config: Final = MistralTextToSpeechConfig()
voice, params = config.map_openai_params(
model="voxtral-mini-tts-2603",
optional_params={"response_format": "wav", "speed": 1.5, "instructions": "sound cheerful"},
voice="en_paul_neutral",
)
assert voice == "en_paul_neutral"
assert params == {"response_format": "wav"}
def test_map_openai_params_accepts_voice_dict_and_ref_audio():
config: Final = MistralTextToSpeechConfig()
voice, params = config.map_openai_params(
model="voxtral-mini-tts-2603",
optional_params={},
voice={"voice_id": "1f3a8b0c-voice-uuid"},
kwargs={"ref_audio": "bXktdm9pY2Utc2FtcGxl"},
)
assert voice == "1f3a8b0c-voice-uuid"
assert params == {"ref_audio": "bXktdm9pY2Utc2FtcGxl"}
def test_transform_request_builds_mistral_body():
config: Final = MistralTextToSpeechConfig()
data: Final = config.transform_text_to_speech_request(
model="voxtral-mini-tts-2603",
input="hello from litellm",
voice="en_paul_neutral",
optional_params={"response_format": "wav"},
litellm_params={},
headers={},
)
assert data["dict_body"] == {
"model": "voxtral-mini-tts-2603",
"input": "hello from litellm",
"voice_id": "en_paul_neutral",
"response_format": "wav",
}
assert data["headers"] == {"Content-Type": "application/json"}
def test_transform_request_omits_voice_for_ref_audio_cloning():
config: Final = MistralTextToSpeechConfig()
data: Final = config.transform_text_to_speech_request(
model="voxtral-mini-tts-2603",
input="clone me",
voice=None,
optional_params={"ref_audio": "bXktdm9pY2Utc2FtcGxl"},
litellm_params={},
headers={},
)
assert data["dict_body"] == {
"model": "voxtral-mini-tts-2603",
"input": "clone me",
"ref_audio": "bXktdm9pY2Utc2FtcGxl",
}
def test_get_complete_url_default_base():
config: Final = MistralTextToSpeechConfig()
url: Final = config.get_complete_url(model="voxtral-mini-tts-2603", api_base=None, litellm_params={})
assert url == SPEECH_URL
@pytest.mark.parametrize(
"api_base",
["https://custom.api.example.com/v1/", "https://custom.api.example.com/v1", "https://custom.api.example.com"],
)
def test_get_complete_url_custom_base_always_versioned(api_base: str):
config: Final = MistralTextToSpeechConfig()
url: Final = config.get_complete_url(model="voxtral-mini-tts-2603", api_base=api_base, litellm_params={})
assert url == "https://custom.api.example.com/v1/audio/speech"
def test_validate_environment_sets_bearer_header():
config: Final = MistralTextToSpeechConfig()
headers: Final = config.validate_environment(
headers={"x-custom": "1"},
model="voxtral-mini-tts-2603",
api_key="sk-mistral-test",
)
assert headers == {
"x-custom": "1",
"Authorization": "Bearer sk-mistral-test",
"Content-Type": "application/json",
}
def test_validate_environment_requires_key(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("MISTRAL_API_KEY", raising=False)
config: Final = MistralTextToSpeechConfig()
with pytest.raises(MistralTextToSpeechException, match="MISTRAL_API_KEY"):
config.validate_environment(headers={}, model="voxtral-mini-tts-2603")
def test_transform_response_decodes_base64_audio():
config: Final = MistralTextToSpeechConfig()
audio_bytes: Final = b"RIFF-fake-wav-bytes"
raw_response: Final = httpx.Response(
200,
json={"audio_data": base64.b64encode(audio_bytes).decode()},
headers={"x-request-id": "req-123"},
request=httpx.Request(
"POST",
SPEECH_URL,
json={"model": "voxtral-mini-tts-2603", "input": "hi", "response_format": "wav"},
),
)
result: Final = config.transform_text_to_speech_response(
model="voxtral-mini-tts-2603",
raw_response=raw_response,
logging_obj=MagicMock(),
)
assert result.content == audio_bytes
assert result.response.headers["content-type"] == "audio/wav"
assert result.response.headers["content-length"] == str(len(audio_bytes))
assert result.response.headers["x-request-id"] == "req-123"
def test_transform_response_missing_audio_data_raises():
config: Final = MistralTextToSpeechConfig()
raw_response: Final = httpx.Response(
200,
json={"detail": "unexpected"},
request=httpx.Request("POST", SPEECH_URL, json={"model": "voxtral-mini-tts-2603", "input": "hi"}),
)
with pytest.raises(MistralTextToSpeechException, match="audio_data"):
config.transform_text_to_speech_response(
model="voxtral-mini-tts-2603",
raw_response=raw_response,
logging_obj=MagicMock(),
)
def test_map_openai_params_maps_openai_voice_aliases():
config: Final = MistralTextToSpeechConfig()
alloy_voice, _ = config.map_openai_params(
model="voxtral-mini-tts-2603",
optional_params={},
voice="alloy",
)
nova_voice, _ = config.map_openai_params(
model="voxtral-mini-tts-2603",
optional_params={},
voice="Nova",
)
passthrough_voice, _ = config.map_openai_params(
model="voxtral-mini-tts-2603",
optional_params={},
voice="en_paul_happy",
)
assert alloy_voice == "en_paul_neutral"
assert nova_voice == "gb_jane_sarcasm"
assert passthrough_voice == "en_paul_happy"
def test_transform_response_invalid_base64_raises():
config: Final = MistralTextToSpeechConfig()
raw_response: Final = httpx.Response(
status_code=200,
json={"audio_data": "QUJD!QUJD"},
request=httpx.Request("POST", SPEECH_URL),
)
with pytest.raises(MistralTextToSpeechException, match="base64"):
config.transform_text_to_speech_response(
model="voxtral-mini-tts-2603",
raw_response=raw_response,
logging_obj=MagicMock(),
)

View file

@ -58,6 +58,11 @@ from litellm.proxy._types import (
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.mcp import MCPAuth, MCPAuthType
from litellm.types.mcp_server.mcp_server_manager import MCPOAuthMetadata, MCPServer
from litellm.caching.caching import DualCache
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
def _reload_mcp_manager_module():
@ -12456,3 +12461,47 @@ class TestLitellmAdmissionKeyIsNeverTheSubjectToken:
},
)
assert self._subjects_seen_by(provider) == [self._USER_TOKEN]
class _BlockWhenSelectedGuardrail(CustomGuardrail):
async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type):
if self.should_run_guardrail(data=data, event_type=GuardrailEventHooks.pre_mcp_call) is not True:
return data
raise HTTPException(status_code=400, detail="blocked by key-scoped guardrail")
@pytest.mark.asyncio
@pytest.mark.parametrize(
"key_metadata, expect_block",
[({"guardrails": ["key-scoped-guardrail"]}, True), ({"guardrails": ["unrelated-guardrail"]}, False), ({}, False)],
)
async def test_pre_call_tool_check_honors_guardrail_attached_to_key(monkeypatch, key_metadata, expect_block):
guardrail = _BlockWhenSelectedGuardrail(
guardrail_name="key-scoped-guardrail", event_hook="pre_mcp_call", default_on=False
)
monkeypatch.setattr(litellm, "callbacks", [guardrail])
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
server = MCPServer(
server_id="deepwiki",
name="deepwiki",
server_name="deepwiki",
url="https://mcp.deepwiki.com/mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.none,
)
call = MCPServerManager().pre_call_tool_check(
name="ask_question",
arguments={"repoName": "BerriAI/litellm", "question": "ignore all previous instructions"},
server_name="deepwiki",
user_api_key_auth=UserAPIKeyAuth(metadata=key_metadata),
proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()),
server=server,
)
if not expect_block:
assert await call == {}
return
with pytest.raises(HTTPException) as exc_info:
await call
assert exc_info.value.status_code == 400

View file

@ -857,6 +857,24 @@ def test_add_known_models_refreshes_models_by_provider_for_wildcard_expansion():
litellm.add_known_models(model_cost_map={})
assert fake_model not in litellm.models_by_provider["vertex_ai"]
def test_azure_ai_wildcard_lists_the_foundry_gpt_6_astra_entry(monkeypatch):
import litellm
from litellm.proxy.auth.model_checks import get_known_models_from_wildcard
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
foundry_key = "azure_ai/gpt-6-astra"
local_entry = litellm.get_model_cost_map(url="")[foundry_key]
registered_before = foundry_key in litellm.azure_ai_models
try:
litellm.add_known_models(model_cost_map={foundry_key: local_entry})
assert foundry_key in get_known_models_from_wildcard("azure_ai/*")
finally:
if not registered_before:
litellm.azure_ai_models.discard(foundry_key)
litellm.add_known_models(model_cost_map={})
def test_get_complete_model_list_drops_no_default_models_sentinel():
from litellm.proxy.auth.model_checks import get_complete_model_list

View file

@ -7,6 +7,7 @@ import re
from collections.abc import Callable
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, call, patch
import pytest
@ -2936,7 +2937,9 @@ async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkey
"call_type, expects_flush",
[("aresponses", True), ("responses", True), ("acompletion", False)],
)
async def test_insert_spend_log_asks_for_an_immediate_flush_on_responses_calls(call_type: str, expects_flush: bool):
async def test_insert_spend_log_asks_for_an_immediate_flush_on_rows_other_workers_read_back(
call_type: str, expects_flush: bool
):
"""
A `previous_response_id` chained straight off the previous turn reads the DB, so a
Responses row cannot sit in this worker's queue until the monitor's next poll.
@ -2957,6 +2960,303 @@ async def test_insert_spend_log_asks_for_an_immediate_flush_on_responses_calls(c
PrismaClient.spend_log_flush_requested.clear()
def _batch_cost_payload() -> dict:
return {
**_minimal_spend_payload(),
"request_id": "batch_abc_batch_cost",
"call_type": "aretrieve_batch",
"status": "success",
}
def _spend_logs_prisma(inserted: int, existing: object, taken_over: int = 1) -> MagicMock:
prisma = _tool_usage_prisma()
prisma.jsonify_object = lambda data: dict(data)
prisma.db.litellm_spendlogs.create_many = AsyncMock(return_value=inserted)
prisma.db.litellm_spendlogs.find_unique = AsyncMock(return_value=existing)
prisma.db.litellm_spendlogs.update_many = AsyncMock(return_value=taken_over)
return prisma
async def _update_database_with(
db_writer: DBSpendUpdateWriter,
prisma: MagicMock,
payload: dict,
disable_spend_logs: bool = False,
response_cost: float = 0.25,
) -> bool:
with (
patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam
"litellm.proxy.proxy_server.disable_spend_logs", disable_spend_logs
),
patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam
"litellm.proxy.proxy_server.prisma_client", prisma
),
patch( # test-quality-ok: update_database reads this proxy_server global at call time, no seam
"litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"
),
patch( # test-quality-ok: update_database imports the payload builder inside its body, no seam
"litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload",
return_value=payload,
),
):
charged = await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id=None,
team_id=None,
org_id=None,
kwargs={"model": "gpt-5.6-luna", "call_type": "aretrieve_batch"},
completion_response=None,
start_time=datetime.now(timezone.utc),
end_time=datetime.now(timezone.utc),
response_cost=response_cost,
)
await asyncio.sleep(0)
return charged
@pytest.mark.asyncio
@pytest.mark.parametrize(
("inserted", "existing", "charged"),
[
(1, None, True),
(0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25), False),
(0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0), True),
(0, SimpleNamespace(call_type="aretrieve_batch", status="failure", spend=0.0), True),
(0, SimpleNamespace(call_type="aembedding", status="success", spend=0.25), True),
(0, None, True),
],
ids=[
"first_retrieve_owns_the_row",
"another_retrieve_already_charged",
"an_older_proxy_left_a_zero_row_while_the_batch_ran",
"failed_retrieve_holds_the_row",
"client_chosen_call_id_holds_the_row",
"row_gone_between_insert_and_lookup",
],
)
async def test_update_database_charges_a_batch_only_from_the_retrieve_that_wrote_its_row(
inserted: int, existing: object, charged: bool
):
"""
Every retrieve of one batch shares one spend row, so the insert that lands first is
the charge and every later retrieve must leave the counters alone (LIT-7048). A row
that recorded no charge must not be able to take the charge away: neither one a
client planted under the batch id, nor the $0 row a pre-upgrade proxy wrote every
time it polled the batch while it was still running.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(inserted, existing)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is charged
claimed_rows = prisma.db.litellm_spendlogs.create_many.await_args.kwargs
assert claimed_rows["skip_duplicates"] is True
assert [(row["request_id"], row["spend"]) for row in claimed_rows["data"]] == [("batch_abc_batch_cost", 0.25)]
assert prisma.spend_log_transactions == []
assert db_writer._batch_database_updates.await_count == (1 if charged else 0)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("taken_over", "charged"),
[(1, True), (0, False)],
ids=["this_retrieve_takes_it_over", "another_one_got_there_first"],
)
async def test_update_database_charges_a_batch_whose_row_a_pre_upgrade_poll_left_at_zero(
taken_over: int, charged: bool
):
"""
A proxy without this fix wrote the batch's row at $0 on every poll of a running batch,
and the row outlives the upgrade, so the charge has to land on the row itself. Charging
without writing it there would charge again on every later retrieve (LIT-7048).
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0)
prisma = _spend_logs_prisma(0, existing, taken_over)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is charged
taken = prisma.db.litellm_spendlogs.update_many.await_args.kwargs
assert taken["where"] == {
"request_id": "batch_abc_batch_cost",
"call_type": "aretrieve_batch",
"status": "success",
"spend": 0.0,
}
assert taken["data"]["spend"] == 0.25
assert "request_id" not in taken["data"]
assert db_writer._batch_database_updates.await_count == (1 if charged else 0)
@pytest.mark.asyncio
async def test_update_database_leaves_a_batch_whose_zero_row_it_could_not_take_over_to_the_next_retrieve():
"""
A DB that refuses the takeover leaves the row reading $0, so charging here would charge
the batch again on every later retrieve. The retrieve that does take the row over is the
one that charges.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0)
prisma = _spend_logs_prisma(0, existing)
prisma.db.litellm_spendlogs.update_many = AsyncMock(side_effect=RuntimeError("db unreachable"))
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is False
assert db_writer._batch_database_updates.await_count == 0
@pytest.mark.asyncio
async def test_update_database_leaves_a_batch_that_cost_nothing_to_the_retrieve_that_wrote_its_row():
"""
A batch every line of which failed costs $0, so its row reads $0 for the honest reason
and the retrieve that wrote it is still the one that accounted it. Taking that row over
on every later retrieve would count one batch as many requests.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
existing = SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.0)
prisma = _spend_logs_prisma(0, existing)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload(), response_cost=0.0) is False
prisma.db.litellm_spendlogs.update_many.assert_not_called()
assert db_writer._batch_database_updates.await_count == 0
@pytest.mark.asyncio
@pytest.mark.parametrize(
("inserted", "existing", "charged"),
[
(1, None, True),
(0, SimpleNamespace(call_type="aretrieve_batch", status="success", spend=0.25), False),
],
ids=["first_retrieve_owns_the_row", "another_retrieve_already_charged"],
)
async def test_update_database_charges_a_batch_once_even_with_spend_logs_disabled(
inserted: int, existing: object, charged: bool
):
"""
disable_spend_logs drops the per-request logs, not the batch's charge, so the one row
that makes a batch chargeable exactly once is still written and still read back.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(inserted, existing)
assert await _update_database_with(db_writer, prisma, _batch_cost_payload(), True) is charged
assert prisma.db.litellm_spendlogs.create_many.await_count == 1
assert db_writer._batch_database_updates.await_count == (1 if charged else 0)
@pytest.mark.asyncio
async def test_update_database_writes_no_ordinary_spend_row_with_spend_logs_disabled():
"""The batch carve-out above stays a carve-out: every other row still goes unwritten."""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(1, None)
payload = {**_batch_cost_payload(), "call_type": "acompletion"}
assert await _update_database_with(db_writer, prisma, payload, True) is True
prisma.db.litellm_spendlogs.create_many.assert_not_called()
assert prisma.spend_log_transactions == []
assert db_writer._batch_database_updates.await_count == 1
_BATCH_CLAIM_FIELDS = {"request_id", "call_type", "status", "spend", "startTime", "endTime"}
def _logged_batch_cost_payload() -> dict:
return {
**_batch_cost_payload(),
"api_key": "0e5b0e9e5f",
"model": "gpt-5.6-luna",
"user": "test-user",
"metadata": '{"batch_models": ["gpt-5.6-luna"]}',
"requester_ip_address": "127.0.0.1",
"proxy_server_request": '{"headers": {"user-agent": "litellm-batch-cost-check"}}',
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
("disable_spend_logs", "logs_the_request"),
[(False, True), (True, False)],
ids=["spend_logs_on", "spend_logs_off"],
)
async def test_update_database_claims_a_batch_without_logging_the_request_that_polled_it(
disable_spend_logs: bool, logs_the_request: bool
):
"""
disable_spend_logs has to keep meaning that no request gets logged, and the batch's cost
row is the one row it cannot drop, so with logging off that row carries only what tells
the retrieves apart: no metadata, no requester IP, no key, model, or token counts.
"""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(1, None)
payload = _logged_batch_cost_payload()
assert await _update_database_with(db_writer, prisma, payload, disable_spend_logs) is True
claimed = prisma.db.litellm_spendlogs.create_many.await_args.kwargs["data"][0]
assert set(claimed) == (set(payload) if logs_the_request else _BATCH_CLAIM_FIELDS)
assert claimed["spend"] == 0.25
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
async def test_update_database_queues_only_the_claim_for_a_batch_it_could_not_write_with_logs_disabled():
"""A refused claim is retried through the queue, so what it queues has to stay unlogged too."""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(0, None)
prisma.db.litellm_spendlogs.create_many = AsyncMock(side_effect=RuntimeError("db unreachable"))
assert await _update_database_with(db_writer, prisma, _logged_batch_cost_payload(), True) is True
assert [set(row) for row in prisma.spend_log_transactions] == [_BATCH_CLAIM_FIELDS]
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
async def test_update_database_queues_a_batch_cost_row_it_could_not_claim():
"""An unreachable DB must not drop the batch's only spend row, nor its charge."""
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(0, None)
prisma.db.litellm_spendlogs.create_many = AsyncMock(side_effect=RuntimeError("db unreachable"))
assert await _update_database_with(db_writer, prisma, _batch_cost_payload()) is True
assert [row["request_id"] for row in prisma.spend_log_transactions] == ["batch_abc_batch_cost"]
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize(
"payload",
[{**_batch_cost_payload(), "call_type": "acompletion"}, {**_batch_cost_payload(), "status": "failure"}],
ids=["not_a_batch_retrieve", "failed_batch_retrieve"],
)
async def test_update_database_queues_every_other_spend_row_for_the_next_flush(payload: dict):
db_writer = DBSpendUpdateWriter()
db_writer._batch_database_updates = AsyncMock()
prisma = _spend_logs_prisma(1, None)
assert await _update_database_with(db_writer, prisma, payload) is True
prisma.db.litellm_spendlogs.create_many.assert_not_called()
assert prisma.spend_log_transactions == [payload]
assert db_writer._batch_database_updates.await_count == 1
@pytest.mark.asyncio
@pytest.mark.parametrize(
"injected_deployment, attributed",

View file

@ -1,4 +1,4 @@
import asyncio
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
@ -70,9 +70,7 @@ async def test_async_post_call_failure_hook():
# Check that metadata was properly updated
assert "litellm_params" in call_args["kwargs"]
assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == {
"request_id": "test_request_id"
}
assert call_args["kwargs"]["litellm_params"]["proxy_server_request"] == {"request_id": "test_request_id"}
metadata = call_args["kwargs"]["litellm_params"]["metadata"]
assert metadata["user_api_key"] == "test_api_key"
assert metadata["status"] == "failure"
@ -336,9 +334,7 @@ async def test_should_continue_failure_tracking_when_budget_release_fails():
)
assert mock_invalidate_budget_reservation_counters.await_count == 1
assert (
mock_invalidate_budget_reservation_counters.await_args.kwargs[
"budget_reservation"
]
mock_invalidate_budget_reservation_counters.await_args.kwargs["budget_reservation"]
is user_api_key_dict.budget_reservation
)
assert user_api_key_dict.budget_reservation["finalized"] is True
@ -433,36 +429,21 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object():
"entries": [{"counter_key": "spend:key:test_api_key"}],
}
assert _get_budget_reservation_from_metadata(metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}) is None
assert (
_get_budget_reservation_from_metadata(
metadata={"user_api_key_auth": dict(UserAPIKeyAuth())}
)
is None
)
assert (
_get_budget_reservation_from_metadata(
metadata={
"user_api_key_auth": UserAPIKeyAuth(
budget_reservation=budget_reservation
)
}
metadata={"user_api_key_auth": UserAPIKeyAuth(budget_reservation=budget_reservation)}
)
== budget_reservation
)
assert (
_get_budget_reservation_from_metadata(
metadata={
"user_api_key_auth": dict(
UserAPIKeyAuth(budget_reservation=budget_reservation)
)
}
metadata={"user_api_key_auth": dict(UserAPIKeyAuth(budget_reservation=budget_reservation))}
)
== budget_reservation
)
assert (
_get_budget_reservation_from_metadata(
metadata={"user_api_key_budget_reservation": budget_reservation}
)
_get_budget_reservation_from_metadata(metadata={"user_api_key_budget_reservation": budget_reservation})
is budget_reservation
)
@ -470,9 +451,7 @@ def test_get_budget_reservation_from_metadata_handles_dict_auth_object():
@pytest.mark.asyncio
async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails():
proxy_logging_obj = MagicMock()
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
side_effect=Exception("db unavailable")
)
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=Exception("db unavailable"))
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
@ -508,9 +487,7 @@ async def test_update_database_and_spend_counters_releases_reservation_when_db_u
async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails():
proxy_logging_obj = MagicMock()
db_exception = RuntimeError("db unavailable")
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(
side_effect=db_exception
)
proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock(side_effect=db_exception)
increment_spend_counters = AsyncMock()
budget_reservation = {"reserved_cost": 0.5, "entries": []}
@ -554,12 +531,8 @@ async def test_update_database_and_spend_counters_preserves_db_exception_when_re
budget_reservation=budget_reservation,
)
assert mock_log_exception.call_count == 2
mock_log_exception.assert_any_call(
"Failed to release budget reservation after database update failed"
)
mock_log_exception.assert_any_call(
"Failed to invalidate budget reservation counters after release failed"
)
mock_log_exception.assert_any_call("Failed to release budget reservation after database update failed")
mock_log_exception.assert_any_call("Failed to invalidate budget reservation counters after release failed")
increment_spend_counters.assert_not_awaited()
@ -778,6 +751,107 @@ async def test_track_cost_callback_defers_in_progress_background_interaction():
mock_proxy_logging.failed_tracking_alert.assert_not_called()
def _batch_retrieve_kwargs(call_type: str, reservation: dict | None = None) -> dict:
metadata = {
"user_api_key": "hashed_key",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
**({"user_api_key_budget_reservation": reservation} if reservation is not None else {}),
}
return {
"call_type": call_type,
"model": "gpt-5.6-luna",
"litellm_call_id": "test-call-id",
"litellm_params": {"metadata": metadata},
"standard_logging_object": {"response_cost": 0.0, "request_tags": None},
"stream": False,
}
def _retrieved_batch(status: str, output_file_id: str | None):
from litellm.types.utils import LiteLLMBatch
return LiteLLMBatch(
id="batch_abc",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="file-in",
object="batch",
status=status,
output_file_id=output_file_id,
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("call_type", "status", "output_file_id", "row_claimed", "spend_written", "charged"),
[
("aretrieve_batch", "in_progress", None, True, False, False),
("aretrieve_batch", "completed", None, True, False, False),
("aretrieve_batch", "completed", "file-out", False, True, False),
("aretrieve_batch", "completed", "file-out", True, True, True),
("aretrieve_batch", "failed", None, True, True, True),
("acreate_batch", "validating", None, True, True, True),
],
ids=[
"retrieve_before_final",
"retrieve_completed_without_output_yet",
"retrieve_after_another_retrieve_charged",
"retrieve_first_final",
"retrieve_failed_batch",
"create_before_final",
],
)
async def test_track_cost_callback_charges_a_batch_once_and_only_when_final( # test-quality-ok: whether the spend writer runs, whether the counters move, and whether the poll's reservation is handed back is the whole observable contract of the gate
call_type, status, output_file_id, row_claimed, spend_written, charged
):
"""
A poll before the batch is final used to pin its shared spend row at $0, and every
completed retrieve after the first charged the key again (LIT-7048). Only retrieves
are gated, since creating a batch is its own billable request, and a retrieve that
charges nothing hands its budget reservation back instead.
"""
logger = _ProxyDBLogger()
budget_reservation = None if charged else {"reserved_cost": 0.5, "entries": []}
kwargs = _batch_retrieve_kwargs(call_type, reservation=budget_reservation)
with (
patch( # test-quality-ok: increment_spend_counters is a proxy_server global the callback reads lazily, no seam
"litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock
) as mock_increment_spend_counters,
patch( # test-quality-ok: update_cache is a proxy_server global the callback reads lazily, no seam
"litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock
) as mock_update_cache,
patch( # test-quality-ok: callback imports proxy_logging_obj off proxy_server in its body, no seam
"litellm.proxy.proxy_server.proxy_logging_obj"
) as mock_proxy_logging,
patch( # test-quality-ok: the release is imported inside the callback's helper, no seam
"litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", new_callable=AsyncMock
) as mock_release_budget_reservation,
):
mock_proxy_logging.failed_tracking_alert = AsyncMock()
mock_proxy_logging.db_spend_update_writer.update_database = AsyncMock(return_value=row_claimed)
mock_proxy_logging.slack_alerting_instance.customer_spend_alert = AsyncMock()
await logger._PROXY_track_cost_callback(
kwargs=kwargs,
completion_response=_retrieved_batch(status, output_file_id),
start_time=datetime.now(),
end_time=datetime.now(),
)
await asyncio.sleep(0)
mock_proxy_logging.failed_tracking_alert.assert_not_called()
assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if spend_written else 0)
assert mock_increment_spend_counters.await_count == (1 if charged else 0)
assert mock_update_cache.await_count == (1 if charged else 0)
if charged:
mock_release_budget_reservation.assert_not_awaited()
else:
mock_release_budget_reservation.assert_awaited_once_with(budget_reservation=budget_reservation)
def _in_progress_interaction_kwargs(reservation: dict) -> dict:
return {
"call_type": "acreate_interaction",
@ -1101,10 +1175,7 @@ async def test_async_post_call_failure_hook_propagates_trace_id_from_logging_obj
# standard_logging_object should have been propagated from logging obj
assert call_kwargs.get("standard_logging_object") is not None
assert (
call_kwargs["standard_logging_object"]["trace_id"]
== "trace-id-from-logging-obj"
)
assert call_kwargs["standard_logging_object"]["trace_id"] == "trace-id-from-logging-obj"
# litellm_trace_id should also be propagated as a fallback
assert call_kwargs.get("litellm_trace_id") == "trace-id-from-logging-obj"
@ -1691,9 +1762,7 @@ async def test_async_post_call_failure_hook_records_recovered_partial_spend():
"metadata": {},
"proxy_server_request": {"request_id": "rid"},
"response_cost": 3.5e-05,
"combined_usage_object": Usage(
prompt_tokens=30, completion_tokens=1, total_tokens=31
),
"combined_usage_object": Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31),
}
with patch(
@ -1772,15 +1841,10 @@ async def test_track_cost_callback_enriches_user_id_for_mcp_style_metadata():
assert mock_increment.call_args.kwargs["team_id"] == "team-123"
assert mock_increment.call_args.kwargs["org_id"] == "org-456"
update_kwargs = (
mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs
)
update_kwargs = mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs
assert update_kwargs["user_id"] == "mcp-user@example.com"
assert update_kwargs["team_id"] == "team-123"
assert (
kwargs["litellm_params"]["metadata"]["user_api_key_user_id"]
== "mcp-user@example.com"
)
assert kwargs["litellm_params"]["metadata"]["user_api_key_user_id"] == "mcp-user@example.com"
@pytest.mark.asyncio
@ -1875,9 +1939,7 @@ def test_should_track_cost_callback_pass_through_without_owner(call_type, expect
],
)
@pytest.mark.asyncio
async def test_track_cost_callback_logs_unauthenticated_pass_through_request(
call_type, expect_spend_log
):
async def test_track_cost_callback_logs_unauthenticated_pass_through_request(call_type, expect_spend_log):
"""Regression for LIT-3782: a pass-through request with auth=false reaches the
cost callback with no key/user/team/end-user. Before the fix the spend-log
write was skipped and the request never appeared in request/usage logs. It
@ -1923,9 +1985,7 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request(
end_time=datetime.now(),
)
assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (
1 if expect_spend_log else 0
)
assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == (1 if expect_spend_log else 0)
class _FakeDeploymentLookup:

View file

@ -1922,6 +1922,34 @@ async def test_proxy_only_error_5xx_keeps_traceback_and_runs_sync_callbacks(monk
assert "test_proxy_utils" in captured["async_traceback"]
@pytest.mark.parametrize(
"key_metadata, team_metadata, expected_to_run",
[
({"guardrails": ["key-scoped-guardrail"]}, None, True),
({}, {"guardrails": ["key-scoped-guardrail"]}, True),
({"guardrails": ["some-other-guardrail"]}, None, False),
({}, None, False),
],
)
def test_convert_mcp_to_llm_format_carries_key_and_team_guardrails(key_metadata, team_metadata, expected_to_run):
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
guardrail = CustomGuardrail(guardrail_name="key-scoped-guardrail", event_hook="pre_mcp_call", default_on=False)
kwargs = {
"name": "ask_question",
"arguments": {"question": "hello"},
"server_name": "deepwiki",
"user_api_key_auth": UserAPIKeyAuth(metadata=key_metadata, team_metadata=team_metadata),
}
request_obj = proxy_logging._create_mcp_request_object_from_kwargs(kwargs)
with patch( # test-quality-ok: the key-guardrail premium gate reads this proxy_server module global and has no injection seam
"litellm.proxy.proxy_server.premium_user", True
):
synthetic = proxy_logging._convert_mcp_to_llm_format(request_obj, kwargs)
assert guardrail.should_run_guardrail(synthetic, GuardrailEventHooks.pre_mcp_call) is expected_to_run
class _TracebackRecordingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()

View file

@ -389,14 +389,24 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
"max",
)
@pytest.mark.parametrize("model", ["azure/gpt-6-astra", "azure/us/gpt-6-astra"])
def test_a_foundry_deployment_also_advertises_none(self, local_model_cost_map, model):
"""Microsoft Foundry serves the same model but its API accepts reasoning_effort none
(verified live: 200 with zero reasoning tokens, and it unlocks temperature), which
OpenAI's rejects, so an Azure deployment offers none on top of low through max."""
@pytest.mark.parametrize(
"model,custom_llm_provider",
[
("azure/gpt-6-astra", "azure"),
("azure/us/gpt-6-astra", "azure"),
("azure_ai/gpt-6-astra", "azure_ai"),
],
)
def test_an_azure_hosted_deployment_advertises_none_but_not_max(
self, local_model_cost_map, model, custom_llm_provider
):
"""Microsoft hosts the same model with a different level set than OpenAI does. Verified live
on both Azure routes: none returns 200 with zero reasoning tokens and unlocks temperature,
which OpenAI's API rejects, while max returns 400 unsupported_value naming none through
xhigh as the levels it does take."""
from litellm.utils import _get_model_info_helper
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider="azure"))
model_info = dict(_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider))
assert resolve_supported_reasoning_efforts(model_info, deployment_is_mapped=True) == (
"none",
@ -404,5 +414,4 @@ class TestGpt6AstraAdvertisesItsDocumentedLevels:
"medium",
"high",
"xhigh",
"max",
)

View file

@ -4526,6 +4526,18 @@ def test_explicit_pricing_precedes_private_provider_response_model(
assert selected == expected
def test_cost_per_token_mistral_voxtral_tts_bills_per_input_character(_local_model_cost_map):
prompt_usd, completion_usd = cost_per_token(
model="voxtral-mini-tts-2603",
custom_llm_provider="mistral",
call_type="speech",
prompt_characters=1000,
)
assert prompt_usd == pytest.approx(1000 * 1.6e-05)
assert completion_usd == 0.0
def test_batch_cost_calculator_gpt_6_astra_bills_half_the_standard_rate(_local_model_cost_map):
"""gpt-6-astra batch pricing is 50% off the standard $10 input and $50 output rates per 1M tokens."""
from litellm.cost_calculator import batch_cost_calculator

View file

@ -3351,6 +3351,52 @@ def test_stream_chunk_builder_leaves_xai_reported_cost_to_the_calculator(monkeyp
assert logging_obj._response_cost_calculator(result=response) == pytest.approx(0.63)
def test_speech_mistral_dispatches_and_decodes_audio(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
audio_bytes: Final = b"ID3-fake-mp3-bytes"
mock_route: Final = respx_mock.post("https://api.mistral.ai/v1/audio/speech").mock(
return_value=httpx.Response(200, json={"audio_data": base64.b64encode(audio_bytes).decode()})
)
response: Final = litellm.speech(
model="mistral/voxtral-mini-tts-2603",
input="hello from litellm",
voice="en_paul_neutral",
response_format="wav",
speed=2,
instructions="sound cheerful",
)
assert mock_route.called
request_body: Final = json.loads(mock_route.calls.last.request.content)
assert request_body == {
"model": "voxtral-mini-tts-2603",
"input": "hello from litellm",
"voice_id": "en_paul_neutral",
"response_format": "wav",
}
assert mock_route.calls.last.request.headers["authorization"] == "Bearer sk-mistral-test"
assert response.content == audio_bytes
def test_speech_mistral_routes_to_configured_api_base(respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
audio_bytes: Final = b"ID3-gateway-bytes"
gateway_route: Final = respx_mock.post("https://mistral.gateway.internal/v1/audio/speech").mock(
return_value=httpx.Response(200, json={"audio_data": base64.b64encode(audio_bytes).decode()})
)
response: Final = litellm.speech(
model="mistral/voxtral-mini-tts-2603",
input="hello from litellm",
voice="en_paul_neutral",
api_base="https://mistral.gateway.internal",
)
assert gateway_route.called
assert response.content == audio_bytes
FOUNDRY_HOST: Final = "https://my-project.services.ai.azure.com"

View file

@ -39,6 +39,7 @@ from litellm.router import (
_anthropic_stream_should_drop_pre_content_ping,
_is_retriable_anthropic_status,
)
from litellm.router_strategy import simple_shuffle
from litellm.types.router import DeploymentTypedDict
@ -12805,6 +12806,82 @@ class TestTierParamsTheTargetAccepts:
assert accepted == {"reasoning_effort": "max"}
@pytest.mark.asyncio
async def test_router_aspeech_without_voice_dispatches_ref_audio_cloning(respx_mock, monkeypatch):
import base64
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
audio_bytes = b"RIFFfake-wav-bytes"
respx_mock.post("https://api.mistral.ai/v1/audio/speech").respond(
json={"audio_data": base64.b64encode(audio_bytes).decode()}
)
router = Router(
model_list=[
{
"model_name": "voxtral-tts",
"litellm_params": {"model": "mistral/voxtral-mini-tts-2603"},
}
]
)
response = await router.aspeech(model="voxtral-tts", input="clone me", ref_audio="ZmFrZQ==")
request_body = json.loads(respx_mock.calls.last.request.content)
assert request_body == {"model": "voxtral-mini-tts-2603", "input": "clone me", "ref_audio": "ZmFrZQ=="}
assert response.content == audio_bytes
@pytest.mark.asyncio
async def test_router_aspeech_without_voice_keeps_deployment_default_voice(respx_mock, monkeypatch):
import base64
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
audio_bytes = b"RIFFfake-wav-bytes"
respx_mock.post("https://api.mistral.ai/v1/audio/speech").respond(
json={"audio_data": base64.b64encode(audio_bytes).decode()}
)
router = Router(
model_list=[
{
"model_name": "voxtral-tts",
"litellm_params": {"model": "mistral/voxtral-mini-tts-2603", "voice": "en_paul_neutral"},
}
]
)
await router.aspeech(model="voxtral-tts", input="use my default")
request_body = json.loads(respx_mock.calls.last.request.content)
assert request_body["voice_id"] == "en_paul_neutral"
@pytest.mark.asyncio
async def test_router_aspeech_request_voice_overrides_deployment_default(respx_mock, monkeypatch):
import base64
monkeypatch.setenv("MISTRAL_API_KEY", "sk-mistral-test")
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
audio_bytes = b"RIFFfake-wav-bytes"
respx_mock.post("https://api.mistral.ai/v1/audio/speech").respond(
json={"audio_data": base64.b64encode(audio_bytes).decode()}
)
router = Router(
model_list=[
{
"model_name": "voxtral-tts",
"litellm_params": {"model": "mistral/voxtral-mini-tts-2603", "voice": "en_paul_neutral"},
}
]
)
await router.aspeech(model="voxtral-tts", input="override me", voice="gb_oliver_neutral")
request_body = json.loads(respx_mock.calls.last.request.content)
assert request_body["voice_id"] == "gb_oliver_neutral"
class TestRequestReasoningEffortOverride:
def test_drop_effort_from_nested_carrier_preserves_other_nested_values(self):
params: dict[str, object] = {"output_config": {"effort": "high", "format": "json"}}
@ -13116,6 +13193,7 @@ async def test_prompt_management_factory_marks_injection_for_every_deployment(mo
({"DefaultRetries": 0}, 502, litellm.BadGatewayError, 1),
({"DefaultRetries": 0, "ServiceUnavailableErrorRetries": 1}, 503, litellm.ServiceUnavailableError, 2),
({"ServiceUnavailableErrorRetries": 0}, 502, litellm.BadGatewayError, 3),
({"BadRequestErrorRetries": 2}, 400, litellm.BadRequestError, 3),
],
)
async def test_router_retry_policy_controls_upstream_attempt_count(
@ -13152,6 +13230,323 @@ async def test_router_retry_policy_controls_upstream_attempt_count(
assert upstream.call_count == expected_upstream_calls
@pytest.mark.asyncio
@pytest.mark.parametrize(
"retry_policy,upstream_error",
[
(
{"BadRequestErrorRetries": 2},
{
"message": "This model's maximum context length is 16385 tokens",
"type": "invalid_request_error",
"code": "context_length_exceeded",
},
),
(
{"ContentPolicyViolationErrorRetries": 2},
{
"message": "Your request was rejected as a result of our safety system",
"type": "invalid_request_error",
"code": "content_policy_violation",
},
),
],
)
async def test_router_retry_policy_400_retries_on_sibling_deployment(
monkeypatch: pytest.MonkeyPatch, retry_policy, upstream_error
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
router = litellm.Router(
model_list=[
{
"model_name": "gpt-5.6",
"litellm_params": {
"model": "openai/gpt-5.6",
"api_key": "sk-fake",
"api_base": "https://rejecting.local/v1",
"weight": 1,
},
"model_info": {"id": "rejecting"},
},
{
"model_name": "gpt-5.6",
"litellm_params": {
"model": "openai/gpt-5.6",
"api_key": "sk-fake",
"api_base": "https://accepting.local/v1",
"weight": 0,
},
"model_info": {"id": "accepting"},
},
],
num_retries=2,
retry_policy=retry_policy,
disable_cooldowns=True,
)
with respx.mock(assert_all_called=True) as respx_mock:
rejecting = respx_mock.post("https://rejecting.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": upstream_error})
)
accepting = respx_mock.post("https://accepting.local/v1/chat/completions").mock(
return_value=httpx.Response(
200,
json={
"id": "chatcmpl-lit-7036",
"object": "chat.completion",
"created": 1,
"model": "gpt-5.6",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "hi back"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3},
},
)
)
response = await router.acompletion(model="gpt-5.6", messages=[{"role": "user", "content": "hi"}])
assert rejecting.call_count == 1
assert accepting.call_count == 1
assert response.choices[0].message.content == "hi back"
assert response._hidden_params["additional_headers"]["x-litellm-attempted-retries"] == 1
_UPSTREAM_400 = {"message": "upstream refused this request", "type": "invalid_request_error", "code": "bad_request"}
def _retry_skip_deployment(deployment_id, host, litellm_params=None, model_info=None):
return {
"model_name": "gpt-5.6",
"litellm_params": {
"model": "openai/gpt-5.6",
"api_key": "sk-fake",
"api_base": f"https://{host}.local/v1",
**(litellm_params or {}),
},
"model_info": {"id": deployment_id, **(model_info or {})},
}
@pytest.mark.parametrize(
"status_code,failed_deployment_id,already_skipped,expected",
[
(400, "rejecting", None, ("rejecting",)),
(403, "rejecting", None, ("rejecting",)),
(400, "second", ("first",), ("first", "second")),
(400, "first", ("first",), ("first",)),
(429, "rejecting", None, ()),
(503, "rejecting", None, ()),
(408, "rejecting", None, ()),
(400, None, None, ()),
(None, "rejecting", None, ()),
("400", "rejecting", None, ()),
(400, "second", 7, ("second",)),
(400, "second", "first", ("second",)),
(400, "second", ["first"], ("second",)),
(400, "second", ("first", 7), ("first", "second")),
],
)
def test_router_deployment_ids_to_skip_on_retry(status_code, failed_deployment_id, already_skipped, expected):
exception = Exception("upstream refused this request")
exception.status_code = status_code
exception.failed_deployment_id = failed_deployment_id
assert litellm.Router._deployment_ids_to_skip_on_retry(exception, already_skipped) == expected
@pytest.mark.parametrize(
"value,expected",
[
(("first", "second"), ("first", "second")),
((), ()),
(("first", 7, None, "second"), ("first", "second")),
(None, ()),
(7, ()),
("first", ()),
(["first"], ()),
({"first": True}, ()),
(object(), ()),
],
)
def test_router_as_retry_skipped_deployment_ids_keeps_only_a_tuple_of_strings(value, expected):
from litellm.router import _as_retry_skipped_deployment_ids
assert _as_retry_skipped_deployment_ids(value) == expected
@pytest.mark.parametrize(
"deployment_ids,skipped,expected",
[
(["rejecting", "sibling"], ("rejecting",), ["sibling"]),
(["rejecting"], ("rejecting",), ["rejecting"]),
(["rejecting", "sibling"], ("rejecting", "sibling"), ["rejecting", "sibling"]),
(["rejecting", "sibling"], (), ["rejecting", "sibling"]),
(["rejecting", "sibling"], None, ["rejecting", "sibling"]),
(["rejecting", "sibling"], ("absent",), ["rejecting", "sibling"]),
(["rejecting", "sibling"], 7, ["rejecting", "sibling"]),
(["rejecting", "sibling"], "rejecting", ["rejecting", "sibling"]),
(["rejecting", "sibling"], ["rejecting"], ["rejecting", "sibling"]),
(["rejecting", "sibling"], {"rejecting": True}, ["rejecting", "sibling"]),
(["rejecting", "sibling"], ("rejecting", 7), ["sibling"]),
],
)
@pytest.mark.asyncio
async def test_router_healthy_deployments_keep_the_last_candidate_a_retry_skipped(deployment_ids, skipped, expected):
router = litellm.Router(
model_list=[_retry_skip_deployment(deployment_id, deployment_id) for deployment_id in deployment_ids],
disable_cooldowns=True,
)
request_kwargs = {"_retry_skipped_deployment_ids": skipped}
healthy_deployments = await router.async_get_healthy_deployments(model="gpt-5.6", request_kwargs=request_kwargs)
assert sorted(deployment["model_info"]["id"] for deployment in healthy_deployments) == sorted(expected)
assert "_retry_skipped_deployment_ids" not in request_kwargs
@pytest.mark.parametrize("client_supplied", [7, "rejecting", ["rejecting"], {"rejecting": True}, object()])
@pytest.mark.asyncio
async def test_router_retry_policy_400_keeps_upstream_error_when_a_client_forges_the_skip_list(
monkeypatch: pytest.MonkeyPatch, client_supplied
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
router = litellm.Router(
model_list=[_retry_skip_deployment("rejecting", "rejecting"), _retry_skip_deployment("sibling", "sibling")],
num_retries=2,
retry_policy={"BadRequestErrorRetries": 2},
disable_cooldowns=True,
)
with respx.mock as respx_mock:
respx_mock.post("https://rejecting.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
respx_mock.post("https://sibling.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
with pytest.raises(litellm.BadRequestError) as raised:
await router.acompletion(
model="gpt-5.6",
messages=[{"role": "user", "content": "hi"}],
_retry_skipped_deployment_ids=client_supplied,
)
assert "upstream refused this request" in str(raised.value)
@pytest.mark.asyncio
async def test_router_retry_policy_400_keeps_upstream_error_on_order_fallback_hop(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
router = litellm.Router(
model_list=[
_retry_skip_deployment("order1", "order1", litellm_params={"order": 1}),
_retry_skip_deployment("order2", "order2", litellm_params={"order": 2}),
],
num_retries=2,
retry_policy={"BadRequestErrorRetries": 2},
disable_cooldowns=True,
)
with respx.mock as respx_mock:
order1 = respx_mock.post("https://order1.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
order2 = respx_mock.post("https://order2.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
with pytest.raises(litellm.BadRequestError) as raised:
await router.acompletion(model="gpt-5.6", messages=[{"role": "user", "content": "hi"}])
assert "upstream refused this request" in str(raised.value)
assert "No deployments available" not in str(raised.value)
assert order1.call_count >= 1
assert order2.call_count >= 1
@pytest.mark.asyncio
async def test_router_retry_policy_400_keeps_upstream_error_when_tags_narrow_the_group(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
router = litellm.Router(
model_list=[
_retry_skip_deployment(
"tagged", "tagged", litellm_params={"tags": ["free"]}, model_info={"enable_tag_filtering": True}
),
_retry_skip_deployment("untagged", "untagged", model_info={"enable_tag_filtering": True}),
],
num_retries=2,
retry_policy={"BadRequestErrorRetries": 2},
disable_cooldowns=True,
)
with respx.mock as respx_mock:
tagged = respx_mock.post("https://tagged.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
untagged = respx_mock.post("https://untagged.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
with pytest.raises(litellm.BadRequestError) as raised:
await router.acompletion(
model="gpt-5.6",
messages=[{"role": "user", "content": "hi"}],
metadata={"tags": ["free"]},
)
assert "upstream refused this request" in str(raised.value)
assert "No deployments available" not in str(raised.value)
assert tagged.call_count == 3
assert untagged.call_count == 0
@pytest.mark.asyncio
async def test_router_retry_policy_400_never_returns_to_a_deployment_that_already_refused(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(simple_shuffle.random, "choice", lambda deployments: deployments[0])
router = litellm.Router(
model_list=[
_retry_skip_deployment("first-refuser", "first-refuser", litellm_params={"weight": 1}),
_retry_skip_deployment("second-refuser", "second-refuser", litellm_params={"weight": 0}),
_retry_skip_deployment("accepting", "accepting", litellm_params={"weight": 0}),
],
num_retries=3,
retry_policy={"BadRequestErrorRetries": 3},
disable_cooldowns=True,
)
with respx.mock as respx_mock:
first = respx_mock.post("https://first-refuser.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
second = respx_mock.post("https://second-refuser.local/v1/chat/completions").mock(
return_value=httpx.Response(400, json={"error": _UPSTREAM_400})
)
accepting = respx_mock.post("https://accepting.local/v1/chat/completions").mock(
return_value=httpx.Response(
200,
json={
"id": "chatcmpl-lit-7036",
"object": "chat.completion",
"created": 1,
"model": "gpt-5.6",
"choices": [
{"index": 0, "message": {"role": "assistant", "content": "hi back"}, "finish_reason": "stop"}
],
"usage": {"prompt_tokens": 1, "completion_tokens": 2, "total_tokens": 3},
},
)
)
response = await router.acompletion(model="gpt-5.6", messages=[{"role": "user", "content": "hi"}])
assert first.call_count == 1
assert second.call_count == 1
assert accepting.call_count == 1
assert response.choices[0].message.content == "hi back"
def _make_failure_logging_obj():
return LiteLLMLogging(
model="gpt-5.6",

View file

@ -18,10 +18,10 @@ from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from litellm.router_utils.prompt_caching_cache import PromptCachingCache
from litellm.types.router import RouterRateLimitError
from litellm.utils import _get_deployment_order, _get_order_filtered_deployments
from litellm.utils import _get_deployment_order, get_order_filtered_deployments
# ---------------------------------------------------------------------------
# Unit tests for _get_order_filtered_deployments
# Unit tests for get_order_filtered_deployments
# ---------------------------------------------------------------------------
@ -42,7 +42,7 @@ class TestGetOrderFilteredDeployments:
self._make_deployment(2, "b"),
self._make_deployment(1, "c"),
]
result = _get_order_filtered_deployments(deps)
result = get_order_filtered_deployments(deps)
assert len(result) == 2
assert all(d["model_info"]["id"] in ("a", "c") for d in result)
@ -52,7 +52,7 @@ class TestGetOrderFilteredDeployments:
self._make_deployment(2, "b"),
self._make_deployment(3, "c"),
]
result = _get_order_filtered_deployments(deps, target_order=2)
result = get_order_filtered_deployments(deps, target_order=2)
assert len(result) == 1
assert result[0]["model_info"]["id"] == "b"
@ -61,7 +61,7 @@ class TestGetOrderFilteredDeployments:
self._make_deployment(1, "a"),
self._make_deployment(2, "b"),
]
result = _get_order_filtered_deployments(deps, target_order=99)
result = get_order_filtered_deployments(deps, target_order=99)
assert result == []
def test_target_order_no_match_does_not_reselect_lower_order(self):
@ -70,7 +70,7 @@ class TestGetOrderFilteredDeployments:
self._make_deployment(2, "b"),
]
remaining_after_pre_call = [deps[0]]
result = _get_order_filtered_deployments(remaining_after_pre_call, target_order=2)
result = get_order_filtered_deployments(remaining_after_pre_call, target_order=2)
assert result == []
def test_no_order_set_returns_all(self):
@ -78,11 +78,11 @@ class TestGetOrderFilteredDeployments:
self._make_deployment(None, "a"),
self._make_deployment(None, "b"),
]
result = _get_order_filtered_deployments(deps)
result = get_order_filtered_deployments(deps)
assert len(result) == 2
def test_empty_list(self):
result = _get_order_filtered_deployments([])
result = get_order_filtered_deployments([])
assert result == []
def test_single_order_returns_all_with_that_order(self):
@ -90,7 +90,7 @@ class TestGetOrderFilteredDeployments:
self._make_deployment(1, "a"),
self._make_deployment(1, "b"),
]
result = _get_order_filtered_deployments(deps)
result = get_order_filtered_deployments(deps)
assert len(result) == 2

View file

@ -15,11 +15,11 @@ import pytest
import litellm
from litellm import Router
from litellm.utils import _get_excluded_filtered_deployments
from litellm.utils import get_excluded_filtered_deployments
# ---------------------------------------------------------------------------
# Unit tests for _get_excluded_filtered_deployments
# Unit tests for get_excluded_filtered_deployments
# ---------------------------------------------------------------------------
@ -37,17 +37,17 @@ def _make_dep(dep_id: str, weight: Optional[int] = None) -> dict:
class TestGetExcludedFilteredDeployments:
def test_no_excluded_returns_all(self):
deps = [_make_dep("a"), _make_dep("b")]
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=None)
result = get_excluded_filtered_deployments(deps, excluded_deployment_ids=None)
assert len(result) == 2
def test_empty_excluded_returns_all(self):
deps = [_make_dep("a"), _make_dep("b")]
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=[])
result = get_excluded_filtered_deployments(deps, excluded_deployment_ids=[])
assert len(result) == 2
def test_drops_excluded(self):
deps = [_make_dep("a"), _make_dep("b"), _make_dep("c")]
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=["b"])
result = get_excluded_filtered_deployments(deps, excluded_deployment_ids=["b"])
ids = sorted(d["model_info"]["id"] for d in result)
assert ids == ["a", "c"]
@ -57,12 +57,12 @@ class TestGetExcludedFilteredDeployments:
# error. Returning the original list here would re-include the
# just-failed deployment and let weighted failover re-pick it.
deps = [_make_dep("a"), _make_dep("b")]
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=["a", "b"])
result = get_excluded_filtered_deployments(deps, excluded_deployment_ids=["a", "b"])
assert result == []
def test_excluded_set_with_unknown_ids(self):
deps = [_make_dep("a"), _make_dep("b")]
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=["zzz"])
result = get_excluded_filtered_deployments(deps, excluded_deployment_ids=["zzz"])
assert len(result) == 2
def test_handles_missing_model_info(self):
@ -70,7 +70,7 @@ class TestGetExcludedFilteredDeployments:
{"model_name": "x", "litellm_params": {"model": "gpt-4o"}}, # no model_info
_make_dep("b"),
]
result = _get_excluded_filtered_deployments(deps, excluded_deployment_ids=["b"])
result = get_excluded_filtered_deployments(deps, excluded_deployment_ids=["b"])
assert len(result) == 1

View file

@ -1,11 +1,9 @@
"""Tests for scripts/test_quality_gate.py.
The gate's whole value is that it blames a change only for what it adds, that a limit
can never rise, and that a limit cannot stay above a count the branch pushed below it.
All three live in pure functions, so they are tested directly: `evaluate` for the blame
rule, `ratcheted_budget` for the one-way ratchet, `unratcheted` for the ceiling a branch
left behind, and `parse_changed_lines` for the diff scan that turns a breach into
file:line.
The gate's whole value is that it blames a change only for what it adds and that a
limit can never rise. Both live in pure functions, so they are tested directly:
`evaluate` for the blame rule, `ratcheted_budget` for the one-way ratchet, and
`parse_changed_lines` for the diff scan that turns a breach into file:line.
"""
import importlib.util
@ -73,29 +71,6 @@ def test_ratchet_lowers_a_rule_introduced_on_this_branch_like_any_other():
assert updated["TQ001"]["limit"] == 4
def test_a_branch_that_cleared_violations_must_lower_the_ceiling():
stale = gate.unratcheted({"TQ001": 6}, {"TQ001": 10}, _BUDGET)
assert [(b.rule, b.total, b.cap, b.added) for b in stale] == [("TQ001", 6, 10, -4)]
def test_headroom_already_in_the_base_is_not_blamed_on_this_branch():
assert gate.unratcheted({"TQ001": 6}, {"TQ001": 6}, _BUDGET) == ()
def test_a_branch_that_cleared_down_to_the_ceiling_exactly_is_clean():
assert gate.unratcheted({"TQ001": 10}, {"TQ001": 12}, _BUDGET) == ()
def test_a_branch_that_added_violations_is_not_a_ratchet_finding():
assert gate.unratcheted({"TQ001": 14}, {"TQ001": 10}, _BUDGET) == ()
def test_the_ratchet_finding_survives_the_update_that_answers_it():
cleared = {"TQ001": 6}
updated = gate.ratcheted_budget(_BUDGET, cleared, {"TQ001": 10})
assert gate.unratcheted(cleared, {"TQ001": 10}, updated) == ()
def test_parse_changed_lines_groups_hunks_under_their_own_file():
diff = (
"diff --git a/tests/a.py b/tests/a.py\n"

View file

@ -12,7 +12,7 @@ import {
} from "@/components/mcp_tools/types";
import { AUTH_TYPES_REQUIRING_CREDENTIALS } from "./createServerPayload";
import { TOOL_DISPLAY_NAME_PATTERN, normalizeEnvVars } from "./utils";
import { buildEditServerPayload, type EditServerUiState } from "./editServerPayload";
import { buildEditServerPayload, type EditServerFormValues, type EditServerUiState } from "./editServerPayload";
import { CASES, baseUi } from "./editServerPayload.differential.cases";
// GENERATED by scratchpad/emit_test.py. The body below is machine-extracted from
@ -317,6 +317,47 @@ describe("buildEditServerPayload matches the pre-extraction handleSave body", ()
});
});
const EDIT_FORM_VALUES: EditServerFormValues = {
server_name: "srv",
alias: "srv_alias",
description: "a server",
transport: "http",
url: "https://example.com/mcp",
auth_type: "none",
mcp_access_groups: [],
extra_headers: [],
static_headers: [],
env_vars: [],
allow_all_keys: false,
available_on_public_internet: true,
};
describe("buildEditServerPayload wire contract", () => {
it("carries an edited alias and the server identifier onto the wire", () => {
const result = buildEditServerPayload({ ...EDIT_FORM_VALUES, alias: "renamed" }, baseUi);
expect(result).toMatchObject({ kind: "ok", payload: { server_id: "srv_1", alias: "renamed" } });
});
it.fails(
"sends description as an explicit null when the field is cleared (expected to fail until the forms revamp, tri-state PATCH tracker: today the cleared field reaches the wire as an empty string)",
() => {
const result = buildEditServerPayload({ ...EDIT_FORM_VALUES, description: "" }, baseUi);
expect(result).toMatchObject({ kind: "ok", payload: { description: null } });
},
);
it.fails(
"sends only the server identifier and the edited alias (expected to fail until the forms revamp, tri-state PATCH tracker)",
() => {
const result = buildEditServerPayload({ ...EDIT_FORM_VALUES, alias: "renamed" }, baseUi);
expect(result).toStrictEqual({ kind: "ok", payload: { server_id: "srv_1", alias: "renamed" } });
},
);
});
void ADMIN_CONFIG_CREDENTIAL_KEYS;
void AUTH_TYPE;
void AUTH_TYPES_REQUIRING_CREDENTIALS;

View file

@ -1729,5 +1729,77 @@ describe("ModelInfoView", () => {
expect(payload.litellm_params.cache_control_injection_points).toEqual([{ location: "message", index: "2" }]);
});
});
const setInputCost = (value: string) => {
fireEvent.change(screen.getByPlaceholderText("Enter input cost"), { target: { value } });
};
it("carries an edited input cost and the model identifier onto the wire", async () => {
const user = userEvent.setup();
await enterEditMode(user);
setInputCost("5");
const payload = await save(user);
expect(mockModelPatchUpdateCall.mock.calls[0][2]).toBe("123");
expect(payload.litellm_params.input_cost_per_token).toBe(5 / 1_000_000);
});
it.fails(
"sends only the edited input cost (expected to fail until the forms revamp, tri-state PATCH tracker)",
async () => {
const user = userEvent.setup();
await enterEditMode(user);
setInputCost("5");
const payload = await save(user);
expect(payload).toStrictEqual({ litellm_params: { input_cost_per_token: 5 / 1_000_000 } });
},
);
const savePayloadAfterCostEditOnResolvedModel = async () => {
const resolved = {
...defaultModelData,
model_info: {
...defaultModelData.model_info,
max_input_tokens: 128_000,
mode: "chat",
supports_vision: true,
supports_function_calling: true,
},
};
mockUseModelsInfo.mockReturnValue({ data: { data: [resolved] }, isLoading: false, error: null });
mockModelInfoV1Call.mockResolvedValue({ data: [resolved] });
const user = userEvent.setup();
await enterEditMode(user);
setInputCost("5");
return save(user);
};
it.fails(
"leaves max_input_tokens off the wire when only the input cost is edited (expected to fail until the forms revamp, tri-state PATCH tracker)",
async () => {
const payload = await savePayloadAfterCostEditOnResolvedModel();
expect(payload.model_info).not.toHaveProperty("max_input_tokens");
},
);
it.fails(
"leaves mode off the wire when only the input cost is edited (expected to fail until the forms revamp, tri-state PATCH tracker)",
async () => {
const payload = await savePayloadAfterCostEditOnResolvedModel();
expect(payload.model_info).not.toHaveProperty("mode");
},
);
it.fails(
"leaves every supports_ capability off the wire when only the input cost is edited (expected to fail until the forms revamp, tri-state PATCH tracker)",
async () => {
const payload = await savePayloadAfterCostEditOnResolvedModel();
expect(Object.keys(payload.model_info).filter((key) => key.startsWith("supports_"))).toStrictEqual([]);
},
);
});
});

View file

@ -544,6 +544,37 @@ describe("CreateKey", () => {
expect((await createdPayload()).metadata).toBe('{"team":"research"}');
});
it("carries the typed key alias and the chosen team onto the wire", async () => {
state.teams = [{ team_id: "team-1", team_alias: "Team One", models: [] }];
await openModal({ teams: state.teams as unknown as Team[] });
await nameTheKey("wire-alias");
await userEvent.click(await screen.findByLabelText("Team"));
await userEvent.click(await screen.findByRole("option", { name: /Team One/ }));
await submit();
expect(await createdPayload()).toMatchObject({ key_alias: "wire-alias", team_id: "team-1" });
});
it("sends team_id as an explicit null when no team is chosen", async () => {
await openModal();
await nameTheKey();
await submit();
expect(await createdPayload()).toHaveProperty("team_id", null);
});
it.fails(
"adds no keys for an Optional Settings section the user opened but never filled (expected to fail until the forms revamp, tri-state PATCH tracker)",
async () => {
await openModal();
await nameTheKey();
await openSection(/Optional Settings/i);
await submit();
expect(await createdPayload()).toStrictEqual(ALL_CLOSED_PAYLOAD);
},
);
});
describe("key ownership", () => {

View file

@ -2300,5 +2300,57 @@ describe("KeyEditView", () => {
});
expect(onSubmitMock.mock.calls[0][0]).toHaveProperty("tag_rpm_limit", { "test-tag": 7 });
});
const setRpmLimit = (value: string) => {
fireEvent.change(screen.getByLabelText("RPM Limit"), { target: { value } });
};
it("carries an edited RPM limit and the key identifier onto the wire", async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
renderForPayload(onSubmitMock);
await screen.findByRole("button", { name: /save changes/i });
setRpmLimit("25");
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalledTimes(1);
});
expect(onSubmitMock.mock.calls[0][0]).toMatchObject({ token: "test-token-123", rpm_limit: "25" });
});
it.fails(
"sends max_budget as an explicit null when the field is cleared (expected to fail until the forms revamp, tri-state PATCH tracker: today the view hands KeyInfoView an empty string and handleKeyUpdate maps it to null)",
async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
renderForPayload(onSubmitMock);
await screen.findByRole("button", { name: /save changes/i });
await userEvent.clear(screen.getByLabelText("Max Budget (USD)"));
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalledTimes(1);
});
expect(onSubmitMock.mock.calls[0][0]).toHaveProperty("max_budget", null);
},
);
it.fails(
"sends only the key identifier and the edited RPM limit (expected to fail until the forms revamp, tri-state PATCH tracker)",
async () => {
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
renderForPayload(onSubmitMock);
await screen.findByRole("button", { name: /save changes/i });
setRpmLimit("25");
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmitMock).toHaveBeenCalledTimes(1);
});
expect(onSubmitMock.mock.calls[0][0]).toStrictEqual({ token: "test-token-123", rpm_limit: "25" });
},
);
});
});

View file

@ -1015,6 +1015,16 @@ describe("KeyInfoView", () => {
expect(keyUpdateCall).toHaveBeenCalledWith(expect.anything(), expect.objectContaining({ policies: [] }));
});
it("puts the key identifier and an explicit null max_budget on the wire when the edit view hands over a cleared budget", async () => {
await enterEditMode({ ...MOCK_KEY_DATA, user_id: "proxy-admin-user" } as KeyResponse);
await editViewMocks.onSubmit!({ token: MOCK_KEY_DATA.token, max_budget: "" });
expect(keyUpdateCall).toHaveBeenCalledWith(
expect.anything(),
expect.objectContaining({ key: "test-token-123", max_budget: null }),
);
});
});
describe("MCP tool permissions on save", () => {