merge: litellm_internal_staging into litellm_slot_leak_stream_logging

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
milan 2026-08-20 00:59:09 +00:00
commit 66dbf2101c
274 changed files with 16615 additions and 9447 deletions

2
.github/CODEOWNERS vendored
View file

@ -1,3 +1,5 @@
/ui/ @yuneng-jiang @ryan-crabbe-berri
/litellm/proxy/_experimental/out/ @yuneng-jiang @ryan-crabbe-berri
/ui/litellm-dashboard/src/lib/http/schema.d.ts
/model_prices_and_context_window.json @mateo-berri
/litellm/model_prices_and_context_window_backup.json @mateo-berri

View file

@ -3,9 +3,18 @@ description: >-
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
and expose decision=run|skip. decision=skip means only ui/**, **.md or **.mdx files
changed, so callers can short-circuit expensive steps while the job still completes
successfully and satisfies its required status check. The decision defaults to run for
any non pull_request event or whenever the changed set cannot be resolved, so tests are
never skipped when the classification is uncertain.
successfully and satisfies its required status check. The file list comes from the
pull request itself rather than from a git diff, because the checked-out merge ref is
recomputed as the base branch advances and would otherwise attribute the base
branch's own commits to the pull request. The decision defaults to run for any non
pull_request event or whenever the changed set cannot be resolved, so tests are never
skipped when the classification is uncertain.
inputs:
github-token:
description: "Token used to list the pull request's files; needs pull-requests: read"
required: false
default: ${{ github.token }}
outputs:
decision:
@ -18,31 +27,8 @@ runs:
- id: classify
shell: bash
env:
BASE_SHA: ${{ github.event.pull_request.base.sha }}
run: |
set -uo pipefail
if [ -z "${BASE_SHA:-}" ]; then
echo "detect-backend-changes: not a pull_request event; running job"
echo "decision=run" >> "${GITHUB_OUTPUT}"
exit 0
fi
if ! git fetch --no-tags --depth=1 origin "${BASE_SHA}" >/dev/null 2>&1; then
echo "detect-backend-changes: could not fetch base ${BASE_SHA}; running job"
echo "decision=run" >> "${GITHUB_OUTPUT}"
exit 0
fi
changed="$(git diff --name-only "${BASE_SHA}" HEAD 2>/dev/null)" || {
echo "detect-backend-changes: git diff failed; running job"
echo "decision=run" >> "${GITHUB_OUTPUT}"
exit 0
}
if [ -z "${changed}" ]; then
echo "detect-backend-changes: no changed files vs ${BASE_SHA}; skipping job"
echo "decision=skip" >> "${GITHUB_OUTPUT}"
exit 0
fi
echo "detect-backend-changes: changed files vs ${BASE_SHA}:"
printf '%s\n' "${changed}" | sed 's/^/ /'
decision="$(printf '%s\n' "${changed}" | bash .circleci/scripts/classify_changes.sh backend)" || decision="run"
echo "detect-backend-changes: decision=${decision}"
echo "decision=${decision}" >> "${GITHUB_OUTPUT}"
GH_TOKEN: ${{ inputs.github-token }}
REPO: ${{ github.repository }}
PR_NUMBER: ${{ github.event.pull_request.number }}
CHANGED_FILE_COUNT: ${{ github.event.pull_request.changed_files }}
run: bash "${GITHUB_ACTION_PATH}/../../scripts/detect_backend_changes.sh"

View file

@ -53,7 +53,8 @@ After: the same request comes back with real token counts, so the dashboard show
**Please complete all items before asking a LiteLLM maintainer to review your PR**
- [ ] I have added meaningful tests
- [ ] My PR passes all CI/CD checks (e.g., lint, format, unit tests)
- [ ] The handful of test files covering my change pass locally, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. Leave the suites (`make test-unit-*`, `make test-unit`) to CI: it finishes in ~15 minutes where a laptop takes an hour or more
- [ ] My PR passes all required CI/CD checks (e.g., lint, schema.d.ts sync check, etc.)
- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem
- [ ] I have received a Greptile **Confidence Score of at least 4/5** before requesting a maintainer review (Greptile reviews automatically once the PR is opened; only comment `@greptileai` to re-request a review after pushing changes)

41
.github/scripts/detect_backend_changes.sh vendored Executable file
View file

@ -0,0 +1,41 @@
#!/usr/bin/env bash
set -uo pipefail
readonly API_FILE_CEILING=3000
decide() {
echo "detect-backend-changes: decision=$1"
[ -z "${GITHUB_OUTPUT:-}" ] || echo "decision=$1" >>"${GITHUB_OUTPUT}"
exit 0
}
run_full() {
echo "detect-backend-changes: $1; running job"
decide run
}
here="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
classify="${here}/../../.circleci/scripts/classify_changes.sh"
[ -n "${PR_NUMBER:-}" ] || run_full "not a pull_request event"
[ -n "${REPO:-}" ] || run_full "no repository in the environment"
case "${CHANGED_FILE_COUNT:-}" in
'' | *[!0-9]*) run_full "the event payload carries no changed_files count" ;;
esac
[ "${CHANGED_FILE_COUNT}" -le "${API_FILE_CEILING}" ] ||
run_full "PR #${PR_NUMBER} changes ${CHANGED_FILE_COUNT} files, past the ${API_FILE_CEILING}-file listing ceiling"
changed="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate --jq '.[].filename')" ||
run_full "could not list the files on PR #${PR_NUMBER}"
[ -n "${changed}" ] || run_full "the API listed no files on PR #${PR_NUMBER}"
echo "detect-backend-changes: files changed by PR #${PR_NUMBER}:"
printf '%s\n' "${changed}" | sed 's/^/ /'
decision="$(printf '%s\n' "${changed}" | bash "${classify}" backend)" ||
run_full "classify_changes.sh failed"
case "${decision}" in
run | skip) decide "${decision}" ;;
*) run_full "classify_changes.sh printed an unexpected decision: ${decision}" ;;
esac

View file

@ -60,6 +60,9 @@ jobs:
name: Run tests
runs-on: ubuntu-latest
timeout-minutes: ${{ inputs.job-timeout-minutes }}
permissions:
contents: read
pull-requests: read
outputs:
decision: ${{ steps.changes.outputs.decision }}

View file

@ -23,6 +23,9 @@ jobs:
documentation:
runs-on: ubuntu-latest
timeout-minutes: 10
permissions:
contents: read
pull-requests: read
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0

1
.gitignore vendored
View file

@ -5,6 +5,7 @@ tests/e2e/.fixtures/
.venv_policy_test
.env
.claude
CLAUDE.local.md
.newenv
newenv/*
litellm/proxy/myenv/*

View file

@ -13,8 +13,8 @@ Here are the core requirements for any PR submitted to LiteLLM:
- [ ] **Add testing** - Adding at least 1 test is a hard requirement - [see details](#adding-testing)
- [ ] **Ensure your PR passes all checks**:
- [ ] [Unit Tests](#running-unit-tests) - `make test-unit`
- [ ] [Linting / Formatting](#running-linting-and-formatting-checks) - `make lint`
- [ ] [The tests covering your change](#running-unit-tests) pass, e.g. `uv run pytest tests/test_litellm/<your_test_file>.py -v`. CI runs the full unit test matrix, so you don't need to run the whole suite locally
#### UI PRs
@ -71,8 +71,8 @@ make format
# Run all linting checks (matches CI exactly)
make lint
# Run unit tests to ensure nothing is broken
make test-unit
# Run the tests covering your change (CI runs the full suite)
uv run pytest tests/test_litellm/<your_test_file>.py -v
# Commit your changes (must follow Conventional Commits — see above)
git add .
@ -123,12 +123,13 @@ def test_your_feature():
### Running Unit Tests
Run all unit tests (uses parallel execution for speed):
Run the tests covering your change:
```bash
make test-unit
uv run pytest tests/test_litellm/test_your_file.py -v
```
`tests/test_litellm` holds thousands of tests, so running all of it locally takes a long time. CI runs it as a parallel matrix (`make test-unit-llms`, `make test-unit-proxy-core`, and the other `test-unit-*` targets) on beefier boxes, so if, for whatever reason, you must run the whole suite, it's better to rely on CI to do that.
If you're running broader test suites, proxy tests, or anything that touches PostgreSQL-backed fixtures/plugins, install the full local test environment first:
```bash
@ -137,11 +138,6 @@ make install-test-deps
This syncs the locked test environment used across the repo, including `psycopg` v3 plus `psycopg-binary` (used by `pytest-postgresql`), `psycopg2-binary` (used by some proxy E2E tests), and a generated Prisma client for DB-backed proxy tests, so pytest startup matches CI without manual package installs.
Run specific test files:
```bash
uv run pytest tests/test_litellm/test_your_file.py -v
```
### Running Linting and Formatting Checks
Run all linting checks (matches CI exactly):

View file

@ -57,7 +57,7 @@
"limit": 5663
},
"reportMissingTypeArgument": {
"limit": 15557
"limit": 15555
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,13 +105,13 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 39043
"limit": 39017
},
"reportUnknownParameterType": {
"limit": 19887
"limit": 19885
},
"reportUnknownVariableType": {
"limit": 30574
"limit": 30572
},
"reportUnnecessaryCast": {
"limit": 117

View file

@ -145,6 +145,11 @@ NUMBER_KEYS: dict[str, JsonSchema] = {
"minimum": 1,
"description": "Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%).",
},
"regional_endpoint_uplift_multiplier": {
"type": "number",
"minimum": 1,
"description": "Multiplier applied to all token costs when served from a non-global Vertex AI endpoint (e.g. 1.10 = +10%).",
},
}
COST_DESCRIPTIONS: dict[str, str] = {

View file

@ -18,7 +18,7 @@ type: application
# This is the chart version. This version number should be incremented each time you make changes
# to the chart and its templates, including the app version.
# Versions are expected to follow Semantic Versioning (https://semver.org/)
version: 1.1.1
version: 1.1.2
# This is the version number of the application being deployed. This version number should be
# incremented each time you make changes to the application. Versions are not expected to

View file

@ -29,7 +29,7 @@ If `db.useStackgresOperator` is used (not yet implemented):
| `masterkey` | The Master API Key for LiteLLM. If not specified, a random key in the `sk-...` format is generated. | N/A |
| `environmentSecrets` | An optional array of Secret object names. The keys and values in these secrets will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `environmentConfigMaps` | An optional array of ConfigMap object names. The keys and values in these configmaps will be presented to the LiteLLM proxy pod as environment variables. See below for an example Secret object. | `[]` |
| `image.repository` | LiteLLM Proxy image repository | `docker.litellm.ai/berriai/litellm` |
| `image.repository` | LiteLLM Proxy image repository | `ghcr.io/berriai/litellm` |
| `image.pullPolicy` | LiteLLM Proxy image pull policy | `IfNotPresent` |
| `image.tag` | Overrides the image tag whose default the latest version of LiteLLM at the time this chart was published. | `""` |
| `imagePullSecrets` | Registry credentials for the LiteLLM and initContainer images. | `[]` |

View file

@ -15,7 +15,7 @@ tests:
pattern: -litellm$
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-database:test
value: ghcr.io/berriai/litellm:test
- it: should work with tolerations
template: deployment.yaml
set:
@ -337,7 +337,7 @@ tests:
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
extraInitContainers:
- name: init-tpl
@ -348,7 +348,7 @@ tests:
path: spec.template.spec.initContainers
content:
name: init-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
command: ["echo", "hello"]
- it: should work with extraContainers
template: deployment.yaml
@ -366,7 +366,7 @@ tests:
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
extraContainers:
- name: sidecar-tpl
@ -376,12 +376,12 @@ tests:
path: spec.template.spec.containers
content:
name: sidecar-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
- it: should support tpl in podAnnotations
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
# Mirrors the real-world scenario this feature unblocks:
# user disables the built-in ConfigMap (and its built-in checksum/config
@ -398,7 +398,7 @@ tests:
value: "test"
- equal:
path: spec.template.metadata.annotations["example.com/some-key"]
value: "ghcr.io/berriai/litellm-database"
value: "ghcr.io/berriai/litellm"
- equal:
path: spec.template.metadata.annotations["example.com/literal"]
value: "plain-string-value"

View file

@ -208,7 +208,7 @@ tests:
template: migrations-job.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
migrationJob:
enabled: true
@ -221,7 +221,7 @@ tests:
path: spec.template.spec.initContainers
content:
name: init-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
command: ["echo", "hello"]
- it: should work with extraContainers
template: migrations-job.yaml
@ -241,7 +241,7 @@ tests:
template: migrations-job.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
repository: ghcr.io/berriai/litellm
tag: test
migrationJob:
enabled: true
@ -253,7 +253,7 @@ tests:
path: spec.template.spec.containers
content:
name: sidecar-tpl
image: "ghcr.io/berriai/litellm-database:test"
image: "ghcr.io/berriai/litellm:test"
- it: should render the pod-level securityContext from podSecurityContext
template: migrations-job.yaml
set:

View file

@ -6,8 +6,9 @@ replicaCount: 1
# numWorkers: 2
image:
# Use "ghcr.io/berriai/litellm-database" for optimized image with database
repository: ghcr.io/berriai/litellm-database
# Bundles the prisma CLI and engines, which is what lets the migrations job
# and the proxy's own schema check run without network access.
repository: ghcr.io/berriai/litellm
pullPolicy: Always
# Overrides the image tag whose default is the chart appVersion.
# tag: "latest"

View file

@ -0,0 +1,9 @@
-- CreateTable
CREATE TABLE "LiteLLM_ProxyWorkerHeartbeat" (
"worker_id" TEXT NOT NULL,
"hostname" TEXT NOT NULL,
"started_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"last_heartbeat_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
CONSTRAINT "LiteLLM_ProxyWorkerHeartbeat_pkey" PRIMARY KEY ("worker_id")
);

View file

@ -0,0 +1,4 @@
UPDATE "LiteLLM_SpendLogs"
SET "created_at" = "endTime",
"updated_at" = "endTime"
WHERE "created_at" > "endTime" + interval '1 hour';

View file

@ -947,6 +947,17 @@ model LiteLLM_DailyTagSpend {
}
// One row per live proxy worker process. Workers upsert their row on a fixed
// heartbeat; counting rows with a recent heartbeat tells how many workers share
// this database, which lets the Admin UI hide its "no Redis" warning for
// deployments that are provably a single worker.
model LiteLLM_ProxyWorkerHeartbeat {
worker_id String @id
hostname String
started_at DateTime @default(now())
last_heartbeat_at DateTime @default(now())
}
// Track the status of cron jobs running. Only allow one pod to run the job at a time
model LiteLLM_CronJob {
cronjob_id String @id @default(cuid()) // Unique ID for the record

View file

@ -1,5 +1,5 @@
import json
from collections.abc import Iterable, Iterator
from collections.abc import Iterable, Iterator, Mapping
from dataclasses import dataclass
from typing import Any, Final, Literal
@ -87,7 +87,7 @@ async def _handle_completed_batch(
return batch_cost, batch_usage, [model_name]
return _aggregate_batch_cost_usage_models(
entries=_iter_batch_input_entries(file_content),
entries=_iter_batch_output_entries(file_content),
custom_llm_provider=custom_llm_provider,
model_name=model_name,
model_info=model_info,
@ -111,43 +111,91 @@ def _iter_successful_output_line_stats(
model_name: str | None,
model_info: ModelInfo | None,
) -> Iterator[_BatchOutputLineStats]:
for entry in entries:
stats = _safe_output_line_stats(entry, custom_llm_provider, model_name, model_info)
if stats is not None:
yield stats
def _safe_output_line_stats(
entry: Mapping[str, Any],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
model_info: ModelInfo | None,
) -> _BatchOutputLineStats | None:
"""Return the stats for one batch output line, or None for a line that is
unsuccessful or cannot be costed, so a single bad line never aborts the
whole batch's cost accounting."""
custom_id: Final = entry.get("custom_id") if isinstance(entry, dict) else None
try:
if not _batch_response_was_successful(entry, custom_llm_provider):
return None
return _compute_output_line_stats(entry, custom_llm_provider, model_name, model_info)
except Exception as e: # noqa: BLE001 # any single line's costing failure must not abort the whole batch
verbose_logger.warning(
"batch output line could not be costed, so it is billed at $0 and the rest of the batch "
"is still billed. custom_id=%s error=%s",
custom_id,
str(e),
)
return None
def _compute_output_line_stats(
entry: Mapping[str, Any],
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
model_info: ModelInfo | None,
) -> _BatchOutputLineStats:
response_body: Final = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
usage: Final = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
prompt_details: Final = parse_prompt_tokens_details(usage)
raw_model: Final = response_body.get("model")
response_model: Final = raw_model if isinstance(raw_model, str) and raw_model else None
return _BatchOutputLineStats(
cost=_output_line_cost(
response_body=response_body,
usage=usage,
custom_llm_provider=custom_llm_provider,
model_name=model_name,
response_model=response_model,
model_info=model_info,
),
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
cache_read_tokens=prompt_details["cache_hit_tokens"],
cache_creation_tokens=prompt_details["cache_creation_tokens"],
model=response_model,
)
def _output_line_cost(
response_body: Mapping[str, Any],
usage: Usage,
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm", "anthropic", "bedrock"],
model_name: str | None,
response_model: str | None,
model_info: ModelInfo | None,
) -> float:
from litellm.cost_calculator import batch_cost_calculator
for entry in entries:
if not _batch_response_was_successful(entry, custom_llm_provider):
continue
response_body = _get_response_from_batch_job_output_file(entry, custom_llm_provider)
usage = _get_batch_job_usage_from_response_body(response_body, custom_llm_provider)
prompt_details = parse_prompt_tokens_details(usage)
raw_model = response_body.get("model")
response_model = raw_model if isinstance(raw_model, str) and raw_model else None
if model_info is not None or custom_llm_provider in ("anthropic", "bedrock"):
if custom_llm_provider == "bedrock" and model_name:
cost_model = model_name
else:
cost_model = response_model or model_name or ""
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=cost_model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
line_cost = prompt_cost + completion_cost
else:
line_cost = litellm.completion_cost(
completion_response=response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
yield _BatchOutputLineStats(
cost=line_cost,
prompt_tokens=usage.prompt_tokens,
completion_tokens=usage.completion_tokens,
total_tokens=usage.total_tokens,
cache_read_tokens=prompt_details["cache_hit_tokens"],
cache_creation_tokens=prompt_details["cache_creation_tokens"],
model=response_model,
if model_info is None and custom_llm_provider not in ("anthropic", "bedrock"):
return litellm.completion_cost(
completion_response=response_body,
custom_llm_provider=custom_llm_provider,
call_type=CallTypes.aretrieve_batch.value,
)
cost_model: Final = (
model_name if custom_llm_provider == "bedrock" and model_name else response_model or model_name or ""
)
prompt_cost, completion_cost = batch_cost_calculator(
usage=usage,
model=cost_model,
custom_llm_provider=custom_llm_provider,
model_info=model_info,
)
return prompt_cost + completion_cost
def _aggregate_batch_cost_usage_models(
@ -338,9 +386,10 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
def _get_file_content_as_dictionary(file_content: bytes) -> list[dict]:
"""
Get the file content as a list of dictionaries from JSON Lines format
Get the file content as a list of dictionaries from JSON Lines format,
skipping malformed lines
"""
return list(_iter_batch_input_entries(file_content))
return list(_iter_batch_output_entries(file_content))
def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
@ -361,15 +410,29 @@ def _iter_batch_input_lines(file_content: bytes) -> Iterator[bytes]:
yield line
def _iter_batch_input_entries(file_content: bytes) -> Iterator[dict]:
def _iter_batch_output_entries(file_content: bytes) -> Iterator[dict]:
"""
Yield parsed batch input JSONL entries one at a time without materializing the
whole file as a list, so peak memory stays bounded. Raises on a malformed line;
callers that must survive bad rows should iterate ``_iter_batch_input_lines``
and parse per-row instead.
Yield parsed batch output JSONL entries one at a time without materializing
the whole file as a list, so peak memory stays bounded. A malformed or
non-object line is skipped with a warning so one bad line never aborts the
whole batch's cost accounting.
"""
for line in _iter_batch_input_lines(file_content):
yield json.loads(line)
entry = _parse_batch_output_line(line)
if entry is not None:
yield entry
def _parse_batch_output_line(line: bytes) -> dict | None:
try:
parsed: Final = json.loads(line)
except ValueError as e:
verbose_logger.warning("skipping malformed batch output line: %s", str(e))
return None
if isinstance(parsed, dict):
return parsed
verbose_logger.warning("skipping non-object batch output line of type %s", type(parsed).__name__)
return None
# A batch request's input tokens scale roughly with its serialized size, so this
@ -440,7 +503,9 @@ def _count_prompt_or_input_tokens(model: str, value: Any) -> int:
return 0
def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_provider: str = "openai") -> Usage:
def _get_batch_job_usage_from_response_body(
response_body: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> Usage:
"""
Get the tokens of a batch job from the response body
"""
@ -472,7 +537,7 @@ def _get_batch_job_usage_from_response_body(response_body: dict, custom_llm_prov
return usage
def _get_anthropic_result_from_batch_results_line(batch_results_line: dict) -> dict:
def _get_anthropic_result_from_batch_results_line(batch_results_line: Mapping[str, Any]) -> dict:
"""
Get the ``result`` object from a line of an Anthropic message batch results JSONL file.
@ -482,7 +547,9 @@ def _get_anthropic_result_from_batch_results_line(batch_results_line: dict) -> d
return batch_results_line.get("result", None) or {}
def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom_llm_provider: str = "openai") -> Any:
def _get_response_from_batch_job_output_file(
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> Any:
"""
Get the response from the batch job output file
"""
@ -495,7 +562,9 @@ def _get_response_from_batch_job_output_file(batch_job_output_file: dict, custom
return _response_body
def _batch_response_was_successful(batch_job_output_file: dict, custom_llm_provider: str = "openai") -> bool:
def _batch_response_was_successful(
batch_job_output_file: Mapping[str, Any], custom_llm_provider: str = "openai"
) -> bool:
"""
Check if the batch job response was successful

View file

@ -327,6 +327,8 @@ def cost_per_token(
service_tier: str | None = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
response: Any | None = None,
### REQUEST MODEL ###
request_model: str | None = None, # original request model for router detection
@ -587,6 +589,7 @@ def cost_per_token(
prompt_characters=prompt_characters,
completion_characters=completion_characters,
usage=usage_block,
vertex_location=vertex_location,
)
elif cost_router == "cost_per_token":
return google_cost_per_token(
@ -594,6 +597,7 @@ def cost_per_token(
custom_llm_provider=custom_llm_provider,
usage=usage_block,
service_tier=service_tier,
vertex_location=vertex_location,
)
elif custom_llm_provider == "anthropic":
return anthropic_cost_per_token(model=model, usage=usage_block, service_tier=service_tier)
@ -1071,6 +1075,7 @@ def _store_cost_breakdown_in_logging_obj(
reasoning_cost: float | None = None,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
) -> None:
"""
Helper function to store cost breakdown in the logging object.
@ -1090,6 +1095,7 @@ def _store_cost_breakdown_in_logging_obj(
margin_total_amount: Total margin added in USD
service_tier: Tier the costs above were priced on, already resolved
data_residency: Region uplift the costs above were priced on, already resolved
vertex_location: Vertex AI location the costs above were priced on, already resolved
"""
if litellm_logging_obj is None:
return
@ -1113,6 +1119,7 @@ def _store_cost_breakdown_in_logging_obj(
reasoning_cost=reasoning_cost,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
except Exception as breakdown_error:
@ -1149,6 +1156,8 @@ def completion_cost(
service_tier: str | None = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
) -> float:
"""
Calculate the cost of a given completion call fot GPT-3.5-turbo, llama2, any litellm supported llm.
@ -1577,6 +1586,7 @@ def completion_cost(
rerank_billed_units=rerank_billed_units,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
response=completion_response,
request_model=request_model_for_cost,
)
@ -1664,6 +1674,7 @@ def completion_cost(
usage=cost_per_token_usage_object,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
_reasoning_cost = _token_type_breakdown.reasoning_cost
_cache_read_cost = _token_type_breakdown.cache_read_cost
@ -1686,6 +1697,7 @@ def completion_cost(
reasoning_cost=_reasoning_cost,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
return _final_cost
@ -1765,6 +1777,8 @@ def response_cost_calculator(
service_tier: str | None = None, # for OpenAI service tier pricing
### DATA RESIDENCY ###
data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us")
### VERTEX LOCATION ###
vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global")
) -> float:
"""
Returns
@ -1797,6 +1811,7 @@ def response_cost_calculator(
litellm_logging_obj=litellm_logging_obj,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
return response_cost
except Exception as e:

View file

@ -42,6 +42,7 @@ from litellm.types.utils import (
# OpenTelemetry imports moved to individual functions to avoid import errors when not installed
if TYPE_CHECKING:
from opentelemetry.sdk.resources import Resource as _Resource
from opentelemetry.sdk.trace import TracerProvider as _SDKTracerProvider
from opentelemetry.sdk.trace.export import SpanExporter as _SpanExporter
from opentelemetry.trace import Context as _Context
@ -389,6 +390,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self._tracer_provider_cache: OrderedDict[str, _CachedTracerProvider] = OrderedDict()
self._tracer_provider_cache_lock: Final = threading.Lock()
self._max_dynamic_tracer_providers: Final = max(1, max_dynamic_tracer_providers)
self._litellm_resource_memo: _Resource | None = None
self._init_tracing(tracer_provider)
_debug_otel: Final = str(os.getenv("DEBUG_OTEL", "False")).lower()
@ -414,7 +416,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
self._init_otel_logger_on_litellm_proxy()
@staticmethod
def _get_litellm_resource(config: OpenTelemetryConfig):
def _get_litellm_resource(config: OpenTelemetryConfig) -> "_Resource":
"""Create an OpenTelemetry Resource using config-driven defaults."""
from opentelemetry.sdk.resources import OTELResourceDetector, Resource
@ -429,6 +431,21 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
env_resource: Final = otel_resource_detector.detect()
return base_resource.merge(env_resource)
def _litellm_resource(self) -> "_Resource":
"""The Resource every provider on this logger is built with, frozen at first use.
``Resource.create`` scans every installed distribution's entry points, roughly 3ms and
200 file opens, and the dynamic providers reach it from the async logging path. Freezing
also keeps them consistent with whatever this logger built at startup. ``cached_property``
locks class-wide before 3.12, which this file still supports.
"""
memo: Final = self._litellm_resource_memo
if memo is not None:
return memo
built: Final = self._get_litellm_resource(self.config)
self._litellm_resource_memo = built
return built
def _init_otel_logger_on_litellm_proxy(self):
"""
Initializes OpenTelemetry for litellm proxy server
@ -596,7 +613,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
from opentelemetry.trace import SpanKind
def create_tracer_provider():
provider: Final = TracerProvider(resource=self._get_litellm_resource(self.config))
provider: Final = TracerProvider(resource=self._litellm_resource())
provider.add_span_processor(self._get_span_processor())
return provider
@ -634,7 +651,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
metric_reader: Final = self._get_metric_reader()
return MeterProvider(
metric_readers=[metric_reader],
resource=self._get_litellm_resource(self.config),
resource=self._litellm_resource(),
)
meter_provider = self._get_or_create_provider(
@ -692,7 +709,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
def create_logger_provider():
provider: Final = OTLoggerProvider(resource=self._get_litellm_resource(self.config))
provider: Final = OTLoggerProvider(resource=self._litellm_resource())
log_exporter: Final = self._get_log_exporter()
provider.add_log_record_processor(BatchLogRecordProcessor(log_exporter))
return provider
@ -1144,9 +1161,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
owns_exporter: Final = _provider_owns_exporter(dynamic_config.exporter)
def _build() -> "_SDKTracerProvider":
provider: Final = TracerProvider(
resource=self._get_litellm_resource(self.config), shutdown_on_exit=owns_exporter
)
provider: Final = TracerProvider(resource=self._litellm_resource(), shutdown_on_exit=owns_exporter)
provider.add_span_processor(self._get_span_processor(config_override=dynamic_config))
return provider
@ -1162,9 +1177,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
owns_exporter: Final = _provider_owns_exporter(self.OTEL_EXPORTER)
def _build() -> "_SDKTracerProvider":
provider: Final = TracerProvider(
resource=self._get_litellm_resource(self.config), shutdown_on_exit=owns_exporter
)
provider: Final = TracerProvider(resource=self._litellm_resource(), shutdown_on_exit=owns_exporter)
provider.add_span_processor(self._get_span_processor(dynamic_headers=dynamic_headers))
return provider

View file

@ -13,7 +13,7 @@ import traceback
from collections.abc import Callable, Mapping, Sequence
from datetime import datetime as dt_object
from functools import lru_cache
from types import TracebackType
from types import MappingProxyType, TracebackType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from httpx import Response
@ -372,6 +372,35 @@ def _published_pricing(deployment_model: str | None) -> ModelInfo | None:
return None
def _resolve_vertex_location_for_cost(
custom_llm_provider: str | None,
litellm_params: Mapping[str, object] | None,
optional_params: Mapping[str, object] | None,
model: str,
) -> str | None:
"""
The Vertex AI location a request was served from, resolved the same way
dispatch resolves it, so regional deployments price with the
regional-endpoint uplift. None for non-Vertex providers.
Chat dispatch reads the location from request kwargs, which reach this
logging object through optional_params: on the proxy the logging object is
created before the router picks a deployment, so the deployment's location
never lands in litellm_params.
"""
if custom_llm_provider is None or not custom_llm_provider.startswith("vertex_ai"):
return None
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
empty: Final[Mapping[str, object]] = MappingProxyType({})
configured_location: Final = (
VertexBase.explicit_vertex_ai_location(optional_params or empty)
or VertexBase.explicit_vertex_ai_location(litellm_params or empty)
or VertexBase.safe_get_vertex_ai_location(empty)
)
return VertexBase.get_vertex_region(configured_location, model)
class Logging(LiteLLMLoggingBaseClass):
global \
supabaseClient, \
@ -1432,6 +1461,7 @@ class Logging(LiteLLMLoggingBaseClass):
reasoning_cost: float | None = None,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
) -> None:
"""
Helper method to store cost breakdown in the logging object.
@ -1450,6 +1480,7 @@ class Logging(LiteLLMLoggingBaseClass):
margin_total_amount: Total margin added in USD
service_tier: Tier the costs above were priced on, already resolved
data_residency: Region uplift the costs above were priced on, already resolved
vertex_location: Vertex AI location the costs above were priced on, already resolved
"""
self.cost_breakdown = CostBreakdown(
@ -1459,6 +1490,7 @@ class Logging(LiteLLMLoggingBaseClass):
tool_usage_cost=cost_for_built_in_tools_cost_usd_dollar,
service_tier=service_tier,
data_residency=data_residency,
vertex_location=vertex_location,
)
if cache_read_cost is not None and cache_read_cost > 0:
self.cost_breakdown["cache_read_cost"] = cache_read_cost
@ -1574,6 +1606,12 @@ class Logging(LiteLLMLoggingBaseClass):
if hasattr(self, "litellm_params") and self.litellm_params
else None
),
"vertex_location": _resolve_vertex_location_for_cost(
custom_llm_provider=self.model_call_details.get("custom_llm_provider", None),
litellm_params=(self.litellm_params if hasattr(self, "litellm_params") else None),
optional_params=self.optional_params,
model=litellm_model_name or self.model,
),
}
except Exception as e: # error creating kwargs for cost calculation
debug_info = StandardLoggingModelCostFailureDebugInformation(

View file

@ -757,6 +757,33 @@ def _get_regional_uplift_multiplier(model_info: ModelInfo, data_residency: str |
return 1.0
def get_vertex_regional_endpoint_uplift(model_info: ModelInfo, vertex_location: str | None) -> float:
"""
Resolve the per-model uplift multiplier for Vertex AI non-global (regional and
multi-region) endpoints.
Google prices every non-global endpoint at a flat premium over the global
endpoint (e.g. 1.10 = +10%) on all token types for the models that carry
regional pricing. The multiplier is stored on the model entry as
``regional_endpoint_uplift_multiplier``.
Returns 1.0 (no uplift) when ``vertex_location`` is ``None`` or ``"global"``,
or when the model has no multiplier configured.
"""
if vertex_location is None or vertex_location.lower() == "global":
return 1.0
multiplier: Final = model_info.get("regional_endpoint_uplift_multiplier")
if multiplier is None:
return 1.0
try:
return float(cast(float, multiplier))
except (TypeError, ValueError):
verbose_logger.exception(
"Invalid regional_endpoint_uplift_multiplier for model; defaulting to 1.0",
)
return 1.0
def get_provider_specific_geo_multiplier(model_info: ModelInfo, usage: Usage) -> float:
"""
Resolve the provider-specific regional pricing multiplier for the geo the
@ -798,6 +825,7 @@ def generic_cost_per_token(
service_tier: str | None = None,
data_residency: str | None = None,
model_info: ModelInfo | None = None,
vertex_location: str | None = None,
) -> tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -809,6 +837,9 @@ def generic_cost_per_token(
- usage: LiteLLM Usage block, containing anthropic caching information
- data_residency: optional OpenAI data-residency region (e.g. "eu", "us"),
used to apply the per-model regional-processing uplift multiplier.
- vertex_location: optional Vertex AI location the request was served from
(e.g. "us-east5", "global"), used to apply the per-model
regional-endpoint uplift multiplier when non-global.
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -968,6 +999,11 @@ def generic_cost_per_token(
prompt_cost *= uplift
completion_cost *= uplift
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
if vertex_uplift != 1.0:
prompt_cost *= vertex_uplift
completion_cost *= vertex_uplift
return prompt_cost, completion_cost
@ -988,6 +1024,7 @@ def get_token_type_cost_breakdown(
usage: Usage,
service_tier: str | None = None,
data_residency: str | None = None,
vertex_location: str | None = None,
) -> TokenTypeCostBreakdown:
"""
Provider-agnostic cost of reasoning and cache tokens, derived from the usage
@ -1069,6 +1106,12 @@ def get_token_type_cost_breakdown(
cache_read_cost *= uplift
cache_creation_cost *= uplift
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
if vertex_uplift != 1.0:
reasoning_cost *= vertex_uplift
cache_read_cost *= vertex_uplift
cache_creation_cost *= vertex_uplift
# Mirror the provider-specific geo uplift (e.g. Anthropic us: 1.1) the totals
# apply, so cache and reasoning line items stay reconciled with them.
geo_multiplier: Final = get_provider_specific_geo_multiplier(model_info=model_info, usage=usage)

View file

@ -723,7 +723,7 @@ class ChunkProcessor:
for choice in response.choices:
if (
hasattr(cast(Choices, choice).message, "reasoning_content")
and cast(Choices, choice).message.reasoning_content is not None
and cast(Choices, choice).message.reasoning_content
):
if reasoning_tokens is None:
reasoning_tokens = 0
@ -987,7 +987,12 @@ class ChunkProcessor:
returned_usage.completion_tokens_details is not None
and returned_usage.completion_tokens_details.reasoning_tokens is None
):
returned_usage.completion_tokens_details.reasoning_tokens = reasoning_tokens
capped_reasoning_tokens: Final = min(max(0, reasoning_tokens), returned_usage.completion_tokens)
returned_usage.completion_tokens_details.reasoning_tokens = capped_reasoning_tokens
if returned_usage.completion_tokens_details.text_tokens is None:
returned_usage.completion_tokens_details.text_tokens = (
returned_usage.completion_tokens - capped_reasoning_tokens
)
if prompt_tokens_details is not None:
returned_usage.prompt_tokens_details = prompt_tokens_details

View file

@ -5,6 +5,7 @@ from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, NoReturn, cast
import httpx
from pydantic import ValidationError
import litellm
from litellm.constants import (
@ -39,6 +40,7 @@ from litellm.types.llms.anthropic import (
AnthropicMessagesTool,
AnthropicMessagesToolChoice,
AnthropicOutputSchema,
AnthropicOutputTokensDetails,
AnthropicSystemMessageContent,
AnthropicThinkingParam,
AnthropicWebSearchTool,
@ -2104,6 +2106,68 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
compaction_blocks,
)
@staticmethod
def _thinking_tokens_from_usage(usage_object: Mapping[str, object]) -> int | None:
details: Final = usage_object.get("output_tokens_details")
if not isinstance(details, Mapping):
return None
try:
return AnthropicOutputTokensDetails.model_validate(details).thinking_tokens
except ValidationError:
return None
@staticmethod
def _response_has_thinking_block(completion_response: Mapping[str, object] | None) -> bool:
if completion_response is None:
return False
content: Final = completion_response.get("content")
if not isinstance(content, list):
return False
return any(
isinstance(block, Mapping) and block.get("type") in ("thinking", "redacted_thinking") for block in content
)
def _build_completion_token_details(
self,
usage_object: Mapping[str, object],
iterations: Sequence[object] | None,
completion_tokens: int,
reasoning_content: str | None,
completion_response: Mapping[str, object] | None,
) -> CompletionTokensDetailsWrapper:
iteration_thinking_tokens: Final = self._sum_iteration_thinking_tokens(iterations) if iterations else None
reported_thinking_tokens: Final = (
iteration_thinking_tokens
if iteration_thinking_tokens is not None
else self._thinking_tokens_from_usage(usage_object)
)
if reported_thinking_tokens is not None:
capped_reported: Final = min(max(0, reported_thinking_tokens), completion_tokens)
return CompletionTokensDetailsWrapper(
reasoning_tokens=capped_reported,
text_tokens=completion_tokens - capped_reported,
)
if reasoning_content:
estimated: Final = min(
token_counter(text=reasoning_content, count_response_tokens=True),
completion_tokens,
)
return CompletionTokensDetailsWrapper(
reasoning_tokens=max(0, estimated),
text_tokens=completion_tokens - max(0, estimated),
)
if self._response_has_thinking_block(completion_response):
return CompletionTokensDetailsWrapper(reasoning_tokens=None, text_tokens=None)
return CompletionTokensDetailsWrapper(reasoning_tokens=0, text_tokens=completion_tokens)
def _sum_iteration_thinking_tokens(self, iterations: Sequence[object]) -> int | None:
per_iteration: Final = tuple(
self._thinking_tokens_from_usage(iteration) if isinstance(iteration, Mapping) else None
for iteration in iterations
)
reported: Final = tuple(tokens for tokens in per_iteration if tokens is not None)
return sum(reported) if len(reported) == len(per_iteration) else None
@staticmethod
def is_anthropic_usage_object(usage_object: dict) -> bool:
"""Anthropic reports prompt cache tokens as top-level ``cache_read_input_tokens`` /
@ -2222,14 +2286,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
cache_creation_token_details=cache_creation_token_details,
text_tokens=raw_input_tokens,
)
# Always populate completion_token_details, not just when there's reasoning_content
estimated_reasoning_tokens: Final = (
token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0
)
reasoning_tokens: Final = min(estimated_reasoning_tokens, completion_tokens)
completion_token_details: Final = CompletionTokensDetailsWrapper(
reasoning_tokens=max(0, reasoning_tokens),
text_tokens=(completion_tokens - reasoning_tokens if reasoning_tokens > 0 else completion_tokens),
completion_token_details: Final = self._build_completion_token_details(
usage_object=_usage,
iterations=iterations,
completion_tokens=completion_tokens,
reasoning_content=reasoning_content,
completion_response=completion_response,
)
total_tokens: Final = prompt_tokens + completion_tokens

View file

@ -183,6 +183,29 @@ class BaseSearchConfig:
"""
return headers
def sign_request(
self,
headers: dict[str, str], # mutable-ok: matches the request header dict every other hook on this base takes
optional_params: dict[str, object], # mutable-ok: matches every other hook on this base
request_data: dict[str, object] | list[dict[str, object]], # mutable-ok: transform_search_request's body
api_base: str,
api_key: str | None = None,
) -> tuple[dict[str, str], bytes | None]: # mutable-ok: the handler passes these headers straight to httpx
"""
OPTIONAL
Sign the request. Providers like Bedrock AgentCore need to SigV4-sign
the request before sending it to the API.
For all other providers, this is a no-op and we just return the headers.
Returns:
Tuple of (headers, signed_json_body). When signed_json_body is not
None, the handler MUST send it verbatim as the request body —
re-serializing the payload would invalidate the signature.
"""
return headers, None
def get_complete_url(
self,
api_base: str | None,

View file

@ -1834,6 +1834,7 @@ class AmazonConverseConfig(BaseConfig):
self,
usage: ConverseTokenUsageBlock,
reasoning_content: str | None = None,
thinking_ran: bool = False,
) -> Usage:
input_tokens = usage["inputTokens"]
output_tokens: Final = usage["outputTokens"]
@ -1854,10 +1855,19 @@ class AmazonConverseConfig(BaseConfig):
cache_creation_tokens=cache_creation_input_tokens,
text_tokens=raw_input_tokens,
)
reasoning_tokens = token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0
completion_tokens_details: Final = CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens,
text_tokens=(output_tokens - reasoning_tokens if reasoning_tokens > 0 else output_tokens),
reasoning_tokens: Final = (
token_counter(text=reasoning_content, count_response_tokens=True) if reasoning_content else 0
)
completion_tokens_details: Final = (
CompletionTokensDetailsWrapper(
reasoning_tokens=reasoning_tokens,
text_tokens=output_tokens - reasoning_tokens,
)
if reasoning_tokens > 0
else CompletionTokensDetailsWrapper(
reasoning_tokens=None if thinking_ran else 0,
text_tokens=None if thinking_ran else output_tokens,
)
)
openai_usage: Final = Usage(
prompt_tokens=input_tokens,
@ -2254,6 +2264,7 @@ class AmazonConverseConfig(BaseConfig):
usage: Final = self.transform_usage(
completion_response["usage"],
reasoning_content=chat_completion_message.get("reasoning_content"),
thinking_ran=reasoningContentBlocks is not None,
)
## HANDLE TOOL CALLS

View file

@ -330,6 +330,7 @@ class AWSEventStreamDecoder:
self.response_id: str | None = None
self.json_mode = json_mode
self._current_tool_name: str | None = None
self._thinking_ran = False
def check_empty_tool_call_args(self) -> bool:
"""
@ -559,7 +560,12 @@ class AWSEventStreamDecoder:
elif "stopReason" in chunk_data:
finish_reason = map_finish_reason(chunk_data.get("stopReason", "stop"))
elif "usage" in chunk_data:
usage = converse_config.transform_usage(chunk_data.get("usage", {}))
usage = converse_config.transform_usage(
chunk_data.get("usage", {}),
thinking_ran=self._thinking_ran,
)
if thinking_blocks:
self._thinking_ran = True
model_response_provider_specific_fields: Final = {}
if "trace" in chunk_data:

View file

View file

@ -0,0 +1,455 @@
"""
Calls an Amazon Bedrock AgentCore Gateway web-search target (MCP protocol) to search the web.
Web Search on Amazon Bedrock AgentCore exposes Amazon's managed web index through
an AgentCore Gateway MCP endpoint.
AWS docs: https://docs.aws.amazon.com/bedrock-agentcore/latest/devguide/gateway-target-connector-web-search-tool.html
Authentication (matches the gateway's inbound authorizer type):
- AWS_IAM gateway: the request is SigV4-signed. Credentials come from explicit
params (aws_access_key_id / aws_secret_access_key / aws_session_token /
aws_region_name, also settable in a proxy search_tools entry) or the
standard AWS credential chain (env / profile / IRSA / assumed role)
- CUSTOM_JWT gateway: pass the OAuth2 bearer token (e.g. Cognito
client_credentials) as api_key, or set AGENTCORE_GATEWAY_TOKEN
Setup:
1. Create an AgentCore Gateway with a web-search connector target
2. Set AGENTCORE_GATEWAY_URL (or pass api_base) to the gateway MCP endpoint, e.g.
https://<gateway-id>.gateway.bedrock-agentcore.<region>.amazonaws.com/mcp
3. AWS_IAM: ensure the credentials allow bedrock-agentcore:InvokeGateway
CUSTOM_JWT: set AGENTCORE_GATEWAY_TOKEN (or pass api_key)
Usage:
response = litellm.search(
query="latest AI developments",
search_provider="agentcore",
max_results=5,
aws_access_key_id="...", # optional, omit to use the default chain
aws_secret_access_key="...",
)
"""
import json
import re
from collections.abc import Iterator, Mapping, Sequence
from typing import Final
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.search.transformation import (
BaseSearchConfig,
SearchResponse,
SearchResult,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.secret_managers.main import get_secret_str
# AgentCore web-search rejects queries longer than 200 characters
AGENTCORE_MAX_QUERY_LENGTH: Final = 200
# The provider contract documents a default of 10 results, send it explicitly
# so the gateway can't silently apply a different default.
AGENTCORE_DEFAULT_MAX_RESULTS: Final = 10
# Default MCP tool name for a gateway web-search connector target:
# "<target-name>___<tool-name>". Override with AGENTCORE_SEARCH_TOOL_NAME
# or optional_params["tool_name"] when the target uses a custom name.
AGENTCORE_DEFAULT_TOOL_NAME: Final = "web-search-tool___WebSearch"
# All web-search connector tools share this suffix; rejecting other names keeps
# a caller-supplied tool_name from invoking unrelated tools on the same gateway
# with the proxy's credentials.
AGENTCORE_TOOL_NAME_SUFFIX: Final = "___WebSearch"
# MCP revision this provider speaks. Sent on every request because the gateway is
# called statelessly, without an initialize handshake to negotiate a version.
# AgentCore gateways whose protocolConfiguration leaves supportedVersions unset
# accept only 2025-03-26 and reject anything newer with a -32600 error, so that
# is the default; a gateway pinned to another version needs
# AGENTCORE_MCP_PROTOCOL_VERSION set to match.
AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION: Final = "2025-03-26"
# Matched against the URL host so a crafted path or query string can't pass for
# a gateway hostname.
_GATEWAY_HOST_PATTERN: Final = re.compile(r"[a-z0-9-]+\.gateway\.bedrock-agentcore\.([a-z0-9-]+)\.amazonaws\.com")
_SSE_EVENT_SEPARATOR: Final = re.compile(r"\r?\n[ \t]*\r?\n")
_SSE_LINE_PREFIXES: Final = ("event:", "data:", ":", "id:", "retry:")
def _gateway_host_match(api_base: str) -> re.Match[str] | None:
return _GATEWAY_HOST_PATTERN.fullmatch(httpx.URL(api_base).host)
_LOOPBACK_HOSTS: Final = frozenset({"localhost", "127.0.0.1", "::1"})
def _credential_safe_transport(api_base: str) -> bool:
url: Final = httpx.URL(api_base)
return url.scheme == "https" or url.host in _LOOPBACK_HOSTS
def _string_field(item: Mapping[str, object], *keys: str) -> str | None:
return next(
(value for key in keys if isinstance(value := item.get(key), str) and value),
None,
)
def _to_search_result(item: Mapping[str, object]) -> SearchResult:
return SearchResult(
title=_string_field(item, "title") or "",
url=_string_field(item, "url") or "",
snippet=_string_field(item, "text", "snippet") or "",
date=_string_field(item, "publishedDate", "date"),
last_updated=None,
)
def _result_items(parsed: object) -> tuple[Mapping[str, object], ...]:
items: Final = parsed.get("results", ()) if isinstance(parsed, Mapping) else parsed
if not isinstance(items, Sequence) or isinstance(items, (str, bytes)):
return ()
return tuple(item for item in items if isinstance(item, Mapping))
def _parse_result_items(raw_text: object) -> tuple[Mapping[str, object], ...]:
"""
Parse one MCP text block into the search result objects it carries.
A block holds either a JSON list of results or a {"results": [...]} object;
anything unparseable is skipped rather than failing the whole response.
"""
if not isinstance(raw_text, str):
return ()
try:
parsed: Final = json.loads(raw_text)
except json.JSONDecodeError:
return ()
return _result_items(parsed)
def _iter_sse_events(text: str) -> Iterator[Mapping[str, object]]:
"""
Yield the JSON payload of each SSE event in a Streamable HTTP MCP response.
Per the SSE spec an event's data is the concatenation of all its ``data:``
lines (joined with newlines), and a stream may carry several events, e.g.
progress notifications before the JSON-RPC response.
"""
for chunk in _SSE_EVENT_SEPARATOR.split(text):
payload = "\n".join(line[len("data:") :].lstrip() for line in chunk.splitlines() if line.startswith("data:"))
if not payload:
continue
try:
parsed = json.loads(payload)
except json.JSONDecodeError:
continue
if isinstance(parsed, dict):
yield parsed
class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM):
def __init__(self) -> None:
BaseSearchConfig.__init__(self)
BaseAWSLLM.__init__(self)
@staticmethod
def ui_friendly_name() -> str:
return "Web Search on Amazon Bedrock"
def validate_environment(
self,
headers: dict, # mutable-ok: BaseSearchConfig hands providers the mutable request header dict
api_key: str | None = None,
api_base: str | None = None,
**kwargs: object, # kwargs-ok: BaseSearchConfig.validate_environment forwards provider-specific extras
) -> dict: # mutable-ok: the handler passes these headers straight to httpx, which wants a dict
"""
Set MCP transport headers. Per the MCP Streamable HTTP transport spec,
the client MUST accept both application/json and text/event-stream, and
declare its protocol revision with MCP-Protocol-Version.
Authentication itself happens in sign_request(): bearer token for
CUSTOM_JWT gateways, AWS SigV4 for AWS_IAM gateways.
"""
return { # mutable-ok: httpx request headers are a dict
**headers,
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
"MCP-Protocol-Version": get_secret_str("AGENTCORE_MCP_PROTOCOL_VERSION")
or AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION,
}
def get_complete_url(
self,
api_base: str | None,
optional_params: dict, # mutable-ok: BaseSearchConfig passes optional params as a dict
data: dict | list[dict] | None = None, # mutable-ok: BaseSearchConfig request bodies are JSON dicts
**kwargs: object, # kwargs-ok: BaseSearchConfig.get_complete_url forwards provider-specific extras
) -> str:
gateway_url: Final = api_base or get_secret_str("AGENTCORE_GATEWAY_URL")
if not gateway_url:
raise ValueError(
"AGENTCORE_GATEWAY_URL is not set. Set it to your AgentCore Gateway MCP "
"endpoint (https://<gateway-id>.gateway.bedrock-agentcore.<region>"
".amazonaws.com/mcp) or pass api_base."
)
return gateway_url
def transform_search_request(
self,
query: str | list[str], # mutable-ok: BaseSearchConfig accepts a list of queries
optional_params: dict, # mutable-ok: BaseSearchConfig passes optional params as a dict
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_request forwards provider-specific extras
) -> dict: # mutable-ok: the JSON-RPC body is serialized as a JSON object
"""
Transform Search request to an MCP tools/call request.
Args:
query: Search query (string or list of strings). AgentCore only
supports single string queries; lists are joined with spaces.
optional_params: Optional parameters for the request
- max_results: Maximum number of results (1-25), default 10
- tool_name: Override the MCP tool name of the gateway target
Returns:
Dict with the JSON-RPC 2.0 request body
"""
joined_query: Final = " ".join(query) if isinstance(query, list) else query
tool_name: Final = (
optional_params.get("tool_name")
or get_secret_str("AGENTCORE_SEARCH_TOOL_NAME")
or AGENTCORE_DEFAULT_TOOL_NAME
)
if not tool_name.endswith(AGENTCORE_TOOL_NAME_SUFFIX):
raise ValueError(
f"Invalid AgentCore search tool_name '{tool_name}': must end with "
f"'{AGENTCORE_TOOL_NAME_SUFFIX}' (a web-search connector tool). "
"Other gateway tools cannot be invoked through this provider."
)
return { # mutable-ok: JSON-RPC request bodies are JSON objects
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": { # mutable-ok: JSON-RPC request bodies are JSON objects
"name": tool_name,
"arguments": { # mutable-ok: JSON-RPC request bodies are JSON objects
"query": joined_query[:AGENTCORE_MAX_QUERY_LENGTH],
"maxResults": optional_params.get("max_results", AGENTCORE_DEFAULT_MAX_RESULTS),
},
},
}
def sign_request(
self,
headers: dict[str, str], # mutable-ok: BaseSearchConfig hands providers the mutable request header dict
optional_params: dict[str, object], # mutable-ok: BaseSearchConfig passes optional params as a dict
request_data: dict[str, object] | list[dict[str, object]], # mutable-ok: request bodies are JSON dicts
api_base: str,
api_key: str | None = None,
) -> tuple[dict[str, str], bytes | None]: # mutable-ok: BaseSearchConfig.sign_request returns httpx headers
"""
Authenticate the MCP request.
CUSTOM_JWT gateways: attach the caller's OAuth2 bearer token (api_key
or AGENTCORE_GATEWAY_TOKEN), no AWS credentials involved.
AWS_IAM gateways: SigV4-sign with the bedrock-agentcore service name.
"""
if not isinstance(request_data, dict):
raise TypeError("AgentCore search expects a single dict request body")
if not _credential_safe_transport(api_base):
raise ValueError(
f"Refusing to send AgentCore credentials over plaintext HTTP to '{api_base}': a bearer "
"token or SigV4 signature would be readable in transit. Use an https gateway URL "
"(plain http is allowed only for localhost)."
)
# Server-managed credentials only go to a trusted host, otherwise an
# authenticated caller could point api_base at their own server (e.g. via
# /search_tools/test_connection) and collect AGENTCORE_GATEWAY_TOKEN or a
# SigV4 signature with the proxy's credential scope and session token.
gateway_host_match: Final = _gateway_host_match(api_base)
bearer_token: Final = self.resolve_server_api_key(
caller_api_key=api_key,
caller_api_base=api_base,
key_env_vars=("AGENTCORE_GATEWAY_TOKEN",),
base_env_var="AGENTCORE_GATEWAY_URL",
default_api_base=api_base if gateway_host_match else None,
)
if bearer_token:
bearer_headers: Final = { # mutable-ok: httpx request headers are a dict
**headers,
"Authorization": f"Bearer {bearer_token}",
}
return bearer_headers, json.dumps(request_data).encode()
if gateway_host_match is None and not self._is_configured_gateway(api_base):
raise ValueError(
f"Refusing to send SigV4-signed AgentCore requests to '{api_base}': it is neither an "
"AgentCore gateway hostname nor the host in AGENTCORE_GATEWAY_URL. Set "
"AGENTCORE_GATEWAY_URL to authorize a custom gateway hostname."
)
signing_params: Final = (
optional_params
if optional_params.get("aws_region_name") is not None
else { # mutable-ok: BaseAWSLLM._sign_request takes optional params as a dict
**optional_params,
"aws_region_name": self._signing_region(api_base),
}
)
# api_key="" (not None, but falsy) disables BaseAWSLLM's fallback to the
# AWS_BEARER_TOKEN_BEDROCK env var: that token is a Bedrock Runtime
# credential and must not be sent to an AgentCore gateway.
return self._sign_request(
service_name="bedrock-agentcore",
headers=headers,
optional_params=signing_params,
request_data=request_data,
api_base=api_base,
api_key="",
)
@staticmethod
def _is_configured_gateway(api_base: str) -> bool:
configured: Final = get_secret_str("AGENTCORE_GATEWAY_URL")
if not configured:
return False
return httpx.URL(configured).host == httpx.URL(api_base).host
@staticmethod
def _signing_region(api_base: str) -> str:
"""
Resolve the SigV4 signing region, which must match the gateway's region.
Standard gateway hostnames carry it, so callers don't have to set
aws_region_name to a region different from their default. For custom or
private hostnames, defer to the AWS configuration chain (env vars and
the shared config / profile region), and error out when that yields
nothing rather than silently signing for a guessed region the gateway
would reject with a confusing auth error.
"""
match: Final = _gateway_host_match(api_base)
if match:
return match.group(1)
# boto3's session resolution covers env vars AND the AWS shared config
# (profile region), unlike BaseAWSLLM's helper, which silently defaults
# to us-west-2 when nothing is configured.
import boto3
configured_region: Final = boto3.Session().region_name
if configured_region:
return configured_region
raise ValueError(
f"Cannot derive the SigV4 signing region from api_base '{api_base}' "
"or the AWS configuration chain. Set aws_region_name (or AWS_DEFAULT_REGION / "
"a profile region) to the gateway's region when using a custom hostname."
)
def transform_search_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_response forwards provider-specific extras
) -> SearchResponse:
"""
Transform an MCP tools/call response to LiteLLM unified SearchResponse.
The gateway returns JSON-RPC (as plain JSON or a single-message SSE
stream) whose result.content[] text blocks contain a JSON list of
{title, url, date/publishedDate, text} entries. Web-search connector
1.1.0 and later repeat that list in result.structuredContent, which is
the only machine-readable copy when the text block holds prose instead.
"""
response_json: Final = self._parse_mcp_body(raw_response)
error: Final = response_json.get("error")
if error is not None:
raise BedrockError(
status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
message=f"AgentCore gateway MCP error: {error}",
)
# A failed tools/call is reported in-band, as HTTP 200 with result.isError
# and the failure text where the results would be.
result: Final = response_json.get("result")
if isinstance(result, dict) and result.get("isError"):
raise BedrockError(
status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
message=f"AgentCore web search tool error: {self._tool_error_message(response_json)}",
)
text_items: Final = tuple(
item for block in self._text_blocks(response_json) for item in _parse_result_items(block.get("text"))
)
structured: Final = result.get("structuredContent") if isinstance(result, Mapping) else None
items: Final = text_items or _result_items(structured)
results: Final = [_to_search_result(item) for item in items] # mutable-ok: pydantic list field
return SearchResponse(results=results, object="search")
def _tool_error_message(self, response_json: Mapping[str, object]) -> str:
texts: Final = tuple(
text for block in self._text_blocks(response_json) if isinstance(text := block.get("text"), str)
)
return " ".join(texts) if texts else json.dumps(response_json.get("result"))[:500]
@staticmethod
def _text_blocks(response_json: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
result: Final = response_json.get("result")
content: Final = result.get("content") if isinstance(result, dict) else None
if not isinstance(content, Sequence) or isinstance(content, (str, bytes)):
return ()
return tuple(block for block in content if isinstance(block, dict) and block.get("type") == "text")
@staticmethod
def _parse_mcp_body(raw_response: httpx.Response) -> Mapping[str, object]:
"""
Parse a JSON or SSE-framed (Streamable HTTP transport) MCP response.
Return the event whose payload carries the JSON-RPC response, i.e. one
containing ``result`` or ``error``, falling back to the last event when
the stream carries only notifications.
"""
text: Final = raw_response.text
if not text.lstrip().startswith(_SSE_LINE_PREFIXES):
return raw_response.json()
events: Final = tuple(_iter_sse_events(text))
response_event: Final = next(
(event for event in events if "result" in event or "error" in event),
None,
)
if response_event is not None:
return response_event
if events:
return events[-1]
raise BedrockError(
status_code=502,
message=f"AgentCore gateway returned SSE without a JSON data frame: {text[:200]}",
)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict, # mutable-ok: BaseSearchConfig.get_error_class takes the response headers as a dict
) -> Exception:
return BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -5,7 +5,7 @@ import ssl
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from contextlib import asynccontextmanager
from functools import lru_cache
from types import ModuleType
from types import MappingProxyType, ModuleType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
@ -1788,6 +1788,14 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
headers=headers,
optional_params=optional_params,
request_data=data,
api_base=complete_url,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
input=query if isinstance(query, str) else str(query),
@ -1811,14 +1819,15 @@ class BaseLLMHTTPHandler:
# Note: timeout is set on the client itself, not per-request for GET
response = client.get(
url=complete_url,
headers=headers,
headers=signed_headers,
)
else:
# Make POST request with JSON data
# A signed body must be sent verbatim, re-serializing it would break the signature
response = client.post(
url=complete_url,
headers=headers,
json=data,
headers=signed_headers,
data=signed_json_body,
json=data if signed_json_body is None else None,
timeout=timeout,
)
except Exception as e:
@ -1872,6 +1881,14 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
headers=headers,
optional_params=optional_params,
request_data=data,
api_base=complete_url,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
input=query if isinstance(query, str) else str(query),
@ -1900,14 +1917,15 @@ class BaseLLMHTTPHandler:
# Note: timeout is set on the client itself, not per-request for GET
response = await async_httpx_client.get(
url=complete_url,
headers=headers,
headers=signed_headers,
)
else:
# Make async POST request with JSON data
# A signed body must be sent verbatim, re-serializing it would break the signature
response = await async_httpx_client.post(
url=complete_url,
headers=headers,
json=data,
headers=signed_headers,
data=signed_json_body,
json=data if signed_json_body is None else None,
timeout=timeout,
)
except Exception as e:
@ -2069,6 +2087,14 @@ class BaseLLMHTTPHandler:
if anthropic_messages_provider_config.should_filter_anthropic_beta_headers():
headers = update_headers_with_filtered_beta(headers=headers, provider=custom_llm_provider)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
explicit_vertex_location: Final = VertexBase.explicit_vertex_ai_location(MappingProxyType(dict(litellm_params)))
vertex_location_params: Final = (
MappingProxyType({"vertex_location": explicit_vertex_location})
if explicit_vertex_location
else MappingProxyType({})
)
logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
@ -2077,6 +2103,7 @@ class BaseLLMHTTPHandler:
"preset_cache_key": None,
"stream_response": {},
"model_info": kwargs.get("model_info"),
**vertex_location_params,
**anthropic_messages_optional_request_params,
},
custom_llm_provider=custom_llm_provider,

View file

@ -7,6 +7,7 @@ from litellm import verbose_logger
from litellm.litellm_core_utils.llm_cost_calc.utils import (
_is_above_128k,
generic_cost_per_token,
get_vertex_regional_endpoint_uplift,
)
from litellm.types.utils import ModelInfo, Usage
@ -63,6 +64,7 @@ def cost_per_character(
usage: Usage,
prompt_characters: float | None = None,
completion_characters: float | None = None,
vertex_location: str | None = None,
) -> tuple[float, float]:
"""
Calculates the cost per character for a given VertexAI model, input messages, and response object.
@ -72,6 +74,8 @@ def cost_per_character(
- custom_llm_provider: str, "vertex_ai-*"
- prompt_characters: float, the number of input characters
- completion_characters: float, the number of output characters
- vertex_location: the Vertex AI location serving the request; non-global
locations apply the model's regional-endpoint uplift multiplier
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -79,8 +83,6 @@ def cost_per_character(
Raises:
Exception if model requires >128k pricing, but model cost not mapped
"""
model_info = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
## GET MODEL INFO
model_info = litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
@ -162,7 +164,8 @@ def cost_per_character(
usage=usage,
)
return prompt_cost, completion_cost
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
return prompt_cost * vertex_uplift, completion_cost * vertex_uplift
def _handle_128k_pricing(
@ -196,6 +199,7 @@ def cost_per_token(
custom_llm_provider: str,
usage: Usage,
service_tier: str | None = None,
vertex_location: str | None = None,
) -> tuple[float, float]:
"""
Calculates the cost per token for a given model, prompt tokens, and completion tokens.
@ -207,6 +211,8 @@ def cost_per_token(
- completion_tokens: float, the number of output tokens
- service_tier: optional tier derived from Gemini trafficType
("priority" for ON_DEMAND_PRIORITY, "flex" for FLEX/batch).
- vertex_location: the Vertex AI location serving the request; non-global
locations apply the model's regional-endpoint uplift multiplier
Returns:
Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd
@ -222,14 +228,17 @@ def cost_per_token(
input_cost_per_token_above_128k_tokens: Final = model_info.get("input_cost_per_token_above_128k_tokens")
output_cost_per_token_above_128k_tokens: Final = model_info.get("output_cost_per_token_above_128k_tokens")
if input_cost_per_token_above_128k_tokens is not None or output_cost_per_token_above_128k_tokens is not None:
return _handle_128k_pricing(
prompt_cost_128k, completion_cost_128k = _handle_128k_pricing(
model_info=model_info,
usage=usage,
)
vertex_uplift: Final = get_vertex_regional_endpoint_uplift(model_info, vertex_location)
return prompt_cost_128k * vertex_uplift, completion_cost_128k * vertex_uplift
return generic_cost_per_token(
model=model,
custom_llm_provider=custom_llm_provider,
usage=usage,
service_tier=service_tier,
vertex_location=vertex_location,
)

View file

@ -8,6 +8,7 @@ import asyncio
import json
import os
import threading
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final, Literal
from urllib.parse import urlparse
@ -68,7 +69,8 @@ class VertexBase:
# re-acquire it without deadlocking the current thread.
self._sync_refresh_lock = threading.RLock()
def get_vertex_region(self, vertex_region: str | None, model: str) -> str:
@staticmethod
def get_vertex_region(vertex_region: str | None, model: str) -> str:
import litellm
# Try to get supported_regions directly from model_cost
@ -1191,7 +1193,18 @@ class VertexBase:
)
@staticmethod
def safe_get_vertex_ai_location(litellm_params: dict) -> str | None:
def explicit_vertex_ai_location(params: Mapping[str, object]) -> str | None:
"""
The location explicitly configured in the given params, without any
module-level or environment fallback. None when not configured.
"""
for configured in (params.get("vertex_location"), params.get("vertex_ai_location")):
if isinstance(configured, str) and configured:
return configured
return None
@staticmethod
def safe_get_vertex_ai_location(litellm_params: Mapping[str, object]) -> str | None:
"""
Safely get Vertex AI location without mutating the litellm_params dict.
@ -1205,8 +1218,7 @@ class VertexBase:
Vertex AI location/region or None
"""
return (
litellm_params.get("vertex_location")
or litellm_params.get("vertex_ai_location")
VertexBase.explicit_vertex_ai_location(litellm_params)
or litellm.vertex_location
or get_secret_str("VERTEXAI_LOCATION")
or get_secret_str("VERTEX_LOCATION")

View file

@ -16792,6 +16792,14 @@
"notes": "APISerpent deep search (/api/search), multi-engine (Google, Bing, Yahoo, DuckDuckGo). Pricing: $0.60/1k searches."
}
},
"agentcore/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "agentcore",
"mode": "search",
"metadata": {
"notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway"
}
},
"tinyfish/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "tinyfish",
@ -19527,6 +19535,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -19584,6 +19593,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-preview-customtools": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -19738,6 +19748,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.5-flash": {
"prompt_cache_min_tokens": 4096,
"deprecation_date": "2027-05-19",
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 1.5e-06,
@ -19749,6 +19760,7 @@
"mode": "chat",
"output_cost_per_reasoning_token": 9e-06,
"output_cost_per_token": 9e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19790,6 +19802,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -19804,6 +19817,7 @@
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19844,6 +19858,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -19858,6 +19873,7 @@
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19898,6 +19914,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -19955,6 +19972,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.1-pro-preview-customtools": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -20689,8 +20707,8 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.1-flash-image": {
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_batches": 1.25e-07,
"input_cost_per_token": 5e-07,
"input_cost_per_token_batches": 2.5e-07,
"litellm_provider": "gemini",
"max_input_tokens": 65536,
"max_output_tokens": 32768,
@ -20698,8 +20716,8 @@
"mode": "image_generation",
"output_cost_per_image": 0.045,
"output_cost_per_image_token": 6e-05,
"output_cost_per_token": 1.5e-06,
"output_cost_per_token_batches": 7.5e-07,
"output_cost_per_token": 3e-06,
"output_cost_per_token_batches": 1.5e-06,
"rpm": 1000,
"tpm": 4000000,
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image",
@ -20732,8 +20750,8 @@
},
"gemini/gemini-3.1-flash-image-preview": {
"deprecation_date": "2026-06-25",
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_batches": 1.25e-07,
"input_cost_per_token": 5e-07,
"input_cost_per_token_batches": 2.5e-07,
"litellm_provider": "gemini",
"max_input_tokens": 65536,
"max_output_tokens": 32768,
@ -20741,8 +20759,8 @@
"mode": "image_generation",
"output_cost_per_image": 0.045,
"output_cost_per_image_token": 6e-05,
"output_cost_per_token": 1.5e-06,
"output_cost_per_token_batches": 7.5e-07,
"output_cost_per_token": 3e-06,
"output_cost_per_token_batches": 1.5e-06,
"rpm": 1000,
"tpm": 4000000,
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image-preview",
@ -21464,6 +21482,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.5-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 1.5e-06,
@ -21518,6 +21537,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -21575,6 +21595,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -21665,6 +21686,7 @@
"tpm": 800000
},
"gemini/gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 2e-06,
@ -21722,6 +21744,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.1-pro-preview-customtools": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 2e-06,
@ -21860,6 +21883,7 @@
"supports_vision": true
},
"gemini-3.5-flash": {
"prompt_cache_min_tokens": 4096,
"deprecation_date": "2027-05-19",
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_audio_token": 1e-06,
@ -21913,6 +21937,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -21968,6 +21993,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -38857,6 +38883,7 @@
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -38881,6 +38908,7 @@
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -39100,6 +39128,7 @@
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"regional_endpoint_uplift_multiplier": 1.1,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@ -39129,6 +39158,7 @@
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"regional_endpoint_uplift_multiplier": 1.1,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@ -39149,6 +39179,7 @@
},
"vertex_ai/claude-opus-4-6": {
"deprecation_date": "2027-02-05",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39180,6 +39211,7 @@
},
"vertex_ai/claude-opus-4-6@default": {
"deprecation_date": "2027-02-05",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39211,6 +39243,7 @@
},
"vertex_ai/claude-opus-4-7": {
"deprecation_date": "2027-04-16",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39243,6 +39276,7 @@
},
"vertex_ai/claude-opus-4-7@default": {
"deprecation_date": "2027-04-16",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39275,6 +39309,7 @@
},
"vertex_ai/claude-fable-5": {
"deprecation_date": "2027-06-08",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
@ -39307,6 +39342,7 @@
},
"vertex_ai/claude-fable-5@default": {
"deprecation_date": "2027-06-08",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
@ -39339,6 +39375,7 @@
},
"vertex_ai/claude-opus-5": {
"deprecation_date": "2027-01-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39372,6 +39409,7 @@
},
"vertex_ai/claude-opus-5@default": {
"deprecation_date": "2027-01-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39405,6 +39443,7 @@
},
"vertex_ai/claude-opus-4-8": {
"deprecation_date": "2027-05-28",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39438,6 +39477,7 @@
},
"vertex_ai/claude-opus-4-8@default": {
"deprecation_date": "2027-05-28",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39487,6 +39527,7 @@
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -39500,6 +39541,7 @@
},
"vertex_ai/claude-sonnet-5": {
"deprecation_date": "2026-12-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_1hr": 4e-06,
@ -39532,6 +39574,7 @@
"prompt_cache_min_tokens": 1024
},
"vertex_ai/claude-sonnet-4-6": {
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -39579,6 +39622,7 @@
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -39992,6 +40036,7 @@
"output_cost_per_token_batches": 7.5e-07,
"output_cost_per_token_flex": 7.5e-07,
"output_cost_per_token_priority": 2.7e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@ -40048,6 +40093,7 @@
"output_cost_per_token_batches": 1.25e-06,
"output_cost_per_token_flex": 1.25e-06,
"output_cost_per_token_priority": 4.5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@ -47219,6 +47265,7 @@
},
"vertex_ai/claude-sonnet-5@default": {
"deprecation_date": "2026-12-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_1hr": 4e-06,
@ -47251,6 +47298,7 @@
"prompt_cache_min_tokens": 1024
},
"vertex_ai/claude-sonnet-4-6@default": {
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -48180,15 +48228,15 @@
},
"deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"cache_read_input_token_cost": 1.4e-08,
"input_cost_per_token": 4.4e-07,
"input_cost_per_token_cache_hit": 1.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_cost_per_token": 1.32e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
@ -48206,15 +48254,15 @@
},
"deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"cache_read_input_token_cost": 4.4e-08,
"input_cost_per_token": 1.32e-06,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"output_cost_per_token": 3.96e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
@ -48232,15 +48280,15 @@
},
"deepseek/deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"cache_read_input_token_cost": 1.4e-08,
"input_cost_per_token": 4.4e-07,
"input_cost_per_token_cache_hit": 1.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_cost_per_token": 1.32e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
@ -48258,15 +48306,15 @@
},
"deepseek/deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"cache_read_input_token_cost": 4.4e-08,
"input_cost_per_token": 1.32e-06,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"output_cost_per_token": 3.96e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"

View file

@ -2439,6 +2439,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="max request size in MB, if a request is larger than this size it will be rejected",
)
max_batch_file_size_mb: int | None = Field(
None,
description="max batch input file size in MB for /v1/files uploads with purpose=batch, if a file is larger than this size it will be rejected before being forwarded to the provider",
)
max_response_size_mb: int | None = Field(
None,
description="max response size in MB, if a response is larger than this size it will be rejected",

View file

@ -521,6 +521,20 @@ This is a one-time file patch and restore, not a live traffic interceptor. A Cla
Cursor is not supported: it has no equivalent file-based config to hot-patch this way, since its model routing lives in its own app storage and is configured through its GUI.
#### Making It Permanent at Login
`lite up` holds the patch only for as long as it runs. To wire Claude Code up once and leave it that way, pass `--config-claude` to `lite login`:
```bash
lite --base-url https://your-proxy.example.com login --config-claude
```
It writes the same two settings `lite up` does, `env.ANTHROPIC_BASE_URL` and `apiKeyHelper`, but persistently: there is no backup, nothing to restore, and no foreground process to keep alive. Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag.
Because the credential is reached through `apiKeyHelper` rather than copied into the file, a later `lite login` refreshes it with no further action: Claude Code re-runs the helper on every request and picks up whatever token the most recent login stored. Nothing secret is written to `settings.json`.
Run it again to point Claude Code at a different proxy; the base URL and the helper are both rewritten. `lite up` and `--config-claude` manage the same file, so the flag refuses to run while a `lite up` session holds a backup, and tells you to run `lite down` first, rather than writing settings that `lite up` would silently revert when it stops.
### QA Complexity-Based Auto-Routing Against Your Real Proxy
`lite autoroute` lets you try LiteLLM's complexity-based auto-routing -- picking a cheaper or more expensive model depending on how complex a prompt looks -- against models your key already has access to on your real, running proxy, without editing that proxy's `config.yaml` and without any real request ever bypassing it. It builds a second, throwaway proxy locally that forwards every request back to your real proxy, and points Claude Code at that local proxy for the duration of the session.

View file

@ -16,6 +16,12 @@ from typing_extensions import NotRequired, TypedDict
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
from .claude_settings import (
CLAUDE_SETTINGS_PATH,
SETTINGS_FILE_OWNERS,
ClaudeSettingsError,
write_claude_settings,
)
from .private_json import write_private_json
@ -629,9 +635,28 @@ def _render_and_prompt_for_team_selection(teams: list[CliTeam]) -> str | None:
return None
def _configure_claude_code(base_url: str) -> None:
"""Point Claude Code at base_url by patching ~/.claude/settings.json."""
try:
write_claude_settings(base_url, CLAUDE_SETTINGS_PATH, SETTINGS_FILE_OWNERS)
except ClaudeSettingsError as e:
raise click.ClickException(f"Logged in, but could not configure Claude Code: {e}")
click.echo(f"\nConfigured Claude Code: {CLAUDE_SETTINGS_PATH} now routes through {base_url.rstrip('/')}.")
click.echo("Your other Claude Code settings were left untouched. Restart Claude Code to pick this up.")
@click.command(name="login")
@click.option(
"--config-claude",
is_flag=True,
default=False,
help=(
"After logging in, update ~/.claude/settings.json so Claude Code routes through this proxy. "
"Unrelated settings are preserved."
),
)
@click.pass_context
def login(ctx: click.Context):
def login(ctx: click.Context, config_claude: bool):
"""Login to LiteLLM proxy using SSO authentication"""
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.client.cli.interface import show_commands
@ -683,6 +708,9 @@ def login(ctx: click.Context):
click.echo(f"JWT Token: {api_key[:20]}...")
click.echo("You can now use the CLI without specifying --api-key")
if config_claude:
_configure_claude_code(base_url)
# Show available commands after successful login
click.echo("\n" + "=" * 60)
show_commands()
@ -698,6 +726,10 @@ def login(ctx: click.Context):
except KeyboardInterrupt:
click.echo("\nAuthentication cancelled by user.")
return
except click.ClickException:
# Login itself already succeeded; only the post-login step failed, so this
# must not be relabelled as an authentication failure by the handler below.
raise
except Exception as e:
click.echo(f"Authentication failed: {e}")
return

View file

@ -10,11 +10,16 @@ import click
import yaml
from pydantic import JsonValue, TypeAdapter, ValidationError
from ..up import CLAUDE_SETTINGS_PATH, UpError, load_json_or_empty, restore_claude_settings, write_backup
from ..claude_settings import (
AUTOROUTE_BACKUP_PATH,
CLAUDE_SETTINGS_PATH,
ClaudeSettingsError,
load_json_or_empty,
)
from ..up import BackupRecord as ClaudeBackupRecord
from ..up import restore_claude_settings, write_backup
from .config import master_key_from_config
from .process import (
AUTOROUTE_DIR,
CONFIG_PATH,
DEFAULT_AUTOROUTE_PORT,
LOG_PATH,
@ -35,8 +40,6 @@ from .process import (
from .settings import merge_claude_settings_static_token
from .wizard import run_configure_wizard
AUTOROUTE_BACKUP_PATH: Final = AUTOROUTE_DIR / "claude_settings_backup.json"
_GENERATED_CONFIG_ADAPTER: Final = TypeAdapter(dict[str, JsonValue])
@ -108,7 +111,7 @@ def up(port: int) -> None:
try:
existing_pid: Final = read_pid_record()
except UpError as e:
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
if existing_pid is not None and is_running(existing_pid.pid):
raise click.ClickException(
@ -157,7 +160,7 @@ def up(port: int) -> None:
CLAUDE_SETTINGS_PATH.parent.mkdir(parents=True, exist_ok=True)
with secure_create(CLAUDE_SETTINGS_PATH) as f:
json.dump(merged, f, indent=2)
except UpError as e:
except ClaudeSettingsError as e:
terminate(process.pid)
clear_pid_record()
raise click.ClickException(str(e))
@ -175,7 +178,7 @@ def up(port: int) -> None:
clear_pid_record()
try:
restore_claude_settings(CLAUDE_SETTINGS_PATH, AUTOROUTE_BACKUP_PATH)
except UpError as e:
except ClaudeSettingsError as e:
# Runs from atexit/a signal handler too, outside Click's own exception
# handling -- raising here would only produce an unhandled-exception
# warning on stderr, not a clean message.
@ -207,7 +210,7 @@ def down() -> None:
"""Restore Claude Code settings and stop a leftover ephemeral proxy, if any"""
try:
record: PidRecord | None = read_pid_record()
except UpError as e:
except ClaudeSettingsError as e:
# down is the crash-recovery path -- a corrupt pid record must not block it; clear the
# unusable record and keep going rather than leaving the user with no way to clean up.
click.echo(f"{e} Clearing it and continuing cleanup.", err=True)
@ -219,7 +222,7 @@ def down() -> None:
try:
restored: Final = restore_claude_settings(CLAUDE_SETTINGS_PATH, AUTOROUTE_BACKUP_PATH)
except UpError as e:
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
if restored is None:
click.echo("Nothing to restore.")

View file

@ -0,0 +1,155 @@
"""Shared handling of Claude Code's ~/.claude/settings.json.
`lite up` patches this file temporarily and restores it on exit; `lite login
--config-claude` patches it persistently. Both need the same merge and the same
apiKeyHelper command, and `up` already imports from `auth`, so the shared parts
live here rather than in either command module.
"""
import shlex
import shutil
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final
from pydantic import JsonValue, TypeAdapter, ValidationError
from .private_json import write_private_json
ENV_KEY: Final = "env"
API_KEY_HELPER_KEY: Final = "apiKeyHelper"
ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
AUTOROUTE_BACKUP_PATH: Final = Path.home() / ".litellm" / "autorouter" / "claude_settings_backup.json"
@dataclass(frozen=True, slots=True)
class SettingsFileOwner:
"""A command that takes temporary ownership of CLAUDE_SETTINGS_PATH and restores it later."""
backup_path: Path
start_command: str
stop_command: str
SETTINGS_FILE_OWNERS: Final = (
SettingsFileOwner(BACKUP_PATH, "lite up", "lite down"),
SettingsFileOwner(AUTOROUTE_BACKUP_PATH, "lite autoroute up", "lite autoroute down"),
)
_SETTINGS_ADAPTER: Final = TypeAdapter(dict[str, JsonValue])
class ClaudeSettingsError(Exception):
"""Raised for any user-actionable failure while reading or writing Claude Code settings."""
def load_json_or_empty(path: Path) -> dict[str, JsonValue]:
try:
content: Final = path.read_bytes() if path.exists() else b""
except OSError as e:
raise ClaudeSettingsError(f"Could not read {path}: {e}") from e
if not content.strip():
return {}
try:
return _SETTINGS_ADAPTER.validate_json(content)
except ValidationError:
raise ClaudeSettingsError(
f"{path} contains invalid JSON (or its root is not an object); cannot proceed safely."
)
def merge_claude_settings(
settings: Mapping[str, JsonValue], base_url: str, api_key_helper: str
) -> dict[str, JsonValue]:
"""Return a new settings dict wired to route Claude Code through the proxy.
Only env.ANTHROPIC_BASE_URL and the top-level apiKeyHelper are overridden; a
stray env.ANTHROPIC_API_KEY is dropped so it cannot outrank the helper-issued
token (same reasoning as build_agent_env in agents.py). Every other key is
preserved untouched.
"""
raw_env: Final = settings.get(ENV_KEY, {})
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
env: Final = {
**{key: value for key, value in base_env.items() if key != ANTHROPIC_API_KEY_KEY},
ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/"),
}
return {**settings, ENV_KEY: env, API_KEY_HELPER_KEY: api_key_helper}
def resolve_api_key_helper(base_url: str) -> str:
"""Build the shell command Claude Code should run for its apiKeyHelper.
Resolves `lite` to an absolute path so the helper works regardless of the
PATH visible to whatever subprocess Claude Code spawns it from. Passing
--base-url explicitly (rather than relying on the bare invocation Claude
Code would otherwise use) makes `print-token` enforce that the cached
token was actually issued for this proxy -- without it, a token minted
for a different, previously-logged-into proxy would be handed to
whichever server the settings currently point at.
--base-url belongs to the top-level `lite` group, so it has to precede the
subcommand; click rejects it outright after `print-token`.
"""
lite_path: Final = shutil.which("lite")
if lite_path is None:
raise ClaudeSettingsError(
"Could not find `lite` on your PATH. Claude Code's apiKeyHelper needs an absolute path to it."
)
return f"{shlex.quote(lite_path)} --base-url {shlex.quote(base_url)} auth print-token"
def write_claude_settings(base_url: str, settings_path: Path, owners: Sequence[SettingsFileOwner]) -> None:
"""Persistently point Claude Code at base_url, preserving every unrelated setting.
Refuses while any owner holds a backup: each restores its backup when it
stops, which would silently undo this write.
"""
for owner in owners:
if owner.backup_path.exists():
raise ClaudeSettingsError(
f"`{owner.start_command}` is currently managing {settings_path} (backup at "
f"{owner.backup_path}) and will restore it when it stops. "
f"Run `{owner.stop_command}` first, then retry."
)
normalized_base_url: Final = base_url.rstrip("/")
api_key_helper: Final = resolve_api_key_helper(normalized_base_url)
existing: Final = load_json_or_empty(settings_path)
raw_env: Final = existing.get(ENV_KEY)
if raw_env is not None and not isinstance(raw_env, dict):
raise ClaudeSettingsError(
f'{settings_path} has a non-object "{ENV_KEY}" value, which this would discard. '
"Fix or remove it, then retry."
)
merged: Final = merge_claude_settings(existing, normalized_base_url, api_key_helper)
# os.replace() swaps the symlink itself for a regular file, silently detaching a
# settings.json that is symlinked into a dotfiles repo. There is no backup to undo
# that here, unlike `lite up`, so write through to the link's target instead.
target: Final = settings_path.resolve() if settings_path.is_symlink() else settings_path
try:
write_private_json(str(target), merged)
except OSError as e:
raise ClaudeSettingsError(f"Could not write {target}: {e}") from e
__all__ = (
"ANTHROPIC_API_KEY_KEY",
"ANTHROPIC_BASE_URL_KEY",
"API_KEY_HELPER_KEY",
"AUTOROUTE_BACKUP_PATH",
"BACKUP_PATH",
"CLAUDE_SETTINGS_PATH",
"ENV_KEY",
"SETTINGS_FILE_OWNERS",
"ClaudeSettingsError",
"SettingsFileOwner",
"load_json_or_empty",
"merge_claude_settings",
"resolve_api_key_helper",
"write_claude_settings",
)

View file

@ -2,12 +2,10 @@ import atexit
import contextlib
import json
import os
import shlex
import shutil
import signal
import sys
import threading
from collections.abc import Iterator, Mapping
from collections.abc import Iterator
from dataclasses import dataclass
from pathlib import Path
from types import FrameType
@ -20,17 +18,17 @@ from litellm.litellm_core_utils.cli_token_utils import is_cli_token_fresh
from .agents import AgentRunError, resolve_api_key, verify_proxy_key
from .auth import load_token, login
ENV_KEY: Final = "env"
API_KEY_HELPER_KEY: Final = "apiKeyHelper"
ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL"
ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY"
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
from .claude_settings import (
BACKUP_PATH,
CLAUDE_SETTINGS_PATH,
ClaudeSettingsError,
load_json_or_empty,
merge_claude_settings,
resolve_api_key_helper,
)
class UpError(Exception):
class UpError(ClaudeSettingsError):
"""Raised for any user-actionable failure while starting/stopping interception."""
@ -42,40 +40,9 @@ class BackupRecord:
content: dict[str, JsonValue] | None
_SETTINGS_ADAPTER: Final = TypeAdapter(dict[str, JsonValue])
_BACKUP_RECORD_ADAPTER: Final = TypeAdapter(BackupRecord)
def load_json_or_empty(path: Path) -> dict[str, JsonValue]:
if not path.exists():
return {}
with open(path, "r") as f:
content: Final = f.read()
if not content.strip():
return {}
try:
return _SETTINGS_ADAPTER.validate_json(content)
except ValidationError:
raise UpError(f"{path} contains invalid JSON (or its root is not an object); cannot proceed safely.")
def merge_claude_settings(
settings: Mapping[str, JsonValue], base_url: str, api_key_helper: str
) -> dict[str, JsonValue]:
"""Return a new settings dict wired to route Claude Code through the proxy.
Only env.ANTHROPIC_BASE_URL and the top-level apiKeyHelper are overridden; a
stray env.ANTHROPIC_API_KEY is dropped so it cannot outrank the helper-issued
token (same reasoning as build_agent_env in agents.py). Every other key is
preserved untouched.
"""
raw_env: Final = settings.get(ENV_KEY, {})
base_env: Final = raw_env if isinstance(raw_env, dict) else {}
env: Final = {**base_env, ANTHROPIC_BASE_URL_KEY: base_url.rstrip("/")}
env.pop(ANTHROPIC_API_KEY_KEY, None)
return {**settings, ENV_KEY: env, API_KEY_HELPER_KEY: api_key_helper}
@contextlib.contextmanager
def secure_create(path: Path) -> Iterator[IO[str]]:
"""Open path for writing with mode 0600 fixed up before any content is written.
@ -136,26 +103,6 @@ def restore_claude_settings(settings_path: Path | None = None, backup_path: Path
return record
def resolve_api_key_helper(base_url: str) -> str:
"""Build the shell command Claude Code should run for its apiKeyHelper.
Resolves `lite` to an absolute path so the helper works regardless of the
PATH visible to whatever subprocess Claude Code spawns it from. Passing
--base-url explicitly (rather than relying on the bare invocation Claude
Code would otherwise use) makes `print-token` enforce that the cached
token was actually issued for this proxy -- without it, a token minted
for a different, previously-logged-into proxy would be handed to
whichever server `up` currently points at.
"""
lite_path: Final = shutil.which("lite")
if lite_path is None:
raise UpError(
"Could not find `lite` on your PATH. Claude Code's apiKeyHelper needs "
"an absolute path to it, so `lite up` cannot continue."
)
return f"{shlex.quote(lite_path)} auth print-token --base-url {shlex.quote(base_url)}"
def _ensure_fresh_login(ctx: click.Context) -> None:
base_url: Final = ctx.obj["base_url"].rstrip("/")
token_data = load_token()
@ -224,7 +171,7 @@ def up(ctx: click.Context) -> None:
merged: Final = merge_claude_settings(original_settings, base_url, api_key_helper)
with open(CLAUDE_SETTINGS_PATH, "w") as f:
json.dump(merged, f, indent=2)
except (AgentRunError, UpError) as e:
except (AgentRunError, ClaudeSettingsError) as e:
raise click.ClickException(str(e))
click.echo(f"litellm: routing Claude Code through proxy at {base_url.rstrip('/')}")
@ -241,7 +188,7 @@ def up(ctx: click.Context) -> None:
return
try:
_restore_and_report()
except UpError as e:
except ClaudeSettingsError as e:
# Runs from atexit/a signal handler, outside Click's own exception
# handling -- raising here would only produce an unhandled-exception
# warning on stderr, not a clean message.
@ -264,7 +211,7 @@ def down() -> None:
"""
try:
_restore_and_report()
except UpError as e:
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
@ -272,6 +219,7 @@ __all__ = [
"BACKUP_PATH",
"CLAUDE_SETTINGS_PATH",
"BackupRecord",
"ClaudeSettingsError",
"UpError",
"down",
"load_json_or_empty",

View file

@ -1204,6 +1204,29 @@ class DBSpendUpdateWriter:
except Exception as e:
verbose_proxy_logger.debug("_flush_tool_discovery_queue error (non-blocking): %s", e)
@staticmethod
async def _handle_spend_update_failure(
e: Exception,
attempt: int,
n_retry_times: int,
start_time: float,
proxy_logging_obj: ProxyLogging,
) -> None:
"""Retry a failed spend-update transaction on connection errors or deadlocks, else re-raise."""
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.utils import _raise_failed_update_spend_exception
is_retryable = isinstance(e, DB_RETRY_SAFE_ERROR_TYPES) or PrismaDBExceptionHandler.is_deadlock_error(e)
if not is_retryable or attempt >= n_retry_times:
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
verbose_proxy_logger.warning(
"Retrying spend update after retryable DB error (attempt %s/%s): %s",
attempt + 1,
n_retry_times,
e,
)
await asyncio.sleep(random.uniform(2**attempt, 2 ** (attempt + 1)))
async def _commit_spend_updates_to_db(
self,
prisma_client: PrismaClient,
@ -1215,10 +1238,7 @@ class DBSpendUpdateWriter:
Commits all the spend `UPDATE` transactions to the Database
"""
from litellm.proxy.utils import (
ProxyUpdateSpend,
_raise_failed_update_spend_exception,
)
from litellm.proxy.utils import ProxyUpdateSpend
### UPDATE USER TABLE ###
user_list_transactions: Final = db_spend_update_transactions["user_list_transactions"]
@ -1238,18 +1258,13 @@ class DBSpendUpdateWriter:
data={"spend": {"increment": response_cost}},
)
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times: # If we've reached the maximum number of retries
_raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# Optionally, sleep for a bit before retrying
await asyncio.sleep(2**i) # Exponential backoff
except Exception as e:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
await self._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
### UPDATE END-USER TABLE ###
@ -1281,18 +1296,13 @@ class DBSpendUpdateWriter:
},
)
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times: # If we've reached the maximum number of retries
_raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# Optionally, sleep for a bit before retrying
await asyncio.sleep(2**i) # Exponential backoff
except Exception as e:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
await self._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
### UPDATE TEAM TABLE ###
@ -1314,18 +1324,13 @@ class DBSpendUpdateWriter:
data={"spend": {"increment": response_cost}},
)
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times: # If we've reached the maximum number of retries
_raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# Optionally, sleep for a bit before retrying
await asyncio.sleep(2**i) # Exponential backoff
except Exception as e:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
await self._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
### UPDATE TEAM Membership TABLE with spend ###
@ -1361,18 +1366,13 @@ class DBSpendUpdateWriter:
)
# Transaction succeeded, break out of retry loop
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times: # If we've reached the maximum number of retries
_raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# Optionally, sleep for a bit before retrying
await asyncio.sleep(2**i) # Exponential backoff
except Exception as e:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
await self._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# Invalidate cache for updated team memberships
@ -1403,25 +1403,13 @@ class DBSpendUpdateWriter:
data={"spend": {"increment": response_cost}},
)
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times: # If we've reached the maximum number of retries
_raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# Optionally, sleep for a bit before retrying
await asyncio.sleep(
# Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are
# cancelled basically at the same time, so if they wait the same time they will also retry at the same time
# and thus they are more likely to deadlock again.
# Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of
# repeated deadlocks, and therefore of exceeding the retry limit.
random.uniform(2**i, 2 ** (i + 1))
)
except Exception as e:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
await self._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
### UPDATE TAG TABLE ###
@ -1470,8 +1458,6 @@ class DBSpendUpdateWriter:
prisma_client: Prisma client instance
proxy_logging_obj: Proxy logging object
"""
from litellm.proxy.utils import _raise_failed_update_spend_exception
verbose_proxy_logger.debug("%s Spend transactions: %s", entity_name, transactions)
if transactions is not None and len(transactions.keys()) > 0:
for i in range(n_retry_times + 1):
@ -1493,17 +1479,13 @@ class DBSpendUpdateWriter:
data={"spend": {"increment": response_cost}},
)
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times:
_raise_failed_update_spend_exception(
e=e,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
await asyncio.sleep(2**i) # Exponential backoff
except Exception as e:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
await DBSpendUpdateWriter._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
# fmt: off
@ -1672,7 +1654,16 @@ class DBSpendUpdateWriter:
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
except Exception as e:
from litellm.proxy.db.exception_handler import (
PrismaDBExceptionHandler,
)
is_retryable = isinstance(
e, DB_RETRY_SAFE_ERROR_TYPES
) or PrismaDBExceptionHandler.is_deadlock_error(e)
if not is_retryable:
raise
if i >= n_retry_times:
_raise_failed_update_spend_exception(
e=e,

View file

@ -166,6 +166,22 @@ class PrismaDBExceptionHandler:
return True
return False
@staticmethod
def is_deadlock_error(e: Exception) -> bool:
"""True iff ``e`` is a Postgres deadlock (P2034 / 40P01) surfaced through prisma."""
import prisma
if not isinstance(e, prisma.errors.PrismaError):
return False
if getattr(e, "code", None) == "P2034":
return True
error_message = str(e).lower()
return (
"deadlock detected" in error_message
or "40p01" in error_message
or "write conflict or a deadlock" in error_message
)
@staticmethod
def is_prisma_engine_internal_error(e: Exception) -> bool:
"""True iff ``e`` is a non-``PrismaError`` exception raised from inside

View file

@ -0,0 +1,93 @@
"""
Live proxy worker census, one row per worker process.
Every uvicorn worker upserts its own row on a fixed heartbeat, so counting
rows with a recent heartbeat answers "how many workers share this database?"
without any coordination. The Admin UI's "no Redis" banner uses that count to
hide itself for deployments that are provably a single worker, where per-worker
rate limits, budgets, and router state are already global. All timestamps are
written and compared with the database's own clock, so pods with skewed clocks
still agree.
"""
from __future__ import annotations
import socket
from typing import TYPE_CHECKING, Final
from pydantic import TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient
PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS: Final = 60
PROXY_WORKER_LIVENESS_WINDOW_SECONDS: Final = 3 * PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS
STALE_ROW_RETENTION_SECONDS: Final = 3600
BEAT_SQL: Final = """
INSERT INTO "LiteLLM_ProxyWorkerHeartbeat" (worker_id, hostname, last_heartbeat_at)
VALUES ($1, $2, NOW())
ON CONFLICT (worker_id) DO UPDATE SET last_heartbeat_at = NOW()
"""
PRUNE_SQL: Final = """
DELETE FROM "LiteLLM_ProxyWorkerHeartbeat"
WHERE last_heartbeat_at < NOW() - make_interval(secs => $1)
"""
COUNT_SQL: Final = """
SELECT COUNT(*)::int AS live_workers FROM "LiteLLM_ProxyWorkerHeartbeat"
WHERE last_heartbeat_at > NOW() - make_interval(secs => $1)
"""
DEREGISTER_SQL: Final = """
DELETE FROM "LiteLLM_ProxyWorkerHeartbeat" WHERE worker_id = $1
"""
class _LiveWorkerCountRow(TypedDict):
live_workers: ReadOnly[int]
_COUNT_ROWS_ADAPTER: Final = TypeAdapter(tuple[_LiveWorkerCountRow, ...])
class ProxyWorkerHeartbeat:
def __init__(self, prisma_client: PrismaClient, worker_id: str | None = None) -> None:
self.prisma_client: Final = prisma_client
self.worker_id: Final[str] = worker_id or str(uuid.uuid4())
self.hostname: Final = socket.gethostname()
async def beat(self) -> None:
try:
await self.prisma_client.db.execute_raw(BEAT_SQL, self.worker_id, self.hostname)
await self.prisma_client.db.execute_raw(PRUNE_SQL, STALE_ROW_RETENTION_SECONDS)
except Exception as beat_err: # noqa: BLE001 # a missed heartbeat must never take down the worker
verbose_proxy_logger.debug("Proxy worker heartbeat write failed: %s", beat_err)
async def deregister(self) -> None:
try:
await self.prisma_client.db.execute_raw(DEREGISTER_SQL, self.worker_id)
except Exception as deregister_err: # noqa: BLE001 # best-effort cleanup; the liveness window ages the row out anyway
verbose_proxy_logger.debug("Proxy worker heartbeat deregister failed: %s", deregister_err)
async def count_live_proxy_workers(prisma_client: PrismaClient) -> int | None:
"""
The number of workers with a recent heartbeat, or None when the database
cannot answer. Callers must treat None as "unknown", not as zero. Always
counts on the primary: a lagging read replica must never undercount.
"""
try:
db: Final = prisma_client.db
primary_db: Final = db.writer if isinstance(db, RoutingPrismaWrapper) else db
rows: Final = await primary_db.query_raw(COUNT_SQL, PROXY_WORKER_LIVENESS_WINDOW_SECONDS)
return _COUNT_ROWS_ADAPTER.validate_python(rows)[0]["live_workers"]
except Exception as count_err: # noqa: BLE001 # an unknown count must degrade to "warn", never to a 503
verbose_proxy_logger.debug("Live proxy worker count unavailable: %s", count_err)
return None

View file

@ -0,0 +1,40 @@
# Claude Code / Anthropic-native web search on Bedrock, backed by
# Amazon Bedrock AgentCore Web Search (AWS-managed web index, no third-party
# search API). See litellm/llms/bedrock/search/transformation.py for details.
model_list:
- model_name: claude-sonnet
litellm_params:
model: bedrock/us.anthropic.claude-sonnet-5
aws_region_name: us-east-1
search_tools:
- search_tool_name: agentcore-search
litellm_params:
search_provider: agentcore
# Your AgentCore Gateway MCP endpoint (gateway must have a `web-search`
# connector target). Alternatively set the AGENTCORE_GATEWAY_URL env var.
api_base: https://<gateway-id>.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp
# The gateway exposes the connector as "<target-name>___WebSearch".
# Default is "web-search-tool___WebSearch", matching the target name used
# in the AWS docs' boto3/CLI setup examples. If your target was created
# with a different name (misconfiguration surfaces as an MCP "tool not
# found" error), set the AGENTCORE_SEARCH_TOOL_NAME env var or pass
# tool_name in the request body. The search router forwards only
# search_provider / api_key / api_base from this litellm_params block,
# so a tool_name set here would be silently ignored.
# AWS_IAM gateway (default): SigV4-signed using the standard AWS
# credential chain (env / profile / IRSA / instance role). Explicit
# aws_access_key_id / aws_secret_access_key set here would be silently
# ignored for the same reason; pass them per request instead.
# CUSTOM_JWT gateway alternative — OAuth2 bearer token instead of SigV4:
# api_key: os.environ/AGENTCORE_GATEWAY_TOKEN
litellm_settings:
callbacks: ["websearch_interception"]
websearch_interception_params:
enabled_providers: ["bedrock"]
search_tool_name: agentcore-search

View file

@ -34,6 +34,7 @@ from litellm.proxy.auth.auth_utils import (
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.db.proxy_worker_heartbeat import count_live_proxy_workers
from litellm.proxy.health_check import (
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS,
_clean_endpoint_data,
@ -1451,7 +1452,7 @@ def callback_name(callback):
DISABLE_NO_REDIS_WARNING_ENV_VAR: Final = "LITELLM_DISABLE_NO_REDIS_WARNING"
def _show_no_redis_warning() -> bool:
async def _show_no_redis_warning() -> bool:
"""
Whether the UI should warn that no Redis is configured.
@ -1461,16 +1462,22 @@ def _show_no_redis_warning() -> bool:
coordination cache (from a Redis response cache, general_settings.
coordination_redis, or the REDIS_* env fallback) and the router's own
Redis (router_settings.redis_host), which backs cooldowns and usage-based
routing on its own. Operators who know they run one worker can silence the
warning with LITELLM_DISABLE_NO_REDIS_WARNING=true.
routing on its own. A deployment whose worker-heartbeat census proves it
is exactly one worker needs no cross-worker coordination, so it never
warns; when the census is unavailable or shows more than one worker, the
warning stands unless LITELLM_DISABLE_NO_REDIS_WARNING=true silences it.
"""
from litellm.proxy.proxy_server import llm_router, redis_usage_cache
from litellm.proxy.proxy_server import llm_router, prisma_client, redis_usage_cache
if redis_usage_cache is not None:
return False
if llm_router is not None and llm_router.cache.redis_cache is not None:
return False
return get_secret_bool(DISABLE_NO_REDIS_WARNING_ENV_VAR, False) is not True
if get_secret_bool(DISABLE_NO_REDIS_WARNING_ENV_VAR, False) is True:
return False
if prisma_client is None:
return True
return await count_live_proxy_workers(prisma_client) != 1
async def _get_health_readiness_details(
@ -1513,7 +1520,7 @@ async def _get_health_readiness_details(
# check log level
log_level_name: Final = logging.getLevelName(verbose_logger.getEffectiveLevel())
is_detailed_debug: Final = verbose_logger.isEnabledFor(logging.DEBUG)
show_no_redis_warning: Final = _show_no_redis_warning()
show_no_redis_warning: Final = await _show_no_redis_warning()
# check DB
if prisma_client is not None: # if db passed in, check if it's connected

View file

@ -0,0 +1,182 @@
import json
from collections.abc import Iterator
from dataclasses import dataclass
from itertools import chain
from typing import BinaryIO, Final, NoReturn, assert_never
from litellm.proxy._types import ProxyException
BATCH_LINE_REQUIRED_KEYS: Final = ("custom_id", "method", "url", "body")
_MB: Final = 1024 * 1024
@dataclass(frozen=True, slots=True)
class BatchFileTooLarge:
size_bytes: int
limit_mb: int
@dataclass(frozen=True, slots=True)
class BatchFileWrongExtension:
filename: str
@dataclass(frozen=True, slots=True)
class BatchFileEmpty:
pass
@dataclass(frozen=True, slots=True)
class BatchFileInvalidJsonLine:
line_number: int
@dataclass(frozen=True, slots=True)
class BatchFileLineNotObject:
line_number: int
@dataclass(frozen=True, slots=True)
class BatchFileMissingLineKey:
line_number: int
key: str
BatchFileValidationFailure = (
BatchFileTooLarge
| BatchFileWrongExtension
| BatchFileEmpty
| BatchFileInvalidJsonLine
| BatchFileLineNotObject
| BatchFileMissingLineKey
)
def _file_size_bytes(file_source: bytes | BinaryIO) -> int:
if isinstance(file_source, bytes):
return len(file_source)
file_source.seek(0, 2)
size: Final = file_source.tell()
file_source.seek(0)
return size
def _iter_lines(file_source: bytes | BinaryIO) -> Iterator[bytes]:
if isinstance(file_source, bytes):
return iter(file_source.splitlines())
file_source.seek(0)
return iter(file_source)
def _check_line(line_number: int, raw_line: bytes) -> BatchFileValidationFailure | None:
try:
parsed: Final = json.loads(raw_line)
except (json.JSONDecodeError, UnicodeDecodeError):
return BatchFileInvalidJsonLine(line_number=line_number)
if not isinstance(parsed, dict):
return BatchFileLineNotObject(line_number=line_number)
missing: Final = next((key for key in BATCH_LINE_REQUIRED_KEYS if key not in parsed), None)
if missing is None:
return None
return BatchFileMissingLineKey(line_number=line_number, key=missing)
def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | None:
content_lines: Final = (
(line_number, raw_line)
for line_number, raw_line in enumerate(_iter_lines(file_source), start=1)
if raw_line.strip()
)
first_line: Final = next(content_lines, None)
if first_line is None:
return BatchFileEmpty()
return next(
(
failure
for line_number, raw_line in chain((first_line,), content_lines)
for failure in (_check_line(line_number, raw_line),)
if failure is not None
),
None,
)
def check_batch_file_upload(
filename: str | None,
file_source: bytes | BinaryIO,
max_batch_file_size_mb: int | None,
) -> BatchFileValidationFailure | None:
if filename is None or not filename.lower().endswith(".jsonl"):
return BatchFileWrongExtension(filename=filename or "")
if max_batch_file_size_mb is not None and max_batch_file_size_mb > 0:
size_bytes: Final = _file_size_bytes(file_source)
if size_bytes > max_batch_file_size_mb * _MB:
return BatchFileTooLarge(size_bytes=size_bytes, limit_mb=max_batch_file_size_mb)
scan_failure: Final = _scan_lines(file_source)
if not isinstance(file_source, bytes):
file_source.seek(0)
return scan_failure
def raise_batch_file_validation_failure(failure: BatchFileValidationFailure) -> NoReturn:
match failure:
case BatchFileTooLarge(size_bytes=size_bytes, limit_mb=limit_mb):
raise ProxyException(
message=(
f"Batch input file is {size_bytes / _MB:.1f} MB, which exceeds the configured "
f"max_batch_file_size_mb of {limit_mb} MB. The file was not forwarded to the provider."
),
type="invalid_request_error",
param="file",
code=413,
)
case BatchFileWrongExtension(filename=filename):
raise ProxyException(
message=(
f"Invalid file format for Batch API: '{filename}'. "
"Batch input files must be .jsonl files. The file was not forwarded to the provider."
),
type="invalid_request_error",
param="file",
code=400,
)
case BatchFileEmpty():
raise ProxyException(
message="Batch input file has no request lines. The file was not forwarded to the provider.",
type="invalid_request_error",
param="file",
code=400,
)
case BatchFileInvalidJsonLine(line_number=line_number):
raise ProxyException(
message=(
f"Batch input file line {line_number} is not valid JSON. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",
param="file",
code=400,
)
case BatchFileLineNotObject(line_number=line_number):
raise ProxyException(
message=(
f"Batch input file line {line_number} must be a JSON object. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",
param="file",
code=400,
)
case BatchFileMissingLineKey(line_number=line_number, key=key):
raise ProxyException(
message=(
f"Missing required parameter: '{key}' (batch input file line {line_number}). "
f"Each line must be a JSON object with keys {', '.join(BATCH_LINE_REQUIRED_KEYS)}. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",
param=key,
code=400,
)
case _:
assert_never(failure)

View file

@ -21,6 +21,7 @@ from fastapi import (
UploadFile,
status,
)
from pydantic import TypeAdapter
import litellm
from litellm import CreateFileRequest, get_secret_str
@ -41,6 +42,10 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
get_custom_llm_provider_from_request_headers,
get_custom_llm_provider_from_request_query,
)
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
check_batch_file_upload,
raise_batch_file_validation_failure,
)
from litellm.proxy.openai_files_endpoints.common_utils import (
_is_base64_encoded_unified_file_id,
add_internal_model_credentials,
@ -65,6 +70,8 @@ from litellm.types.llms.openai import (
router: Final = APIRouter()
_MAX_BATCH_FILE_SIZE_MB_ADAPTER: Final = TypeAdapter(int | None)
files_config = None
@ -361,18 +368,27 @@ async def create_file(
# Prepare the data for forwarding
# Replace with:
valid_purposes: Final = get_args(OpenAIFilesPurpose)
if purpose not in valid_purposes:
raise HTTPException(
status_code=400,
detail={
"error": f"Invalid purpose: {purpose}. Must be one of: {valid_purposes}",
},
raise ProxyException(
message=f"Invalid purpose: {purpose}. Must be one of: {valid_purposes}",
type="invalid_request_error",
param="purpose",
code=400,
)
# Cast purpose to OpenAIFilesPurpose type
purpose = cast(OpenAIFilesPurpose, purpose)
if purpose == "batch":
batch_file_failure: Final = await asyncio.to_thread(
check_batch_file_upload,
file.filename,
file_source,
_MAX_BATCH_FILE_SIZE_MB_ADAPTER.validate_python(general_settings.get("max_batch_file_size_mb")),
)
if batch_file_failure is not None:
raise_batch_file_validation_failure(batch_file_failure)
data = {}
# Parse expires_after if provided
@ -552,6 +568,8 @@ async def create_file(
user_api_key_dict=user_api_key_dict, original_exception=e, request_data=data
)
verbose_proxy_logger.exception("litellm.proxy.proxy_server.create_file(): Exception occured - %s", e)
if isinstance(e, ProxyException):
raise e
if isinstance(e, HTTPException):
raise ProxyException(
message=getattr(e, "message", str(e.detail)),

View file

@ -10,6 +10,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import VERTEX_BATCH_PREDICTION_JOBS_ROUTE
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.vertex_ai.common_utils import get_vertex_location_from_url
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator as VertexModelResponseIterator,
)
@ -60,6 +61,9 @@ class VertexPassthroughLoggingHandler:
request_body: dict | None = None,
**kwargs,
) -> PassThroughEndpointLoggingTypedDict:
vertex_location: Final = get_vertex_location_from_url(url_route)
if vertex_location is not None:
logging_obj.optional_params["vertex_location"] = vertex_location
if "predictLongRunning" in url_route:
model = VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
@ -82,6 +86,7 @@ class VertexPassthroughLoggingHandler:
model=model,
custom_llm_provider="vertex_ai",
call_type="create_video",
vertex_location=vertex_location,
)
# Set response_cost in _hidden_params to prevent recalculation
@ -123,6 +128,7 @@ class VertexPassthroughLoggingHandler:
end_time=end_time,
logging_obj=logging_obj,
custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route),
vertex_location=vertex_location,
)
return {
@ -190,6 +196,7 @@ class VertexPassthroughLoggingHandler:
end_time=end_time,
logging_obj=logging_obj,
custom_llm_provider="vertex_ai",
vertex_location=vertex_location,
)
return {
@ -206,6 +213,7 @@ class VertexPassthroughLoggingHandler:
model="vertex_ai/search_api",
custom_llm_provider="vertex_ai",
call_type="vector_store_search",
vertex_location=vertex_location,
)
standard_pass_through_response_object: Final[StandardPassThroughResponseObject] = {
@ -302,6 +310,7 @@ class VertexPassthroughLoggingHandler:
completion_response=litellm_prediction_response,
model=model,
custom_llm_provider="vertex_ai",
vertex_location=get_vertex_location_from_url(url_route),
)
kwargs["response_cost"] = response_cost
@ -381,6 +390,7 @@ class VertexPassthroughLoggingHandler:
completion_response=litellm_embedding_response,
model=model,
custom_llm_provider=custom_llm_provider,
vertex_location=get_vertex_location_from_url(url_route),
)
kwargs["response_cost"] = response_cost
@ -413,6 +423,9 @@ class VertexPassthroughLoggingHandler:
- Logs in litellm callbacks
"""
kwargs: dict[str, Any] = {}
vertex_location: Final = get_vertex_location_from_url(url_route)
if vertex_location is not None:
litellm_logging_obj.optional_params["vertex_location"] = vertex_location
model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route)
complete_streaming_response: Final = VertexPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=all_chunks,
@ -438,6 +451,7 @@ class VertexPassthroughLoggingHandler:
end_time=end_time,
logging_obj=litellm_logging_obj,
custom_llm_provider=VertexPassthroughLoggingHandler._get_custom_llm_provider_from_url(url_route),
vertex_location=vertex_location,
)
return {
@ -591,6 +605,7 @@ class VertexPassthroughLoggingHandler:
end_time: datetime,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: str,
vertex_location: str | None,
) -> dict:
"""
Create the standard logging object for Vertex passthrough generateContent (streaming and non-streaming)
@ -601,6 +616,7 @@ class VertexPassthroughLoggingHandler:
completion_response=litellm_model_response,
model=model,
custom_llm_provider="vertex_ai",
vertex_location=vertex_location,
)
kwargs["response_cost"] = response_cost

View file

@ -384,6 +384,10 @@ from litellm.proxy.db.gateway_request_tracking import (
GatewayRequestAccumulator,
flush_gateway_requests,
)
from litellm.proxy.db.proxy_worker_heartbeat import (
PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS,
ProxyWorkerHeartbeat,
)
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
from litellm.proxy.discovery_endpoints import ui_discovery_endpoints_router
from litellm.proxy.fine_tuning_endpoints.endpoints import router as fine_tuning_router
@ -635,6 +639,7 @@ from litellm.secret_managers.main import (
get_secret_bool,
get_secret_str,
normalize_nonempty_secret_str,
secret_manager_would_be_consulted,
str_to_bool,
)
from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs
@ -874,9 +879,11 @@ async def _flush_spend_logs_queue_on_shutdown() -> None:
verbose_proxy_logger.exception("Error flushing spend logs queue on shutdown: %s", e)
async def proxy_shutdown_event() -> None:
async def proxy_shutdown_event(worker_heartbeat: ProxyWorkerHeartbeat | None = None) -> None:
global prisma_client, master_key, user_custom_auth, user_custom_key_generate, user_custom_key_update
verbose_proxy_logger.info("Shutting down LiteLLM Proxy Server")
if worker_heartbeat is not None and prisma_client:
await worker_heartbeat.deregister()
if prisma_client:
# Drain the SGR fold first: it lives in memory, so an un-drained interval
# is lost, and a write attempted after disconnect raises
@ -1210,7 +1217,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
)
### START BATCH WRITING DB + CHECKING NEW MODELS###
if prisma_client is not None:
worker_heartbeat: Final = (
await ProxyStartupEvent.initialize_scheduled_background_jobs(
general_settings=general_settings,
prisma_client=prisma_client,
@ -1219,7 +1226,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
proxy_batch_write_at=proxy_batch_write_at,
proxy_logging_obj=proxy_logging_obj,
)
if prisma_client is not None
else None
)
if prisma_client is not None:
await ProxyStartupEvent._update_default_team_member_budget()
## SYNC UI SETTINGS ##
@ -1290,7 +1300,7 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
await proxy_config.stop_auth_cache_invalidation_subscriber()
await proxy_shutdown_event()
await proxy_shutdown_event(worker_heartbeat=worker_heartbeat)
def _generate_stable_operation_id(route: "APIRoute") -> str:
@ -4371,9 +4381,55 @@ class ProxyConfig:
item = self._check_for_os_environ_vars(config=item, depth=depth + 1, max_depth=max_depth)
# if the value is a string and starts with "os.environ/" - then it's an environment variable
elif isinstance(value, str) and value.startswith("os.environ/"):
config[key] = get_secret(value)
resolved = get_secret(value)
if resolved is None and secret_manager_would_be_consulted(value):
verbose_proxy_logger.warning("%s is absent from the configured secret manager", value)
config[key] = resolved
return config
def _initialize_secret_manager_from_raw_config(
self, config: Mapping[str, object], config_file_path: str | None
) -> None:
"""
Bring the secret manager up before `os.environ/<KEY>` references are resolved.
`_check_for_os_environ_vars` writes whatever it resolves back into the config, so a key
held only by the secret manager would otherwise become a permanent `None` that the later
fallbacks in `load_config` can no longer recover from.
`get_config` also runs on management-endpoint request paths, so this returns early once a
manager exists rather than rebuilding the client on every request.
The manager's own settings can only come from real environment variables, so they are
resolved against a throwaway copy and the config is left untouched for the main pass.
"""
if litellm.secret_manager_client is not None:
return
general_settings: Final = config.get("general_settings")
if not isinstance(general_settings, dict):
return
raw_system: Final = general_settings.get("key_management_system")
key_management_system: Final = (
get_secret(raw_system)
if isinstance(raw_system, str) and raw_system.startswith("os.environ/")
else raw_system
)
if not isinstance(key_management_system, str):
return
raw_settings: Final = general_settings.get("key_management_settings")
if isinstance(raw_settings, dict):
litellm._key_management_settings = KeyManagementSettings(
**self._check_for_os_environ_vars(config=copy.deepcopy(raw_settings))
)
self.initialize_secret_manager(
key_management_system=key_management_system,
config_file_path=config_file_path,
)
def _get_team_config(self, team_id: str, all_teams_config: list[dict]) -> dict:
team_config: dict = {}
for team in all_teams_config:
@ -4544,6 +4600,8 @@ class ProxyConfig:
printed_yaml: Final = copy.deepcopy(config)
printed_yaml.pop("environment_variables", None)
self._initialize_secret_manager_from_raw_config(config=config, config_file_path=config_file_path)
config = self._check_for_os_environ_vars(config=config)
self.update_config_state(config=config)
@ -5114,17 +5172,14 @@ class ProxyConfig:
key: general_settings[key] for key in SPEND_LOG_CLEANUP_BOUND_SETTINGS if key in general_settings
}
### LOAD KEY MANAGEMENT SETTINGS FIRST (needed for custom secret manager) ###
### LOAD KEY MANAGEMENT SETTINGS ###
# The secret manager itself is brought up by get_config(), which runs before the
# `os.environ/` references in this config were resolved. Re-reading the settings here
# picks up any of them that were themselves secret-manager backed.
key_management_settings: Final = general_settings.get("key_management_settings", None)
if key_management_settings is not None:
litellm._key_management_settings = KeyManagementSettings(**key_management_settings)
### LOAD SECRET MANAGER ###
key_management_system: Final = general_settings.get("key_management_system", None)
self.initialize_secret_manager(
key_management_system=key_management_system,
config_file_path=config_file_path,
)
### [DEPRECATED] LOAD FROM GOOGLE KMS ### old way of loading from google kms
use_google_kms: Final = general_settings.get("use_google_kms", False)
load_google_kms(use_google_kms=use_google_kms)
@ -6316,6 +6371,9 @@ class ProxyConfig:
if "global_max_parallel_requests" in _general_settings:
general_settings["global_max_parallel_requests"] = _general_settings["global_max_parallel_requests"]
if "max_batch_file_size_mb" not in self._yaml_general_settings_keys:
general_settings["max_batch_file_size_mb"] = _general_settings.get("max_batch_file_size_mb")
## ALERTING ARGS ##
if "alerting_args" in _general_settings:
general_settings["alerting_args"] = _general_settings["alerting_args"]
@ -8733,7 +8791,7 @@ class ProxyStartupEvent:
proxy_budget_rescheduler_max_time: int,
proxy_batch_write_at: int,
proxy_logging_obj: ProxyLogging,
):
) -> ProxyWorkerHeartbeat:
"""Initializes scheduled background jobs"""
global store_model_in_db, scheduler
@ -8778,6 +8836,18 @@ class ProxyStartupEvent:
# Ensure minimum interval of 30 seconds for batch writing to prevent memory issues
batch_writing_interval: Final = proxy_batch_write_at + random.randint(0, 5)
### PROXY WORKER HEARTBEAT ###
worker_heartbeat: Final = ProxyWorkerHeartbeat(prisma_client=prisma_client)
await worker_heartbeat.beat()
scheduler.add_job(
worker_heartbeat.beat,
"interval",
seconds=PROXY_WORKER_HEARTBEAT_INTERVAL_SECONDS,
id="proxy_worker_heartbeat_job",
replace_existing=True,
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
)
### RESET BUDGET ###
if general_settings.get("disable_reset_budget", False) is False:
budget_reset_job: Final = ResetBudgetJob(
@ -9117,6 +9187,7 @@ class ProxyStartupEvent:
"APScheduler started with memory leak prevention settings: removed jitter, increased intervals, misfire_grace_time=%s",
APSCHEDULER_MISFIRE_GRACE_TIME,
)
return worker_heartbeat
@classmethod
async def _initialize_spend_tracking_background_jobs(cls, scheduler: AsyncIOScheduler):
@ -15689,6 +15760,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"max_parallel_requests": "Integer",
"global_max_parallel_requests": "Integer",
"max_request_size_mb": "Integer",
"max_batch_file_size_mb": "Integer",
"max_response_size_mb": "Integer",
"proxy_config_reload_interval_seconds": "Integer",
"pass_through_endpoints": "PydanticModel",

View file

@ -9,7 +9,8 @@ effects.
Instead we reuse ``ProxyConfig.get_config`` — the actual config reader — so the
gateway inherits the same heavy lifting the proxy does: ``include:`` merging,
``os.environ/`` + secret-manager resolution, and DB-stored models (when a DB is
configured). It has no proxy-setup side effects. Returns the resolved
configured). Its only proxy-setup side effect is bringing up the configured
secret manager, which is what makes that resolution work. Returns the resolved
``model_list``; the Rust side deserializes each entry into its ``Deployment``.
"""

View file

@ -947,6 +947,17 @@ model LiteLLM_DailyTagSpend {
}
// One row per live proxy worker process. Workers upsert their row on a fixed
// heartbeat; counting rows with a recent heartbeat tells how many workers share
// this database, which lets the Admin UI hide its "no Redis" warning for
// deployments that are provably a single worker.
model LiteLLM_ProxyWorkerHeartbeat {
worker_id String @id
hostname String
started_at DateTime @default(now())
last_heartbeat_at DateTime @default(now())
}
// Track the status of cron jobs running. Only allow one pod to run the job at a time
model LiteLLM_CronJob {
cronjob_id String @id @default(cuid()) // Unique ID for the record

View file

@ -130,6 +130,7 @@ class PricingBasis(NamedTuple):
service_tier: str | None = None
data_residency: str | None = None
vertex_location: str | None = None
_STANDARD_RATES: Final = PricingBasis()
@ -141,8 +142,8 @@ def _pricing_basis(cost_breakdown: Mapping[str, object] | None) -> PricingBasis:
Rows written before this field shipped carry neither key, and there is no backfill:
they price at standard rates, which is what they already did.
Both values survive a JSON round trip on the way here, so neither is guaranteed to be
a string. `generic_cost_per_token` calls `.lower()` on both without a type check, and
These values survive a JSON round trip on the way here, so none is guaranteed to be
a string. `generic_cost_per_token` calls `.lower()` on them without a type check, and
the resulting `AttributeError` would be swallowed into a silent zero by the caller's
`except`, so anything that is not a string is dropped here instead.
"""
@ -150,9 +151,11 @@ def _pricing_basis(cost_breakdown: Mapping[str, object] | None) -> PricingBasis:
return _STANDARD_RATES
service_tier: Final = cost_breakdown.get("service_tier")
data_residency: Final = cost_breakdown.get("data_residency")
vertex_location: Final = cost_breakdown.get("vertex_location")
return PricingBasis(
service_tier=service_tier if isinstance(service_tier, str) else None,
data_residency=data_residency if isinstance(data_residency, str) else None,
vertex_location=vertex_location if isinstance(vertex_location, str) else None,
)
@ -193,6 +196,7 @@ def _cost_of_usage(
service_tier=basis.service_tier,
data_residency=basis.data_residency,
model_info=model_info,
vertex_location=basis.vertex_location,
)
except Exception as e: # noqa: BLE001 # get_model_info raises bare Exception for unmapped models; degrade to zero savings
verbose_proxy_logger.debug(

View file

@ -30,7 +30,6 @@ from litellm.constants import (
SPEND_LOG_WRITE_BATCH_MAX_BYTES,
)
from litellm.proxy._types import (
DB_RETRY_SAFE_ERROR_TYPES,
CommonProxyErrors,
ProxyErrorTypes,
ProxyException,
@ -5960,15 +5959,14 @@ class ProxyUpdateSpend:
)
break
except DB_RETRY_SAFE_ERROR_TYPES as e:
if i >= n_retry_times: # If we've reached the maximum number of retries
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
)
# Optionally, sleep for a bit before retrying
await asyncio.sleep(2**i) # Exponential backoff
except Exception as e:
_raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj)
await DBSpendUpdateWriter._handle_spend_update_failure(
e=e,
attempt=i,
n_retry_times=n_retry_times,
start_time=start_time,
proxy_logging_obj=proxy_logging_obj,
)
@staticmethod
async def update_spend_logs(

View file

@ -1996,6 +1996,12 @@ class LiteLLMCompletionResponsesConfig:
output_items.append(item)
return output_items
@staticmethod
def _encode_thinking_blocks(message: Message) -> str | None:
thinking_blocks: Final[Sequence[Mapping[str, object]]] = getattr(message, "thinking_blocks", None) or ()
preserved: Final = tuple(block for block in thinking_blocks if block.get("signature") or block.get("data"))
return json.dumps(preserved, separators=(",", ":")) if preserved else None
@staticmethod
def _extract_reasoning_output_items(
chat_completion_response: ModelResponse,
@ -2004,12 +2010,14 @@ class LiteLLMCompletionResponsesConfig:
for choice in choices:
if hasattr(choice, "message") and choice.message:
message = choice.message
if hasattr(message, "reasoning_content") and message.reasoning_content:
reasoning_content = getattr(message, "reasoning_content", None) or ""
encrypted_content = LiteLLMCompletionResponsesConfig._encode_thinking_blocks(message)
if reasoning_content or encrypted_content:
# Only check the first choice for reasoning content
return [
GenericResponseOutputItem(
type="reasoning",
id=f"rs_{hash(str(message.reasoning_content))}",
id=f"rs_{hash(reasoning_content or encrypted_content)}",
status=LiteLLMCompletionResponsesConfig._map_chat_completion_finish_reason_to_responses_status(
choice.finish_reason
),
@ -2017,10 +2025,13 @@ class LiteLLMCompletionResponsesConfig:
content=[
OutputText(
type="output_text",
text=message.reasoning_content,
text=text,
annotations=[],
)
for text in (reasoning_content,)
if text
],
encrypted_content=encrypted_content,
)
]
return []
@ -2292,18 +2303,19 @@ class LiteLLMCompletionResponsesConfig:
# Translate completion_tokens_details to output_tokens_details
if hasattr(usage, "completion_tokens_details") and usage.completion_tokens_details is not None:
completion_details: Final = usage.completion_tokens_details
output_details_dict: Final[dict[str, int]] = {}
if hasattr(completion_details, "reasoning_tokens") and completion_details.reasoning_tokens is not None:
output_details_dict["reasoning_tokens"] = completion_details.reasoning_tokens
if hasattr(completion_details, "text_tokens") and completion_details.text_tokens is not None:
output_details_dict["text_tokens"] = completion_details.text_tokens
if hasattr(completion_details, "image_tokens") and completion_details.image_tokens is not None:
output_details_dict["image_tokens"] = completion_details.image_tokens
if output_details_dict:
response_usage.output_tokens_details = OutputTokensDetails(**output_details_dict)
reasoning_token_count: Final = getattr(completion_details, "reasoning_tokens", None)
optional_output_details: Final[dict[str, int]] = {
field: value
for field, value in (
("text_tokens", getattr(completion_details, "text_tokens", None)),
("image_tokens", getattr(completion_details, "image_tokens", None)),
)
if value is not None
}
response_usage.output_tokens_details = OutputTokensDetails(
reasoning_tokens=reasoning_token_count if reasoning_token_count is not None else 0,
**optional_output_details,
)
return response_usage

View file

@ -10253,11 +10253,13 @@ class Router:
returned_models.extend(self.get_model_list_from_routing_groups(model_name=model_name))
if len(returned_models) == 0: # check if wildcard route
potential_wildcard_models: Final = self.pattern_router.route(model_name) or []
potential_wildcard_models: Final = self.pattern_router.get_deployments_by_pattern(model=model_name or "")
## check for team-specific wildcard models
if team_id is not None and team_id in self.team_pattern_routers:
potential_team_only_wildcard_models: Final = self.team_pattern_routers[team_id].route(model_name) or []
potential_team_only_wildcard_models: Final = self.team_pattern_routers[
team_id
].get_deployments_by_pattern(model=model_name or "")
potential_wildcard_models.extend(potential_team_only_wildcard_models)
if model_name is not None and potential_wildcard_models is not None:

View file

@ -165,7 +165,7 @@ response = litellm.completion(
### Reasoning Override
If 2+ reasoning markers are detected in the user message, the request is automatically routed to the REASONING tier regardless of the weighted score. This ensures complex reasoning tasks get the appropriate model.
If 2+ reasoning markers are detected in the user message, the request is promoted to the REASONING tier even when the weighted score maps lower, so complex reasoning tasks get the appropriate model. The promotion requires the score to reach `reasoning_override_min_score`, which tracks `tier_boundaries.simple_medium` unless set, so stock phrases on an otherwise trivial prompt cannot buy the top tier. Set it to `0` to promote on the markers alone.
### System Prompt Handling

View file

@ -1021,10 +1021,10 @@ class ComplexityRouter(CustomLogger):
weighted_score: Final = sum(d.score * weights.get(d.name, 0) for d in dimensions)
boundaries: Final = self._effective_tier_boundaries()
scored_above_simple: Final = weighted_score >= boundaries["simple_medium"]
clears_override_floor: Final = weighted_score >= self._effective_reasoning_override_min_score()
# Reuse match count from _score_keyword_match to avoid scanning twice
if reasoning_match_count >= 2 and scored_above_simple:
if reasoning_match_count >= 2 and clears_override_floor:
return ComplexityTier.REASONING, weighted_score, tuple(signals), "reasoning_override"
# Map score to tier
@ -1039,6 +1039,18 @@ class ComplexityRouter(CustomLogger):
return tier, weighted_score, tuple(signals), "heuristic_scorer"
def _effective_reasoning_override_min_score(self) -> float:
"""The score a request must reach before the reasoning-marker override may promote it.
Unset tracks the SIMPLE/MEDIUM boundary, so moving that boundary moves this floor with it
and the override still cannot rescue a request the mapping would call SIMPLE. An explicit
0 is a real floor, not an absent one, so the comparison is against None.
"""
configured: Final = self.config.reasoning_override_min_score
if configured is None:
return self._effective_tier_boundaries()["simple_medium"]
return configured
def _effective_tier_boundaries(self) -> StandardLoggingRoutingDecisionTierBoundaries:
"""The tier boundaries in effect, with the documented defaults filled in.
@ -1095,6 +1107,7 @@ class ComplexityRouter(CustomLogger):
if score is not None:
decision["score"] = score
decision["tier_boundaries"] = self._effective_tier_boundaries()
decision["reasoning_override_min_score"] = self._effective_reasoning_override_min_score()
if signals:
# Stored as a list because this record is serialized to JSON for the spend
# log and read back as an array by the dashboard; a sequence type that only

View file

@ -481,6 +481,15 @@ class ComplexityRouterConfig(BaseModel):
),
)
reasoning_override_min_score: float | None = Field(
default=None,
description=(
"Minimum weighted score a request must reach before 2+ reasoning markers may promote it to the "
"reasoning tier. Unset tracks tier_boundaries.simple_medium, so the override never rescues a "
"request the scorer placed in the cheapest tier; 0 restores the unconditional override"
),
)
# Token count thresholds
token_thresholds: dict[str, int] = Field(
default_factory=lambda: DEFAULT_TOKEN_THRESHOLDS.copy(),

View file

@ -365,6 +365,22 @@ def get_secret(
raise e
def secret_manager_would_be_consulted(secret_name: str) -> bool:
"""
Returns True if a `get_secret` read for `secret_name` would actually reach the hosted manager.
Mirrors the gating `get_secret` applies below: the manager has to be up and readable, and
`hosted_keys`, when set, is an allowlist of the names it is consulted for. Callers use this to
tell "the manager does not have this key" apart from "the manager was never asked".
"""
if not _should_read_secret_from_secret_manager():
return False
key_management_settings: Final = litellm._key_management_settings
if key_management_settings is None or key_management_settings.hosted_keys is None:
return True
return secret_name.removeprefix("os.environ/") in key_management_settings.hosted_keys
def _should_read_secret_from_secret_manager() -> bool:
"""
Returns True if the secret manager should be used to read the secret, False otherwise
@ -373,11 +389,7 @@ def _should_read_secret_from_secret_manager() -> bool:
- If the `_key_management_settings` access mode is "read_only" or "read_and_write", return True
- Otherwise, return False
"""
if litellm.secret_manager_client is not None:
if litellm._key_management_settings is not None:
if (
litellm._key_management_settings.access_mode == "read_only"
or litellm._key_management_settings.access_mode == "read_and_write"
):
return True
return False
key_management_settings: Final = litellm._key_management_settings
if litellm.secret_manager_client is None or key_management_settings is None:
return False
return key_management_settings.access_mode in ("read_only", "read_and_write")

View file

@ -620,6 +620,12 @@ class AnthropicResponseUsageBlock(BaseModel):
output_tokens: int
class AnthropicOutputTokensDetails(BaseModel):
model_config = ConfigDict(extra="allow")
thinking_tokens: int | None = None
AnthropicFinishReason = Literal["end_turn", "max_tokens", "stop_sequence", "tool_use"]

View file

@ -248,6 +248,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
regional_processing_uplift_multiplier_us: (
float | None
) # OpenAI US data-residency uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
regional_endpoint_uplift_multiplier: ReadOnly[
float | None
] # Vertex AI non-global (regional) endpoint uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
output_cost_per_character: float | None # only for vertex ai models
output_cost_per_audio_token: float | None
output_cost_per_token_above_128k_tokens: float | None # only for vertex ai models
@ -2836,6 +2839,7 @@ class StandardLoggingRoutingDecision(TypedDict, total=False):
classifier_cost: float
escalated: bool
tier_boundaries: StandardLoggingRoutingDecisionTierBoundaries
reasoning_override_min_score: ReadOnly[float]
conversation_continuing: bool
savings_baseline_model: str
savings_baseline_deployment_id: str
@ -2860,6 +2864,7 @@ DERIVED_ROUTING_DECISION_FIELDS: Final[frozenset[str]] = frozenset(
"classifier_cost",
"escalated",
"tier_boundaries",
"reasoning_override_min_score",
"conversation_continuing",
"savings_baseline_model",
"savings_baseline_deployment_id",
@ -3113,16 +3118,17 @@ class CostBreakdown(TypedDict, total=False):
"""
Detailed cost breakdown for a request.
``service_tier`` and ``data_residency`` record the pricing basis the cost was
computed on, not what the caller asked for. A consumer that has to price a
counterfactual against this request (what another model would have charged for
it) needs the same basis to compare like with like, and re-deriving it from the
request is not possible after the fact: the tier the biller used comes from
``optional_params``, which no log record carries.
``service_tier``, ``data_residency``, and ``vertex_location`` record the pricing
basis the cost was computed on, not what the caller asked for. A consumer that has
to price a counterfactual against this request (what another model would have
charged for it) needs the same basis to compare like with like, and re-deriving it
from the request is not possible after the fact: the tier the biller used comes
from ``optional_params``, which no log record carries.
"""
service_tier: str | None
data_residency: str | None
vertex_location: ReadOnly[str | None]
input_cost: float # Cost of raw (non-cached) input tokens only
cache_read_cost: float # Cost of cache-read tokens (discounted rate)
cache_creation_cost: float # Cost of cache-write tokens (premium rate)
@ -3388,6 +3394,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
annotation_cost_per_page: float | None = None
regional_processing_uplift_multiplier_eu: float | None = None
regional_processing_uplift_multiplier_us: float | None = None
regional_endpoint_uplift_multiplier: float | None = None
@classmethod
def strip_custom_pricing_fields(cls, model_info: dict[str, Any]) -> dict[str, Any]:
@ -3818,6 +3825,7 @@ class SearchProviders(str, Enum):
YOU_COM = "you_com"
APISERPENT = "apiserpent"
TINYFISH = "tinyfish"
AGENTCORE = "agentcore"
NIMBLE = "nimble"

View file

@ -5662,6 +5662,7 @@ def _get_model_info_helper(
regional_processing_uplift_multiplier_us=_model_info.get(
"regional_processing_uplift_multiplier_us", None
),
regional_endpoint_uplift_multiplier=_model_info.get("regional_endpoint_uplift_multiplier", None),
output_cost_per_audio_token=_model_info.get("output_cost_per_audio_token", None),
output_cost_per_character=_model_info.get("output_cost_per_character", None),
output_cost_per_reasoning_token=_model_info.get("output_cost_per_reasoning_token", None),
@ -9077,6 +9078,7 @@ class ProviderConfigManager:
from litellm.llms.apiserpent.search.transformation import (
APISerpentSearchConfig,
)
from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig
from litellm.llms.brave.search.transformation import BraveSearchConfig
from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig
from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig
@ -9115,6 +9117,7 @@ class ProviderConfigManager:
SearchProviders.YOU_COM: YouComSearchConfig,
SearchProviders.APISERPENT: APISerpentSearchConfig,
SearchProviders.TINYFISH: TinyfishSearchConfig,
SearchProviders.AGENTCORE: AgentCoreSearchConfig,
SearchProviders.NIMBLE: NimbleSearchConfig,
}
config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None)

View file

@ -16792,6 +16792,14 @@
"notes": "APISerpent deep search (/api/search), multi-engine (Google, Bing, Yahoo, DuckDuckGo). Pricing: $0.60/1k searches."
}
},
"agentcore/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "agentcore",
"mode": "search",
"metadata": {
"notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway"
}
},
"tinyfish/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "tinyfish",
@ -19527,6 +19535,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -19584,6 +19593,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.1-pro-preview-customtools": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -19738,6 +19748,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.5-flash": {
"prompt_cache_min_tokens": 4096,
"deprecation_date": "2027-05-19",
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_token": 1.5e-06,
@ -19749,6 +19760,7 @@
"mode": "chat",
"output_cost_per_reasoning_token": 9e-06,
"output_cost_per_token": 9e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19790,6 +19802,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -19804,6 +19817,7 @@
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19844,6 +19858,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -19858,6 +19873,7 @@
"output_cost_per_token": 3.75e-06,
"output_cost_per_token_batches": 1.875e-06,
"output_cost_per_token_flex": 1.875e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_endpoints": [
"/v1/chat/completions",
@ -19898,6 +19914,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -19955,6 +19972,7 @@
"web_search_billing_unit": "per_query"
},
"vertex_ai/gemini-3.1-pro-preview-customtools": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"cache_creation_input_token_cost_above_200k_tokens": 2.5e-07,
@ -20689,8 +20707,8 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.1-flash-image": {
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_batches": 1.25e-07,
"input_cost_per_token": 5e-07,
"input_cost_per_token_batches": 2.5e-07,
"litellm_provider": "gemini",
"max_input_tokens": 65536,
"max_output_tokens": 32768,
@ -20698,8 +20716,8 @@
"mode": "image_generation",
"output_cost_per_image": 0.045,
"output_cost_per_image_token": 6e-05,
"output_cost_per_token": 1.5e-06,
"output_cost_per_token_batches": 7.5e-07,
"output_cost_per_token": 3e-06,
"output_cost_per_token_batches": 1.5e-06,
"rpm": 1000,
"tpm": 4000000,
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image",
@ -20732,8 +20750,8 @@
},
"gemini/gemini-3.1-flash-image-preview": {
"deprecation_date": "2026-06-25",
"input_cost_per_token": 2.5e-07,
"input_cost_per_token_batches": 1.25e-07,
"input_cost_per_token": 5e-07,
"input_cost_per_token_batches": 2.5e-07,
"litellm_provider": "gemini",
"max_input_tokens": 65536,
"max_output_tokens": 32768,
@ -20741,8 +20759,8 @@
"mode": "image_generation",
"output_cost_per_image": 0.045,
"output_cost_per_image_token": 6e-05,
"output_cost_per_token": 1.5e-06,
"output_cost_per_token_batches": 7.5e-07,
"output_cost_per_token": 3e-06,
"output_cost_per_token_batches": 1.5e-06,
"rpm": 1000,
"tpm": 4000000,
"source": "https://ai.google.dev/gemini-api/docs/pricing#gemini-3.1-flash-image-preview",
@ -21464,6 +21482,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.5-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_audio_token": 1e-06,
"input_cost_per_token": 1.5e-06,
@ -21518,6 +21537,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -21575,6 +21595,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -21665,6 +21686,7 @@
"tpm": 800000
},
"gemini/gemini-3.1-pro-preview": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 2e-06,
@ -21722,6 +21744,7 @@
"web_search_billing_unit": "per_query"
},
"gemini/gemini-3.1-pro-preview-customtools": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 2e-07,
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
"input_cost_per_token": 2e-06,
@ -21860,6 +21883,7 @@
"supports_vision": true
},
"gemini-3.5-flash": {
"prompt_cache_min_tokens": 4096,
"deprecation_date": "2027-05-19",
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_audio_token": 1e-06,
@ -21913,6 +21937,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.6-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -21968,6 +21993,7 @@
"web_search_billing_unit": "per_query"
},
"gemini-3.7-flash": {
"prompt_cache_min_tokens": 4096,
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 3.75e-08,
"input_cost_per_token": 7.5e-07,
@ -38857,6 +38883,7 @@
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -38881,6 +38908,7 @@
"max_tokens": 8192,
"mode": "chat",
"output_cost_per_token": 5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/haiku-4-5",
"supports_assistant_prefill": true,
"supports_function_calling": true,
@ -39100,6 +39128,7 @@
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"regional_endpoint_uplift_multiplier": 1.1,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@ -39129,6 +39158,7 @@
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 2.5e-05,
"regional_endpoint_uplift_multiplier": 1.1,
"search_context_cost_per_query": {
"search_context_size_high": 0.01,
"search_context_size_low": 0.01,
@ -39149,6 +39179,7 @@
},
"vertex_ai/claude-opus-4-6": {
"deprecation_date": "2027-02-05",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39180,6 +39211,7 @@
},
"vertex_ai/claude-opus-4-6@default": {
"deprecation_date": "2027-02-05",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39211,6 +39243,7 @@
},
"vertex_ai/claude-opus-4-7": {
"deprecation_date": "2027-04-16",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39243,6 +39276,7 @@
},
"vertex_ai/claude-opus-4-7@default": {
"deprecation_date": "2027-04-16",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
@ -39275,6 +39309,7 @@
},
"vertex_ai/claude-fable-5": {
"deprecation_date": "2027-06-08",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
@ -39307,6 +39342,7 @@
},
"vertex_ai/claude-fable-5@default": {
"deprecation_date": "2027-06-08",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
@ -39339,6 +39375,7 @@
},
"vertex_ai/claude-opus-5": {
"deprecation_date": "2027-01-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39372,6 +39409,7 @@
},
"vertex_ai/claude-opus-5@default": {
"deprecation_date": "2027-01-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39405,6 +39443,7 @@
},
"vertex_ai/claude-opus-4-8": {
"deprecation_date": "2027-05-28",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39438,6 +39477,7 @@
},
"vertex_ai/claude-opus-4-8@default": {
"deprecation_date": "2027-05-28",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 6.25e-06,
@ -39487,6 +39527,7 @@
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -39500,6 +39541,7 @@
},
"vertex_ai/claude-sonnet-5": {
"deprecation_date": "2026-12-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_1hr": 4e-06,
@ -39532,6 +39574,7 @@
"prompt_cache_min_tokens": 1024
},
"vertex_ai/claude-sonnet-4-6": {
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -39579,6 +39622,7 @@
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"output_cost_per_token_batches": 7.5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"supports_assistant_prefill": true,
"supports_computer_use": true,
"supports_function_calling": true,
@ -39992,6 +40036,7 @@
"output_cost_per_token_batches": 7.5e-07,
"output_cost_per_token_flex": 7.5e-07,
"output_cost_per_token_priority": 2.7e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@ -40048,6 +40093,7 @@
"output_cost_per_token_batches": 1.25e-06,
"output_cost_per_token_flex": 1.25e-06,
"output_cost_per_token_priority": 4.5e-06,
"regional_endpoint_uplift_multiplier": 1.1,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models",
"supported_endpoints": [
"/v1/chat/completions",
@ -47219,6 +47265,7 @@
},
"vertex_ai/claude-sonnet-5@default": {
"deprecation_date": "2026-12-24",
"regional_endpoint_uplift_multiplier": 1.1,
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_1hr": 4e-06,
@ -47251,6 +47298,7 @@
"prompt_cache_min_tokens": 1024
},
"vertex_ai/claude-sonnet-4-6@default": {
"regional_endpoint_uplift_multiplier": 1.1,
"supports_adaptive_thinking": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
@ -48180,15 +48228,15 @@
},
"deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"cache_read_input_token_cost": 1.4e-08,
"input_cost_per_token": 4.4e-07,
"input_cost_per_token_cache_hit": 1.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_cost_per_token": 1.32e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
@ -48206,15 +48254,15 @@
},
"deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"cache_read_input_token_cost": 4.4e-08,
"input_cost_per_token": 1.32e-06,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"output_cost_per_token": 3.96e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
@ -48232,15 +48280,15 @@
},
"deepseek/deepseek-v4-flash": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 2.8e-09,
"input_cost_per_token": 1.4e-07,
"input_cost_per_token_cache_hit": 2.8e-09,
"cache_read_input_token_cost": 1.4e-08,
"input_cost_per_token": 4.4e-07,
"input_cost_per_token_cache_hit": 1.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_cost_per_token": 1.32e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"
@ -48258,15 +48306,15 @@
},
"deepseek/deepseek-v4-pro": {
"cache_creation_input_token_cost": 0.0,
"cache_read_input_token_cost": 3.625e-09,
"input_cost_per_token": 4.35e-07,
"input_cost_per_token_cache_hit": 3.625e-09,
"cache_read_input_token_cost": 4.4e-08,
"input_cost_per_token": 1.32e-06,
"input_cost_per_token_cache_hit": 4.4e-08,
"litellm_provider": "deepseek",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"mode": "chat",
"output_cost_per_token": 8.7e-07,
"output_cost_per_token": 3.96e-06,
"source": "https://api-docs.deepseek.com/quick_start/pricing",
"supported_endpoints": [
"/v1/chat/completions"

View file

@ -514,6 +514,11 @@
"type": "object",
"description": "Provider-internal routing hints (e.g. bedrock_invocation_schema)."
},
"regional_endpoint_uplift_multiplier": {
"type": "number",
"minimum": 1,
"description": "Multiplier applied to all token costs when served from a non-global Vertex AI endpoint (e.g. 1.10 = +10%)."
},
"regional_processing_uplift_multiplier_eu": {
"type": "number",
"minimum": 1,

View file

@ -18,7 +18,7 @@
"limit": 711
},
"ANN205": {
"limit": 113
"limit": 112
},
"ANN206": {
"limit": 133

View file

@ -947,6 +947,17 @@ model LiteLLM_DailyTagSpend {
}
// One row per live proxy worker process. Workers upsert their row on a fixed
// heartbeat; counting rows with a recent heartbeat tells how many workers share
// this database, which lets the Admin UI hide its "no Redis" warning for
// deployments that are provably a single worker.
model LiteLLM_ProxyWorkerHeartbeat {
worker_id String @id
hostname String
started_at DateTime @default(now())
last_heartbeat_at DateTime @default(now())
}
// Track the status of cron jobs running. Only allow one pod to run the job at a time
model LiteLLM_CronJob {
cronjob_id String @id @default(cuid()) // Unique ID for the record

View file

@ -75,11 +75,11 @@ Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure cover
`E2E_FIXTURE_MODE` selects the transport every client is built on: `live` (the default, and what an unset variable means: nothing changes), `record` (run against the live proxy and write every interaction to a fixture bundle), or `replay` (serve every interaction back from the bundle with no HTTP at all, so a replay run needs no proxy and cannot bill a provider). The seam is `select_transport` in `fixture_transport.py`, applied inside `build_proxy_client`; both transports fulfil the same `Transport` protocol, so no test or client changes shape in any mode
A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per transport call in call order (`0000-post-chat-completions.json`). Auth header values are redacted on write, and file uploads store a sha256 digest instead of the bytes; response bodies are stored verbatim (a /key/generate response keeps the ephemeral virtual key it minted), which is part of why bundles are gitignored. `fixture_bundle.py` owns the format
A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per transport call in call order (`0000-post-chat-completions.json`). Auth header values and credential request fields (`api_key`, `*_secret_key`, `static_headers`, and the like; the list is `fixture_canonical.py`'s) are redacted on write, and file uploads store a sha256 digest instead of the bytes; response bodies are stored verbatim (a /key/generate response keeps the ephemeral virtual key it minted), which is part of why bundles are gitignored. `fixture_bundle.py` owns the format
Replay matches calls per test by transport verb and path in recorded order and raises `ReplayMiss` on any drift, naming the recorded and the actual call; a passed test must also consume its whole recording, or teardown fails it naming the first leftover interaction. Either way the fix is always to re-record with `E2E_FIXTURE_MODE=record`. Record starts fresh every time: it wipes the previous bundle (refusing to wipe a directory that is not a bundle) and never reads it. A replay bundle whose manifest is older than seven days hard-fails at collection time naming the bundle's age, so replay can never certify against fixtures that have drifted more than a week from the live proxy
Replay matches calls per test by canonical key: `fixture_canonical.py` canonicalizes the recorded request (volatile headers and credential fields out, unique markers, generated ids, uuids, and timestamps replaced with fixed placeholders, object keys sorted) and the key is the method, path, and a content hash, so identity survives re-records and machine changes while any real content drift is a `ReplayMiss` that names the computed key, the closest recorded key with its file, and a content diff, and never falls through to a live call. Matching is order-independent across distinct keys (concurrent calls may interleave) and FIFO within one key (a poll loop replays its responses in recorded order); a passed test must also consume its whole recording, or teardown fails it naming a leftover key. Either way the fix is always to re-record with `E2E_FIXTURE_MODE=record`. Every rewrite rule lives in `fixture_canonical.py`, so a new volatile header, credential field name, or generated-id shape is one edit there. Record starts fresh every time: it wipes the previous bundle (refusing to wipe a directory that is not a bundle) and never reads it. A replay bundle whose manifest is older than seven days hard-fails at collection time naming the bundle's age, so replay can never certify against fixtures that have drifted more than a week from the live proxy
Deliberately not here yet: canonical content-based match keys (LIT-5741), streaming chunk fidelity (LIT-5742), and scoping record/replay to provider-bound traffic (LIT-5745)
Deliberately not here yet: streaming chunk fidelity (LIT-5742) and scoping record/replay to provider-bound traffic (LIT-5745)
## Typing

View file

@ -2,12 +2,13 @@
from __future__ import annotations
import time
from dataclasses import dataclass
from pydantic import BaseModel, ValidationError
from proxy_client import ProxyClient
from e2e_http import StreamingResponse
from e2e_http import NoBody, StreamingResponse, is_ok, unwrap
from models import (
ChatBody,
ChatMessage,
@ -15,9 +16,16 @@ from models import (
LiteLLMParamsBody,
ModelInfoBody,
ModelNewBody,
TeamDeleteBody,
TeamInfoParams,
TeamInfoResponse,
TeamNewBody,
TeamNewResponse,
TeamUpdateBody,
)
MODEL_ACCESS_DENIED_MARKER = "key_model_access_denied"
TEAM_MODEL_ACCESS_DENIED_MARKER = "team_model_access_denied"
ROUTE_NOT_ALLOWED_MARKER = "not allowed to call this route"
@ -31,6 +39,14 @@ class ApiErrorEnvelope(BaseModel):
error: ApiErrorDetail
class AccessGroupInfoResponse(BaseModel):
"""GET /access_group/{name}/info: the deployments a model access group grants."""
access_group: str
model_names: list[str]
deployment_count: int
def error_envelope(body: str) -> ApiErrorEnvelope | None:
"""The OpenAI-shaped `{"error": {...}}` a client parses, or None if absent."""
try:
@ -51,15 +67,75 @@ class AccessControlClient:
def delete_key(self, key: str) -> None:
self.proxy.delete_key(key)
def chat_status(self, key: str, model: str, content: str) -> StreamingResponse:
def chat_status(
self, key: str, model: str, content: str, max_completion_tokens: int | None = None
) -> StreamingResponse:
return self.proxy.transport.send(
"/chat/completions",
headers=self.proxy.transport.bearer(key),
json=ChatBody(
model=model, messages=[ChatMessage(role="user", content=content)]
model=model,
messages=[ChatMessage(role="user", content=content)],
max_completion_tokens=max_completion_tokens,
),
)
def create_team(self, team_alias: str, models: list[str]) -> str:
team_id = unwrap(
self.proxy.transport.post(
"/team/new",
headers=self.proxy.transport.master,
json=TeamNewBody(team_alias=team_alias, models=models),
response_type=TeamNewResponse,
)
).team_id
self._await_team(team_id)
return team_id
def set_team_models(self, team_id: str, team_alias: str, models: list[str]) -> None:
"""Replace the team's allow-list. /model/new appends a team-scoped deployment's
public name to it, so a test that means to grant only an access group has to
put the allow-list back afterwards."""
_ = unwrap(
self.proxy.transport.post(
"/team/update",
headers=self.proxy.transport.master,
json=TeamUpdateBody(team_id=team_id, team_alias=team_alias, models=models),
response_type=NoBody,
)
)
def delete_team(self, team_id: str) -> None:
_ = self.proxy.transport.post(
"/team/delete",
headers=self.proxy.transport.master,
json=TeamDeleteBody(team_ids=[team_id]),
response_type=NoBody,
)
def access_group_info(self, access_group: str) -> AccessGroupInfoResponse | None:
result = self.proxy.transport.get(
f"/access_group/{access_group}/info",
headers=self.proxy.transport.master,
params=NoBody(),
response_type=AccessGroupInfoResponse,
)
return unwrap(result) if is_ok(result) else None
def _await_team(self, team_id: str) -> None:
deadline = time.monotonic() + self.proxy.poll_timeout
while time.monotonic() < deadline:
result = self.proxy.transport.get(
"/team/info",
headers=self.proxy.transport.master,
params=TeamInfoParams(team_id=team_id),
response_type=TeamInfoResponse,
)
if is_ok(result):
return
time.sleep(self.proxy.poll_interval)
raise AssertionError(f"/team/info never resolved team {team_id!r} created by /team/new")
def create_model_status(self, key: str, model_name: str) -> StreamingResponse:
return self.proxy.transport.send(
"/model/new",

View file

@ -0,0 +1,277 @@
"""Live e2e: a model access group as the grant on a key and on a team.
Whoever holds the group can call every deployment in it and nothing else, whether
the request names a deployment exactly, names a model that a wildcard deployment
in the group covers, or spells that model with its provider prefix. The bare-name
spelling is the LIT-5813 regression: the group-membership lookup skipped the
provider-prefix retry every other model-resolution path performs, so a group
holding `openai/gpt-5.4*` denied `gpt-5.4-nano` while allowing `openai/gpt-5.4-nano`.
"""
from __future__ import annotations
import os
import time
from collections.abc import Callable, Iterator
from dataclasses import dataclass
from typing import Final
import pytest
from access_control_client import (
AccessControlClient,
MODEL_ACCESS_DENIED_MARKER,
TEAM_MODEL_ACCESS_DENIED_MARKER,
)
from e2e_config import unique_marker
from lifecycle import ResourceManager
from models import (
ChatResponse,
KeyGenerateBody,
LiteLLMParamsBody,
ModelInfoBody,
ModelNewBody,
)
pytestmark = pytest.mark.e2e
WILDCARD_PATTERN: Final = "openai/gpt-5.4*"
WILDCARD_BARE_MODEL: Final = "gpt-5.4-nano"
WILDCARD_PREFIXED_MODEL: Final = "openai/gpt-5.4-nano"
GROUP_BACKEND: Final = "openai/gpt-5.4-nano"
UNCOVERED_OPENAI_MODEL: Final = "gpt-5.2"
TEAM_WILDCARD_PATTERN: Final = "openai/gpt-5.6*"
TEAM_WILDCARD_BARE_MODEL: Final = "gpt-5.6-luna"
MAX_COMPLETION_TOKENS: Final = 16
PROMPT: Final = "Reply with exactly: OK"
@dataclass(frozen=True, slots=True)
class GroupedDeployments:
"""A wildcard deployment and an exactly-named one inside `access_group`, plus a
deployment left out of it."""
access_group: str
member_model: str
outsider_model: str
@dataclass(frozen=True, slots=True)
class TeamGrant:
"""A team whose whole allow-list is `access_group`, holding one team-scoped
wildcard deployment, and a key that belongs to it."""
access_group: str
team_id: str
key: str
ModelSelector = Callable[[GroupedDeployments], str]
ALLOWED: Final[tuple[tuple[str, ModelSelector], ...]] = (
("bare name the group's wildcard covers", lambda grouped: WILDCARD_BARE_MODEL),
("provider-prefixed name the group's wildcard covers", lambda grouped: WILDCARD_PREFIXED_MODEL),
("exactly-named deployment in the group", lambda grouped: grouped.member_model),
)
DENIED: Final[tuple[tuple[str, ModelSelector], ...]] = (
("deployment outside the group", lambda grouped: grouped.outsider_model),
("provider model outside the group's wildcard", lambda grouped: UNCOVERED_OPENAI_MODEL),
("name no provider claims", lambda grouped: f"e2e-ag-unknown-{unique_marker()}"),
)
def _provider_key(env_var: str) -> str:
return os.environ.get(env_var) or f"os.environ/{env_var}"
def _grouped_model(model_name: str, backend: str, access_groups: list[str] | None) -> ModelNewBody:
return ModelNewBody(
model_name=model_name,
litellm_params=LiteLLMParamsBody(model=backend, api_key=_provider_key("OPENAI_API_KEY")),
model_info=ModelInfoBody(access_groups=access_groups),
)
def _await_group_members(client: AccessControlClient, access_group: str, expected: frozenset[str]) -> None:
"""The grant under test is the group's membership, so prove the proxy recorded it
before asserting on what the group lets through."""
deadline = time.monotonic() + client.proxy.poll_timeout
listed: list[str] = []
while time.monotonic() < deadline:
info = client.access_group_info(access_group)
listed = info.model_names if info is not None else []
if expected.issubset(listed):
return
time.sleep(client.proxy.poll_interval)
pytest.fail(
f"/access_group/{access_group}/info never listed {sorted(expected)} as members; last read {listed}"
)
def _await_team_allowlist(client: AccessControlClient, grant_key: str, access_group: str) -> None:
"""Registering a team-scoped deployment appends its public name to the team's
allow-list, and a wildcard sitting there directly would grant the model under test
on its own. Poll a denial until the message enumerates the allow-list the test
means to exercise: the group, and nothing else."""
allowlist: Final = f"models=['{access_group}']"
deadline = time.monotonic() + client.proxy.poll_timeout
body = ""
while time.monotonic() < deadline:
body = client.chat_status(
grant_key, UNCOVERED_OPENAI_MODEL, f"{PROMPT} {unique_marker()}", MAX_COMPLETION_TOKENS
).body
if allowlist in body:
return
time.sleep(client.proxy.poll_interval)
pytest.fail(f"the team's allow-list never settled to {allowlist}; last denial read {body[:300]}")
@pytest.fixture(scope="module")
def grouped(client: AccessControlClient) -> Iterator[GroupedDeployments]:
marker: Final = unique_marker()
deployments: Final = GroupedDeployments(
access_group=f"e2e-ag-{marker}",
member_model=f"e2e-ag-member-{marker}",
outsider_model=f"e2e-ag-outsider-{marker}",
)
registrations: Final = (
_grouped_model(WILDCARD_PATTERN, WILDCARD_PATTERN, [deployments.access_group]),
_grouped_model(deployments.member_model, GROUP_BACKEND, [deployments.access_group]),
_grouped_model(deployments.outsider_model, GROUP_BACKEND, None),
)
created: Final = tuple(client.proxy.register_model(body) for body in registrations)
try:
_await_group_members(
client,
deployments.access_group,
frozenset({WILDCARD_PATTERN, deployments.member_model}),
)
yield deployments
finally:
for model_id in created:
client.proxy.delete_model(model_id)
@pytest.fixture(scope="module")
def team_grant(client: AccessControlClient) -> Iterator[TeamGrant]:
marker: Final = unique_marker()
access_group: Final = f"e2e-agt-{marker}"
team_alias: Final = f"e2e-ag-team-{marker}"
team_id: Final = client.create_team(team_alias, [access_group])
key: Final = client.proxy.generate_key(KeyGenerateBody(models=[], team_id=team_id))
model_id: Final = client.proxy.register_model(
ModelNewBody(
model_name=TEAM_WILDCARD_PATTERN,
litellm_params=LiteLLMParamsBody(
model=TEAM_WILDCARD_PATTERN, api_key=_provider_key("OPENAI_API_KEY")
),
model_info=ModelInfoBody(team_id=team_id, access_groups=[access_group]),
),
listed_for=key,
)
client.set_team_models(team_id, team_alias, [access_group])
try:
_await_team_allowlist(client, key, access_group)
yield TeamGrant(access_group=access_group, team_id=team_id, key=key)
finally:
client.proxy.delete_model(model_id)
client.proxy.delete_key(key)
client.delete_team(team_id)
class TestKeyScopedToAccessGroup:
@pytest.mark.covers(
"other.auth.model_access_group.wildcard_bare_name_allowed",
"other.auth.model_access_group.member_allowed",
)
@pytest.mark.parametrize(("case", "select_model"), ALLOWED)
def test_group_grants_every_deployment_in_it(
self,
case: str,
select_model: ModelSelector,
client: AccessControlClient,
resources: ResourceManager,
grouped: GroupedDeployments,
) -> None:
key = resources.key(models=[grouped.access_group])
model = select_model(grouped)
result = client.chat_status(
key, model, f"{PROMPT} {unique_marker()}", MAX_COMPLETION_TOKENS
)
assert result.status_code == 200, (
f"a key holding access group {grouped.access_group!r} must be able to call "
f"{model!r} ({case}), got {result.status_code}: {result.body[:300]}"
)
assert ChatResponse.model_validate_json(result.body).choices, (
f"200 must carry a real completion, not an error envelope: {result.body[:300]}"
)
@pytest.mark.covers("other.auth.model_access_group.non_member_denied")
@pytest.mark.parametrize(("case", "select_model"), DENIED)
def test_group_grants_nothing_outside_it(
self,
case: str,
select_model: ModelSelector,
client: AccessControlClient,
resources: ResourceManager,
grouped: GroupedDeployments,
) -> None:
key = resources.key(models=[grouped.access_group])
model = select_model(grouped)
result = client.chat_status(
key, model, f"{PROMPT} {unique_marker()}", MAX_COMPLETION_TOKENS
)
assert result.status_code == 403, (
f"a key holding only access group {grouped.access_group!r} must be denied 403 on "
f"{model!r} ({case}), got {result.status_code}: {result.body[:300]}"
)
assert MODEL_ACCESS_DENIED_MARKER in result.body, (
f"403 body must be a key model-access denial, got: {result.body[:300]}"
)
class TestTeamScopedToAccessGroup:
@pytest.mark.covers("other.auth.model_access_group.team_wildcard_bare_name_allowed")
def test_group_grants_the_teams_own_wildcard(
self, client: AccessControlClient, team_grant: TeamGrant
) -> None:
result = client.chat_status(
team_grant.key,
TEAM_WILDCARD_BARE_MODEL,
f"{PROMPT} {unique_marker()}",
MAX_COMPLETION_TOKENS,
)
assert result.status_code == 200, (
f"a team whose allow-list is access group {team_grant.access_group!r} must be able to "
f"call {TEAM_WILDCARD_BARE_MODEL!r} through its team-scoped {TEAM_WILDCARD_PATTERN!r} "
f"deployment, got {result.status_code}: {result.body[:300]}"
)
assert ChatResponse.model_validate_json(result.body).choices, (
f"200 must carry a real completion, not an error envelope: {result.body[:300]}"
)
@pytest.mark.covers("other.auth.model_access_group.team_non_member_denied")
def test_group_grants_the_team_nothing_outside_it(
self, client: AccessControlClient, team_grant: TeamGrant
) -> None:
model = f"e2e-ag-unknown-{unique_marker()}"
result = client.chat_status(
team_grant.key, model, f"{PROMPT} {unique_marker()}", MAX_COMPLETION_TOKENS
)
assert result.status_code == 403, (
f"a team holding only access group {team_grant.access_group!r} must be denied 403 on "
f"{model!r}, got {result.status_code}: {result.body[:300]}"
)
assert TEAM_MODEL_ACCESS_DENIED_MARKER in result.body, (
f"403 body must be a team model-access denial, got: {result.body[:300]}"
)

View file

@ -12,6 +12,11 @@
- {id: other.auth.jwt.valid_token_allows, module: other, tier: P0, area: auth, assertions: [valid_token_allows], source: "handle_jwt.py:77-150", rationale: "Valid JWT with correct issuer + claims grants access"}
- {id: other.auth.jwt.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "handle_jwt.py:125-135", rationale: "Expired JWT rejected even with valid signature"}
- {id: other.auth.jwt.invalid_signature_denied, module: other, tier: P0, area: auth, assertions: [invalid_signature_denied], source: "handle_jwt.py:145-150", rationale: "Bad/missing signature fails verification"}
- {id: other.auth.model_access_group.wildcard_bare_name_allowed, module: other, tier: P0, area: auth, assertions: [wildcard_bare_name_allowed], source: "auth_checks.py:3232 / LIT-5813", fail_before_fix: proven, rationale: "A grant of a group holding a wildcard deployment covers the bare model names callers actually send, not only the provider-prefixed spelling"}
- {id: other.auth.model_access_group.member_allowed, module: other, tier: P0, area: auth, assertions: [member_allowed], source: "auth_checks.py:3232", rationale: "A key whose allow-list is a model access group can call the deployments in that group"}
- {id: other.auth.model_access_group.non_member_denied, module: other, tier: P0, area: auth, assertions: [non_member_denied], source: "auth_checks.py:3232", rationale: "That same grant reaches nothing outside the group, including provider models the group's wildcard does not cover"}
- {id: other.auth.model_access_group.team_wildcard_bare_name_allowed, module: other, tier: P1, area: auth, assertions: [team_wildcard_bare_name_allowed], source: "auth_checks.py:3232 / LIT-5813", fail_before_fix: proven, rationale: "The same bare-name grant holds when the wildcard deployment is team-scoped and the team's allow-list is the group"}
- {id: other.auth.model_access_group.team_non_member_denied, module: other, tier: P1, area: auth, assertions: [team_non_member_denied], source: "auth_checks.py:3232", rationale: "A team-level group grant reaches nothing outside the group"}
- {id: other.auth.virtual_key.route_permission_enforced, module: other, tier: P0, area: auth, assertions: [route_permission_enforced], source: "route_checks.py:89-151", rationale: "allowed_routes whitelist denies disallowed routes"}
- {id: other.auth.virtual_key.route_group_allowed, module: other, tier: P1, area: auth, assertions: [route_group_allowed], source: "route_checks.py:106-128", rationale: "allowed_routes=[llm_api_routes] grants all LLM endpoints"}
- {id: other.auth.passthrough.model_allowlist_enforced, module: other, tier: P1, area: auth, assertions: [model_allowlist_enforced], source: "route_checks.py:135-151", rationale: "Passthrough enforces per-key model allow-lists"}

View file

@ -8,10 +8,10 @@ green replay run can never certify against fixtures that have drifted more than
a week from the live proxy.
This module owns the format only. The transports that produce and consume it
live in fixture_transport.py; canonical request matching, streaming chunk
fidelity, and provider-scoping are follow-ups (LIT-5741/5742/5745) and are
deliberately absent here, which is why every interaction file stores the full
redacted request even though replay today matches by call order.
live in fixture_transport.py and the canonical match keys they compute live in
fixture_canonical.py (LIT-5741); streaming chunk fidelity and provider-scoping
are follow-ups (LIT-5742/5745). Every interaction file stores the full redacted
request because replay matches on its canonicalized content.
"""
from __future__ import annotations
@ -54,12 +54,13 @@ class Manifest(BaseModel):
class RecordedRequest(BaseModel):
"""The request as the transport saw it, auth header values redacted.
"""The request as the transport saw it, auth header values and credential
body/form fields redacted.
Replay today only matches ``method`` (the transport verb, not the HTTP verb)
and ``path`` in call order; the rest is stored so LIT-5741 can move to
content-based match keys without re-recording. File uploads store a content
digest instead of the bytes."""
Replay matches on the canonical content key fixture_canonical.py computes
over ``method`` (the transport verb, not the HTTP verb), ``path``, and the
canonicalized headers, params, body, form, and file identity. File uploads
store a content digest instead of the bytes."""
method: str
path: str

View file

@ -0,0 +1,150 @@
"""Canonical request identity for replay matching (LIT-5741).
Matching a replayed call against the raw recorded request never hits: unique
markers salt prompts, model names, and tags; every run mints fresh virtual
keys; request ids and timestamps differ on every call. Matching on transport
verb + path alone collides: two different requests to the same route silently
swap responses, which passes when it should miss. The canonicalizer strips
exactly the volatile material (volatile headers, credential fields, markers,
generated ids, timestamps) and hashes what remains with sorted object keys, so
identity is content-based and stable across runs and machines.
Every rewrite rule lives in this module, next to the transports that apply it:
a new volatile header, credential field name, or generated-id shape is one
edit here, never a per-suite change.
"""
from __future__ import annotations
import hashlib
import json
import re
from dataclasses import dataclass
from functools import reduce
from typing import Final
from pydantic import JsonValue
from fixture_bundle import RecordedRequest
VOLATILE_HEADER_NAMES: Final[frozenset[str]] = frozenset(
{
"authorization",
"x-litellm-api-key",
"x-api-key",
"x-goog-api-key",
"x-request-id",
"traceparent",
"tracestate",
}
)
SECRET_FIELD_NAMES: Final[frozenset[str]] = frozenset(
{"api_key", "aws_access_key_id", "static_headers", "vertex_credentials"}
)
SECRET_FIELD_SUFFIXES: Final[tuple[str, ...]] = (
"_api_key",
"_secret_key",
"_secret_access_key",
"_session_token",
"_credentials",
"_password",
)
SECRET_PLACEHOLDER: Final = "<secret>"
PLACEHOLDER_RULES: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
(re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{64}(?![0-9a-fA-F])"), "<sha256>"),
(
re.compile(r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}"),
"<uuid>",
),
(re.compile(r"sk-[A-Za-z0-9_-]{16,}"), "<key>"),
(
re.compile(r"\d{4}-\d{2}-\d{2}[T ]\d{2}:\d{2}:\d{2}(?:\.\d+)?(?:Z|[+-]\d{2}:?\d{2})?"),
"<timestamp>",
),
(re.compile(r"(?<!\d)\d{4}-\d{2}-\d{2}(?!\d)"), "<date>"),
(
re.compile(r"\b(?:chatcmpl|msgbatch|msg|resp|batch|call|req|ftjob|gen|file)[-_][A-Za-z0-9]{8,}\b"),
"<id>",
),
(re.compile(r"(?<![0-9a-fA-F])[0-9a-f]{12}(?![0-9a-fA-F])"), "<marker>"),
)
def is_secret_field(name: str) -> bool:
lowered: Final = name.lower()
return lowered in SECRET_FIELD_NAMES or lowered.endswith(SECRET_FIELD_SUFFIXES)
def canonical_string(value: str) -> str:
return reduce(lambda acc, rule: rule[0].sub(rule[1], acc), PLACEHOLDER_RULES, value)
def _canonical_flat(fields: dict[str, str]) -> dict[str, JsonValue]:
return {
key: SECRET_PLACEHOLDER if is_secret_field(key) else canonical_string(value)
for key, value in fields.items()
}
def _canonical_value(value: JsonValue) -> JsonValue:
match value:
case str():
return canonical_string(value)
case dict():
return {
key: SECRET_PLACEHOLDER
if is_secret_field(key) and item is not None
else _canonical_value(item)
for key, item in value.items()
}
case list():
return [_canonical_value(item) for item in value]
case _:
return value
@dataclass(frozen=True, slots=True)
class CanonicalRequest:
method: str
path: str
content: str
@property
def key(self) -> str:
digest: Final = hashlib.sha256(
f"{self.method} {self.path}\n{self.content}".encode()
).hexdigest()[:16]
return f"{self.method} {self.path} #{digest}"
def pretty_content(self) -> str:
return json.dumps(json.loads(self.content), indent=2, sort_keys=True)
def canonicalize(request: RecordedRequest) -> CanonicalRequest:
file_identity: Final[JsonValue | None] = (
None
if request.file_name is None and request.file_sha256 is None
else {
"name": None if request.file_name is None else canonical_string(request.file_name),
"sha256": request.file_sha256,
"bytes": request.file_bytes,
}
)
content: Final[dict[str, JsonValue]] = {
"headers": {
name.lower(): canonical_string(value)
for name, value in request.headers.items()
if name.lower() not in VOLATILE_HEADER_NAMES
},
"params": _canonical_flat(request.params),
"body": _canonical_value(request.body),
"form": None if request.form is None else _canonical_flat(request.form),
"file": file_identity,
}
return CanonicalRequest(
method=request.method,
path=canonical_string(request.path),
content=json.dumps(content, sort_keys=True, separators=(",", ":")),
)

View file

@ -7,23 +7,30 @@ HTTP, no proxy, no provider spend. Because both fulfil ``Transport``, no test
or client changes shape; ``build_proxy_client`` picks the transport from
``E2E_FIXTURE_MODE`` (live | record | replay, default live).
Replay matches each call by test node id and call order, verifying transport
verb + path and failing hard on any drift (``ReplayMiss``). Canonical
content-based match keys are LIT-5741; streaming chunk fidelity is LIT-5742;
scoping record/replay to provider-bound traffic is LIT-5745.
Replay matches each call by test node id and canonical content key
(fixture_canonical.py, LIT-5741): volatile headers, credential fields, unique
markers, generated ids, and timestamps are canonicalized out before hashing, so
matching is order-independent across distinct keys, FIFO within a key, and a
miss fails hard (``ReplayMiss``) printing the computed key and the closest
recorded key without ever falling through to a live call. Streaming chunk
fidelity is LIT-5742; scoping record/replay to provider-bound traffic is
LIT-5745.
"""
from __future__ import annotations
import difflib
import functools
import hashlib
import os
from collections import deque
from dataclasses import dataclass, field
from datetime import datetime
from itertools import islice
from pathlib import Path
from typing import Final, Literal, assert_never
from pydantic import BaseModel
from pydantic import BaseModel, JsonValue
from e2e_http import AuthHeaders, BinaryStream, ProbeResult, Result, StreamingResponse
from fixture_bundle import (
@ -43,12 +50,14 @@ from fixture_bundle import (
check_freshness,
format_age,
from_result,
interaction_filename,
load_bundle,
prepare_bundle,
slug_for_test,
to_json_value,
to_result,
)
from fixture_canonical import CanonicalRequest, canonicalize, is_secret_field
from transport import Transport
type FixtureMode = Literal["live", "record", "replay"]
@ -118,6 +127,27 @@ def _redact(headers: dict[str, str]) -> dict[str, str]:
}
def _redact_secret_fields(value: JsonValue) -> JsonValue:
match value:
case dict():
return {
key: REDACTED_VALUE
if is_secret_field(key) and item is not None
else _redact_secret_fields(item)
for key, item in value.items()
}
case list():
return [_redact_secret_fields(item) for item in value]
case _:
return value
def _redact_flat(fields: dict[str, str]) -> dict[str, str]:
return {
key: REDACTED_VALUE if is_secret_field(key) else value for key, value in fields.items()
}
def recorded_request(
method: str,
path: str,
@ -133,9 +163,9 @@ def recorded_request(
method=method,
path=path,
headers=_redact(_dump_flat(headers)),
params=_dump_flat(params),
body=None if body is None else to_json_value(body),
form=None if form is None else _dump_flat(form),
params=_redact_flat(_dump_flat(params)),
body=None if body is None else _redact_secret_fields(to_json_value(body)),
form=None if form is None else _redact_flat(_dump_flat(form)),
file_name=file_name,
file_sha256=None if file_content is None else hashlib.sha256(file_content).hexdigest(),
file_bytes=None if file_content is None else len(file_content),
@ -304,46 +334,112 @@ class RecordingTransport:
return response
def _build_pool(recorded: tuple[Interaction, ...]) -> dict[str, deque[Interaction]]:
keys: Final = tuple(canonicalize(interaction.request).key for interaction in recorded)
return {
key: deque(
interaction
for candidate_key, interaction in zip(keys, recorded, strict=True)
if candidate_key == key
)
for key in dict.fromkeys(keys)
}
def _closest_recorded(
canonical: CanonicalRequest, recorded: tuple[Interaction, ...]
) -> tuple[CanonicalRequest, str]:
candidates: Final = tuple(canonicalize(interaction.request) for interaction in recorded)
ratios: Final = tuple(
difflib.SequenceMatcher(
None, f"{canonical.method} {canonical.path}\n{canonical.content}",
f"{candidate.method} {candidate.path}\n{candidate.content}",
).ratio()
for candidate in candidates
)
best: Final = max(range(len(candidates)), key=lambda index: ratios[index])
return candidates[best], interaction_filename(best, recorded[best].request)
def _miss_message(test_key: str, slug: str, canonical: CanonicalRequest, bundle: LoadedBundle) -> str:
recorded: Final = bundle.interactions.get(slug, ())
if not recorded:
return (
f"replay miss for {test_key}: computed key {canonical.key} but nothing is recorded "
f"under {slug}; re-record with E2E_FIXTURE_MODE=record"
)
closest, closest_file = _closest_recorded(canonical, recorded)
diff: Final = "\n".join(
islice(
difflib.unified_diff(
closest.pretty_content().splitlines(),
canonical.pretty_content().splitlines(),
fromfile=f"closest recorded ({closest_file})",
tofile="test made",
lineterm="",
),
60,
)
)
return (
f"replay miss for {test_key}: no recorded interaction matches key {canonical.key}; "
f"closest recorded key is {closest.key} ({closest_file})\n{diff}\n"
"re-record with E2E_FIXTURE_MODE=record"
)
@dataclass(slots=True)
class ReplaySource:
"""One shared cursor set over a loaded bundle, so every client built in the
session consumes the same recorded sequence per test."""
"""One shared pool per test over a loaded bundle, so every client built in
the session consumes the same recorded interactions. Every pool is built
once at construction and per-key consumption is a single atomic deque pop,
so concurrent replay calls never race. Calls match by canonical content
key: order-independent across distinct keys (concurrent tests interleave
calls nondeterministically), FIFO within one key (a poll loop replays its
recorded responses in recorded order)."""
bundle: LoadedBundle
_cursors: dict[str, int] = field(default_factory=dict)
_pools: dict[str, dict[str, deque[Interaction]]] = field(init=False)
def next_interaction(self, method: str, path: str) -> Interaction:
test_key = current_test_key()
slug = slug_for_test(test_key)
recorded = self.bundle.interactions.get(slug, ())
index = self._cursors.get(slug, 0)
if index >= len(recorded):
def __post_init__(self) -> None:
self._pools = {
slug: _build_pool(recorded) for slug, recorded in self.bundle.interactions.items()
}
def _pool(self, slug: str) -> dict[str, deque[Interaction]]:
return self._pools.get(slug, {})
def next_interaction(self, request: RecordedRequest) -> Interaction:
test_key: Final = current_test_key()
slug: Final = slug_for_test(test_key)
pool: Final = self._pool(slug)
canonical: Final = canonicalize(request)
queue: Final = pool.get(canonical.key)
if queue is None:
raise ReplayMiss(_miss_message(test_key, slug, canonical, self.bundle))
try:
return queue.popleft()
except IndexError:
raise ReplayMiss(
f"replay exhausted for {test_key}: call #{index + 1} ({method} {path}) has no recorded "
f"interaction ({len(recorded)} recorded under {slug}); re-record with E2E_FIXTURE_MODE=record"
)
interaction = recorded[index]
if interaction.request.method != method or interaction.request.path != path:
raise ReplayMiss(
f"replay mismatch for {test_key} at call #{index + 1}: recorded "
f"{interaction.request.method} {interaction.request.path}, test made {method} {path}; "
"re-record with E2E_FIXTURE_MODE=record"
)
self._cursors[slug] = index + 1
return interaction
f"replay exhausted for {test_key}: every recorded interaction for key "
f"{canonical.key} is already consumed; re-record with E2E_FIXTURE_MODE=record"
) from None
def leftover_error(self, test_key: str) -> str | None:
"""Non-None when the test consumed fewer interactions than were recorded,
meaning a passing replay proved less than the bundle claims."""
slug = slug_for_test(test_key)
recorded = self.bundle.interactions.get(slug, ())
consumed = self._cursors.get(slug, 0)
if consumed >= len(recorded):
slug: Final = slug_for_test(test_key)
recorded: Final = self.bundle.interactions.get(slug, ())
if not recorded:
return None
leftover: Final = tuple(
interaction for queue in self._pool(slug).values() for interaction in queue
)
if not leftover:
return None
pending = recorded[consumed]
return (
f"replay incomplete for {test_key}: {len(recorded) - consumed} of {len(recorded)} recorded "
f"interactions never consumed, next is {pending.request.method} {pending.request.path}; "
f"replay incomplete for {test_key}: {len(leftover)} of {len(recorded)} recorded "
f"interactions never consumed, e.g. {canonicalize(leftover[0].request).key}; "
"re-record with E2E_FIXTURE_MODE=record"
)
@ -386,7 +482,12 @@ class ReplayTransport:
def post[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return to_result(_expect_result(self.source.next_interaction("post", path)), response_type)
return to_result(
_expect_result(
self.source.next_interaction(recorded_request("post", path, headers=headers, body=json))
),
response_type,
)
def get[R: BaseModel](
self,
@ -397,7 +498,12 @@ class ReplayTransport:
response_type: type[R],
timeout: float | None = None,
) -> Result[R]:
return to_result(_expect_result(self.source.next_interaction("get", path)), response_type)
return to_result(
_expect_result(
self.source.next_interaction(recorded_request("get", path, headers=headers, params=params))
),
response_type,
)
def delete[R: BaseModel](
self,
@ -408,25 +514,46 @@ class ReplayTransport:
response_type: type[R],
params: BaseModel | None = None,
) -> Result[R]:
return to_result(_expect_result(self.source.next_interaction("delete", path)), response_type)
return to_result(
_expect_result(
self.source.next_interaction(
recorded_request("delete", path, headers=headers, body=json, params=params)
)
),
response_type,
)
def patch[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return to_result(_expect_result(self.source.next_interaction("patch", path)), response_type)
return to_result(
_expect_result(
self.source.next_interaction(recorded_request("patch", path, headers=headers, body=json))
),
response_type,
)
def put[R: BaseModel](
self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R]
) -> Result[R]:
return to_result(_expect_result(self.source.next_interaction("put", path)), response_type)
return to_result(
_expect_result(
self.source.next_interaction(recorded_request("put", path, headers=headers, body=json))
),
response_type,
)
def stream(self, path: str, *, headers: BaseModel, json: BaseModel) -> StreamingResponse:
return _expect_streaming(self.source.next_interaction("stream", path))
return _expect_streaming(
self.source.next_interaction(recorded_request("stream", path, headers=headers, body=json))
)
def stream_binary(
self, path: str, *, headers: BaseModel, json: BaseModel, chunk_size: int = 8192
) -> BinaryStream:
interaction = self.source.next_interaction("stream_binary", path)
interaction = self.source.next_interaction(
recorded_request("stream_binary", path, headers=headers, body=json)
)
match interaction.response:
case RecordedBinary(payload=payload):
return payload
@ -444,10 +571,16 @@ class ReplayTransport:
params: BaseModel | None = None,
stream: bool = False,
) -> StreamingResponse:
return _expect_streaming(self.source.next_interaction("send", path))
return _expect_streaming(
self.source.next_interaction(
recorded_request("send", path, headers=headers, body=json, params=params)
)
)
def probe(self, path: str, *, params: BaseModel) -> ProbeResult:
interaction = self.source.next_interaction("probe", path)
interaction = self.source.next_interaction(
recorded_request("probe", path, headers=self.master, params=params)
)
match interaction.response:
case RecordedProbe(payload=payload):
return payload
@ -467,10 +600,27 @@ class ReplayTransport:
params: BaseModel | None = None,
response_type: type[R],
) -> Result[R]:
return to_result(_expect_result(self.source.next_interaction("upload", path)), response_type)
return to_result(
_expect_result(
self.source.next_interaction(
recorded_request(
"upload",
path,
headers=headers,
params=params,
form=form,
file_name=filename,
file_content=content,
)
)
),
response_type,
)
def download(self, path: str, *, headers: BaseModel) -> StreamingResponse:
return _expect_streaming(self.source.next_interaction("download", path))
return _expect_streaming(
self.source.next_interaction(recorded_request("download", path, headers=headers))
)
@functools.lru_cache(maxsize=8)

View file

@ -766,6 +766,8 @@ class ModelInfoBody(BaseModel):
# constraint when a prior run's teardown had not removed the row.
id: str | None = None
mode: ModelMode | None = None
access_groups: list[str] | None = None
team_id: str | None = None
class ModelNewBody(BaseModel):
@ -861,6 +863,7 @@ class TeamNewResponse(BaseModel):
class TeamUpdateBody(BaseModel):
team_id: str
team_alias: str
models: list[str] | None = None
class TeamInfoParams(BaseModel):

View file

@ -278,7 +278,21 @@ class ProxyClient:
mode: ModelMode | None = None,
) -> str:
"""Register a deployment under `model_name` and return its proxy-assigned
model_id, once the model is actually servable on the data plane.
model_id, once the model is actually servable on the data plane."""
return self.register_model(
ModelNewBody(
model_name=model_name,
litellm_params=litellm_params,
model_info=ModelInfoBody(mode=mode),
)
)
def register_model(self, body: ModelNewBody, listed_for: str | None = None) -> str:
"""`create_model` for deployments that carry more than a mode: access groups,
team scoping, a pinned id. `listed_for` is the virtual key whose /v1/models
view must list the deployment before it counts as servable, because a
team-scoped deployment is listed to its own team and to nobody else, master
key included; leave it unset for a proxy-wide model.
/model/new is a control-plane route; the data plane (which serves /chat,
/ocr, ...) only picks the new model up on its next DB reload, so a call
@ -296,25 +310,22 @@ class ProxyClient:
self.transport.post(
"/model/new",
headers=self.transport.master,
json=ModelNewBody(
model_name=model_name,
litellm_params=litellm_params,
model_info=ModelInfoBody(mode=mode),
),
json=body,
response_type=ModelNewResponse,
)
).model_id
written_at = time.monotonic()
self._await_model_servable(model_name)
self._await_model_servable(body.model_name, listed_for)
settle_propagation(written_at)
return model_id
def _await_model_servable(self, model_name: str) -> None:
def _await_model_servable(self, model_name: str, listed_for: str | None = None) -> None:
"""Block until the data plane lists `model_name`, or fail at model_servable_timeout."""
headers = self.transport.master if listed_for is None else self.transport.bearer(listed_for)
outcome = await_servable(
lambda poll_timeout: self.transport.get(
"/v1/models",
headers=self.transport.master,
headers=headers,
params=NoBody(),
response_type=ModelsListResponse,
timeout=poll_timeout,

View file

@ -0,0 +1,163 @@
"""Harness coverage for canonical request identity (LIT-5741).
No proxy and no ``e2e`` marker: pure functions over ``RecordedRequest``. Pins
the two failure modes match keys must avoid: keying on volatile material so
nothing ever matches (markers, virtual keys, ids, timestamps, volatile
headers), and keying on too little so different requests collide and a test
silently asserts against another request's response.
"""
from __future__ import annotations
import pytest
from pydantic import JsonValue
from fixture_bundle import RecordedRequest
from fixture_canonical import CanonicalRequest, canonical_string, canonicalize, is_secret_field
def request(
method: str = "post",
path: str = "/chat/completions",
*,
headers: dict[str, str] | None = None,
params: dict[str, str] | None = None,
body: JsonValue | None = None,
form: dict[str, str] | None = None,
file_name: str | None = None,
file_sha256: str | None = None,
file_bytes: int | None = None,
) -> RecordedRequest:
return RecordedRequest(
method=method,
path=path,
headers=headers or {},
params=params or {},
body=body,
form=form,
file_name=file_name,
file_sha256=file_sha256,
file_bytes=file_bytes,
)
class TestPlaceholders:
@pytest.mark.parametrize(
("raw", "expected"),
[
("Reply ok. 4d5152a995b7", "Reply ok. <marker>"),
("e2e-chat-stream-4d5152a995b7", "e2e-chat-stream-<marker>"),
("sk-3mCXCTGmYuEEIU2i2qmVE3Xq6tSK1O0X6ZIRP1Lpw8ZlbNjt", "<key>"),
("9f1c8a2e-4b3d-4f6a-8f2f-0a1b2c3d4e5f", "<uuid>"),
("z" * 64, "z" * 64),
("0123456789abcdef" * 4, "<sha256>"),
("2026-08-19T20:57:13.363499+00:00", "<timestamp>"),
("2026-08-19", "<date>"),
("chatcmpl-C0LO6rRkfJlpJ2mqW9BHYo4Sm8FWl", "<id>"),
("batch_688a8b7f9a08819096e0f7c88fcd07c5", "<id>"),
("file-XyZ12345abc", "<id>"),
("gpt-4o-mini", "gpt-4o-mini"),
("max_tokens", "max_tokens"),
("sk-1234", "sk-1234"),
],
)
def test_rewrites_exactly_the_volatile_shapes(self, raw: str, expected: str) -> None:
assert canonical_string(raw) == expected
class TestSecretFields:
@pytest.mark.parametrize(
("name", "secret"),
[
("api_key", True),
("openai_api_key", True),
("aws_secret_access_key", True),
("aws_session_token", True),
("vertex_credentials", True),
("static_headers", True),
("langfuse_secret_key", True),
("model", False),
("max_completion_tokens", False),
("api_base", False),
],
)
def test_names_that_carry_credentials(self, name: str, secret: bool) -> None:
assert is_secret_field(name) is secret
class TestKeyStability:
def test_volatile_material_does_not_change_the_key(self) -> None:
"""Acceptance: a suite recorded on one machine (fresh keys, that day's
dates, that run's markers) replays on another with no misses."""
first = request(
headers={"authorization": "Bearer sk-run-one-aaaaaaaaaaaaaaaa", "x-request-id": "req-1"},
params={"start_date": "2026-08-18"},
body={
"model": "e2e-chat-4d5152a995b7",
"messages": [{"role": "user", "content": "Reply ok. 4d5152a995b7"}],
"api_key": "sk-live-one-aaaaaaaaaaaaaaaa",
},
)
second = request(
headers={"authorization": "Bearer sk-run-two-bbbbbbbbbbbbbbbb", "x-request-id": "req-2"},
params={"start_date": "2026-08-19"},
body={
"model": "e2e-chat-1a2b3c4d5e6f",
"messages": [{"role": "user", "content": "Reply ok. 1a2b3c4d5e6f"}],
"api_key": "os.environ/OPENAI_API_KEY",
},
)
assert canonicalize(first).key == canonicalize(second).key
def test_serialization_order_is_not_identity(self) -> None:
ordered = request(body={"model": "m", "stream": True})
reversed_order = request(body={"stream": True, "model": "m"})
assert canonicalize(ordered).key == canonicalize(reversed_order).key
def test_generated_ids_in_the_path_do_not_change_the_key(self) -> None:
first = request("get", "/v1/batches/batch_688a8b7f9a08819096e0f7c88fcd07c5")
second = request("get", "/v1/batches/batch_770b9c8f0b19920107f1f8d99fde18d6")
assert canonicalize(first).key == canonicalize(second).key
class TestKeyDistinctness:
def test_requests_differing_only_inside_canonicalized_fields_stay_distinct(self) -> None:
"""Acceptance: a naive verb+path hash collides these; the content key
must not, or one test silently asserts against the other's response."""
first = request(body={"messages": [{"content": "Reply ok. 4d5152a995b7"}]})
second = request(body={"messages": [{"content": "Count to three. 4d5152a995b7"}]})
naive = (first.method, first.path)
assert naive == (second.method, second.path)
assert canonicalize(first).key != canonicalize(second).key
def test_a_kept_header_is_identity(self) -> None:
first = request(headers={"x-litellm-tags": "prod"})
second = request(headers={"x-litellm-tags": "shadow"})
assert canonicalize(first).key != canonicalize(second).key
def test_a_volatile_header_is_not_identity(self) -> None:
first = request(headers={"traceparent": "00-aa-bb-01", "x-api-key": "one"})
second = request(headers={"traceparent": "00-cc-dd-01", "x-api-key": "two"})
assert canonicalize(first).key == canonicalize(second).key
def test_secret_set_versus_unset_stays_distinct(self) -> None:
with_key = request(body={"api_key": "sk-live-aaaaaaaaaaaaaaaa"})
without_key = request(body={"api_key": None})
assert canonicalize(with_key).key != canonicalize(without_key).key
def test_file_content_is_identity(self) -> None:
first = request(
"upload", "/v1/files", file_name="batch.jsonl", file_sha256="a" * 64, file_bytes=10
)
second = request(
"upload", "/v1/files", file_name="batch.jsonl", file_sha256="b" * 64, file_bytes=10
)
assert canonicalize(first).key != canonicalize(second).key
class TestKeyShape:
def test_key_names_method_path_and_digest(self) -> None:
canonical = canonicalize(request("post", "/model/new", body={"model_name": "m"}))
assert isinstance(canonical, CanonicalRequest)
assert canonical.key.startswith("post /model/new #")
assert len(canonical.key.rsplit("#", 1)[1]) == 16

View file

@ -5,17 +5,22 @@ the live one (dependency injection, no monkeypatching): recording must pass
every value through unchanged while writing one redacted interaction file per
call, and replay must serve identical values from the bundle alone - the
fake's call log proves nothing reaches the inner transport - failing hard
(``ReplayMiss``) on any drift in order, verb, or path. The collection-time
gate and report header are pinned here too, including the stale message that
names the bundle's age.
(``ReplayMiss``) on any content drift, printing the computed canonical key and
the closest recorded key (LIT-5741; the pure canonicalizer is pinned in
test_fixture_canonical.py). The collection-time gate and report header are
pinned here too, including the stale message that names the bundle's age.
"""
from __future__ import annotations
import hashlib
import sys
import threading
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from datetime import datetime, timedelta, timezone
from pathlib import Path
from uuid import uuid4
import pytest
from pydantic import BaseModel
@ -35,10 +40,12 @@ from fixture_bundle import (
Interaction,
LoadedBundle,
Manifest,
RecordedResult,
load_bundle,
prepare_bundle,
slug_for_test,
)
from fixture_canonical import canonicalize
from fixture_transport import (
InvalidFixtureMode,
RecordingTransport,
@ -50,6 +57,7 @@ from fixture_transport import (
fixture_mode_collection_error,
fixture_report_lines,
parse_fixture_mode,
recorded_request,
replay_leftover_error,
select_transport,
)
@ -70,6 +78,17 @@ class Query(BaseModel):
q: str
class DeployParams(BaseModel):
model: str
api_key: str | None = None
aws_secret_access_key: str | None = None
class DeployBody(BaseModel):
model_name: str
litellm_params: DeployParams
STREAMING = StreamingResponse(
status_code=200,
body="",
@ -272,6 +291,28 @@ class TestRecordingTransport:
}
assert "sk-secret" not in this_tests_files(root)[0].read_text(encoding="utf-8")
def test_redacts_credential_body_fields_in_the_recorded_request(self, tmp_path: Path) -> None:
fake = FakeTransport()
root = tmp_path / "bundle"
recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root))
recording.post(
"/model/new",
headers=fake.master,
json=DeployBody(
model_name="m",
litellm_params=DeployParams(model="openai/gpt", api_key="sk-live-provider-secret-123456"),
),
response_type=Payload,
)
raw = this_tests_files(root)[0].read_text(encoding="utf-8")
interaction = Interaction.model_validate_json(raw)
assert "sk-live-provider-secret-123456" not in raw
assert isinstance(interaction.request.body, dict)
params = interaction.request.body["litellm_params"]
assert isinstance(params, dict)
assert params["api_key"] == "<redacted>"
assert params["aws_secret_access_key"] is None
def test_upload_records_a_content_digest_not_the_bytes(self, tmp_path: Path) -> None:
fake = FakeTransport()
root = tmp_path / "bundle"
@ -335,25 +376,175 @@ class TestReplayTransport:
)
assert fake.calls == calls_after_record
def test_mismatched_call_names_recorded_and_actual(self, tmp_path: Path) -> None:
def test_miss_names_the_computed_key_and_the_closest_recorded_key(self, tmp_path: Path) -> None:
fake = FakeTransport()
root = tmp_path / "bundle"
recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root))
recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload)
replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234")
with pytest.raises(ReplayMiss, match=r"recorded post /model/new, test made get /v1/models"):
with pytest.raises(ReplayMiss) as excinfo:
replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload)
message = str(excinfo.value)
assert "no recorded interaction matches key get /v1/models #" in message
assert "closest recorded key is post /model/new #" in message
assert "0000-post-model-new.json" in message
assert "re-record with E2E_FIXTURE_MODE=record" in message
def test_exhausted_recording_names_the_call_count(self, tmp_path: Path) -> None:
def test_content_drift_on_the_same_route_misses_with_no_live_call(self, tmp_path: Path) -> None:
"""The naive verb+path match replayed a stale response for a request
whose content had changed, silently passing; a content key must miss,
print both canonical forms' diff, and never reach the inner transport."""
fake = FakeTransport()
root = tmp_path / "bundle"
recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root))
recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload)
calls_after_record = list(fake.calls)
replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234")
with pytest.raises(ReplayMiss) as excinfo:
replay.post("/model/new", headers=replay.master, json=Body(prompt="y"), response_type=Payload)
message = str(excinfo.value)
assert "no recorded interaction matches key post /model/new #" in message
assert "closest recorded key is post /model/new #" in message
assert '- "prompt": "x"' in message
assert '+ "prompt": "y"' in message
assert fake.calls == calls_after_record
def test_exhausted_key_names_the_key(self, tmp_path: Path) -> None:
fake = FakeTransport()
root = tmp_path / "bundle"
recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root))
recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload)
replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234")
replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload)
with pytest.raises(ReplayMiss, match=r"call #2 \(post /model/new\) has no recorded interaction \(1 recorded"):
with pytest.raises(
ReplayMiss, match=r"every recorded interaction for key post /model/new #\w{16} is already consumed"
):
replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload)
def test_replays_out_of_recorded_order_across_distinct_keys(self, tmp_path: Path) -> None:
"""Concurrent tests interleave independent calls nondeterministically
(e.g. a burst of parallel chat calls), so replay matches by content,
never by recorded position."""
fake = FakeTransport()
root = tmp_path / "bundle"
recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root))
recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload)
recording.post("/key/generate", headers=fake.master, json=Body(prompt="k"), response_type=Payload)
source = replay_source(root)
replay: Transport = ReplayTransport(source=source, master_key="sk-1234")
replay.post("/key/generate", headers=replay.master, json=Body(prompt="k"), response_type=Payload)
replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload)
assert source.leftover_error(current_test_key()) is None
def test_identical_requests_replay_their_responses_in_recorded_order(self, tmp_path: Path) -> None:
"""A poll loop makes the same request repeatedly and asserts on the
progression, so duplicates under one key stay FIFO."""
root = tmp_path / "bundle"
recorder = make_recorder(root)
recorder.record(
test_key=current_test_key(),
request=recorded_request(
"get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all")
),
response=RecordedResult(kind="success", status_code=200, data={"value": "first"}),
)
recorder.record(
test_key=current_test_key(),
request=recorded_request(
"get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all")
),
response=RecordedResult(kind="success", status_code=200, data={"value": "second"}),
)
replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234")
first = replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload)
second = replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload)
assert first == Success(status_code=200, data=Payload(value="first"))
assert second == Success(status_code=200, data=Payload(value="second"))
def test_concurrent_replays_of_one_key_serve_each_recording_exactly_once(self, tmp_path: Path) -> None:
"""A burst of parallel identical calls consumes one shared pool: no
response duplicated, none forgotten, nothing left over at teardown.
The tiny switch interval forces thread preemption inside pool setup
and consumption, so a non-atomic pool build or pop fails this test."""
root = tmp_path / "bundle"
recorder = make_recorder(root)
for ordinal in range(32):
recorder.record(
test_key=current_test_key(),
request=recorded_request(
"get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all")
),
response=RecordedResult(kind="success", status_code=200, data={"value": f"v{ordinal:02d}"}),
)
source = replay_source(root)
replay: Transport = ReplayTransport(source=source, master_key="sk-1234")
barrier = threading.Barrier(8)
def consume_one() -> str:
result = replay.get(
"/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload
)
assert isinstance(result, Success)
return result.data.value
def consume(_: int) -> tuple[str, ...]:
barrier.wait()
return tuple(consume_one() for _call in range(4))
previous_interval = sys.getswitchinterval()
sys.setswitchinterval(1e-6)
try:
with ThreadPoolExecutor(max_workers=8) as executor:
served = sorted(value for values in executor.map(consume, range(8)) for value in values)
finally:
sys.setswitchinterval(previous_interval)
assert served == [f"v{ordinal:02d}" for ordinal in range(32)]
assert source.leftover_error(current_test_key()) is None
class TestRecordedKeySets:
def test_two_separate_recordings_of_one_flow_produce_identical_key_sets(
self, tmp_path: Path
) -> None:
"""Everything a run randomizes (markers, virtual keys, dates) must
canonicalize out, so separately recorded runs of the same suite agree
on every match key and a bundle recorded elsewhere replays here."""
def record_flow(root: Path, run_date: str) -> list[str]:
fake = FakeTransport()
recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root))
marker = deterministic_marker()
recording.post(
"/model/new",
headers=fake.master,
json=DeployBody(
model_name=f"e2e-chat-{marker}",
litellm_params=DeployParams(model="openai/gpt", api_key=f"sk-live-{uuid4().hex}"),
),
response_type=Payload,
)
recording.post(
"/chat/completions",
headers=recording.bearer(f"sk-{uuid4().hex}"),
json=Body(prompt=f"Reply with the single word ok. {marker}"),
response_type=Payload,
)
recording.get(
"/spend/logs", headers=fake.master, params=Query(q=run_date), response_type=Payload
)
loaded = load_bundle(root)
assert isinstance(loaded, LoadedBundle)
return sorted(
canonicalize(interaction.request).key
for interactions in loaded.interactions.values()
for interaction in interactions
)
first_keys = record_flow(tmp_path / "one", "2026-08-18")
second_keys = record_flow(tmp_path / "two", "2026-08-19")
assert first_keys == second_keys
assert len(first_keys) == 3
class TestReplayLeftover:
def test_fully_consumed_recording_leaves_nothing(self, tmp_path: Path) -> None:
@ -378,7 +569,7 @@ class TestReplayLeftover:
error = source.leftover_error(current_test_key())
assert error is not None
assert "1 of 2 recorded interactions never consumed" in error
assert "next is probe /health/liveliness" in error
assert "e.g. probe /health/liveliness #" in error
assert "re-record with E2E_FIXTURE_MODE=record" in error
def test_test_without_recordings_has_no_leftover(self, tmp_path: Path) -> None:

View file

@ -150,13 +150,13 @@ def test_parse_jsonl_empty_content_is_empty_list():
assert bu._get_file_content_as_dictionary(b"") == []
def test_parse_jsonl_malformed_raises():
with pytest.raises(Exception):
bu._get_file_content_as_dictionary(b"not valid json")
def test_parse_jsonl_malformed_lines_skipped():
content = b'{"a": 1}\nnot valid json\n{"b": 2}\n'
assert bu._get_file_content_as_dictionary(content) == [{"a": 1}, {"b": 2}]
# =========================================================================== #
# _iter_batch_input_lines / _iter_batch_input_entries (JSONL parsing)
# _iter_batch_input_lines / _iter_batch_output_entries (JSONL parsing)
# =========================================================================== #
@ -173,19 +173,22 @@ def test_iter_input_lines_empty():
assert list(bu._iter_batch_input_lines(b"")) == []
def test_iter_input_entries_parses_each_row():
def test_iter_output_entries_parses_each_row():
content = b'{"body": {"model": "gpt-4o"}}\n{"body": {"model": "claude-3"}}\n'
assert list(bu._iter_batch_input_entries(content)) == [
assert list(bu._iter_batch_output_entries(content)) == [
{"body": {"model": "gpt-4o"}},
{"body": {"model": "claude-3"}},
]
def test_iter_input_entries_raises_on_malformed_line():
# _iter_batch_input_entries raises on a bad row; callers that must survive
# bad rows iterate _iter_batch_input_lines and parse per-row instead.
with pytest.raises(Exception):
list(bu._iter_batch_input_entries(b'{"ok":1}\nnot-json\n'))
def test_iter_output_entries_skips_malformed_and_non_object_lines():
content = b'{"ok": 1}\nnot-json\n[1, 2]\n{"ok": 2}\n'
assert list(bu._iter_batch_output_entries(content)) == [{"ok": 1}, {"ok": 2}]
def test_iter_output_entries_skips_undecodable_line():
content = b'{"ok": 1}\n{"note": "\xff-bad"}\n{"ok": 2}\n'
assert list(bu._iter_batch_output_entries(content)) == [{"ok": 1}, {"ok": 2}]
# =========================================================================== #
@ -471,6 +474,25 @@ def test_cost_from_content_completion_cost_path(monkeypatch):
assert len(calls) == 2 # failed row not costed
def test_empty_body_line_does_not_zero_whole_batch():
"""A status-200 row with an empty body makes litellm.completion_cost raise;
that line must be skipped instead of zeroing the whole batch."""
rows = [
_success_row(usage=_usage(10, 5)),
{
"custom_id": "request-poison-empty",
"response": {"status_code": 200, "request_id": "inject-empty-body", "body": {}},
},
_success_row(usage=_usage(20, 10)),
]
cost, usage, models = bu._aggregate_batch_cost_usage_models(entries=rows, custom_llm_provider="openai")
assert cost > 0.0
assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (30, 15, 45)
assert models == ["gpt-4o", "gpt-4o"]
def test_cost_from_content_model_info_path(monkeypatch):
# model_info set -> batch_cost_calculator(prompt_cost, completion_cost).
import litellm.cost_calculator as cc

View file

@ -6213,6 +6213,32 @@ class TestDynamicTracerProviderCache(unittest.TestCase):
self.assertTrue(entry.owns_exporter)
self.assertIsNotNone(entry.provider._atexit_handler)
def test_dynamic_providers_share_one_resource(self):
"""Building the Resource scans every installed distribution's entry points, and the
dynamic providers reach it from the async logging path, so one logger builds it once."""
logger = self._logger(cap=8)
for i in range(4):
logger._get_tracer_with_dynamic_headers({"authorization": f"Basic tenant-{i}"})
entries = list(logger._tracer_provider_cache.values())
self.assertEqual(len(entries), 4)
self.assertEqual(len({id(entry.provider.resource) for entry in entries}), 1)
self.assertIs(entries[0].provider.resource, logger._litellm_resource())
def test_resource_is_memoized_per_logger_not_shared(self):
"""Two loggers must not share a Resource; the second's service.name would be wrong."""
first = self._logger()
second = OpenTelemetry(
config=OpenTelemetryConfig(exporter="console", skip_set_global=True, service_name="svc-second")
)
self.addCleanup(second._tracer_provider.shutdown)
self.assertIsNot(first._litellm_resource(), second._litellm_resource())
self.assertEqual(second._litellm_resource().attributes.get("service.name"), "svc-second")
class TestOpenTelemetryDatabaseSemconvAttributes(unittest.TestCase):
"""A Postgres service span must name the PostgreSQL server it reached.

View file

@ -2324,6 +2324,87 @@ def test_data_residency_composes_with_service_tier(_local_model_cost_map):
assert priority_eu_total == pytest.approx(priority_base_total * 1.10, rel=1e-9)
@pytest.mark.parametrize("model", ["gemini-3.5-flash", "claude-haiku-4-5@20251001"])
@pytest.mark.parametrize("vertex_location", ["us-central1", "us-east5", "europe-west1", "asia-southeast1"])
def test_vertex_regional_location_applies_uplift(vertex_location, model, _local_model_cost_map):
"""Google bills every non-global Vertex endpoint at 1.1x the global rate for GA
Gemini 3+ and regional-pricing Claude models, so a request served from a regional
location must cost 1.1x what the same usage costs on the global endpoint."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="vertex_ai")
regional = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="vertex_ai",
vertex_location=vertex_location,
)
base_total = base[0] + base[1]
regional_total = regional[0] + regional[1]
assert base_total > 0
assert regional_total == pytest.approx(base_total * 1.10, rel=1e-9)
assert regional[0] == pytest.approx(base[0] * 1.10, rel=1e-9)
assert regional[1] == pytest.approx(base[1] * 1.10, rel=1e-9)
@pytest.mark.parametrize("vertex_location", [None, "global", "GLOBAL"])
def test_vertex_global_or_absent_location_no_uplift(vertex_location, _local_model_cost_map):
"""The global endpoint prices at the base rate, whatever the casing, and an
unresolved location must never uplift."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(
model="claude-haiku-4-5@20251001", usage=usage, custom_llm_provider="vertex_ai"
)
located = generic_cost_per_token(
model="claude-haiku-4-5@20251001",
usage=usage,
custom_llm_provider="vertex_ai",
vertex_location=vertex_location,
)
assert base == located
@pytest.mark.parametrize("model", ["claude-opus-4-1", "gemini-2.0-flash-001"])
def test_vertex_location_no_uplift_for_uniformly_priced_model(model, _local_model_cost_map):
"""Models Google prices uniformly across endpoints (Gemini 2.x, Claude Opus 4.1
and older) carry no multiplier and must not move with the location."""
from litellm.types.utils import Usage
usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500)
base = generic_cost_per_token(model=model, usage=usage, custom_llm_provider="vertex_ai")
regional = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="vertex_ai",
vertex_location="us-east5",
)
assert base == regional, f"{model} should not have a regional-endpoint uplift"
def test_vertex_uplift_invalid_multiplier_defaults_to_one():
"""A malformed multiplier in the cost map degrades to base pricing, never raises."""
from litellm.litellm_core_utils.llm_cost_calc.utils import (
get_vertex_regional_endpoint_uplift,
)
assert (
get_vertex_regional_endpoint_uplift(
{"regional_endpoint_uplift_multiplier": "not-a-number"}, "us-east5"
)
== 1.0
)
def test_priority_service_tier_above_threshold_uses_priority_tier_rates_for_cached_tokens(
_local_model_cost_map,
):
@ -2877,6 +2958,57 @@ def test_token_type_cost_breakdown_applies_regional_uplift():
assert text_input_cost + eu.cache_read_cost == pytest.approx(prompt_cost)
def test_token_type_cost_breakdown_applies_vertex_regional_uplift():
"""
Non-global Vertex endpoints apply a flat 1.1x uplift to every token cost. The
per-type breakdown must apply the same uplift via vertex_location so it stays
reconciled with the uplifted input_cost/output_cost totals, instead of being
logged at the global rate.
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model = "claude-haiku-4-5@20251001"
custom_llm_provider = "vertex_ai"
usage = Usage(
prompt_tokens=1000,
completion_tokens=500,
total_tokens=1500,
prompt_tokens_details=PromptTokensDetailsWrapper(
cached_tokens=400, text_tokens=600
),
)
model_info = litellm.get_model_info(
model=model, custom_llm_provider=custom_llm_provider
)
uplift = model_info["regional_endpoint_uplift_multiplier"]
assert uplift > 1.0
base = get_token_type_cost_breakdown(
model=model, custom_llm_provider=custom_llm_provider, usage=usage
)
regional = get_token_type_cost_breakdown(
model=model,
custom_llm_provider=custom_llm_provider,
usage=usage,
vertex_location="us-east5",
)
assert base.cache_read_cost > 0
assert regional.cache_read_cost == pytest.approx(base.cache_read_cost * uplift)
# The uplifted breakdown must still reconcile with the uplifted totals.
prompt_cost, _completion_cost = generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider=custom_llm_provider,
vertex_location="us-east5",
)
text_input_cost = 600 * model_info["input_cost_per_token"] * uplift
assert text_input_cost + regional.cache_read_cost == pytest.approx(prompt_cost)
def test_token_type_cost_breakdown_applies_anthropic_geo_multiplier(monkeypatch):
"""
Anthropic's regional (geo) uplift lives in provider_specific_entry and is

View file

@ -5026,3 +5026,150 @@ async def test_streaming_success_callbacks_survive_standard_logging_payload_fail
await logging_obj.async_success_handler(result=_assembled_stream_result())
releasing.async_log_success_event.assert_awaited_once()
def _resolve(custom_llm_provider, litellm_params, optional_params, model):
from litellm.litellm_core_utils.litellm_logging import (
_resolve_vertex_location_for_cost,
)
return _resolve_vertex_location_for_cost(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
optional_params=optional_params,
model=model,
)
def test_resolve_vertex_location_for_cost():
"""Vertex requests resolve the serving location the way dispatch does; other providers get None."""
assert _resolve("openai", {"vertex_location": "us-east5"}, None, "gpt-4o") is None
assert _resolve(None, {}, None, "gemini-3.5-flash") is None
assert _resolve("vertex_ai", {"vertex_location": "us-east5"}, None, "gemini-3.5-flash") == "us-east5"
assert _resolve("vertex_ai", {"vertex_location": "global"}, None, "gemini-3.5-flash") == "global"
assert (
_resolve("vertex_ai_beta", {"vertex_ai_location": "europe-west1"}, None, "claude-haiku-4-5@20251001")
== "europe-west1"
)
def test_resolve_vertex_location_for_cost_reads_optional_params(monkeypatch):
"""
On the proxy the logging object predates deployment selection, so the deployment's
configured location only reaches it through optional_params. A configured global
location must beat the environment fallback, or every proxy call gets the regional uplift.
"""
monkeypatch.setenv("VERTEXAI_LOCATION", "us-east5")
monkeypatch.setattr(litellm, "vertex_location", None)
assert _resolve("vertex_ai", {}, {"vertex_location": "global"}, "gemini-3.5-flash") == "global"
assert _resolve("vertex_ai", None, {"vertex_location": "europe-west1"}, "gemini-3.5-flash") == "europe-west1"
assert (
_resolve(
"vertex_ai",
{"vertex_location": "us-east5"},
{"vertex_location": "global"},
"gemini-3.5-flash",
)
== "global"
)
assert _resolve("vertex_ai", {"vertex_location": "global"}, {}, "gemini-3.5-flash") == "global"
assert _resolve("vertex_ai", {}, {}, "gemini-3.5-flash") == "us-east5"
def test_resolve_vertex_location_for_cost_default_region(monkeypatch):
"""With no location configured anywhere, resolution lands on the dispatch default us-central1."""
monkeypatch.delenv("VERTEXAI_LOCATION", raising=False)
monkeypatch.delenv("VERTEX_LOCATION", raising=False)
monkeypatch.setattr(litellm, "vertex_location", None)
assert _resolve("vertex_ai", {}, None, "gemini-3.5-flash") == "us-central1"
assert _resolve("vertex_ai", None, None, "gemini-3.5-flash") == "us-central1"
def test_response_cost_calculator_prices_proxy_vertex_calls_on_the_configured_location(monkeypatch):
"""
Proxy-shaped logging objects (created before the router picks a deployment) carry the
deployment's vertex_location only in optional_params. A global deployment must price at
base rates even when the environment points at a regional location, and a regional one
must price with the uplift.
"""
from datetime import datetime
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", get_model_cost_map(url=""))
monkeypatch.setenv("VERTEXAI_LOCATION", "us-east5")
monkeypatch.setattr(litellm, "vertex_location", None)
def cost_at(location):
logging_obj = LitellmLogging(
model="gemini-3.5-flash",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="completion",
start_time=datetime.now(),
litellm_call_id=f"vertex-loc-{location}",
function_id="f",
)
logging_obj.update_environment_variables(
model="gemini-3.5-flash",
user="",
optional_params={"vertex_location": location},
litellm_params={"api_base": ""},
custom_llm_provider="vertex_ai",
)
response = ModelResponse(
id="resp-1",
model="gemini-3.5-flash",
choices=[{"message": {"role": "assistant", "content": "hello"}, "index": 0, "finish_reason": "stop"}],
usage={"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
)
return logging_obj._response_cost_calculator(result=response)
info = litellm.model_cost["vertex_ai/gemini-3.5-flash"]
expected_global = 10 * info["input_cost_per_token"] + 5 * info["output_cost_per_token"]
assert cost_at("global") == pytest.approx(expected_global)
assert cost_at("us-east5") == pytest.approx(info["regional_endpoint_uplift_multiplier"] * expected_global)
def test_set_cost_breakdown_stores_vertex_location():
"""vertex_location is recorded in the pricing basis, None for non-vertex requests."""
from datetime import datetime
logging_obj = LitellmLogging(
model="vertex_ai/claude-haiku-4-5@20251001",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="completion",
start_time=datetime.now(),
litellm_call_id="vertex-location-set",
function_id="f",
)
logging_obj.set_cost_breakdown(
input_cost=0.001,
output_cost=0.002,
total_cost=0.003,
cost_for_built_in_tools_cost_usd_dollar=0.0,
vertex_location="us-east5",
)
assert logging_obj.cost_breakdown["vertex_location"] == "us-east5"
no_location = LitellmLogging(
model="gpt-4o",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="completion",
start_time=datetime.now(),
litellm_call_id="vertex-location-absent",
function_id="f",
)
no_location.set_cost_breakdown(
input_cost=0.001,
output_cost=0.002,
total_cost=0.003,
cost_for_built_in_tools_cost_usd_dollar=0.0,
)
assert no_location.cost_breakdown.get("vertex_location") is None

View file

@ -1262,3 +1262,83 @@ def test_get_combined_tool_content_joins_many_custom_tool_input_fragments_in_ord
assert isinstance(combined[1], ChatCompletionMessageCustomToolCall)
assert combined[1].custom.name == "run_script"
assert combined[1].custom.input == "".join(object_fragments)
def _reasoning_stream_chunk() -> ModelResponseStream:
return ModelResponseStream(
id="chatcmpl-reasoning",
model="claude-opus-4-8",
choices=[StreamingChoices(finish_reason=None, index=0, delta=Delta(content="10", role="assistant"))],
)
def test_count_reasoning_tokens_returns_none_for_signature_only_thinking():
from litellm.types.utils import Choices, Message, ModelResponse
processor = ChunkProcessor(chunks=[_reasoning_stream_chunk()])
response = ModelResponse(
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(content="10", role="assistant", reasoning_content=""),
)
]
)
assert processor.count_reasoning_tokens(response) is None
def test_count_reasoning_tokens_counts_visible_reasoning():
from litellm.types.utils import Choices, Message, ModelResponse
processor = ChunkProcessor(chunks=[_reasoning_stream_chunk()])
response = ModelResponse(
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
content="10",
role="assistant",
reasoning_content="let me count the primes under thirty",
),
)
]
)
assert processor.count_reasoning_tokens(response) > 0
@pytest.mark.parametrize(
"estimated_reasoning_tokens, expected_reasoning_tokens, expected_text_tokens",
[(40, 40, 60), (250, 100, 0)],
)
def test_calculate_usage_fills_unknown_split_from_reasoning_estimate(
estimated_reasoning_tokens, expected_reasoning_tokens, expected_text_tokens
):
from litellm.types.utils import CompletionTokensDetailsWrapper
chunk = ModelResponseStream(
id="chatcmpl-unknown-split",
model="claude-opus-4-8",
choices=[StreamingChoices(finish_reason="stop", index=0, delta=Delta(content=None, role=None))],
usage=Usage(
prompt_tokens=50,
completion_tokens=100,
total_tokens=150,
completion_tokens_details=CompletionTokensDetailsWrapper(reasoning_tokens=None, text_tokens=None),
),
)
processor = ChunkProcessor(chunks=[chunk])
usage = processor.calculate_usage(
chunks=[chunk],
model="claude-opus-4-8",
completion_output="10",
reasoning_tokens=estimated_reasoning_tokens,
)
assert usage.completion_tokens == 100
assert usage.completion_tokens_details.reasoning_tokens == expected_reasoning_tokens
assert usage.completion_tokens_details.text_tokens == expected_text_tokens

View file

@ -221,6 +221,162 @@ def test_calculate_usage_clamps_text_tokens_when_reasoning_estimate_exceeds_outp
assert usage.completion_tokens_details.text_tokens == 0
def test_calculate_usage_prefers_provider_reported_thinking_tokens():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={
"input_tokens": 32,
"output_tokens": 421,
"output_tokens_details": {"thinking_tokens": 372},
},
reasoning_content="",
completion_response={
"content": [
{"type": "thinking", "thinking": "", "signature": "sig"},
{"type": "text", "text": "10"},
]
},
)
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 372
assert usage.completion_tokens_details.text_tokens == 49
def test_calculate_usage_provider_thinking_tokens_win_over_visible_reasoning_estimate():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={
"input_tokens": 50,
"output_tokens": 811,
"output_tokens_details": {"thinking_tokens": 747},
},
reasoning_content="short visible reasoning that tokenizes to far fewer than 747 tokens",
)
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 747
assert usage.completion_tokens_details.text_tokens == 64
def test_calculate_usage_sums_provider_thinking_tokens_across_iterations():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={
"input_tokens": 10,
"output_tokens": 300,
"iterations": [
{"input_tokens": 5, "output_tokens": 100, "output_tokens_details": {"thinking_tokens": 60}},
{"input_tokens": 5, "output_tokens": 200, "output_tokens_details": {"thinking_tokens": 90}},
],
},
reasoning_content=None,
)
assert usage.completion_tokens == 300
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 150
assert usage.completion_tokens_details.text_tokens == 150
def test_calculate_usage_falls_back_when_only_some_iterations_report_thinking_tokens():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={
"input_tokens": 10,
"output_tokens": 300,
"output_tokens_details": {"thinking_tokens": 240},
"iterations": [
{"input_tokens": 5, "output_tokens": 100, "output_tokens_details": {"thinking_tokens": 60}},
{"input_tokens": 5, "output_tokens": 200},
],
},
reasoning_content=None,
)
assert usage.completion_tokens == 300
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 240
assert usage.completion_tokens_details.text_tokens == 60
def test_calculate_usage_reports_unknown_split_when_only_some_iterations_report_thinking_tokens():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={
"input_tokens": 10,
"output_tokens": 300,
"iterations": [
{"input_tokens": 5, "output_tokens": 100, "output_tokens_details": {"thinking_tokens": 60}},
{"input_tokens": 5, "output_tokens": 200},
],
},
reasoning_content="",
completion_response={"content": [{"type": "thinking", "thinking": "", "signature": "sig"}]},
)
assert usage.completion_tokens == 300
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens is None
assert usage.completion_tokens_details.text_tokens is None
def test_calculate_usage_reports_unknown_split_when_thinking_ran_without_a_count():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={"input_tokens": 32, "output_tokens": 580},
reasoning_content="",
completion_response={
"content": [
{"type": "redacted_thinking", "data": "encrypted"},
{"type": "text", "text": "10"},
]
},
)
assert usage.completion_tokens == 580
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens is None
assert usage.completion_tokens_details.text_tokens is None
def test_calculate_usage_without_thinking_reports_all_output_as_text():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={"input_tokens": 32, "output_tokens": 171},
reasoning_content=None,
completion_response={"content": [{"type": "text", "text": "10"}]},
)
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 0
assert usage.completion_tokens_details.text_tokens == 171
def test_calculate_usage_ignores_malformed_provider_thinking_tokens():
config = AnthropicConfig()
usage = config.calculate_usage(
usage_object={
"input_tokens": 32,
"output_tokens": 100,
"output_tokens_details": {"thinking_tokens": "not-a-number"},
},
reasoning_content=None,
)
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 0
assert usage.completion_tokens_details.text_tokens == 100
def test_calculate_usage_handles_mocked_output_tokens_with_reasoning_content():
config = AnthropicConfig()

View file

@ -6005,6 +6005,87 @@ def test_adaptive_thinking_dropped_when_max_tokens_too_small_converse():
assert "thinking" not in optional_params
def test_converse_usage_reports_unknown_split_for_signature_only_thinking():
config = AmazonConverseConfig()
usage = config.transform_usage(
ConverseTokenUsageBlock(inputTokens=32, outputTokens=581, totalTokens=613),
reasoning_content="",
thinking_ran=True,
)
assert usage.completion_tokens == 581
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens is None
assert usage.completion_tokens_details.text_tokens is None
def test_converse_usage_estimates_split_for_visible_thinking():
config = AmazonConverseConfig()
usage = config.transform_usage(
ConverseTokenUsageBlock(inputTokens=32, outputTokens=581, totalTokens=613),
reasoning_content="Let me think about how many primes there are under thirty.",
thinking_ran=True,
)
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens > 0
assert (
usage.completion_tokens_details.reasoning_tokens + usage.completion_tokens_details.text_tokens
== usage.completion_tokens
)
def test_converse_usage_without_thinking_reports_all_output_as_text():
config = AmazonConverseConfig()
usage = config.transform_usage(ConverseTokenUsageBlock(inputTokens=32, outputTokens=171, totalTokens=203))
assert usage.completion_tokens_details is not None
assert usage.completion_tokens_details.reasoning_tokens == 0
assert usage.completion_tokens_details.text_tokens == 171
def test_converse_transform_response_signature_only_thinking_reports_unknown_split():
config = AmazonConverseConfig()
raw_response = MagicMock(status_code=200)
raw_response.text = json.dumps(
{
"output": {
"message": {
"role": "assistant",
"content": [
{"reasoningContent": {"reasoningText": {"text": "", "signature": "sig"}}},
{"text": "10"},
],
}
},
"stopReason": "end_turn",
"usage": {"inputTokens": 32, "outputTokens": 581, "totalTokens": 613},
}
)
raw_response.json.return_value = json.loads(raw_response.text)
response = config._transform_response(
model="bedrock/global.anthropic.claude-opus-4-8",
response=raw_response,
model_response=ModelResponse(),
stream=False,
logging_obj=None,
optional_params={},
api_key=None,
data={},
messages=[],
encoding=None,
)
assert response.choices[0].message.reasoning_content == ""
assert response.usage.completion_tokens_details.reasoning_tokens is None
assert response.usage.completion_tokens_details.text_tokens is None
def test_is_converse_usage_shape_distinguishes_camel_case_from_anthropic():
config = AmazonConverseConfig()
assert config.is_converse_usage_shape({"inputTokens": 1, "outputTokens": 2}) is True

View file

@ -0,0 +1,637 @@
"""
Tests for Amazon Bedrock AgentCore Web Search integration.
Mirror of tests/search_tests/test_agentcore_search.py placed in the
test_litellm tree so the AgentCoreSearchConfig transformation is exercised by
the sharded CI (coverage collection runs against this tree).
"""
import json
import os
import pytest
from unittest.mock import AsyncMock, patch, MagicMock
import litellm
from litellm.llms.bedrock.search.transformation import (
AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION,
AgentCoreSearchConfig,
)
GATEWAY_URL = "https://testgateway-abc123.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp"
MCP_RESULTS = [
{
"title": "Test Result 1",
"url": "https://example.com/1",
"text": "Snippet for result 1",
"publishedDate": "2026-06-16",
},
{
"title": "Test Result 2",
"url": "https://example.com/2",
"text": "Snippet for result 2",
},
]
def _mcp_response_body() -> dict:
return {
"jsonrpc": "2.0",
"id": 1,
"result": {"content": [{"type": "text", "text": json.dumps(MCP_RESULTS)}]},
}
def _make_mock_response(json_body: dict = None, text: str = None) -> MagicMock:
mock_response = MagicMock()
mock_response.status_code = 200
if text is not None:
mock_response.text = text
else:
mock_response.text = json.dumps(json_body)
mock_response.json.return_value = json_body
return mock_response
class TestAgentCoreSearch:
"""
Tests for AgentCore Web Search functionality with mocked network/signing.
"""
@pytest.mark.asyncio
async def test_agentcore_search_request_payload(self):
"""Validates the MCP tools/call payload and SigV4 signing without real AWS calls."""
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
mock_response = _make_mock_response(_mcp_response_body())
with (
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post,
patch.object(
AgentCoreSearchConfig,
"_sign_request",
return_value=(
{"Authorization": "AWS4-HMAC-SHA256 test", "Content-Type": "application/json"},
json.dumps({"signed": True}).encode(),
),
) as mock_sign,
):
mock_post.return_value = mock_response
response = await litellm.asearch(
query="latest developments in AI",
search_provider="agentcore",
max_results=5,
)
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
assert call_kwargs["url"] == GATEWAY_URL
# Signed body must be sent verbatim
assert call_kwargs["data"] == json.dumps({"signed": True}).encode()
assert call_kwargs["json"] is None
# Signing was invoked with the MCP request
mock_sign.assert_called_once()
sign_kwargs = mock_sign.call_args.kwargs
request_data = sign_kwargs["request_data"]
assert request_data["method"] == "tools/call"
assert request_data["params"]["name"] == "web-search-tool___WebSearch"
assert request_data["params"]["arguments"]["query"] == "latest developments in AI"
assert request_data["params"]["arguments"]["maxResults"] == 5
assert sign_kwargs["service_name"] == "bedrock-agentcore"
assert len(response.results) == 2
assert response.results[0].title == "Test Result 1"
assert response.results[0].url == "https://example.com/1"
assert response.results[0].snippet == "Snippet for result 1"
assert response.results[0].date == "2026-06-16"
def test_transform_search_request_query_truncation(self):
"""AgentCore rejects queries > 200 chars; the request must truncate."""
config = AgentCoreSearchConfig()
long_query = "a" * 300
data = config.transform_search_request(query=long_query, optional_params={})
assert len(data["params"]["arguments"]["query"]) == 200
def test_transform_search_request_joins_list_queries(self):
config = AgentCoreSearchConfig()
data = config.transform_search_request(query=["foo", "bar"], optional_params={})
assert data["params"]["arguments"]["query"] == "foo bar"
def test_transform_search_request_custom_tool_name(self):
config = AgentCoreSearchConfig()
data = config.transform_search_request(query="q", optional_params={"tool_name": "my-target___WebSearch"})
assert data["params"]["name"] == "my-target___WebSearch"
def test_transform_search_request_rejects_non_websearch_tool_name(self):
"""A caller-supplied tool_name must not reach other tools on the gateway."""
config = AgentCoreSearchConfig()
with pytest.raises(ValueError, match="must end with"):
config.transform_search_request(query="q", optional_params={"tool_name": "admin-target___DeleteUser"})
def test_transform_search_request_sends_documented_default_max_results(self):
"""The documented default of 10 is sent explicitly, not left to the gateway."""
config = AgentCoreSearchConfig()
data = config.transform_search_request(query="q", optional_params={})
assert data["params"]["arguments"]["maxResults"] == 10
def test_get_complete_url_requires_gateway_url(self):
config = AgentCoreSearchConfig()
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
with pytest.raises(ValueError, match="AGENTCORE_GATEWAY_URL"):
config.get_complete_url(api_base=None, optional_params={})
def test_get_complete_url_prefers_api_base(self):
config = AgentCoreSearchConfig()
assert config.get_complete_url(api_base=GATEWAY_URL, optional_params={}) == GATEWAY_URL
def test_validate_environment_sets_mcp_headers(self):
"""MCP Streamable HTTP requires accepting both JSON and SSE, and declaring
the protocol revision the client speaks."""
config = AgentCoreSearchConfig()
headers = config.validate_environment(headers={})
assert headers["Accept"] == "application/json, text/event-stream"
assert headers["Content-Type"] == "application/json"
assert headers["MCP-Protocol-Version"] == AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION
def test_default_protocol_version_is_the_agentcore_gateway_default(self):
"""A default AgentCore gateway supports only 2025-03-26 and answers
-32600 to anything newer, so that exact revision must be the default."""
assert AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION == "2025-03-26"
def test_protocol_version_env_override_wins(self):
"""A gateway pinned to a newer supportedVersions list needs the header
to match, so AGENTCORE_MCP_PROTOCOL_VERSION must override the default."""
config = AgentCoreSearchConfig()
with patch.dict(os.environ, {"AGENTCORE_MCP_PROTOCOL_VERSION": "2025-06-18"}):
headers = config.validate_environment(headers={})
assert headers["MCP-Protocol-Version"] == "2025-06-18"
def test_protocol_version_header_survives_signing(self):
"""Both auth paths must keep the MCP-Protocol-Version header on the wire."""
config = AgentCoreSearchConfig()
headers = config.validate_environment(headers={})
bearer_headers, _ = config.sign_request(
headers=headers,
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
api_key="test-jwt-token",
)
assert bearer_headers["MCP-Protocol-Version"] == AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION
with patch.dict(
os.environ,
{
"AWS_ACCESS_KEY_ID": "AKIAIOSFODNN7EXAMPLE",
"AWS_SECRET_ACCESS_KEY": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
},
):
signed_headers, _ = config.sign_request(
headers=headers,
optional_params={"aws_region_name": "us-east-1"},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
assert signed_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
assert signed_headers["MCP-Protocol-Version"] == AGENTCORE_DEFAULT_MCP_PROTOCOL_VERSION
def test_transform_search_response_parses_sse_frame(self):
"""Gateway may answer with an SSE-framed JSON-RPC message."""
config = AgentCoreSearchConfig()
body = _mcp_response_body()
sse_text = f"event: message\ndata: {json.dumps(body)}\n\n"
mock_response = _make_mock_response(text=sse_text)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
assert len(response.results) == 2
assert response.results[1].url == "https://example.com/2"
def test_transform_search_response_parses_multiline_sse_data(self):
"""SSE data may be split across several data: lines (joined per spec)."""
config = AgentCoreSearchConfig()
pretty = json.dumps(_mcp_response_body(), indent=2)
sse_text = "event: message\n" + "\n".join(f"data: {line}" for line in pretty.splitlines()) + "\n\n"
mock_response = _make_mock_response(text=sse_text)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
assert len(response.results) == 2
def test_transform_search_response_skips_progress_events(self):
"""A progress notification before the JSON-RPC result must not shadow it."""
config = AgentCoreSearchConfig()
progress = {"jsonrpc": "2.0", "method": "notifications/progress", "params": {"progress": 1}}
sse_text = (
f"event: message\ndata: {json.dumps(progress)}\n\n"
f"event: message\ndata: {json.dumps(_mcp_response_body())}\n\n"
)
mock_response = _make_mock_response(text=sse_text)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
assert len(response.results) == 2
def test_transform_search_response_raises_on_mcp_error(self):
config = AgentCoreSearchConfig()
mock_response = _make_mock_response(
{"jsonrpc": "2.0", "id": 1, "error": {"code": -32601, "message": "tool not found"}}
)
with pytest.raises(Exception, match="tool not found"):
config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
def test_transform_search_response_raises_on_tool_error(self):
"""A failed tools/call comes back as HTTP 200 with result.isError; it must not be
reported to the caller as a successful search with zero results."""
config = AgentCoreSearchConfig()
mock_response = _make_mock_response(
{
"jsonrpc": "2.0",
"id": 1,
"result": {
"isError": True,
"content": [{"type": "text", "text": "AccessDeniedException: not authorized"}],
},
}
)
with pytest.raises(Exception, match="AccessDeniedException"):
config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
def test_transform_search_response_reads_structured_content(self):
"""Connector 1.1.0+ puts the machine-readable results in structuredContent and may
leave the text block as prose, which must not come back as an empty result list."""
config = AgentCoreSearchConfig()
mock_response = _make_mock_response(
{
"jsonrpc": "2.0",
"id": 1,
"result": {
"content": [{"type": "text", "text": "Here is a prose summary of what I found."}],
"structuredContent": {"id": "824f89d0", "results": MCP_RESULTS},
},
}
)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
assert [result.title for result in response.results] == ["Test Result 1", "Test Result 2"]
assert response.results[0].url == "https://example.com/1"
assert response.results[0].snippet == "Snippet for result 1"
assert response.results[0].date == "2026-06-16"
def test_transform_search_response_does_not_duplicate_structured_content(self):
"""1.1.0+ repeats the same results in both places, so parsing both would double them."""
config = AgentCoreSearchConfig()
body = _mcp_response_body()
body["result"]["structuredContent"] = {"results": MCP_RESULTS}
mock_response = _make_mock_response(body)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
assert len(response.results) == 2
def test_transform_search_response_parses_crlf_framed_sse(self):
"""SSE streams may be CRLF framed; events must still split into separate events."""
config = AgentCoreSearchConfig()
progress = {"jsonrpc": "2.0", "method": "notifications/progress", "params": {"progress": 1}}
sse_text = (
f"event: message\r\ndata: {json.dumps(progress)}\r\n\r\n"
f"event: message\r\ndata: {json.dumps(_mcp_response_body())}\r\n\r\n"
)
mock_response = _make_mock_response(text=sse_text)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
assert len(response.results) == 2
assert response.results[0].title == "Test Result 1"
def test_sign_request_uses_bearer_token_when_api_key_set(self):
"""CUSTOM_JWT gateways: api_key is sent as a bearer token, no SigV4."""
config = AgentCoreSearchConfig()
request_data = {"jsonrpc": "2.0", "id": 1}
headers, signed_body = config.sign_request(
headers={"Content-Type": "application/json"},
optional_params={},
request_data=request_data,
api_base=GATEWAY_URL,
api_key="test-jwt-token",
)
assert headers["Authorization"] == "Bearer test-jwt-token"
assert signed_body == json.dumps(request_data).encode()
def test_sign_request_uses_bearer_token_from_env(self):
"""Server token is attached when the request targets the configured gateway host."""
config = AgentCoreSearchConfig()
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
try:
headers, _ = config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
assert headers["Authorization"] == "Bearer env-jwt-token"
finally:
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
def test_sign_request_refuses_server_token_to_untrusted_host(self):
"""Server-managed token must not be sent to a caller-chosen api_base."""
config = AgentCoreSearchConfig()
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
try:
with pytest.raises(ValueError, match="Refusing to send"):
config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base="https://attacker.example.com/mcp",
)
finally:
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
def test_sign_request_uses_env_token_for_gateway_api_base_without_gateway_url(self):
"""api_base pointing at a real gateway is a trusted destination for the env token,
so operators configuring api_base in yaml don't also need AGENTCORE_GATEWAY_URL."""
config = AgentCoreSearchConfig()
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
try:
headers, _ = config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
assert headers["Authorization"] == "Bearer env-jwt-token"
finally:
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
@pytest.mark.parametrize(
"untrusted_api_base",
[
"https://attacker.example.com/mcp",
# gateway hostname in the path/query must not pass for the host
"https://attacker.example.com/gw.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp",
],
)
def test_sign_request_refuses_sigv4_to_untrusted_host(self, untrusted_api_base):
"""A SigV4 signature carries the proxy's credential scope and session token, so it
must never be sent to a host that is not the operator's gateway."""
config = AgentCoreSearchConfig()
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
try:
with patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign:
with pytest.raises(ValueError, match="Refusing to send"):
config.sign_request(
headers={},
optional_params={"aws_region_name": "us-east-1"},
request_data={"jsonrpc": "2.0"},
api_base=untrusted_api_base,
)
mock_base_sign.assert_not_called()
finally:
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
@pytest.mark.parametrize(
"plaintext_api_base",
[
"http://gw.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp",
"http://internal-gateway.corp/mcp",
],
)
def test_sign_request_refuses_server_token_over_plaintext_http(self, plaintext_api_base):
"""A trusted hostname over plain http would expose the bearer token to
network observers, so credentials only ride https (or localhost)."""
config = AgentCoreSearchConfig()
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
os.environ["AGENTCORE_GATEWAY_URL"] = plaintext_api_base
try:
with pytest.raises(ValueError, match="plaintext"):
config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=plaintext_api_base,
)
finally:
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
def test_sign_request_refuses_sigv4_over_plaintext_http(self):
"""Same for SigV4: a signature over plain http is replayable by observers."""
config = AgentCoreSearchConfig()
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
with patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign:
with pytest.raises(ValueError, match="plaintext"):
config.sign_request(
headers={},
optional_params={"aws_region_name": "us-east-1"},
request_data={"jsonrpc": "2.0"},
api_base="http://gw.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp",
)
mock_base_sign.assert_not_called()
def test_sign_request_allows_plain_http_for_localhost(self):
"""Local development against an MCP stub on 127.0.0.1 keeps working."""
config = AgentCoreSearchConfig()
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
os.environ["AGENTCORE_GATEWAY_URL"] = "http://127.0.0.1:8931/mcp"
try:
headers, _ = config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base="http://127.0.0.1:8931/mcp",
)
assert headers["Authorization"] == "Bearer env-jwt-token"
finally:
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
def test_sign_request_does_not_leak_bedrock_bearer_token(self):
"""AWS_BEARER_TOKEN_BEDROCK is a Bedrock Runtime credential — it must not
replace SigV4 on requests to an AgentCore gateway."""
config = AgentCoreSearchConfig()
with patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign:
config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
# api_key="" (falsy, not None) disables the base class's
# AWS_BEARER_TOKEN_BEDROCK env fallback.
assert mock_base_sign.call_args.kwargs["api_key"] == ""
def test_sign_request_custom_hostname_requires_region(self):
"""Custom hostname + empty AWS config chain → clear error, no guessed region."""
config = AgentCoreSearchConfig()
custom_url = "https://gateway.internal.example.com/mcp"
os.environ["AGENTCORE_GATEWAY_URL"] = custom_url
mock_session = MagicMock()
mock_session.region_name = None # nothing configured anywhere
try:
with patch("boto3.Session", return_value=mock_session):
with pytest.raises(ValueError, match="signing region"):
config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=custom_url,
)
finally:
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
def test_sign_request_custom_hostname_uses_shared_config_region(self):
"""Custom hostname + region from AWS shared config (profile) must be honored."""
config = AgentCoreSearchConfig()
custom_url = "https://gateway.internal.example.com/mcp"
os.environ["AGENTCORE_GATEWAY_URL"] = custom_url
mock_session = MagicMock()
mock_session.region_name = "eu-west-1" # e.g. from ~/.aws/config profile
try:
with (
patch("boto3.Session", return_value=mock_session),
patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign,
):
config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=custom_url,
)
assert mock_base_sign.call_args.kwargs["optional_params"]["aws_region_name"] == "eu-west-1"
finally:
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
def test_sign_request_passes_explicit_aws_credentials(self):
"""Explicit aws_* params (e.g. from a proxy search_tools entry) reach the signer."""
config = AgentCoreSearchConfig()
with patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign:
config.sign_request(
headers={},
optional_params={
"aws_access_key_id": "AKIATEST",
"aws_secret_access_key": "secret",
"aws_session_token": "token",
},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
passed = mock_base_sign.call_args.kwargs["optional_params"]
assert passed["aws_access_key_id"] == "AKIATEST"
assert passed["aws_secret_access_key"] == "secret"
assert passed["aws_session_token"] == "token"
def test_sign_request_derives_region_from_gateway_url(self):
"""Signing region must come from the gateway URL, not the caller's default region."""
config = AgentCoreSearchConfig()
eu_url = "https://gw-x.gateway.bedrock-agentcore.eu-central-1.amazonaws.com/mcp"
with patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign:
config.sign_request(
headers={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=eu_url,
)
assert mock_base_sign.call_args.kwargs["optional_params"]["aws_region_name"] == "eu-central-1"
class TestAgentCoreSearchEdgeCases:
"""Branch coverage for response parsing and error mapping."""
def test_transform_search_response_skips_non_text_and_bad_json_blocks(self):
"""Non-text blocks and unparseable text blocks are skipped, not fatal."""
config = AgentCoreSearchConfig()
body = {
"jsonrpc": "2.0",
"id": 1,
"result": {
"content": [
{"type": "image", "data": "..."},
{"type": "text", "text": "not-json"},
{"type": "text", "text": json.dumps(["scalar", {"title": "T", "url": "u", "text": "s"}])},
]
},
}
mock_response = _make_mock_response(body)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
# only the one dict item survives; non-dict list entries are skipped
assert len(response.results) == 1
assert response.results[0].title == "T"
def test_parse_mcp_body_sse_without_json_frame_raises(self):
"""An SSE stream carrying no parseable JSON object is a 502."""
config = AgentCoreSearchConfig()
mock_response = _make_mock_response(text="event: ping\ndata: not-json\n\n")
with pytest.raises(Exception, match="SSE without a JSON data frame"):
config._parse_mcp_body(mock_response)
def test_parse_mcp_body_returns_last_event_when_no_result_frame(self):
"""A stream of only notifications returns the last parsed event."""
config = AgentCoreSearchConfig()
note = {"jsonrpc": "2.0", "method": "notifications/progress"}
mock_response = _make_mock_response(text=f"data: {json.dumps(note)}\n\n")
assert config._parse_mcp_body(mock_response) == note
def test_sign_request_rejects_list_request_body(self):
config = AgentCoreSearchConfig()
with pytest.raises(TypeError, match="single dict"):
config.sign_request(
headers={},
optional_params={},
request_data=[{"jsonrpc": "2.0"}],
api_base=GATEWAY_URL,
)
def test_get_error_class_maps_status_and_message(self):
config = AgentCoreSearchConfig()
err = config.get_error_class(error_message="boom", status_code=503, headers={})
assert getattr(err, "status_code", None) == 503
assert "boom" in str(err)
def test_search_cost_lookup_is_mapped(self, monkeypatch):
"""Assert against the map in this checkout: the remote cost map litellm loads by
default only carries providers already released."""
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
from litellm.search.cost_calculator import search_provider_cost_per_query
monkeypatch.setattr(litellm, "model_cost", GetModelCostMap.load_local_model_cost_map())
assert search_provider_cost_per_query(model="agentcore/search", custom_llm_provider="agentcore") == (0.0, 0.0)

View file

@ -2226,3 +2226,78 @@ def test_direct_vector_store_search_debug_log_omits_stored_credentials(caplog, i
logged = "\n".join(record.getMessage() for record in caplog.records)
assert "sup3r-s3cret-valkey-pw" not in logged
assert "sk-embedding-s3cret" not in logged
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_carries_deployment_vertex_location_for_pricing(monkeypatch):
"""
The proxy pre-creates the logging object before the router picks a deployment, so the
native /v1/messages path must copy the deployment's vertex_location into the logging
params it updates; otherwise cost resolution falls back to the environment and every
call on this surface prices with the regional uplift (#34393).
"""
import contextlib
from datetime import datetime
from litellm.litellm_core_utils.litellm_logging import (
Logging,
_resolve_vertex_location_for_cost,
)
monkeypatch.setenv("VERTEXAI_LOCATION", "us-east5")
monkeypatch.setattr(litellm, "vertex_location", None)
handler = BaseLLMHTTPHandler()
async def logging_obj_after_handler(generic_params):
logging_obj = Logging(
model="vertex_ai/claude-haiku-4-5@20251001",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="anthropic_messages",
start_time=datetime.now(),
litellm_call_id="vertex-messages-location",
function_id="f",
)
logging_obj.update_environment_variables(
model="vertex_ai/claude-haiku-4-5@20251001",
user="",
optional_params={},
litellm_params={"api_base": ""},
custom_llm_provider="vertex_ai",
)
mock_config = Mock()
mock_config.validate_anthropic_messages_environment = Mock(
return_value=({"authorization": "Bearer t"}, "https://us-east5-aiplatform.googleapis.com")
)
mock_config.transform_anthropic_messages_request = Mock(
return_value={"model": "claude-haiku-4-5@20251001", "messages": []}
)
with contextlib.suppress(Exception):
await handler.async_anthropic_messages_handler(
model="claude-haiku-4-5@20251001",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=mock_config,
anthropic_messages_optional_request_params={"max_tokens": 10},
custom_llm_provider="vertex_ai",
litellm_params=generic_params,
logging_obj=logging_obj,
client=AsyncMock(),
kwargs={},
)
return logging_obj
global_deployment = await logging_obj_after_handler(GenericLiteLLMParams(vertex_location="global"))
assert global_deployment.litellm_params["vertex_location"] == "global"
assert (
_resolve_vertex_location_for_cost(
custom_llm_provider="vertex_ai",
litellm_params=global_deployment.litellm_params,
optional_params=global_deployment.optional_params,
model="claude-haiku-4-5@20251001",
)
== "global"
)
unconfigured_deployment = await logging_obj_after_handler(GenericLiteLLMParams())
assert "vertex_location" not in unconfigured_deployment.litellm_params

View file

@ -6392,3 +6392,168 @@ def test_is_user_proxy_admin_rejects_view_only_admin():
assert _is_user_proxy_admin(user_obj=viewer) is False
assert _is_user_proxy_admin(user_obj=admin) is True
assert _is_user_proxy_admin(user_obj=None) is False
def _make_wildcard_access_group_router():
"""
`openai/*` tagged into an access group, plus an untagged `azure/*`, mirroring a
proxy that fronts a whole provider behind one wildcard deployment.
"""
from litellm import Router
return Router(
model_list=[
{
"model_name": "openai/*",
"litellm_params": {"model": "openai/*", "api_key": "fake"},
"model_info": {
"id": "wildcard-openai",
"access_groups": ["default-models"],
},
},
{
"model_name": "azure/*",
"litellm_params": {"model": "azure/*", "api_key": "fake"},
"model_info": {"id": "wildcard-azure"},
},
]
)
def test_can_object_call_model_access_group_wildcard_accepts_bare_model_name():
"""
Regression: a key holding only the access group name was denied for `gpt-4o`
while `openai/gpt-4o` was allowed, because group membership resolved through the
pattern router's raw regex and skipped the `{provider}/{model}` retry that both
routing and the direct-wildcard grant already perform.
"""
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_wildcard_access_group_router()
assert (
_can_object_call_model(
model="gpt-4o",
llm_router=router,
models=["default-models"],
object_type="key",
)
is True
)
def test_can_object_call_model_access_group_wildcard_accepts_prefixed_model_name():
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_wildcard_access_group_router()
assert (
_can_object_call_model(
model="openai/gpt-4o",
llm_router=router,
models=["default-models"],
object_type="key",
)
is True
)
@pytest.mark.parametrize(
"model",
[
"totally-made-up-model-zzz", # no provider can be inferred
"azure/some-deployment", # wildcard exists but carries no access group
],
)
def test_can_object_call_model_access_group_wildcard_does_not_over_grant(model):
"""The bare-name retry must not turn an access group into a blanket grant."""
from litellm.proxy._types import ProxyException
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = _make_wildcard_access_group_router()
with pytest.raises(ProxyException):
_can_object_call_model(
model=model,
llm_router=router,
models=["default-models"],
object_type="key",
)
def test_can_object_call_model_access_group_rejects_unconsumed_namespace():
"""
`bedrockz/...` infers provider `bedrock` from a fragment of the name, so
re-prefixing would smuggle an unrecognized namespace through a `bedrock/*` group.
"""
from litellm import Router
from litellm.proxy._types import ProxyException
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = Router(
model_list=[
{
"model_name": "bedrock/*",
"litellm_params": {"model": "bedrock/*"},
"model_info": {
"id": "wildcard-bedrock",
"access_groups": ["bedrock-models"],
},
}
]
)
assert (
_can_object_call_model(
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_router=router,
models=["bedrock-models"],
object_type="key",
)
is True
)
with pytest.raises(ProxyException):
_can_object_call_model(
model="bedrockz/anthropic.claude-3-5-sonnet-20240620-v1:0",
llm_router=router,
models=["bedrock-models"],
object_type="key",
)
def test_can_object_call_model_team_scoped_wildcard_accepts_bare_model_name():
"""
Same regression as the proxy-wide wildcard, but for a team-scoped deployment
whose public name is a wildcard: those live in a separate per-team pattern
index that needed the same `{provider}/{model}` retry.
"""
from litellm import Router
from litellm.proxy.auth.auth_checks import _can_object_call_model
router = Router(
model_list=[
{
"model_name": "openai/*_team-a_abc",
"litellm_params": {"model": "openai/*", "api_key": "fake"},
"model_info": {
"id": "team-byok-wildcard",
"team_id": "team-a",
"team_public_model_name": "openai/*",
"access_groups": ["team-models"],
},
}
]
)
for model in ("gpt-4o", "openai/gpt-4o"):
assert (
_can_object_call_model(
model=model,
llm_router=router,
models=["team-models"],
object_type="team",
team_id="team-a",
)
is True
)

View file

@ -25,6 +25,7 @@ from litellm.proxy.client.cli.commands.auth import (
save_token,
whoami,
)
from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner
def _mock_cli_sso_start_response(
@ -267,7 +268,7 @@ class TestTokenUtilities:
def test_load_token_io_error(self):
"""Test loading token with IO error"""
with (
patch("builtins.open", side_effect=IOError("Permission denied")),
patch("builtins.open", side_effect=OSError("Permission denied")),
patch("litellm.proxy.client.cli.commands.auth.get_token_file_path") as mock_path,
patch("os.path.exists", return_value=True),
):
@ -1029,3 +1030,83 @@ class TestSaveTokenPrivateWrite:
assert json.loads(token_file.read_text()) == {"key": "sk-original", "timestamp": 1234567890}
assert list(token_file.parent.glob(".tmp-*")) == []
class TestLoginConfigClaude:
"""`lite login --config-claude` wiring into ~/.claude/settings.json"""
def setup_method(self):
self.runner = CliRunner()
def _run_login(self, tmp_path, args, base_url="https://test.example.com"):
settings_path = tmp_path / "claude" / "settings.json"
backup_path = tmp_path / "claude_settings_backup.json"
poll_response = Mock()
poll_response.status_code = 200
poll_response.json.return_value = {
"status": "ready",
"key": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt",
"user_id": "test-user-123",
"team_id": "team-1",
"teams": ["team-1"],
}
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=poll_response),
patch("litellm.proxy.client.cli.commands.auth.save_token"),
patch("litellm.proxy.client.cli.interface.show_commands"),
patch("litellm.proxy.client.cli.commands.auth.CLAUDE_SETTINGS_PATH", settings_path),
patch(
"litellm.proxy.client.cli.commands.auth.SETTINGS_FILE_OWNERS",
(SettingsFileOwner(backup_path, "lite up", "lite down"),),
),
patch(
"litellm.proxy.client.cli.commands.claude_settings.shutil.which",
return_value="/usr/local/bin/lite",
),
):
result = self.runner.invoke(login, args, obj={"base_url": base_url})
return result, settings_path, backup_path
def test_default_login_does_not_touch_claude_settings(self, tmp_path):
result, settings_path, _backup_path = self._run_login(tmp_path, [])
assert result.exit_code == 0
assert "Login successful!" in result.output
assert not settings_path.exists()
assert "Configured Claude Code" not in result.output
def test_flag_writes_the_settings_file_and_reports_success(self, tmp_path):
result, settings_path, _backup_path = self._run_login(tmp_path, ["--config-claude"])
assert result.exit_code == 0
written = json.loads(settings_path.read_text())
assert written["env"]["ANTHROPIC_BASE_URL"] == "https://test.example.com"
assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://test.example.com auth print-token"
assert "Configured Claude Code" in result.output
def test_flag_preserves_unrelated_settings_on_an_existing_file(self, tmp_path):
settings_path = tmp_path / "claude" / "settings.json"
settings_path.parent.mkdir(parents=True)
settings_path.write_text(json.dumps({"theme": "dark", "env": {"KEEP": "me"}}))
result, _settings_path, _backup_path = self._run_login(tmp_path, ["--config-claude"])
assert result.exit_code == 0
written = json.loads(settings_path.read_text())
assert written["theme"] == "dark"
assert written["env"]["KEEP"] == "me"
def test_settings_failure_is_reported_without_claiming_login_failed(self, tmp_path):
settings_path = tmp_path / "claude" / "settings.json"
settings_path.parent.mkdir(parents=True)
settings_path.write_text("not json at all {{{")
result, _settings_path, _backup_path = self._run_login(tmp_path, ["--config-claude"])
assert result.exit_code != 0
assert "Login successful!" in result.output
assert "could not configure Claude Code" in result.output
assert "invalid JSON" in result.output
assert "Authentication failed" not in result.output

View file

@ -0,0 +1,269 @@
import json
import shlex
import stat
import time
from unittest.mock import patch
import pytest
from click.testing import CliRunner
from litellm.proxy.client.cli import cli
from litellm.proxy.client.cli.commands.claude_settings import (
AUTOROUTE_BACKUP_PATH,
BACKUP_PATH,
SETTINGS_FILE_OWNERS,
ClaudeSettingsError,
SettingsFileOwner,
resolve_api_key_helper,
write_claude_settings,
)
def _owners(*backup_paths):
"""Stand-in owners for the real `lite up` / `lite autoroute up` registry."""
return tuple(SettingsFileOwner(path, "lite up", "lite down") for path in backup_paths)
CLAUDE_SETTINGS_MODULE = "litellm.proxy.client.cli.commands.claude_settings"
AUTH_MODULE = "litellm.proxy.client.cli.commands.auth"
@pytest.fixture
def paths(tmp_path):
return tmp_path / "claude" / "settings.json", tmp_path / "backup.json"
@pytest.fixture
def lite_on_path():
with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value="/usr/local/bin/lite"):
yield
class TestWriteClaudeSettings:
def test_creates_the_file_and_its_parent_when_missing(self, paths, lite_on_path):
settings_path, backup_path = paths
assert not settings_path.parent.exists()
write_claude_settings("https://proxy.example.com/", settings_path, _owners(backup_path))
written = json.loads(settings_path.read_text())
assert written["env"]["ANTHROPIC_BASE_URL"] == "https://proxy.example.com"
assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://proxy.example.com auth print-token"
def test_updates_an_existing_file_preserving_unrelated_settings(self, paths, lite_on_path):
settings_path, backup_path = paths
settings_path.parent.mkdir(parents=True)
settings_path.write_text(
json.dumps(
{
"theme": "dark",
"permissions": {"allow": ["Bash"]},
"env": {"SOME_OTHER_VAR": "keep-me", "ANTHROPIC_BASE_URL": "https://old.example.com"},
"apiKeyHelper": "old-helper",
}
)
)
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
written = json.loads(settings_path.read_text())
assert written["theme"] == "dark"
assert written["permissions"] == {"allow": ["Bash"]}
assert written["env"]["SOME_OTHER_VAR"] == "keep-me"
assert written["env"]["ANTHROPIC_BASE_URL"] == "https://proxy.example.com"
assert written["apiKeyHelper"] != "old-helper"
def test_rerunning_against_a_new_proxy_refreshes_both_base_url_and_helper(self, paths, lite_on_path):
settings_path, backup_path = paths
write_claude_settings("https://first.example.com", settings_path, _owners(backup_path))
write_claude_settings("https://second.example.com", settings_path, _owners(backup_path))
written = json.loads(settings_path.read_text())
assert written["env"]["ANTHROPIC_BASE_URL"] == "https://second.example.com"
assert "second.example.com" in written["apiKeyHelper"]
assert "first.example.com" not in written["apiKeyHelper"]
def test_drops_a_stray_static_api_key_so_the_helper_token_wins(self, paths, lite_on_path):
settings_path, backup_path = paths
settings_path.parent.mkdir(parents=True)
settings_path.write_text(json.dumps({"env": {"ANTHROPIC_API_KEY": "sk-leaked"}}))
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
assert "ANTHROPIC_API_KEY" not in json.loads(settings_path.read_text())["env"]
def test_written_file_is_owner_only(self, paths, lite_on_path):
settings_path, backup_path = paths
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
assert stat.S_IMODE(settings_path.stat().st_mode) == 0o600
def test_refuses_while_lite_up_holds_a_backup(self, paths, lite_on_path):
settings_path, backup_path = paths
backup_path.write_text("{}")
with pytest.raises(ClaudeSettingsError, match="lite down"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
assert not settings_path.exists()
def test_refuses_on_corrupt_existing_settings_without_touching_the_file(self, paths, lite_on_path):
settings_path, backup_path = paths
settings_path.parent.mkdir(parents=True)
settings_path.write_text("not json at all {{{")
with pytest.raises(ClaudeSettingsError, match="invalid JSON"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
assert settings_path.read_text() == "not json at all {{{"
def test_reports_an_actionable_error_when_lite_is_not_on_path(self, paths):
settings_path, backup_path = paths
with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value=None):
with pytest.raises(ClaudeSettingsError, match="Could not find `lite`"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
assert not settings_path.exists()
def test_reports_an_actionable_error_on_a_non_utf8_file(self, paths, lite_on_path):
"""Bytes that are not valid UTF-8 must not escape as UnicodeDecodeError.
UnicodeDecodeError is a ValueError, not an OSError, so a decode-side catch
is easy to miss; login's broad `except Exception` would then relabel it as
an authentication failure and exit 0.
"""
settings_path, backup_path = paths
settings_path.parent.mkdir(parents=True)
settings_path.write_bytes(b'{"theme": "\xff\xfe"}')
with pytest.raises(ClaudeSettingsError, match="invalid JSON"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
def test_reports_an_actionable_error_when_the_file_cannot_be_read(self, paths, lite_on_path):
"""An unreadable settings file must not surface as "Authentication failed".
login wraps the whole flow in a broad `except Exception`, so any OSError
escaping this function gets relabelled as an auth failure and sends the
user looking at their SSO config instead of at file permissions.
"""
settings_path, backup_path = paths
settings_path.parent.mkdir(parents=True)
settings_path.mkdir()
with pytest.raises(ClaudeSettingsError, match="Could not read"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
def test_reports_an_actionable_error_when_the_file_cannot_be_written(self, paths, lite_on_path):
settings_path, backup_path = paths
with patch(
f"{CLAUDE_SETTINGS_MODULE}.write_private_json",
side_effect=OSError("Read-only file system"),
):
with pytest.raises(ClaudeSettingsError, match="Read-only file system"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
class TestApiKeyHelperIsActuallyInvocable:
"""The helper string is executed verbatim by Claude Code, so it has to parse.
Asserting only on its text is what let a malformed command (`--base-url`, a
top-level group option, placed after the `print-token` subcommand) ship: click
rejects it with "No such option" and every Claude Code request loses its token.
"""
def _helper_args(self, base_url):
with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value="/usr/local/bin/lite"):
return shlex.split(resolve_api_key_helper(base_url))[1:]
def test_the_generated_command_parses(self):
result = CliRunner().invoke(cli, self._helper_args("http://localhost:4000"))
assert "No such option" not in result.output
assert result.exit_code != 2
def test_the_generated_command_reaches_print_token(self):
with patch(f"{AUTH_MODULE}.load_token", return_value=None):
result = CliRunner().invoke(cli, self._helper_args("http://localhost:4000"))
assert "Not authenticated" in result.output
def test_the_generated_command_carries_the_base_url_through(self):
stale = {
"base_url": "http://other-proxy.example.com",
"key": "sk-stale",
"timestamp": time.time(),
}
with patch(f"{AUTH_MODULE}.load_token", return_value=stale):
result = CliRunner().invoke(cli, self._helper_args("http://localhost:4000"))
assert "Not authenticated for this server" in result.output
class TestConflictingOwnersOfTheSettingsFile:
"""Both `lite up` and `lite autoroute up` restore a backup when they stop.
Guarding only one of them leaves the other free to silently revert this
write, which is the exact hazard the guard exists to prevent.
"""
def test_any_owner_holding_a_backup_blocks_the_write(self, tmp_path, lite_on_path):
settings_path = tmp_path / "claude" / "settings.json"
for index, owner in enumerate(SETTINGS_FILE_OWNERS):
backup = tmp_path / f"backup-{index}.json"
backup.write_text("{}")
stand_in = SettingsFileOwner(backup, owner.start_command, owner.stop_command)
with pytest.raises(ClaudeSettingsError, match="currently managing"):
write_claude_settings("https://proxy.example.com", settings_path, (stand_in,))
backup.unlink()
assert not settings_path.exists()
def test_the_error_names_the_owner_that_actually_holds_the_file(self, tmp_path, lite_on_path):
settings_path = tmp_path / "claude" / "settings.json"
backup = tmp_path / "auto.json"
backup.write_text("{}")
autoroute = SettingsFileOwner(backup, "lite autoroute up", "lite autoroute down")
with pytest.raises(ClaudeSettingsError, match="`lite autoroute up` is currently managing"):
write_claude_settings("https://proxy.example.com", settings_path, (autoroute,))
with pytest.raises(ClaudeSettingsError, match="Run `lite autoroute down` first"):
write_claude_settings("https://proxy.example.com", settings_path, (autoroute,))
def test_the_registry_matches_the_paths_the_commands_actually_use(self):
"""A second definition of the autoroute dir must not drift from this one."""
from litellm.proxy.client.cli.commands.autoroute.process import AUTOROUTE_DIR
assert AUTOROUTE_BACKUP_PATH == AUTOROUTE_DIR / "claude_settings_backup.json"
assert {o.backup_path for o in SETTINGS_FILE_OWNERS} == {BACKUP_PATH, AUTOROUTE_BACKUP_PATH}
assert {o.stop_command for o in SETTINGS_FILE_OWNERS} == {"lite down", "lite autoroute down"}
class TestDoesNotDestroyUserOwnedStructure:
def test_writes_through_a_symlinked_settings_file(self, tmp_path, lite_on_path):
"""os.replace() swaps the symlink for a regular file, detaching a dotfiles repo.
There is no backup here to undo that, so the link must survive and its
target must be the thing that gets updated.
"""
real = tmp_path / "dotfiles" / "settings.json"
real.parent.mkdir()
real.write_text(json.dumps({"theme": "dark"}))
link = tmp_path / "claude" / "settings.json"
link.parent.mkdir()
link.symlink_to(real)
write_claude_settings("https://proxy.example.com", link, ())
assert link.is_symlink()
assert json.loads(real.read_text())["env"]["ANTHROPIC_BASE_URL"] == "https://proxy.example.com"
assert json.loads(real.read_text())["theme"] == "dark"
def test_refuses_rather_than_discarding_a_non_object_env(self, paths, lite_on_path):
"""merge coerces a non-dict env to {}; that is silent data loss on a persistent write."""
settings_path, backup_path = paths
settings_path.parent.mkdir(parents=True)
settings_path.write_text(json.dumps({"theme": "dark", "env": "not-an-object"}))
with pytest.raises(ClaudeSettingsError, match="non-object"):
write_claude_settings("https://proxy.example.com", settings_path, _owners(backup_path))
assert json.loads(settings_path.read_text())["env"] == "not-an-object"

View file

@ -10,6 +10,7 @@ from click.testing import CliRunner
from litellm.proxy.client.cli.commands import up as up_module
from litellm.proxy.client.cli.commands.agents import AgentRunError
from litellm.proxy.client.cli.commands.claude_settings import ClaudeSettingsError
from litellm.proxy.client.cli.commands.up import (
BackupRecord,
UpError,
@ -25,6 +26,7 @@ from litellm.proxy.client.cli.commands.up import (
)
UP_MODULE = "litellm.proxy.client.cli.commands.up"
AUTH_MODULE = "litellm.proxy.client.cli.commands.auth"
def _patch_paths(monkeypatch, tmp_path):
@ -92,13 +94,13 @@ class TestLoadJsonOrEmpty:
def test_raises_clean_error_on_invalid_json(self, tmp_path):
path = tmp_path / "settings.json"
path.write_text("not json at all {{{")
with pytest.raises(UpError, match="invalid JSON"):
with pytest.raises(ClaudeSettingsError, match="invalid JSON"):
load_json_or_empty(path)
def test_raises_clean_error_on_non_object_root(self, tmp_path):
path = tmp_path / "settings.json"
path.write_text(json.dumps([1, 2, 3]))
with pytest.raises(UpError, match="invalid JSON"):
with pytest.raises(ClaudeSettingsError, match="invalid JSON"):
load_json_or_empty(path)
@ -201,16 +203,16 @@ class TestResolveApiKeyHelper:
def test_returns_helper_command_bound_to_the_selected_proxy(self, monkeypatch):
monkeypatch.setattr(shutil, "which", lambda name: "/usr/local/bin/lite")
helper = resolve_api_key_helper("http://localhost:4000")
assert helper == "/usr/local/bin/lite auth print-token --base-url http://localhost:4000"
assert helper == "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"
def test_quotes_a_base_url_containing_shell_metacharacters(self, monkeypatch):
monkeypatch.setattr(shutil, "which", lambda name: "/usr/local/bin/lite")
helper = resolve_api_key_helper("http://example.com/path; rm -rf /")
assert helper == "/usr/local/bin/lite auth print-token --base-url 'http://example.com/path; rm -rf /'"
assert helper == "/usr/local/bin/lite --base-url 'http://example.com/path; rm -rf /' auth print-token"
def test_raises_when_lite_not_on_path(self, monkeypatch):
monkeypatch.setattr(shutil, "which", lambda name: None)
with pytest.raises(UpError, match="Could not find `lite`"):
with pytest.raises(ClaudeSettingsError, match="Could not find `lite`"):
resolve_api_key_helper("http://localhost:4000")
@ -396,3 +398,50 @@ class TestDownCommand:
assert result.exit_code != 0
assert result.exception is None or isinstance(result.exception, SystemExit)
assert "invalid or unexpected JSON" in result.output
class TestUpCanInvokeTheRealLoginCommand:
"""`lite up` calls ctx.invoke(login) on the real command object.
Every other test in this file monkeypatches `up_module.login` with a fake, so
none of them would notice a login parameter that ctx.invoke cannot supply.
"""
def test_ctx_invoke_supplies_every_login_parameter(self):
from litellm.proxy.client.cli.commands.auth import login as real_login
reached = []
@click.command()
@click.pass_context
def driver(ctx):
ctx.obj = {"base_url": "http://127.0.0.1:9"}
ctx.invoke(real_login)
with patch(
f"{AUTH_MODULE}._start_cli_sso_flow",
side_effect=lambda base_url: reached.append(base_url) or RuntimeError("stop"),
):
result = CliRunner().invoke(driver, [], standalone_mode=False)
assert not isinstance(result.exception, TypeError), result.exception
assert reached == ["http://127.0.0.1:9"]
def test_ctx_invoke_leaves_claude_settings_alone(self, tmp_path):
from litellm.proxy.client.cli.commands.auth import login as real_login
settings_path = tmp_path / "settings.json"
@click.command()
@click.pass_context
def driver(ctx):
ctx.obj = {"base_url": "http://127.0.0.1:9"}
ctx.invoke(real_login)
with (
patch(f"{AUTH_MODULE}.CLAUDE_SETTINGS_PATH", settings_path),
patch(f"{AUTH_MODULE}._start_cli_sso_flow", side_effect=RuntimeError("stop")),
):
CliRunner().invoke(driver, [], standalone_mode=False)
assert not settings_path.exists()

View file

@ -2533,3 +2533,187 @@ async def test_daily_transaction_internal_call_keeps_spend_but_not_request_count
assert internal["autorouter_savings_spend"] == 0.0
assert user_sent["api_requests"] == 1
assert user_sent["successful_requests"] == 1
def _deadlock_error():
from prisma.errors import RawQueryError
return RawQueryError(
data={"user_facing_error": {"error_code": "P2034", "meta": {"table": "LiteLLM_VerificationToken"}}}
)
def _empty_spend_transactions(**overrides):
base = {
"user_list_transactions": {},
"end_user_list_transactions": {},
"key_list_transactions": {},
"team_list_transactions": {},
"team_member_list_transactions": {},
"org_list_transactions": {},
"tag_list_transactions": {},
"agent_list_transactions": {},
}
return {**base, **overrides}
def _good_tx(mock_batcher):
tx = AsyncMock()
tx.__aenter__ = AsyncMock(return_value=tx)
tx.__aexit__ = AsyncMock(return_value=False)
tx.batch_ = MagicMock(
return_value=AsyncMock(
__aenter__=AsyncMock(return_value=mock_batcher),
__aexit__=AsyncMock(return_value=False),
)
)
return tx
def _failing_tx(error):
tx = MagicMock()
tx.__aenter__ = AsyncMock(side_effect=error)
tx.__aexit__ = AsyncMock(return_value=False)
return tx
@pytest.mark.asyncio
async def test_commit_spend_updates_retries_deadlock_then_commits(monkeypatch):
"""Regression: a deadlock on the key-spend UPDATE is retried and commits the increment exactly once."""
slept = []
monkeypatch.setattr(
"litellm.proxy.db.db_spend_update_writer.asyncio.sleep",
AsyncMock(side_effect=lambda s: slept.append(s)),
)
mock_batcher = MagicMock()
mock_prisma_client = MagicMock()
mock_prisma_client.db.tx = MagicMock(side_effect=[_failing_tx(_deadlock_error()), _good_tx(mock_batcher)])
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
await DBSpendUpdateWriter()._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=3,
proxy_logging_obj=proxy_logging,
db_spend_update_transactions=_empty_spend_transactions(key_list_transactions={"sk-abc": 0.5}),
)
assert mock_prisma_client.db.tx.call_count == 2
mock_batcher.litellm_verificationtoken.update_many.assert_called_once()
call_kwargs = mock_batcher.litellm_verificationtoken.update_many.call_args[1]
assert call_kwargs["where"] == {"token": "sk-abc"}
assert call_kwargs["data"]["spend"] == {"increment": 0.5}
assert len(slept) == 1
proxy_logging.failure_handler.assert_not_called()
@pytest.mark.asyncio
async def test_commit_spend_updates_raises_after_exhausting_deadlock_retries(monkeypatch):
"""A deadlock that never clears must surface after the retry budget is spent, not loop or swallow."""
monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", AsyncMock(return_value=None))
mock_prisma_client = MagicMock()
mock_prisma_client.db.tx = MagicMock(side_effect=lambda *a, **k: _failing_tx(_deadlock_error()))
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
from prisma.errors import RawQueryError
with pytest.raises(RawQueryError):
await DBSpendUpdateWriter()._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=2,
proxy_logging_obj=proxy_logging,
db_spend_update_transactions=_empty_spend_transactions(key_list_transactions={"sk-abc": 0.5}),
)
assert mock_prisma_client.db.tx.call_count == 3
@pytest.mark.asyncio
async def test_commit_spend_updates_does_not_retry_non_deadlock_data_error(monkeypatch):
"""A non-retryable data-layer error raises on the first attempt, never retried against the increment."""
monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", AsyncMock(return_value=None))
from prisma.errors import UniqueViolationError
non_deadlock = UniqueViolationError(
data={"user_facing_error": {"error_code": "P2002", "meta": {"table": "LiteLLM_VerificationToken"}}}
)
mock_prisma_client = MagicMock()
mock_prisma_client.db.tx = MagicMock(side_effect=lambda *a, **k: _failing_tx(non_deadlock))
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
with pytest.raises(UniqueViolationError):
await DBSpendUpdateWriter()._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=3,
proxy_logging_obj=proxy_logging,
db_spend_update_transactions=_empty_spend_transactions(key_list_transactions={"sk-abc": 0.5}),
)
mock_prisma_client.db.tx.assert_called_once()
@pytest.mark.asyncio
async def test_update_daily_spend_retries_deadlock(monkeypatch):
"""The daily-spend upsert path retries a deadlock on the bulk upsert and then drains successfully."""
mock_prisma_client = MagicMock()
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[_deadlock_error(), None])
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", AsyncMock(return_value=None))
daily_spend_transactions = {"k1": _daily_txn()}
await DBSpendUpdateWriter._update_daily_spend(
n_retry_times=3,
prisma_client=mock_prisma_client,
proxy_logging_obj=proxy_logging,
daily_spend_transactions=daily_spend_transactions,
entity_type="user",
entity_id_field="user_id",
)
assert mock_prisma_client.db.execute_raw.call_count == 2
assert daily_spend_transactions == {}
proxy_logging.failure_handler.assert_not_called()
@pytest.mark.parametrize(
"transactions_key, sample_key",
[
("user_list_transactions", "user-1"),
("team_list_transactions", "team-1"),
("team_member_list_transactions", "team_id::team-1::user_id::user-1"),
("org_list_transactions", "org-1"),
("tag_list_transactions", "tag-1"),
("agent_list_transactions", "agent-1"),
],
)
@pytest.mark.asyncio
async def test_commit_spend_updates_retries_deadlock_on_every_entity_path(monkeypatch, transactions_key, sample_key):
"""Every per-entity spend path, not just keys, retries a deadlock instead of dropping the increment."""
monkeypatch.setattr("litellm.proxy.db.db_spend_update_writer.asyncio.sleep", AsyncMock(return_value=None))
mock_batcher = MagicMock()
mock_prisma_client = MagicMock()
mock_prisma_client.db.tx = MagicMock(side_effect=[_failing_tx(_deadlock_error()), _good_tx(mock_batcher)])
proxy_logging = MagicMock()
proxy_logging.failure_handler = AsyncMock()
proxy_logging.call_details = {}
await DBSpendUpdateWriter()._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=3,
proxy_logging_obj=proxy_logging,
db_spend_update_transactions=_empty_spend_transactions(**{transactions_key: {sample_key: 0.5}}),
)
assert mock_prisma_client.db.tx.call_count == 2
proxy_logging.failure_handler.assert_not_called()

View file

@ -549,3 +549,36 @@ def test_handle_db_exception_surfaces_a_permanent_fault_even_when_degraded_mode_
with pytest.raises(BinaryNotFoundError):
PrismaDBExceptionHandler.handle_db_exception(BinaryNotFoundError("query engine binary not found"))
@pytest.mark.parametrize(
"error",
[
RawQueryError(data={"user_facing_error": {"error_code": "P2034", "meta": {"table": "t"}}}),
PrismaError("Transaction failed due to a write conflict or a deadlock. Please retry your transaction"),
RawQueryError(data={"user_facing_error": {"message": "deadlock detected", "meta": {"table": "t"}}}),
RawQueryError(
data={"user_facing_error": {"message": "ERROR: 40P01: deadlock detected", "meta": {"table": "t"}}}
),
],
)
def test_is_deadlock_error_matches_postgres_deadlock(error):
"""A Postgres deadlock surfaced through prisma (P2034 or 40P01 / "deadlock detected" text) is recognized."""
assert PrismaDBExceptionHandler.is_deadlock_error(error) is True
@pytest.mark.parametrize(
"error",
[
UniqueViolationError(data={"user_facing_error": {"error_code": "P2002", "meta": {"table": "t"}}}),
RecordNotFoundError(data={"user_facing_error": {"meta": {"table": "t"}}}),
PrismaError("validation failed on query"),
PrismaError("can't reach database server"),
httpx.ConnectError("connection refused"),
RuntimeError("deadlock detected"),
ValueError("40P01"),
],
)
def test_is_deadlock_error_excludes_non_deadlocks(error):
"""Non-deadlock prisma errors, connectivity failures, and non-prisma exceptions are not treated as deadlocks."""
assert PrismaDBExceptionHandler.is_deadlock_error(error) is False

View file

@ -0,0 +1,94 @@
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.db.proxy_worker_heartbeat import (
BEAT_SQL,
COUNT_SQL,
DEREGISTER_SQL,
PROXY_WORKER_LIVENESS_WINDOW_SECONDS,
PRUNE_SQL,
STALE_ROW_RETENTION_SECONDS,
ProxyWorkerHeartbeat,
count_live_proxy_workers,
)
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
def _prisma():
prisma = MagicMock()
prisma.db.execute_raw = AsyncMock()
prisma.db.query_raw = AsyncMock()
return prisma
@pytest.mark.asyncio
async def test_beat_upserts_own_row_then_prunes_stale_rows():
prisma = _prisma()
heartbeat = ProxyWorkerHeartbeat(prisma_client=prisma, worker_id="worker-1")
await heartbeat.beat()
calls = prisma.db.execute_raw.call_args_list
assert calls[0].args == (BEAT_SQL, "worker-1", heartbeat.hostname)
assert calls[1].args == (PRUNE_SQL, STALE_ROW_RETENTION_SECONDS)
@pytest.mark.asyncio
async def test_beat_survives_a_database_error():
prisma = _prisma()
prisma.db.execute_raw = AsyncMock(side_effect=RuntimeError("db down"))
await ProxyWorkerHeartbeat(prisma_client=prisma).beat()
def test_each_worker_process_gets_its_own_id():
prisma = _prisma()
first = ProxyWorkerHeartbeat(prisma_client=prisma)
second = ProxyWorkerHeartbeat(prisma_client=prisma)
assert first.worker_id != second.worker_id
@pytest.mark.asyncio
async def test_deregister_deletes_only_its_own_row():
prisma = _prisma()
await ProxyWorkerHeartbeat(prisma_client=prisma, worker_id="worker-1").deregister()
assert prisma.db.execute_raw.call_args.args == (DEREGISTER_SQL, "worker-1")
@pytest.mark.asyncio
async def test_deregister_survives_a_database_error():
prisma = _prisma()
prisma.db.execute_raw = AsyncMock(side_effect=RuntimeError("db down"))
await ProxyWorkerHeartbeat(prisma_client=prisma, worker_id="worker-1").deregister()
@pytest.mark.asyncio
async def test_count_reads_workers_within_the_liveness_window():
prisma = _prisma()
prisma.db.query_raw.return_value = [{"live_workers": 3}]
assert await count_live_proxy_workers(prisma) == 3
assert prisma.db.query_raw.call_args.args == (COUNT_SQL, PROXY_WORKER_LIVENESS_WINDOW_SECONDS)
@pytest.mark.asyncio
async def test_count_reads_from_the_primary_when_reads_route_to_a_replica():
writer = MagicMock()
writer.query_raw = AsyncMock(return_value=[{"live_workers": 2}])
reader = MagicMock()
reader.query_raw = AsyncMock(return_value=[{"live_workers": 1}])
prisma = MagicMock()
prisma.db = RoutingPrismaWrapper(writer=writer, reader=reader)
assert await count_live_proxy_workers(prisma) == 2
reader.query_raw.assert_not_awaited()
@pytest.mark.asyncio
async def test_count_returns_unknown_when_the_query_fails():
prisma = _prisma()
prisma.db.query_raw.side_effect = RuntimeError("db down")
assert await count_live_proxy_workers(prisma) is None
@pytest.mark.asyncio
async def test_count_returns_unknown_for_a_malformed_row():
prisma = _prisma()
prisma.db.query_raw.return_value = [{"unexpected": "shape"}]
assert await count_live_proxy_workers(prisma) is None

View file

@ -2467,61 +2467,140 @@ class TestNoRedisWarning:
def _router(redis_cache):
return SimpleNamespace(cache=SimpleNamespace(redis_cache=redis_cache))
def test_warns_when_no_redis_is_configured(self, monkeypatch):
@staticmethod
def _prisma_with_workers(live_workers=None, error=None):
prisma = MagicMock()
if error is not None:
prisma.db.query_raw = AsyncMock(side_effect=error)
else:
prisma.db.query_raw = AsyncMock(return_value=[{"live_workers": live_workers}])
return prisma
@pytest.mark.asyncio
async def test_warns_when_no_redis_and_no_db_to_count_workers(self, monkeypatch):
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
assert _show_no_redis_warning() is True
assert await _show_no_redis_warning() is True
def test_warns_when_there_is_no_router_at_all(self, monkeypatch):
@pytest.mark.asyncio
async def test_warns_when_there_is_no_router_at_all(self, monkeypatch):
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", None),
patch("litellm.proxy.proxy_server.prisma_client", None),
):
assert _show_no_redis_warning() is True
assert await _show_no_redis_warning() is True
def test_stays_quiet_when_a_coordination_redis_is_configured(self, monkeypatch):
@pytest.mark.asyncio
async def test_stays_quiet_for_a_confirmed_single_worker(self, monkeypatch):
"""One live worker needs no cross-worker coordination, so no env var is needed."""
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(1)),
):
assert await _show_no_redis_warning() is False
@pytest.mark.asyncio
@pytest.mark.parametrize("live_workers", [2, 5])
async def test_warns_when_multiple_workers_share_the_db(self, monkeypatch, live_workers):
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(live_workers)),
):
assert await _show_no_redis_warning() is True
@pytest.mark.asyncio
async def test_warns_when_the_worker_census_is_empty(self, monkeypatch):
"""Zero rows means the census cannot CONFIRM a single worker, so warn."""
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(0)),
):
assert await _show_no_redis_warning() is True
@pytest.mark.asyncio
async def test_warns_when_the_worker_census_query_fails(self, monkeypatch):
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch(
"litellm.proxy.proxy_server.prisma_client",
self._prisma_with_workers(error=RuntimeError("db down")),
),
):
assert await _show_no_redis_warning() is True
@pytest.mark.asyncio
async def test_stays_quiet_when_a_coordination_redis_is_configured(self, monkeypatch):
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
prisma = self._prisma_with_workers(5)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", MagicMock()),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", prisma),
):
assert _show_no_redis_warning() is False
assert await _show_no_redis_warning() is False
prisma.db.query_raw.assert_not_called()
def test_stays_quiet_when_only_the_router_has_redis(self, monkeypatch):
@pytest.mark.asyncio
async def test_stays_quiet_when_only_the_router_has_redis(self, monkeypatch):
"""router_settings.redis_host alone backs cooldowns and usage-based routing."""
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(MagicMock())),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(5)),
):
assert _show_no_redis_warning() is False
assert await _show_no_redis_warning() is False
@pytest.mark.asyncio
@pytest.mark.parametrize("value", ["true", "True"])
def test_env_var_suppresses_the_warning(self, monkeypatch, value):
async def test_env_var_suppresses_the_warning_despite_multiple_workers(self, monkeypatch, value):
monkeypatch.setenv("LITELLM_DISABLE_NO_REDIS_WARNING", value)
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(5)),
):
assert _show_no_redis_warning() is False
assert await _show_no_redis_warning() is False
def test_env_var_set_false_keeps_the_warning(self, monkeypatch):
@pytest.mark.asyncio
async def test_env_var_set_false_keeps_the_warning_for_multiple_workers(self, monkeypatch):
monkeypatch.setenv("LITELLM_DISABLE_NO_REDIS_WARNING", "false")
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(2)),
):
assert _show_no_redis_warning() is True
assert await _show_no_redis_warning() is True
@pytest.mark.asyncio
async def test_env_var_set_false_does_not_force_the_warning_for_a_single_worker(self, monkeypatch):
monkeypatch.setenv("LITELLM_DISABLE_NO_REDIS_WARNING", "false")
with (
patch("litellm.proxy.proxy_server.redis_usage_cache", None),
patch("litellm.proxy.proxy_server.llm_router", self._router(None)),
patch("litellm.proxy.proxy_server.prisma_client", self._prisma_with_workers(1)),
):
assert await _show_no_redis_warning() is False
@pytest.mark.asyncio
@pytest.mark.parametrize("has_prisma_client", [True, False])
async def test_readiness_details_carries_the_flag(self, monkeypatch, has_prisma_client):
monkeypatch.delenv("LITELLM_DISABLE_NO_REDIS_WARNING", raising=False)
prisma_client = MagicMock() if has_prisma_client else None
prisma_client = self._prisma_with_workers(2) if has_prisma_client else None
with (
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
patch("litellm.proxy.proxy_server.redis_usage_cache", None),

View file

@ -1739,18 +1739,18 @@ def _make_batch_input_bytes(n_rows: int, padding: int = 200) -> bytes:
return ("\n".join(rows)).encode("utf-8")
def test_iter_batch_input_entries_matches_dict_list():
def test_iter_batch_output_entries_matches_dict_list():
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
_iter_batch_input_entries,
_iter_batch_output_entries,
)
raw = _make_batch_input_bytes(50)
streamed = list(_iter_batch_input_entries(raw))
streamed = list(_iter_batch_output_entries(raw))
assert streamed == _get_file_content_as_dictionary(raw)
assert streamed[0]["custom_id"] == "request-0"
# tolerant of blank lines and a missing trailing newline
assert list(_iter_batch_input_entries(raw + b"\n\n")) == streamed
assert list(_iter_batch_output_entries(raw + b"\n\n")) == streamed
def test_streaming_count_peak_below_dict_list():
@ -1759,7 +1759,7 @@ def test_streaming_count_peak_below_dict_list():
from litellm.batches.batch_utils import (
_get_file_content_as_dictionary,
_iter_batch_input_entries,
_iter_batch_output_entries,
)
raw = _make_batch_input_bytes(8000)
@ -1777,7 +1777,7 @@ def test_streaming_count_peak_below_dict_list():
def _stream():
count = 0
models: set = set()
for entry in _iter_batch_input_entries(raw):
for entry in _iter_batch_output_entries(raw):
count += 1
model = (entry.get("body") or {}).get("model")
if model:

View file

@ -0,0 +1,176 @@
import io
import pytest
from litellm.proxy._types import ProxyException
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
BATCH_LINE_REQUIRED_KEYS,
BatchFileEmpty,
BatchFileInvalidJsonLine,
BatchFileLineNotObject,
BatchFileMissingLineKey,
BatchFileTooLarge,
BatchFileWrongExtension,
check_batch_file_upload,
raise_batch_file_validation_failure,
)
VALID_LINE = (
b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions",'
b' "body": {"model": "gpt-4.1-nano", "messages": [{"role": "user", "content": "hi"}]}}'
)
def test_valid_bytes_pass():
assert check_batch_file_upload("batch.jsonl", VALID_LINE + b"\n" + VALID_LINE + b"\n", 10) is None
def test_valid_binaryio_passes_and_resets_position():
handle = io.BytesIO(VALID_LINE + b"\n" + VALID_LINE + b"\n")
handle.seek(17)
assert check_batch_file_upload("batch.jsonl", handle, 10) is None
assert handle.tell() == 0
def test_uppercase_extension_accepted():
assert check_batch_file_upload("BATCH.JSONL", VALID_LINE, None) is None
@pytest.mark.parametrize("filename", ["batch.csv", "batch.json", "batch", None])
def test_wrong_extension_rejected(filename):
assert check_batch_file_upload(filename, VALID_LINE, None) == BatchFileWrongExtension(filename=filename or "")
def test_size_over_cap_rejected_for_bytes():
content = b"x" * (2 * 1024 * 1024)
assert check_batch_file_upload("batch.jsonl", content, 1) == BatchFileTooLarge(
size_bytes=len(content), limit_mb=1
)
def test_size_over_cap_rejected_for_binaryio():
content = b"x" * (2 * 1024 * 1024)
assert check_batch_file_upload("batch.jsonl", io.BytesIO(content), 1) == BatchFileTooLarge(
size_bytes=len(content), limit_mb=1
)
def test_size_exactly_at_cap_allowed():
line = VALID_LINE + b"\n"
padding_key = b'{"custom_id": "pad", "method": "POST", "url": "/v1/chat/completions", "body": {"note": "'
pad_line = padding_key + b"a" * (1024 * 1024 - len(line) - len(padding_key) - len(b'"}}\n')) + b'"}}\n'
content = line + pad_line
assert len(content) == 1024 * 1024
assert check_batch_file_upload("batch.jsonl", content, 1) is None
def test_no_cap_skips_size_check():
content = (VALID_LINE + b"\n") * 5000
assert check_batch_file_upload("batch.jsonl", content, None) is None
@pytest.mark.parametrize("cap", [0, -3])
def test_nonpositive_cap_disables_size_check(cap):
content = (VALID_LINE + b"\n") * 5000
assert check_batch_file_upload("batch.jsonl", content, cap) is None
@pytest.mark.parametrize("content", [b"", b"\n\n", b" \n\t\n"])
def test_empty_file_rejected(content):
assert check_batch_file_upload("batch.jsonl", content, None) == BatchFileEmpty()
def test_invalid_json_line_rejected_with_line_number():
content = VALID_LINE + b"\n" + b"not json at all\n" + VALID_LINE + b"\n"
assert check_batch_file_upload("batch.jsonl", content, None) == BatchFileInvalidJsonLine(line_number=2)
def test_non_utf8_line_rejected_as_invalid_json():
assert check_batch_file_upload("batch.jsonl", b"\xff\xfe\x00\x01\n", None) == BatchFileInvalidJsonLine(
line_number=1
)
def test_non_object_line_rejected():
content = VALID_LINE + b"\n" + b'["custom_id", "method"]\n'
assert check_batch_file_upload("batch.jsonl", content, None) == BatchFileLineNotObject(line_number=2)
@pytest.mark.parametrize("missing_key", BATCH_LINE_REQUIRED_KEYS)
def test_missing_required_key_rejected(missing_key):
import json
line_dict = {
"custom_id": "req-1",
"method": "POST",
"url": "/v1/chat/completions",
"body": {"model": "gpt-4.1-nano"},
}
del line_dict[missing_key]
content = VALID_LINE + b"\n" + json.dumps(line_dict).encode() + b"\n"
assert check_batch_file_upload("batch.jsonl", content, None) == BatchFileMissingLineKey(
line_number=2, key=missing_key
)
def test_blank_lines_do_not_shift_line_numbers():
content = b"\n" + VALID_LINE + b"\n\n" + b"broken\n"
assert check_batch_file_upload("batch.jsonl", content, None) == BatchFileInvalidJsonLine(line_number=4)
def test_failed_scan_leaves_handle_open_and_reset():
handle = io.BytesIO(b"not json\n" + VALID_LINE + b"\n")
assert check_batch_file_upload("batch.jsonl", handle, None) == BatchFileInvalidJsonLine(line_number=1)
assert not handle.closed
assert handle.tell() == 0
def test_scan_stops_at_first_failure():
class ExplodingLines(io.BytesIO):
def __init__(self):
super().__init__(b"not json\n" + VALID_LINE + b"\n")
self.lines_read = 0
def __next__(self):
self.lines_read += 1
return super().__next__()
handle = ExplodingLines()
assert check_batch_file_upload("batch.jsonl", handle, None) == BatchFileInvalidJsonLine(line_number=1)
assert handle.lines_read == 1
@pytest.mark.parametrize(
"failure, expected_code, expected_param, expected_fragments",
[
(
BatchFileTooLarge(size_bytes=220200960, limit_mb=10),
"413",
"file",
("210.0 MB", "max_batch_file_size_mb", "10 MB", "not forwarded"),
),
(
BatchFileWrongExtension(filename="batch.csv"),
"400",
"file",
("batch.csv", ".jsonl", "not forwarded"),
),
(BatchFileEmpty(), "400", "file", ("no request lines", "not forwarded")),
(BatchFileInvalidJsonLine(line_number=3), "400", "file", ("line 3", "not valid JSON")),
(BatchFileLineNotObject(line_number=2), "400", "file", ("line 2", "JSON object")),
(
BatchFileMissingLineKey(line_number=5, key="method"),
"400",
"method",
("'method'", "line 5", "custom_id, method, url, body"),
),
],
)
def test_failures_map_to_openai_shaped_proxy_exceptions(failure, expected_code, expected_param, expected_fragments):
with pytest.raises(ProxyException) as exc_info:
raise_batch_file_validation_failure(failure)
assert exc_info.value.code == expected_code
assert exc_info.value.type == "invalid_request_error"
assert exc_info.value.param == expected_param
for fragment in expected_fragments:
assert fragment in exc_info.value.message

View file

@ -31,6 +31,11 @@ from litellm.caching.caching import DualCache
from litellm.proxy.proxy_server import hash_token
from litellm.proxy.utils import ProxyLogging
VALID_BATCH_LINE = (
b'{"custom_id": "req-1", "method": "POST", "url": "/v1/chat/completions",'
b' "body": {"model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "hi"}]}}\n'
)
@pytest.fixture
def llm_router() -> Router:
@ -225,7 +230,10 @@ def test_invalid_purpose(mocker: MockerFixture, monkeypatch, llm_router: Router)
assert response.status_code == 400
print(f"response: {response.json()}")
assert "Invalid purpose: my-bad-purpose" in response.json()["error"]["message"]
error = response.json()["error"]
assert "Invalid purpose: my-bad-purpose" in error["message"]
assert error["type"] == "invalid_request_error"
assert error["param"] == "purpose"
def test_get_file_content_rejects_raw_cloud_storage_uri(llm_router: Router):
@ -1599,7 +1607,7 @@ def _post_file_with_team_metadata(
user_key = UserAPIKeyAuth(api_key="test-key", team_metadata=team_metadata)
app.dependency_overrides[user_api_key_auth] = lambda: user_key
test_file = ("mydata.jsonl", b'{"prompt": "Hello"}', "application/json")
test_file = ("mydata.jsonl", VALID_BATCH_LINE, "application/jsonl")
try:
response = client.post(
"/v1/files",
@ -1703,7 +1711,7 @@ def _post_file_raw(
user_key = UserAPIKeyAuth(api_key="test-key", team_metadata=team_metadata)
app.dependency_overrides[user_api_key_auth] = lambda: user_key
test_file = ("mydata.jsonl", b'{"prompt": "Hello"}', "application/json")
test_file = ("mydata.jsonl", VALID_BATCH_LINE, "application/jsonl")
try:
response = client.post(
"/v1/files",
@ -2749,7 +2757,7 @@ def test_create_file_provider_only_resolves_named_vertex_credentials(
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", b"{}", "application/jsonl")},
files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")},
data={"purpose": "batch"},
headers={
"Authorization": "Bearer test-key",
@ -2991,7 +2999,7 @@ def test_create_file_provider_only_skips_other_team_vertex_deployment(
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", b"{}", "application/jsonl")},
files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")},
data={"purpose": "batch"},
headers={
"Authorization": "Bearer test-key",
@ -3341,3 +3349,170 @@ def test_raw_provider_file_id_retrieve_allowed_when_managed_files_not_required(
assert response.status_code == 200, response.text
mock_retrieve.assert_called_once()
def _setup_batch_upload_endpoint(monkeypatch, llm_router: Router) -> list:
import litellm.proxy.proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.openai_files_endpoints import files_endpoints as fe
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", llm_router)
setup_proxy_logging_object(monkeypatch, llm_router)
forwarded_calls: list = []
async def fake_route_create_file(**kwargs):
forwarded_calls.append(kwargs)
return OpenAIFileObject(
id="dummy-id",
object="file",
bytes=0,
created_at=1234567890,
filename="batch.jsonl",
purpose="batch",
status="uploaded",
)
monkeypatch.setattr(fe, "route_create_file", fake_route_create_file)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test-user"
)
return forwarded_calls
def _teardown_batch_upload_endpoint():
import litellm.proxy.proxy_server as ps
app.dependency_overrides.pop(ps.user_api_key_auth, None)
def test_create_file_batch_over_max_batch_file_size_mb_rejected_before_forwarding(
monkeypatch, llm_router: Router
):
import litellm.proxy.proxy_server as ps
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
monkeypatch.setitem(ps.general_settings, "max_batch_file_size_mb", 1)
oversized = VALID_BATCH_LINE * (2 * 1024 * 1024 // len(VALID_BATCH_LINE) + 1)
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", oversized, "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
finally:
_teardown_batch_upload_endpoint()
assert response.status_code == 413, response.text
error = response.json()["error"]
assert error["type"] == "invalid_request_error"
assert error["param"] == "file"
assert "max_batch_file_size_mb" in error["message"]
assert "1 MB" in error["message"]
assert forwarded_calls == []
def test_create_file_batch_under_max_batch_file_size_mb_forwards(monkeypatch, llm_router: Router):
import litellm.proxy.proxy_server as ps
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
monkeypatch.setitem(ps.general_settings, "max_batch_file_size_mb", 1)
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", VALID_BATCH_LINE, "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
finally:
_teardown_batch_upload_endpoint()
assert response.status_code == 200, response.text
assert len(forwarded_calls) == 1
def test_create_file_batch_wrong_extension_rejected_before_forwarding(monkeypatch, llm_router: Router):
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
try:
response = client.post(
"/v1/files",
files={"file": ("batch.csv", VALID_BATCH_LINE, "text/csv")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
finally:
_teardown_batch_upload_endpoint()
assert response.status_code == 400, response.text
error = response.json()["error"]
assert error["type"] == "invalid_request_error"
assert error["param"] == "file"
assert "batch.csv" in error["message"]
assert ".jsonl" in error["message"]
assert forwarded_calls == []
def test_create_file_batch_missing_line_key_rejected_before_forwarding(monkeypatch, llm_router: Router):
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
bad_line = b'{"custom_id": "req-1", "url": "/v1/chat/completions", "body": {"model": "gpt-3.5-turbo"}}\n'
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", VALID_BATCH_LINE + bad_line, "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
finally:
_teardown_batch_upload_endpoint()
assert response.status_code == 400, response.text
error = response.json()["error"]
assert error["type"] == "invalid_request_error"
assert error["param"] == "method"
assert "line 2" in error["message"]
assert forwarded_calls == []
def test_create_file_batch_invalid_json_line_rejected_before_forwarding(monkeypatch, llm_router: Router):
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
try:
response = client.post(
"/v1/files",
files={"file": ("batch.jsonl", b"this is not jsonl\n", "application/jsonl")},
data={"purpose": "batch"},
headers={"Authorization": "Bearer test-key"},
)
finally:
_teardown_batch_upload_endpoint()
assert response.status_code == 400, response.text
error = response.json()["error"]
assert error["param"] == "file"
assert "line 1" in error["message"]
assert "not valid JSON" in error["message"]
assert forwarded_calls == []
def test_create_file_non_batch_purpose_skips_batch_validation(monkeypatch, llm_router: Router):
forwarded_calls = _setup_batch_upload_endpoint(monkeypatch, llm_router)
try:
response = client.post(
"/v1/files",
files={"file": ("notes.txt", b"plain text, not jsonl", "text/plain")},
data={"purpose": "user_data"},
headers={"Authorization": "Bearer test-key"},
)
finally:
_teardown_batch_upload_endpoint()
assert response.status_code == 200, response.text
assert len(forwarded_calls) == 1

View file

@ -719,6 +719,7 @@ class TestVertexAIPassThroughHandler:
# Create mock logging object
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.optional_params = {}
mock_logging_obj.litellm_call_id = "test-call-id-123"
mock_logging_obj.model_call_details = {}
@ -895,6 +896,7 @@ class TestVertexAIPassThroughHandler:
# Create mock logging object
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.optional_params = {}
mock_logging_obj.litellm_call_id = "test-call-id-123"
mock_logging_obj.model_call_details = {}
@ -965,6 +967,7 @@ class TestVertexAIPassThroughHandler:
mock_httpx_response.status_code = 200
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.optional_params = {}
mock_logging_obj.litellm_call_id = "test-call-id-embed"
mock_logging_obj.model_call_details = {}
@ -1023,6 +1026,7 @@ class TestVertexAIPassThroughHandler:
mock_httpx_response.status_code = 200
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.optional_params = {}
mock_logging_obj.litellm_call_id = "test-call-id-batch"
mock_logging_obj.model_call_details = {}
@ -1079,6 +1083,7 @@ class TestVertexAIPassThroughHandler:
mock_httpx_response.status_code = 200
mock_logging_obj = Mock(spec=LiteLLMLoggingObj)
mock_logging_obj.optional_params = {}
mock_logging_obj.litellm_call_id = "test-call-id-gemini-studio"
mock_logging_obj.model_call_details = {}
@ -1109,6 +1114,117 @@ class TestVertexAIPassThroughHandler:
assert result["kwargs"].get("model") == "gemini-embedding-2-preview"
mock_completion_cost.assert_called_once()
@pytest.mark.parametrize("streaming", [False, True])
def test_vertex_passthrough_handler_prices_regional_endpoint_with_uplift(self, monkeypatch, streaming):
"""
Both cost computations for a passthrough call must price on the URL's serving location:
the handler-computed cost, and the async success recompute, which re-resolves the
location from the logging object and previously fell through empty optional_params to
the us-central1 default, billing the regional uplift on global traffic too (#34393).
"""
import datetime
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import (
VertexPassthroughLoggingHandler,
)
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(
litellm,
"model_cost",
{
**litellm.get_model_cost_map(url=""),
"vertex_ai/gemini-fake-regional": {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
"regional_endpoint_uplift_multiplier": 1.1,
},
},
)
response_body: Final = {
"candidates": [
{
"content": {"parts": [{"text": "hello"}], "role": "model"},
"finishReason": "STOP",
}
],
"usageMetadata": {
"promptTokenCount": 10,
"candidatesTokenCount": 20,
"totalTokenCount": 30,
},
}
def costs_for(location: str) -> tuple[float, float]:
url_route: Final = (
f"https://{location}-aiplatform.googleapis.com/v1/projects/p/locations/{location}"
"/publishers/google/models/gemini-fake-regional:"
f"{'streamGenerateContent' if streaming else 'generateContent'}"
)
start_time: Final = datetime.datetime.now()
end_time: Final = datetime.datetime.now()
logging_obj: Final = Logging(
model="gemini-fake-regional",
messages=[{"role": "user", "content": "hi"}],
stream=streaming,
call_type="pass_through_endpoint",
start_time=start_time,
litellm_call_id="call-id",
function_id="fn-id",
)
logging_obj.update_environment_variables(
model="gemini-fake-regional",
user="unknown",
optional_params={},
litellm_params={},
call_type="pass_through_endpoint",
)
if streaming:
result = VertexPassthroughLoggingHandler._handle_logging_vertex_collected_chunks(
litellm_logging_obj=logging_obj,
passthrough_success_handler_obj=Mock(),
url_route=url_route,
request_body={},
endpoint_type="vertex_ai",
start_time=start_time,
all_chunks=[json.dumps(response_body)],
model=None,
end_time=end_time,
)
else:
mock_httpx_response: Final = Mock()
mock_httpx_response.json.return_value = response_body
mock_httpx_response.headers = {}
mock_httpx_response.status_code = 200
result = VertexPassthroughLoggingHandler.vertex_passthrough_handler(
httpx_response=mock_httpx_response,
logging_obj=logging_obj,
url_route=url_route,
result="test-result",
start_time=start_time,
end_time=end_time,
cache_hit=False,
)
recomputed: Final = logging_obj._response_cost_calculator(result=result["result"])
return result["kwargs"]["response_cost"], recomputed
global_handler_cost, global_recomputed_cost = costs_for("global")
regional_handler_cost, regional_recomputed_cost = costs_for("us-east5")
plain_cost: Final = 10 * 1e-06 + 20 * 2e-06
assert global_handler_cost == pytest.approx(plain_cost, rel=1e-9)
assert regional_handler_cost == pytest.approx(plain_cost * 1.10, rel=1e-9), (
"regional Vertex passthrough traffic must bill at 1.1x the global rate"
)
assert global_recomputed_cost == pytest.approx(plain_cost, rel=1e-9), (
"the logging recompute must not price global passthrough traffic as regional"
)
assert regional_recomputed_cost == pytest.approx(plain_cost * 1.10, rel=1e-9)
class TestVertexAIDiscoveryPassThroughHandler:
"""

View file

@ -769,6 +769,214 @@ async def test_ProxyConfig_get_config_missing_file_raises(monkeypatch):
await pc.get_config(config_file_path="/no/such/path.yaml")
# ---------------------------------------------------------------------------
# ProxyConfig._initialize_secret_manager_from_raw_config
# ---------------------------------------------------------------------------
VAULT_SECRET_MANAGER_MODULE = '''
import os
from litellm.integrations.custom_secret_manager import CustomSecretManager
VAULT = {"LITELLM_MASTER_KEY": "master-from-vault", "MY_PROVIDER_KEY": "provider-from-vault"}
class VaultSecretManager(CustomSecretManager):
def __init__(self):
super().__init__()
# The loader re-executes this module on every construction, so an in-module counter
# would reset. Append to a file instead, to count constructions across the whole load.
with open(os.environ["VAULT_CONSTRUCTION_LOG"], "a") as f:
f.write("constructed\\n")
def sync_read_secret(self, secret_name, optional_params=None, timeout=None, **kwargs):
return VAULT.get(secret_name)
async def async_read_secret(self, secret_name, optional_params=None, timeout=None, **kwargs):
return VAULT.get(secret_name)
'''
VAULT_BACKED_CONFIG = """
model_list:
- model_name: my-model
litellm_params:
model: openai/gpt-4o-mini
api_key: os.environ/MY_PROVIDER_KEY
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
key_management_system: custom
key_management_settings:
custom_secret_manager: vault_secret_manager.VaultSecretManager
hosted_keys:
- LITELLM_MASTER_KEY
- MY_PROVIDER_KEY
"""
def _write_vault_backed_config(tmp_path, monkeypatch, config_yaml: str) -> str:
"""Write a config whose secrets live only in a custom secret manager, never in the env."""
(tmp_path / "vault_secret_manager.py").write_text(VAULT_SECRET_MANAGER_MODULE)
config_file = tmp_path / "c.yaml"
config_file.write_text(config_yaml)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
monkeypatch.delenv("LITELLM_MASTER_KEY", raising=False)
monkeypatch.delenv("MY_PROVIDER_KEY", raising=False)
monkeypatch.setenv("VAULT_CONSTRUCTION_LOG", str(tmp_path / "constructions.log"))
monkeypatch.setattr(litellm, "secret_manager_client", None)
return str(config_file)
def _construction_count(tmp_path) -> int:
log = tmp_path / "constructions.log"
return len(log.read_text().splitlines()) if log.exists() else 0
@pytest.mark.asyncio
async def test_ProxyConfig_get_config_resolves_keys_held_only_by_the_secret_manager(tmp_path, monkeypatch):
"""Regression for GH #35239.
get_config() used to resolve every ``os.environ/<KEY>`` reference and write the result
back into the config before the secret manager was initialized, so any key that lived
only in the manager became a permanent ``None``.
"""
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, VAULT_BACKED_CONFIG)
cfg = await ProxyConfig().get_config(config_file_path=config_file_path)
assert {
"master_key": cfg["general_settings"]["master_key"],
"api_key": cfg["model_list"][0]["litellm_params"]["api_key"],
"hosted_keys": litellm._key_management_settings.hosted_keys,
} == {
"master_key": "master-from-vault",
"api_key": "provider-from-vault",
"hosted_keys": ["LITELLM_MASTER_KEY", "MY_PROVIDER_KEY"],
}
@pytest.mark.asyncio
async def test_ProxyConfig_load_config_builds_the_secret_manager_exactly_once(tmp_path, monkeypatch):
"""The full startup path must not build the manager, then throw it away and build another.
A discarded client costs a Vault/CyberArk re-auth and leaks a gRPC channel on Google KMS.
"""
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, VAULT_BACKED_CONFIG)
_router, _model_list, general_settings = await ProxyConfig().load_config(
router=None, config_file_path=config_file_path
)
assert {
"constructions": _construction_count(tmp_path),
"master_key": general_settings["master_key"],
} == {"constructions": 1, "master_key": "master-from-vault"}
@pytest.mark.asyncio
async def test_ProxyConfig_get_config_reuses_an_already_initialized_secret_manager(tmp_path, monkeypatch):
"""get_config() also runs on management-endpoint request paths.
Rebuilding the client on every call would re-execute the custom manager module, drop the
Vault/CyberArk token caches, and leak a gRPC channel per request on Google KMS.
"""
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, VAULT_BACKED_CONFIG)
await ProxyConfig().get_config(config_file_path=config_file_path)
first_client = litellm.secret_manager_client
second = await ProxyConfig().get_config(config_file_path=config_file_path)
assert {
"client_reused": litellm.secret_manager_client is first_client,
"master_key": second["general_settings"]["master_key"],
} == {"client_reused": True, "master_key": "master-from-vault"}
@pytest.mark.asyncio
async def test_ProxyConfig_get_config_without_key_management_system_leaves_secret_manager_unset(
tmp_path, monkeypatch
):
"""No ``key_management_system`` means no manager, an unresolvable reference stays None, and
nothing is warned about: with no manager there is nothing to have been absent from."""
config_yaml = VAULT_BACKED_CONFIG.replace(" key_management_system: custom\n", "")
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, config_yaml)
warn = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.verbose_proxy_logger.warning", warn)
cfg = await ProxyConfig().get_config(config_file_path=config_file_path)
assert {
"master_key": cfg["general_settings"]["master_key"],
"api_key": cfg["model_list"][0]["litellm_params"]["api_key"],
"client": litellm.secret_manager_client,
"warned_about": [call.args[1] for call in warn.call_args_list],
} == {"master_key": None, "api_key": None, "client": None, "warned_about": []}
@pytest.mark.asyncio
async def test_ProxyConfig_get_config_warns_when_a_reference_is_missing_from_the_secret_manager(
tmp_path, monkeypatch
):
"""A reference the manager cannot resolve is logged, instead of silently becoming None."""
config_yaml = VAULT_BACKED_CONFIG.replace("MY_PROVIDER_KEY", "NOT_IN_VAULT")
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, config_yaml)
warn = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.verbose_proxy_logger.warning", warn)
cfg = await ProxyConfig().get_config(config_file_path=config_file_path)
assert {
"api_key": cfg["model_list"][0]["litellm_params"]["api_key"],
"warned_about": [call.args[1] for call in warn.call_args_list],
} == {"api_key": None, "warned_about": ["os.environ/NOT_IN_VAULT"]}
@pytest.mark.asyncio
async def test_ProxyConfig_get_config_does_not_warn_for_a_name_outside_hosted_keys(tmp_path, monkeypatch):
"""``hosted_keys`` is an allowlist, so a name outside it is never looked up in the manager.
Warning about it would claim a lookup that never happened, on every optional env-only
reference, on every config reload.
"""
config_yaml = VAULT_BACKED_CONFIG.replace("api_key: os.environ/MY_PROVIDER_KEY", "api_key: os.environ/ENV_ONLY")
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, config_yaml)
warn = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.verbose_proxy_logger.warning", warn)
cfg = await ProxyConfig().get_config(config_file_path=config_file_path)
assert {
"api_key": cfg["model_list"][0]["litellm_params"]["api_key"],
"client_is_up": litellm.secret_manager_client is not None,
"warned_about": [call.args[1] for call in warn.call_args_list],
} == {"api_key": None, "client_is_up": True, "warned_about": []}
@pytest.mark.asyncio
async def test_ProxyConfig_get_config_does_not_warn_under_write_only_access_mode(tmp_path, monkeypatch):
"""``write_only`` means reads never reach the manager, so an absent name is not its fault.
That mode exists so the manager can store virtual keys while config secrets stay in the
environment, which makes env-only references the expected state rather than an error.
"""
config_yaml = VAULT_BACKED_CONFIG.replace(
" key_management_settings:\n", " key_management_settings:\n access_mode: write_only\n"
)
config_file_path = _write_vault_backed_config(tmp_path, monkeypatch, config_yaml)
warn = MagicMock()
monkeypatch.setattr("litellm.proxy.proxy_server.verbose_proxy_logger.warning", warn)
cfg = await ProxyConfig().get_config(config_file_path=config_file_path)
assert {
"master_key": cfg["general_settings"]["master_key"],
"client_is_up": litellm.secret_manager_client is not None,
"warned_about": [call.args[1] for call in warn.call_args_list],
} == {"master_key": None, "client_is_up": True, "warned_about": []}
# ---------------------------------------------------------------------------
# ProxyConfig.update_config_state / get_config_state
# ---------------------------------------------------------------------------
@ -2442,6 +2650,43 @@ async def test_ProxyConfig__update_general_settings_updates_max_parallel(monkeyp
}
@pytest.mark.asyncio
async def test_ProxyConfig__update_general_settings_applies_db_max_batch_file_size_mb(monkeypatch):
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
pc = ProxyConfig()
await pc._update_general_settings({"max_batch_file_size_mb": 5})
from litellm.proxy import proxy_server as ps
assert ps.general_settings.get("max_batch_file_size_mb") == 5
@pytest.mark.asyncio
async def test_ProxyConfig__update_general_settings_yaml_max_batch_file_size_mb_wins_over_db(monkeypatch):
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"max_batch_file_size_mb": 3},
)
pc = ProxyConfig()
pc._yaml_general_settings_keys = {"max_batch_file_size_mb"}
await pc._update_general_settings({"max_batch_file_size_mb": 5})
from litellm.proxy import proxy_server as ps
assert ps.general_settings.get("max_batch_file_size_mb") == 3
@pytest.mark.asyncio
async def test_ProxyConfig__update_general_settings_cleared_db_max_batch_file_size_mb_lifts_cap(monkeypatch):
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings",
{"max_batch_file_size_mb": 8},
)
pc = ProxyConfig()
await pc._update_general_settings({"max_parallel_requests": 1})
from litellm.proxy import proxy_server as ps
assert ps.general_settings.get("max_batch_file_size_mb") is None
@pytest.mark.asyncio
async def test_ProxyConfig__update_general_settings_none_input_noop():
pc = ProxyConfig()

View file

@ -841,6 +841,44 @@ def test_the_baseline_is_priced_on_the_basis_the_request_was_billed_at(basis, ex
assert reported == pytest.approx(expected_multiplier * baseline - served)
def test_the_baseline_is_priced_on_the_vertex_location_the_request_was_billed_at(monkeypatch):
"""A request served from a regional Vertex endpoint was billed with the
regional-endpoint uplift, so the counterfactual single-model operator would
have paid it too. The served model carries no uplift field, so only the
baseline moves with the recorded location."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
gemini = litellm.get_model_info("gemini-3.5-flash", "vertex_ai")
haiku = litellm.get_model_info("claude-haiku-4-5", "anthropic")
assert gemini.get("regional_endpoint_uplift_multiplier") == 1.1
assert haiku.get("regional_endpoint_uplift_multiplier") is None, "served model must not move with the basis"
usage = _usage(fresh=20_000, cached=0, written=0, out=1_000)
served = 20_000 * haiku["input_cost_per_token"] + 1_000 * haiku["output_cost_per_token"]
baseline = 20_000 * gemini["input_cost_per_token"] + 1_000 * gemini["output_cost_per_token"]
regional = compute_autorouter_savings(
baseline_model="vertex_ai/gemini-3.5-flash",
selected_model="claude-haiku-4-5",
selected_provider="anthropic",
usage=usage,
conversation_continuing=False,
cost_breakdown=_breakdown(served, vertex_location="us-east5"),
)
global_endpoint = compute_autorouter_savings(
baseline_model="vertex_ai/gemini-3.5-flash",
selected_model="claude-haiku-4-5",
selected_provider="anthropic",
usage=usage,
conversation_continuing=False,
cost_breakdown=_breakdown(served, vertex_location="global"),
)
assert regional == pytest.approx(1.1 * baseline - served)
assert global_endpoint == pytest.approx(baseline - served)
def test_a_baseline_recorded_on_the_decision_turns_the_driver_on():
"""An operator who configures nothing still sees the driver work."""
result = compute_savings_spend(

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