diff --git a/.circleci/scripts/classify_changes.sh b/.circleci/scripts/classify_changes.sh index 2ca2654a207..7aa0c3544ee 100755 --- a/.circleci/scripts/classify_changes.sh +++ b/.circleci/scripts/classify_changes.sh @@ -1,15 +1,17 @@ #!/usr/bin/env bash set -uo pipefail -category="${1:?usage: classify_changes.sh }" +category="${1:?usage: classify_changes.sh }" 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 ;; diff --git a/.github/CODEOWNERS b/.github/CODEOWNERS index 51d489459d9..118e5491939 100644 --- a/.github/CODEOWNERS +++ b/.github/CODEOWNERS @@ -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 diff --git a/.github/actions/detect-backend-changes/action.yml b/.github/actions/detect-backend-changes/action.yml deleted file mode 100644 index 6e6019e4af9..00000000000 --- a/.github/actions/detect-backend-changes/action.yml +++ /dev/null @@ -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" diff --git a/.github/actions/detect-changes/action.yml b/.github/actions/detect-changes/action.yml new file mode 100644 index 00000000000..9b22d2c23a8 --- /dev/null +++ b/.github/actions/detect-changes/action.yml @@ -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" diff --git a/.github/scripts/detect_backend_changes.sh b/.github/scripts/detect_changes.sh similarity index 80% rename from .github/scripts/detect_backend_changes.sh rename to .github/scripts/detect_changes.sh index 7be49992b37..2d427c92fb5 100755 --- a/.github/scripts/detect_backend_changes.sh +++ b/.github/scripts/detect_changes.sh @@ -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}" ;; diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 6ad240717e8..4f4339a360a 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -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: diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 69495cff896..7a0ae0faaa0 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -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) diff --git a/.github/workflows/test-litellm-ui-build.yml b/.github/workflows/test-litellm-ui-build.yml index 618b0195b5a..b3a07a6e0ff 100644 --- a/.github/workflows/test-litellm-ui-build.yml +++ b/.github/workflows/test-litellm-ui-build.yml @@ -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 diff --git a/.github/workflows/test-litellm-ui-unit.yml b/.github/workflows/test-litellm-ui-unit.yml index 69cbc082d98..a7329432be5 100644 --- a/.github/workflows/test-litellm-ui-unit.yml +++ b/.github/workflows/test-litellm-ui-unit.yml @@ -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 }} diff --git a/.github/workflows/test-mcp.yml b/.github/workflows/test-mcp.yml index 05cc13d0af2..95187ef2835 100644 --- a/.github/workflows/test-mcp.yml +++ b/.github/workflows/test-mcp.yml @@ -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 diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml index 00566e13f53..cb8035aafa1 100644 --- a/.github/workflows/test-unit-documentation.yml +++ b/.github/workflows/test-unit-documentation.yml @@ -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: | diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index 59d56a3f63d..b4c324a2c4c 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -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 diff --git a/ci_cd/generate_model_prices_schema.py b/ci_cd/generate_model_prices_schema.py index 153fbc0fdc2..252e3675329 100644 --- a/ci_cd/generate_model_prices_schema.py +++ b/ci_cd/generate_model_prices_schema.py @@ -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] = { diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260819000000_backfill_spend_log_timestamps/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260819000000_backfill_spend_log_timestamps/migration.sql new file mode 100644 index 00000000000..10003afa9db --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260819000000_backfill_spend_log_timestamps/migration.sql @@ -0,0 +1,4 @@ +UPDATE "LiteLLM_SpendLogs" +SET "created_at" = "endTime", + "updated_at" = "endTime" +WHERE "created_at" > "endTime" + interval '1 hour'; diff --git a/litellm/__init__.py b/litellm/__init__.py index 1ecb04b6e54..00f67ea0ff5 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 ############# diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 7d7380665d3..8f7cd09d364 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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: diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 946110abf9e..91a312b4f45 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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( diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index f73c4942a1c..0793fe20b21 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -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) diff --git a/litellm/litellm_core_utils/ptu_pricing.py b/litellm/litellm_core_utils/ptu_pricing.py new file mode 100644 index 00000000000..a1f8bb36e27 --- /dev/null +++ b/litellm/litellm_core_utils/ptu_pricing.py @@ -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, + ), + } + ) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index f999eae1be6..922769dbbfd 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -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() diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 0161f4fadc9..f74a290d773 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -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( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 46f687c5763..cc522aed1ee 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/llms/vertex_ai/cost_calculator.py b/litellm/llms/vertex_ai/cost_calculator.py index 86a5bb207ec..23cb1e5b580 100644 --- a/litellm/llms/vertex_ai/cost_calculator.py +++ b/litellm/llms/vertex_ai/cost_calculator.py @@ -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, ) diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 445e34966a9..75098515deb 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -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") diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 1c8c267430d..d0eca17272d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index 6231563450b..88803d6442d 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -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: diff --git a/litellm/proxy/hooks/user_management_event_hooks.py b/litellm/proxy/hooks/user_management_event_hooks.py index 929df2a778c..6d978929c05 100644 --- a/litellm/proxy/hooks/user_management_event_hooks.py +++ b/litellm/proxy/hooks/user_management_event_hooks.py @@ -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 ( diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 53d03bc7ba6..385073edc90 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index 12c99477d3c..dde0751d98d 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py index fe9a613656d..86ce336c7a3 100644 --- a/litellm/proxy/management_endpoints/coordination_redis_endpoints.py +++ b/litellm/proxy/management_endpoints/coordination_redis_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 2b88658e1b4..99a85e02b52 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index ab343cbde2d..5f823c5b20d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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( diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 06c32af2dc2..54a591a5e1a 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index ade24d194d2..1b49e2455e4 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -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: diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index 472e25bbc28..0e8c4e1825d 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -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) diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 95632d7cb35..0211d86b83d 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -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( diff --git a/litellm/proxy/management_helpers/audit_logs.py b/litellm/proxy/management_helpers/audit_logs.py index 2b714f06413..ecd6abea3c3 100644 --- a/litellm/proxy/management_helpers/audit_logs.py +++ b/litellm/proxy/management_helpers/audit_logs.py @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 621b3ff9c83..ddcca1d372b 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 5adee476f33..af082f04706 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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/` 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) diff --git a/litellm/proxy/read_model_list.py b/litellm/proxy/read_model_list.py index cdd6680aa40..a1830e7f2bc 100644 --- a/litellm/proxy/read_model_list.py +++ b/litellm/proxy/read_model_list.py @@ -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``. """ diff --git a/litellm/proxy/spend_tracking/ptu_feature_flag.py b/litellm/proxy/spend_tracking/ptu_feature_flag.py index 9078079b676..7f52dfa155d 100644 --- a/litellm/proxy/spend_tracking/ptu_feature_flag.py +++ b/litellm/proxy/spend_tracking/ptu_feature_flag.py @@ -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") diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index 25de9f6d065..381641be96d 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -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__ = ( diff --git a/litellm/proxy/spend_tracking/savings.py b/litellm/proxy/spend_tracking/savings.py index 448723ab3bc..997180efdde 100644 --- a/litellm/proxy/spend_tracking/savings.py +++ b/litellm/proxy/spend_tracking/savings.py @@ -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( diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 3146d8bccfb..0b56f0d8246 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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 ), diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 41187af2bd8..1d042e2521b 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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) diff --git a/litellm/router.py b/litellm/router.py index b25b4f92467..26158ae0a56 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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( diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index 65f1029bf55..cf7bde93360 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -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: diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 5d6e13c7fc0..0cb50cf3a3d 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -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, ), ) diff --git a/litellm/router_strategy/complexity_router/config.py b/litellm/router_strategy/complexity_router/config.py index e82ae991100..73f1378e5f7 100644 --- a/litellm/router_strategy/complexity_router/config.py +++ b/litellm/router_strategy/complexity_router/config.py @@ -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: diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index d1e4b3bb2ce..e89fbbdab65 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -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") diff --git a/litellm/types/router.py b/litellm/types/router.py index 7d1dd1358d5..99a4603ae49 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index bcb6b5e6ec3..96b9343353d 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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]: diff --git a/litellm/utils.py b/litellm/utils.py index fb8e7a9c8f0..d1b0cb882ac 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 1c8c267430d..d0eca17272d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index cd02fde595f..82854a3b717 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -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, diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 1826f56d667..06be96fefdf 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 54016470f8b..c2d73ea467d 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -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 diff --git a/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py b/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py new file mode 100644 index 00000000000..270c59f595f --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py @@ -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 diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py index 6ea9098c228..5c1cd88835f 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_streaming_iterator.py @@ -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") diff --git a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index ce81edf3101..1a3fdb9f652 100644 --- a/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -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 diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 69f2312f203..9e9242137e6 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -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 diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py index bdea37f1358..3bc62d549b0 100644 --- a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py +++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py @@ -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 diff --git a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py index 787e5776897..fa7320b2bc6 100644 --- a/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py +++ b/tests/test_litellm/proxy/hooks/test_key_management_event_hooks.py @@ -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.""" diff --git a/tests/test_litellm/proxy/hooks/test_send_invite_email.py b/tests/test_litellm/proxy/hooks/test_send_invite_email.py index c916af5c128..83fb1f5faba 100644 --- a/tests/test_litellm/proxy/hooks/test_send_invite_email.py +++ b/tests/test_litellm/proxy/hooks/test_send_invite_email.py @@ -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", diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index db39fcd3799..c20b6062e7a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -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. diff --git a/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py b/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py index 48e4353966b..2d54d249713 100644 --- a/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py +++ b/tests/test_litellm/proxy/management_helpers/test_audit_log_callbacks.py @@ -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): diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index b56a8da7c66..a9454854948 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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: """ diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 37a03617a65..58465772b3b 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -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/`` 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 # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py index f41756d2b87..333fd597b49 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_flat_cost_rollup.py @@ -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 diff --git a/tests/test_litellm/proxy/spend_tracking/test_savings.py b/tests/test_litellm/proxy/spend_tracking/test_savings.py index 1435547c434..9006288bdae 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_savings.py +++ b/tests/test_litellm/proxy/spend_tracking/test_savings.py @@ -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( diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index e5add059260..9710dc44e99 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -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"] == "" diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index d6cf0e30139..07877514b69 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -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. diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index 7586bbf551e..006fb79b08f 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -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"} diff --git a/tests/test_litellm/secret_managers/test_secret_managers_main.py b/tests/test_litellm/secret_managers/test_secret_managers_main.py index 3631f640136..acc91b691d4 100644 --- a/tests/test_litellm/secret_managers/test_secret_managers_main.py +++ b/tests/test_litellm/secret_managers/test_secret_managers_main.py @@ -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 diff --git a/tests/test_litellm/test_circleci_path_filter.py b/tests/test_litellm/test_circleci_path_filter.py index b8e763b4979..b427a1a3bd8 100644 --- a/tests/test_litellm/test_circleci_path_filter.py +++ b/tests/test_litellm/test_circleci_path_filter.py @@ -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: diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 75c90d793fe..98938dee62e 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -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 diff --git a/tests/test_litellm/test_detect_backend_changes.py b/tests/test_litellm/test_detect_changes.py similarity index 73% rename from tests/test_litellm/test_detect_backend_changes.py rename to tests/test_litellm/test_detect_changes.py index 527d4143a9c..d8feb2371a3 100644 --- a/tests/test_litellm/test_detect_backend_changes.py +++ b/tests/test_litellm/test_detect_changes.py @@ -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" diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index dfe46d54ab8..4674a8b1dfa 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -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 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 1d9477432ee..60453931595 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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"}, diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 31e8f81f3b3..2d68ea2aa68 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -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 diff --git a/ui/litellm-dashboard/eslint.config.mjs b/ui/litellm-dashboard/eslint.config.mjs index 35603123736..3cbb93f9a48 100644 --- a/ui/litellm-dashboard/eslint.config.mjs +++ b/ui/litellm-dashboard/eslint.config.mjs @@ -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.", + }, ], }, ], diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 7069450a75c..154a19da7f3 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -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", diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 962cabcba7b..3f9c29d2afe 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -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", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx index 8f51177bafe..4fc51910161 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx @@ -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 (
- = ({ visible, onClose, accessTok }`} onClick={() => handleAgentTypeChange(CUSTOM_AGENT_TYPE)} > - +
Custom / Other @@ -846,7 +846,8 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
{/* Agent name chip */}
- } color="purple" className="px-3 py-1 text-sm"> + + {agentName}
@@ -884,7 +885,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
- + Create a new key for this agent

A dedicated key scoped to this agent.

@@ -920,7 +921,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok
- + Assign an existing key

Re-assign a key you already have to this agent.

@@ -958,10 +959,11 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok const renderReadyStep = () => (
- +

Agent Created!

- } color="purple" className="px-3 py-1 text-sm"> + + {createdAgentName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx index 9aeedb29ecc..ad0f04f77aa 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/budgets/_components/budget_panel.tsx @@ -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 = ({ accessToken }) => { return (
- } title="Budgets" subtitle="Spend, TPM and RPM limits you can assign to customers." diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx index 619ccd8d974..271de78c272 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails-monitor/_components/GuardrailConfig.tsx @@ -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
@@ -214,7 +208,7 @@ export function GuardrailConfig({ guardrailName, guardrailType, provider }: Guar

- + Custom Code Override

Replace the built-in guardrail with custom evaluation code

@@ -247,13 +241,13 @@ export function GuardrailConfig({ guardrailName, guardrailType, provider }: Guar
{rerunStatus === "success" && ( - 7/10 would now pass with new config + 7/10 would now pass with new config )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx index d81eaa5a4fc..25f63ca2902 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.test.tsx @@ -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) => , -})); - const baseCard: GuardrailCardInfo = { id: "test-guard", name: "Test Guardrail", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx index b6ee130d50a..4f1ac3e7d7a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.test.tsx @@ -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(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx index f89e8a3248c..8f3ec6185fd 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info.tsx @@ -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 = ({ guardrailId, onClose, value === "MASK" ? "text-blue-600" : "text-red-600" }`} > - {value === "MASK" ? : } + {value === "MASK" ? : } {String(value)}

@@ -667,7 +666,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose,

Guardrail Settings

{isConfigGuardrail && ( - + )} {!isEditing && diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/AwsSigV4Fields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/AwsSigV4Fields.tsx index 09016bda266..77a77596bed 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/AwsSigV4Fields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/AwsSigV4Fields.tsx @@ -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 }) => ( {label} - - - + + + ); @@ -71,10 +73,10 @@ const AwsSigV4Fields: React.FC = () => ( }} > {(control) => ( - )} @@ -94,10 +96,10 @@ const AwsSigV4Fields: React.FC = () => ( }} > {(control) => ( - )} @@ -106,10 +108,10 @@ const AwsSigV4Fields: React.FC = () => ( name={["credentials", "aws_session_token"]} > {(control) => ( - )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx index 11f7a5f7f90..6ab11d1c9ea 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.integration.test.tsx @@ -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(); - 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(); - 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(); - 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(); - 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(); - 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(); - 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(); - 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(); - 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(); - 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(); - 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(); - 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(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.permissions.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.permissions.integration.test.tsx index 788be106d8a..c73a0d63252 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.permissions.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.permissions.integration.test.tsx @@ -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 => { + 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 => { + 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(); 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(); 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 () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx index 5233081fc39..75680d2e03f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/CreateMCPServer.tsx @@ -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 = ({ const [toolNameToDescription, setToolNameToDescription] = useState>({}); const [transportType, setTransportType] = useState(""); const [keyTools, setKeyTools] = useState([]); - const [searchValue, setSearchValue] = useState(""); const [oauthAccessToken, setOauthAccessToken] = useState(null); const [logoUrl, setLogoUrl] = useState(undefined); const [oauthDocsUrl, setOauthDocsUrl] = useState(null); @@ -174,7 +178,6 @@ const CreateMCPServer: React.FC = ({ costConfig, allowedTools, hasToolAllowlistInteraction, - searchValue, aliasManuallyEdited, logoUrl, authorizedIdentity, @@ -340,9 +343,6 @@ const CreateMCPServer: React.FC = ({ 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 = ({ setFormValues(allFieldsValue(form)); }; - // Generate options with existing groups and potential new group - const getAccessGroupOptions = () => { - const existingOptions = availableAccessGroups.map((group: string) => ({ - value: group, - label: ( -
-
- {group} -
- ), - })); - - // 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: ( -
-
- {searchValue} - create new group -
- ), - }); - } - - 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 = ({ !open && handleCancel()}> -
+
{onBackToDiscovery && ( - + )} - MCP Logo - + MCP Logo + {isAdmin ? "Add New MCP Server" : "Submit MCP Server for Review"}
+
@@ -687,9 +651,9 @@ const CreateMCPServer: React.FC = ({ label={ MCP Server Name - - - + + + } name="server_name" @@ -708,9 +672,9 @@ const CreateMCPServer: React.FC = ({ label={ Alias - - - + + + } name="alias" @@ -765,19 +729,20 @@ const CreateMCPServer: React.FC = ({ > {(control) => ( )} @@ -796,7 +761,7 @@ const CreateMCPServer: React.FC = ({ }} > {(control) => ( - = ({ label={ Max Concurrent Requests (optional) - - - + + + } name="max_concurrent_requests" > {(control) => ( - )} {/* Authentication - show for HTTP, SSE, and OpenAPI */} {transportType !== "stdio" && transportType !== "" && ( - Authentication, - children: ( - <> - - {(control) => ( - - )} - + + + Authentication settings + + + + + {(control) => ( + + )} + - + - + + {shouldShowAuthValueField && ( + + Authentication Value + + + + + } + name={["credentials", "auth_value"]} + rules={{ + validate: { + notWhitespace: notOnlyWhitespace("Authentication value cannot be empty whitespace"), + }, + }} + > + {(control) => ( + + )} + + )} - {shouldShowAuthValueField && ( - - Authentication Value - - - - - } - name={["credentials", "auth_value"]} - rules={{ - validate: { - notWhitespace: notOnlyWhitespace( - "Authentication value cannot be empty whitespace", - ), - }, - }} - > - {(control) => ( - - )} - - )} + {isOAuthAuthType && ( + + )} - {isOAuthAuthType && ( - - )} + {isTokenExchangeAuthType && } - {isTokenExchangeAuthType && } - - {isIdJagAuthType && } - - ), - }, - ]} - /> + {isIdJagAuthType && } + + )} {transportType !== "stdio" && transportType !== "" && isAwsSigV4AuthType && } @@ -974,9 +918,6 @@ const CreateMCPServer: React.FC = ({ availableAccessGroups={availableAccessGroups} mcpServer={null} mountedAuthType={authSectionMounted ? watchedAuthType : undefined} - searchValue={searchValue} - setSearchValue={setSearchValue} - getAccessGroupOptions={getAccessGroupOptions} />
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx index 35b23f9873c..acf65992ce1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/DcrBridgeToggle.tsx @@ -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={ Gateway-hosted sign-in (DCR bridge) - - - + + + } name="dcr_bridge" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx index 23a55d2dc1a..f6f724b666d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/EnvVarsSection.tsx @@ -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 (
- - Variables - - Variables + Define variables you can interpolate in Static Headers or Authentication using{" "} {"${VAR_NAME}"}.
@@ -55,15 +54,15 @@ const EnvVarsSection: React.FC = () => { } > - -
+ +
- + Reference these in Static Headers or Authentication as {"${VAR_NAME}"}. For example:{" "} {"${DB_PROTOCOL}://${CORP_USERNAME}:${CORP_PASSWORD}@${DB_HOSTNAME}"} - +
{fields.length > 0 && ( @@ -97,18 +96,31 @@ const EnvVarsSection: React.FC = () => {
- {(control) => (control)} items={SCOPE_OPTIONS}> + + + + + {SCOPE_OPTIONS.map((option) => ( + + {option.label} + + ))} + + + )}
- remove(index)} - className="text-gray-500 hover:text-red-500 cursor-pointer" + className="size-4 text-gray-500 hover:text-red-500 cursor-pointer" />
))}
@@ -125,19 +137,17 @@ const ScopedValueOrDescription: React.FC<{ index: number }> = ({ index }) => { return ( {(control) => ( - + + + - + Hint - - } - placeholder="e.g. Your DB username" - styles={{ input: { color: "#9ca3af" } }} - /> + + + + )} ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/IdJagFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/IdJagFormFields.tsx index 9a6f4ab571f..d4760a65e6c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/IdJagFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/IdJagFormFields.tsx @@ -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 }) => ( {label} - - - + + + ); @@ -75,10 +79,10 @@ const IdJagFormFields: React.FC = ({ isEditing = false }) rules={requiredWhenCreating("Client ID is required for ID-JAG")} > {(control) => ( - )} @@ -105,10 +109,10 @@ const IdJagFormFields: React.FC = ({ isEditing = false }) } > {(control) => ( - )} @@ -122,7 +126,7 @@ const IdJagFormFields: React.FC = ({ isEditing = false }) name={PRIVATE_KEY_PATH} > {(control) => ( - = ({ isEditing = false }) label={} name={["credentials", "scopes"]} > - {(control) => ( - )} @@ -65,23 +80,17 @@ const StaticHeadersFieldArray: React.FC = () => { rules={{ validate: { required: antdRequired("Header value is required") } }} > {(valueControl) => ( - + )} - remove(index)} - className="text-gray-500 hover:text-red-500 cursor-pointer" + className="size-4 text-gray-500 hover:text-red-500 cursor-pointer" /> - +
))}
@@ -91,9 +100,6 @@ const StaticHeadersFieldArray: React.FC = () => { const MCPPermissionManagement: React.FC = ({ availableAccessGroups, mcpServer, - searchValue, - setSearchValue, - getAccessGroupOptions, mountedAuthType, }) => { const { setValue } = useFormContext(); @@ -175,36 +181,35 @@ const MCPPermissionManagement: React.FC = ({ }, [canEnableOAuthPassthrough, setValue]); return ( - - -
-
-

Permission Management / Access Control

-
-

Configure access permissions and security settings (Optional)

-
- } - key="permissions" - className="border-0" - forceRender - > + + + + + + Permission Management / Access Control + + + Configure access permissions and security settings (Optional) + + + + +
Allow All LiteLLM Keys - - - + + +

Enable if this server should be "public" to all keys.

- {(control) => } + {(control) => }
@@ -212,16 +217,16 @@ const MCPPermissionManagement: React.FC = ({
Internal network only - - - + + +

Turn on to restrict access to callers within your internal network only.

- {(control) => } + {(control) => }
@@ -230,9 +235,9 @@ const MCPPermissionManagement: React.FC = ({
Delegate auth to upstream (PKCE passthrough) - - - + + +

Bypass LiteLLM auth so clients authenticate directly with the upstream OAuth MCP server. @@ -243,7 +248,9 @@ const MCPPermissionManagement: React.FC = ({ defaultValue={mcpServer?.delegate_auth_to_upstream ?? false} className="mb-0" > - {(control) => } + {(control) => ( + + )}

)} @@ -253,9 +260,9 @@ const MCPPermissionManagement: React.FC = ({
OAuth pass-through - - - + + +

Forward upstream OAuth discovery and 401 challenges so clients negotiate OAuth directly with the @@ -267,7 +274,7 @@ const MCPPermissionManagement: React.FC = ({ defaultValue={mcpServer?.oauth_passthrough ?? false} className="mb-0" > - {(control) => } + {(control) => }

)} @@ -288,27 +295,20 @@ const MCPPermissionManagement: React.FC = ({ label={ MCP Access Groups - - - + + + } name="mcp_access_groups" className="mb-4" > {(control) => ( - 0 ? `Currently: ${mcpServer.extra_headers.join(", ")}` : "Enter header names (e.g., Authorization, X-Custom-Header)" } className="rounded-lg" - size="large" - tokenSeparators={[","]} - allowClear /> )} @@ -350,16 +346,16 @@ const MCPPermissionManagement: React.FC = ({ Static Headers - - - + + +
- - + + ); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx index 545efb910c4..e88c78271d3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OAuthFormFields.tsx @@ -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 }) => ( {label} - - - + + + ); @@ -78,19 +94,24 @@ const OAuthFormFields: React.FC = ({ {...(initialFlowType ? { defaultValue: initialFlowType } : {})} > {(control) => ( - (control)} items={OAUTH_FLOW_ITEMS}> + + + + + +
+ Machine-to-Machine (M2M) + server-to-server, no user interaction +
+
+ +
+ Interactive (PKCE) + browser-based user authorization +
+
+
)} @@ -104,10 +125,10 @@ const OAuthFormFields: React.FC = ({ rules={requiredWhenCreating("Client ID is required for M2M OAuth")} > {(control) => ( - )} @@ -120,10 +141,10 @@ const OAuthFormFields: React.FC = ({ rules={requiredWhenCreating("Client Secret is required for M2M OAuth")} > {(control) => ( - )} @@ -151,16 +172,7 @@ const OAuthFormFields: React.FC = ({ } name={["credentials", "scopes"]} > - {(control) => ( - - )} + {(control) => } = ({ rules={{ validate: { json: parsesAsJson("Must be valid JSON") } }} > {(control) => ( - = ({ name="token_storage_ttl_seconds" > {(control) => ( - + )} {oauthFlow && ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx index 78c8bbe73a9..8907595d92a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenAPIFormSection.tsx @@ -1,6 +1,7 @@ +import { Info } from "lucide-react"; import React, { useState } from "react"; -import { Input, Tooltip } from "antd"; -import { InfoCircleOutlined } from "@ant-design/icons"; +import { SimpleTooltip } from "@/components/ui/tooltip"; +import { Input } from "@/components/ui/input"; import { AUTH_TYPE, OAUTH_FLOW } from "@/components/mcp_tools/types"; import { MountedFormField } from "@/components/common_components/MountedFormField"; import { antdRequired } from "@/components/common_components/antdFormRules"; @@ -68,9 +69,9 @@ const OpenAPIFormSection: React.FC = ({ label={ OpenAPI Spec URL - - - + + + } name="spec_path" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenApiByokFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenApiByokFields.tsx index 83c841439f0..982998c0aae 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenApiByokFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/OpenApiByokFields.tsx @@ -1,10 +1,13 @@ +import { Info } from "lucide-react"; import React from "react"; -import { Input, Select, Switch, Tooltip } from "antd"; -import { InfoCircleOutlined } from "@ant-design/icons"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { SimpleTooltip } from "@/components/ui/tooltip"; import { useWatch } from "react-hook-form"; +import { Input } from "@/components/ui/input"; +import { Switch } from "@/components/ui/switch"; import { MountedFormField } from "@/components/common_components/MountedFormField"; -import { selectControl, switchControl, textControl } from "./mcpFieldRules"; +import { switchControl, tagsControl, textControl } from "./mcpFieldRules"; const AUTH_HEADER_FORMATS: Readonly> = { bearer_token: "Authorization: Bearer {key}", @@ -25,9 +28,9 @@ const OpenApiByokFields: React.FC = () => { label={ BYOK (Bring Your Own Key) - - - + + + } name="is_byok" @@ -39,7 +42,7 @@ const OpenApiByokFields: React.FC = () => { <> {hasAuthType && (
- + User keys will be sent as:{" "} @@ -50,7 +53,7 @@ const OpenApiByokFields: React.FC = () => { )} {!authType && (
- + Set the Authentication Type below to specify how user keys are sent (e.g., Bearer Token, API Key header). @@ -61,20 +64,18 @@ const OpenApiByokFields: React.FC = () => { label={ Access Description - - - + + + } name="byok_description" > {(control) => ( - - )} + {(control) => { + const placeholder = isEditing + ? "Leave blank to keep existing (default Client Secret Post)" + : "Default (Client Secret Post)"; + return ( + + ); + }} ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx index ba213b20655..5aaa013aba4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/TokenExchangeFormFields.tsx @@ -1,11 +1,15 @@ +import { Info } from "lucide-react"; import React from "react"; -import { Input, 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 { useWatch } from "react-hook-form"; import { MountedFormField } from "@/components/common_components/MountedFormField"; import { antdRequired } from "@/components/common_components/antdFormRules"; -import { selectControl, textControl } from "./mcpFieldRules"; +import { PasswordInput } from "@/components/shared/PasswordInput"; +import { Input } from "@/components/ui/input"; +import { selectControl, selectTriggerControl, tagsControl, textControl } from "./mcpFieldRules"; interface TokenExchangeFormFieldsProps { isEditing?: boolean; @@ -13,12 +17,17 @@ interface TokenExchangeFormFieldsProps { const fieldClassName = "rounded-lg border-gray-300 focus:border-blue-500 focus:ring-blue-500"; +const TOKEN_EXCHANGE_PROFILE_ITEMS = [ + { value: "rfc8693", label: "RFC 8693 (standard)" }, + { value: "entra_obo", label: "Microsoft Entra OBO" }, +]; + const FieldLabel: React.FC<{ label: string; tooltip: string }> = ({ label, tooltip }) => ( {label} - - - + + + ); @@ -41,13 +50,17 @@ const TokenExchangeFormFields: React.FC = ({ isEdi {...(isEditing ? {} : { defaultValue: "rfc8693" })} > {(control) => ( - (control)} items={TOKEN_EXCHANGE_PROFILE_ITEMS}> + + + + + {TOKEN_EXCHANGE_PROFILE_ITEMS.map((item) => ( + + {item.label} + + ))} + )} @@ -80,10 +93,10 @@ const TokenExchangeFormFields: React.FC = ({ isEdi rules={requiredWhenCreating("Client ID is required for token exchange")} > {(control) => ( - )} @@ -99,10 +112,10 @@ const TokenExchangeFormFields: React.FC = ({ isEdi rules={requiredWhenCreating("Client Secret is required for token exchange")} > {(control) => ( - )} @@ -164,13 +177,10 @@ const TokenExchangeFormFields: React.FC = ({ isEdi } > {(control) => ( - {!field.required && } - {field.prop.enum.map((option) => ( + {prop.enum.map((option) => ( @@ -87,20 +92,20 @@ const ToolArgumentControl: React.FC<{ ); } - if (field.prop.type === "number" || field.prop.type === "integer") { + if (prop.type === "number" || prop.type === "integer") { return ( ); } - if (field.prop.type === "boolean") { + if (prop.type === "boolean") { return ( ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx index 3fb5d602402..26e66c8338c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/ToolTestPanel.test.tsx @@ -1,5 +1,5 @@ import React from "react"; -import { render, screen } from "@testing-library/react"; +import { fireEvent, render, screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import type { UserEvent } from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; @@ -389,3 +389,85 @@ describe("ToolTestPanel schema changes under a stable tool name", () => { expect(screen.getByLabelText("message")).toHaveValue("typed by hand"); }); }); + +describe("ToolTestPanel optional union-typed parameters", () => { + const qaEchoSchema: InputSchema = { + type: "object", + properties: { + message: { type: "string" }, + repeat: { type: "integer", default: 1 }, + loud: { type: "boolean", default: false }, + tags: { anyOf: [{ type: "array", items: { type: "string" } }, { type: "null" }], default: null }, + }, + required: ["message"], + }; + + const runPanel = async (drive: () => void) => { + const onSubmit = vi.fn(); + render( + , + ); + drive(); + await userEvent.setup().click(screen.getByRole("button", { name: "Call Tool" })); + return onSubmit; + }; + + it("renders the JSON textarea for an optional array parameter, not a plain text input", () => { + renderPanel(qaEchoSchema); + + const tags = screen.getByTestId("textarea-tags"); + expect(tags.tagName).toBe("TEXTAREA"); + expect(tags).toHaveValue(""); + expect(screen.getByPlaceholderText("Enter JSON array for tags")).toBe(tags); + expect(screen.queryByPlaceholderText("Enter tags")).not.toBeInTheDocument(); + }); + + it("renders the JSON textarea for an optional object parameter, not a plain text input", () => { + renderPanel({ + type: "object", + properties: { + payload: { anyOf: [{ type: "object", properties: { id: { type: "string" } } }, { type: "null" }] }, + }, + }); + + const payload = screen.getByTestId("textarea-payload"); + expect(payload).toHaveValue(JSON.stringify({ id: "" }, null, 2)); + expect(screen.getByPlaceholderText("Enter JSON object for payload")).toBe(payload); + expect(screen.queryByPlaceholderText("Enter payload")).not.toBeInTheDocument(); + }); + + it("sends an optional array parameter as a real array", async () => { + const onSubmit = await runPanel(() => { + fireEvent.change(screen.getByPlaceholderText("Enter message"), { target: { value: "hi" } }); + fireEvent.change(screen.getByTestId("textarea-tags"), { target: { value: '["a","b"]' } }); + }); + + const expected = { message: "hi", repeat: 1, loud: false, tags: ["a", "b"] }; + expect(onSubmit).toHaveBeenCalledWith(expected); + }); + + it("omits an optional array parameter the user never filled in", async () => { + const onSubmit = await runPanel(() => { + fireEvent.change(screen.getByPlaceholderText("Enter message"), { target: { value: "hi" } }); + }); + + expect(onSubmit).toHaveBeenCalledWith({ message: "hi", repeat: 1, loud: false }); + }); + + it("blocks the call instead of sending comma-separated text as a raw string", async () => { + const onSubmit = await runPanel(() => { + fireEvent.change(screen.getByPlaceholderText("Enter message"), { target: { value: "hi" } }); + fireEvent.change(screen.getByTestId("textarea-tags"), { target: { value: "a,b" } }); + }); + + expect(await screen.findByText("Invalid JSON")).toBeInTheDocument(); + expect(onSubmit).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx index e494af11c04..ea2addcd5f1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx @@ -1,7 +1,5 @@ import React from "react"; -import { Spin, Tag, Typography } from "antd"; import { CircleAlert, Info } from "lucide-react"; -import { Alert, AlertTitle } from "@/components/shared/Alert"; import { useMutation, useQuery } from "@tanstack/react-query"; import { z } from "zod/v4"; import { MCPServer, MCPUserEnvVarsStatus, MCPUserEnvVarSpec } from "@/components/mcp_tools/types"; @@ -9,14 +7,14 @@ import { getMCPUserEnvVars, storeMCPUserEnvVars } from "@/components/networking" import { toast } from "@/lib/toast"; import { FieldGroup } from "@/components/shared/form/field"; import { FormField } from "@/components/shared/form/FormField"; +import { Alert, AlertTitle } from "@/components/shared/Alert"; import { PasswordInput } from "@/components/shared/PasswordInput"; +import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { useZodForm } from "@/lib/forms/useZodForm"; -const { Text } = Typography; - interface UserEnvVarsModalProps { server: MCPServer | null; open: boolean; @@ -57,7 +55,7 @@ const UserEnvVarsForm: React.FC = ({ required, isSaving, o label={ {spec.name} - {spec.is_set && Set} + {spec.is_set && Set} } > @@ -130,21 +128,20 @@ const UserEnvVarsModal: React.FC = ({ server, open, acces const isSaving = saveMutation.isPending; return ( - !nextOpen && onClose()}> + !opened && onClose()}>
- Set your credentials - Per-user + Set your credentials + Per-user
- - {displayName} - + {displayName}
-
+ +
{isLoading ? (
- +
) : isError ? ( @@ -158,11 +155,11 @@ const UserEnvVarsModal: React.FC = ({ server, open, acces ) : ( <> - + These values are private to you. Your admin configured this MCP server to require these per-user credentials. Saved values are never shown back; leave an already-set field blank to keep it, or enter a value to set or change it. - + )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.test.ts index 94544ec6e2d..97373939ad1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.test.ts @@ -13,7 +13,6 @@ const fullSnapshot: CreateUiSnapshot = { costConfig: { default_cost_per_query: 0.02 }, allowedTools: ["search"], hasToolAllowlistInteraction: true, - searchValue: "group-a", aliasManuallyEdited: true, logoUrl: "https://cdn/logo.png", authorizedIdentity: "identity-abc", @@ -82,9 +81,8 @@ describe("createOAuthUiState", () => { }); it("omits falsy scalars so a restore never blanks freshly mounted state", () => { - seedRaw({ searchValue: "", logoUrl: "", transportType: "", modalVisible: false }); + seedRaw({ logoUrl: "", transportType: "", modalVisible: false }); const restored = readCreateUiSnapshot(); - expect(restored).not.toHaveProperty("searchValue"); expect(restored).not.toHaveProperty("logoUrl"); expect(restored).not.toHaveProperty("transportType"); expect(restored).not.toHaveProperty("modalVisible"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.ts index f6475b97830..b50ba0841c2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/createOAuthUiState.ts @@ -14,7 +14,6 @@ export interface CreateUiSnapshot { readonly costConfig: MCPServerCostInfo; readonly allowedTools: readonly string[]; readonly hasToolAllowlistInteraction: boolean; - readonly searchValue: string; readonly aliasManuallyEdited: boolean; readonly logoUrl: string | undefined; readonly authorizedIdentity: string | undefined; @@ -29,7 +28,6 @@ export type RestoredUiSnapshot = { readonly costConfig?: MCPServerCostInfo; readonly allowedTools?: readonly string[]; readonly hasToolAllowlistInteraction?: boolean; - readonly searchValue?: string; readonly aliasManuallyEdited?: boolean; readonly logoUrl?: string; readonly authorizedIdentity?: string; @@ -83,7 +81,6 @@ export const readCreateUiSnapshot = (): RestoredUiSnapshot | null => { ...(typeof parsed.hasToolAllowlistInteraction === "boolean" ? { hasToolAllowlistInteraction: parsed.hasToolAllowlistInteraction } : {}), - ...(parsed.searchValue ? { searchValue: parsed.searchValue } : {}), ...(typeof parsed.aliasManuallyEdited === "boolean" ? { aliasManuallyEdited: parsed.aliasManuallyEdited } : {}), ...(parsed.logoUrl ? { logoUrl: parsed.logoUrl } : {}), }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.test.ts new file mode 100644 index 00000000000..52ab74f7669 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.test.ts @@ -0,0 +1,49 @@ +import { describe, expect, it, vi } from "vitest"; +import type { MountedFieldControlProps } from "@/components/common_components/MountedFormField"; +import { tagsControl } from "./mcpFieldRules"; + +const controlWith = (value: unknown, onChange = vi.fn()): MountedFieldControlProps => + ({ id: "field", name: "field", value, onChange, onBlur: vi.fn() }) as unknown as MountedFieldControlProps; + +// These fields were antd Selects with tokenSeparators={[","]}, so a comma commits a tag as an admin +// types. MultiSelect owns that rule now, which leaves this adapter one job: hand the stored value to +// the input and the edited value back, without rewriting either. Stdio args are process argv, so a +// comma inside one and a deliberately repeated flag both have to survive a round trip. +describe("tagsControl", () => { + it("keeps a stored argument that contains a comma as one argument", () => { + expect(tagsControl(controlWith(["--filter=a,b", "--verbose"])).value).toStrictEqual(["--filter=a,b", "--verbose"]); + }); + + it("keeps a repeated stdio flag rather than collapsing it to one", () => { + expect(tagsControl(controlWith(["-v", "-v"])).value).toStrictEqual(["-v", "-v"]); + }); + + it("offers each repeated tag once, since the dropdown keys its entries by value", () => { + expect(tagsControl(controlWith(["-v", "-v"])).options).toStrictEqual([{ label: "-v", value: "-v" }]); + }); + + it("stores the edited tags exactly as the input committed them", () => { + const onChange = vi.fn(); + tagsControl(controlWith(["npx"], onChange)).onValueChange(["npx", "--header=X-Trace: on"]); + + expect(onChange).toHaveBeenCalledWith(["npx", "--header=X-Trace: on"]); + }); + + it("renders a stored list unchanged", () => { + expect(tagsControl(controlWith(["run", "--flag a"])).value).toStrictEqual(["run", "--flag a"]); + }); + + it("clears to an empty list rather than a blank tag when the field is emptied", () => { + expect(tagsControl(controlWith("")).value).toStrictEqual([]); + expect(tagsControl(controlWith([""])).value).toStrictEqual([]); + }); + + it("ignores a stored value that is neither a string nor a list", () => { + expect(tagsControl(controlWith(null)).value).toStrictEqual([]); + expect(tagsControl(controlWith(42)).value).toStrictEqual([]); + }); + + it("wraps a bare stored string into the single tag the antd field would have shown", () => { + expect(tagsControl(controlWith("Authorization")).value).toStrictEqual(["Authorization"]); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.ts b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.ts index 76ff74eac52..9046cb6d039 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcpFieldRules.ts @@ -1,3 +1,4 @@ +import type React from "react"; import type { Validate } from "react-hook-form"; import type { MountedFieldControlProps, MountedFormValues } from "@/components/common_components/MountedFormField"; @@ -20,27 +21,58 @@ export const textControl = (control: MountedFieldControlProps) => ({ }); export const selectControl = (control: MountedFieldControlProps) => ({ - ...ariaOf(control), - value: control.value as TValue, - onChange: control.onChange, + value: (control.value ?? null) as TValue | null, + onValueChange: control.onChange, }); -export const numberControl = (control: MountedFieldControlProps) => ({ +export const selectTriggerControl = (control: MountedFieldControlProps) => ariaOf(control); + +// These fields were antd Selects carrying tokenSeparators={[","]}, and MultiSelect applies that same +// rule to what an admin types. So this only adapts the stored value: splitting or deduping it here +// would rewrite stdio argv entries and repeated flags that nobody edited. +const toTags = (value: unknown): string[] => + (Array.isArray(value) ? value : [value]).filter( + (entry): entry is string => typeof entry === "string" && entry !== "", + ); + +export const tagsControl = (control: MountedFieldControlProps) => { + const value = toTags(control.value); + return { + id: control.id, + options: [...new Set(value)].map((tag) => ({ label: tag, value: tag })), + value, + onValueChange: control.onChange, + emptyText: "Type to add", + allowCustomValues: true, + }; +}; + +const toNumberOrNull = (raw: string, precision: number | undefined): number | null => { + if (raw.trim() === "") return null; + const parsed = Number(raw); + if (!Number.isFinite(parsed)) return null; + return precision === undefined ? parsed : Number(parsed.toFixed(precision)); +}; + +export const numberControl = (control: MountedFieldControlProps, precision?: number) => ({ ...ariaOf(control), - value: control.value as number | null | undefined, - onChange: control.onChange, + name: control.name, + type: "number" as const, + value: control.value === null || control.value === undefined ? "" : String(control.value), + onChange: (event: React.ChangeEvent) => + control.onChange(toNumberOrNull(event.target.value, precision)), }); export const switchControl = (control: MountedFieldControlProps) => ({ ...ariaOf(control), checked: control.value === true, - onChange: control.onChange, + onCheckedChange: (checked: boolean) => control.onChange(checked), }); export const invertedSwitchControl = (control: MountedFieldControlProps) => ({ ...ariaOf(control), checked: control.value !== true, - onChange: (checked: boolean) => control.onChange(!checked), + onCheckedChange: (checked: boolean) => control.onChange(!checked), }); export const valueAt = (values: MountedFormValues, path: readonly string[]): unknown => diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx index 8aeb94404f2..663fd4c0f5c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx @@ -1,9 +1,11 @@ /* eslint-disable react/no-unescaped-entities */ -import React, { useState } from "react"; -import { Card, Typography, Space, Switch } from "antd"; +import React, { useId, useState } from "react"; import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; import { Button } from "@/components/ui/button"; +import { Card, CardContent } from "@/components/ui/card"; +import { Label } from "@/components/ui/label"; +import { Switch } from "@/components/ui/switch"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { CopyIcon, @@ -20,8 +22,6 @@ import { import { getProxyBaseUrl } from "@/components/networking"; import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; -const { Title, Text } = Typography; - interface CodeBlockProps { code: string; title?: string; @@ -47,6 +47,7 @@ const FeatureCard: React.FC = ({ accessGroups = ["dev-group"], }) => { const [useServerHeader, setUseServerHeader] = useState(false); + const serverHeaderToggleId = useId(); const getHeadersConfig = () => { const headers: Record = { @@ -62,58 +63,65 @@ const FeatureCard: React.FC = ({ }; return ( - -
- {icon} -
- - {title} - - {description} -
-
- {serverName && (title === "Implementation Example" || title === "Configuration") && ( -
-
- - - Limit tools to specific MCP servers or MCP groups by passing the x-mcp-servers header - + + +
+ {icon} +
+
{title}
+ {description}
- {useServerHeader && ( - - - Two Options - -

- Option 1: Get a specific server: "{serverName.replace(/\s+/g, "_")}" -

-

- Option 2: Get a group of MCPs: "dev-group" -

-

- You can also mix both: "Server1,dev-group" -

-
-
- )}
- )} - {React.Children.map(children, (child) => { - if ( - React.isValidElement(child) && - child.props.hasOwnProperty("code") && - child.props.hasOwnProperty("copyKey") - ) { - const code = child.props.code; - if (code && code.includes('"headers":')) { - return React.cloneElement(child, { - code: code.replace(/"headers":\s*{[^}]*}/, `"headers": ${JSON.stringify(getHeadersConfig(), null, 8)}`), - }); + {serverName && (title === "Implementation Example" || title === "Configuration") && ( +
+
+ + +
+ {useServerHeader && ( + + + Two Options + +
+

+ Option 1: Get a specific server: "{serverName.replace(/\s+/g, "_")}" +

+

+ Option 2: Get a group of MCPs: "dev-group" +

+

+ You can also mix both: "Server1,dev-group" +

+
+
+
+ )} +
+ )} + {React.Children.map(children, (child) => { + if ( + React.isValidElement(child) && + child.props.hasOwnProperty("code") && + child.props.hasOwnProperty("copyKey") + ) { + const code = child.props.code; + if (code && code.includes('"headers":')) { + return React.cloneElement(child, { + code: code.replace(/"headers":\s*{[^}]*}/, `"headers": ${JSON.stringify(getHeadersConfig(), null, 8)}`), + }); + } } - } - return child; - })} + return child; + })} +
); }; @@ -147,25 +155,25 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] {title && (
- - {title} - + {title}
)} - - -
{code}
+ + + +
{code}
+
); @@ -182,40 +190,36 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = []
- - {title} - + {title} {children}
); const LiteLLMProxyTab = () => ( - +
- - LiteLLM Proxy API Integration - +

LiteLLM Proxy API Integration

- + Connect to LiteLLM Proxy Responses API for seamless tool integration with multiple model providers - +
- +
} title="Virtual Key Setup" description="Configure your LiteLLM Proxy Virtual Key for authentication" > - +
- Get your Virtual Key from your LiteLLM Proxy dashboard or contact your administrator + Get your Virtual Key from your LiteLLM Proxy dashboard or contact your administrator
- +
= ({ currentServerAccessGroups = [] className="text-xs" /> - - +
+
); const OpenAITab = () => ( - +
- + <h4 className="mb-0 text-xl font-semibold text-blue-900 dark:text-blue-100"> OpenAI Responses API Integration - +
- + Connect OpenAI Responses API to your LiteLLM MCP server for seamless tool integration - +
- +
} title="API Key Setup" description="Configure your OpenAI API key for authentication" > - +
- {/* eslint-disable-next-line react/no-unescaped-entities */} - + Get your API key from the{" "} = ({ currentServerAccessGroups = [] > OpenAI platform - +
- +
= ({ currentServerAccessGroups = [] className="text-xs" /> - - +
+
); const CursorTab = () => ( - +
- - Cursor IDE Integration - +

Cursor IDE Integration

- + Use tools directly from Cursor IDE with LiteLLM MCP. Enable your AI assistant to perform real-world tasks without leaving your coding environment. - +
- - - Setup Instructions - - - - - Use the keyboard shortcut ⇧+⌘+J (Mac) or{" "} - Ctrl+Shift+J (Windows/Linux) - - + + +
Setup Instructions
+
+ + + Use the keyboard shortcut ⇧+⌘+J (Mac) or{" "} + Ctrl+Shift+J (Windows/Linux) + + - - Go to the "MCP Tools" tab and click "New MCP Server" - + + Go to the "MCP Tools" tab and click "New MCP Server" + - - - Copy the JSON configuration below and paste it into Cursor, then save with{" "} - Cmd+S or{" "} - Ctrl+S - - } - title="Configuration" - description="Cursor MCP configuration" - serverName="Zapier Gmail" - accessGroups={["dev-group"]} - > - + + Copy the JSON configuration below and paste it into Cursor, then save with{" "} + Cmd+S or{" "} + Ctrl+S + + } + title="Configuration" + description="Cursor MCP configuration" + serverName="Zapier Gmail" + accessGroups={["dev-group"]} + > + - - - + }`} + copyKey="cursor-config" + className="text-xs" + /> + + +
+
-
+
); const StreamableHTTPTab = () => ( - +
- - Streamable HTTP Transport - +

Streamable HTTP Transport

- + Connect to LiteLLM MCP using HTTP transport. Compatible with any MCP client that supports HTTP streaming. - +
= ({ currentServerAccessGroups = [] title="Universal MCP Connection" description="Use this URL with any MCP client that supports HTTP transport" > - +
- + Each MCP client supports different transports. Refer to your client documentation to determine the appropriate transport method. - +
= ({ currentServerAccessGroups = [] Learn more about MCP transports
-
+
-
+
); return (
- +

Connect to your MCP client

@@ -524,7 +523,7 @@ const MCPConnect: React.FC = ({ currentServerAccessGroups = [] - +

); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx index 96685bb8359..438caa2f5e6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.test.tsx @@ -6,7 +6,7 @@ import MCPServerEdit, { EDIT_OAUTH_UI_STATE_KEY } from "./mcp_server_edit"; import { setSecureItem } from "@/utils/secureStorage"; import * as networking from "@/components/networking"; import { toast } from "@/lib/toast"; -import { selectAntOption } from "./testUtils"; +import { selectOption } from "./testUtils"; vi.mock("@/components/networking", () => ({ updateMCPServer: vi.fn(), @@ -370,7 +370,7 @@ describe("MCPServerEdit (true passthrough warning)", () => { />, ); - await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + await selectOption("Authentication", "True Passthrough (no LiteLLM auth)"); await waitFor(() => { expect(mockOauth.getTemporaryPayload).toBeTruthy(); @@ -407,7 +407,7 @@ describe("MCPServerEdit (auth type switch)", () => { />, ); - await selectAntOption("Authentication", "OAuth Token Exchange (OBO)"); + await selectOption("Authentication", "OAuth Token Exchange (OBO)"); const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); await act(async () => { @@ -484,7 +484,7 @@ describe("MCPServerEdit OAuth token invalidation", () => { }); mockOauth.reset.mockClear(); - await selectAntOption("Transport Type", "Standard Input/Output (stdio)"); + await selectOption("Transport Type", "Standard Input/Output (stdio)"); await waitFor(() => expect(mockOauth.reset).toHaveBeenCalled()); expect(mockRemoveToken).toHaveBeenCalledWith("oauth_server_1", undefined); @@ -602,7 +602,7 @@ describe("MCPServerEdit OAuth token invalidation", () => { }); mockOauth.reset.mockClear(); - await selectAntOption("Transport Type", "Server-Sent Events (SSE)"); + await selectOption("Transport Type", "Server-Sent Events (SSE)"); expect(mockOauth.reset).not.toHaveBeenCalled(); expect(mockRemoveToken).not.toHaveBeenCalled(); @@ -861,7 +861,7 @@ describe("MCPServerEdit (interactive OAuth)", () => { expect(screen.getByText("Token Endpoint Auth Method (optional)")).toBeInTheDocument(); }); - await selectAntOption("Token Endpoint Auth Method (optional)", "Client Secret Basic"); + await selectOption("Token Endpoint Auth Method (optional)", "Client Secret Basic"); const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); await act(async () => { @@ -1695,9 +1695,9 @@ describe("MCPServerEdit (OAuth token persistence on save)", () => { />, ); - await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); + await selectOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); mockOauth.tokenResponse = { access_token: "fresh-tok", token_type: "bearer" }; - await selectAntOption("Authentication", "True Passthrough (no LiteLLM auth)"); + await selectOption("Authentication", "True Passthrough (no LiteLLM auth)"); await waitFor(() => { const withHeaders = vi @@ -1838,7 +1838,7 @@ describe("MCPServerEdit oauth2_flow selector", () => { />, ); - await selectAntOption("OAuth Flow Type", "Machine-to-Machine (M2M)"); + await selectOption("OAuth Flow Type", "Machine-to-Machine (M2M)"); const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); await act(async () => { @@ -1870,7 +1870,7 @@ describe("MCPServerEdit oauth2_flow selector", () => { />, ); - await selectAntOption("OAuth Flow Type", "Interactive (PKCE)"); + await selectOption("OAuth Flow Type", "Interactive (PKCE)"); const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); await act(async () => { @@ -1968,7 +1968,7 @@ describe("MCPServerEdit OAuth flow prefill display", () => { expect(screen.getByText("This server has no OAuth flow set")).toBeInTheDocument(); - await selectAntOption("OAuth Flow Type", "Machine-to-Machine (M2M)"); + await selectOption("OAuth Flow Type", "Machine-to-Machine (M2M)"); await waitFor(() => { expect(screen.queryByText("This server has no OAuth flow set")).not.toBeInTheDocument(); @@ -2062,7 +2062,7 @@ describe("MCPServerEdit (dcr_bridge toggle)", () => { mockOauth.tokenResponse = null; }); - const getDcrToggle = () => document.getElementById("dcr_bridge"); + const getDcrToggle = () => screen.queryByRole("switch", { name: /Gateway-hosted sign-in \(DCR bridge\)/ }); function renderEdit(server: Record) { render( @@ -2184,7 +2184,7 @@ describe("MCPServerEdit (dcr_bridge toggle)", () => { expect(getDcrToggle()).toBeInTheDocument(); }); - await selectAntOption("Authentication", "API Key"); + await selectOption("Authentication", "API Key"); await waitFor(() => { expect(getDcrToggle()).not.toBeInTheDocument(); }); @@ -2210,7 +2210,7 @@ describe("MCPServerEdit (dcr_bridge toggle)", () => { // The field stays mounted across the two client-forwarded modes, so the live toggle value is // preserved rather than forced false by the switch. - await selectAntOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); + await selectOption("Authentication", "OAuth Delegate (client-supplied upstream token)"); await waitFor(() => { expect(getDcrToggle()).toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx index 4b092aa3f80..b42d9ec5f40 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx @@ -1,10 +1,14 @@ import React, { useState, useEffect } from "react"; -import { Select, Tooltip, Input, InputNumber } from "antd"; -import { TriangleAlert } from "lucide-react"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { SimpleTooltip } from "@/components/ui/tooltip"; +import { Info, TriangleAlert } from "lucide-react"; import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; import { FormProvider, useForm } from "react-hook-form"; -import { InfoCircleOutlined } from "@ant-design/icons"; +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 { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { AUTH_TYPE, @@ -21,6 +25,8 @@ import { MCPServer, MCPServerCostInfo, TRANSPORT, + TRANSPORT_ITEMS, + AUTH_TYPE_ITEMS, getMcpOAuthMode, oauth2FlowToFormValue, } from "@/components/mcp_tools/types"; @@ -63,7 +69,15 @@ import { singleBranchChange, useMountedValues, } from "./mcpFormStore"; -import { numberControl, notOnlyWhitespace, parsesAsJsonObject, selectControl, textControl } from "./mcpFieldRules"; +import { + numberControl, + notOnlyWhitespace, + parsesAsJsonObject, + selectControl, + selectTriggerControl, + tagsControl, + textControl, +} from "./mcpFieldRules"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; interface MCPServerEditProps { @@ -151,7 +165,6 @@ const MCPServerEdit: React.FC = ({ const [tools, setTools] = useState([]); const [isLoadingTools, setIsLoadingTools] = useState(false); const [toolsError, setToolsError] = useState(null); - const [searchValue, setSearchValue] = useState(""); const [aliasManuallyEdited, setAliasManuallyEdited] = useState(false); const [removeStoredApp, setRemoveStoredApp] = useState(false); // Set when the upstream identity (url/endpoints) changed while a declared app is present, so the @@ -213,7 +226,6 @@ const MCPServerEdit: React.FC = ({ costConfig, allowedTools, hasToolAllowlistInteraction, - searchValue, aliasManuallyEdited, }), ); @@ -398,9 +410,6 @@ const MCPServerEdit: React.FC = ({ if (typeof parsed.hasToolAllowlistInteraction === "boolean") { setHasToolAllowlistInteraction(parsed.hasToolAllowlistInteraction); } - if (parsed.searchValue) { - setSearchValue(parsed.searchValue); - } if (typeof parsed.aliasManuallyEdited === "boolean") { setAliasManuallyEdited(parsed.aliasManuallyEdited); } @@ -611,38 +620,6 @@ const MCPServerEdit: React.FC = ({ } }; - // Generate options with existing groups and potential new group - const getAccessGroupOptions = () => { - const existingOptions = availableAccessGroups.map((group: string) => ({ - value: group, - label: ( -
-
- {group} -
- ), - })); - - // 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: ( -
-
- {searchValue} - create new group -
- ), - }); - } - - return existingOptions; - }; - const handleTransportChange = (value: string) => { // Clear fields that are not relevant for the selected transport. if (value === "stdio") { @@ -680,6 +657,14 @@ const MCPServerEdit: React.FC = ({ } }; + const handleTransportSelected = + (onChange: (value: string) => void) => + (value: string | null): void => { + if (value === null) return; + onChange(value); + handleTransportChange(value); + }; + const valuesChangeRef = React.useRef(handleFormValuesChange); valuesChangeRef.current = handleFormValuesChange; @@ -831,16 +816,20 @@ const MCPServerEdit: React.FC = ({ > {(control) => ( )} @@ -874,9 +863,9 @@ const MCPServerEdit: React.FC = ({ label={ OpenAPI Spec URL - - - + + + } name="spec_path" @@ -897,21 +886,20 @@ const MCPServerEdit: React.FC = ({ label={ Max Concurrent Requests (optional) - - - + + + } name="max_concurrent_requests" > {(control) => ( - )} @@ -926,20 +914,17 @@ const MCPServerEdit: React.FC = ({ rules={{ validate: { required: antdRequired("Authentication is required") } }} > {(control) => ( - (control)} items={AUTH_TYPE_ITEMS}> + + + + + {AUTH_TYPE_ITEMS.map((item) => ( + + {item.label} + + ))} + )} @@ -985,11 +970,8 @@ const MCPServerEdit: React.FC = ({ {(control) => ( - trigger.parentElement || document.body} - filterOption={(input, option) => (option?.label ?? "").toLowerCase().includes(input.toLowerCase())} + ({ label: m, value: m }))} + value={group.primaryModel ?? ""} + onValueChange={handlePrimaryChange} + placeholder="Select primary model" + emptyText="No models found" + disabled={disablePrimaryModel} + className="h-12" /> {!disablePrimaryModel && !group.primaryModel && (
@@ -112,43 +111,16 @@ export function FallbackGroupConfig({
{/* Add Fallback Input */}
- + handleTierDistancePenaltyChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={0} step={0.1} - style={{ width: "100%" }} + className="w-full" /> - + Score penalty applied per tier-step away from the classified tier. - +
)}
diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx index 26bd9af94a8..340fd811e97 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx @@ -1,6 +1,5 @@ import { renderHook, screen, waitFor, renderWithProviders } from "../../../tests/test-utils"; import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; -import type { UploadProps } from "antd/es/upload"; import { describe, expect, it, vi } from "vitest"; import type { Team } from "../key_team_helpers/key_list"; import type { CredentialItem } from "../networking"; @@ -157,11 +156,6 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi }, ]; - const uploadProps: UploadProps = { - beforeUpload: () => false, - showUploadList: false, - }; - return { form, registry, @@ -176,7 +170,6 @@ const createTestProps = (userRole = "proxy_admin", userId = "user-1", isTeamAdmi showAdvancedSettings: false, teams, credentials, - uploadProps, userRole, userId, }; @@ -318,6 +311,35 @@ describe("AddModelForm", () => { expect(await screen.findByRole("button", { name: "Add Model" })).toBeInTheDocument(); }); + describe("the enterprise gate on the Team-BYOK switch", () => { + const renderForm = async (premiumUser: boolean) => { + const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); + mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", premiumUser)); + renderWithProviders(); + return screen.findByRole("switch", { name: "Team-BYOK Model" }); + }; + + it("explains the gate on hover even though the switch it sits on is disabled", async () => { + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + const teamOnlySwitch = await renderForm(false); + expect(teamOnlySwitch).toHaveAttribute("aria-disabled", "true"); + + await user.hover(teamOnlySwitch); + + expect(await screen.findByText(/enterprise-only feature/)).toBeInTheDocument(); + }); + + it("says nothing on hover once the user is premium", async () => { + const user = userEvent.setup(); + const teamOnlySwitch = await renderForm(true); + expect(teamOnlySwitch).not.toHaveAttribute("aria-disabled", "true"); + + await user.hover(teamOnlySwitch); + + expect(screen.queryByText(/enterprise-only feature/)).not.toBeInTheDocument(); + }); + }); + describe("cache control bindings reach the parent form store", () => { const renderWithForm = async () => { const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); @@ -331,11 +353,11 @@ describe("AddModelForm", () => { user, openCacheControl: async () => { await user.click(await screen.findByText("Advanced Settings")); - await user.click(screen.getByLabelText("Cache Control Injection Points")); + await user.click(screen.getByRole("switch", { name: "Cache Control Injection Points" })); await screen.findByText("Add Injection Point"); }, closeCacheControl: async () => { - await user.click(screen.getByLabelText("Cache Control Injection Points")); + await user.click(screen.getByRole("switch", { name: "Cache Control Injection Points" })); await waitFor(() => expect(screen.queryByText("Add Injection Point")).not.toBeInTheDocument()); }, mountedValues: async (): Promise> => props.mountedValues(), diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index 65efd6aed51..b6dddf43588 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -9,7 +9,6 @@ import { Select as AntdSelect, Card, Col, Row, Tooltip, Typography } from "antd" import { Info } from "lucide-react"; import { Alert, AlertDescription, AlertTitle } from "@/components/shared/Alert"; import { Button } from "@/components/ui/button"; -import type { UploadProps } from "antd/es/upload"; import React, { useEffect, useMemo, useState } from "react"; import { FormProvider, useWatch, type UseFormReturn } from "react-hook-form"; import TeamDropdown from "../common_components/team_dropdown"; @@ -44,7 +43,6 @@ interface AddModelFormProps { providerModels: string[]; setProviderModelsFn: (provider: Providers) => void; getPlaceholder: (provider: Providers) => string; - uploadProps: UploadProps; showAdvancedSettings: boolean; setShowAdvancedSettings: (show: boolean) => void; teams: Team[] | null; @@ -71,7 +69,6 @@ const AddModelForm: React.FC = ({ providerModels, setProviderModelsFn, getPlaceholder, - uploadProps, showAdvancedSettings, setShowAdvancedSettings, teams, @@ -311,7 +308,7 @@ const AddModelForm: React.FC = ({ OR
- + )}
diff --git a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx index be4573f8a85..3c384aed1b3 100644 --- a/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ClassificationMethodConfig.tsx @@ -1,6 +1,13 @@ -import { InfoCircleOutlined } from "@ant-design/icons"; +import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select as AntdSelect, Card, InputNumber, Radio, Space, Switch, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Card, CardContent } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { Label } from "@/components/ui/label"; +import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group"; +import { Switch } from "@/components/ui/switch"; import React from "react"; import ClassifierPromptEditor from "./ClassifierPromptEditor"; import HeuristicScoringConfig from "./HeuristicScoringConfig"; @@ -21,8 +28,6 @@ import { effectiveTierLabel, } from "./ComplexityRouterConfig"; -const { Text } = Typography; - const DEFAULT_SCORING_EXPLANATION = "The router scores each request across 7 dimensions: token count, code presence, reasoning markers, technical " + "terms, simple indicators, multi-step patterns, and question complexity. The weighted score determines the tier:"; @@ -87,36 +92,35 @@ const HowClassificationWorks: React.FC<{ value: ComplexityRouterConfigValue }> = return ( - - How Classification Works - - - {scoringExplanation(value)} - - {ranges && ( -
    -
  • - {effectiveTierLabel("SIMPLE", value.tier_labels)}: Score < {ranges.simpleMedium} -
  • -
  • - {effectiveTierLabel("MEDIUM", value.tier_labels)}: Score {ranges.simpleMedium} -{" "} - {ranges.mediumComplex} -
  • -
  • - {effectiveTierLabel("COMPLEX", value.tier_labels)}: Score {ranges.mediumComplex} -{" "} - {ranges.complexReasoning} -
  • -
  • - {effectiveTierLabel("REASONING", value.tier_labels)}: Score > {ranges.complexReasoning}{" "} - (or 2+ reasoning markers with a score of at least {ranges.reasoningOverrideFloor}) -
  • -
- )} - {!ranges && isError && ( - - The tier score ranges could not be loaded from the proxy. - - )} + + How Classification Works + {scoringExplanation(value)} + {ranges && ( +
    +
  • + {effectiveTierLabel("SIMPLE", value.tier_labels)}: Score < {ranges.simpleMedium} +
  • +
  • + {effectiveTierLabel("MEDIUM", value.tier_labels)}: Score {ranges.simpleMedium} -{" "} + {ranges.mediumComplex} +
  • +
  • + {effectiveTierLabel("COMPLEX", value.tier_labels)}: Score {ranges.mediumComplex} -{" "} + {ranges.complexReasoning} +
  • +
  • + {effectiveTierLabel("REASONING", value.tier_labels)}: Score >{" "} + {ranges.complexReasoning} (or 2+ reasoning markers with a score of at least{" "} + {ranges.reasoningOverrideFloor}) +
  • +
+ )} + {!ranges && isError && ( + + The tier score ranges could not be loaded from the proxy. + + )} +
); }; @@ -247,91 +251,103 @@ const ClassificationMethodConfig: React.FC = ({ return ( <> - handleClassifierTypeChange(e.target.value)} + onValueChange={(classifierType: unknown) => handleClassifierTypeChange(classifierType as ClassifierType)} className="w-full" > - - - Heuristic{" "} - (default) — rule-based scoring, no API calls, <1ms latency - - - LLM Classifier{" "} - — use a model to decide the tier (e.g. a small/fast model) - - - +
+ + +
+ {value.classifier_type === "llm" && (
- - Classifier Model - - Classifier Model + - {classifierModelMissing && ( - - A classifier model is required - - )} + {classifierModelMissing && A classifier model is required}
- - Timeout (ms) - - Timeout (ms) + + handleClassifierTimeoutChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={1} - style={{ width: "100%" }} + className="w-full" /> - + How long the classifier call has before it fails and the fallback below takes over. - +
- Classification Rubric + Classification Rubric - +
- ({ + - + {usesCustomPrompt ? "Not in use: the custom prompt below is the classifier's entire rubric." : CLASSIFICATION_RUBRIC_DESCRIPTIONS[classificationRubric].description} - +
- - Classifier Prompt - + Classifier Prompt = ({ />
- - If the classifier fails - - If the classifier fails + handleClassifierFallbackChange(e.target.value)} + onValueChange={(fallback: unknown) => handleClassifierFallbackChange(fallback as ClassifierFallback)} > - - - Score with the heuristic{" "} - — right when the classifier grades complexity too - - +
+ + +
+
+ Applies when the classifier call errors, times out, or returns an unparseable response. - +
- - Context Window Size - - Context Window Size + + handleClassifierContextWindowSizeChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={0} - style={{ width: "100%" }} + className="w-full" /> - + Number of prior user turns (tool output and harness reminders excluded) sent to the classifier as context, so a referring follow-up like "now do the same for the streaming path" is classified against what it refers to. Set to 0 to send only the current message. - +
- - Context Per-Turn Character Limit - - Context Per-Turn Character Limit + + handleClassifierContextPerTurnCharsChange(event.target.value === "" ? null : event.target.valueAsNumber) + } min={1} - style={{ width: "100%" }} + className="w-full" /> - - Prior turns longer than this are truncated. - + Prior turns longer than this are truncated.
- Include Assistant Turns + Include Assistant Turns - +
- + Let the classifier read the assistant's replies, so difficulty the model stated rather than the user stays visible: a plan the assistant calls complex, approved with "yes", is classified on the work being approved. Context Window Size then counts the last N turns across both roles rather than the last N user turns. - +
)} @@ -429,25 +449,29 @@ const ClassificationMethodConfig: React.FC = ({ {value.classifier_type === "heuristic" && (
- Custom Technical Keywords + Custom Technical Keywords - +
- + Optional: Add terms to the built-in list to improve classification accuracy on the technical dimension. (e.g., udp, kafka, terraform). - - + ({ label: keyword, value: keyword }))} value={customTechnicalKeywords ?? []} - onChange={(keywords: string[]) => onCustomTechnicalKeywordsChange?.(keywords)} - placeholder="Type a keyword and press Enter, or paste a comma-separated list" - tokenSeparators={[","]} - open={false} - suffixIcon={null} - style={{ width: "100%" }} - allowClear + onValueChange={(keywords: string[]) => + onCustomTechnicalKeywordsChange?.( + Array.from( + new Set(keywords.flatMap((keyword) => keyword.split(",").map((part) => part.trim())).filter(Boolean)), + ), + ) + } + placeholder="Type a keyword and press Enter" + emptyText="Type to add a keyword" + allowCustomValues + className="w-full" />
)} diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index d0fda9962e5..5f5ae703b0e 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -280,11 +280,28 @@ describe("ComplexityRouterConfig", () => { ); fireEvent.click(screen.getByText("Advanced: Classification Method")); const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement; - const input = within(keywordsSection).getByRole("combobox"); - fireEvent.change(input, { target: { value: "udp," } }); + await user.type(within(keywordsSection).getByRole("combobox"), "udp"); + await user.click(await screen.findByText('Create "udp"')); expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp"]); }); + it("splits a comma-separated keyword entry into one keyword per token", async () => { + const user = userEvent.setup(); + const onCustomTechnicalKeywordsChange = vi.fn(); + renderWithProviders( + , + ); + fireEvent.click(screen.getByText("Advanced: Classification Method")); + const keywordsSection = screen.getByText("Custom Technical Keywords").closest("div")?.parentElement as HTMLElement; + await user.type(within(keywordsSection).getByRole("combobox"), "udp, kafka ,terraform"); + await user.click(await screen.findByText('Create "udp, kafka ,terraform"')); + expect(onCustomTechnicalKeywordsChange).toHaveBeenCalledWith(["udp", "kafka", "terraform"]); + }); + it("should render an empty state when no keyword tier rules exist", () => { renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); @@ -314,9 +331,7 @@ describe("ComplexityRouterConfig", () => { expect(newRules[0]).toMatchObject({ keywords: [], tier: "COMPLEX" }); }); - // The dropdown is closed, so antd has nothing for Enter to select and the word would only land - // on blur. Submitting used to provide that blur; it no longer can while the row reads as empty. - it("commits a typed keyword on Enter, with the dropdown closed", async () => { + it("commits a typed keyword to the rule it was typed into", async () => { const user = userEvent.setup(); const onKeywordTierRulesChange = vi.fn(); renderWithProviders( @@ -329,7 +344,8 @@ describe("ComplexityRouterConfig", () => { fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); const field = screen.getByText("Keywords 1").closest("div") as HTMLElement; - await user.type(within(field).getByRole("combobox"), "invoice{enter}"); + await user.type(within(field).getByRole("combobox"), "invoice"); + await user.click(await screen.findByText('Create "invoice"')); expect(onKeywordTierRulesChange).toHaveBeenCalledWith([{ id: "rule-1", keywords: ["invoice"], tier: "COMPLEX" }]); }); @@ -482,7 +498,7 @@ describe("ComplexityRouterConfig classifier fallback", () => { }; renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Route to the default model/ })).toBeDisabled(); + expect(screen.getByRole("radio", { name: /Route to the default model/ })).toHaveAttribute("aria-disabled", "true"); }); it("hides the fallback choice for the heuristic classifier, which has nothing to fall back from", () => { @@ -584,8 +600,8 @@ describe("ComplexityRouterConfig classifier rubric", () => { it("records the chat preset the operator picks", async () => { const onChange = openClassificationPanel(llmValue); - fireEvent.mouseDown(screen.getByRole("combobox", { name: "Classification Rubric" })); - await userEvent.click(await screen.findByTitle("Chat")); + await userEvent.click(screen.getByRole("combobox", { name: "Classification Rubric" })); + await userEvent.click(await screen.findByRole("option", { name: "Chat" })); expect(onChange).toHaveBeenCalledWith( expect.objectContaining({ classifier_llm_config: expect.objectContaining({ classification_rubric: "chat" }) }), ); @@ -678,7 +694,7 @@ describe("ComplexityRouterConfig tier labels", () => { />, ); fireEvent.click(screen.getByText("Advanced: Keyword/Semantic Matching")); - expect(screen.getByTitle("Deep")).toBeInTheDocument(); + expect(screen.getByRole("combobox", { name: "Route keyword rule 1 to tier" })).toHaveTextContent("Deep"); }); }); @@ -716,7 +732,7 @@ describe("ComplexityRouterConfig default model", () => { it("shows what the tiers currently imply, so an untouched router still names its default", () => { renderWithProviders(); - expect(screen.getByText("Derived from tiers: gpt-3.5-turbo")).toBeInTheDocument(); + expect(getDefaultModelSelect()).toHaveAttribute("placeholder", "Derived from tiers: gpt-3.5-turbo"); }); it("asks for a model rather than naming a derived one when no tier holds one", () => { @@ -725,7 +741,7 @@ describe("ComplexityRouterConfig default model", () => { tiers: { SIMPLE: [], MEDIUM: [], COMPLEX: [], REASONING: [] }, }; renderWithProviders(); - expect(screen.getByText("Add a model to the Simple or Medium tier")).toBeInTheDocument(); + expect(getDefaultModelSelect()).toHaveAttribute("placeholder", "Add a model to the Simple or Medium tier"); }); it("records a pinned model", async () => { @@ -734,7 +750,7 @@ describe("ComplexityRouterConfig default model", () => { renderWithProviders(); await user.click(getDefaultModelSelect()); - await user.click((await screen.findAllByTitle("claude-3-opus")).slice(-1)[0]); + await user.click(await screen.findByRole("option", { name: "claude-3-opus" })); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ default_model: "claude-3-opus" })); }); @@ -745,8 +761,7 @@ describe("ComplexityRouterConfig default model", () => { const pinned: ComplexityRouterConfigValue = { ...defaultValue, default_model: "claude-3-opus" }; renderWithProviders(); - // eslint-disable-next-line local/no-antd-class-selectors -- antd marks the clear affordance aria-hidden, so no accessible query reaches it - await user.click(document.querySelector(".ant-select-clear") as HTMLElement); + await user.click(screen.getByRole("button", { name: "Clear" })); expect(onChange).toHaveBeenCalledWith(expect.objectContaining({ default_model: undefined })); }); @@ -754,10 +769,7 @@ describe("ComplexityRouterConfig default model", () => { it("shows a pinned model as the selection instead of the tier-derived one", () => { const pinned: ComplexityRouterConfigValue = { ...defaultValue, default_model: "claude-3-opus" }; renderWithProviders(); - expect( - // eslint-disable-next-line local/no-antd-class-selectors -- the tier selects show the same model as a tag, so the assertion has to scope to this select's root, which antd exposes only as a class - within(getDefaultModelSelect().closest(".ant-select") as HTMLElement).getByTitle("claude-3-opus"), - ).toBeInTheDocument(); + expect(getDefaultModelSelect()).toHaveValue("claude-3-opus"); }); it("unlocks the default model fallback on a pin alone, with no tier to derive from", () => { @@ -770,7 +782,7 @@ describe("ComplexityRouterConfig default model", () => { }; renderWithProviders(); fireEvent.click(screen.getByText("Advanced: Classification Method")); - expect(screen.getByRole("radio", { name: /Route to the default model/ })).toBeEnabled(); + expect(screen.getByRole("radio", { name: /Route to the default model/ })).not.toHaveAttribute("aria-disabled"); }); it("names the resolved default on the fallback option, so the destination is not a guess", () => { @@ -827,9 +839,9 @@ describe("plan-mode override", () => { />, ); openPanel(); - fireEvent.mouseDown(await screen.findByRole("combobox", { name: "Plan-mode minimum tier" })); - expect(await screen.findByTitle("Medium")).toBeInTheDocument(); - expect(screen.queryByTitle("Reasoning")).not.toBeInTheDocument(); + await userEvent.click(await screen.findByRole("combobox", { name: "Plan-mode minimum tier" })); + expect(await screen.findByRole("option", { name: "Medium" })).toBeInTheDocument(); + expect(screen.queryByRole("option", { name: "Reasoning" })).not.toBeInTheDocument(); }); it("disables the toggle until some tier has models", async () => { @@ -840,6 +852,6 @@ describe("plan-mode override", () => { />, ); openPanel(); - expect(await screen.findByRole("switch", { name: switchName })).toBeDisabled(); + expect(await screen.findByRole("switch", { name: switchName })).toHaveAttribute("aria-disabled", "true"); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 744c74cecb0..ab6b1d401ce 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -1,6 +1,13 @@ -import { InfoCircleOutlined } from "@ant-design/icons"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select as AntdSelect, Card, Collapse, Divider, Input, Space, Switch, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { ChevronRight, Info, X } from "lucide-react"; +import { Switch } from "@/components/ui/switch"; +import { Card, CardContent } from "@/components/ui/card"; +import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; +import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; +import { Separator } from "@/components/ui/separator"; import React from "react"; import { ModelGroup } from "@/components/llm_calls/fetch_models"; import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig"; @@ -13,8 +20,6 @@ import { type DimensionWeights, type TierBoundaries, type TokenThresholds } from export type { DimensionWeights, TierBoundaries, TokenThresholds }; -const { Text } = Typography; - export const DEFAULT_CLASSIFIER_TIMEOUT_MS = 3000; export const DEFAULT_TIER_DISTANCE_PENALTY = 0.5; export const DEFAULT_CLASSIFIER_CONTEXT_WINDOW_SIZE = 3; @@ -218,6 +223,9 @@ const ComplexityRouterConfig: React.FC = ({ showValidationErrors = false, }) => { const planModeTiers = planModeEligibleTiers(value.tiers); + const planModeTierOptions = tierOptions(value.tier_labels).filter((option) => + (planModeTiers as string[]).includes(option.value), + ); const derivedDefaultModel = resolveComplexityDefaultModel(value.tiers); const defaultModel = resolveComplexityDefaultModel(value.tiers, value.default_model); @@ -251,128 +259,119 @@ const ComplexityRouterConfig: React.FC = ({ return (
- - - Complexity Tier Configuration - +
+

Complexity Tier Configuration

- + - +
- + The complexity router automatically classifies requests by complexity using rule-based scoring (no API calls, <1ms latency). Configure which model(s) handle each tier. - + - + Rename a tier to use your own vocabulary in the dashboard and your spend logs. Renaming doesn't change how requests are classified, and callers never see these names. {value.classifier_type === "llm" && " Your classifier model reads these names, so clearer ones can sharpen its choices."} - + - {TIER_KEYS.map((tier, index) => { - const tierInfo = TIER_DESCRIPTIONS[tier]; - const label = effectiveTierLabel(tier, value.tier_labels); - const tierMissing = showValidationErrors && value.tiers[tier].length === 0; - return ( -
- {index > 0 && } -
-
- - {label} Tier - - - - - - Tier {index + 1} of {TIER_KEYS.length} · {tier} - + + {TIER_KEYS.map((tier, index) => { + const tierInfo = TIER_DESCRIPTIONS[tier]; + const label = effectiveTierLabel(tier, value.tier_labels); + const tierMissing = showValidationErrors && value.tiers[tier].length === 0; + return ( +
+ {index > 0 && } +
+
+ {label} Tier + + + + + Tier {index + 1} of {TIER_KEYS.length} · {tier} + +
+ Examples: {tierInfo.examples} + + handleTierLabelChange(tier, event.target.value)} + placeholder={`Display name (default: ${tierInfo.label})`} + aria-label={`Display name for the ${tierInfo.label} tier`} + /> + {value.tier_labels?.[tier] && ( + + handleTierLabelChange(tier, "")} + > + + + + )} + + handleTierChange(tier, models)} + placeholder={`Select model(s) for ${label.toLowerCase()} queries`} + emptyText="No models found" + className={tierMissing ? "w-full border-destructive" : "w-full"} + /> + {value.tiers[tier].length > 1 && ( + + Multiple models selected — the router randomly picks among them per request (or Thompson-samples + within the pool when adaptive routing is on). + + )} + {tierMissing && The {label} tier is required}
- - Examples: {tierInfo.examples} - - handleTierLabelChange(tier, event.target.value)} - placeholder={`Display name (default: ${tierInfo.label})`} - aria-label={`Display name for the ${tierInfo.label} tier`} - style={{ marginBottom: 8 }} - allowClear - /> - handleTierChange(tier, models)} - placeholder={`Select model(s) for ${label.toLowerCase()} queries`} - showSearch - style={{ width: "100%" }} - options={modelOptions} - status={tierMissing ? "error" : undefined} - /> - {value.tiers[tier].length > 1 && ( - - Multiple models selected — the router randomly picks among them per request (or Thompson-samples - within the pool when adaptive routing is on). - - )} - {tierMissing && ( - - The {label} tier is required - - )}
-
- ); - })} - + ); + })} + -
-
- - Default Model - - - - +
+
+ Default Model + + + +
+ + + Used when the tier the request lands in has no model, and when the classifier fails with "Route to + the default model" selected. +
- - - Used when the tier the request lands in has no model, and when the classifier fails with "Route to the - default model" selected. - -
+ - + - + {[ { key: "classifier", - label: ( - - Advanced: Classification Method - - ), + label: Advanced: Classification Method, children: ( = ({ }, { key: "adaptive", - label: ( - - Advanced: Adaptive Routing - - ), + label: Advanced: Adaptive Routing, children: , }, { key: "affinity", - label: ( - - Advanced: Affinity - - ), + label: Advanced: Affinity, children: ( <>
onChange({ ...value, deployment_affinity: deploymentAffinity })} + onCheckedChange={(deploymentAffinity) => + onChange({ ...value, deployment_affinity: deploymentAffinity }) + } aria-label="Pin a session to one deployment per model group" /> - Pin a session to one deployment per model group + Pin a session to one deployment per model group
- + Keeps a session on the same deployment within a group, so provider prompt caches stay warm. Turn off to load-balance every turn. - +
onChange({ ...value, session_affinity: sessionAffinity })} + onCheckedChange={(sessionAffinity) => onChange({ ...value, session_affinity: sessionAffinity })} aria-label="Pin a session to its first model" /> - Pin a session to its first model + Pin a session to its first model
- + Keeps a session on its first turn's model instead of re-classifying each turn. Also pins the deployment. - + ), }, { key: "plan-mode", - label: ( - - Advanced: Plan-Mode Override - - ), + label: Advanced: Plan-Mode Override, children: ( <>
+ onCheckedChange={(enabled) => onChange({ ...value, plan_mode_min_tier: enabled ? planModeTiers.at(-1) : undefined }) } aria-label="Route plan-mode requests to a minimum tier" /> - Route plan-mode requests to a minimum tier + Route plan-mode requests to a minimum tier
- + Requests from coding agents in plan mode (Claude Code, GitHub Copilot) route to at least this tier. The classifier still wins when it picks higher, and the override only lasts while plan mode is active. {planModeTiers.length === 0 && " Add models to a tier to enable this."} - + {value.plan_mode_min_tier !== undefined && (
- - (planModeTiers as string[]).includes(option.value), - )} - onChange={(tier: string) => onChange({ ...value, plan_mode_min_tier: tier })} - /> + onValueChange={(tier: string | null) => tier && onChange({ ...value, plan_mode_min_tier: tier })} + > + + + + + {planModeTierOptions.map((option) => ( + + {option.label} + + ))} + +
)} @@ -473,23 +469,22 @@ const ComplexityRouterConfig: React.FC = ({ }, { key: "response", - label: ( - - Advanced: Response Format - - ), + label: Advanced: Response Format, children: ( <>
onChange({ ...value, return_raw_model_name: returnRawModelName })} + onCheckedChange={(returnRawModelName) => + onChange({ ...value, return_raw_model_name: returnRawModelName }) + } + aria-label="Return raw model name" /> - Return raw model name + Return raw model name
- + Return the resolved underlying model name in responses instead of the autorouter alias. - + ), }, @@ -497,11 +492,7 @@ const ComplexityRouterConfig: React.FC = ({ ? [ { key: "escalation", - label: ( - - Advanced: Escalation Keywords - - ), + label: Advanced: Escalation Keywords, children: , }, ] @@ -510,11 +501,7 @@ const ComplexityRouterConfig: React.FC = ({ ? [ { key: "keyword-semantic", - label: ( - - Advanced: Keyword/Semantic Matching - - ), + label: Advanced: Keyword/Semantic Matching, children: ( <> {onKeywordTierRulesChange && ( @@ -524,9 +511,7 @@ const ComplexityRouterConfig: React.FC = ({ tierLabels={value.tier_labels} /> )} - {onKeywordTierRulesChange && onSemanticMatchingEnabledChange && ( - - )} + {onKeywordTierRulesChange && onSemanticMatchingEnabledChange && } {onSemanticMatchingEnabledChange && ( = ({ }, ] : []), - ]} - /> + ].map(({ key, label, children }) => ( + + + + {label} + + {children} + + ))} +
); }; diff --git a/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx b/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx index 5d5e2c7b02d..b1bb25deb13 100644 --- a/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx +++ b/ui/litellm-dashboard/src/components/add_model/EscalationKeywords.tsx @@ -1,10 +1,8 @@ -import { InfoCircleOutlined } from "@ant-design/icons"; +import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Select as AntdSelect, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; import React from "react"; -const { Text } = Typography; - export const DEFAULT_ESCALATION_KEYWORDS = ["LITELLM ESCALATE"]; interface EscalationKeywordsProps { @@ -16,28 +14,24 @@ const EscalationKeywords: React.FC = ({ keywords, onCha return (
- - Escalation Keywords - +

Escalation Keywords

- +
- + Optional: when a user message contains one of these phrases, the request is bumped one tier higher than it would otherwise route to. Matching is case-sensitive, so "LITELLM ESCALATE" only fires on the exact, shouted form. Leave empty to disable. - - + ({ label: keyword, value: keyword }))} value={keywords} - onChange={onChange} + onValueChange={onChange} placeholder="e.g., LITELLM ESCALATE" - tokenSeparators={[","]} - open={false} - suffixIcon={null} - style={{ width: "100%" }} - allowClear + emptyText="Type to add a phrase" + allowCustomValues + className="w-full" />
); diff --git a/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx b/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx index 24e4c33a49d..f7edd26f0d8 100644 --- a/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx +++ b/ui/litellm-dashboard/src/components/add_model/KeywordTierRules.tsx @@ -1,14 +1,14 @@ -import { DeleteOutlined, InfoCircleOutlined, PlusOutlined } from "@ant-design/icons"; +import { Inbox, Info, Plus, Trash2 } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { Card, Empty, Select as AntdSelect, Typography } from "antd"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Card, CardContent } from "@/components/ui/card"; import { Button } from "@/components/ui/button"; import React from "react"; import { emptyKeywordTierRuleIndexes } from "./complexity_router_keywords"; import { tierOptions } from "./complexity_router_tiers"; -const { Text } = Typography; - export type ComplexityTier = "SIMPLE" | "MEDIUM" | "COMPLEX" | "REASONING"; export interface KeywordTierRule { @@ -28,29 +28,9 @@ interface KeywordTierRulesProps { // there is no failed attempt left to surface it. const KeywordTierRules: React.FC = ({ rules, onChange, tierLabels }) => { const emptyRuleIndexes = new Set(emptyKeywordTierRuleIndexes(rules)); - const [drafts, setDrafts] = React.useState>({}); - - const setDraft = (id: string, text: string) => setDrafts((current) => ({ ...current, [id]: text })); - - // The dropdown is kept closed, which leaves antd nothing for Enter to select, so a typed keyword - // would only become a tag on blur. Submitting used to supply that blur; the button is disabled - // while the row reads as empty, so Enter has to commit the word itself or the row cannot be filled. - const commitDraft = (rule: KeywordTierRule) => { - const keyword = (drafts[rule.id] ?? "").trim(); - if (!keyword) return; - updateRule(rule.id, { keywords: [...rule.keywords, keyword] }); - setDraft(rule.id, ""); - }; - - const commitDraftOnEnter = (rule: KeywordTierRule) => (event: React.KeyboardEvent) => { - if (event.key !== "Enter") return; - event.preventDefault(); - commitDraft(rule); - }; const replaceKeywords = (rule: KeywordTierRule) => (keywords: string[]) => { updateRule(rule.id, { keywords }); - setDraft(rule.id, ""); }; const addRule = () => { @@ -69,79 +49,81 @@ const KeywordTierRules: React.FC = ({ rules, onChange, ti
- - Keyword Tier Overrides - +

Keyword Tier Overrides

- +
- + Optional: route requests containing specific keywords directly to a tier, e.g. route "invoice, refund, billing" to the medium tier. - + {rules.length === 0 ? ( - + +
+
+
) : (
{rules.map((rule, index) => ( - -
-
- - Keywords {index + 1} - - setDraft(rule.id, text)} - onInputKeyDown={commitDraftOnEnter(rule)} - onBlur={() => commitDraft(rule)} - placeholder="e.g., invoice, refund, billing" - tokenSeparators={[","]} - open={false} - suffixIcon={null} - style={{ width: "100%" }} - allowClear - status={emptyRuleIndexes.has(index) ? "error" : undefined} - /> - {emptyRuleIndexes.has(index) && ( - - At least one keyword is required - - )} + + +
+
+ Keywords {index + 1} + ({ label: keyword, value: keyword }))} + value={rule.keywords} + onValueChange={replaceKeywords(rule)} + placeholder="e.g., invoice, refund, billing" + emptyText="Type to add a keyword" + allowCustomValues + className={emptyRuleIndexes.has(index) ? "w-full border-destructive" : "w-full"} + /> + {emptyRuleIndexes.has(index) && ( + At least one keyword is required + )} +
+
+ Route to tier + +
+
-
- - Route to tier - - updateRule(rule.id, { tier })} - options={tierOptions(tierLabels)} - style={{ width: "100%" }} - /> -
- -
+ ))}
diff --git a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx index 867f6a901e8..9da0bf394e5 100644 --- a/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx +++ b/ui/litellm-dashboard/src/components/add_model/SemanticKeywordMatching.tsx @@ -1,11 +1,11 @@ -import { InfoCircleOutlined } from "@ant-design/icons"; +import { Info } from "lucide-react"; import { SimpleTooltip } from "@/components/ui/tooltip"; -import { InputNumber, Select as AntdSelect, Switch, Typography } from "antd"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Input } from "@/components/ui/input"; +import { Switch } from "@/components/ui/switch"; import React from "react"; import { ModelGroup } from "@/components/llm_calls/fetch_models"; -const { Text } = Typography; - const DEFAULT_MATCH_THRESHOLD = 0.5; interface SemanticKeywordMatchingProps { @@ -41,49 +41,49 @@ const SemanticKeywordMatching: React.FC = ({
- Semantic keyword matching + Semantic keyword matching - +
- + Uses same keyword-tier pairs as above and overrides direct keyword matching. Adds latency based on embedding model network request. - +
- +
{enabled && (
- Embedding model - Embedding model + - {embeddingModelMissing && ( - - An embedding model is required - - )} + {embeddingModelMissing && An embedding model is required}
- Minimum match score - Minimum match score + onMatchThresholdChange(value ?? DEFAULT_MATCH_THRESHOLD)} + onChange={(event) => + onMatchThresholdChange(event.target.value === "" ? DEFAULT_MATCH_THRESHOLD : event.target.valueAsNumber) + } min={0} max={1} step={0.05} - style={{ width: "100%" }} + className="w-full" /> - Match only at or above this similarity score. + Match only at or above this similarity score.
)} diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx index e3a7482b02d..e45422dee08 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.test.tsx @@ -27,7 +27,7 @@ const ALL_FAMILY_MODELS: ModelGroup[] = [ const ANTHROPIC_ONLY_MODEL = ANTHROPIC_TIERS.COMPLEX[0]; const openTemplateDropdown = (): void => { - fireEvent.mouseDown(within(screen.getByTestId("template-selector")).getByRole("combobox")); + fireEvent.click(screen.getByTestId("template-selector")); }; // Detailed Configuration is collapsed by default, so any test reaching into it (a tier select, an @@ -36,15 +36,32 @@ const expandDetailedConfiguration = (): void => { fireEvent.click(screen.getByTestId("detailed-configuration-toggle")); }; -const visibleOptions = (): HTMLElement[] => - // eslint-disable-next-line local/no-antd-class-selectors -- antd puts role="option" only on a hidden mirror list of raw values; the visible options carry no role, no aria-disabled, and only a tooltip in title - Array.from(document.querySelectorAll(".ant-select-item-option")); +const visibleOptions = (): HTMLElement[] => screen.queryAllByRole("option"); const optionByLabel = (label: string): HTMLElement | undefined => visibleOptions().find((el) => el.textContent?.startsWith(label)); -// eslint-disable-next-line local/no-antd-class-selectors -- antd signals option disabled state only through this class -const isOptionDisabled = (option: HTMLElement): boolean => option.classList.contains("ant-select-item-option-disabled"); +const isOptionDisabled = (option: HTMLElement): boolean => option.getAttribute("aria-disabled") === "true"; + +const selectTemplate = async (label: string): Promise => { + await userEvent.click(optionByLabel(label)!); +}; + +// Opens the dropdown only when it is closed, since openTemplateDropdown toggles: waiting on a +// second preset in the same test would otherwise close the list out from under the poll. +const waitForPresetEnabled = async (label: string) => { + if (visibleOptions().length === 0) openTemplateDropdown(); + await waitFor(() => { + expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); + }); +}; + +// The keyword field is a combobox that offers whatever is typed as a "Create ..." entry, so a +// keyword only lands on the rule once that entry is picked. +const addKeyword = async (user: ReturnType, field: HTMLElement, keyword: string) => { + await user.type(within(field).getByRole("combobox"), keyword); + await user.click(await screen.findByText(`Create "${keyword}"`)); +}; const { mockFetchAvailableModels, mockFetchAllModelDeployments } = vi.hoisted(() => ({ mockFetchAvailableModels: vi.fn(), @@ -202,10 +219,7 @@ describe("AddAutoRouterTab", () => { await user.click(screen.getByRole("button", { name: /add keyword rule/i })); expect(screen.getByRole("button", { name: /add auto router/i })).toBeDisabled(); - await user.type( - within(screen.getByText("Keywords 1").closest("div") as HTMLElement).getByRole("combobox"), - "invoice{enter}", - ); + await addKeyword(user, screen.getByText("Keywords 1").closest("div") as HTMLElement, "invoice"); expect(screen.getByRole("button", { name: /add auto router/i })).toBeEnabled(); expect(screen.queryByText("At least one keyword is required")).not.toBeInTheDocument(); @@ -221,10 +235,7 @@ describe("AddAutoRouterTab", () => { expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); - await user.type( - within(screen.getByText("Keywords 1").closest("div") as HTMLElement).getByRole("combobox"), - "invoice{enter}", - ); + await addKeyword(user, screen.getByText("Keywords 1").closest("div") as HTMLElement, "invoice"); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); expect(await screen.findAllByText("At least one keyword is required")).toHaveLength(1); @@ -242,7 +253,7 @@ describe("AddAutoRouterTab", () => { await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); const keywordsField = screen.getByText("Keywords 1").closest("div") as HTMLElement; - await user.type(within(keywordsField).getByRole("combobox"), "invoice{enter}"); + await addKeyword(user, keywordsField, "invoice"); await user.click(screen.getByRole("button", { name: /add auto router/i })); await waitFor(() => expect(handleAddAutoRouterSubmit).toHaveBeenCalled()); @@ -407,7 +418,7 @@ describe("AddAutoRouterTab", () => { await user.click(screen.getByText("Advanced: Keyword/Semantic Matching")); await user.click(screen.getByRole("button", { name: /add keyword rule/i })); const keywordsField = screen.getByText("Keywords 1").closest("div") as HTMLElement; - await user.type(within(keywordsField).getByRole("combobox"), "invoice{enter}"); + await addKeyword(user, keywordsField, "invoice"); await user.click(screen.getByTestId("auto-router-test-routing-btn")); await user.type(await screen.findByTestId("auto-router-routing-test-prompt"), "reconcile this invoice"); @@ -452,17 +463,6 @@ describe("AddAutoRouterTab", () => { }); describe("template presets", () => { - // Opens the dropdown once, then waits out the useQuery load: an open antd Select re-renders its - // already-mounted options in place as state changes, so polling only re-reads the DOM here. - // Re-firing the open/close mousedown on every poll (calling openTemplateDropdown inside the - // waitFor callback) fights the dropdown's own open/close animation and hangs the test. - const waitForPresetEnabled = async (label: string) => { - openTemplateDropdown(); - await waitFor(() => { - expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); - }); - }; - it("disables every preset while the model list is loading", async () => { let resolveModels: (models: ModelGroup[]) => void = () => {}; mockFetchAvailableModels.mockImplementation( @@ -537,7 +537,7 @@ describe("AddAutoRouterTab", () => { renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expect(screen.queryByText("Advanced: Keyword/Semantic Matching")).not.toBeInTheDocument(); expect( @@ -548,11 +548,11 @@ describe("AddAutoRouterTab", () => { ).toBeInTheDocument(); }); - it("expands detailed configuration when Custom Configuration is chosen", () => { + it("expands detailed configuration when Custom Configuration is chosen", async () => { renderWithProviders(); openTemplateDropdown(); - fireEvent.click(optionByLabel("Custom Configuration")!); + await selectTemplate("Custom Configuration"); expect(screen.getByText("Advanced: Keyword/Semantic Matching")).toBeInTheDocument(); }); @@ -561,7 +561,7 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expect(screen.queryByText("Advanced: Keyword/Semantic Matching")).not.toBeInTheDocument(); fireEvent.click(screen.getByTestId("detailed-configuration-toggle")); @@ -578,7 +578,7 @@ describe("AddAutoRouterTab", () => { renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); await user.type(screen.getByPlaceholderText(/smart_router/i), "anthropic-router"); await user.click(screen.getByRole("button", { name: /add auto router/i })); @@ -600,7 +600,7 @@ describe("AddAutoRouterTab", () => { const { container } = renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); fireEvent.change(screen.getByPlaceholderText(/smart_router/i), { target: { value: "stale-model-router" } }); expect(screen.getByRole("button", { name: /add auto router/i })).toBeEnabled(); @@ -622,26 +622,16 @@ describe("AddAutoRouterTab", () => { describe("default model pin", () => { const PINNED_MODEL = "pinned-default-model"; - const waitForPresetEnabled = async (label: string) => { - openTemplateDropdown(); - await waitFor(() => { - expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); - }); - }; - const applyPresetAndPin = async (user: ReturnType) => { await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); // Applying a preset collapses Detailed Configuration, so the default model row is behind it. expandDetailedConfiguration(); const defaultModel = screen.getByRole("combobox", { name: "Default model" }); await user.click(defaultModel); - // antd virtualizes the option list and jsdom gives every row zero height, so options past - // the first window never render. Typing filters the list down to the pin instead of relying - // on its index, which adding a preset to the bundled JSON shifts. await user.type(defaultModel, PINNED_MODEL); - await user.click((await screen.findAllByTitle(PINNED_MODEL)).slice(-1)[0]); + await user.click(await screen.findByRole("option", { name: PINNED_MODEL })); }; beforeEach(() => { @@ -692,19 +682,12 @@ describe("AddAutoRouterTab", () => { mockFetchAvailableModels.mockResolvedValue(ALL_FAMILY_MODELS); }); - const waitForPresetEnabled = async (label: string) => { - openTemplateDropdown(); - await waitFor(() => { - expect(isOptionDisabled(optionByLabel(label)!)).toBe(false); - }); - }; - it("omits plan_mode_min_tier from the payload when never touched", async () => { const user = userEvent.setup(); renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); await user.type(screen.getByPlaceholderText(/smart_router/i), "no-plan-router"); await user.click(screen.getByRole("button", { name: /add auto router/i })); @@ -719,7 +702,7 @@ describe("AddAutoRouterTab", () => { renderWithProviders(); await waitForPresetEnabled("Anthropic Family"); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expandDetailedConfiguration(); await user.click(screen.getByText("Advanced: Plan-Mode Override")); await user.click(await screen.findByRole("switch", { name: "Route plan-mode requests to a minimum tier" })); @@ -773,7 +756,7 @@ describe("AddAutoRouterTab", () => { await waitFor(() => { expect(isOptionDisabled(optionByLabel("Anthropic Family")!)).toBe(false); }); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); expect(screen.getByText("Advanced: Keyword/Semantic Matching")).toBeInTheDocument(); @@ -862,7 +845,7 @@ describe("AddAutoRouterTab", () => { await waitFor(() => { expect(isOptionDisabled(optionByLabel("Anthropic Family")!)).toBe(false); }); - fireEvent.click(optionByLabel("Anthropic Family")!); + await selectTemplate("Anthropic Family"); await user.type(screen.getByPlaceholderText(/smart_router/i), "wildcard-router"); await user.click(screen.getByRole("button", { name: /add auto router/i })); diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index a071b92fcd5..64d1519f915 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -1,13 +1,14 @@ import React, { useEffect, useState } from "react"; import { useQuery } from "@tanstack/react-query"; import { useWatch } from "react-hook-form"; -import { Card, Select as AntdSelect } from "antd"; +import { Card } from "antd"; import { ChevronDown, ChevronRight, CircleHelp } from "lucide-react"; import { z } from "zod/v4"; import { FieldGroup } from "@/components/shared/form/field"; import { FormField } from "@/components/shared/form/FormField"; import { Button } from "@/components/ui/button"; import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { useZodForm } from "@/lib/forms/useZodForm"; @@ -289,6 +290,14 @@ const AddAutoRouterTab: React.FC = ({ [presetAvailability], ); + const templateItems = React.useMemo( + () => [ + ...sortedPresetOptions.map(({ preset }) => ({ value: preset.key, label: preset.label })), + { value: "custom", label: "Custom Configuration" }, + ], + [sortedPresetOptions], + ); + const applyPrefill = (prefill: PresetPrefill) => { setComplexityRouterConfig(prefill.complexityRouterConfig); setCustomTechnicalKeywords(prefill.customTechnicalKeywords); @@ -486,49 +495,52 @@ const AddAutoRouterTab: React.FC = ({
- handlePresetChange(presetKey ?? undefined)} > - {sortedPresetOptions.map(({ preset, availability: presetState }) => { - const disabledHint = presetDisabledHint(presetState); - const isDisabled = disabledHint !== null; - const hintClass = isPresetHintAlarming(presetState) - ? "text-red-500 dark:text-red-400" - : "text-muted-foreground"; - const matchedHint = - presetState.kind === "available" && presetState.viaDeployments ? "Matches your deployments" : null; + + + + + {sortedPresetOptions.map(({ preset, availability: presetState }) => { + const disabledHint = presetDisabledHint(presetState); + const hintClass = isPresetHintAlarming(presetState) + ? "text-red-500 dark:text-red-400" + : "text-muted-foreground"; + const matchedHint = + presetState.kind === "available" && presetState.viaDeployments + ? "Matches your deployments" + : null; - return ( - -
-
{preset.label}
-
{preset.description}
- {disabledHint &&
{disabledHint}
} - {matchedHint && ( -
{matchedHint}
- )} -
-
- ); - })} - -
-
Custom Configuration
-
Define your auto router from scratch
-
-
-
+ return ( + +
+
{preset.label}
+
{preset.description}
+ {disabledHint &&
{disabledHint}
} + {matchedHint && ( +
{matchedHint}
+ )} +
+
+ ); + })} + +
+
Custom Configuration
+
Define your auto router from scratch
+
+
+ + {modelsUnverifiable && (
Could not load available models.{" "} diff --git a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx index 140b4363327..5bc0dcd143d 100644 --- a/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx +++ b/ui/litellm-dashboard/src/components/add_model/advanced_settings.tsx @@ -1,15 +1,18 @@ import React from "react"; -import { Switch, Select, Tooltip, DatePicker } from "antd"; -import { ChevronDown } from "lucide-react"; +import { MultiSelect } from "@/components/shared/MultiSelect"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Switch } from "@/components/ui/switch"; +import { SimpleTooltip } from "@/components/ui/tooltip"; +import type { Dayjs } from "dayjs"; +import { ChevronDown, Info } from "lucide-react"; import { Collapsible, CollapsibleContent, CollapsibleTrigger } from "@/components/ui/collapsible"; import { Input } from "@/components/ui/input"; -import { Row, Col, Typography } from "antd"; -import TextArea from "antd/es/input/TextArea"; -import { InfoCircleOutlined } from "@ant-design/icons"; +import { Textarea } from "@/components/ui/textarea"; import { Team } from "../key_team_helpers/key_list"; import { antdRules } from "../common_components/antdFormRules"; import { labelWithHint } from "@/components/shared/form/LabelWithHint"; import { MountedFormField } from "../common_components/MountedFormField"; +import { UtcDateTimeInput } from "@/components/shared/form/UtcDateTimeInput"; import CacheControlInjectionPoints, { CACHE_CONTROL_LABEL, CACHE_CONTROL_TOOLTIP, @@ -31,7 +34,6 @@ import { PTU_END_FIELD, } from "../../utils/ptuValidation"; import { usePtuCostAttributionEnabled } from "@/app/(dashboard)/hooks/uiSettings/usePtuCostAttributionEnabled"; -const { Link } = Typography; interface AdvancedSettingsProps { showAdvancedSettings: boolean; @@ -52,6 +54,11 @@ const USAGE_COST_FIELDS = [ const REVALIDATED_WHEN_PTU_COUNT_CHANGES = [PTU_RATE_FIELD, PTU_START_FIELD, ...USAGE_COST_FIELDS]; +const PRICING_MODEL_ITEMS = [ + { value: "per_token", label: "Per Million Tokens" }, + { value: "per_second", label: "Per Second" }, +] as const; + const validateNumber = (_: unknown, value: unknown) => { if (!value) { return Promise.resolve(); @@ -80,6 +87,14 @@ const AdvancedSettings: React.FC = ({ const [showCacheControl, setShowCacheControl] = React.useState(false); const ptuCostAttributionEnabled = usePtuCostAttributionEnabled(); + const handlePricingModelChange = + (onChange: (value: string) => void) => + (value: "per_token" | "per_second" | null): void => { + if (value === null) return; + onChange(value); + setPricingModel(value); + }; + return ( <> @@ -94,11 +109,10 @@ const AdvancedSettings: React.FC = ({ { + onCheckedChange={(checked) => { control.onChange(checked); setCustomPricing(checked); }} - className="bg-gray-600" /> )} @@ -108,16 +122,16 @@ const AdvancedSettings: React.FC = ({ label={ Attached Knowledge Bases (RAG){" "} - + e.stopPropagation()} > - + - + } className="mt-4" @@ -138,50 +152,48 @@ const AdvancedSettings: React.FC = ({ label={ Guardrails{" "} - + e.stopPropagation()} // Prevent accordion from collapsing when clicking link > - + - + } className="mt-4" help="Select existing guardrails. Go to 'Guardrails' tab to create new guardrails." > {(control) => ( - ({ value: tag.name, label: tag.name, - title: tag.description || tag.name, + description: tag.description || undefined, }))} + allowCustomValues /> )} @@ -250,11 +262,9 @@ const AdvancedSettings: React.FC = ({ className="mb-4" > {(control) => ( - @@ -274,11 +284,9 @@ const AdvancedSettings: React.FC = ({ className="mb-4" > {(control) => ( - @@ -292,19 +300,21 @@ const AdvancedSettings: React.FC = ({ {(control) => ( )} @@ -402,20 +412,20 @@ const AdvancedSettings: React.FC = ({ "Use in pass through routes", Allow using these credentials in pass through routes.{" "} - + Learn more - + , )} className="mb-4 mt-4" > {(control) => ( - + )} @@ -428,11 +438,10 @@ const AdvancedSettings: React.FC = ({ { + onCheckedChange={(checked) => { control.onChange(checked); setShowCacheControl(checked); }} - className="bg-gray-600" /> )} @@ -457,7 +466,7 @@ const AdvancedSettings: React.FC = ({ rules={{ validate: antdRules({ validator: formItemValidateJSON }) }} > {(control) => ( -