mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_/guardrail-automation-testing-3ecd3d
This commit is contained in:
commit
5561d2e476
48 changed files with 2452 additions and 462 deletions
129
.github/workflows/report-rust-release-wheel.yml
vendored
129
.github/workflows/report-rust-release-wheel.yml
vendored
|
|
@ -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,
|
||||
});
|
||||
}
|
||||
111
.github/workflows/test-rust.yml
vendored
111
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
210
litellm/llms/mistral/audio_speech/transformation.py
Normal file
210
litellm/llms/mistral/audio_speech/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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})")
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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__()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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([]);
|
||||
},
|
||||
);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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" });
|
||||
},
|
||||
);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue