Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_cli_refresh_tokens

# Conflicts:
#	basedpyright-code-budget.json
This commit is contained in:
mateo-berri 2026-08-19 19:50:26 -07:00
commit 65d834f01d
210 changed files with 6697 additions and 3341 deletions

View file

@ -1,15 +1,17 @@
#!/usr/bin/env bash
set -uo pipefail
category="${1:?usage: classify_changes.sh <backend|client>}"
category="${1:?usage: classify_changes.sh <backend|client|ui>}"
has_client=false
has_backend=false
has_ci=false
while IFS= read -r file || [ -n "$file" ]; do
[ -n "$file" ] || continue
case "$file" in
ui/* | tests/e2e/ui/*) has_client=true ;;
docs/* | *.md | *.mdx) : ;;
.github/* | .circleci/*) has_ci=true; has_backend=true ;;
*) has_backend=true ;;
esac
done
@ -21,6 +23,9 @@ case "$category" in
client)
{ [ "$has_client" = true ] || [ "$has_backend" = true ]; } && echo run || echo skip
;;
ui)
{ [ "$has_client" = true ] || [ "$has_ci" = true ]; } && echo run || echo skip
;;
*)
echo run
;;

2
.github/CODEOWNERS vendored
View file

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

View file

@ -1,34 +0,0 @@
name: "Detect backend-relevant changes"
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 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:
description: "run when backend-relevant files changed, otherwise skip"
value: ${{ steps.classify.outputs.decision }}
runs:
using: composite
steps:
- id: classify
shell: bash
env:
GH_TOKEN: ${{ inputs.github-token }}
REPO: ${{ github.repository }}
PR_NUMBER: ${{ github.event.pull_request.number }}
CHANGED_FILE_COUNT: ${{ github.event.pull_request.changed_files }}
run: bash "${GITHUB_ACTION_PATH}/../../scripts/detect_backend_changes.sh"

View file

@ -0,0 +1,41 @@
name: "Detect relevant changes"
description: >-
Classify the pull request's changed files with .circleci/scripts/classify_changes.sh
and expose decision=run|skip for one category. backend means anything outside ui/,
docs/ and markdown; ui means the dashboard sources alone. decision=skip lets callers
short-circuit expensive steps while the job still completes successfully and satisfies
its required status check, which a paths: filter cannot do because a workflow that
never starts never reports. 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 jobs are never skipped when the classification
is uncertain.
inputs:
category:
description: "Which classification to apply: backend, client or ui"
required: false
default: backend
github-token:
description: "Token used to list the pull request's files; needs pull-requests: read"
required: false
default: ${{ github.token }}
outputs:
decision:
description: "run when category-relevant files changed, otherwise skip"
value: ${{ steps.classify.outputs.decision }}
runs:
using: composite
steps:
- id: classify
shell: bash
env:
GH_TOKEN: ${{ inputs.github-token }}
CATEGORY: ${{ inputs.category }}
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_changes.sh"

View file

@ -2,15 +2,16 @@
set -uo pipefail
readonly API_FILE_CEILING=3000
readonly CATEGORY="${CATEGORY:-backend}"
decide() {
echo "detect-backend-changes: decision=$1"
echo "detect-changes[${CATEGORY}]: decision=$1"
[ -z "${GITHUB_OUTPUT:-}" ] || echo "decision=$1" >>"${GITHUB_OUTPUT}"
exit 0
}
run_full() {
echo "detect-backend-changes: $1; running job"
echo "detect-changes[${CATEGORY}]: $1; running job"
decide run
}
@ -30,10 +31,10 @@ changed="$(gh api "repos/${REPO}/pulls/${PR_NUMBER}/files" --paginate --jq '.[].
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}:"
echo "detect-changes[${CATEGORY}]: files changed by PR #${PR_NUMBER}:"
printf '%s\n' "${changed}" | sed 's/^/ /'
decision="$(printf '%s\n' "${changed}" | bash "${classify}" backend)" ||
decision="$(printf '%s\n' "${changed}" | bash "${classify}" "${CATEGORY}")" ||
run_full "classify_changes.sh failed"
case "${decision}" in
run | skip) decide "${decision}" ;;

View file

@ -72,24 +72,27 @@ jobs:
with:
persist-credentials: false
- name: Detect backend-relevant changes
- name: Detect relevant changes
id: changes
timeout-minutes: 2
uses: ./.github/actions/detect-backend-changes
uses: ./.github/actions/detect-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 3
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
timeout-minutes: 5
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:

View file

@ -24,6 +24,7 @@ jobs:
# re-running basedpyright over the merge-base tree.
permissions:
contents: read
pull-requests: read
actions: read
steps:
@ -37,7 +38,12 @@ jobs:
clean: true
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Fetch gate base (merge-base with target branch)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
BASE_SHA: ${{ github.event.pull_request.base.sha }}
@ -50,39 +56,47 @@ jobs:
echo "GATE_BASE_SHA=$MERGE_BASE" >> "$GITHUB_ENV"
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Clean Python cache
if: steps.changes.outputs.decision != 'skip'
run: |
find . -type d -name "__pycache__" -exec rm -rf {} + || true
find . -name "*.pyc" -delete || true
- name: Check uv.lock is up to date
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check || (echo "❌ uv.lock is out of sync with pyproject.toml. Run 'uv lock' locally and commit the result." && exit 1)
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv sync --frozen --group proxy-dev --group e2e-dev
- name: Cache Prisma binaries
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/cache-prisma-binaries
# basedpyright resolves Prisma's generated client (litellm/proxy/schema.prisma)
# only after `prisma generate` writes prisma/client.py et al. Without this the
# DB wrappers typed against the generated client would degrade to Unknown.
- name: Generate Prisma client
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- name: Check ruff format
if: steps.changes.outputs.decision != 'skip'
run: |
git diff --name-only --diff-filter=ACMR "$GATE_BASE_SHA" HEAD -- 'litellm/**/*.py' | grep -v '^litellm/enterprise/' > "$RUNNER_TEMP/ruff_format_files.txt" || true
if [ ! -s "$RUNNER_TEMP/ruff_format_files.txt" ]; then
@ -92,6 +106,7 @@ jobs:
xargs uv run --no-sync ruff format --check --exclude '/enterprise/' < "$RUNNER_TEMP/ruff_format_files.txt"
- name: Debug - Check file state
if: steps.changes.outputs.decision != 'skip'
run: |
echo "Current branch:"
git branch --show-current
@ -101,30 +116,36 @@ jobs:
head -50 litellm/litellm_core_utils/custom_logger_registry.py | tail -10
- name: Run Ruff linting
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync ruff check .
cd ..
- name: Check strict-rule budget (delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/ruff_strict_gate.py --base "$GATE_BASE_SHA"
- name: Check type-discipline budget (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs base)
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA"
- name: Print OpenAI version
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')"
- name: Check basedpyright budget (delta vs base)
if: steps.changes.outputs.decision != 'skip'
env:
GH_TOKEN: ${{ github.token }}
run: |
uv run --no-sync python scripts/type_check_gate.py --base "$GATE_BASE_SHA"
- name: Check tests/e2e basedpyright (zero errors)
if: steps.changes.outputs.decision != 'skip'
run: |
if git diff --name-only --diff-filter=ACMRD "$GATE_BASE_SHA" HEAD -- 'tests/e2e/**/*.py' | grep -q .; then
uv run --no-sync basedpyright tests/e2e
@ -133,12 +154,14 @@ jobs:
fi
- name: Check for circular imports
if: steps.changes.outputs.decision != 'skip'
run: |
cd litellm
uv run --no-sync python ../tests/documentation_tests/test_circular_imports.py
cd ..
- name: Check import safety
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)

View file

@ -1,6 +1,7 @@
name: UI Build Check
permissions:
contents: read
pull-requests: read
on:
pull_request:
@ -28,7 +29,14 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
@ -36,7 +44,9 @@ jobs:
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Build
if: steps.changes.outputs.decision != 'skip'
run: npm run build

View file

@ -1,6 +1,7 @@
name: UI Unit Tests
permissions:
contents: read
pull-requests: read
on:
pull_request:
@ -32,7 +33,14 @@ jobs:
fetch-depth: 1
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
with:
category: ui
- name: Setup Node.js
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version-file: ui/litellm-dashboard/.nvmrc
@ -40,14 +48,17 @@ jobs:
cache-dependency-path: ui/litellm-dashboard/package-lock.json
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: npm ci
- name: Run UI type tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
run: npm run test:types
- name: Run UI unit tests (Vitest)
if: steps.changes.outputs.decision != 'skip'
env:
CI: "true"
GH_TOKEN: ${{ github.token }}

View file

@ -10,6 +10,7 @@ on:
permissions:
contents: read
pull-requests: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
@ -25,26 +26,34 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Thank You Message
run: |
echo "### 🙏 Thank you for contributing to LiteLLM!" >> $GITHUB_STEP_SUMMARY
echo "Your PR is being tested now. We appreciate your help in making LiteLLM better!" >> $GITHUB_STEP_SUMMARY
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Install dependencies
if: steps.changes.outputs.decision != 'skip'
run: |
uv lock --check
.github/scripts/uv_sync_with_retries.sh --frozen --group proxy-dev --extra proxy --extra semantic-router
- name: Run MCP tests
if: steps.changes.outputs.decision != 'skip'
run: |
uv run --no-sync pytest tests/mcp_tests -x -vv -n 4 --cov=./litellm --cov-report=xml --durations=5

View file

@ -32,28 +32,32 @@ jobs:
with:
persist-credentials: false
- name: Detect relevant changes
id: changes
uses: ./.github/actions/detect-changes
- name: Checkout litellm-docs into docs/my-website (for documentation_tests)
if: steps.changes.outputs.decision != 'skip'
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
repository: BerriAI/litellm-docs
path: docs/my-website
persist-credentials: false
- name: Detect backend-relevant changes
id: changes
uses: ./.github/actions/detect-backend-changes
- name: Set up Python
if: steps.changes.outputs.decision != 'skip'
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
if: steps.changes.outputs.decision != 'skip'
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Cache uv dependencies
if: steps.changes.outputs.decision != 'skip'
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
with:
path: |

View file

@ -57,7 +57,7 @@
"limit": 5663
},
"reportMissingTypeArgument": {
"limit": 15556
"limit": 15555
},
"reportMissingTypeStubs": {
"limit": 40
@ -105,13 +105,13 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 39042
"limit": 39017
},
"reportUnknownParameterType": {
"limit": 19886
"limit": 19885
},
"reportUnknownVariableType": {
"limit": 30571
"limit": 30572
},
"reportUnnecessaryCast": {
"limit": 117

View file

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

View file

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

View file

@ -221,7 +221,7 @@ overwrite_user_with_key_hash: bool = (
bedrock_request_metadata_fields: Optional[Sequence[str]] = (
None # allow-list of `user_api_key_*` fields (+ `spend_logs_metadata`) sent as Bedrock `requestMetadata`
)
store_audit_logs = False # Enterprise feature, allow users to see audit logs
store_audit_logs: bool | None = None
skip_system_message_in_guardrail: bool = False
skip_tool_message_in_guardrail: bool = False
### end of callbacks #############

View file

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

View file

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

View file

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

View file

@ -0,0 +1,149 @@
"""Which deployments accrue PTU flat cost, and what that costs them per token.
Reserved provisioned throughput is billed by the hour whether or not requests are sent, so
a deployment that accrues flat cost must not also bill per token. The two halves live here
together because they have to agree: a deployment the rollup declines to charge but the
router prices at zero serves its traffic for free.
"""
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Final
from litellm.secret_managers.main import get_secret_bool
from litellm.types.router import ModelInfo
from litellm.types.utils import CustomPricingLiteLLMParams, MirroredPricingParams
PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION"
def is_ptu_cost_attribution_enabled() -> bool:
"""Whether PTU flat-cost attribution is turned on for this process."""
return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True
PTU_ZEROED_PRICING_FIELDS: Final = tuple(f for f in MirroredPricingParams.model_fields if f != "tiered_pricing") + (
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost_above_200k_tokens",
)
# tiered_pricing is emptied rather than zeroed: its tiers outrank the zeros written beside
# them, so a zero here would leave the cost map's tiers billing the traffic the reserved
# capacity already covers.
PTU_EMPTIED_PRICING_FIELDS: Final = frozenset(("tiered_pricing",))
# search_context_cost_per_query holds its rates in a table keyed by context size, and an
# absent table means the provider's own default rather than free, so it is zeroed in place
# and written on every PTU deployment rather than only where a table is already stored.
PTU_ZEROED_TABLE_FIELDS: Final = frozenset(("search_context_cost_per_query",))
SEARCH_CONTEXT_SIZES: Final = ("search_context_size_low", "search_context_size_medium", "search_context_size_high")
# Rate fields only. CustomPricingLiteLLMParams also carries settings that are not charges,
# and zeroing one of those would destroy the deployment's configuration rather than stop a
# charge.
CUSTOM_PRICING_FIELDS: Final = frozenset(f for f in CustomPricingLiteLLMParams.model_fields if "cost" in f)
PTU_ZEROED_PRICING: Final[Mapping[str, float | tuple[()] | Mapping[str, float]]] = MappingProxyType(
{
**dict.fromkeys(PTU_ZEROED_PRICING_FIELDS, 0.0),
**dict.fromkeys(PTU_EMPTIED_PRICING_FIELDS, ()),
**dict.fromkeys(PTU_ZEROED_TABLE_FIELDS, MappingProxyType(dict.fromkeys(SEARCH_CONTEXT_SIZES, 0.0))),
}
)
@dataclass(frozen=True, slots=True)
class PTUTerms:
"""The reservation a deployment declares, once every field has been validated."""
team_id: str
ptu_count: int
cost_per_ptu_per_hour: float
effective_from: datetime
effective_to: datetime | None
def _to_utc(parsed: datetime) -> datetime:
"""``parsed`` as UTC, reading a naive value as UTC rather than local time."""
return parsed.replace(tzinfo=timezone.utc) if parsed.tzinfo is None else parsed.astimezone(timezone.utc)
def _as_utc(value: object) -> datetime | None:
"""A model_info datetime as UTC, parsing an ISO string, else None."""
if isinstance(value, datetime):
return _to_utc(value)
if not isinstance(value, str):
return None
try:
return _to_utc(datetime.fromisoformat(value.replace("Z", "+00:00")))
except ValueError:
return None
def ptu_terms(model_info: Mapping[str, object]) -> PTUTerms | None:
"""The reservation this deployment accrues flat cost for, else None.
A start is required rather than inferred because flat cost accrues from it, and a
present but unparseable bound would read as no bound and widen the window to the whole
day, so either one leaves the deployment unpriced until the config is fixed.
"""
ptu_count: Final = model_info.get("ptu_count")
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
team_id: Final = model_info.get("team_id")
if ptu_count is None or cost_per_hour is None or not team_id:
return None
try:
ptu_count_int: Final = int(ptu_count)
cost_per_hour_float: Final = float(cost_per_hour)
except (TypeError, ValueError, OverflowError):
return None
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
return None
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
return None
raw_from: Final = model_info.get("ptu_effective_from")
raw_to: Final = model_info.get("ptu_effective_to")
effective_from: Final = _as_utc(raw_from)
effective_to: Final = _as_utc(raw_to)
if effective_from is None or (raw_to is not None and effective_to is None):
return None
if effective_to is not None and effective_to <= effective_from:
return None
return PTUTerms(
team_id=str(team_id),
ptu_count=ptu_count_int,
cost_per_ptu_per_hour=cost_per_hour_float,
effective_from=effective_from,
effective_to=effective_to,
)
def zeroed_ptu_pricing(
model_info: Mapping[str, object], declared: Mapping[str, object]
) -> Mapping[str, float | tuple[()] | Mapping[str, float]] | None:
"""The pricing a deployment accruing flat cost must carry, else None.
Both conditions hold or nothing is zeroed. Without the flag no flat cost accrues, so
zeroing would leave the deployment serving for free with nothing charged in its place,
which is what an SDK user who happens to carry ptu_count would otherwise get. The terms
are checked first only because they are a few dict reads, while the flag can resolve
through a configured secret manager, and this runs for every deployment registered.
Any further rate the deployment itself declares is zeroed alongside the standing set,
since one left standing bills the traffic the reserved capacity already paid for.
"""
if ptu_terms(model_info) is None:
return None
if not is_ptu_cost_attribution_enabled():
return None
return MappingProxyType(
{
**PTU_ZEROED_PRICING,
**dict.fromkeys(
CUSTOM_PRICING_FIELDS.intersection(declared)
.difference(PTU_ZEROED_TABLE_FIELDS)
.difference(PTU_EMPTIED_PRICING_FIELDS),
0.0,
),
}
)

View file

@ -10,6 +10,7 @@ from typing_extensions import TypedDict
from litellm.litellm_core_utils.core_helpers import process_response_headers
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
@ -134,8 +135,11 @@ class BaseAnthropicMessagesStreamingIterator:
if self.completion_start_time is not None:
self.litellm_logging_obj.completion_start_time = self.completion_start_time
self.litellm_logging_obj.model_call_details["completion_start_time"] = self.completion_start_time
asyncio.create_task(
PassThroughStreamingHandler._route_streaming_logging_to_handler(
# Enqueue on the rooted logging worker rather than asyncio.create_task:
# this also runs during generator teardown after a client disconnect,
# where an unrooted task could be garbage-collected before it bills.
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
async_coroutine=PassThroughStreamingHandler._route_streaming_logging_to_handler(
litellm_logging_obj=self.litellm_logging_obj,
passthrough_success_handler_obj=GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ,
url_route="/v1/messages",
@ -197,13 +201,21 @@ class BaseAnthropicMessagesStreamingIterator:
collected_chunks: Final = []
saw_terminal_event = False
async for chunk in completion_stream:
if self.completion_start_time is None:
self.completion_start_time = datetime.now()
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
collected_chunks.append(encoded_chunk)
yield encoded_chunk
try:
async for chunk in completion_stream:
if self.completion_start_time is None:
self.completion_start_time = datetime.now()
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
collected_chunks.append(encoded_chunk)
yield encoded_chunk
except (GeneratorExit, asyncio.CancelledError):
# A client disconnect tears the generator down at the yield, so the
# post-loop logging below never runs and the tokens already streamed
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
if collected_chunks:
await self._handle_streaming_logging(collected_chunks)
raise
if not saw_terminal_event:
yield _incomplete_stream_error_sse_event()

View file

@ -842,8 +842,15 @@ class AmazonAnthropicClaudeMessagesConfig(
patched_stream: Final = self._promote_message_stop_usage(completion_stream)
async for chunk in handler.async_sse_wrapper(patched_stream):
yield chunk
sse_stream: Final = handler.async_sse_wrapper(patched_stream)
try:
async for chunk in sse_stream:
yield chunk
finally:
# Close the inner generator deterministically so a client disconnect
# (GeneratorExit here) reaches async_sse_wrapper's partial-spend logging
# now instead of at garbage collection. See LIT-5839.
await sse_stream.aclose()
@staticmethod
def _merge_message_start_cache_into_delta_usage(

View file

@ -5,7 +5,7 @@ import ssl
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from contextlib import asynccontextmanager
from functools import lru_cache
from types import ModuleType
from types import MappingProxyType, ModuleType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
@ -2087,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,
@ -2095,6 +2103,7 @@ class BaseLLMHTTPHandler:
"preset_cache_key": None,
"stream_response": {},
"model_info": kwargs.get("model_info"),
**vertex_location_params,
**anthropic_messages_optional_request_params,
},
custom_llm_provider=custom_llm_provider,

View file

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

View file

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

View file

@ -19760,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",
@ -19816,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",
@ -19871,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",
@ -38880,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,
@ -38904,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,
@ -39123,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,
@ -39152,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,
@ -39172,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,
@ -39203,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,
@ -39234,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,
@ -39266,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,
@ -39298,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,
@ -39330,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,
@ -39362,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,
@ -39395,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,
@ -39428,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,
@ -39461,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,
@ -39510,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,
@ -39523,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,
@ -39555,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,
@ -39602,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,
@ -40015,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",
@ -40071,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",
@ -47242,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,
@ -47274,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,

View file

@ -43,6 +43,7 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -53,8 +54,7 @@ class KeyManagementEventHooks:
except Exception as e:
verbose_proxy_logger.warning("Failed to send key created email: %s", e)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
_updated_values: Final = response.model_dump_json(exclude_none=True)
asyncio.create_task(
create_audit_log_for_update(
@ -103,11 +103,11 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
_updated_values: Final = json.dumps(data.json(exclude_none=True), default=str)
_before_value = existing_key_row.json(exclude_none=True)
@ -144,6 +144,7 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
@ -180,7 +181,7 @@ class KeyManagementEventHooks:
verbose_proxy_logger.warning("Failed to send key rotated email: %s", e)
# store the audit log
if litellm.store_audit_logs is True and existing_key_row.token is not None:
if is_audit_logging_enabled() and existing_key_row.token is not None:
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -218,12 +219,12 @@ class KeyManagementEventHooks:
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
# we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes
if litellm.store_audit_logs is True and data.keys is not None:
if is_audit_logging_enabled() and data.keys is not None:
# make an audit log for each key deleted
for key in keys_being_deleted:
if key.token is None:

View file

@ -20,7 +20,10 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
WebhookEvent,
)
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.repositories.user_repository import UserRepository
@ -203,7 +206,7 @@ class UserManagementEventHooks:
- user_api_key_dict: UserAPIKeyAuth - The user api key dictionary.
- litellm_proxy_admin_name: Optional[str] - The name of the proxy admin.
"""
if not litellm.store_audit_logs:
if not is_audit_logging_enabled():
return
from litellm.proxy.management_helpers.audit_logs import (

View file

@ -17,7 +17,6 @@ from typing import TYPE_CHECKING, Any, Final, Protocol
from fastapi import APIRouter, Depends, Header, HTTPException
from pydantic import BaseModel, Field
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._redis import _redis_kwargs_from_environment
from litellm._uuid import uuid
@ -299,14 +298,15 @@ async def _emit_cache_settings_audit_log(
exception. Captured under ``LiteLLM_CacheConfig`` so the row
co-locates with the table it mutates.
"""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
task: Final = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -100,16 +100,15 @@ async def _emit_hashicorp_vault_audit_log(
``LiteLLM_ConfigOverrides`` so the row co-locates with the table it
mutates.
"""
import litellm
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
task: Final = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -243,12 +243,15 @@ async def _emit_coordination_redis_audit_log(
litellm_changed_by: str | None,
) -> None:
"""Emit an audit-log row for a /coordination_redis/settings mutation."""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
task: Final = asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -2220,6 +2220,7 @@ async def delete_user(
)
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -2298,9 +2299,8 @@ async def delete_user(
},
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
# we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
# make an audit log for each team deleted
_user_row = user_row.json(exclude_none=True)

View file

@ -6245,6 +6245,7 @@ async def block_key(
"""
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -6291,7 +6292,7 @@ async def block_key(
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -6358,6 +6359,7 @@ async def unblock_key(
"""
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -6404,7 +6406,7 @@ async def unblock_key(
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(

View file

@ -64,7 +64,10 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
)
from litellm.proxy.management_helpers.audit_logs import get_audit_log_changed_by
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.repositories.table_repositories import (
MCPServerRepository,
MCPUserCredentialsRepository,
@ -2018,7 +2021,7 @@ if MCP_AVAILABLE:
await global_mcp_server_manager.reload_servers_from_database()
# TODO: Enterprise: Finish audit log trail
if litellm.store_audit_logs:
if is_audit_logging_enabled():
pass
# TODO: Delete from virtual keys
@ -2613,7 +2616,7 @@ if MCP_AVAILABLE:
)
# TODO: Enterprise: Finish audit log trail
if litellm.store_audit_logs:
if is_audit_logging_enabled():
pass
return _redact_mcp_credentials(mcp_server_record_updated)

View file

@ -24,6 +24,13 @@ from pydantic import BaseModel, ConfigDict, Field, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import LITELLM_PROXY_ADMIN_NAME
from litellm.litellm_core_utils.ptu_pricing import (
CUSTOM_PRICING_FIELDS,
PTU_EMPTIED_PRICING_FIELDS,
PTU_ZEROED_PRICING_FIELDS,
PTU_ZEROED_TABLE_FIELDS,
SEARCH_CONTEXT_SIZES,
)
from litellm.proxy._types import (
BlockModelRequest,
CommonProxyErrors,
@ -89,7 +96,6 @@ from litellm.types.router import (
ModelInfo,
updateDeployment,
)
from litellm.types.utils import CustomPricingLiteLLMParams
from litellm.utils import get_utc_datetime
router: Final = APIRouter()
@ -346,12 +352,8 @@ def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None:
# tiered_pricing is the one mirrored field that is a table of ranges, not a rate, so it is stored
# empty (see _PTU_EMPTIED_PRICING_FIELDS): its tiers outrank the zeros written beside them, so
# dropping it would leave the cost map's tiers billing the traffic the reserved capacity covers.
_PTU_ZEROED_PRICING_FIELDS: Final = tuple(f for f in SPECIAL_MODEL_INFO_PARAMS if f != "tiered_pricing") + (
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost_above_200k_tokens",
)
_PTU_EMPTIED_PRICING_FIELDS: Final = frozenset({"tiered_pricing"})
_PTU_ZEROED_PRICING_FIELDS: Final = PTU_ZEROED_PRICING_FIELDS
_PTU_EMPTIED_PRICING_FIELDS: Final = PTU_EMPTIED_PRICING_FIELDS
_PTU_ZEROED_PRICING: Final[Mapping[str, float | tuple[()]]] = MappingProxyType(
{
**dict.fromkeys(_PTU_ZEROED_PRICING_FIELDS, 0.0),
@ -363,13 +365,13 @@ _EMPTY_MODEL_INFO: Final[Mapping[str, object]] = _NO_PRICING_OVERRIDE
# Rate fields only. CustomPricingLiteLLMParams also carries settings that are not charges
# (an embedding's output_vector_size, the regional uplift multipliers), and zeroing one of
# those would destroy the deployment's configuration rather than stop a charge.
_CUSTOM_PRICING_FIELDS: Final = frozenset(f for f in CustomPricingLiteLLMParams.model_fields if "cost" in f)
_CUSTOM_PRICING_FIELDS: Final = CUSTOM_PRICING_FIELDS
# search_context_cost_per_query holds its rates in a table keyed by context size, and an absent
# table means the provider's own default rate rather than free (litellm/llms/gemini/cost_calculator
# falls back to $0.035), so it is zeroed in place rather than emptied like tiered_pricing, and
# written on every PTU deployment rather than only where a table is already stored.
_PTU_ZEROED_TABLE_FIELDS: Final = frozenset({"search_context_cost_per_query"})
_SEARCH_CONTEXT_SIZES: Final = ("search_context_size_low", "search_context_size_medium", "search_context_size_high")
_PTU_ZEROED_TABLE_FIELDS: Final = PTU_ZEROED_TABLE_FIELDS
_SEARCH_CONTEXT_SIZES: Final = SEARCH_CONTEXT_SIZES
def _is_nonzero_rate(value: object) -> bool:

View file

@ -13,7 +13,6 @@ from typing import Annotated, Any, Final
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.proxy._types import (
@ -182,14 +181,15 @@ async def _emit_team_callback_audit_log(
Callback secrets are redacted before serialization so the audit table
cannot itself become a credential-harvest sink.
"""
if litellm.store_audit_logs is not True:
return
from litellm.proxy.management_helpers.audit_logs import (
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import litellm_proxy_admin_name
if not is_audit_logging_enabled():
return
redacted_before: Final = _redact_callback_secrets(before_metadata)
redacted_after: Final = _redact_callback_secrets(after_metadata)

View file

@ -1249,6 +1249,7 @@ async def new_team(
try:
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
_license_check,
@ -1551,8 +1552,7 @@ async def new_team(
litellm_proxy_admin_name=litellm_proxy_admin_name,
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
_updated_values = complete_team_data.json(exclude_none=True)
_updated_values = json.dumps(_updated_values, default=str)
@ -1944,6 +1944,7 @@ async def update_team(
```
"""
try:
from litellm.proxy.management_helpers.audit_logs import is_audit_logging_enabled
from litellm.proxy.proxy_server import (
litellm_proxy_admin_name,
llm_router,
@ -2246,8 +2247,7 @@ async def update_team(
proxy_logging_obj=proxy_logging_obj,
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
await _create_team_update_audit_log(
existing_team_row=existing_team_row,
updated_kv=updated_kv,
@ -3712,6 +3712,7 @@ async def delete_team(
"""
from litellm.proxy.management_helpers.audit_logs import (
get_audit_log_changed_by,
is_audit_logging_enabled,
)
from litellm.proxy.proxy_server import (
create_audit_log_for_update,
@ -3756,9 +3757,8 @@ async def delete_team(
litellm_changed_by=litellm_changed_by,
)
# Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True
# we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes
if litellm.store_audit_logs is True:
if is_audit_logging_enabled():
# make an audit log for each team deleted
for team_id in data.team_ids:
team_row: LiteLLM_TeamTable | None = await prisma_client.get_data(

View file

@ -24,6 +24,22 @@ _audit_log_callback_cache: Final[dict[str, CustomLogger]] = {}
ALLOW_LITELLM_CHANGED_BY_HEADER_METADATA_KEY: Final = "allow_litellm_changed_by_header"
def is_audit_logging_enabled(store_audit_logs: bool | None = None) -> bool:
from litellm.secret_managers.main import get_secret_bool
configured_value: Final[bool | None] = litellm.store_audit_logs if store_audit_logs is None else store_audit_logs
if configured_value is not None:
return configured_value
environment_value: Final[bool | None] = get_secret_bool("LITELLM_STORE_AUDIT_LOGS")
if environment_value is not None:
return environment_value
from litellm.proxy.proxy_server import premium_user
return premium_user is True
def _allows_litellm_changed_by_header(user_api_key_dict: UserAPIKeyAuth) -> bool:
for admin_metadata in (user_api_key_dict.metadata, user_api_key_dict.team_metadata):
if (
@ -164,11 +180,7 @@ async def create_object_audit_log(
- user_api_key_dict: UserAPIKeyAuth - The user api key dictionary.
- litellm_proxy_admin_name: Optional[str] - The name of the proxy admin.
"""
from litellm.secret_managers.main import get_secret_bool
_store_audit_logs: Final[bool | None] = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS")
if _store_audit_logs is not True:
if not is_audit_logging_enabled():
return
_changed_by: Final = get_audit_log_changed_by(
@ -196,10 +208,7 @@ async def create_audit_log_for_update(request_data: LiteLLM_AuditLogs):
"""
Create an audit log for an object.
"""
from litellm.secret_managers.main import get_secret_bool
_store_audit_logs: Final[bool | None] = litellm.store_audit_logs or get_secret_bool("LITELLM_STORE_AUDIT_LOGS")
if _store_audit_logs is not True:
if not is_audit_logging_enabled():
return
from litellm.proxy.proxy_server import premium_user, prisma_client

View file

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

View file

@ -639,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
@ -4380,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:
@ -4553,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)
@ -4986,6 +5035,7 @@ class ProxyConfig:
)
elif key == "audit_log_callbacks":
from litellm.proxy.management_helpers.audit_logs import (
is_audit_logging_enabled,
reset_audit_log_callback_cache,
)
@ -5004,14 +5054,14 @@ class ProxyConfig:
litellm.audit_log_callbacks.append(callback)
_store_audit_logs = litellm_settings.get("store_audit_logs", litellm.store_audit_logs)
if _store_audit_logs:
if is_audit_logging_enabled(store_audit_logs=_store_audit_logs):
print( # noqa: T201
f"{blue_color_code} Initialized Audit Log Callbacks - {litellm.audit_log_callbacks} {reset_color_code}"
)
else:
verbose_proxy_logger.warning(
"'audit_log_callbacks' is configured but 'store_audit_logs' is not enabled. "
"Audit log callbacks will not fire until 'store_audit_logs: true' is added to litellm_settings."
"'audit_log_callbacks' is configured but audit logging is not enabled. "
"Audit log callbacks will not fire."
)
elif key == "cache_params":
# this is set in the cache branch
@ -5123,17 +5173,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)

View file

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

View file

@ -1,18 +1,12 @@
"""Opt-in flag for PTU (provisioned throughput unit) flat-cost attribution.
"""Re-exported from ``litellm.litellm_core_utils.ptu_pricing``.
The whole feature is inert unless an operator sets
``LITELLM_ENABLE_PTU_COST_ATTRIBUTION``: the daily rollup is not scheduled, the
model endpoints reject PTU config, the daily activity read path reports zero flat
cost, and the model form hides the PTU inputs.
The flag lives in core because the router reads it while registering a deployment, and
router code cannot import from the proxy.
"""
from typing import Final
from litellm.litellm_core_utils.ptu_pricing import (
PTU_COST_ATTRIBUTION_ENV_VAR,
is_ptu_cost_attribution_enabled,
)
from litellm.secret_managers.main import get_secret_bool
PTU_COST_ATTRIBUTION_ENV_VAR: Final = "LITELLM_ENABLE_PTU_COST_ATTRIBUTION"
def is_ptu_cost_attribution_enabled() -> bool:
"""Report whether this deployment opted into PTU flat-cost attribution."""
return get_secret_bool(PTU_COST_ATTRIBUTION_ENV_VAR, False) is True
__all__ = ("PTU_COST_ATTRIBUTION_ENV_VAR", "is_ptu_cost_attribution_enabled")

View file

@ -14,6 +14,7 @@ and share the existing unique constraint.
import asyncio
import json
import sys
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from datetime import date, datetime, time, timedelta, timezone
@ -29,14 +30,15 @@ from litellm.constants import (
PTU_ROLLUP_MAX_BACKFILL_DAYS,
PTU_SENTINEL_API_KEY,
)
from litellm.litellm_core_utils.ptu_pricing import ptu_terms
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.types.router import ModelInfo
if TYPE_CHECKING:
from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager
from litellm.proxy.utils import PrismaClient
_HOURS_PER_DAY: Final = 24
_PRUNE_ID_CHUNK_SIZE: Final = 5_000
_UPSERT_ATTEMPTS: Final = 3
_UPSERT_RETRY_BACKOFF_SECONDS: Final = 0.5
@ -72,28 +74,6 @@ class PTUModel:
effective_to: datetime | None = None
def _parse_utc_datetime(value: object) -> datetime | None:
"""Parse a model_info datetime (ISO string or datetime) into a UTC-aware datetime, else None."""
parsed: Final = _coerce_datetime(value)
if parsed is None:
return None
if parsed.tzinfo is None:
return parsed.replace(tzinfo=timezone.utc)
return parsed.astimezone(timezone.utc)
def _coerce_datetime(value: object) -> datetime | None:
"""``value`` as a datetime, parsing an ISO string, else None."""
if isinstance(value, datetime):
return value
if not isinstance(value, str):
return None
try:
return datetime.fromisoformat(value.replace("Z", "+00:00"))
except ValueError:
return None
def _public_model_name(row: object, model_info: Mapping[str, object]) -> str:
"""The name an operator recognises for this deployment.
@ -167,46 +147,20 @@ def _parse_ptu_model(row: object) -> PTUModel | None:
Valid means model_info has a positive ptu_count, a non-negative
cost_per_ptu_per_hour, and a team_id (1 model -> 1 team).
"""
raw_model_info: Final = getattr(row, "model_info", None)
model_info: Final = _decode_model_info(raw_model_info)
model_info: Final = _decode_model_info(getattr(row, "model_info", None))
if model_info is None:
return None
ptu_count: Final = model_info.get("ptu_count")
cost_per_hour: Final = model_info.get("cost_per_ptu_per_hour")
team_id: Final = model_info.get("team_id")
if ptu_count is None or cost_per_hour is None or not team_id:
return None
try:
ptu_count_int: Final = int(ptu_count)
cost_per_hour_float: Final = float(cost_per_hour)
except (TypeError, ValueError, OverflowError):
return None
if not 0 < ptu_count_int <= ModelInfo.MAX_PTU_COUNT:
return None
if not 0 <= cost_per_hour_float <= ModelInfo.MAX_COST_PER_PTU_PER_HOUR:
return None
if model_info.get("ptu_effective_from") is None:
# The endpoints require a start; a row without one predates that rule or was
# written around them, and inferring one would bill days the deployment did not exist
return None
raw_from: Final = model_info.get("ptu_effective_from")
raw_to: Final = model_info.get("ptu_effective_to")
effective_from: Final = _parse_utc_datetime(raw_from)
effective_to: Final = _parse_utc_datetime(raw_to)
# A present-but-unparseable bound would read as "no bound" and silently widen the
# window to the whole day, so the deployment is skipped until the config is fixed
if (raw_from is not None and effective_from is None) or (raw_to is not None and effective_to is None):
return None
if effective_from is not None and effective_to is not None and effective_to <= effective_from:
terms: Final = ptu_terms(model_info)
if terms is None:
return None
return PTUModel(
model_id=str(getattr(row, "model_id", "") or ""),
model_name=_public_model_name(row, model_info),
team_id=str(team_id),
ptu_count=ptu_count_int,
cost_per_ptu_per_hour=cost_per_hour_float,
effective_from=effective_from,
effective_to=effective_to,
team_id=terms.team_id,
ptu_count=terms.ptu_count,
cost_per_ptu_per_hour=terms.cost_per_ptu_per_hour,
effective_from=terms.effective_from,
effective_to=terms.effective_to,
)
@ -358,10 +312,70 @@ async def _upsert_charge_with_retry(
return False
async def _load_ptu_models(prisma_client: "PrismaClient") -> tuple[PTUModel, ...]:
"""Every model deployment currently carrying valid manual PTU config."""
@dataclass(frozen=True, slots=True)
class _LoadedDeployments:
"""The deployments a run will price, and every deployment id it looked at.
The id set is deliberately wider than the priced set. A deployment whose PTU config
was removed produces no charge and still has to be prunable, so bounding the prune on
what priced would strand its old rows forever. It is also a guaranteed superset of the
priced set, or a run could write a charge that falls outside its own delete filter.
"""
models: tuple[PTUModel, ...]
scanned_ids: frozenset[str]
config_sourced: bool
def _running_router() -> object | None:
"""The proxy's router, or None outside a running proxy.
Read out of ``sys.modules`` rather than imported, so a rollup driven from a test or a
script does not pull the whole proxy server in behind it.
"""
proxy_server: Final = sys.modules.get("litellm.proxy.proxy_server")
return getattr(proxy_server, "llm_router", None) if proxy_server is not None else None
def _config_deployments(router: object | None, *, owned_by_db: frozenset[str]) -> tuple[_PTUDeployment, ...]:
"""Deployments the router holds that no ``LiteLLM_ProxyModelTable`` row owns.
``db_model`` is forced True on every deployment loaded from that table and defaults to
False on ModelInfo, so the complement is what config.yaml declared. A per-request
credential clone carries ``original_model_id`` and reuses its source's PTU config under
a fresh id, so pricing it would bill one reservation once per distinct client key.
"""
entries: Final = tuple(getattr(router, "model_list", None) or ())
records: Final = tuple(_router_deployment(entry) for entry in entries)
return tuple(
record
for record in records
if record is not None
and record.model_info.get("db_model") is not True
and record.model_info.get("original_model_id") is None
and record.model_id not in owned_by_db
)
async def _load_ptu_models(prisma_client: "PrismaClient") -> _LoadedDeployments:
"""Every deployment carrying valid manual PTU config, and every id the scan saw.
Reserved capacity is billed by the provider whichever file declared it, so a
deployment the proxy only knows from config.yaml accrues alongside the stored ones.
"""
rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many()
return tuple(parsed for parsed in (_parse_ptu_model(row) for row in rows) if parsed is not None)
db_ids: Final = frozenset(model_id for row in rows if (model_id := str(getattr(row, "model_id", "") or "")))
config_records: Final = _config_deployments(_running_router(), owned_by_db=db_ids)
models: Final = tuple(
parsed for parsed in (_parse_ptu_model(row) for row in (*rows, *config_records)) if parsed is not None
)
return _LoadedDeployments(
models=models,
config_sourced=bool(config_records),
scanned_ids=db_ids
| frozenset(record.model_id for record in config_records)
| frozenset(model.model_id for model in models),
)
async def run_ptu_flat_cost_rollup(
@ -378,8 +392,10 @@ async def run_ptu_flat_cost_rollup(
The prune predicate is ``updated_at < run_started`` rather than "not in the charge
set I computed", which matters under concurrency: whether a row is garbage becomes a
property of the row instead of one run's in-memory config snapshot, so a run can
never delete a row a concurrent run just wrote. It is still skipped when any charge
failed to write, since a row whose replacement never landed would look unrefreshed.
never delete a row a concurrent run just wrote. It is bounded to the deployments this
run looked at, so a row it cannot account for is out of reach either way. It is still
skipped when any charge failed to write, since a row whose replacement never landed
would look unrefreshed.
"""
day: Final = target_date or (datetime.now(timezone.utc).date() - timedelta(days=1))
@ -390,7 +406,8 @@ async def run_ptu_flat_cost_rollup(
date_str: Final = day.isoformat()
run_started: Final = datetime.now(timezone.utc)
ptu_models: Final = await _load_ptu_models(prisma_client)
loaded: Final = await _load_ptu_models(prisma_client)
ptu_models: Final = loaded.models
charges: Final = _aggregate_charges(ptu_models, day)
landed: Final = tuple(
@ -415,7 +432,12 @@ async def run_ptu_flat_cost_rollup(
date_str,
)
else:
await _prune_unrefreshed_sentinel_rows(prisma_client, date_str=date_str, run_started=run_started)
await _prune_unrefreshed_sentinel_rows(
prisma_client,
date_str=date_str,
run_started=run_started,
scanned_ids=loaded.scanned_ids if loaded.config_sourced else None,
)
verbose_proxy_logger.info(
"PTU rollup for %s: %d PTU models processed, %d rows written, %d rows failed",
@ -524,7 +546,7 @@ async def run_ptu_flat_cost_backfill(
verbose_proxy_logger.warning("PTU backfill: prisma_client is None, skipping")
return BackfillResult(start=end, end=end, days_scanned=0, rows_written=0)
ptu_models: Final = await _load_ptu_models(prisma_client)
ptu_models: Final = (await _load_ptu_models(prisma_client)).models
days: Final = _backfill_window(ptu_models, end)
if not days:
@ -707,26 +729,61 @@ async def _prune_unrefreshed_sentinel_rows(
*,
date_str: str,
run_started: datetime,
scanned_ids: frozenset[str] | None,
) -> None:
"""Delete the day's PTU sentinel rows this run did not refresh.
"""Delete the day's PTU sentinel rows this run looked at and did not refresh.
Every charge the run wrote bumps ``updated_at`` past ``run_started``, so anything
left below that mark is a (team, model) the current config no longer prices. The mark
is pulled back by ``PTU_PRUNE_SKEW_GRACE_SECONDS`` because the two timestamps come
from different hosts: a stale row is hours old, a concurrently written one is seconds
old, and the grace separates them without waiting on clocks agreeing. The
predicate reads only the row, never the caller's config snapshot, which is what
makes it safe to run twice, out of order, or beside another pod: a row written
after this run began is out of reach of its delete. Mirrors the retention predicate
``SpendLogCleanup`` deletes by."""
Two conditions, and a row survives unless it meets both. It must be stale: every
charge the run wrote bumps ``updated_at`` past ``run_started``, so anything left below
that mark is a (team, model) the current config no longer prices. The mark is pulled
back by ``PTU_PRUNE_SKEW_GRACE_SECONDS`` because the two timestamps come from
different hosts, and the grace separates a row that is hours old from one written
seconds ago without waiting on clocks agreeing.
A run that priced a deployment only its own host declares must also name the
deployments it scanned. Staleness alone is sufficient while every run derives its
charges from the same table, because then any two runs compute the same set, so a
database-only run still sweeps by timestamp exactly as it always has. Once one host's
charges come from a file the others cannot read, a row it never considered is not
evidence of anything, and deleting it drops a charge that host is responsible for.
Where the bound applies the ids go out in chunks, because each is one bind variable and
the server rejects a statement carrying more than 32767 of them, which a proxy holding
that many deployments would otherwise hit every night with no handler above here.
"""
cutoff: Final = run_started - timedelta(seconds=PTU_PRUNE_SKEW_GRACE_SECONDS)
await prisma_client.db.litellm_dailyteamspend.delete_many(
where={ # mutable-ok: prisma delete filter
"date": date_str,
"api_key": PTU_SENTINEL_API_KEY,
"updated_at": {"lt": cutoff}, # mutable-ok: prisma comparison filter
}
unbounded: Final = { # mutable-ok: prisma delete filter
"date": date_str,
"api_key": PTU_SENTINEL_API_KEY,
"updated_at": {"lt": cutoff}, # mutable-ok: prisma comparison filter
}
ordered: Final = () if scanned_ids is None else tuple(sorted(scanned_ids))
filters: Final = (
(unbounded,)
if scanned_ids is None
else tuple(
MappingProxyType(
{
**unbounded,
"model": { # mutable-ok: prisma membership filter
"in": ordered[start : start + _PRUNE_ID_CHUNK_SIZE]
},
}
)
for start in range(0, len(ordered), _PRUNE_ID_CHUNK_SIZE)
)
)
deletions: Final = tuple(
[await prisma_client.db.litellm_dailyteamspend.delete_many(where=where) for where in filters]
)
deleted: Final = sum(deletions)
if deleted:
verbose_proxy_logger.info(
"PTU rollup for %s: pruned %s stale sentinel row(s) across %s deployment(s)",
date_str,
deleted,
"every" if scanned_ids is None else len(scanned_ids),
)
__all__ = (

View file

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

View file

@ -216,6 +216,15 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d
return {}
def _sl_attribution_fallback(
standard_logging_payload: StandardLoggingPayload | None,
field: Literal["model_id", "model_group", "api_base", "custom_llm_provider"],
) -> str:
if standard_logging_payload is None:
return ""
return standard_logging_payload.get(field) or ""
def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload:
if kwargs is None:
kwargs = {}
@ -288,8 +297,15 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
): # use 'tags' from standard logging payload instead
request_tags = safe_dumps(standard_logging_payload["request_tags"])
_model_id: Final = metadata.get("model_info", {}).get("id", "")
_model_group: Final = metadata.get("model_group", "")
_model_id: Final = metadata.get("model_info", {}).get("id", "") or _sl_attribution_fallback(
standard_logging_payload, "model_id"
)
_model_group: Final = metadata.get("model_group", "") or _sl_attribution_fallback(
standard_logging_payload, "model_group"
)
_api_base: Final = litellm_params.get("api_base", "") or _sl_attribution_fallback(
standard_logging_payload, "api_base"
)
# Extract overhead from hidden_params if available
litellm_overhead_time_ms = None
@ -389,7 +405,11 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
# Extract agent_id for A2A requests (set directly on model_call_details)
agent_id: Final[str | None] = kwargs.get("agent_id") or metadata.get("agent_id")
custom_llm_provider: Final = kwargs.get("custom_llm_provider")
custom_llm_provider: Final = (
kwargs.get("custom_llm_provider")
or _sl_attribution_fallback(standard_logging_payload, "custom_llm_provider")
or None
)
raw_model: Final = cast(str, kwargs.get("model") or "")
model_name: Final = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})
@ -414,13 +434,13 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
completion_tokens=usage.get("completion_tokens", standard_logging_completion_tokens),
request_tags=request_tags,
end_user=end_user_id or "",
api_base=litellm_params.get("api_base", ""),
api_base=_api_base,
model_group=_model_group,
model_id=_model_id,
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
agent_id=agent_id,
requester_ip_address=clean_metadata.get("requester_ip_address", None),
custom_llm_provider=kwargs.get("custom_llm_provider", ""),
custom_llm_provider=custom_llm_provider or "",
messages=_get_messages_for_spend_logs_payload(
standard_logging_payload=standard_logging_payload, metadata=metadata
),

View file

@ -516,6 +516,34 @@ def _failure_usage_to_lift(
return estimated_usage, 0.0
_EMPTY_LIFT: Final = MappingProxyType({})
def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]:
"""Failure-path callbacks run after ``litellm_logging_obj`` is popped from
request_data (it is not serialisable), so the caller merges these fields
onto request_data first: the first-handoff instant for preprocessing
latency, recovered or estimated usage for token counts, and the standard
logging object for deployment attribution on failed-request spend logs."""
_logging_obj: Final = request_data.get("litellm_logging_obj")
if _logging_obj is None:
return _EMPTY_LIFT
_model_call_details: Final = getattr(_logging_obj, "model_call_details", {})
_first_handoff: Final = _model_call_details.get("first_api_call_start_time")
_usage_to_lift: Final = _failure_usage_to_lift(
model_call_details=_model_call_details,
request_body=request_data,
dispatched=_first_handoff is not None,
)
_entries: Final = (
("first_api_call_start_time", _first_handoff),
("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]),
("response_cost", None if _usage_to_lift is None else (_usage_to_lift[1] or 0.0)),
("standard_logging_object", _model_call_details.get("standard_logging_object")),
)
return MappingProxyType({key: value for key, value in _entries if value is not None})
@dataclass(frozen=True)
class _CallbackCapabilities:
"""Cached per-hook capability flags derived from ``litellm.callbacks``.
@ -2280,6 +2308,11 @@ class ProxyLogging:
)
)
# Auth and pass-through failure bodies are unstripped client input, and
# the logging handler below flattens body keys into model_call_details,
# so drop the key before it can masquerade as the built payload.
request_data.pop("standard_logging_object", None)
### LOGGING ###
if self._is_proxy_only_llm_api_error(
original_exception=original_exception,
@ -2293,29 +2326,7 @@ class ProxyLogging:
original_exception=original_exception,
)
# Lift the first-handoff instant onto request_data (top-level
# internal key, not metadata) so failure-path callbacks can still
# compute preprocessing latency after the logging object is popped.
_logging_obj: Final = request_data.get("litellm_logging_obj")
if _logging_obj is not None:
_model_call_details: Final = getattr(_logging_obj, "model_call_details", {})
_first_handoff: Final = _model_call_details.get("first_api_call_start_time")
if _first_handoff is not None:
request_data["first_api_call_start_time"] = _first_handoff
# Lift recovered partial-stream usage, or an estimated input-side
# usage for a dispatched failure, onto request_data so the
# failure-path spend callbacks (which run after the logging object
# is popped) record real token counts instead of zero.
_usage_to_lift: Final = _failure_usage_to_lift(
model_call_details=_model_call_details,
request_body=request_data,
dispatched=_first_handoff is not None,
)
if _usage_to_lift is not None:
_lifted_usage, _lifted_cost = _usage_to_lift
request_data["combined_usage_object"] = _lifted_usage
request_data["response_cost"] = _lifted_cost
request_data.update(_failure_fields_to_lift(request_data))
# Remove before callbacks iterate — not serialisable
request_data.pop("litellm_logging_obj", None)

View file

@ -64,6 +64,7 @@ from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.litellm_core_utils.ptu_pricing import zeroed_ptu_pricing
from litellm.litellm_core_utils.request_timeout_resolver import (
get_configured_request_timeout,
)
@ -7695,7 +7696,16 @@ class Router:
- None: If the deployment is not active for the current environment (if 'supported_environments' is set in litellm_params)
"""
try:
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**_litellm_params)
zeroed_pricing: Final = (
zeroed_ptu_pricing(_model_info, _litellm_params) if _model_info.get("db_model") is not True else None
)
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(
**(
_litellm_params
if zeroed_pricing is None
else MappingProxyType({**_litellm_params, **zeroed_pricing})
)
)
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
deployment = Deployment(
**deployment_info,
@ -11191,6 +11201,8 @@ class Router:
if pre_routing_hook_response is not None:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
if pre_routing_hook_response.litellm_params:
request_kwargs.update(pre_routing_hook_response.litellm_params)
#########################################################
# Resolve the strategy and logger AFTER the pre-routing hook, since
@ -11300,6 +11312,8 @@ class Router:
if pre_routing_hook_response is not None:
model = pre_routing_hook_response.model
messages = pre_routing_hook_response.messages
if pre_routing_hook_response.litellm_params:
request_kwargs.update(pre_routing_hook_response.litellm_params)
# 2. Get healthy deployments
healthy_deployments: Final = await self.async_get_healthy_deployments(

View file

@ -53,6 +53,21 @@ model_list:
REASONING: o1-preview
```
Each tier can also use a model entry with request parameter overrides. A tier value may be
a model string, a single object, or a list mixing strings and objects. Object entries must
contain a model name and may contain any LiteLLM request parameters. The model name must
still resolve to a deployment in `model_list`; this configuration does not create one
```yaml
tiers:
COMPLEX: opus
REASONING:
- model_name: opus
litellm_params:
reasoning_effort: xhigh
- abc
```
### Renaming the tiers
`tier_labels` puts your own vocabulary on the four tiers:

View file

@ -30,6 +30,7 @@ from litellm.constants import EMPTY_MAPPING, RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import forwarded_internal_call_metadata
from litellm.litellm_core_utils.sensitive_data_masker import mask_credentials_in_payload
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.types.utils import (
AUTOROUTER_CLASSIFIER_CALL_ORIGIN,
@ -663,6 +664,35 @@ class ClassificationOutcome(NamedTuple):
classifier_cost: float | None = None
class _SessionAffinityPin(NamedTuple):
model: str
tier: ComplexityTier | None
def _parse_session_affinity_pin(value: object) -> _SessionAffinityPin | None:
if isinstance(value, str):
return _SessionAffinityPin(model=value, tier=None)
parts: Final[tuple[object, object] | None] = (
(value.get("model"), value.get("tier"))
if isinstance(value, Mapping)
else (value[0], value[1])
if isinstance(value, (list, tuple)) and len(value) == 2
else None
)
if parts is None:
return None
model, tier_value = parts
if not isinstance(model, str):
return None
tier: Final = ComplexityTier(tier_value) if isinstance(tier_value, str) else None
return _SessionAffinityPin(model=model, tier=tier)
def _session_affinity_cache_value(model: str, tier: ComplexityTier | str | None) -> Mapping[str, str | None]:
tier_value: Final = _tier_name(tier) if tier is not None else None
return {"model": model, "tier": tier_value} # mutable-ok: cache requires JSON mapping
class ComplexityRouter(CustomLogger):
"""
Complexity router that classifies requests and routes to appropriate models.
@ -1078,6 +1108,7 @@ class ComplexityRouter(CustomLogger):
classifier_model: str | None = None,
classifier_cost: float | None = None,
conversation_continuing: bool = True,
tier_litellm_params: Mapping[str, object] | None = None,
) -> StandardLoggingRoutingDecision:
"""Assemble the per-request provenance record for this router's decision.
@ -1127,6 +1158,10 @@ class ComplexityRouter(CustomLogger):
decision["classifier_model"] = classifier_model
if classifier_cost is not None:
decision["classifier_cost"] = classifier_cost
if tier_litellm_params:
masked_tier_litellm_params: Final = mask_credentials_in_payload(tier_litellm_params)
if isinstance(masked_tier_litellm_params, Mapping):
decision["tier_litellm_params"] = masked_tier_litellm_params
return decision
async def aclassify(
@ -1457,6 +1492,13 @@ class ComplexityRouter(CustomLogger):
raise ValueError(f"No model configured for tier {tier_key} and no default_model set")
def _litellm_params_for_model(self, tier: ComplexityTier | str | None, model: str) -> Mapping[str, object]:
if tier is None:
return MappingProxyType({})
entries: Final = self.config.tier_model_configs.get(_tier_name(tier), ())
entry: Final = next((candidate for candidate in entries if candidate.model_name == model), None)
return entry.litellm_params if entry is not None else MappingProxyType({})
@staticmethod
def _pick_from_tier_value(model: str | list[str], tier_key: str) -> str:
if isinstance(model, str):
@ -2068,9 +2110,10 @@ class ComplexityRouter(CustomLogger):
cache_key = self._get_session_affinity_cache_key(session_id, request_kwargs) if session_id is not None else None
if cache_key is not None:
pinned_model: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
if isinstance(pinned_model, str):
routed_model: str | None = pinned_model
pinned_value: Final = await self.litellm_router_instance.cache.async_get_cache(key=cache_key)
pinned_pin: Final = _parse_session_affinity_pin(pinned_value)
if pinned_pin is not None:
routed_model: str | None = pinned_pin.model
pin_escalation_keyword: str | None = None
if self.escalation_keywords:
user_message: Final = (
@ -2079,16 +2122,21 @@ class ComplexityRouter(CustomLogger):
if user_message is not None:
pin_escalation_keyword = self._matched_escalation_keyword(user_message)
if pin_escalation_keyword is not None:
routed_model = self._escalated_pin(pinned_model)
routed_model = self._escalated_pin(pinned_pin.model)
if routed_model is not None:
escalated: Final = routed_model != pinned_model
escalated: Final = routed_model != pinned_pin.model
resolved_pin_tier: Final = (
pinned_pin.tier
if not escalated and pinned_pin.tier is not None
else self._tier_for_model(routed_model)
)
# The floor outranks the pin because plan mode is a transient state of the
# session, not a request to move it: the turns carrying the sentinel route at
# the floor, and the stored pin deliberately keeps the session's own model so
# the first turn after plan mode exits auto-routes exactly as it would have.
# Escalation is the opposite on purpose -- an explicit ask to re-pin higher.
pin_plan_sentinel: Final = self._matched_plan_mode_signal(request_kwargs, resolved_messages)
pinned_tier: Final = self._tier_for_model(routed_model) if pin_plan_sentinel is not None else None
pinned_tier: Final = resolved_pin_tier if pin_plan_sentinel is not None else None
plan_floored: Final = (
pinned_tier is not None and self._apply_plan_mode_floor(pinned_tier) != pinned_tier
)
@ -2099,7 +2147,7 @@ class ComplexityRouter(CustomLogger):
# pin mid-conversation just because it outlives the original write.
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=session_model,
value=_session_affinity_cache_value(session_model, resolved_pin_tier),
ttl=self.config.session_affinity_ttl_seconds,
)
if self.config.adaptive:
@ -2118,19 +2166,23 @@ class ComplexityRouter(CustomLogger):
verbose_router_logger.info(
"ComplexityRouter: routing decision cause=%s, routed_model=%s", cause, routed_model
)
routed_pin_tier: Final = self._tier_for_model(routed_model) if plan_floored else resolved_pin_tier
session_tier_litellm_params: Final = self._litellm_params_for_model(routed_pin_tier, routed_model)
has_original_messages: Final = messages is not None and len(messages) > 0
return self._with_session_deployment_affinity(
PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=session_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
cause=cause,
tier=self._tier_for_model(routed_model),
tier=routed_pin_tier,
matched_keyword=pin_plan_sentinel if plan_floored else None,
escalation_keyword=pin_escalation_keyword,
escalated=escalated,
conversation_continuing=conversation_continuing,
tier_litellm_params=session_tier_litellm_params,
),
)
)
@ -2157,7 +2209,10 @@ class ComplexityRouter(CustomLogger):
if pinnable and cache_key is not None and response is not None:
await self.litellm_router_instance.cache.async_set_cache(
key=cache_key,
value=response.model,
value=_session_affinity_cache_value(
response.model,
response.routing_decision.get("tier") if response.routing_decision is not None else None,
),
ttl=self.config.session_affinity_ttl_seconds,
)
return self._with_session_deployment_affinity(response)
@ -2271,6 +2326,7 @@ class ComplexityRouter(CustomLogger):
)
keyword_plan_floored: Final = routed_tier != escalated_tier
routed_model = await self._pick_model_for_tier(routed_tier, messages, resolved_messages, request_kwargs)
keyword_tier_litellm_params: Final = self._litellm_params_for_model(routed_tier, routed_model)
keyword_cause: Final[RoutingDecisionCause] = (
"plan_mode"
if keyword_plan_floored
@ -2286,6 +2342,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=keyword_tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
conversation_continuing=conversation_continuing,
@ -2294,6 +2351,7 @@ class ComplexityRouter(CustomLogger):
matched_keyword=plan_mode_sentinel if keyword_plan_floored else override.matched_keyword,
escalation_keyword=escalation_keyword,
escalated=keyword_escalated,
tier_litellm_params=keyword_tier_litellm_params,
),
)
@ -2380,6 +2438,7 @@ class ComplexityRouter(CustomLogger):
routed_model,
)
tier_litellm_params: Final = self._litellm_params_for_model(tier, routed_model)
classifier_model: Final = (
self.config.classifier_llm_config.model
if outcome.cause == "llm_classifier" and self.config.classifier_llm_config is not None
@ -2405,6 +2464,7 @@ class ComplexityRouter(CustomLogger):
return PreRoutingHookResponse(
model=routed_model,
messages=messages if has_original_messages else None,
litellm_params=tier_litellm_params,
routing_decision=self._build_routing_decision(
routed_model=routed_model,
conversation_continuing=conversation_continuing,
@ -2417,5 +2477,6 @@ class ComplexityRouter(CustomLogger):
escalated=escalated,
classifier_model=classifier_model,
classifier_cost=outcome.classifier_cost,
tier_litellm_params=tier_litellm_params,
),
)

View file

@ -5,10 +5,12 @@ Contains default keyword lists, weights, tier boundaries, and configuration clas
All values are configurable via proxy config.yaml.
"""
from collections.abc import Mapping
from enum import Enum
from typing import Final, Literal
from types import MappingProxyType
from typing import Annotated, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from pydantic import BaseModel, ConfigDict, Field, SkipValidation, field_serializer, field_validator, model_validator
from litellm.types.router import AdaptiveRouterWeights, ClassifierPlugin, RoutingPlugin
@ -159,6 +161,44 @@ class ReminderMarkerPair(BaseModel):
return self
class ComplexityTierModel(BaseModel):
model_config = ConfigDict(frozen=True)
model_name: str
litellm_params: Annotated[Mapping[str, object], SkipValidation()] = Field(
default_factory=lambda: MappingProxyType({})
)
@field_validator("litellm_params", mode="before")
@classmethod
def _freeze_litellm_params(cls, value: Mapping[str, object]) -> Mapping[str, object]:
return MappingProxyType(dict(value))
@field_serializer("litellm_params")
def _serialize_litellm_params(self, value: Mapping[str, object]) -> Mapping[str, object]:
return dict(value) # mutable-ok: Pydantic JSON serialization requires a concrete mapping
def _normalize_tier_entries(
raw_value: object,
tier: str,
) -> tuple[str | list[str], tuple[ComplexityTierModel, ...]]:
raw_entries: Final = raw_value if isinstance(raw_value, (list, tuple)) else (raw_value,)
entries: Final = tuple(
ComplexityTierModel(model_name=entry) if isinstance(entry, str) else ComplexityTierModel.model_validate(entry)
for entry in raw_entries
)
model_names: Final = tuple(entry.model_name for entry in entries)
if len(model_names) != len(frozenset(model_names)):
raise ValueError(f"tier {tier} contains duplicate model_name values; each pool entry needs distinct parameters")
normalized: Final = (
entries[0].model_name
if not isinstance(raw_value, (list, tuple))
else list(model_names) # mutable-ok: config.tiers must preserve its existing list contract
)
return normalized, entries
# ─── Default Keyword Lists ───
# Note: Keywords should be full words/phrases to avoid substring false positives.
# The matching logic uses word boundary detection for single-word keywords.
@ -425,6 +465,9 @@ class ComplexityRouterConfig(BaseModel):
"A list is randomly picked from when adaptive=False, and used as a soft-floor home pool when adaptive=True"
),
)
tier_model_configs: Mapping[str, tuple[ComplexityTierModel, ...]] = Field(
default_factory=dict,
)
tier_definitions: tuple[TierDefinition, ...] | None = Field(
default=None,
@ -777,6 +820,55 @@ class ComplexityRouterConfig(BaseModel):
coerced[key] = item
return coerced
@model_validator(mode="before")
@classmethod
def _normalize_tier_model_configs(cls, value: object) -> object:
if not isinstance(value, dict):
return value
raw_tiers: Final = value.get("tiers")
if not isinstance(raw_tiers, dict):
return value
existing_configs: Final = value.get("tier_model_configs")
normalized_entries: Final = MappingProxyType(
{tier: _normalize_tier_entries(raw_value, tier) for tier, raw_value in raw_tiers.items()}
)
normalized_tiers: Final = MappingProxyType(
{tier: normalized for tier, (normalized, _) in normalized_entries.items()}
)
incoming_params: Final = (
MappingProxyType(
{
(tier, entry.model_name): entry.litellm_params
for tier, entries in existing_configs.items()
for entry in (ComplexityTierModel.model_validate(item) for item in entries)
}
)
if isinstance(existing_configs, dict)
else MappingProxyType({})
)
tier_model_configs: Final = MappingProxyType(
{
tier: tuple(
entry.model_copy(
update=MappingProxyType(
{
"litellm_params": incoming_params.get((tier, entry.model_name), entry.litellm_params),
}
)
)
for entry in entries
)
for tier, (_, entries) in normalized_entries.items()
if any(entry.litellm_params for entry in entries)
or (isinstance(existing_configs, dict) and tier in existing_configs)
}
)
return { # mutable-ok: Pydantic before-validator requires a concrete mapping
**value,
"tiers": normalized_tiers,
"tier_model_configs": tier_model_configs,
}
@field_validator("escalation_keywords")
@classmethod
def _normalize_escalation_keywords(cls, value: list[str] | None) -> list[str] | None:

View file

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

View file

@ -4,6 +4,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc
import datetime
import enum
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints
@ -897,6 +898,7 @@ class PreRoutingHookResponse(BaseModel):
messages: list[dict[str, Any]] | None
routing_decision: StandardLoggingRoutingDecision | None = None
session_affinity_ttl_seconds: int | None = None
litellm_params: Mapping[str, object] | None = None
_PreRoutingStrategyT_co = TypeVar("_PreRoutingStrategyT_co", covariant=True)

View file

@ -248,6 +248,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
regional_processing_uplift_multiplier_us: (
float | None
) # OpenAI US data-residency uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
regional_endpoint_uplift_multiplier: ReadOnly[
float | None
] # Vertex AI non-global (regional) endpoint uplift multiplier applied to all token costs (e.g. 1.10 = +10%)
output_cost_per_character: float | None # only for vertex ai models
output_cost_per_audio_token: float | None
output_cost_per_token_above_128k_tokens: float | None # only for vertex ai models
@ -2840,6 +2843,7 @@ class StandardLoggingRoutingDecision(TypedDict, total=False):
conversation_continuing: bool
savings_baseline_model: str
savings_baseline_deployment_id: str
tier_litellm_params: Mapping[str, object] # writable-ok: Pydantic warns on ReadOnly TypedDict fields
# Fields whose values quote the caller's prompt. Dropped when an operator turns message
@ -2865,6 +2869,7 @@ DERIVED_ROUTING_DECISION_FIELDS: Final[frozenset[str]] = frozenset(
"conversation_continuing",
"savings_baseline_model",
"savings_baseline_deployment_id",
"tier_litellm_params",
}
)
@ -3115,16 +3120,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)
@ -3390,6 +3396,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]:

View file

@ -5662,6 +5662,7 @@ def _get_model_info_helper(
regional_processing_uplift_multiplier_us=_model_info.get(
"regional_processing_uplift_multiplier_us", None
),
regional_endpoint_uplift_multiplier=_model_info.get("regional_endpoint_uplift_multiplier", None),
output_cost_per_audio_token=_model_info.get("output_cost_per_audio_token", None),
output_cost_per_character=_model_info.get("output_cost_per_character", None),
output_cost_per_reasoning_token=_model_info.get("output_cost_per_reasoning_token", None),

View file

@ -19760,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",
@ -19816,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",
@ -19871,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",
@ -38880,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,
@ -38904,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,
@ -39123,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,
@ -39152,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,
@ -39172,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,
@ -39203,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,
@ -39234,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,
@ -39266,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,
@ -39298,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,
@ -39330,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,
@ -39362,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,
@ -39395,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,
@ -39428,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,
@ -39461,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,
@ -39510,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,
@ -39523,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,
@ -39555,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,
@ -39602,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,
@ -40015,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",
@ -40071,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",
@ -47242,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,
@ -47274,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,

View file

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

View file

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

View file

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

View file

@ -0,0 +1,163 @@
"""Tests for the shared PTU rules: which deployments accrue flat cost, and what that zeroes."""
import os
from datetime import datetime, timezone
from unittest.mock import patch
import pytest
from litellm.litellm_core_utils.ptu_pricing import (
CUSTOM_PRICING_FIELDS,
PTU_EMPTIED_PRICING_FIELDS,
PTU_ZEROED_PRICING_FIELDS,
PTU_ZEROED_TABLE_FIELDS,
SEARCH_CONTEXT_SIZES,
ptu_terms,
zeroed_ptu_pricing,
)
from litellm.types.router import ModelInfo
_VALID = {
"team_id": "team-alpha",
"ptu_count": 100,
"cost_per_ptu_per_hour": 0.02,
"ptu_effective_from": "2026-01-01T00:00:00Z",
}
def _with_flag(model_info, declared=None, enabled=True):
with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True" if enabled else ""}, clear=False):
return zeroed_ptu_pricing(model_info, declared or {})
def test_a_complete_reservation_is_accepted():
terms = ptu_terms(_VALID)
assert terms is not None
assert terms.team_id == "team-alpha"
assert terms.ptu_count == 100
assert terms.effective_from == datetime(2026, 1, 1, tzinfo=timezone.utc)
assert terms.effective_to is None
@pytest.mark.parametrize(
"override",
[
{"team_id": None},
{"team_id": ""},
{"ptu_count": None},
{"cost_per_ptu_per_hour": None},
{"ptu_count": 0},
{"ptu_count": -1},
{"ptu_count": ModelInfo.MAX_PTU_COUNT + 1},
{"cost_per_ptu_per_hour": -0.01},
{"cost_per_ptu_per_hour": ModelInfo.MAX_COST_PER_PTU_PER_HOUR + 1},
{"ptu_count": "not-a-number"},
{"ptu_effective_from": None},
{"ptu_effective_from": "not-a-date"},
{"ptu_effective_to": "not-a-date"},
{"ptu_effective_to": "2025-01-01T00:00:00Z"},
{"ptu_effective_to": "2026-01-01T00:00:00Z"},
],
ids=[
"no team",
"blank team",
"no count",
"no rate",
"zero count",
"negative count",
"count over the cap",
"negative rate",
"rate over the cap",
"count not a number",
"no start",
"unparseable start",
"unparseable end",
"end before start",
"end equal to start",
],
)
def test_an_incomplete_reservation_accrues_nothing(override):
"""Anything the rollup declines to charge must also decline to be zeroed, or the
deployment serves its traffic for free with nothing charged in its place."""
assert ptu_terms({**_VALID, **override}) is None
assert _with_flag({**_VALID, **override}) is None
def test_a_naive_start_is_read_as_utc():
"""config.yaml is hand-typed, and pydantic hands back a naive datetime for a date with
no offset."""
terms = ptu_terms({**_VALID, "ptu_effective_from": datetime(2026, 5, 1, 12, 0)})
assert terms is not None
assert terms.effective_from == datetime(2026, 5, 1, 12, 0, tzinfo=timezone.utc)
def test_an_offset_start_is_converted_rather_than_relabelled():
terms = ptu_terms({**_VALID, "ptu_effective_from": "2026-05-01T12:00:00-05:00"})
assert terms is not None
assert terms.effective_from == datetime(2026, 5, 1, 17, 0, tzinfo=timezone.utc)
def test_nothing_is_zeroed_while_the_feature_is_off():
"""No flat cost accrues with the flag off, so zeroing would serve the traffic free."""
assert _with_flag(_VALID, enabled=False) is None
def test_the_standing_rates_are_all_zeroed():
override = _with_flag(_VALID)
assert override is not None
assert [field for field in PTU_ZEROED_PRICING_FIELDS if override[field] != 0.0] == []
def test_tiered_pricing_is_emptied_rather_than_zeroed():
"""A tier outranks the flat rates written beside it, so a zero there would leave the
cost map's tiers billing the traffic the reserved capacity already covers."""
override = _with_flag(_VALID, declared={"tiered_pricing": [{"range": [0, 1000], "input_cost_per_token": 0.003}]})
assert override is not None
for field in PTU_EMPTIED_PRICING_FIELDS:
assert override[field] == ()
def test_the_search_context_table_is_zeroed_in_place_on_every_deployment():
"""An absent table means the provider's own default rather than free, so it is written
even when the deployment never declared one."""
override = _with_flag(_VALID)
assert override is not None
for field in PTU_ZEROED_TABLE_FIELDS:
assert dict(override[field]) == dict.fromkeys(SEARCH_CONTEXT_SIZES, 0.0)
def test_a_declared_table_does_not_become_a_scalar():
"""Zeroing it as a plain 0.0 would leave the provider's reader without a table to
consult, which is the same as absent."""
override = _with_flag(_VALID, declared={"search_context_cost_per_query": {"search_context_size_medium": 0.05}})
assert override is not None
assert dict(override["search_context_cost_per_query"]) == dict.fromkeys(SEARCH_CONTEXT_SIZES, 0.0)
def test_a_rate_the_deployment_declares_itself_is_zeroed_too():
"""The standing set covers the mirrored rates. Anything else the operator wrote would
otherwise survive and bill the traffic the hourly charge already paid for."""
extra = "input_cost_per_token_above_200k_tokens"
assert extra in CUSTOM_PRICING_FIELDS
assert extra not in PTU_ZEROED_PRICING_FIELDS
override = _with_flag(_VALID, declared={extra: 9e-06})
assert override is not None
assert override[extra] == 0.0
def test_a_setting_that_is_not_a_charge_is_left_alone():
"""CustomPricingLiteLLMParams also carries configuration, and zeroing one of those
would break the deployment rather than stop a charge."""
override = _with_flag(_VALID, declared={"output_vector_size": 1536})
assert override is not None
assert "output_vector_size" not in override

View file

@ -1,3 +1,4 @@
import asyncio
import json
import os
import sys
@ -20,9 +21,11 @@ class _RecordingLoggingIterator(BaseAnthropicMessagesStreamingIterator):
def __init__(self, litellm_logging_obj: LiteLLMLoggingObj, request_body: dict):
super().__init__(litellm_logging_obj=litellm_logging_obj, request_body=request_body)
self.logged_chunks: list = []
self.logging_call_count: int = 0
async def _handle_streaming_logging(self, collected_chunks):
self.logged_chunks = list(collected_chunks)
self.logging_call_count += 1
def _make_logging_obj(test_name: str) -> LiteLLMLoggingObj:
@ -233,6 +236,70 @@ async def test_async_sse_wrapper_excludes_synthetic_error_event_from_logged_chun
assert not any(chunk.startswith(b"event: error\n") for chunk in iterator.logged_chunks)
async def _events_then_hang(events):
for event in events:
yield event
await asyncio.Event().wait()
@pytest.mark.asyncio
async def test_async_sse_wrapper_logs_partial_chunks_on_client_disconnect():
"""
Regression test for LIT-5839: a client disconnect tears the generator
down with GeneratorExit at the yield, which used to skip the post-loop
logging dispatch entirely, so the partial output tokens the provider
already generated (and billed) never reached spend tracking.
"""
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_disconnect_logs_partial_chunks"),
request_body={},
)
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
streamed = [await wrapped.__anext__() for _ in range(len(TRUNCATED_TOOL_USE_EVENTS))]
assert iterator.logging_call_count == 0
await wrapped.aclose()
assert iterator.logging_call_count == 1
assert iterator.logged_chunks == streamed
@pytest.mark.asyncio
async def test_async_sse_wrapper_logs_partial_chunks_on_cancellation():
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_cancellation_logs_partial_chunks"),
request_body={},
)
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
streamed = [await wrapped.__anext__() for _ in range(len(TRUNCATED_TOOL_USE_EVENTS))]
consume_task = asyncio.ensure_future(wrapped.__anext__())
await asyncio.sleep(0.01)
consume_task.cancel()
with pytest.raises(asyncio.CancelledError):
await consume_task
assert iterator.logging_call_count == 1
assert iterator.logged_chunks == streamed
@pytest.mark.asyncio
async def test_async_sse_wrapper_skips_logging_on_disconnect_before_first_chunk():
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_disconnect_before_first_chunk"),
request_body={},
)
wrapped = iterator.async_sse_wrapper(_events_then_hang(()))
consume_task = asyncio.ensure_future(wrapped.__anext__())
await asyncio.sleep(0.01)
consume_task.cancel()
with pytest.raises(asyncio.CancelledError):
await consume_task
assert iterator.logging_call_count == 0
def test_incomplete_stream_error_sse_event_is_valid_anthropic_error():
event = _incomplete_stream_error_sse_event().decode()
lines = event.split("\n")

View file

@ -2948,3 +2948,38 @@ def test_bedrock_invoke_messages_allows_converted_websearch_function_tool():
headers={},
)
assert result["tools"][0]["name"] == "litellm_web_search"
@pytest.mark.asyncio
async def test_bedrock_sse_wrapper_dispatches_logging_on_client_disconnect():
"""
Regression test for LIT-5839: closing the outer bedrock_sse_wrapper
mid-stream (what the proxy does on a client disconnect) must close the
inner async_sse_wrapper deterministically so the partial-stream logging
fires. `completion_start_time` is only stamped on the logging object by
that dispatch, so it observing a value proves the whole chain ran.
"""
cfg = AmazonAnthropicClaudeMessagesConfig()
async def _hanging_stream():
yield {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 25, "output_tokens": 1}}}
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}}
await asyncio.Event().wait()
logging_obj = LiteLLMLoggingObj(
model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="chat",
start_time=datetime.now(),
litellm_call_id="test_bedrock_sse_wrapper_disconnect_logging",
function_id="test_bedrock_sse_wrapper_disconnect_logging",
)
wrapped = cfg.bedrock_sse_wrapper(_hanging_stream(), litellm_logging_obj=logging_obj, request_body={})
await wrapped.__anext__()
await wrapped.__anext__()
assert logging_obj.completion_start_time is None
await wrapped.aclose()
assert logging_obj.completion_start_time is not None

View file

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

View file

@ -25,6 +25,11 @@ from litellm.proxy._types import (
from litellm.proxy.common_utils.key_rotation_manager import KeyRotationManager
@pytest.fixture
def disable_audit_logging_for_mocked_key(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("litellm.store_audit_logs", False)
class TestKeyRotationManagerPassesKeyAlias:
"""
Regression tests to ensure KeyRotationManager passes key_alias
@ -155,7 +160,10 @@ class TestKeyRotationSecretNamingStability:
"""
@pytest.mark.asyncio
async def test_rotation_hook_uses_initial_secret_name_fallback(self):
async def test_rotation_hook_uses_initial_secret_name_fallback(
self,
disable_audit_logging_for_mocked_key,
):
"""
GIVEN: A key WITHOUT an alias (has an initial_secret_name based on token ID)
WHEN: The key is rotated
@ -206,7 +214,10 @@ class TestKeyRotationSecretNamingStability:
), f"Secret name drift! Expected {initial_secret_name}, got {call_kwargs['new_secret_name']}. This causes secret sprawl."
@pytest.mark.asyncio
async def test_rotation_hook_pre_rotation_alias_consistency(self):
async def test_rotation_hook_pre_rotation_alias_consistency(
self,
disable_audit_logging_for_mocked_key,
):
"""
GIVEN: A key WITH an alias
WHEN: The key is rotated

View file

@ -4,6 +4,7 @@ Tests for KeyManagementEventHooks.
Validates that email and secret manager operations are independent and non-blocking.
"""
import asyncio
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -155,6 +156,44 @@ class TestKeyManagementEventHooksIndependentOperations:
assert email_called["called"] is True
@pytest.mark.parametrize(
("premium_user", "expected_audit_log_calls"),
((True, 1), (False, 0)),
)
@pytest.mark.asyncio
async def test_key_generated_audit_log_uses_license_default(
monkeypatch: pytest.MonkeyPatch,
premium_user: bool,
expected_audit_log_calls: int,
):
from litellm.proxy._types import GenerateKeyRequest, GenerateKeyResponse, UserAPIKeyAuth
monkeypatch.setattr("litellm.store_audit_logs", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", premium_user)
monkeypatch.delenv("LITELLM_STORE_AUDIT_LOGS", raising=False)
response = GenerateKeyResponse(key="sk-test-key", token_id="token-123")
with (
patch(
"litellm.proxy.management_helpers.audit_logs.create_audit_log_for_update",
new_callable=AsyncMock,
) as mock_create_audit_log,
patch.object(
KeyManagementEventHooks,
"_store_virtual_key_in_secret_manager",
new_callable=AsyncMock,
),
):
await KeyManagementEventHooks.async_key_generated_hook(
data=GenerateKeyRequest(),
response=response,
user_api_key_dict=UserAPIKeyAuth(api_key="sk-admin-key", user_id="admin"),
)
await asyncio.sleep(0.01)
assert mock_create_audit_log.await_count == expected_audit_log_calls
class TestRotateVirtualKeyInSecretManager:
"""Tests for _rotate_virtual_key_in_secret_manager with team_id support."""

View file

@ -212,8 +212,9 @@ async def test_v1_key_generation_sends_email_when_send_invite_email_true():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.slack_alerting_instance = mock_slack_alerting
with patch.object(
KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email
with (
patch("litellm.store_audit_logs", False),
patch.object(KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email),
):
with patch(
"litellm.logging_callback_manager.get_custom_loggers_for_type",
@ -257,8 +258,9 @@ async def test_v1_key_generation_no_email_when_send_invite_email_false():
mock_proxy_logging_obj = MagicMock()
mock_proxy_logging_obj.slack_alerting_instance = mock_slack_alerting
with patch.object(
KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email
with (
patch("litellm.store_audit_logs", False),
patch.object(KeyManagementEventHooks, "_send_key_created_email", mock_send_key_created_email),
):
with patch(
"litellm.logging_callback_manager.get_custom_loggers_for_type",

View file

@ -91,6 +91,8 @@ mock_prisma_client = MagicMock()
mock_prisma_client.db = MagicMock()
mock_prisma_client.db.litellm_teamtable = MagicMock()
mock_prisma_client.db.litellm_teamtable.update = AsyncMock()
mock_prisma_client.db.litellm_auditlog = MagicMock()
mock_prisma_client.db.litellm_auditlog.create = AsyncMock()
# Fixture to provide the mock prisma client
@ -103,6 +105,11 @@ def mock_db_client():
mock_prisma_client.reset_mock()
@pytest.fixture
def disable_audit_logging_for_mocked_team(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr("litellm.store_audit_logs", False)
# Fixture to provide a mock admin user auth object
@pytest.fixture
def mock_admin_auth():
@ -2060,7 +2067,9 @@ async def test_team_model_add_delete_refresh_team_cache(endpoint_name):
@pytest.mark.asyncio
async def test_update_team_team_member_budget_not_passed_to_db():
async def test_update_team_team_member_budget_not_passed_to_db(
disable_audit_logging_for_mocked_team,
):
"""
Test that 'team_member_budget' is never passed to prisma_client.db.litellm_teamtable.update
regardless of whether the value is set or None.
@ -2498,7 +2507,9 @@ async def test_upsert_team_member_budget_table_no_existing_budget():
@pytest.mark.asyncio
async def test_update_team_with_team_member_budget_duration():
async def test_update_team_with_team_member_budget_duration(
disable_audit_logging_for_mocked_team,
):
"""
Test that team/update endpoint properly handles team_member_budget_duration.
"""
@ -5171,7 +5182,9 @@ async def test_update_team_standalone_budget_raise_blocked_for_team_admin():
@pytest.mark.asyncio
async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin():
async def test_update_team_standalone_budget_raise_allowed_for_proxy_admin(
disable_audit_logging_for_mocked_team,
):
"""
Test that a proxy admin CAN raise a standalone team's budget on /team/update.
@ -5325,7 +5338,9 @@ async def test_update_team_standalone_budget_removal_blocked_for_team_admin():
@pytest.mark.asyncio
async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed():
async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed(
disable_audit_logging_for_mocked_team,
):
"""
When a team currently has NO cap (max_budget=None / unlimited), a team admin
setting a finite max_budget is a RESTRICTION, not a raise, and is
@ -5407,7 +5422,9 @@ async def test_update_team_standalone_uncapped_team_admin_sets_finite_allowed():
@pytest.mark.asyncio
async def test_update_team_standalone_unchanged_budget_allowed():
async def test_update_team_standalone_unchanged_budget_allowed(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for a standalone team does NOT compare against the
caller's personal max_budget when the budget is unchanged.
@ -5508,7 +5525,9 @@ async def test_update_team_standalone_unchanged_budget_allowed():
@pytest.mark.asyncio
async def test_update_team_standalone_lower_budget_allowed():
async def test_update_team_standalone_lower_budget_allowed(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for a standalone team allows lowering the budget
below the team's current value even when the new value still exceeds the
@ -5691,7 +5710,9 @@ async def test_update_team_org_scoped_budget_exceeds_org_limit():
@pytest.mark.asyncio
async def test_update_team_standalone_models_not_gated_by_user_limit():
async def test_update_team_standalone_models_not_gated_by_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for a standalone team does NOT gate the team's models
by the caller's personal allowed models.
@ -5775,7 +5796,9 @@ async def test_update_team_standalone_models_not_gated_by_user_limit():
@pytest.mark.asyncio
async def test_update_team_org_scoped_budget_bypasses_user_limit():
async def test_update_team_org_scoped_budget_bypasses_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team does NOT validate budget against user's personal max_budget.
@ -5890,7 +5913,9 @@ async def test_update_team_org_scoped_budget_bypasses_user_limit():
@pytest.mark.asyncio
async def test_update_team_org_scoped_models_bypasses_user_limit():
async def test_update_team_org_scoped_models_bypasses_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team does NOT validate models against user's personal models.
@ -6080,7 +6105,9 @@ async def test_update_team_org_scoped_models_not_in_org_models():
@pytest.mark.asyncio
async def test_update_team_org_scoped_models_with_all_proxy_models():
async def test_update_team_org_scoped_models_with_all_proxy_models(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team succeeds when organization has 'all-proxy-models'.
@ -6196,7 +6223,9 @@ async def test_update_team_org_scoped_models_with_all_proxy_models():
@pytest.mark.asyncio
async def test_update_team_tpm_limit_not_gated_by_user_limit():
async def test_update_team_tpm_limit_not_gated_by_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update does NOT gate the team's tpm_limit by the caller's
personal tpm_limit.
@ -6279,7 +6308,9 @@ async def test_update_team_tpm_limit_not_gated_by_user_limit():
@pytest.mark.asyncio
async def test_update_team_rpm_limit_not_gated_by_user_limit():
async def test_update_team_rpm_limit_not_gated_by_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update does NOT gate the team's rpm_limit by the caller's
personal rpm_limit.
@ -6795,7 +6826,9 @@ async def test_update_team_org_scoped_rpm_exceeds_org_limit():
@pytest.mark.asyncio
async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit():
async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit(
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update for an org-scoped team bypasses user's TPM/RPM limits.
@ -6905,7 +6938,9 @@ async def test_update_team_org_scoped_tpm_rpm_bypasses_user_limit():
@pytest.mark.asyncio
async def test_update_team_guardrails_with_org_id():
async def test_update_team_guardrails_with_org_id(
disable_audit_logging_for_mocked_team,
):
"""
Test that updating team guardrails works when team has an organization_id.
The fix ensures 'teams' field is included when fetching organization data.
@ -7242,7 +7277,10 @@ async def test_persist_deleted_team_records():
@pytest.mark.asyncio
async def test_delete_team_persists_deleted_teams(monkeypatch):
async def test_delete_team_persists_deleted_teams(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
from litellm.proxy._types import DeleteTeamRequest
mock_prisma_client = AsyncMock()
@ -7325,7 +7363,10 @@ async def test_delete_team_persists_deleted_teams(monkeypatch):
@pytest.mark.asyncio
async def test_delete_team_sweeps_references_outside_members_with_roles(monkeypatch):
async def test_delete_team_sweeps_references_outside_members_with_roles(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
Regression pin for LIT-5511: a deleted team stayed visible on user records.
@ -7431,7 +7472,10 @@ async def test_delete_team_sweeps_references_outside_members_with_roles(monkeypa
@pytest.mark.asyncio
async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes(monkeypatch):
async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
A virtual key scoped to the team is deleted from the db with the team, but auth resolves a
cached key object without re-reading the team, so leaving the cache entry behind lets that key
@ -7493,7 +7537,10 @@ async def test_delete_team_evicts_the_auth_cache_of_the_keys_it_deletes(monkeypa
@pytest.mark.asyncio
async def test_delete_team_failing_reconcile_sweep_cannot_strand_the_team_in_cache(monkeypatch):
async def test_delete_team_failing_reconcile_sweep_cannot_strand_the_team_in_cache(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
The reconcile sweep runs after the team row is committed deleted. If it ran before cache
eviction, a sweep failure would return an error with the team gone from the db but still
@ -7555,7 +7602,10 @@ async def test_delete_team_failing_reconcile_sweep_cannot_strand_the_team_in_cac
@pytest.mark.asyncio
async def test_delete_team_broadcasts_cache_invalidation_to_other_workers(monkeypatch):
async def test_delete_team_broadcasts_cache_invalidation_to_other_workers(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
Evicting locally only reaches the worker that handled the delete. Without the broadcast, every
other worker keeps serving the deleted team, and the deleted team's keys, out of its own
@ -7619,7 +7669,10 @@ async def test_delete_team_broadcasts_cache_invalidation_to_other_workers(monkey
@pytest.mark.asyncio
async def test_delete_team_survives_a_failing_cache_backend(monkeypatch):
async def test_delete_team_survives_a_failing_cache_backend(
monkeypatch,
disable_audit_logging_for_mocked_team,
):
"""
Cache eviction runs after the reference sweep has already committed, so a cache backend that
is unreachable must not abort the delete. If it did, `/team/delete` would fail with the team
@ -8088,6 +8141,7 @@ async def test_update_team_soft_budget_validation(
expected_soft_budget,
expected_max_budget,
error_message,
disable_audit_logging_for_mocked_team,
):
"""
Test soft_budget validation in /team/update endpoint.
@ -8498,7 +8552,11 @@ async def test_get_team_daily_activity_member_without_permission_filters_by_keys
@pytest.mark.asyncio
async def test_update_team_with_router_settings(mock_db_client, mock_admin_auth):
async def test_update_team_with_router_settings(
mock_db_client,
mock_admin_auth,
disable_audit_logging_for_mocked_team,
):
"""
Test that /team/update correctly handles router_settings by:
1. Accepting router_settings as a dict parameter
@ -11594,7 +11652,9 @@ async def test_update_team_output_token_estimate_lowered_rejected_for_team_admin
@pytest.mark.asyncio
async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edit():
async def test_update_team_output_token_estimate_unchanged_allows_team_admin_edit(
disable_audit_logging_for_mocked_team,
):
"""The team settings form resends every field it renders, so gating on
presence would break a team admin editing an unrelated setting."""
import contextlib
@ -11957,7 +12017,9 @@ class _FakeMirrorDb:
@pytest.mark.asyncio
async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directions():
async def test_update_team_syncs_access_group_assigned_team_ids_in_both_directions(
disable_audit_logging_for_mocked_team,
):
"""
A team-side edit of `access_group_ids` must be mirrored onto every affected access
group's `assigned_team_ids`, in one transaction, in both directions.
@ -12113,7 +12175,9 @@ async def test_sync_reads_the_committed_team_row_rather_than_the_callers_snapsho
@pytest.mark.asyncio
async def test_new_team_and_delete_team_both_drive_the_mirror():
async def test_new_team_and_delete_team_both_drive_the_mirror(
disable_audit_logging_for_mocked_team,
):
"""Every writer of `team.access_group_ids` has to reach the mirror, not just update.
These pin the wiring on the other two paths; the mirror's own behavior is covered above.

View file

@ -19,6 +19,7 @@ from litellm.proxy.management_helpers.audit_logs import (
_build_audit_log_payload,
_dispatch_audit_log_to_callbacks,
create_audit_log_for_update,
is_audit_logging_enabled,
)
from litellm.types.utils import StandardAuditLogPayload
@ -49,6 +50,34 @@ def _make_audit_log(
)
@pytest.mark.parametrize(
("premium_user", "configured_value", "environment_value", "expected"),
(
(True, None, None, True),
(True, False, None, False),
(True, None, "false", False),
(False, None, None, False),
(False, True, None, True),
(True, True, "false", True),
),
)
def test_is_audit_logging_enabled_precedence(
monkeypatch: pytest.MonkeyPatch,
premium_user: bool,
configured_value: bool | None,
environment_value: str | None,
expected: bool,
):
monkeypatch.setattr(litellm, "store_audit_logs", configured_value)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", premium_user)
if environment_value is None:
monkeypatch.delenv("LITELLM_STORE_AUDIT_LOGS", raising=False)
else:
monkeypatch.setenv("LITELLM_STORE_AUDIT_LOGS", environment_value)
assert is_audit_logging_enabled() is expected
class TestBuildAuditLogPayload:
def test_builds_correct_payload(self):
audit_log = _make_audit_log()
@ -185,12 +214,14 @@ class TestCreateAuditLogForUpdateWithCallbacks:
with (
patch("litellm.proxy.proxy_server.premium_user", False),
patch("litellm.store_audit_logs", True),
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
):
audit_log = _make_audit_log()
await create_audit_log_for_update(audit_log)
await asyncio.sleep(0.1)
mock_logger.async_log_audit_log_event.assert_not_called()
mock_prisma.db.litellm_auditlog.create.assert_not_called()
@pytest.mark.asyncio
async def test_no_dispatch_when_store_audit_logs_false(self):

View file

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

View file

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

View file

@ -210,8 +210,10 @@ async def test_rollup_prunes_stale_row_when_config_is_gone():
where = table.delete_many.await_args.kwargs["where"]
assert where["date"] == DAY.isoformat()
assert where["api_key"] == PTU_SENTINEL_API_KEY
# the row is garbage because this run did not refresh it, not because of a key list
# the row is garbage because this run did not refresh it, and it is reachable at all
# because the run scanned the deployment it belongs to
assert "lt" in where["updated_at"]
assert "model" not in where, "a database-only run has no reason to bound the sweep"
@pytest.mark.asyncio
@ -705,13 +707,20 @@ class _FakeSentinelTable:
async def delete_many(self, where):
self.delete_many_calls.append(where)
cutoff = where["updated_at"]["lt"]
# honouring "model" matters: a fake that ignored an unknown clause would delete
# the row the prune-scoping test exists to protect and still report a pass
allowed = where.get("model", {}).get("in")
doomed = [
k
for k, v in self.rows.items()
if k[1] == where["date"] and k[2] == where["api_key"] and v["updated_at"] < cutoff
if k[1] == where["date"]
and k[2] == where["api_key"]
and v["updated_at"] < cutoff
and (allowed is None or k[3] in allowed)
]
for k in doomed:
del self.rows[k]
return len(doomed)
async def find_many(self, where=None):
"""Read back sentinel rows the way prisma would, honouring api_key and a date range."""
@ -785,11 +794,11 @@ async def test_an_older_run_cannot_delete_a_newer_runs_row():
@pytest.mark.asyncio
async def test_a_later_clean_run_clears_the_row_the_race_left_behind():
"""The race can leave a charge for a since-removed deployment in place for a day; the
next run, seeing only the current config, must sweep it."""
"""The race can leave a charge for a no-longer-priced deployment in place for a day;
the next run, seeing only the current config, must sweep it."""
table = _FakeSentinelTable()
ptu = {"ptu_count": 10, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
stale_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-removed")
stale_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-retired")
table.rows[stale_key] = {
"ptu_flat_cost": 480.0,
"model_group": "retired",
@ -797,7 +806,14 @@ async def test_a_later_clean_run_clears_the_row_the_race_left_behind():
}
await run_ptu_flat_cost_rollup(
_prisma_for([_model_row(model_id="dep-live", model_info=ptu)], table), target_date=DAY
_prisma_for(
[
_model_row(model_id="dep-live", model_info=ptu),
_model_row(model_id="dep-retired", model_info={"team_id": "t"}),
],
table,
),
target_date=DAY,
)
assert stale_key not in table.rows
@ -1715,16 +1731,19 @@ async def test_a_run_holding_the_lock_still_prunes():
"""Losing the sweep entirely would leave stale charges forever, so the guarded path,
which is the normal one, keeps it."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-gone", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
table.seed("t", DAY, "dep-unpriced", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
prisma = _prisma_for(
[_model_row(model_id="dep-live", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"})],
[
_model_row(model_id="dep-live", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}),
_model_row(model_id="dep-unpriced", model_info={"team_id": "t"}),
],
table,
)
await run_scheduled_ptu_rollup(prisma, pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert table.delete_many_calls != []
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-gone") not in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-unpriced") not in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-live") in table.rows
@ -1737,7 +1756,7 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts():
just_written = datetime.now(timezone.utc) - timedelta(seconds=30)
table.seed("t", DAY, "dep-concurrent", 480.0, updated_at=just_written)
table.seed("t", DAY, "dep-stale", 480.0, updated_at=datetime.now(timezone.utc) - timedelta(hours=6))
prisma = _prisma_for([], table)
prisma = _prisma_for([_model_row(model_id="dep-concurrent"), _model_row(model_id="dep-stale")], table)
await run_ptu_flat_cost_rollup(prisma, target_date=DAY)
@ -1747,6 +1766,127 @@ async def test_the_prune_cutoff_allows_for_clock_skew_between_hosts():
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-stale") not in table.rows
@pytest.mark.asyncio
async def test_a_run_pricing_config_cannot_prune_a_row_it_did_not_scan(monkeypatch):
"""Staleness alone stops being evidence once two hosts hold different configuration: a
row this run never considered belongs to a deployment another host is pricing from its
own file, and sweeping it drops that charge."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-elsewhere", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
entry = _router_entry(model_id="cfg-here", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-elsewhere") in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-here") in table.rows
assert table.delete_many_calls[-1]["model"]["in"] == ("cfg-here",)
@pytest.mark.asyncio
async def test_a_deployment_deleted_from_the_table_keeps_the_day_it_was_charged(monkeypatch):
"""The accepted cost of bounding the prune, driven through the sequence that produces
it: charge the day while the deployment exists, remove it, run the day again. Nothing
scans it now, so nothing may judge its row, and the amount it was billed stands."""
table = _FakeSentinelTable()
ptu = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
live_row = _model_row(model_id="dep-live", model_info=ptu)
doomed_row = _model_row(model_id="dep-doomed", model_info=ptu)
charged_key = ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-doomed")
monkeypatch.setattr(
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="cfg", model_info=dict(ptu)))
)
await run_scheduled_ptu_rollup(
_prisma_for([live_row, doomed_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
)
billed = table.rows[charged_key]["ptu_flat_cost"]
table.rows[charged_key]["updated_at"] = datetime(2020, 1, 1, tzinfo=timezone.utc)
await run_scheduled_ptu_rollup(
_prisma_for([live_row], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
)
assert table.rows[charged_key]["ptu_flat_cost"] == billed
assert "dep-doomed" not in table.delete_many_calls[-1]["model"]["in"]
@pytest.mark.asyncio
async def test_a_database_only_run_sweeps_exactly_as_it_did_before():
"""The bound exists for charges another host declares. A deployment nobody declares any
more still has its leftover row swept, which is what the table-only sweep always did."""
table = _FakeSentinelTable()
table.seed("t", DAY, "dep-gone", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
prisma = _prisma_for(
[_model_row(model_id="dep-live", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"})],
table,
)
await run_scheduled_ptu_rollup(prisma, pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-gone") not in table.rows
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "dep-live") in table.rows
@pytest.mark.asyncio
async def test_every_deployment_that_prices_is_inside_the_set_that_bounds_the_prune():
"""The bound has to be a superset of what the same run wrote, or a run's own charge
could fall outside its own delete filter and never be reconciled."""
table = _FakeSentinelTable()
prisma = _prisma_for(
[
_model_row(model_id="dep-a", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}),
_model_row(model_id="dep-b", model_info={"ptu_count": 9, "cost_per_ptu_per_hour": 1.0, "team_id": "u"}),
_model_row(model_id="dep-unpriced", model_info={"team_id": "t"}),
],
table,
)
loaded = await ptu_rollup._load_ptu_models(prisma)
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
assert loaded.scanned_ids == {"dep-a", "dep-b", "dep-unpriced"}
@pytest.mark.asyncio
async def test_a_priced_deployment_is_in_the_bound_even_with_an_id_the_scan_skips():
"""The bound is built by construction rather than by coincidence. The row scan drops a
falsy id while the parser still prices one, and a charge outside its own run's delete
filter could never be reconciled by any later run."""
prisma = _prisma_for(
[_model_row(model_id="", model_info={"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"})],
_FakeSentinelTable(),
)
loaded = await ptu_rollup._load_ptu_models(prisma)
assert {model.model_id for model in loaded.models} <= loaded.scanned_ids
@pytest.mark.asyncio
async def test_the_prune_splits_the_id_set_across_statements(monkeypatch):
"""Every id is one bind variable and the server refuses a statement carrying more than
32767, so a proxy with that many deployments would fail the prune outright, and with it
the rest of the scheduled run."""
monkeypatch.setattr(ptu_rollup, "_PRUNE_ID_CHUNK_SIZE", 2)
table = _FakeSentinelTable()
ptu = {"ptu_count": 5, "cost_per_ptu_per_hour": 2.0, "team_id": "t"}
deployments = [_model_row(model_id=f"dep-{n}", model_info=ptu) for n in range(4)]
monkeypatch.setattr(
ptu_rollup, "_running_router", lambda: _router_holding(_router_entry(model_id="dep-4", model_info=dict(ptu)))
)
table.seed("t", DAY, "dep-3", 480.0, updated_at=datetime(2020, 1, 1, tzinfo=timezone.utc))
await run_scheduled_ptu_rollup(
_prisma_for(deployments, table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY
)
chunks = [call["model"]["in"] for call in table.delete_many_calls]
assert len(chunks) == 3
assert all(len(chunk) <= 2 for chunk in chunks)
assert sorted(i for chunk in chunks for i in chunk) == [f"dep-{n}" for n in range(5)]
@pytest.mark.asyncio
async def test_scheduled_rollup_writes_nothing_when_ptu_attribution_is_disabled(monkeypatch):
"""Startup already skips scheduling the cron, so this guards the function itself: a
@ -1760,3 +1900,148 @@ async def test_scheduled_rollup_writes_nothing_when_ptu_attribution_is_disabled(
assert result is None
assert table.rows == {}
assert table.upsert_keys == []
# --- config.yaml deployments reach the rollup through the router ----------------
def _router_holding(*entries):
"""A stand-in for the proxy's router, carrying whatever model_list is passed."""
return types.SimpleNamespace(model_list=list(entries))
@pytest.mark.asyncio
async def test_a_config_declared_deployment_is_priced(monkeypatch):
"""The whole point. A PTU deployment the proxy only knows from config.yaml is not in
LiteLLM_ProxyModelTable, so a DB-only scan bills the provider's reservation to nobody."""
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert [(m.model_id, m.model_name, m.team_id) for m in loaded.models] == [("cfg-1", "gpt-4o-ptu", "t")]
assert "cfg-1" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_a_database_backed_router_entry_is_not_counted_twice(monkeypatch):
"""Every deployment loaded from the table is also in the router, flagged db_model. Pricing
both copies would write two charges for one reservation."""
row = _model_row(model_id="db-1", model_info=dict(_VALID_PTU))
mirrored = _router_entry(model_id="db-1", model_info={**_VALID_PTU, "db_model": True})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(mirrored))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_router_entry_sharing_an_id_with_the_table_is_priced_once(monkeypatch):
"""db_model is data the router carries rather than something this module controls, so the
id anti-join is what actually maps onto the failure: two charges under one id."""
row = _model_row(model_id="both-1", model_info=dict(_VALID_PTU))
unflagged = _router_entry(model_id="both-1", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(unflagged))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([row], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["both-1"]
@pytest.mark.asyncio
async def test_a_client_credential_clone_is_not_priced(monkeypatch):
"""Supplying an api_key on a request mints a clone of the deployment under a fresh id,
carrying the source's PTU config. Pricing it bills one reservation per distinct caller key."""
source = _router_entry(model_id="cfg-1", model_info=dict(_VALID_PTU))
clone = _router_entry(model_id="cfg-1-clone", model_info={**_VALID_PTU, "original_model_id": "cfg-1"})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(source, clone))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert [m.model_id for m in loaded.models] == ["cfg-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_without_ptu_config_is_scanned_but_not_priced(monkeypatch):
"""It has to stay in the scanned set or its leftover sentinel rows become unprunable."""
entry = _router_entry(model_id="cfg-plain", model_info={"team_id": "t"})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert loaded.models == ()
assert "cfg-plain" in loaded.scanned_ids
@pytest.mark.asyncio
async def test_no_router_in_the_process_prices_the_database_alone(monkeypatch):
"""The rollup is importable and callable outside a running proxy."""
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: None)
loaded = await ptu_rollup._load_ptu_models(
_prisma_for([_model_row(model_id="db-1", model_info=dict(_VALID_PTU))], _FakeSentinelTable())
)
assert [m.model_id for m in loaded.models] == ["db-1"]
@pytest.mark.asyncio
async def test_a_config_deployment_is_charged_end_to_end(monkeypatch):
"""Through the scheduled entry point, so the charge lands in a sentinel row rather than
stopping at the loader."""
table = _FakeSentinelTable()
entry = _router_entry(model_id="cfg-1", model_name="gpt-4o-ptu", model_info=dict(_VALID_PTU))
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(entry))
await run_scheduled_ptu_rollup(_prisma_for([], table), pod_lock_manager=_pod_lock(acquired=True), target_date=DAY)
assert ("t", DAY.isoformat(), PTU_SENTINEL_API_KEY, "cfg-1") in table.rows
@pytest.mark.asyncio
async def test_a_stale_database_backed_router_entry_is_not_treated_as_config(monkeypatch):
"""The reconcile can leave a deployment on the router after its row is gone. The id
anti-join cannot see that one, so the flag is what keeps it from being priced as though
config.yaml had declared it."""
stale = _router_entry(model_id="db-gone", model_info={**_VALID_PTU, "db_model": True})
monkeypatch.setattr(ptu_rollup, "_running_router", lambda: _router_holding(stale))
loaded = await ptu_rollup._load_ptu_models(_prisma_for([], _FakeSentinelTable()))
assert loaded.models == ()
def test_the_router_lookup_reads_the_proxys_own_global():
"""Every other config test replaces this helper, so without one test driving the real
body a typo in the module path or the attribute name leaves the whole feature dead in
production with the suite still green."""
import sys
import types as _types
assert ptu_rollup._running_router() is None or "litellm.proxy.proxy_server" in sys.modules
sentinel = object()
stub = _types.SimpleNamespace(llm_router=sentinel)
real = sys.modules.get("litellm.proxy.proxy_server")
sys.modules["litellm.proxy.proxy_server"] = stub
try:
assert ptu_rollup._running_router() is sentinel
del stub.llm_router
assert ptu_rollup._running_router() is None
finally:
if real is None:
del sys.modules["litellm.proxy.proxy_server"]
else:
sys.modules["litellm.proxy.proxy_server"] = real
def test_the_router_lookup_returns_none_outside_a_proxy():
import sys
real = sys.modules.pop("litellm.proxy.proxy_server", None)
try:
assert ptu_rollup._running_router() is None
finally:
if real is not None:
sys.modules["litellm.proxy.proxy_server"] = real

View file

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

View file

@ -3164,3 +3164,80 @@ def test_batch_cost_row_id_is_stable_across_repeated_accounting():
]
assert ids[0] == ids[1] == "batch_same_batch_cost"
def _make_failed_request_standard_logging_payload() -> StandardLoggingPayload:
base: Final = _make_standard_logging_payload_with_usage_object(usage_object={})
return cast(
StandardLoggingPayload,
{
**base,
"status": "failure",
"call_type": "aresponses",
"model_id": "mid-123",
"model_group": "group-x",
"api_base": "https://api.openai.com/v1/responses",
"custom_llm_provider": "openai",
},
)
def test_get_logging_payload_failed_request_falls_back_to_standard_logging_payload():
"""Failed-request kwargs from the proxy failure hook carry no deployment info
(LIT-5795), so the attribution columns must come from the failure-time
standard_logging_object."""
payload = get_logging_payload(
kwargs={
"model": "group-x",
"litellm_params": {"metadata": {"user_api_key": "test-key", "status": "failure"}},
"standard_logging_object": _make_failed_request_standard_logging_payload(),
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["model_id"] == "mid-123"
assert payload["model_group"] == "group-x"
assert payload["api_base"] == "https://api.openai.com/v1/responses"
assert payload["custom_llm_provider"] == "openai"
def test_get_logging_payload_request_kwargs_win_over_standard_logging_payload():
payload = get_logging_payload(
kwargs={
"model": "group-y",
"custom_llm_provider": "anthropic",
"litellm_params": {
"api_base": "https://kwargs.example.com",
"metadata": {
"user_api_key": "test-key",
"model_group": "kwargs-group",
"model_info": {"id": "kwargs-mid"},
},
},
"standard_logging_object": _make_failed_request_standard_logging_payload(),
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["model_id"] == "kwargs-mid"
assert payload["model_group"] == "kwargs-group"
assert payload["api_base"] == "https://kwargs.example.com"
assert payload["custom_llm_provider"] == "anthropic"
def test_get_logging_payload_failed_request_without_standard_logging_payload_leaves_fields_empty():
payload = get_logging_payload(
kwargs={
"model": "group-x",
"litellm_params": {"metadata": {"user_api_key": "test-key", "status": "failure"}},
},
response_obj={},
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
assert payload["model_id"] == ""
assert payload["model_group"] == ""
assert payload["api_base"] == ""
assert payload["custom_llm_provider"] == ""

View file

@ -468,6 +468,23 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend:
assert request_data["response_cost"] == 3.5e-05
assert "litellm_logging_obj" not in request_data
@pytest.mark.asyncio
async def test_recovered_usage_without_cost_clobbers_client_cost_with_zero(self):
from litellm.types.utils import Usage
recovered_usage = Usage(prompt_tokens=30, completion_tokens=1, total_tokens=31)
logging_obj = MagicMock()
logging_obj.model_call_details = {"combined_usage_object": recovered_usage}
request_data = {
"litellm_logging_obj": logging_obj,
"response_cost": 999.0,
"metadata": {},
}
await self._run(request_data)
assert request_data["combined_usage_object"] is recovered_usage
assert request_data["response_cost"] == 0.0
@pytest.mark.asyncio
async def test_no_recovered_usage_is_noop(self):
logging_obj = MagicMock()
@ -478,6 +495,111 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend:
assert "response_cost" not in request_data
class TestPostCallFailureHookLiftsStandardLoggingObject:
"""Failure callbacks read standard_logging_object from request_data, but
post_call_failure_hook pops litellm_logging_obj before they run. The hook
must lift the logging obj's standard_logging_object onto request_data so
failed-request spend logs keep deployment attribution (LIT-5795).
"""
async def _run(self, request_data):
from unittest.mock import AsyncMock, patch
from litellm.proxy._types import UserAPIKeyAuth
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = []
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
await proxy_logging_obj.post_call_failure_hook(
request_data=request_data,
original_exception=Exception("boom"),
user_api_key_dict=UserAPIKeyAuth(),
)
@pytest.mark.asyncio
async def test_lifts_standard_logging_object(self):
sl_object = {"model_id": "mid-123", "model_group": "group-x"}
logging_obj = MagicMock()
logging_obj.model_call_details = {"standard_logging_object": sl_object}
request_data = {"litellm_logging_obj": logging_obj, "metadata": {}}
await self._run(request_data)
assert request_data["standard_logging_object"] is sl_object
assert "litellm_logging_obj" not in request_data
@pytest.mark.asyncio
async def test_logging_obj_value_overwrites_preexisting_key(self):
authoritative = {"model_id": "from-logging-obj"}
logging_obj = MagicMock()
logging_obj.model_call_details = {"standard_logging_object": authoritative}
request_data = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": {"model_id": "client-injected"},
"metadata": {},
}
await self._run(request_data)
assert request_data["standard_logging_object"] is authoritative
@pytest.mark.asyncio
async def test_client_supplied_key_is_stripped_when_logging_obj_supplies_none(self):
spoofed = {"model_id": "client-injected"}
request_data = {"standard_logging_object": spoofed, "metadata": {}}
await self._run(request_data)
assert "standard_logging_object" not in request_data
logging_obj = MagicMock()
logging_obj.model_call_details = {}
request_data_with_obj = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": spoofed,
"metadata": {},
}
await self._run(request_data_with_obj)
assert "standard_logging_object" not in request_data_with_obj
@pytest.mark.asyncio
async def test_pass_through_failure_never_relifts_client_supplied_key(self):
from datetime import datetime
from unittest.mock import AsyncMock, patch
from fastapi import HTTPException
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.proxy._types import UserAPIKeyAuth
logging_obj = Logging(
model="claude-haiku-4-5",
messages=[{"role": "user", "content": "hi"}],
stream=False,
call_type="pass_through_endpoint",
start_time=datetime.now(),
litellm_call_id="test-call-id",
function_id="test-function-id",
)
request_data = {
"litellm_logging_obj": logging_obj,
"standard_logging_object": {"model_id": "client-injected"},
"metadata": {},
}
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging_obj.alert_types = []
with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()):
await proxy_logging_obj.post_call_failure_hook(
request_data=request_data,
original_exception=HTTPException(status_code=401, detail="unauthorized"),
user_api_key_dict=UserAPIKeyAuth(request_route="/v1/chat/completions"),
)
assert "standard_logging_object" not in request_data
assert "standard_logging_object" not in logging_obj.model_call_details
@pytest.mark.asyncio
async def test_no_standard_logging_object_is_noop(self):
logging_obj = MagicMock()
logging_obj.model_call_details = {}
request_data = {"litellm_logging_obj": logging_obj, "metadata": {}}
await self._run(request_data)
assert "standard_logging_object" not in request_data
class TestPostCallFailureHookEstimatesDispatchedInputTokens:
"""A non-stream request that failed after dispatch (timeout, provider
error) consumed provider-billed input tokens but recovered no usage.

View file

@ -2166,6 +2166,38 @@ class TestRouterPreRoutingAliasOverrides:
assert request_kwargs["drop_params"] is True
assert request_kwargs["cache_control_injection_points"] == [{"location": "message", "role": "system"}]
@pytest.mark.asyncio
async def test_tier_litellm_params_are_applied_before_deployment_selection(self):
router = Router(
model_list=[
{
"model_name": "smart-router",
"litellm_params": {
"model": "auto_router/complexity_router",
"complexity_router_config": {
"tiers": {
"SIMPLE": {
"model_name": "gpt-4o-mini",
"litellm_params": {"reasoning_effort": "xhigh"},
}
}
},
},
},
{"model_name": "gpt-4o-mini", "litellm_params": {"model": "openai/gpt-4o-mini"}},
]
)
request_kwargs: Dict = {"reasoning_effort": "low"}
deployment = await router.async_get_available_deployment(
model="smart-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hi"}],
)
assert deployment["model_name"] == "gpt-4o-mini"
assert request_kwargs["reasoning_effort"] == "xhigh"
@pytest.mark.asyncio
async def test_alias_custom_pricing_is_not_applied_to_request_kwargs(self):
"""Custom pricing on the alias prices the alias, not the tier deployment
@ -3851,7 +3883,7 @@ class TestSessionAffinity:
cache.async_set_cache.assert_called_once()
call_kwargs = cache.async_set_cache.call_args.kwargs
assert call_kwargs["ttl"] == 120
assert call_kwargs["value"] == "gpt-4o-mini"
assert call_kwargs["value"] == {"model": "gpt-4o-mini", "tier": "SIMPLE"}
@pytest.mark.asyncio
async def test_ttl_refreshed_on_cache_hit(self, mock_router_instance, basic_config):
@ -3875,7 +3907,7 @@ class TestSessionAffinity:
assert result.model == "o1-preview"
cache.async_set_cache.assert_called_once()
call_kwargs = cache.async_set_cache.call_args.kwargs
assert call_kwargs["value"] == "o1-preview"
assert call_kwargs["value"] == {"model": "o1-preview", "tier": "REASONING"}
assert call_kwargs["ttl"] == 90
@pytest.mark.asyncio
@ -5545,12 +5577,14 @@ class TestRedactedLoggingDropsPromptText:
"tier_boundaries": {"simple_medium": 0.15, "medium_complex": 0.35, "complex_reasoning": 0.6},
"classifier_model": "claude-haiku",
"escalated": True,
"tier_litellm_params": {"reasoning_effort": "xhigh"},
"signals": ["code (python)"],
"matched_keyword": "deploy to k8s",
"escalation_keyword": "LITELLM ESCALATE",
}
kept = Router._redact_prompt_text_if_needed(request_kwargs={}, routing_decision=full)
assert set(full) - set(kept) == {"signals", "matched_keyword", "escalation_keyword"}
assert kept["tier_litellm_params"] == {"reasoning_effort": "xhigh"}
@pytest.mark.asyncio
async def test_redaction_via_request_header_is_honored(self):
@ -7273,8 +7307,6 @@ class TestClassificationRubrics:
},
)
assert config.classifier_llm_config.system_prompt == "Grade the data sensitivity of the request."
def _custom_tier_config(**overrides) -> Dict:
"""A valid operator-defined tier set: two built-in names plus one custom tier."""
return {
@ -8123,3 +8155,243 @@ class TestPlanModeTierFloor:
assert result.model == "gpt-4o"
assert result.routing_decision is not None
assert result.routing_decision["tier"] == "MEDIUM"
def test_tier_model_params_are_normalized_without_changing_model_pools():
config = ComplexityRouterConfig(
tiers={
"SIMPLE": "mini",
"REASONING": [
{"model_name": "opus", "litellm_params": {"reasoning_effort": "xhigh"}},
"abc",
],
}
)
assert config.tiers == {"SIMPLE": "mini", "REASONING": ["opus", "abc"]}
assert config.tier_model_configs["REASONING"][0].litellm_params == {"reasoning_effort": "xhigh"}
rebuilt = ComplexityRouterConfig.model_validate(config.model_dump())
assert rebuilt.tier_model_configs["REASONING"][0].litellm_params == {"reasoning_effort": "xhigh"}
def test_tier_model_params_accept_a_single_object():
config = ComplexityRouterConfig(
tiers={"REASONING": {"model_name": "opus", "litellm_params": {"thinking": {"type": "enabled"}}}}
)
assert config.tiers == {"REASONING": "opus"}
assert config.tier_model_configs["REASONING"][0].model_name == "opus"
@pytest.mark.parametrize(
"tiers",
[
{"REASONING": [{"litellm_params": {"reasoning_effort": "xhigh"}}]},
],
)
def test_tier_model_params_reject_malformed_entries(tiers):
with pytest.raises(ValidationError):
ComplexityRouterConfig(tiers=tiers)
def test_tier_model_params_reject_duplicate_models():
with pytest.raises(ValidationError, match="duplicate model_name"):
ComplexityRouterConfig(
tiers={
"REASONING": [
{"model_name": "opus", "litellm_params": {"reasoning_effort": "xhigh"}},
{"model_name": "opus", "litellm_params": {"reasoning_effort": "low"}},
]
}
)
def test_non_adaptive_empty_tier_pool_remains_valid():
config = ComplexityRouterConfig(tiers={"SIMPLE": []})
assert config.tiers == {"SIMPLE": []}
def test_adaptive_empty_tier_pool_is_rejected():
with pytest.raises(ValidationError, match="adaptive=True"):
ComplexityRouterConfig(adaptive=True, tiers={"SIMPLE": []})
def test_tier_model_params_are_used_by_pools_and_savings_baseline(mock_router_instance):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": "mini",
"REASONING": [{"model_name": "opus", "litellm_params": {"reasoning_effort": "xhigh"}}, "abc"],
}
},
)
assert router._tier_pools() == {"SIMPLE": ["mini"], "REASONING": ["opus", "abc"]}
assert router._hardest_tier_models() == ("opus", "abc")
assert router._litellm_params_for_model(ComplexityTier.REASONING, "opus") == {"reasoning_effort": "xhigh"}
@pytest.mark.asyncio
async def test_tier_model_params_reach_the_hook_response_and_override_client_values(mock_router_instance):
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"REASONING": {
"model_name": "opus",
"litellm_params": {"reasoning_effort": "xhigh", "max_tokens": 512},
}
},
"keyword_tier_rules": [{"keywords": ["reason carefully"], "tier": "REASONING"}],
},
)
request_kwargs = {"reasoning_effort": "low", "metadata": {}}
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "reason carefully about this"}],
)
assert response is not None
assert response.litellm_params == {"reasoning_effort": "xhigh", "max_tokens": 512}
assert response.routing_decision is not None
assert response.routing_decision["tier_litellm_params"] == response.litellm_params
@pytest.mark.asyncio
@pytest.mark.parametrize("route", ["classification", "keyword", "session"])
async def test_tier_params_mask_credentials_in_routing_decision(route, mock_router_instance):
params = {"reasoning_effort": "xhigh", "api_key": "secret-tier-key"}
config = {
"tiers": {
tier.value: {"model_name": "opus", "litellm_params": params}
for tier in ComplexityTier
},
"keyword_tier_rules": [{"keywords": ["reason carefully"], "tier": "REASONING"}]
if route == "keyword"
else None,
"session_affinity": route == "session",
}
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config=config,
)
request_kwargs = {"metadata": {"session_id": "masked-params-session"}}
if route == "session":
mock_router_instance.cache = DualCache()
await mock_router_instance.cache.async_set_cache(
key=router._get_session_affinity_cache_key("masked-params-session", request_kwargs),
value={"model": "opus", "tier": "REASONING"},
)
message = "reason carefully about this" if route == "keyword" else "hello"
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": message}],
)
assert response is not None
assert response.litellm_params == params
assert response.routing_decision is not None
assert response.routing_decision["tier_litellm_params"] == {
"reasoning_effort": "xhigh",
"api_key": "secr*******-key",
}
@pytest.mark.asyncio
async def test_session_pin_outside_tiers_does_not_inherit_medium_params(mock_router_instance):
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": "mini",
"MEDIUM": {"model_name": "medium", "litellm_params": {"reasoning_effort": "low"}},
},
"session_affinity": True,
"default_model": "orphan",
},
)
request_kwargs = {"metadata": {"session_id": "orphan-session"}}
await mock_router_instance.cache.async_set_cache(
key=router._get_session_affinity_cache_key("orphan-session", request_kwargs),
value="orphan",
)
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.model == "orphan"
assert response.litellm_params == {}
@pytest.mark.asyncio
async def test_session_pin_uses_recorded_tier_when_model_is_in_multiple_tiers(mock_router_instance):
mock_router_instance.cache = DualCache()
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": {"model_name": "shared", "litellm_params": {"reasoning_effort": "low"}},
"REASONING": {"model_name": "shared", "litellm_params": {"reasoning_effort": "xhigh"}},
},
"session_affinity": True,
},
)
request_kwargs = {"metadata": {"session_id": "shared-session"}}
await mock_router_instance.cache.async_set_cache(
key=router._get_session_affinity_cache_key("shared-session", request_kwargs),
value={"model": "shared", "tier": "SIMPLE"},
)
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.litellm_params == {"reasoning_effort": "low"}
assert response.routing_decision is not None
assert response.routing_decision["tier"] == "SIMPLE"
@pytest.mark.asyncio
async def test_session_pin_survives_json_list_round_trip(mock_router_instance):
cache = AsyncMock()
cache.async_get_cache = AsyncMock(return_value=["shared", "SIMPLE"])
mock_router_instance.cache = cache
router = ComplexityRouter(
model_name="test-router",
litellm_router_instance=mock_router_instance,
complexity_router_config={
"tiers": {
"SIMPLE": {"model_name": "shared", "litellm_params": {"reasoning_effort": "low"}},
"REASONING": {"model_name": "shared", "litellm_params": {"reasoning_effort": "xhigh"}},
},
"session_affinity": True,
},
)
request_kwargs = {"metadata": {"session_id": "json-round-trip-session"}}
response = await router.async_pre_routing_hook(
model="test-router",
request_kwargs=request_kwargs,
messages=[{"role": "user", "content": "hello"}],
)
assert response is not None
assert response.model == "shared"
assert response.litellm_params == {"reasoning_effort": "low"}
assert cache.async_set_cache.call_args.kwargs["value"] == {"model": "shared", "tier": "SIMPLE"}

View file

@ -7,7 +7,14 @@ from unittest.mock import Mock, patch
import pytest
from litellm.secret_managers.main import get_secret, normalize_nonempty_secret_str
import litellm
from litellm.integrations.custom_secret_manager import CustomSecretManager
from litellm.secret_managers.main import (
get_secret,
normalize_nonempty_secret_str,
secret_manager_would_be_consulted,
)
from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem
# Set up logging for debugging
logging.basicConfig(level=logging.DEBUG)
@ -364,3 +371,63 @@ def test_unsupported_oidc_provider():
)
def test_normalize_nonempty_secret_str(raw, expected):
assert normalize_nonempty_secret_str(raw) == expected
class _SpySecretManager(CustomSecretManager):
"""Records every name the manager is actually asked for."""
def __init__(self, asked):
self.asked = asked
def sync_read_secret(self, secret_name, optional_params=None, timeout=None, **kwargs):
self.asked.append(secret_name)
return "a-value"
async def async_read_secret(self, secret_name, optional_params=None, timeout=None, **kwargs):
self.asked.append(secret_name)
return "a-value"
@pytest.mark.parametrize(
("access_mode", "hosted_keys", "secret_name", "expected"),
[
("read_only", None, "ANY_NAME", True),
("read_only", ["ALLOWED"], "ALLOWED", True),
("read_only", ["ALLOWED"], "NOT_ALLOWED", False),
("read_and_write", ["ALLOWED"], "ALLOWED", True),
("write_only", None, "ANY_NAME", False),
("write_only", ["ALLOWED"], "ALLOWED", False),
],
)
def test_secret_manager_would_be_consulted_matches_get_secret(
monkeypatch, access_mode, hosted_keys, secret_name, expected
):
"""The predicate must agree with what get_secret actually does, not with a reading of it.
Callers use it to tell "the manager does not have this key" apart from "the manager was
never asked", so a predicate that drifts from get_secret's gating makes them state a
lookup that never happened.
"""
asked = []
monkeypatch.setattr(litellm, "secret_manager_client", _SpySecretManager(asked))
monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM)
monkeypatch.setattr(
litellm,
"_key_management_settings",
KeyManagementSettings(access_mode=access_mode, hosted_keys=hosted_keys),
)
monkeypatch.delenv(secret_name, raising=False)
predicted = secret_manager_would_be_consulted(f"os.environ/{secret_name}")
get_secret(f"os.environ/{secret_name}")
assert {"predicted": predicted, "actually_consulted": bool(asked)} == {
"predicted": expected,
"actually_consulted": expected,
}
def test_secret_manager_would_be_consulted_is_false_without_a_client(monkeypatch):
monkeypatch.setattr(litellm, "secret_manager_client", None)
assert secret_manager_would_be_consulted("os.environ/ANY_NAME") is False

View file

@ -7,6 +7,9 @@ category, it prints `run` or `skip`. The gating contract we lock in here:
* docs-only changes (``*.md``, ``*.mdx``, ``docs/``) run nothing
* client-only changes (``ui/``) run client jobs but skip backend jobs
* any backend change runs both client and backend jobs
* ``ui`` tracks ``ui/`` plus CI config, so a backend-only change skips it
where ``client`` would still run, while a change to the workflows that
define the dashboard jobs still exercises them
If this logic silently regresses, real test jobs get skipped, so these cases
are the guardrail against that.
@ -40,6 +43,7 @@ def classify(category: str, changed: list[str]) -> str:
DOCS = ["README.md", "docs/my_website/index.mdx", "litellm/anywhere.md"]
CLIENT = ["ui/litellm-dashboard/src/App.tsx"]
BACKEND = ["litellm/main.py"]
CI = [".github/workflows/test-litellm-ui-unit.yml"]
@pytest.mark.parametrize(
@ -48,20 +52,27 @@ BACKEND = ["litellm/main.py"]
# docs-only: skip everything
("backend", DOCS, "skip"),
("client", DOCS, "skip"),
("ui", DOCS, "skip"),
("backend", [], "skip"),
("client", [], "skip"),
("ui", [], "skip"),
# client-only: backend skips, client runs
("backend", CLIENT, "skip"),
("client", CLIENT, "run"),
("ui", CLIENT, "run"),
("backend", CLIENT + DOCS, "skip"),
("client", CLIENT + DOCS, "run"),
("ui", CLIENT + DOCS, "run"),
# any backend change: both run ("backend runs both")
("backend", BACKEND, "run"),
("client", BACKEND, "run"),
("ui", BACKEND, "skip"),
("backend", BACKEND + DOCS, "run"),
("client", BACKEND + DOCS, "run"),
("ui", BACKEND + DOCS, "skip"),
("backend", BACKEND + CLIENT, "run"),
("client", BACKEND + CLIENT, "run"),
("ui", BACKEND + CLIENT, "run"),
],
)
def test_classify_decisions(category: str, changed: list[str], expected: str) -> None:
@ -71,6 +82,31 @@ def test_classify_decisions(category: str, changed: list[str], expected: str) ->
def test_markdown_under_ui_counts_as_client_not_docs() -> None:
assert classify("client", ["ui/litellm-dashboard/README.md"]) == "run"
assert classify("backend", ["ui/litellm-dashboard/README.md"]) == "skip"
assert classify("ui", ["ui/litellm-dashboard/README.md"]) == "run"
def test_ci_config_changes_reach_every_category() -> None:
"""A workflow edit has to exercise the jobs it defines, otherwise the change
ships unvalidated: the dashboard jobs would skip on the very pull request
that rewrites them."""
assert classify("ui", CI) == "run"
assert classify("backend", CI) == "run"
assert classify("client", CI) == "run"
def test_markdown_under_dot_github_is_still_docs() -> None:
"""`.github/**` counting as CI config must not drag the pull request template
and other markdown back into running the full suite."""
assert classify("ui", [".github/pull_request_template.md"]) == "skip"
assert classify("backend", [".github/pull_request_template.md"]) == "skip"
def test_ui_and_client_diverge_on_a_backend_only_change() -> None:
"""`client` gates CircleCI's dashboard end-to-end jobs, which drive a real
proxy and so must run on backend changes. `ui` gates the dashboard build and
its unit tests, which cannot see the backend at all."""
assert classify("client", BACKEND) == "run"
assert classify("ui", BACKEND) == "skip"
def test_non_docs_directory_with_docs_in_name_is_backend() -> None:

View file

@ -1742,6 +1742,81 @@ def test_azure_ai_cache_cost_calculation():
), f"Output cost mismatch: got {output_cost}, expected {expected_output_cost}"
def test_vertex_regional_deployment_costs_uplift_over_global(monkeypatch):
"""
Regression for https://github.com/BerriAI/litellm/issues/34393: two Vertex
deployments differing only in vertex_location must not price identically.
Google bills non-global endpoints at 1.1x for regional-pricing models, so the
regional request costs 1.1x the global one for the exact same usage, through
both vertex cost routes (Claude via cost_per_token, Gemini via
cost_per_character's token fallback).
"""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
usage = Usage(prompt_tokens=15, completion_tokens=5, total_tokens=20)
for model in ("claude-haiku-4-5@20251001", "gemini-3.5-flash"):
global_prompt, global_completion = cost_per_token(
model=model,
custom_llm_provider="vertex_ai",
usage_object=usage,
vertex_location="global",
)
regional_prompt, regional_completion = cost_per_token(
model=model,
custom_llm_provider="vertex_ai",
usage_object=usage,
vertex_location="us-east5",
)
global_total = global_prompt + global_completion
regional_total = regional_prompt + regional_completion
assert global_total > 0
assert regional_total == pytest.approx(global_total * 1.10, rel=1e-9), (
f"{model}: regional Vertex request must cost 1.1x the global one"
)
def test_vertex_uplift_composes_with_above_128k_pricing(monkeypatch):
"""The regional-endpoint uplift multiplies whatever rate the request priced at,
including the above-128k dynamic rates, so a synthetic model carrying both keys
prices regional above-128k usage at 1.1x the above-128k rate."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(
litellm,
"model_cost",
{
**litellm.get_model_cost_map(url=""),
"vertex_ai/fake-regional-128k-model": {
"litellm_provider": "vertex_ai",
"mode": "chat",
"input_cost_per_token": 1e-06,
"output_cost_per_token": 2e-06,
"input_cost_per_token_above_128k_tokens": 2e-06,
"output_cost_per_token_above_128k_tokens": 4e-06,
"regional_endpoint_uplift_multiplier": 1.1,
},
},
)
usage = Usage(prompt_tokens=200_000, completion_tokens=10, total_tokens=200_010)
global_prompt, global_completion = cost_per_token(
model="fake-regional-128k-model",
custom_llm_provider="vertex_ai",
usage_object=usage,
vertex_location="global",
)
regional_prompt, regional_completion = cost_per_token(
model="fake-regional-128k-model",
custom_llm_provider="vertex_ai",
usage_object=usage,
vertex_location="europe-west1",
)
assert global_prompt == pytest.approx(200_000 * 2e-06, rel=1e-9)
assert regional_prompt == pytest.approx(global_prompt * 1.10, rel=1e-9)
assert regional_completion == pytest.approx(global_completion * 1.10, rel=1e-9)
def test_cost_discount_vertex_ai():
"""
Test that cost discount is applied correctly for Vertex AI provider

View file

@ -1,12 +1,13 @@
"""Regression tests for the GitHub Actions change-based job gating.
`.github/scripts/detect_backend_changes.sh` decides whether a pull request's
backend unit-test jobs do real work. It asks the API which files the pull
request touches and hands them to `classify_changes.sh`. The contract locked in
here:
`.github/scripts/detect_changes.sh` decides whether a pull request's jobs do
real work. It asks the API which files the pull request touches and hands them
to `classify_changes.sh` under one category. The contract locked in here:
* a UI-only pull request skips backend jobs even when the checked-out merge
ref carries backend commits from the base branch
* the ui category is the mirror image: it skips when only backend files
changed, so a backend-only PR stops building and unit-testing the dashboard
* anything the classification cannot resolve (no pull request, an API
failure, a truncated file list, a broken classifier) runs the job
"""
@ -19,7 +20,7 @@ import subprocess
from pathlib import Path
REPO_ROOT = Path(__file__).resolve().parents[2]
SCRIPT = REPO_ROOT / ".github" / "scripts" / "detect_backend_changes.sh"
SCRIPT = REPO_ROOT / ".github" / "scripts" / "detect_changes.sh"
CLASSIFIER = REPO_ROOT / ".circleci" / "scripts" / "classify_changes.sh"
UI_FILE = "ui/litellm-dashboard/src/components/Teams.tsx"
@ -84,6 +85,7 @@ def _run(
changed_file_count: str | None = None,
gh_exit_code: int = 0,
classifier_body: str | None = None,
category: str | None = None,
) -> tuple[str, str]:
"""Run the script against a stubbed `gh`; returns (decision, stdout)."""
bin_dir = tmp_path / "bin"
@ -102,6 +104,10 @@ def _run(
env["REPO"] = "BerriAI/litellm"
env["PR_NUMBER"] = pr_number
env["CHANGED_FILE_COUNT"] = changed_file_count if changed_file_count is not None else str(len(files))
if category is not None:
env["CATEGORY"] = category
else:
env.pop("CATEGORY", None)
tree = _scripts_tree(tmp_path, classifier_body)
result = subprocess.run(
@ -187,3 +193,43 @@ def test_unexpected_classifier_output_runs(tmp_path: Path) -> None:
)
assert decision == "decision=run"
assert "unexpected decision: maybe" in stdout
def test_ui_category_skips_a_backend_only_pr(tmp_path: Path) -> None:
"""The dashboard build and its unit tests cannot be affected by a pull request
that touches no `ui/` file, and the `client` category cannot express that
because it deliberately runs whenever the backend changes."""
decision, _ = _run(tmp_path, files=[BACKEND_FILE], category="ui")
assert decision == "decision=skip"
def test_ui_category_runs_a_ui_only_pr(tmp_path: Path) -> None:
decision, _ = _run(tmp_path, files=[UI_FILE], category="ui")
assert decision == "decision=run"
def test_ui_category_runs_a_mixed_pr(tmp_path: Path) -> None:
decision, _ = _run(tmp_path, files=[UI_FILE, BACKEND_FILE], category="ui")
assert decision == "decision=run"
def test_absent_category_still_runs_a_backend_pr(tmp_path: Path) -> None:
"""Callers that pass no category keep the pre-existing backend behaviour."""
assert _run(tmp_path, files=[BACKEND_FILE])[0] == "decision=run"
def test_absent_category_still_skips_a_ui_pr(tmp_path: Path) -> None:
assert _run(tmp_path, files=[UI_FILE])[0] == "decision=skip"
def test_ui_category_fails_open_when_the_api_fails(tmp_path: Path) -> None:
decision, stdout = _run(tmp_path, files=[], gh_exit_code=1, category="ui")
assert decision == "decision=run"
assert "detect-changes[ui]" in stdout
def test_ui_category_runs_when_the_ui_workflows_themselves_change(tmp_path: Path) -> None:
"""Without this the dashboard jobs would skip on the pull request that edits
them, shipping a workflow change nothing ever exercised."""
decision, _ = _run(tmp_path, files=[".github/workflows/test-litellm-ui-unit.yml"], category="ui")
assert decision == "decision=run"

View file

@ -1584,3 +1584,116 @@ def test_inherit_builtin_tiered_output_rate_leaves_a_user_rate_alone():
)
assert model_info["output_cost_per_token"] == 9e-07
# --- a config.yaml PTU deployment must not also bill per token ------------------
_PTU_MODEL_INFO = {
"team_id": "team-alpha",
"ptu_count": 100,
"cost_per_ptu_per_hour": 0.02,
"ptu_effective_from": "2026-01-01T00:00:00Z",
}
def _ptu_router(model_info=None, litellm_params=None, ptu_enabled=True):
"""A router built the way loading config.yaml builds one."""
with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True" if ptu_enabled else ""}, clear=False):
return Router(
model_list=[
{
"model_name": "gpt-4o-ptu",
"litellm_params": {
"model": "anthropic/claude-sonnet-4-5-20250929",
"api_key": "sk-not-used",
**(litellm_params or {}),
},
"model_info": dict(_PTU_MODEL_INFO if model_info is None else model_info),
}
]
)
def test_a_config_ptu_deployment_bills_nothing_per_token():
"""Reserved capacity is already billed by the hour, so charging its traffic bills the
same tokens twice. Left unset the rate falls back to the public cost map, which makes
the double charge the default rather than an opt-in."""
router = _ptu_router(litellm_params={"input_cost_per_token": 5e-06, "output_cost_per_token": 1.5e-05})
entry = router.model_list[0]
assert entry["litellm_params"]["input_cost_per_token"] == 0.0
assert entry["litellm_params"]["output_cost_per_token"] == 0.0
assert entry["model_info"]["input_cost_per_token"] == 0.0
assert litellm.model_cost[entry["model_info"]["id"]]["input_cost_per_token"] == 0.0
@pytest.mark.parametrize(
"backend",
["anthropic/claude-sonnet-4-5-20250929", "azure/gpt-4o", "gemini/gemini-2.5-flash"],
)
def test_a_config_ptu_deployment_imports_no_cache_rate_from_its_backend(backend):
"""The cache back-fill runs whenever input_cost_per_token is set, and 0.0 is set, so a
partially zeroed deployment would silently inherit the backend model's real cache rates.
Every backend here publishes non-zero ones, which is what makes the assertion mean
something."""
cache_fields = (
"cache_creation_input_token_cost",
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost",
"cache_read_input_token_cost_above_200k_tokens",
)
builtin = litellm.get_model_info(model=backend)
assert any(builtin.get(field) for field in cache_fields), "backend publishes no cache pricing to leak"
router = _ptu_router(litellm_params={"model": backend})
priced = litellm.model_cost[router.model_list[0]["model_info"]["id"]]
assert [field for field in cache_fields if priced.get(field)] == []
def test_zeroing_a_ptu_deployment_leaves_its_backend_model_priced():
"""A sibling deployment on the same backend must keep billing normally."""
backend = "anthropic/claude-sonnet-4-5-20250929"
builtin = litellm.get_model_info(model=backend)["input_cost_per_token"]
assert builtin > 0
_ptu_router(litellm_params={"model": backend})
assert litellm.get_model_info(model=backend)["input_cost_per_token"] == builtin
def test_zeroing_does_not_change_the_deployment_id():
"""The id is a hash of the deployment's params and keys its cooldowns, its budget, and
every spend row already written against it."""
params = {"input_cost_per_token": 5e-06}
priced = _ptu_router(litellm_params=params, ptu_enabled=False).model_list[0]["model_info"]["id"]
zeroed = _ptu_router(litellm_params=params).model_list[0]["model_info"]["id"]
assert priced == zeroed
def test_a_database_backed_deployment_is_left_alone():
"""The write endpoints already zero those, and they answer 400 rather than silently
rewriting a rate the caller sent."""
entry = _ptu_router(model_info={**_PTU_MODEL_INFO, "db_model": True}).model_list[0]
assert entry["litellm_params"].get("input_cost_per_token") is None
def test_nothing_is_zeroed_while_the_feature_is_off():
"""No flat cost accrues with the flag off, so zeroing would serve the traffic free."""
entry = _ptu_router(litellm_params={"input_cost_per_token": 5e-06}, ptu_enabled=False).model_list[0]
assert entry["litellm_params"]["input_cost_per_token"] == 5e-06
@pytest.mark.parametrize("dropped", ["team_id", "ptu_effective_from"], ids=["no team_id", "no ptu_effective_from"])
def test_a_deployment_the_rollup_will_not_charge_is_not_zeroed(dropped):
"""The rollup refuses to price a reservation missing either field, so zeroing on the
looser count-and-rate test alone would leave the deployment serving for free with
nothing charged in its place."""
incomplete = {k: v for k, v in _PTU_MODEL_INFO.items() if k != dropped}
entry = _ptu_router(model_info=incomplete, litellm_params={"input_cost_per_token": 5e-06}).model_list[0]
assert entry["litellm_params"]["input_cost_per_token"] == 5e-06

View file

@ -832,6 +832,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"output_cost_per_token_above_200k_tokens_priority": {"type": "number"},
"output_cost_per_token_above_272k_tokens_priority": {"type": "number"},
"output_cost_per_token_above_272k_tokens_flex": {"type": "number"},
"regional_endpoint_uplift_multiplier": {"type": "number"},
"regional_processing_uplift_multiplier_eu": {"type": "number"},
"regional_processing_uplift_multiplier_us": {"type": "number"},
"input_cost_per_pixel": {"type": "number"},

View file

@ -57,9 +57,6 @@
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 1
}
@ -199,11 +196,6 @@
"count": 1
}
},
"src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/guardrails-monitor/_components/GuardrailDetail.tsx": {
"no-nested-ternary": {
"count": 3
@ -521,35 +513,12 @@
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/AwsSigV4Fields.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx": {
"max-lines": {
"count": 1
},
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/IdJagFormFields.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx": {
@ -557,16 +526,6 @@
"count": 2
}
},
"src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/MCPPermissionManagement.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -583,25 +542,14 @@
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx": {
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/app/(dashboard)/mcp-servers/_components/OpenAPIQuickPicker.tsx": {
@ -609,36 +557,6 @@
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/OpenApiByokFields.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/PassthroughAuthorizeSection.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/StdioConfiguration.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/TokenEndpointAuthMethodField.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -647,9 +565,6 @@
"src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx": {
"no-nested-ternary": {
"count": 2
},
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/mcp-servers/_components/index.tsx": {
@ -661,9 +576,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/static-components": {
"count": 4
}
@ -703,15 +615,6 @@
},
"max-lines": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 5
}
},
"src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx": {
@ -861,11 +764,6 @@
"count": 1
}
},
"src/app/(dashboard)/playground/components/compareUI/components/ModelSelector.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx": {
"local/no-complex-jsx-arrow": {
"count": 2
@ -1066,9 +964,6 @@
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 1
}
@ -1336,11 +1231,6 @@
"count": 1
}
},
"src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx": {
"no-restricted-imports": {
"count": 2
}
},
"src/app/(dashboard)/vector-stores/_components/VectorStoreForm.tsx": {
"no-nested-ternary": {
"count": 2
@ -1475,18 +1365,10 @@
"count": 1
}
},
"src/components/Settings/AdminSettings/LoggingSettings/LoggingSettings.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterSettings.tsx": {
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1506,11 +1388,6 @@
"count": 1
}
},
"src/components/Settings/AdminSettings/SSOSettings/RoleMappings.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/Settings/AdminSettings/UISettings/PageVisibilitySettings.tsx": {
"react-hooks/set-state-in-render": {
"count": 1
@ -1526,18 +1403,7 @@
"count": 1
}
},
"src/components/Settings/RouterSettings/Fallbacks/FallbackGroupConfig.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx": {
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -1577,11 +1443,6 @@
"count": 1
}
},
"src/components/UsagePage/components/EntityUsage/TopKeyView.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/UsagePage/utils/value_formatters.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1590,9 +1451,6 @@
"src/components/VirtualKeysPage/keyTableColumns.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/activity_metrics.tsx": {
@ -1603,16 +1461,6 @@
"count": 1
}
},
"src/components/add_model/AdaptiveRoutingConfig.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/AddModelForm.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/AddModelForm.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
@ -1620,26 +1468,6 @@
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/components/add_model/ClassificationMethodConfig.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/ComplexityRouterConfig.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/EscalationKeywords.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/KeywordTierRules.tsx": {
"no-restricted-imports": {
"count": 1
}
@ -1649,11 +1477,6 @@
"count": 1
}
},
"src/components/add_model/SemanticKeywordMatching.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/add_model/add_auto_router_tab.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1670,9 +1493,6 @@
"src/components/add_model/advanced_settings.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 3
}
},
"src/components/add_model/auto_router_connection_test.tsx": {
@ -1692,9 +1512,6 @@
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 2
}
@ -1715,9 +1532,6 @@
},
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/components/add_model/model_connection_test.tsx": {
@ -1735,9 +1549,6 @@
"no-nested-ternary": {
"count": 3
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/immutability": {
"count": 3
}
@ -1747,15 +1558,7 @@
"count": 1
}
},
"src/components/agent_management/AgentSelector.test.tsx": {
"react/display-name": {
"count": 1
}
},
"src/components/agent_management/AgentSelector.tsx": {
"no-restricted-imports": {
"count": 1
},
"prefer-const": {
"count": 1
}
@ -1846,11 +1649,6 @@
"count": 1
}
},
"src/components/common_components/AccessGroupSelector.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/DeleteResourceModal.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -1861,16 +1659,6 @@
"count": 1
}
},
"src/components/common_components/MetadataKeyValueFields.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/MetadataKeyValueFields.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/ModelAliasManager.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1
@ -1881,11 +1669,6 @@
"count": 1
}
},
"src/components/common_components/RateLimitTypeFormItem.test.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/common_components/budget_duration_dropdown.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -1894,9 +1677,6 @@
"src/components/common_components/check_openapi_schema.tsx": {
"local/filename-pascal-case": {
"count": 1
},
"no-restricted-imports": {
"count": 2
}
},
"src/components/common_components/fetch_teams.tsx": {
@ -1967,16 +1747,6 @@
"count": 1
}
},
"src/components/key_team_helpers/BudgetFallbacksEditor.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/key_team_helpers/BudgetWindowsEditor.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/key_team_helpers/fetch_available_models_team_key.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -2040,14 +1810,6 @@
"count": 1
}
},
"src/components/mcp_server_management/MCPServerSelector.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/mcp_server_management/MCPToolPermissions.tsx": {
"local/no-complex-jsx-arrow": {
"count": 1
@ -2061,9 +1823,6 @@
"src/components/mcp_tools/MCPToolArgumentsForm.tsx": {
"no-nested-ternary": {
"count": 1
},
"no-restricted-imports": {
"count": 1
}
},
"src/components/mcp_tools/McpCrudPermissionPanel.tsx": {
@ -2078,7 +1837,7 @@
},
"src/components/model_add/CredentialModal.tsx": {
"no-restricted-imports": {
"count": 2
"count": 1
}
},
"src/components/model_add/reuse_credentials.tsx": {
@ -2086,11 +1845,6 @@
"count": 1
}
},
"src/components/model_dashboard/ModelSettingsModal/ModelSettingsModal.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/model_filters.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -2244,11 +1998,6 @@
"count": 1
}
},
"src/components/router_settings/RoutingStrategySelector.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/router_settings/index.tsx": {
"local/filename-pascal-case": {
"count": 1
@ -2397,11 +2146,6 @@
"count": 1
}
},
"src/components/team/LoggingSettings.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/components/team/TeamInfo.tsx": {
"max-lines": {
"count": 1
@ -2409,9 +2153,6 @@
"no-nested-ternary": {
"count": 3
},
"no-restricted-imports": {
"count": 1
},
"react-hooks/set-state-in-effect": {
"count": 1
}
@ -2747,11 +2488,6 @@
"count": 2
}
},
"src/contexts/AntdGlobalProvider.tsx": {
"no-restricted-imports": {
"count": 1
}
},
"src/contexts/AuthContext.tsx": {
"react-hooks/set-state-in-effect": {
"count": 1

View file

@ -62,6 +62,10 @@ const eslintConfig = [
message:
"antd is being phased out; build new UI with shadcn/ui primitives instead of adding antd imports.",
},
{
group: ["@ant-design/icons", "@ant-design/icons/*"],
message: "@ant-design/icons is gone from the dashboard; use lucide-react instead.",
},
],
},
],

View file

@ -9,7 +9,6 @@
"version": "0.1.0",
"dependencies": {
"@ant-design/cssinjs": "1.24.0",
"@ant-design/icons": "5.6.1",
"@anthropic-ai/sdk": "0.92.0",
"@base-ui/react": "^1.6.0",
"@headlessui/tailwindcss": "0.2.2",

View file

@ -25,7 +25,6 @@
},
"dependencies": {
"@ant-design/cssinjs": "1.24.0",
"@ant-design/icons": "5.6.1",
"@anthropic-ai/sdk": "0.92.0",
"@base-ui/react": "^1.6.0",
"@headlessui/tailwindcss": "0.2.2",

View file

@ -3,7 +3,7 @@ import { useDeleteAccessGroup } from "@/app/(dashboard)/hooks/accessGroups/useDe
import { Plus, SearchIcon, X } from "lucide-react";
import { useMemo, useState } from "react";
import DeleteResourceModal from "@/components/common_components/DeleteResourceModal";
import { PageHeader } from "@/components/shared/PageHeader";
import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader";
import { Button } from "@/components/ui/button";
import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group";
import { AccessGroupDetail } from "./AccessGroupsDetailsPage";
@ -61,7 +61,7 @@ export function AccessGroupsPage() {
return (
<div className="p-6 px-12">
<div className="mb-4">
<PageHeader
<LegacyPageHeader
title="Access Groups"
subtitle="Manage resource permissions for your organization"
actions={

View file

@ -3,7 +3,7 @@ import { Select, Steps, Tag } from "antd";
import { FormProvider, useForm, useWatch } from "react-hook-form";
import { toast } from "@/lib/toast";
import { Logo } from "@/components/molecules/logo/Logo";
import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined } from "@ant-design/icons";
import { Bot, CircleCheck, Key, LayoutGrid } from "lucide-react";
import CreatedKeyDisplay from "@/components/shared/CreatedKeyDisplay";
import { Button } from "@/components/ui/button";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
@ -707,7 +707,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
}`}
onClick={() => handleAgentTypeChange(CUSTOM_AGENT_TYPE)}
>
<AppstoreOutlined className="text-lg text-amber-600 dark:text-amber-400" />
<LayoutGrid className="size-4.5 text-amber-600 dark:text-amber-400" />
<div className="flex-1">
<div className="flex items-center gap-2">
<span className="font-medium text-amber-700 dark:text-amber-400">Custom / Other</span>
@ -846,7 +846,8 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<div>
{/* Agent name chip */}
<div className="mb-6 flex justify-center">
<Tag icon={<RobotOutlined />} color="purple" className="px-3 py-1 text-sm">
<Tag color="purple" className="inline-flex items-center gap-1.5 px-3 py-1 text-sm">
<Bot className="size-3.5" />
{agentName}
</Tag>
</div>
@ -884,7 +885,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<RadioGroupItem value="create_new" aria-label="Create a new key for this agent" />
<div className="flex-1">
<div className="flex items-center gap-2">
<KeyOutlined className="text-indigo-600 dark:text-indigo-400" />
<Key className="size-4 text-indigo-600 dark:text-indigo-400" />
<span className="font-medium text-foreground">Create a new key for this agent</span>
</div>
<p className="mt-1 text-sm text-muted-foreground">A dedicated key scoped to this agent.</p>
@ -920,7 +921,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
<RadioGroupItem value="existing_key" aria-label="Assign an existing key" />
<div className="flex-1">
<div className="flex items-center gap-2">
<KeyOutlined className="text-muted-foreground" />
<Key className="size-4 text-muted-foreground" />
<span className="font-medium text-foreground">Assign an existing key</span>
</div>
<p className="mt-1 text-sm text-muted-foreground">Re-assign a key you already have to this agent.</p>
@ -958,10 +959,11 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({ visible, onClose, accessTok
const renderReadyStep = () => (
<div className="py-6 text-center">
<CheckCircleFilled className="mb-4 text-5xl text-green-500" style={{ fontSize: 48 }} />
<CircleCheck className="mb-4 size-12 text-green-500" />
<h3 className="mb-2 text-xl font-semibold text-foreground">Agent Created!</h3>
<div className="mb-4 flex justify-center">
<Tag icon={<RobotOutlined />} color="purple" className="px-3 py-1 text-sm">
<Tag color="purple" className="inline-flex items-center gap-1.5 px-3 py-1 text-sm">
<Bot className="size-3.5" />
{createdAgentName}
</Tag>
</div>

View file

@ -6,7 +6,7 @@
import { Plus, Wallet } from "lucide-react";
import React, { useCallback, useState } from "react";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { PageHeader } from "@/components/shared/PageHeader";
import { LegacyPageHeader } from "@/components/shared/LegacyPageHeader";
import { ToolbarSeparator } from "@/components/shared/ToolbarSeparator";
import { Button } from "@/components/ui/button";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
@ -76,7 +76,7 @@ const BudgetPanel: React.FC<BudgetSettingsPageProps> = ({ accessToken }) => {
return (
<div className="flex h-full flex-col gap-4 p-6 px-12">
<PageHeader
<LegacyPageHeader
icon={<Wallet className="size-5" />}
title="Budgets"
subtitle="Spend, TPM and RPM limits you can assign to customers."

View file

@ -1,10 +1,4 @@
import {
CheckCircleOutlined,
CodeOutlined,
PlayCircleOutlined,
RollbackOutlined,
SaveOutlined,
} from "@ant-design/icons";
import { CircleCheck, CirclePlay, Code, Save, Undo2 } from "lucide-react";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
@ -100,11 +94,11 @@ export function GuardrailConfig({ guardrailName, guardrailType, provider }: Guar
</div>
<div className="flex items-center gap-2">
<Button variant="outline">
<RollbackOutlined />
<Undo2 />
Revert
</Button>
<Button>
<SaveOutlined />
<Save />
Save as v{parseInt(version.replace("v", ""), 10) + 1}
</Button>
</div>
@ -214,7 +208,7 @@ export function GuardrailConfig({ guardrailName, guardrailType, provider }: Guar
<div className="flex items-center justify-between mb-4">
<div>
<h3 className="text-base font-semibold text-gray-900 flex items-center gap-2">
<CodeOutlined className="text-gray-500" />
<Code className="size-4 text-gray-500" />
Custom Code Override
</h3>
<p className="text-xs text-gray-500 mt-0.5">Replace the built-in guardrail with custom evaluation code</p>
@ -247,13 +241,13 @@ export function GuardrailConfig({ guardrailName, guardrailType, provider }: Guar
<div className="flex items-center gap-3">
<Button disabled={rerunStatus === "running"} aria-busy={rerunStatus === "running"} onClick={handleRerun}>
{rerunStatus === "running" ? null : <PlayCircleOutlined />}
{rerunStatus === "running" ? null : <CirclePlay />}
{rerunStatus === "running" ? "Running on 10 samples..." : "Re-run on failing logs"}
</Button>
{rerunStatus === "success" && (
<span className="text-sm text-green-600 flex items-center gap-2">
<CheckCircleOutlined /> 7/10 would now pass with new config
<CircleCheck className="size-4" /> 7/10 would now pass with new config
</span>
)}

View file

@ -4,10 +4,6 @@ import userEvent from "@testing-library/user-event";
import GuardrailCard from "./guardrail_garden_card";
import type { GuardrailCardInfo } from "./guardrail_garden_data";
vi.mock("@ant-design/icons", () => ({
CheckCircleFilled: ({ style, ...props }: any) => <span data-testid="check-icon" {...props} />,
}));
const baseCard: GuardrailCardInfo = {
id: "test-guard",
name: "Test Guardrail",

View file

@ -153,7 +153,7 @@ describe("Guardrail Info", () => {
expect(getByText("Guardrail Settings")).toBeInTheDocument();
});
await userEvent.hover(within(container).getByRole("img", { name: "info-circle" }));
await userEvent.hover(within(container).getByRole("img", { name: "Config guardrail details" }));
expect(await findByText("Guardrail is defined in the config file and cannot be edited.")).toBeInTheDocument();
});

View file

@ -5,9 +5,8 @@ import {
updateGuardrailCall,
} from "@/components/networking";
import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils";
import { EyeInvisibleOutlined, InfoCircleOutlined, StopOutlined } from "@ant-design/icons";
import { ArrowLeft, CheckIcon, Code, CopyIcon } from "lucide-react";
import { ArrowLeft, Ban, CheckIcon, Code, CopyIcon, EyeOff, Info } from "lucide-react";
import { Badge } from "@/components/ui/badge";
import { Card } from "@/components/ui/card";
import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs";
@ -607,7 +606,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
value === "MASK" ? "text-blue-600" : "text-red-600"
}`}
>
{value === "MASK" ? <EyeInvisibleOutlined /> : <StopOutlined />}
{value === "MASK" ? <EyeOff className="size-3.5" /> : <Ban className="size-3.5" />}
{String(value)}
</span>
</p>
@ -667,7 +666,7 @@ const GuardrailInfoView: React.FC<GuardrailInfoProps> = ({ guardrailId, onClose,
<h3 className="text-lg font-medium">Guardrail Settings</h3>
{isConfigGuardrail && (
<SimpleTooltip content="Guardrail is defined in the config file and cannot be edited.">
<InfoCircleOutlined />
<Info role="img" aria-label="Config guardrail details" className="size-4 text-muted-foreground" />
</SimpleTooltip>
)}
{!isEditing &&

View file

@ -1,9 +1,11 @@
import { Info } from "lucide-react";
import React from "react";
import { Input, Tooltip } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import { antdRequired } from "@/components/common_components/antdFormRules";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { Input } from "@/components/ui/input";
import { requiredWhenSiblingSet, textControl } from "./mcpFieldRules";
const fieldClassName = "rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500";
@ -11,9 +13,9 @@ const fieldClassName = "rounded-lg border-gray-300 focus:border-blue-500 focus:r
const FieldLabel: React.FC<{ label: string; tooltip: string }> = ({ label, tooltip }) => (
<span className="text-sm font-medium text-gray-700 flex items-center">
{label}
<Tooltip title={tooltip}>
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content={tooltip}>
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
);
@ -71,10 +73,10 @@ const AwsSigV4Fields: React.FC = () => (
}}
>
{(control) => (
<Input.Password
<PasswordInput
{...textControl(control)}
placeholder="AKIA... (optional — uses IAM role if blank)"
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -94,10 +96,10 @@ const AwsSigV4Fields: React.FC = () => (
}}
>
{(control) => (
<Input.Password
<PasswordInput
{...textControl(control)}
placeholder="Enter secret key (optional — uses IAM role if blank)"
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -106,10 +108,10 @@ const AwsSigV4Fields: React.FC = () => (
name={["credentials", "aws_session_token"]}
>
{(control) => (
<Input.Password
<PasswordInput
{...textControl(control)}
placeholder="Enter session token (optional)"
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>

View file

@ -4,7 +4,7 @@ import { beforeEach, describe, expect, it, vi } from "vitest";
import * as networking from "@/components/networking";
import { setToken } from "@/utils/mcpTokenStore";
import CreateMCPServer from "./CreateMCPServer";
import { selectAntOption } from "./testUtils";
import { selectOption } from "./testUtils";
vi.mock("@/components/networking", () => ({
createMCPServer: vi.fn(),
@ -129,7 +129,7 @@ describe("CreateMCPServer", () => {
expect(screen.getByText("Submit MCP Server for Review")).toBeInTheDocument();
expect(screen.queryByText("Add New MCP Server")).not.toBeInTheDocument();
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
@ -141,7 +141,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://example.com/mcp" },
});
});
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
vi.mocked(networking.registerMCPServer).mockResolvedValue({
server_id: "submitted-1",
@ -167,7 +167,7 @@ describe("CreateMCPServer", () => {
it("should show transport type options", async () => {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
// Verify the option was applied by checking the URL field appears
await waitFor(() => {
@ -178,7 +178,7 @@ describe("CreateMCPServer", () => {
describe("when HTTP transport is selected", () => {
async function selectHttpTransport() {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
// Wait for URL field to appear (confirms transport was set)
await waitFor(() => {
@ -201,7 +201,7 @@ describe("CreateMCPServer", () => {
it("should show auth value field when API Key auth type is selected", async () => {
await selectHttpTransport();
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => {
expect(screen.getByText("Authentication Value")).toBeInTheDocument();
@ -211,7 +211,7 @@ describe("CreateMCPServer", () => {
it("should warn that LiteLLM auth is disabled when True Passthrough is selected", async () => {
await selectHttpTransport();
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
await waitFor(() => {
expect(
@ -223,7 +223,7 @@ describe("CreateMCPServer", () => {
it("should not show the True Passthrough warning when OAuth Delegate is selected", async () => {
await selectHttpTransport();
await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await selectOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await waitFor(() => {
expect(screen.getAllByText("OAuth Delegate (client-supplied upstream token)").length).toBeGreaterThan(0);
@ -238,7 +238,7 @@ describe("CreateMCPServer", () => {
async (optionLabel) => {
await selectHttpTransport();
await selectAntOption("Authentication", optionLabel);
await selectOption("Authentication", optionLabel);
await waitFor(() => {
expect(screen.getByRole("button", { name: "Authorize & Fetch Tools (browser-only)" })).toBeInTheDocument();
@ -251,7 +251,7 @@ describe("CreateMCPServer", () => {
it("should not show the browser-only authorize section for API Key auth", async () => {
await selectHttpTransport();
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => {
expect(screen.getByText("Authentication Value")).toBeInTheDocument();
@ -273,7 +273,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
// Select API Key auth type
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => {
expect(screen.getByText("Authentication Value")).toBeInTheDocument();
@ -315,7 +315,7 @@ describe("CreateMCPServer", () => {
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
await selectAntOption("Authentication", "Bearer Token");
await selectOption("Authentication", "Bearer Token");
await waitFor(() => {
expect(screen.getByText("Authentication Value")).toBeInTheDocument();
@ -356,7 +356,7 @@ describe("CreateMCPServer", () => {
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => {
expect(screen.getByText("Authentication Value")).toBeInTheDocument();
@ -402,7 +402,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
// Simulate the browser Authorize & Fetch flow handing back an upstream token.
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
@ -430,7 +430,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", optionLabel);
await selectOption("Authentication", optionLabel);
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
await act(async () => {
@ -496,7 +496,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", optionLabel);
await selectOption("Authentication", optionLabel);
// Admin declares the org's pre-registered upstream app; unlike the browser-authorized
// token, this is config and must survive onto the server row so internal users'
@ -560,7 +560,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
fireEvent.change(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), {
target: { value: "org-app-client-id" },
@ -620,7 +620,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
// The oauth2 onTokenReceived branch writes the fetched token AND the DCR client into
// form.credentials; both are minted for the oauth2 identity.
@ -635,7 +635,7 @@ describe("CreateMCPServer", () => {
// Switching into a client-forwarded mode changes the identity with auth_type in the changed
// values, so the preserve carve-out must NOT apply: the minted material would otherwise ride
// into a mode that now persists credentials onto the server row.
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
const switchedServer = {
server_id: "switched-server",
@ -669,7 +669,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
await act(async () => {
@ -689,14 +689,14 @@ describe("CreateMCPServer", () => {
it("clears the DCR ref and the upstream warning when the modal closes so nothing leaks to the next session", async () => {
const { rerender } = render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
expect(await screen.findByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
const user = userEvent.setup({ delay: null });
fireEvent.change(getServerNameInput(), { target: { value: "Leak_Server" } });
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
await act(async () => {
@ -725,7 +725,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
await act(async () => {
@ -767,7 +767,7 @@ describe("CreateMCPServer", () => {
await selectHttpTransport();
fillText(getServerNameInput(), "CF_Switch_Keep");
fillText(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp");
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
fillText(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "app-id");
fillText(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "app-secret");
@ -776,7 +776,7 @@ describe("CreateMCPServer", () => {
oauthHook.onTokenReceived!({ access_token: "cf-tok", token_type: "Bearer" }, undefined);
});
await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await selectOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
const switched = {
server_id: "cf-switch-keep",
@ -804,7 +804,7 @@ describe("CreateMCPServer", () => {
await selectHttpTransport();
fillText(getServerNameInput(), "CF_Round");
fillText(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp");
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
fillText(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), "app-id");
fillText(screen.getByPlaceholderText("Leave blank for public clients / PKCE"), "app-secret");
@ -813,8 +813,8 @@ describe("CreateMCPServer", () => {
oauthHook.onTokenReceived!({ access_token: "cf-tok", token_type: "Bearer" }, undefined);
});
await selectAntOption("Authentication", "OAuth");
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "OAuth");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
const cfRoundServer = {
server_id: "cf-round",
@ -845,7 +845,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
fireEvent.change(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), {
target: { value: "app-id" },
});
@ -871,7 +871,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
fireEvent.change(screen.getByPlaceholderText("Leave blank to use dynamic client registration"), {
target: { value: "app-id" },
});
@ -919,7 +919,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(screen.getByPlaceholderText("https://your-mcp-server.com"), {
target: { value: "https://example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => expect(oauthHook.onTokenReceived).toBeTruthy());
const firstToken = { access_token: "T1", refresh_token: "R1", scope: "read", token_type: "Bearer" };
@ -939,7 +939,7 @@ describe("CreateMCPServer", () => {
it("should not show auth value field when None auth type is selected", async () => {
await selectHttpTransport();
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
// Auth value field should not appear for "None"
await waitFor(() => {
@ -958,7 +958,7 @@ describe("CreateMCPServer", () => {
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-server-1",
@ -992,20 +992,20 @@ describe("CreateMCPServer", () => {
await selectHttpTransport();
// Plain OAuth must not render the token-exchange section.
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.queryByText("Token Exchange Endpoint (optional)")).not.toBeInTheDocument();
});
expect(screen.queryByText("Subject Token Type (optional)")).not.toBeInTheDocument();
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
await selectOption("Authentication", "OAuth Token Exchange (OBO)");
await waitFor(() => {
expect(screen.getByText("Token Exchange Endpoint (optional)")).toBeInTheDocument();
});
expect(screen.getByText("Subject Token Type (optional)")).toBeInTheDocument();
// Switching away hides the section again.
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => {
expect(screen.queryByText("Token Exchange Endpoint (optional)")).not.toBeInTheDocument();
});
@ -1015,11 +1015,11 @@ describe("CreateMCPServer", () => {
// whole Authentication section (the section-level transport gate), taking the
// token-exchange fields with it — their required client_id/client_secret rules
// cannot block a stdio submit because antd does not validate unmounted fields.
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
await selectOption("Authentication", "OAuth Token Exchange (OBO)");
await waitFor(() => {
expect(screen.getByText("Token Exchange Endpoint (optional)")).toBeInTheDocument();
});
await selectAntOption("Transport Type", "Standard Input/Output");
await selectOption("Transport Type", "Standard Input/Output");
await waitFor(() => {
expect(screen.queryByText("Token Exchange Endpoint (optional)")).not.toBeInTheDocument();
});
@ -1037,7 +1037,7 @@ describe("CreateMCPServer", () => {
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
const limitInput = screen.getByPlaceholderText("e.g. 10");
fireEvent.change(limitInput, { target: { value: "5" } });
@ -1079,7 +1079,7 @@ describe("CreateMCPServer", () => {
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://upstream.example.com/mcp" } });
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
await selectOption("Authentication", "OAuth Token Exchange (OBO)");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://idp.example.com/oauth2/token")).toBeInTheDocument();
@ -1131,20 +1131,20 @@ describe("CreateMCPServer", () => {
await selectHttpTransport();
// The sibling OBO mode must not render the ID-JAG section.
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
await selectOption("Authentication", "OAuth Token Exchange (OBO)");
await waitFor(() => {
expect(screen.queryByText("Org Token Endpoint (leg 1)")).not.toBeInTheDocument();
});
expect(screen.queryByText("Resource Token Endpoint (leg 2)")).not.toBeInTheDocument();
await selectAntOption("Authentication", "ID-JAG (Okta Cross App Access)");
await selectOption("Authentication", "ID-JAG (Okta Cross App Access)");
await waitFor(() => {
expect(screen.getByText("Org Token Endpoint (leg 1)")).toBeInTheDocument();
});
expect(screen.getByText("Resource Token Endpoint (leg 2)")).toBeInTheDocument();
expect(screen.getByText("Client Private Key (PEM)")).toBeInTheDocument();
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => {
expect(screen.queryByText("Org Token Endpoint (leg 1)")).not.toBeInTheDocument();
});
@ -1159,7 +1159,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://upstream.example.com/mcp" },
});
await selectAntOption("Authentication", "ID-JAG (Okta Cross App Access)");
await selectOption("Authentication", "ID-JAG (Okta Cross App Access)");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-org.okta.com/oauth2/v1/token")).toBeInTheDocument();
@ -1222,7 +1222,7 @@ describe("CreateMCPServer", () => {
target: { value: "https://upstream.example.com/mcp" },
});
await selectAntOption("Authentication", "ID-JAG (Okta Cross App Access)");
await selectOption("Authentication", "ID-JAG (Okta Cross App Access)");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-org.okta.com/oauth2/v1/token")).toBeInTheDocument();
@ -1259,12 +1259,12 @@ describe("CreateMCPServer", () => {
target: { value: "https://upstream.example.com/mcp" },
});
await selectAntOption("Authentication", "OAuth Token Exchange (OBO)");
await selectOption("Authentication", "OAuth Token Exchange (OBO)");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://idp.example.com/oauth2/token")).toBeInTheDocument();
});
await selectAntOption("Profile", "Microsoft Entra OBO");
await selectOption("Profile", "Microsoft Entra OBO");
fireEvent.change(screen.getByPlaceholderText("https://idp.example.com/oauth2/token"), {
target: { value: "https://login.microsoftonline.com/tenant/oauth2/v2.0/token" },
@ -1301,7 +1301,7 @@ describe("CreateMCPServer", () => {
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
await act(async () => {
fireEvent.click(screen.getByRole("button", { name: "Disable all tools" }));
@ -1339,13 +1339,13 @@ describe("CreateMCPServer", () => {
/** Select HTTP transport + OAuth auth, then wait for the OAuth form to appear. */
async function setupOAuthInteractive() {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
// Wait for OAuthFormFields to render (OAuth Flow Type selector is the sentinel)
await waitFor(() => {
@ -1376,7 +1376,7 @@ describe("CreateMCPServer", () => {
oauthHook.reset.mockClear();
// Switching the Authentication mode changes the OAuth identity, so the held token is discarded.
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
await waitFor(() => expect(oauthHook.reset).toHaveBeenCalled());
});
@ -1421,7 +1421,7 @@ describe("CreateMCPServer", () => {
});
vi.mocked(networking.testMCPToolsListRequest).mockClear();
await selectAntOption("Authentication", "API Key");
await selectOption("Authentication", "API Key");
await waitFor(() => expect(vi.mocked(networking.testMCPToolsListRequest)).toHaveBeenCalled());
for (const call of vi.mocked(networking.testMCPToolsListRequest).mock.calls) {
@ -1443,7 +1443,7 @@ describe("CreateMCPServer", () => {
});
oauthHook.reset.mockClear();
await selectAntOption("Transport Type", "Server-Sent Events (SSE)");
await selectOption("Transport Type", "Server-Sent Events (SSE)");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
@ -1545,11 +1545,11 @@ describe("CreateMCPServer", () => {
it("invalidates the DCR client and OAuth flow when the OpenAPI spec URL changes after Authorize & Fetch", async () => {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "OpenAPI Spec");
await selectOption("Transport Type", "OpenAPI Spec");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://petstore3.swagger.io/api/v3/openapi.json")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
@ -1601,11 +1601,11 @@ describe("CreateMCPServer", () => {
it("invalidates the DCR client and OAuth flow when the transport changes after Authorize & Fetch", async () => {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "OpenAPI Spec");
await selectOption("Transport Type", "OpenAPI Spec");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://petstore3.swagger.io/api/v3/openapi.json")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
@ -1624,7 +1624,7 @@ describe("CreateMCPServer", () => {
});
oauthHook.reset.mockClear();
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
@ -1688,7 +1688,7 @@ describe("CreateMCPServer", () => {
fireEvent.change(urlInput, { target: { value: "https://example.com/mcp" } });
});
await selectAntOption("Token Endpoint Auth Method (optional)", "Client Secret Basic");
await selectOption("Token Endpoint Auth Method (optional)", "Client Secret Basic");
const submitButton = screen.getByRole("button", { name: "Add MCP Server" });
await act(async () => {
@ -1804,11 +1804,11 @@ describe("CreateMCPServer", () => {
const { rerender } = render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
@ -1858,11 +1858,11 @@ describe("CreateMCPServer", () => {
const { rerender } = render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
@ -1915,7 +1915,7 @@ describe("CreateMCPServer", () => {
it("should not show auth type or URL fields", async () => {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Standard Input/Output");
await selectOption("Transport Type", "Standard Input/Output");
// Auth and URL fields should not be present for stdio
await waitFor(() => {
@ -1984,7 +1984,7 @@ describe("CreateMCPServer oauth2_flow persistence", () => {
async function setupHttpServerForm() {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
@ -2013,7 +2013,7 @@ describe("CreateMCPServer oauth2_flow persistence", () => {
it("persists authorization_code for an interactive OAuth create", async () => {
vi.mocked(networking.createMCPServer).mockResolvedValue(createdServer);
await setupHttpServerForm();
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
@ -2026,11 +2026,11 @@ describe("CreateMCPServer oauth2_flow persistence", () => {
it("persists client_credentials for an M2M OAuth create", async () => {
vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, oauth2_flow: "client_credentials" });
await setupHttpServerForm();
await selectAntOption("Authentication", "OAuth");
await selectOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
await selectAntOption("OAuth Flow Type", "Machine-to-Machine (M2M)");
await selectOption("OAuth Flow Type", "Machine-to-Machine (M2M)");
await waitFor(() => {
expect(screen.getByPlaceholderText("Enter OAuth client ID")).toBeInTheDocument();
});
@ -2074,11 +2074,11 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
updated_by: "user-1",
};
const getDcrToggle = () => document.getElementById("dcr_bridge");
const getDcrToggle = () => screen.queryByRole("switch", { name: /Gateway-hosted sign-in \(DCR bridge\)/ });
async function setupHttpServerForm() {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
@ -2109,7 +2109,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
async (optionLabel) => {
await setupHttpServerForm();
await selectAntOption("Authentication", optionLabel);
await selectOption("Authentication", optionLabel);
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
@ -2122,7 +2122,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
it.each([["None"], ["API Key"], ["OAuth"]])("does not render the toggle for %s", async (optionLabel) => {
await setupHttpServerForm();
await selectAntOption("Authentication", optionLabel);
await selectOption("Authentication", optionLabel);
await waitFor(() => {
expect(screen.queryByText("Gateway-hosted sign-in (DCR bridge)")).not.toBeInTheDocument();
@ -2133,7 +2133,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
it("renders the toggle between the OAuth client fields and the Authorize button", async () => {
await setupHttpServerForm();
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
@ -2151,7 +2151,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
])("sends dcr_bridge: true by default on create for %s", async (authType, optionLabel) => {
vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: authType });
await setupHttpServerForm();
await selectAntOption("Authentication", optionLabel);
await selectOption("Authentication", optionLabel);
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
});
@ -2163,7 +2163,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
it("sends an explicit dcr_bridge: false when the toggle is unchecked", async () => {
vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "oauth_delegate" });
await setupHttpServerForm();
await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await selectOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
});
@ -2184,7 +2184,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
it("forces dcr_bridge: false when the auth type is switched away after toggling", async () => {
vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "none" });
await setupHttpServerForm();
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
});
@ -2192,7 +2192,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
fireEvent.click(getDcrToggle()!);
});
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
await waitFor(() => {
expect(getDcrToggle()).not.toBeInTheDocument();
});
@ -2204,7 +2204,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
it("preserves the toggle value when switching between the two client-forwarded modes", async () => {
vi.mocked(networking.createMCPServer).mockResolvedValue({ ...createdServer, auth_type: "oauth_delegate" });
await setupHttpServerForm();
await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)");
await selectOption("Authentication", "True Passthrough (no LiteLLM auth)");
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
});
@ -2212,7 +2212,7 @@ describe("CreateMCPServer dcr_bridge toggle", () => {
// The field is mounted in both client-forwarded modes, so switching between them keeps the
// live toggle value rather than forcing it back to the default or to false.
await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await selectOption("Authentication", "OAuth Delegate (client-supplied upstream token)");
await waitFor(() => {
expect(getDcrToggle()).toBeInTheDocument();
});

View file

@ -3,7 +3,7 @@ import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import * as networking from "@/components/networking";
import CreateMCPServer from "./CreateMCPServer";
import { selectAntOption } from "./testUtils";
import { selectOption } from "./testUtils";
vi.mock("@/components/networking", () => ({
createMCPServer: vi.fn(),
@ -58,25 +58,29 @@ const defaultProps = {
const getServerNameInput = () => document.getElementById("server_name") as HTMLInputElement;
const switchFor = (labelText: string): HTMLElement => {
const label = screen.getByText(labelText);
const row = label.closest(".flex.items-start.justify-between");
const control = row?.querySelector("button[role='switch']");
if (control === null || control === undefined) {
throw new Error(`no switch found for "${labelText}"`);
// The switches live behind a collapsed panel, so they only reach the accessibility tree once an
// operator expands it.
const expandPermissionPanel = async (): Promise<void> => {
const trigger = screen.getByRole("button", { name: /Permission Management/ });
if (trigger.getAttribute("aria-expanded") !== "true") {
await userEvent.setup({ delay: null }).click(trigger);
}
return control as HTMLElement;
};
const switchFor = async (labelText: string): Promise<HTMLElement> => {
await expandPermissionPanel();
return screen.getByRole("switch", { name: labelText });
};
const fillMinimalHttpServer = async (name: string) => {
await selectAntOption("Transport Type", "Streamable HTTP");
await selectOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
const user = userEvent.setup({ delay: null });
await user.type(getServerNameInput(), name);
await user.type(screen.getByPlaceholderText("https://your-mcp-server.com"), "https://example.com/mcp");
await selectAntOption("Authentication", "None");
await selectOption("Authentication", "None");
};
const submitAndReadPayload = async () => {
@ -120,8 +124,9 @@ describe("CreateMCPServer permission toggles reaching the payload", () => {
render(<CreateMCPServer {...defaultProps} />);
await fillMinimalHttpServer("Perm_Server");
const allowAllKeys = await switchFor("Allow All LiteLLM Keys");
await act(async () => {
fireEvent.click(switchFor("Allow All LiteLLM Keys"));
fireEvent.click(allowAllKeys);
});
const payload = await submitAndReadPayload();
@ -133,7 +138,7 @@ describe("CreateMCPServer permission toggles reaching the payload", () => {
render(<CreateMCPServer {...defaultProps} />);
await fillMinimalHttpServer("Perm_Server");
const internalOnly = switchFor("Internal network only");
const internalOnly = await switchFor("Internal network only");
expect(internalOnly).toHaveAttribute("aria-checked", "false");
await act(async () => {

View file

@ -1,9 +1,13 @@
import React, { useState } from "react";
import { Tooltip, Select, Input as AntdInput, InputNumber, Collapse } from "antd";
import { FormProvider, useForm, useWatch } from "react-hook-form";
import { InfoCircleOutlined } from "@ant-design/icons";
import { ChevronDown, Info } from "lucide-react";
import { Button } from "@/components/ui/button";
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog";
import { Input } from "@/components/ui/input";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { createMCPServer, registerMCPServer, storeMCPOAuthUserCredential } from "@/components/networking";
import { setToken } from "@/utils/mcpTokenStore";
@ -14,6 +18,8 @@ import {
MCPServer,
MCPServerCostInfo,
TRANSPORT,
TRANSPORT_ITEMS,
AUTH_TYPE_ITEMS,
getMcpOAuthMode,
MCP_OAUTH2_FLOW_M2M,
isClientForwardedTokenMode,
@ -59,9 +65,8 @@ import {
} from "@/components/common_components/MountedFormField";
import { antdRequired, antdRules } from "@/components/common_components/antdFormRules";
import { allFieldsValue, mountedPaths, resetFields, setFieldsValue, singleBranchChange } from "./mcpFormStore";
import { numberControl, notOnlyWhitespace, selectControl, textControl } from "./mcpFieldRules";
import { numberControl, notOnlyWhitespace, selectControl, selectTriggerControl, textControl } from "./mcpFieldRules";
import mcpLogo from "../../../../../public/assets/logos/mcp_logo.png";
import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog";
export const mcpLogoImg = mcpLogo.src;
@ -122,7 +127,6 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
const [toolNameToDescription, setToolNameToDescription] = useState<Record<string, string>>({});
const [transportType, setTransportType] = useState<string>("");
const [keyTools, setKeyTools] = useState<OpenAPIKeyTool[]>([]);
const [searchValue, setSearchValue] = useState<string>("");
const [oauthAccessToken, setOauthAccessToken] = useState<string | null>(null);
const [logoUrl, setLogoUrl] = useState<string | undefined>(undefined);
const [oauthDocsUrl, setOauthDocsUrl] = useState<string | null>(null);
@ -174,7 +178,6 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
costConfig,
allowedTools,
hasToolAllowlistInteraction,
searchValue,
aliasManuallyEdited,
logoUrl,
authorizedIdentity,
@ -340,9 +343,6 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
if (restored.hasToolAllowlistInteraction !== undefined) {
setHasToolAllowlistInteraction(restored.hasToolAllowlistInteraction);
}
if (restored.searchValue) {
setSearchValue(restored.searchValue);
}
if (restored.aliasManuallyEdited !== undefined) {
setAliasManuallyEdited(restored.aliasManuallyEdited);
}
@ -527,37 +527,13 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
setFormValues(allFieldsValue(form));
};
// Generate options with existing groups and potential new group
const getAccessGroupOptions = () => {
const existingOptions = availableAccessGroups.map((group: string) => ({
value: group,
label: (
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-green-500 rounded-full"></div>
<span className="font-medium">{group}</span>
</div>
),
}));
// If search value doesn't match any existing group and is not empty, add "create new group" option
if (
searchValue &&
!availableAccessGroups.some((group) => group.toLowerCase().includes(searchValue.toLowerCase()))
) {
existingOptions.push({
value: searchValue,
label: (
<div className="flex items-center gap-2">
<div className="w-2 h-2 bg-blue-500 rounded-full"></div>
<span className="font-medium">{searchValue}</span>
<span className="text-gray-400 text-xs ml-1">create new group</span>
</div>
),
});
}
return existingOptions;
};
const handleTransportSelected =
(onChange: (value: string) => void) =>
(value: string | null): void => {
if (value === null) return;
onChange(value);
handleTransportChange(value);
};
// Auto-populate alias from server_name unless manually edited
React.useEffect(() => {
@ -647,31 +623,19 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
<Dialog open={isModalVisible} onOpenChange={(open) => !open && handleCancel()}>
<DialogContent className="top-8 max-h-[calc(100dvh-4rem)] translate-y-0 overflow-y-auto sm:max-w-[1000px]">
<DialogHeader>
<div className="flex items-center pb-4 border-b border-gray-100" style={{ gap: 12 }}>
<div className="flex items-center gap-3 border-b border-border pb-4">
{onBackToDiscovery && (
<button
onClick={onBackToDiscovery}
className="text-sm text-blue-600 hover:text-blue-800 cursor-pointer bg-transparent border-none"
style={{ flexShrink: 0 }}
>
<Button variant="link" size="sm" className="shrink-0 px-0" onClick={onBackToDiscovery}>
&#8592;
</button>
</Button>
)}
<img
src={mcpLogoImg}
alt="MCP Logo"
className="w-8 h-8 object-contain"
style={{
height: "20px",
width: "20px",
objectFit: "contain",
}}
/>
<DialogTitle className="text-xl font-semibold text-gray-900">
<img src={mcpLogoImg} alt="MCP Logo" className="size-5 object-contain" />
<DialogTitle className="text-xl font-semibold">
{isAdmin ? "Add New MCP Server" : "Submit MCP Server for Review"}
</DialogTitle>
</div>
</DialogHeader>
<div className="mt-6">
<FormProvider {...form}>
<MountedFormProvider value={{ control: form.control, registry }}>
@ -687,9 +651,9 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
MCP Server Name
<Tooltip title="Best practice: Use a descriptive name that indicates the server's purpose (e.g., 'GitHub_MCP', 'Email_Service'). Cannot contain spaces or hyphens; use underscores instead. Names must comply with SEP-986 and will be rejected if invalid (https://modelcontextprotocol.io/specification/2025-11-25/server/tools#tool-names).">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="Best practice: Use a descriptive name that indicates the server's purpose (e.g., 'GitHub_MCP', 'Email_Service'). Cannot contain spaces or hyphens; use underscores instead. Names must comply with SEP-986 and will be rejected if invalid (https://modelcontextprotocol.io/specification/2025-11-25/server/tools#tool-names).">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
}
name="server_name"
@ -708,9 +672,9 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Alias
<Tooltip title="A short, unique identifier for this server. Defaults to the server name if not provided. Cannot contain spaces or hyphens; use underscores instead.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="A short, unique identifier for this server. Defaults to the server name if not provided. Cannot contain spaces or hyphens; use underscores instead.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
}
name="alias"
@ -765,19 +729,20 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
>
{(control) => (
<Select
{...selectControl<string>(control)}
placeholder="Select transport"
className="rounded-lg"
size="large"
onChange={(value: string) => {
control.onChange(value);
handleTransportChange(value);
}}
items={TRANSPORT_ITEMS}
value={(control.value as string | undefined) ?? null}
onValueChange={handleTransportSelected(control.onChange)}
>
<Select.Option value="http">Streamable HTTP (Recommended)</Select.Option>
<Select.Option value="sse">Server-Sent Events (SSE)</Select.Option>
<Select.Option value="stdio">Standard Input/Output (stdio)</Select.Option>
<Select.Option value={TRANSPORT.OPENAPI}>OpenAPI Spec</Select.Option>
<SelectTrigger {...selectTriggerControl(control)} className="w-full rounded-lg">
<SelectValue placeholder="Select transport" />
</SelectTrigger>
<SelectContent>
{TRANSPORT_ITEMS.map((item) => (
<SelectItem key={item.value} value={item.value}>
{item.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</MountedFormField>
@ -796,7 +761,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
}}
>
{(control) => (
<AntdInput
<Input
{...textControl(control)}
placeholder="https://your-mcp-server.com"
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
@ -826,135 +791,114 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Max Concurrent Requests (optional)
<Tooltip title="Maximum number of tool calls LiteLLM will run against this server at the same time. Additional calls wait for a free slot. Leave blank for no limit.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="Maximum number of tool calls LiteLLM will run against this server at the same time. Additional calls wait for a free slot. Leave blank for no limit.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
}
name="max_concurrent_requests"
>
{(control) => (
<InputNumber
{...numberControl(control)}
<Input
{...numberControl(control, 0)}
min={1}
precision={0}
step={1}
placeholder="e.g. 10"
style={{ width: "100%" }}
className="rounded-lg"
className="w-full rounded-lg"
/>
)}
</MountedFormField>
{/* Authentication - show for HTTP, SSE, and OpenAPI */}
{transportType !== "stdio" && transportType !== "" && (
<Collapse
defaultActiveKey={["auth"]}
className="mb-4"
items={[
{
key: "auth",
label: <span className="text-sm font-semibold text-gray-700">Authentication</span>,
children: (
<>
<MountedFormField
name="auth_type"
required
rules={{ validate: { required: antdRequired("Please select an auth type") } }}
>
{(control) => (
<Select
{...selectControl<string>(control)}
placeholder="Select auth type"
className="rounded-lg"
size="large"
virtual={false}
>
<Select.Option value="none">None</Select.Option>
<Select.Option value="api_key">API Key</Select.Option>
<Select.Option value="bearer_token">Bearer Token</Select.Option>
<Select.Option value="token">Token</Select.Option>
<Select.Option value="basic">Basic Auth</Select.Option>
<Select.Option value="oauth2">OAuth</Select.Option>
<Select.Option value="oauth2_token_exchange">
OAuth Token Exchange (OBO)
</Select.Option>
<Select.Option value="oauth2_id_jag">ID-JAG (Okta Cross App Access)</Select.Option>
<Select.Option value="aws_sigv4">AWS SigV4 (Bedrock AgentCore MCPs)</Select.Option>
<Select.Option value="true_passthrough">
True Passthrough (no LiteLLM auth)
</Select.Option>
<Select.Option value="oauth_delegate">
OAuth Delegate (client-supplied upstream token)
</Select.Option>
</Select>
)}
</MountedFormField>
<Collapsible defaultOpen className="mb-4">
<CollapsibleTrigger className="group flex w-full items-center justify-between gap-4 py-2 text-left">
<span className="text-sm font-semibold text-gray-700">Authentication settings</span>
<ChevronDown className="size-4 text-gray-500 transition-transform group-data-[panel-open]:rotate-180" />
</CollapsibleTrigger>
<CollapsibleContent keepMounted className="space-y-6 pt-2">
<MountedFormField
label="Authentication"
name="auth_type"
required
rules={{ validate: { required: antdRequired("Please select an auth type") } }}
>
{(control) => (
<Select {...selectControl<string>(control)} items={AUTH_TYPE_ITEMS}>
<SelectTrigger {...selectTriggerControl(control)} className="w-full rounded-lg">
<SelectValue placeholder="Select auth type" />
</SelectTrigger>
<SelectContent>
{AUTH_TYPE_ITEMS.map((item) => (
<SelectItem key={item.value} value={item.value}>
{item.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</MountedFormField>
<TruePassthroughWarning authType={authType} />
<TruePassthroughWarning authType={authType} />
<PassthroughAuthorizeSection
authType={authType}
dcrBridgeInitialChecked
oauthFlow={{
startOAuthFlow,
status: oauthStatus,
error: oauthError,
tokenResponse: oauthTokenResponse,
}}
appMayNotMatchUpstream={appMayNotMatchUpstream}
<PassthroughAuthorizeSection
authType={authType}
dcrBridgeInitialChecked
oauthFlow={{
startOAuthFlow,
status: oauthStatus,
error: oauthError,
tokenResponse: oauthTokenResponse,
}}
appMayNotMatchUpstream={appMayNotMatchUpstream}
/>
{shouldShowAuthValueField && (
<MountedFormField
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Authentication Value
<SimpleTooltip content="Token, password, or header value to send with each request for the selected auth type.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
}
name={["credentials", "auth_value"]}
rules={{
validate: {
notWhitespace: notOnlyWhitespace("Authentication value cannot be empty whitespace"),
},
}}
>
{(control) => (
<PasswordInput
{...textControl(control)}
placeholder="Enter token or secret"
groupClassName="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
/>
)}
</MountedFormField>
)}
{shouldShowAuthValueField && (
<MountedFormField
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Authentication Value
<Tooltip title="Token, password, or header value to send with each request for the selected auth type.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
</span>
}
name={["credentials", "auth_value"]}
rules={{
validate: {
notWhitespace: notOnlyWhitespace(
"Authentication value cannot be empty whitespace",
),
},
}}
>
{(control) => (
<AntdInput.Password
{...textControl(control)}
placeholder="Enter token or secret"
className="rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"
/>
)}
</MountedFormField>
)}
{isOAuthAuthType && (
<OAuthFormFields
isM2M={isM2MFlow}
initialFlowType={OAUTH_FLOW.INTERACTIVE}
docsUrl={oauthDocsUrl}
oauthFlow={{
startOAuthFlow,
status: oauthStatus,
error: oauthError,
tokenResponse: oauthTokenResponse,
}}
/>
)}
{isOAuthAuthType && (
<OAuthFormFields
isM2M={isM2MFlow}
initialFlowType={OAUTH_FLOW.INTERACTIVE}
docsUrl={oauthDocsUrl}
oauthFlow={{
startOAuthFlow,
status: oauthStatus,
error: oauthError,
tokenResponse: oauthTokenResponse,
}}
/>
)}
{isTokenExchangeAuthType && <TokenExchangeFormFields />}
{isTokenExchangeAuthType && <TokenExchangeFormFields />}
{isIdJagAuthType && <IdJagFormFields />}
</>
),
},
]}
/>
{isIdJagAuthType && <IdJagFormFields />}
</CollapsibleContent>
</Collapsible>
)}
{transportType !== "stdio" && transportType !== "" && isAwsSigV4AuthType && <AwsSigV4Fields />}
@ -974,9 +918,6 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
availableAccessGroups={availableAccessGroups}
mcpServer={null}
mountedAuthType={authSectionMounted ? watchedAuthType : undefined}
searchValue={searchValue}
setSearchValue={setSearchValue}
getAccessGroupOptions={getAccessGroupOptions}
/>
</div>

View file

@ -1,7 +1,8 @@
import { Info } from "lucide-react";
import React from "react";
import { Switch, Tooltip } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { Switch } from "@/components/ui/switch";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import { isClientForwardedTokenMode } from "@/components/mcp_tools/types";
import { switchControl } from "./mcpFieldRules";
@ -28,9 +29,9 @@ export default function DcrBridgeToggle({
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Gateway-hosted sign-in (DCR bridge)
<Tooltip title="Lets OAuth-only clients like Claude Desktop register and sign in through the gateway. Turn off to relay the upstream server's own OAuth metadata instead (for clients pre-registered with the upstream IdP).">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="Lets OAuth-only clients like Claude Desktop register and sign in through the gateway. Turn off to relay the upstream server's own OAuth metadata instead (for clients pre-registered with the upstream IdP).">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
}
name="dcr_bridge"

View file

@ -1,7 +1,10 @@
import { CircleMinus, Info, Plus } from "lucide-react";
import React from "react";
import { Input, Select, Tooltip, Typography } from "antd";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { Button } from "@/components/ui/button";
import { InfoCircleOutlined, MinusCircleOutlined, PlusOutlined } from "@ant-design/icons";
import { Input } from "@/components/ui/input";
import { InputGroup, InputGroupAddon, InputGroupInput } from "@/components/ui/input-group";
import { useFieldArray, useFormContext, useWatch } from "react-hook-form";
import {
@ -10,11 +13,9 @@ import {
type MountedFormValues,
} from "@/components/common_components/MountedFormField";
import { antdRequired } from "@/components/common_components/antdFormRules";
import { matchesPattern, selectControl, textControl } from "./mcpFieldRules";
import { matchesPattern, selectControl, selectTriggerControl, textControl } from "./mcpFieldRules";
import { listControl } from "./mcpFormStore";
const { Text } = Typography;
const SCOPE_OPTIONS = [
{ value: "global", label: "Instance" },
{ value: "user", label: "Per-user" },
@ -40,11 +41,9 @@ const EnvVarsSection: React.FC = () => {
return (
<div className="rounded-lg border border-gray-200 bg-gray-50 p-4">
<div className="flex items-center gap-2 mb-1">
<Text strong className="text-sm">
Variables
</Text>
<Tooltip
title={
<strong className="text-sm font-semibold">Variables</strong>
<SimpleTooltip
content={
<>
Define variables you can interpolate in Static Headers or Authentication using{" "}
<code>{"${VAR_NAME}"}</code>. <br />
@ -55,15 +54,15 @@ const EnvVarsSection: React.FC = () => {
</>
}
>
<InfoCircleOutlined className="text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<Info className="size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</div>
<Text className="text-xs text-gray-600 block mb-3">
<span className="mb-3 block text-xs text-gray-600">
Reference these in Static Headers or Authentication as <code>{"${VAR_NAME}"}</code>. For example:{" "}
<code className="bg-white px-1 rounded-sm border border-gray-200">
{"${DB_PROTOCOL}://${CORP_USERNAME}:${CORP_PASSWORD}@${DB_HOSTNAME}"}
</code>
</Text>
</span>
<div className="space-y-2">
{fields.length > 0 && (
@ -97,18 +96,31 @@ const EnvVarsSection: React.FC = () => {
<ScopedValueOrDescription index={index} />
</div>
<MountedFormField name={["env_vars", String(index), "scope"]} className="mb-0 w-40" defaultValue="global">
{(control) => <Select {...selectControl<string>(control)} options={SCOPE_OPTIONS} />}
{(control) => (
<Select {...selectControl<string>(control)} items={SCOPE_OPTIONS}>
<SelectTrigger {...selectTriggerControl(control)} className="w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
{SCOPE_OPTIONS.map((option) => (
<SelectItem key={option.value} value={option.value}>
{option.label}
</SelectItem>
))}
</SelectContent>
</Select>
)}
</MountedFormField>
<div style={{ width: 24, height: 32 }} className="flex items-center justify-center">
<MinusCircleOutlined
<CircleMinus
onClick={() => remove(index)}
className="text-gray-500 hover:text-red-500 cursor-pointer"
className="size-4 text-gray-500 hover:text-red-500 cursor-pointer"
/>
</div>
</div>
))}
<Button variant="outline" className="w-full border-dashed" onClick={() => append({ scope: "global" })}>
<PlusOutlined />
<Plus />
Add Variable
</Button>
</div>
@ -125,19 +137,17 @@ const ScopedValueOrDescription: React.FC<{ index: number }> = ({ index }) => {
return (
<MountedFormField name={["env_vars", String(index), "description"]} className="mb-0">
{(control) => (
<Input
{...textControl(control)}
addonBefore={
<Tooltip title="Per-user variables have no shared value. This text is only a hint shown to each user when they fill in their own value.">
<InputGroup>
<InputGroupAddon>
<SimpleTooltip content="Per-user variables have no shared value. This text is only a hint shown to each user when they fill in their own value.">
<span className="text-xs text-gray-500 cursor-help whitespace-nowrap">
<InfoCircleOutlined className="mr-1" />
<Info className="mr-1 inline size-3 align-text-bottom" />
Hint
</span>
</Tooltip>
}
placeholder="e.g. Your DB username"
styles={{ input: { color: "#9ca3af" } }}
/>
</SimpleTooltip>
</InputGroupAddon>
<InputGroupInput {...textControl(control)} placeholder="e.g. Your DB username" className="text-gray-400" />
</InputGroup>
)}
</MountedFormField>
);

View file

@ -1,10 +1,14 @@
import { Info } from "lucide-react";
import React from "react";
import { Input, Select, Tooltip } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import { antdRequired } from "@/components/common_components/antdFormRules";
import { requiredUnlessSiblingSet, selectControl, textControl } from "./mcpFieldRules";
import { MultiSelect } from "@/components/shared/MultiSelect";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { requiredUnlessSiblingSet, tagsControl, textControl } from "./mcpFieldRules";
interface IdJagFormFieldsProps {
isEditing?: boolean;
@ -15,9 +19,9 @@ const fieldClassName = "rounded-lg border-gray-300 focus:border-blue-500 focus:r
const FieldLabel: React.FC<{ label: string; tooltip: string }> = ({ label, tooltip }) => (
<span className="text-sm font-medium text-gray-700 flex items-center">
{label}
<Tooltip title={tooltip}>
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content={tooltip}>
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
);
@ -75,10 +79,10 @@ const IdJagFormFields: React.FC<IdJagFormFieldsProps> = ({ isEditing = false })
rules={requiredWhenCreating("Client ID is required for ID-JAG")}
>
{(control) => (
<Input.Password
<PasswordInput
{...textControl(control)}
placeholder={`Enter OAuth client ID${placeholderSuffix}`}
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -105,10 +109,10 @@ const IdJagFormFields: React.FC<IdJagFormFieldsProps> = ({ isEditing = false })
}
>
{(control) => (
<Input.Password
<PasswordInput
{...textControl(control)}
placeholder={`Enter OAuth client secret${placeholderSuffix}`}
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -122,7 +126,7 @@ const IdJagFormFields: React.FC<IdJagFormFieldsProps> = ({ isEditing = false })
name={PRIVATE_KEY_PATH}
>
{(control) => (
<Input.TextArea
<Textarea
{...textControl(control)}
rows={3}
placeholder={`-----BEGIN PRIVATE KEY-----${placeholderSuffix}`}
@ -199,16 +203,7 @@ const IdJagFormFields: React.FC<IdJagFormFieldsProps> = ({ isEditing = false })
label={<FieldLabel label="Scopes (optional)" tooltip="Scopes requested on leg 1 of the exchange." />}
name={["credentials", "scopes"]}
>
{(control) => (
<Select
{...selectControl(control)}
mode="tags"
tokenSeparators={[","]}
placeholder="Add scopes"
className="rounded-lg"
size="large"
/>
)}
{(control) => <MultiSelect {...tagsControl(control)} placeholder="Add scopes" className="rounded-lg" />}
</MountedFormField>
</>
);

View file

@ -31,11 +31,7 @@ describe("MCPPermissionManagement", () => {
it("should default allow_all_keys switch to unchecked for new servers", async () => {
renderWithForm();
await expandPanel();
// Find the switch associated with "Allow All LiteLLM Keys" text
// The first switch in the component is for allow_all_keys
const switches = screen.getAllByRole("switch");
const toggle = switches[0];
expect(toggle).not.toBeChecked();
expect(screen.getByRole("switch", { name: "Allow All LiteLLM Keys" })).not.toBeChecked();
});
const renderWithInitialValues = (initialValues: Record<string, unknown>, props = {}) =>

View file

@ -1,31 +1,28 @@
import React, { useEffect } from "react";
import { Select, Tooltip, Collapse, Input, Space, Switch } from "antd";
import { TriangleAlert } from "lucide-react";
import { MultiSelect } from "@/components/shared/MultiSelect";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { ChevronRight, CircleMinus, Info, Plus, TriangleAlert, X } from "lucide-react";
import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert";
import { Button } from "@/components/ui/button";
import { InfoCircleOutlined, MinusCircleOutlined, PlusOutlined } from "@ant-design/icons";
import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible";
import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group";
import { Switch } from "@/components/ui/switch";
import { useFieldArray, useFormContext, useWatch } from "react-hook-form";
import { MCPServer, AUTH_TYPE } from "@/components/mcp_tools/types";
import {
MountedFormField,
useMountedName,
type MountedFieldControlProps,
type MountedFormValues,
} from "@/components/common_components/MountedFormField";
import { antdRequired } from "@/components/common_components/antdFormRules";
import { Field, FieldLabel } from "@/components/shared/form/field";
import { invertedSwitchControl, selectControl, switchControl, textControl } from "./mcpFieldRules";
import { invertedSwitchControl, switchControl, tagsControl, textControl } from "./mcpFieldRules";
import { listControl } from "./mcpFormStore";
const { Panel } = Collapse;
interface MCPPermissionManagementProps {
availableAccessGroups: string[];
mcpServer: MCPServer | null;
searchValue: string;
setSearchValue: (value: string) => void;
getAccessGroupOptions: () => Array<{
value: string;
label: React.ReactNode;
}>;
/**
* The auth type as seen through the gate that mounts the auth_type field.
* Callers pass undefined whenever that field is unmounted, because both
@ -35,6 +32,26 @@ interface MCPPermissionManagementProps {
mountedAuthType: string | null | undefined;
}
const ClearableInput: React.FC<{
control: MountedFieldControlProps;
placeholder: string;
clearLabel: string;
}> = ({ control, placeholder, clearLabel }) => {
const text = textControl(control);
return (
<InputGroup className="rounded-lg">
<InputGroupInput {...text} placeholder={placeholder} />
{text.value !== "" && (
<InputGroupAddon align="inline-end">
<InputGroupButton size="icon-xs" aria-label={clearLabel} onClick={() => control.onChange("")}>
<X />
</InputGroupButton>
</InputGroupAddon>
)}
</InputGroup>
);
};
const StaticHeadersFieldArray: React.FC = () => {
const { control } = useFormContext<MountedFormValues>();
const { fields, append, remove } = useFieldArray({ control: listControl(control), name: "static_headers" });
@ -43,19 +60,17 @@ const StaticHeadersFieldArray: React.FC = () => {
return (
<div className="space-y-3">
{fields.map((item, index) => (
<Space key={item.id} className="flex w-full" align="baseline" size="middle">
<div key={item.id} className="flex w-full items-baseline gap-4">
<MountedFormField
name={["static_headers", String(index), "header"]}
className="flex-1"
rules={{ validate: { required: antdRequired("Header name is required") } }}
>
{(headerControl) => (
<Input
{...textControl(headerControl)}
size="large"
allowClear
className="rounded-lg"
<ClearableInput
control={headerControl}
placeholder="Header name (e.g., X-API-Key)"
clearLabel="Clear header name"
/>
)}
</MountedFormField>
@ -65,23 +80,17 @@ const StaticHeadersFieldArray: React.FC = () => {
rules={{ validate: { required: antdRequired("Header value is required") } }}
>
{(valueControl) => (
<Input
{...textControl(valueControl)}
size="large"
allowClear
className="rounded-lg"
placeholder="Header value"
/>
<ClearableInput control={valueControl} placeholder="Header value" clearLabel="Clear header value" />
)}
</MountedFormField>
<MinusCircleOutlined
<CircleMinus
onClick={() => remove(index)}
className="text-gray-500 hover:text-red-500 cursor-pointer"
className="size-4 text-gray-500 hover:text-red-500 cursor-pointer"
/>
</Space>
</div>
))}
<Button variant="outline" className="w-full border-dashed" onClick={() => append({})}>
<PlusOutlined />
<Plus />
Add Static Header
</Button>
</div>
@ -91,9 +100,6 @@ const StaticHeadersFieldArray: React.FC = () => {
const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
availableAccessGroups,
mcpServer,
searchValue,
setSearchValue,
getAccessGroupOptions,
mountedAuthType,
}) => {
const { setValue } = useFormContext<MountedFormValues>();
@ -175,36 +181,35 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
}, [canEnableOAuthPassthrough, setValue]);
return (
<Collapse className="bg-gray-50 border border-gray-200 rounded-lg" expandIconPosition="end" ghost={false}>
<Panel
header={
<div className="flex items-center">
<div className="flex items-center space-x-2">
<div className="w-2 h-2 bg-blue-500 rounded-full"></div>
<h3 className="text-lg font-semibold text-gray-900">Permission Management / Access Control</h3>
</div>
<p className="text-sm text-gray-600 ml-4">Configure access permissions and security settings (Optional)</p>
</div>
}
key="permissions"
className="border-0"
forceRender
>
<Collapsible className="bg-gray-50 border border-gray-200 rounded-lg">
<CollapsibleTrigger className="group flex w-full items-center justify-between gap-4 p-4 text-left">
<span className="flex items-center">
<span className="flex items-center space-x-2">
<span className="w-2 h-2 bg-blue-500 rounded-full"></span>
<span className="text-lg font-semibold text-gray-900">Permission Management / Access Control</span>
</span>
<span className="text-sm text-gray-600 ml-4">
Configure access permissions and security settings (Optional)
</span>
</span>
<ChevronRight className="size-4 shrink-0 text-muted-foreground transition-transform group-data-panel-open:rotate-90" />
</CollapsibleTrigger>
<CollapsibleContent keepMounted className="px-4 pb-4">
<div className="space-y-6 pt-4">
<div className="flex items-start justify-between gap-4">
<div>
<span className="text-sm font-medium text-gray-700 flex items-center">
Allow All LiteLLM Keys
<Tooltip title="When enabled, every API key can access this MCP server.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="When enabled, every API key can access this MCP server.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
<p className="text-sm text-gray-600 mt-1">
Enable if this server should be &quot;public&quot; to all keys.
</p>
</div>
<MountedFormField name="allow_all_keys" defaultValue={mcpServer?.allow_all_keys ?? false} className="mb-0">
{(control) => <Switch {...switchControl(control)} />}
{(control) => <Switch aria-label="Allow All LiteLLM Keys" {...switchControl(control)} />}
</MountedFormField>
</div>
@ -212,16 +217,16 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
<div>
<span className="text-sm font-medium text-gray-700 flex items-center">
Internal network only
<Tooltip title="When on, only requests from within your internal network are accepted. Turn off to allow external clients (other clusters, ChatGPT, etc). API key authentication is always required regardless of this setting.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="When on, only requests from within your internal network are accepted. Turn off to allow external clients (other clusters, ChatGPT, etc). API key authentication is always required regardless of this setting.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
<p className="text-sm text-gray-600 mt-1">
Turn on to restrict access to callers within your internal network only.
</p>
</div>
<MountedFormField name="available_on_public_internet" defaultValue={true} className="mb-0">
{(control) => <Switch {...invertedSwitchControl(control)} />}
{(control) => <Switch aria-label="Internal network only" {...invertedSwitchControl(control)} />}
</MountedFormField>
</div>
@ -230,9 +235,9 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
<div>
<span className="text-sm font-medium text-gray-700 flex items-center">
Delegate auth to upstream (PKCE passthrough)
<Tooltip title="When on, LiteLLM skips its own API key/SSO check for this server and lets the client complete PKCE directly with the upstream MCP server. Only honored when Auth Type is oauth2. No spend tracking or per-key rate limiting will run on this route.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="When on, LiteLLM skips its own API key/SSO check for this server and lets the client complete PKCE directly with the upstream MCP server. Only honored when Auth Type is oauth2. No spend tracking or per-key rate limiting will run on this route.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
<p className="text-sm text-gray-600 mt-1">
Bypass LiteLLM auth so clients authenticate directly with the upstream OAuth MCP server.
@ -243,7 +248,9 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
defaultValue={mcpServer?.delegate_auth_to_upstream ?? false}
className="mb-0"
>
{(control) => <Switch {...switchControl(control)} />}
{(control) => (
<Switch aria-label="Delegate auth to upstream (PKCE passthrough)" {...switchControl(control)} />
)}
</MountedFormField>
</div>
)}
@ -253,9 +260,9 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
<div>
<span className="text-sm font-medium text-gray-700 flex items-center">
OAuth pass-through
<Tooltip title="When on, this server is treated as an OAuth pass-through: the gateway proxies the upstream /.well-known/oauth-protected-resource metadata, emits spec-compliant 401 challenges when no bearer is supplied, and propagates upstream 401/403 responses. Only honored when Auth Type is None and 'Authorization' is in Extra Headers.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="When on, this server is treated as an OAuth pass-through: the gateway proxies the upstream /.well-known/oauth-protected-resource metadata, emits spec-compliant 401 challenges when no bearer is supplied, and propagates upstream 401/403 responses. Only honored when Auth Type is None and 'Authorization' is in Extra Headers.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
<p className="text-sm text-gray-600 mt-1">
Forward upstream OAuth discovery and 401 challenges so clients negotiate OAuth directly with the
@ -267,7 +274,7 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
defaultValue={mcpServer?.oauth_passthrough ?? false}
className="mb-0"
>
{(control) => <Switch {...switchControl(control)} />}
{(control) => <Switch aria-label="OAuth pass-through" {...switchControl(control)} />}
</MountedFormField>
</div>
)}
@ -288,27 +295,20 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
MCP Access Groups
<Tooltip title="Specify access groups for this MCP server. Users must be in at least one of these groups to access the server.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="Specify access groups for this MCP server. Users must be in at least one of these groups to access the server.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
}
name="mcp_access_groups"
className="mb-4"
>
{(control) => (
<Select
{...selectControl(control)}
mode="tags"
showSearch
<MultiSelect
{...tagsControl(control)}
options={availableAccessGroups.map((group) => ({ label: group, value: group }))}
placeholder="Select existing groups or type to create new ones"
optionFilterProp="value"
filterOption={(input, option) => (option?.value ?? "").toLowerCase().includes(input.toLowerCase())}
onSearch={(value) => setSearchValue(value)}
tokenSeparators={[","]}
options={getAccessGroupOptions()}
maxTagCount="responsive"
allowClear
className="rounded-lg"
/>
)}
</MountedFormField>
@ -317,9 +317,9 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
label={
<span className="text-sm font-medium text-gray-700 flex items-center">
Extra Headers
<Tooltip title="Forward custom headers from incoming requests to this MCP server (e.g., Authorization, X-Custom-Header, User-Agent)">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="Forward custom headers from incoming requests to this MCP server (e.g., Authorization, X-Custom-Header, User-Agent)">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
{mcpServer?.extra_headers && mcpServer.extra_headers.length > 0 && (
<span className="ml-2 text-xs bg-blue-100 text-blue-700 px-2 py-1 rounded-full">
{mcpServer.extra_headers.length} configured
@ -330,18 +330,14 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
name="extra_headers"
>
{(control) => (
<Select
{...selectControl(control)}
mode="tags"
<MultiSelect
{...tagsControl(control)}
placeholder={
mcpServer?.extra_headers && mcpServer.extra_headers.length > 0
? `Currently: ${mcpServer.extra_headers.join(", ")}`
: "Enter header names (e.g., Authorization, X-Custom-Header)"
}
className="rounded-lg"
size="large"
tokenSeparators={[","]}
allowClear
/>
)}
</MountedFormField>
@ -350,16 +346,16 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
<FieldLabel>
<span className="text-sm font-medium text-gray-700 flex items-center">
Static Headers
<Tooltip title="Send these key-value headers with every request to this MCP server.">
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content="Send these key-value headers with every request to this MCP server.">
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
</FieldLabel>
<StaticHeadersFieldArray />
</Field>
</div>
</Panel>
</Collapse>
</CollapsibleContent>
</Collapsible>
);
};

View file

@ -1,13 +1,24 @@
import { Info } from "lucide-react";
import React from "react";
import { Input as AntdInput, InputNumber, Select, Tooltip } from "antd";
import { InfoCircleOutlined } from "@ant-design/icons";
import { MultiSelect } from "@/components/shared/MultiSelect";
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { SimpleTooltip } from "@/components/ui/tooltip";
import { PasswordInput } from "@/components/shared/PasswordInput";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { OAUTH_FLOW } from "@/components/mcp_tools/types";
import { MountedFormField } from "@/components/common_components/MountedFormField";
import { antdRequired } from "@/components/common_components/antdFormRules";
import TokenEndpointAuthMethodField from "./TokenEndpointAuthMethodField";
import { numberControl, parsesAsJson, selectControl, textControl } from "./mcpFieldRules";
import {
numberControl,
parsesAsJson,
selectControl,
selectTriggerControl,
tagsControl,
textControl,
} from "./mcpFieldRules";
interface OAuthFlowStatus {
startOAuthFlow: () => void;
@ -27,6 +38,11 @@ interface OAuthFormFieldsProps {
const fieldClassName = "rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500";
const OAUTH_FLOW_ITEMS = [
{ value: OAUTH_FLOW.M2M, label: "Machine-to-Machine (M2M)" },
{ value: OAUTH_FLOW.INTERACTIVE, label: "Interactive (PKCE)" },
];
const UPSTREAM_RESOURCE_TOOLTIP =
"RFC 8707 resource indicator sent to the authorization server so it mints a token audienced for this MCP server. " +
"Leave blank to send nothing, which is the default and what most providers expect. Use 'auto' to send this server's " +
@ -37,9 +53,9 @@ const UPSTREAM_RESOURCE_TOOLTIP =
const FieldLabel: React.FC<{ label: string; tooltip: string }> = ({ label, tooltip }) => (
<span className="text-sm font-medium text-gray-700 flex items-center">
{label}
<Tooltip title={tooltip}>
<InfoCircleOutlined className="ml-2 text-blue-400 hover:text-blue-600 cursor-help" />
</Tooltip>
<SimpleTooltip content={tooltip}>
<Info className="ml-2 size-4 text-blue-400 hover:text-blue-600 cursor-help" />
</SimpleTooltip>
</span>
);
@ -78,19 +94,24 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
{...(initialFlowType ? { defaultValue: initialFlowType } : {})}
>
{(control) => (
<Select {...selectControl(control)} placeholder="Select OAuth flow" className="rounded-lg" size="large">
<Select.Option value={OAUTH_FLOW.M2M}>
<div>
<span className="font-medium">Machine-to-Machine (M2M)</span>
<span className="text-gray-400 text-xs ml-2">server-to-server, no user interaction</span>
</div>
</Select.Option>
<Select.Option value={OAUTH_FLOW.INTERACTIVE}>
<div>
<span className="font-medium">Interactive (PKCE)</span>
<span className="text-gray-400 text-xs ml-2">browser-based user authorization</span>
</div>
</Select.Option>
<Select {...selectControl<string>(control)} items={OAUTH_FLOW_ITEMS}>
<SelectTrigger {...selectTriggerControl(control)} className="w-full rounded-lg">
<SelectValue placeholder="Select OAuth flow" />
</SelectTrigger>
<SelectContent>
<SelectItem value={OAUTH_FLOW.M2M}>
<div>
<span className="font-medium">Machine-to-Machine (M2M)</span>
<span className="ml-2 text-xs text-muted-foreground">server-to-server, no user interaction</span>
</div>
</SelectItem>
<SelectItem value={OAUTH_FLOW.INTERACTIVE}>
<div>
<span className="font-medium">Interactive (PKCE)</span>
<span className="ml-2 text-xs text-muted-foreground">browser-based user authorization</span>
</div>
</SelectItem>
</SelectContent>
</Select>
)}
</MountedFormField>
@ -104,10 +125,10 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
rules={requiredWhenCreating("Client ID is required for M2M OAuth")}
>
{(control) => (
<AntdInput.Password
<PasswordInput
{...textControl(control)}
placeholder={`Enter OAuth client ID${placeholderSuffix}`}
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -120,10 +141,10 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
rules={requiredWhenCreating("Client Secret is required for M2M OAuth")}
>
{(control) => (
<AntdInput.Password
<PasswordInput
{...textControl(control)}
placeholder={`Enter OAuth client secret${placeholderSuffix}`}
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -151,16 +172,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
}
name={["credentials", "scopes"]}
>
{(control) => (
<Select
{...selectControl(control)}
mode="tags"
tokenSeparators={[","]}
placeholder="Add scopes"
className="rounded-lg"
size="large"
/>
)}
{(control) => <MultiSelect {...tagsControl(control)} placeholder="Add scopes" className="rounded-lg" />}
</MountedFormField>
<UpstreamResourceField />
</>
@ -189,10 +201,10 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
name={["credentials", "client_id"]}
>
{(control) => (
<AntdInput.Password
<PasswordInput
{...textControl(control)}
placeholder={`Enter client ID${placeholderSuffix}`}
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -206,10 +218,10 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
name={["credentials", "client_secret"]}
>
{(control) => (
<AntdInput.Password
<PasswordInput
{...textControl(control)}
placeholder={`Enter client secret${placeholderSuffix}`}
className={fieldClassName}
groupClassName={fieldClassName}
/>
)}
</MountedFormField>
@ -222,16 +234,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
}
name={["credentials", "scopes"]}
>
{(control) => (
<Select
{...selectControl(control)}
mode="tags"
tokenSeparators={[","]}
placeholder="Add scopes"
className="rounded-lg"
size="large"
/>
)}
{(control) => <MultiSelect {...tagsControl(control)} placeholder="Add scopes" className="rounded-lg" />}
</MountedFormField>
<UpstreamResourceField />
<MountedFormField
@ -305,7 +308,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
rules={{ validate: { json: parsesAsJson("Must be valid JSON") } }}
>
{(control) => (
<AntdInput.TextArea
<Textarea
{...textControl(control)}
placeholder={'{\n "organization": "my-org",\n "team.id": "123"\n}'}
rows={4}
@ -323,13 +326,7 @@ const OAuthFormFields: React.FC<OAuthFormFieldsProps> = ({
name="token_storage_ttl_seconds"
>
{(control) => (
<InputNumber
{...numberControl(control)}
min={1}
placeholder="e.g. 3600"
className="w-full rounded-lg"
style={{ width: "100%" }}
/>
<Input {...numberControl(control)} min={1} placeholder="e.g. 3600" className="w-full rounded-lg" />
)}
</MountedFormField>
{oauthFlow && (

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