mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
commit
66dbf2101c
274 changed files with 16615 additions and 9447 deletions
2
.github/CODEOWNERS
vendored
2
.github/CODEOWNERS
vendored
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
3
.github/pull_request_template.md
vendored
3
.github/pull_request_template.md
vendored
|
|
@ -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
41
.github/scripts/detect_backend_changes.sh
vendored
Executable 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
|
||||
3
.github/workflows/_test-unit-base.yml
vendored
3
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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 }}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
1
.gitignore
vendored
|
|
@ -5,6 +5,7 @@ tests/e2e/.fixtures/
|
|||
.venv_policy_test
|
||||
.env
|
||||
.claude
|
||||
CLAUDE.local.md
|
||||
.newenv
|
||||
newenv/*
|
||||
litellm/proxy/myenv/*
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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] = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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. | `[]` |
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
|
|
@ -0,0 +1,4 @@
|
|||
UPDATE "LiteLLM_SpendLogs"
|
||||
SET "created_at" = "endTime",
|
||||
"updated_at" = "endTime"
|
||||
WHERE "created_at" > "endTime" + interval '1 hour';
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
0
litellm/llms/bedrock/search/__init__.py
Normal file
0
litellm/llms/bedrock/search/__init__.py
Normal file
455
litellm/llms/bedrock/search/transformation.py
Normal file
455
litellm/llms/bedrock/search/transformation.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.")
|
||||
|
|
|
|||
155
litellm/proxy/client/cli/commands/claude_settings.py
Normal file
155
litellm/proxy/client/cli/commands/claude_settings.py
Normal 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",
|
||||
)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
93
litellm/proxy/db/proxy_worker_heartbeat.py
Normal file
93
litellm/proxy/db/proxy_worker_heartbeat.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
182
litellm/proxy/openai_files_endpoints/batch_file_validation.py
Normal file
182
litellm/proxy/openai_files_endpoints/batch_file_validation.py
Normal 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)
|
||||
|
|
@ -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)),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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``.
|
||||
"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@
|
|||
"limit": 711
|
||||
},
|
||||
"ANN205": {
|
||||
"limit": 113
|
||||
"limit": 112
|
||||
},
|
||||
"ANN206": {
|
||||
"limit": 133
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
277
tests/e2e/access_control/test_model_access_group_e2e.py
Normal file
277
tests/e2e/access_control/test_model_access_group_e2e.py
Normal 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]}"
|
||||
)
|
||||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
150
tests/e2e/fixture_canonical.py
Normal file
150
tests/e2e/fixture_canonical.py
Normal 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=(",", ":")),
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
163
tests/e2e/test_fixture_canonical.py
Normal file
163
tests/e2e/test_fixture_canonical.py
Normal 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
269
tests/test_litellm/proxy/client/cli/test_claude_settings.py
Normal file
269
tests/test_litellm/proxy/client/cli/test_claude_settings.py
Normal 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"
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
94
tests/test_litellm/proxy/db/test_proxy_worker_heartbeat.py
Normal file
94
tests/test_litellm/proxy/db/test_proxy_worker_heartbeat.py
Normal 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
|
||||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue