diff --git a/.circleci/config.yml b/.circleci/config.yml index b0a705966a2..2f01b6de4f3 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -2731,7 +2731,7 @@ jobs: - ~/.cache/uv - restore_cache: keys: - - ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }} + - ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }} - run: name: Install Node dependencies and Playwright # The cimg/python:3.12-browsers image already ships the Chromium system @@ -2742,11 +2742,14 @@ jobs: command: | cd ui/litellm-dashboard npm ci + cd ../../tests/e2e/ui + npm ci npx playwright install chromium - save_cache: - key: ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }} + key: ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }} paths: - ui/litellm-dashboard/node_modules + - tests/e2e/ui/node_modules - ~/.cache/ms-playwright - run: name: Build UI from source @@ -2777,10 +2780,10 @@ jobs: name: Seed database command: | PGPASSWORD=e2epassword psql -h localhost -p 5432 -U e2euser -d litellm_e2e \ - -f ui/litellm-dashboard/e2e_tests/fixtures/seed.sql + -f tests/e2e/ui/fixtures/seed.sql - run: name: Start mock LLM server - command: uv run --no-sync python ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py + command: uv run --no-sync python tests/e2e/ui/fixtures/mock_llm_server/server.py background: true - run: name: Start LiteLLM proxy @@ -2798,7 +2801,7 @@ jobs: command: | LITELLM_LICENSE="$LITELLM_LICENSE" \ uv run --no-sync python -m litellm.proxy.proxy_cli \ - --config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \ + --config tests/e2e/ui/fixtures/config.yml \ --port 4000 background: true - run: @@ -2819,15 +2822,15 @@ jobs: # Forward LITELLM_LICENSE so license.spec.ts can detect that the # proxy was launched with a license and assert premium_user=true. command: | - cd ui/litellm-dashboard + cd tests/e2e/ui LITELLM_LICENSE="$LITELLM_LICENSE" \ - npx playwright test --config e2e_tests/playwright.config.ts + npx playwright test --config playwright.config.ts no_output_timeout: 10m - store_artifacts: - path: ui/litellm-dashboard/test-results + path: tests/e2e/ui/test-results destination: e2e-test-results - store_artifacts: - path: ui/litellm-dashboard/playwright-report + path: tests/e2e/ui/playwright-report destination: e2e-playwright-report e2e_ui_testing_server_root_path: @@ -2870,17 +2873,20 @@ jobs: - ~/.cache/uv - restore_cache: keys: - - ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }} + - ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }} - run: name: Install Node dependencies and Playwright command: | cd ui/litellm-dashboard npm ci + cd ../../tests/e2e/ui + npm ci npx playwright install chromium - save_cache: - key: ui-e2e-node-deps-v2-{{ checksum "ui/litellm-dashboard/package-lock.json" }} + key: ui-e2e-node-deps-v3-{{ checksum "ui/litellm-dashboard/package-lock.json" }}-{{ checksum "tests/e2e/ui/package-lock.json" }} paths: - ui/litellm-dashboard/node_modules + - tests/e2e/ui/node_modules - ~/.cache/ms-playwright - run: name: Build UI from source @@ -2902,10 +2908,10 @@ jobs: name: Seed database command: | PGPASSWORD=e2epassword psql -h localhost -p 5432 -U e2euser -d litellm_e2e \ - -f ui/litellm-dashboard/e2e_tests/fixtures/seed.sql + -f tests/e2e/ui/fixtures/seed.sql - run: name: Start mock LLM server - command: uv run --no-sync python ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py + command: uv run --no-sync python tests/e2e/ui/fixtures/mock_llm_server/server.py background: true - run: name: Start LiteLLM proxy under a server root path @@ -2918,7 +2924,7 @@ jobs: command: | LITELLM_LICENSE="$LITELLM_LICENSE" \ uv run --no-sync python -m litellm.proxy.proxy_cli \ - --config ui/litellm-dashboard/e2e_tests/fixtures/config.yml \ + --config tests/e2e/ui/fixtures/config.yml \ --port 4000 background: true - run: @@ -2937,15 +2943,15 @@ jobs: - run: name: Run migration smoke under SERVER_ROOT_PATH command: | - cd ui/litellm-dashboard + cd tests/e2e/ui LITELLM_LICENSE="$LITELLM_LICENSE" \ - npx playwright test --config e2e_tests/migration.serverRootPath.config.ts + npx playwright test --config migration.serverRootPath.config.ts no_output_timeout: 10m - store_artifacts: - path: ui/litellm-dashboard/test-results + path: tests/e2e/ui/test-results destination: e2e-server-root-path-test-results - store_artifacts: - path: ui/litellm-dashboard/playwright-report + path: tests/e2e/ui/playwright-report destination: e2e-server-root-path-playwright-report build_docker_database_image: diff --git a/.circleci/scripts/classify_changes.sh b/.circleci/scripts/classify_changes.sh index 2c15428be6a..2ca2654a207 100755 --- a/.circleci/scripts/classify_changes.sh +++ b/.circleci/scripts/classify_changes.sh @@ -8,7 +8,7 @@ has_backend=false while IFS= read -r file || [ -n "$file" ]; do [ -n "$file" ] || continue case "$file" in - ui/*) has_client=true ;; + ui/* | tests/e2e/ui/*) has_client=true ;; docs/* | *.md | *.mdx) : ;; *) has_backend=true ;; esac diff --git a/.github/ISSUE_TEMPLATE/bug_report.yml b/.github/ISSUE_TEMPLATE/bug_report.yml index bbe4b76775d..665f8456f0b 100644 --- a/.github/ISSUE_TEMPLATE/bug_report.yml +++ b/.github/ISSUE_TEMPLATE/bug_report.yml @@ -30,7 +30,7 @@ body: id: steps-to-reproduce attributes: label: Steps to Reproduce - description: Please provide detailed steps to reproduce this bug(A curl/python code to reproduce the bug) + description: Please provide a numbered list of the exact steps to reproduce this bug (include a curl/python snippet to reproduce it). Number each step (1., 2., 3., ...) in the order you performed them. placeholder: | 1. config.yaml file/ .env file/ etc. 2. Run the following code... diff --git a/.github/pull_request_template.md b/.github/pull_request_template.md index d7e80b32749..1301bfb0e60 100644 --- a/.github/pull_request_template.md +++ b/.github/pull_request_template.md @@ -1,3 +1,18 @@ +## TLDR + + + +Problem this solves: + +- +- ... + +How it solves it: + +- +- ... + ## Relevant issues diff --git a/.github/workflows/image-scan.yml b/.github/workflows/image-scan.yml index 90ede5a653f..4d4a3242399 100644 --- a/.github/workflows/image-scan.yml +++ b/.github/workflows/image-scan.yml @@ -9,6 +9,7 @@ on: - "litellm_**" paths: - docker/Dockerfile.non_root + - tests/proxy_migration_tests/test_offline_image_migration.py - uv.lock - ui/litellm-dashboard/package-lock.json - .github/workflows/image-scan.yml @@ -51,6 +52,23 @@ jobs: - name: Build runtime image run: docker build -f docker/Dockerfile.non_root -t litellm-image-scan:${{ github.sha }} . + # The prisma bake must migrate a fresh DB with no egress as an arbitrary + # non-root uid (OpenShift restricted-v2 / air-gapped / readOnlyRootFilesystem). + # `docker run` as the default uid with network hides a broken bake because + # the migration entrypoint exits 0 even when it applied nothing; asserting + # the schema was created is what catches it. + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Verify offline migration as a non-root uid + env: + LITELLM_IMAGE: litellm-image-scan:${{ github.sha }} + run: | + python -m pip install "pytest==9.0.3" + python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py -v + # Scans the whole shipped artifact: OS/apk plus every language package # baked into the image, including ones no lockfile declares (e.g. prisma's # vendored node engine) that osv-scan cannot see. osv-scan stays the fast @@ -58,6 +76,8 @@ jobs: # free OSS, run as a pinned, checksum-verified binary; no GitHub Action # dependency and no vendor SaaS callout. - name: Scan image for fixable HIGH/CRITICAL CVEs + env: + GRYPE_MATCH_PYTHON_USING_CPES: "true" run: | "$RUNNER_TEMP/grype" litellm-image-scan:${{ github.sha }} \ --only-fixed \ diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index 9d28ca211cf..ae31395521a 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -115,6 +115,9 @@ jobs: - name: check_fastuuid_usage run: uv run --no-sync python ./tests/code_coverage_tests/check_fastuuid_usage.py + - name: check_e2e_no_raw_requests + run: uv run --no-sync python ./tests/code_coverage_tests/check_e2e_no_raw_requests.py + - name: memory_test run: uv run --no-sync python ./tests/code_coverage_tests/memory_test.py diff --git a/.github/workflows/test-litellm-ui-unit.yml b/.github/workflows/test-litellm-ui-unit.yml new file mode 100644 index 00000000000..5374a0059de --- /dev/null +++ b/.github/workflows/test-litellm-ui-unit.yml @@ -0,0 +1,57 @@ +name: UI Unit Tests +permissions: + contents: read + +on: + pull_request: + branches: + - main + - litellm_internal_staging + - litellm_oss_staging + - "litellm_**" + push: + branches: + - litellm_internal_staging + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: true + +jobs: + ui-unit-tests: + runs-on: ubuntu-latest-16-cores + timeout-minutes: 20 + defaults: + run: + working-directory: ui/litellm-dashboard + + steps: + - name: Checkout repository + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + fetch-depth: 0 + persist-credentials: false + + - name: Setup Node.js + uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0 + with: + node-version: "20" + cache: "npm" + cache-dependency-path: ui/litellm-dashboard/package-lock.json + + - name: Install dependencies + run: npm ci + + - name: Run UI unit tests (Vitest) + env: + CI: "true" + BASE_SHA: ${{ github.event.pull_request.base.sha }} + run: | + if [ -n "$BASE_SHA" ]; then + echo "Pull request: running only tests related to changes since $BASE_SHA" + npm run test -- --run --changed "$BASE_SHA" --passWithNoTests \ + --pool forks --poolOptions.forks.maxForks=14 + else + echo "Push to $GITHUB_REF_NAME: running the full suite" + npm run test -- --run --pool forks --poolOptions.forks.maxForks=14 + fi diff --git a/.github/workflows/test-unit-proxy-endpoints.yml b/.github/workflows/test-unit-proxy-endpoints.yml index cbb36eebdb9..b3eb8f79a43 100644 --- a/.github/workflows/test-unit-proxy-endpoints.yml +++ b/.github/workflows/test-unit-proxy-endpoints.yml @@ -46,6 +46,7 @@ jobs: tests/test_litellm/proxy/rag_endpoints tests/test_litellm/proxy/realtime_endpoints tests/test_litellm/proxy/ui_crud_endpoints + tests/test_litellm/proxy/config_resolvers tests/test_litellm/proxy/utils workers: 2 reruns: 2 diff --git a/.github/workflows/test_server_root_path.yml b/.github/workflows/test_server_root_path.yml index f59cee29893..01f70511e79 100644 --- a/.github/workflows/test_server_root_path.yml +++ b/.github/workflows/test_server_root_path.yml @@ -106,8 +106,8 @@ jobs: with: node-version: "20" - - name: Install UI deps and Chromium - working-directory: ui/litellm-dashboard + - name: Install e2e deps and Chromium + working-directory: tests/e2e/ui run: | retry() { local attempt=1 @@ -131,17 +131,17 @@ jobs: retry npx playwright install --with-deps chromium - name: Run SERVER_ROOT_PATH redirect e2e - working-directory: ui/litellm-dashboard + working-directory: tests/e2e/ui env: SERVER_ROOT_PATH: ${{ matrix.root_path }} - run: npx playwright test --config=e2e_tests/serverRootPath.config.ts + run: npx playwright test --config=serverRootPath.config.ts - name: Upload Playwright artifacts on failure if: failure() uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 with: name: playwright-trace-${{ strategy.job-index }} - path: ui/litellm-dashboard/test-results/ + path: tests/e2e/ui/test-results/ retention-days: 7 - name: Cleanup diff --git a/.github/workflows/weekly_load_anomaly.yml b/.github/workflows/weekly_load_anomaly.yml new file mode 100644 index 00000000000..4c2103f026d --- /dev/null +++ b/.github/workflows/weekly_load_anomaly.yml @@ -0,0 +1,81 @@ +name: "Weekly Load Anomaly Check" + +on: + schedule: + - cron: "0 12 * * 6" + workflow_dispatch: + +permissions: + contents: read + +jobs: + weekly-load-anomaly: + if: github.event_name != 'schedule' || github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + timeout-minutes: 45 + services: + postgres: + image: postgres:16.6 + env: + POSTGRES_USER: llmproxy + POSTGRES_PASSWORD: dbpassword9090 + POSTGRES_DB: litellm + ports: + - 5432:5432 + options: >- + --health-cmd "pg_isready -U llmproxy" + --health-interval 5s + --health-timeout 5s + --health-retries 10 + env: + DATABASE_URL: postgresql://llmproxy:dbpassword9090@localhost:5432/litellm + LITELLM_MASTER_KEY: sk-weekly-anomaly-check + ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }} + AWS_BEARER_TOKEN_BEDROCK: ${{ secrets.AWS_BEARER_TOKEN_BEDROCK }} + steps: + - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Set up uv + uses: ./.github/actions/setup-uv-with-retries + with: + version: "0.10.9" + + - name: Install dependencies + run: | + .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy + + - name: Generate Prisma client + env: + PRISMA_BINARY_CACHE_DIR: ${{ runner.temp }}/prisma-cache + run: | + uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma + + - name: Start the proxy + run: | + nohup uv run --no-sync litellm --config tests/e2e/load/weekly_anomaly_config.yml --port 4000 > proxy.log 2>&1 & + for _ in $(seq 1 90); do + if curl -fs http://localhost:4000/health/liveliness > /dev/null; then + exit 0 + fi + sleep 2 + done + echo "proxy never became live" + tail -n 100 proxy.log + exit 1 + + - name: Run the weekly session anomaly test + env: + E2E_WEEKLY_ANOMALY: "1" + run: | + uv run --no-sync pytest tests/e2e/load/test_weekly_session_anomaly_e2e.py -v --tb=short -rA + + - name: Show proxy log on failure + if: failure() + run: tail -n 300 proxy.log diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 02574ca505d..f3a028f5805 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -18,6 +18,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = ( "/team/", "/v2/team/", "/organization/", + "/v2/organization/", "/customer/", "/end_user/", "/sso/", diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index 839f5da565c..8e05f312ba0 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -54,7 +54,6 @@ ENV UV_PROJECT_ENVIRONMENT=/app/.venv \ UV_LINK_MODE=copy \ PATH="/app/.venv/bin:${PATH}" \ LITELLM_NON_ROOT=true \ - PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \ XDG_CACHE_HOME=/app/.cache # Copy dependency metadata first for layer caching @@ -106,7 +105,9 @@ RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \ --python python3; \ fi -RUN prisma generate --schema=./schema.prisma +RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \ + npm_config_cache=/root/.npm \ + prisma generate --schema=./schema.prisma RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \ sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh @@ -127,8 +128,6 @@ RUN for i in 1 2 3; do \ # the rest of the builder's /app is source and build metadata that must not # ship (manifest-scanning tools attribute everything in it to this image). # entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path. -# Prisma caches live under /app/.cache here (XDG_CACHE_HOME / -# PRISMA_BINARY_CACHE_DIR) so the runtime prisma generate finds them. COPY --from=builder /app/.venv /app/.venv COPY --from=builder /app/docker /app/docker COPY --from=builder /app/schema.prisma /app/schema.prisma @@ -138,21 +137,35 @@ COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/pr # enterprise.enterprise_hooks from it) COPY --from=builder /app/enterprise /app/enterprise COPY --from=builder /app/litellm-proxy-extras /app/litellm-proxy-extras -COPY --from=builder /app/.cache /app/.cache +# Prisma CLI + engines are baked under /opt/prisma, a fixed path every runtime +# uid can read and that no cache volume mount shadows (unlike /app/.cache or +# $HOME/.cache under readOnlyRootFilesystem + emptyDir or arbitrary-uid setups). +# PRISMA_CLI_QUERY_ENGINE_TYPE=binary makes the CLI use the baked binary query +# engine directly, so `prisma migrate deploy` on a fresh database needs no npm +# and no network access; without it the CLI looks for the library engine, which +# prisma stopped baking, and falls back to a download that fails offline or as a +# non-writable uid (#33650, #24554). +COPY --from=builder /opt/prisma /opt/prisma COPY --from=builder /var/lib/litellm/ui /var/lib/litellm/ui COPY --from=builder /var/lib/litellm/assets /var/lib/litellm/assets +# XDG_CACHE_HOME is intentionally left unset so it falls back to $HOME/.cache +# (/app/.cache, writable by the runtime uid). The prisma bake at the read-only +# /opt/prisma is anchored by PRISMA_BINARY_CACHE_DIR / PRISMA_CLI_PATH, so +# nothing needs XDG to point there; pointing it at the read-only bake would +# deny any XDG-aware library that writes a cache at runtime. ENV PATH="/app/.venv/bin:${PATH}" \ - PRISMA_BINARY_CACHE_DIR=/app/.cache/prisma-python/binaries \ + PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \ + PRISMA_CLI_PATH=/opt/prisma/binaries/node_modules/.bin/prisma \ + PRISMA_CLI_QUERY_ENGINE_TYPE=binary \ HOME=/app \ LITELLM_NON_ROOT=true \ - XDG_CACHE_HOME=/app/.cache \ PRISMA_SKIP_POSTINSTALL_GENERATE=1 \ PRISMA_HIDE_UPDATE_MESSAGE=1 \ PRISMA_ENGINES_CHECKSUM_IGNORE_MISSING=1 \ PRISMA_OFFLINE_MODE=true -RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \ +RUN mkdir -p /nonexistent /app/.cache /var/lib/litellm/assets /var/lib/litellm/ui && \ chown -R nobody:nogroup /app /var/lib/litellm/ui /var/lib/litellm/assets /nonexistent && \ PRISMA_PATH=$(python -c "import os, prisma; print(os.path.dirname(prisma.__file__))") && \ chown -R nobody:nogroup "$PRISMA_PATH" && \ @@ -165,12 +178,14 @@ RUN mkdir -p /nonexistent /var/lib/litellm/assets /var/lib/litellm/ui && \ [ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g=u "$LITELLM_PROXY_EXTRAS_PATH" || true && \ chmod -R g+w "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets && \ [ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g+w "$LITELLM_PROXY_EXTRAS_PATH" || true && \ - chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets /app/.cache + chmod -R g+rX "$PRISMA_PATH" /var/lib/litellm/ui /var/lib/litellm/assets && \ + chmod -R a+rX /opt/prisma && \ + test -x /opt/prisma/binaries/node_modules/.bin/prisma && \ + test -f /opt/prisma/binaries/node_modules/prisma/build/index.js && \ + ls /opt/prisma/binaries/node_modules/@prisma/engines/query-engine-* >/dev/null 2>&1 USER 65534 -RUN prisma generate --schema=./schema.prisma - EXPOSE 4000/tcp ENTRYPOINT ["/app/docker/prod_entrypoint.sh"] diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql new file mode 100644 index 00000000000..95412df0a96 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260721000000_add_sso_identity_assertion/migration.sql @@ -0,0 +1,9 @@ +-- CreateTable +CREATE TABLE IF NOT EXISTS "LiteLLM_SSOIdentityAssertion" ( + "user_id" TEXT NOT NULL, + "assertion_b64" TEXT NOT NULL, + "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP, + + CONSTRAINT "LiteLLM_SSOIdentityAssertion_pkey" PRIMARY KEY ("user_id") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index b27ddea010b..23a9c086c73 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient { updated_at DateTime @default(now()) @updatedAt @map("updated_at") } +// The enterprise IdP identity assertion captured at SSO login, one row per user. +// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}. +model LiteLLM_SSOIdentityAssertion { + user_id String @id + assertion_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") +} + // Generate Tokens for Proxy model LiteLLM_VerificationToken { token String @id diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 3288f7fd584..ccca88c9996 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.79" +version = "0.4.80" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.79" +version = "0.4.80" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs index 4b906155665..68ecc3f17c1 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs @@ -50,3 +50,15 @@ pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool { .iter() .any(|(key, _)| key.eq_ignore_ascii_case(name)) } + +pub(super) fn has_bearer_auth(headers: &[(String, String)]) -> bool { + headers.iter().any(|(name, value)| { + if !name.eq_ignore_ascii_case("authorization") { + return false; + } + let value = value.trim(); + value.len() > 7 + && value[..7].eq_ignore_ascii_case("bearer ") + && !value[7..].trim().is_empty() + }) +} diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs index 624c3598fb0..9a027490eb6 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs @@ -3,7 +3,7 @@ use litellm_core::CoreResult; use litellm_core::messages::transformation::MessagesAuthStrategy; use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; -use super::common_utils::{has_header, messages_provider_config, string_headers}; +use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; use super::types::{MessagesRequest, ProviderMessagesRequest}; pub(super) fn prepare_messages_call( @@ -33,7 +33,9 @@ pub(super) fn prepare_messages_call( let mut headers = string_headers(request.extra_headers)?; let auth_strategy = config.auth_strategy(); - if !has_header(&headers, auth_strategy.header_name()) { + let already_authorized = has_header(&headers, auth_strategy.header_name()) + || (config.accepts_bearer_auth() && has_bearer_auth(&headers)); + if !already_authorized { let api_key = config.resolve_api_key(request.api_key, &env_lookup)?; let auth_header = match auth_strategy { MessagesAuthStrategy::Bearer => { diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/ai-gateway/src/messages/tests.rs index a2d0f6fae23..23a53e98045 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/tests.rs @@ -6,7 +6,7 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; use super::common_utils::{ - has_header, messages_provider_config, string_headers, truncate_error_body, + has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, }; use super::{MessagesRequest, messages}; @@ -85,6 +85,34 @@ fn has_header_is_case_insensitive() { assert!(!has_header(&headers, "authorization")); } +#[test] +fn has_bearer_auth_requires_a_nonempty_bearer_token() { + assert!(has_bearer_auth(&[( + "Authorization".to_string(), + "Bearer tok".to_string() + )])); + assert!(has_bearer_auth(&[( + "authorization".to_string(), + "bearer tok".to_string() + )])); + assert!(!has_bearer_auth(&[( + "authorization".to_string(), + "Bearer ".to_string() + )])); + assert!(!has_bearer_auth(&[( + "authorization".to_string(), + String::new() + )])); + assert!(!has_bearer_auth(&[( + "authorization".to_string(), + "Basic abc".to_string() + )])); + assert!(!has_bearer_auth(&[( + "x-api-key".to_string(), + "sk".to_string() + )])); +} + #[tokio::test] async fn messages_round_trip_builds_azure_request_and_passes_response_through() { let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); @@ -252,6 +280,112 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() { assert!(!head.contains("rust-fallback-key"), "{head}"); } +#[tokio::test] +async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let request = read_http_request(&mut socket).await; + let response_body = + r#"{"id":"msg_3","type":"message","role":"assistant","content":[],"model":"m"}"#; + socket + .write_all(write_response(response_body).as_bytes()) + .await + .expect("writes response"); + request + }); + + let mut headers = Map::new(); + headers.insert( + "Authorization".to_string(), + Value::String("Bearer entra-token".to_string()), + ); + + messages(MessagesRequest { + model: "claude-sonnet-4-5", + body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), + api_key: None, + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("azure_ai"), + extra_headers: Some(headers), + timeout: Some(Duration::from_secs(5)), + }) + .await + .expect("entra id request succeeds without api key"); + + let request = server.await.expect("server task completes"); + let head = request + .split_once("\r\n\r\n") + .expect("has body") + .0 + .to_ascii_lowercase(); + assert!(head.contains("authorization: bearer entra-token"), "{head}"); + assert!(!head.contains("x-api-key"), "{head}"); +} + +#[tokio::test] +async fn messages_requires_auth_when_no_key_and_no_header() { + let err = messages(MessagesRequest { + model: "claude-sonnet-4-5", + body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), + api_key: None, + api_base: Some("http://127.0.0.1:1"), + custom_llm_provider: Some("azure_ai"), + extra_headers: None, + timeout: Some(Duration::from_millis(50)), + }) + .await + .expect_err("missing auth errors"); + + assert!(matches!(err, CoreError::Auth(_))); +} + +#[tokio::test] +async fn messages_ignores_malformed_authorization_and_uses_api_key() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let request = read_http_request(&mut socket).await; + let response_body = + r#"{"id":"msg_4","type":"message","role":"assistant","content":[],"model":"m"}"#; + socket + .write_all(write_response(response_body).as_bytes()) + .await + .expect("writes response"); + request + }); + + let mut headers = Map::new(); + headers.insert( + "Authorization".to_string(), + Value::String("Bearer ".to_string()), + ); + + messages(MessagesRequest { + model: "claude-sonnet-4-5", + body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}), + api_key: Some("sk-azure"), + api_base: Some(&format!("http://{addr}")), + custom_llm_provider: Some("azure_ai"), + extra_headers: Some(headers), + timeout: Some(Duration::from_secs(5)), + }) + .await + .expect("falls back to api key"); + + let request = server.await.expect("server task completes"); + let head = request + .split_once("\r\n\r\n") + .expect("has body") + .0 + .to_ascii_lowercase(); + assert!(head.contains("x-api-key: sk-azure"), "{head}"); +} + #[tokio::test] async fn messages_maps_provider_error_status_to_http_error() { let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); diff --git a/litellm-rust/crates/core/src/messages/transformation.rs b/litellm-rust/crates/core/src/messages/transformation.rs index 3a34a58de6f..b478e20d24b 100644 --- a/litellm-rust/crates/core/src/messages/transformation.rs +++ b/litellm-rust/crates/core/src/messages/transformation.rs @@ -35,6 +35,10 @@ pub trait AnthropicMessagesProviderConfig: Sync { MessagesAuthStrategy::Header("x-api-key") } + fn accepts_bearer_auth(&self) -> bool { + false + } + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { &[ ("anthropic-version", "2023-06-01"), diff --git a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs index 6935bb4604b..7b958c77ba3 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs @@ -163,6 +163,10 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig { self.anthropic.auth_strategy() } + fn accepts_bearer_auth(&self) -> bool { + true + } + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { self.anthropic.default_headers() } @@ -294,6 +298,11 @@ mod tests { ); } + #[test] + fn accepts_bearer_auth_for_entra_id() { + assert!(AZURE_ANTHROPIC_MESSAGES_CONFIG.accepts_bearer_auth()); + } + #[test] fn default_headers_match_python() { assert_eq!( diff --git a/litellm/__init__.py b/litellm/__init__.py index 2f6643c644c..55821012df9 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -427,7 +427,9 @@ default_team_settings: Optional[List] = None max_user_budget: Optional[float] = None default_max_internal_user_budget: Optional[float] = None max_internal_user_budget: Optional[float] = None -max_ui_session_budget: Optional[float] = 0.25 # $0.25 USD budgets for UI Chat sessions +max_ui_session_budget: Optional[float] = ( + 1.0 # USD budget for each dashboard login session (playground, test connection) +) internal_user_budget_duration: Optional[str] = None tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None max_end_user_budget: Optional[float] = None diff --git a/litellm/constants.py b/litellm/constants.py index 05944c81ea2..9f60c635249 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -264,6 +264,9 @@ MAX_REDIS_BUFFER_DEQUEUE_COUNT = int(os.getenv("MAX_REDIS_BUFFER_DEQUEUE_COUNT", # Bounds asyncio.Queue() instances (log queues, spend update queues, etc.) to prevent unbounded memory growth LITELLM_ASYNCIO_QUEUE_MAXSIZE = int(os.getenv("LITELLM_ASYNCIO_QUEUE_MAXSIZE", 1000)) TOOL_POLICY_CACHE_TTL_SECONDS = int(os.getenv("TOOL_POLICY_CACHE_TTL_SECONDS", 60)) +GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS = int( + os.getenv("GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS", 24 * 60 * 60) +) # Aggregation threshold: default to 80% of the asyncio queue maxsize so the check can always trigger. # Must be < LITELLM_ASYNCIO_QUEUE_MAXSIZE; if set higher the aggregation logic will never fire. MAX_SIZE_IN_MEMORY_QUEUE = int(os.getenv("MAX_SIZE_IN_MEMORY_QUEUE", int(LITELLM_ASYNCIO_QUEUE_MAXSIZE * 0.8))) @@ -1469,6 +1472,7 @@ _batch_polling_env = os.getenv("PROXY_BATCH_POLLING_ENABLED", "true").lower() PROXY_BATCH_POLLING_ENABLED = _batch_polling_env == "true" PROXY_BUDGET_RESCHEDULER_MAX_TIME = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MAX_TIME", 605)) PROXY_BATCH_WRITE_AT = int(os.getenv("PROXY_BATCH_WRITE_AT", 10)) # in seconds, increased from 10 +PROXY_CONFIG_RELOAD_INTERVAL_SECONDS = get_env_int("PROXY_CONFIG_RELOAD_INTERVAL_SECONDS", 30) # APScheduler Configuration - MEMORY LEAK FIX # These settings prevent memory leaks in APScheduler's normalize() and _apply_jitter() functions @@ -1524,6 +1528,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [ # test_general_settings_ui_fields_are_db_overridable enforces that pairing. "enable_anthropic_prompt_caching", "anthropic_prompt_caching_ttl", + "max_ui_session_budget", ] SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"] DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60)) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 856556f7c56..cf9dafcb222 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1,3 +1,4 @@ +import hashlib import os import secrets from datetime import datetime @@ -46,7 +47,10 @@ if TYPE_CHECKING: dc = DualCache() -from litellm.constants import PRE_CALL_EXECUTED_GUARDRAILS_KEY +from litellm.constants import ( + GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS, + PRE_CALL_EXECUTED_GUARDRAILS_KEY, +) from litellm.exceptions import ( BlockedPiiEntityError, GuardrailRaisedException, @@ -113,6 +117,7 @@ class CustomGuardrail(CustomLogger): on_sensitive_data: Optional[str] = None, sensitive_data_route_to_model: Optional[str] = None, sticky_session_routing: bool = True, + only_scan_new_messages: bool = False, **kwargs, ): """ @@ -145,6 +150,7 @@ class CustomGuardrail(CustomLogger): self.on_sensitive_data: Optional[str] = on_sensitive_data self.sensitive_data_route_to_model: Optional[str] = sensitive_data_route_to_model self.sticky_session_routing: bool = sticky_session_routing + self.only_scan_new_messages: bool = only_scan_new_messages if supported_event_hooks: ## validate event_hook is in supported_event_hooks @@ -269,6 +275,100 @@ class CustomGuardrail(CustomLogger): """Extract session_id from request data.""" return get_session_id_from_request_data(request_data) + @staticmethod + def _scanned_text_hash(text: str) -> str: + """Stable content hash for a single scannable text segment. + + Hashing the exact text the provider would receive means an edited earlier + segment produces a different hash and gets re-scanned, while an unchanged + segment repeated on a later turn is skipped. + """ + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + def _scanned_texts_cache_key(self, session_id: str) -> str: + return f"guardrail_scanned_texts:{self.guardrail_name}:{session_id}" + + async def filter_new_texts_for_session( + self, + texts: list[str] | None, + request_data: dict[str, object], + cache: DualCache, + ) -> list[str] | None: + """Return only the text segments not already scanned earlier in this session. + + Returns ``None`` when incremental scanning is inactive (feature off, no + session id, masking enabled, or the cache read failed). ``None`` signals + the caller to fall back to a full scan; a returned list (possibly empty) + signals the caller to scan only that subset and skip masking write-back. + """ + if not self.only_scan_new_messages or not texts: + return None + + if self.mask_request_content or self.mask_response_content: + verbose_logger.warning( + "Guardrail %s: only_scan_new_messages is not supported with masking; scanning full context.", + self.guardrail_name, + ) + return None + + session_id = get_session_id_from_request_data(request_data) + if not session_id: + verbose_logger.debug( + "Guardrail %s: only_scan_new_messages enabled but request has no session id; scanning full context.", + self.guardrail_name, + ) + return None + + try: + cached: object = await cache.async_get_cache(key=self._scanned_texts_cache_key(session_id)) + except Exception as e: # noqa: BLE001 # cache is best-effort; any failure must fall back to a full scan + verbose_logger.warning( + "Guardrail %s: failed to read scanned-message cache (%s); scanning full context.", + self.guardrail_name, + e, + ) + return None + + seen: set[str] = {str(h) for h in cached} if isinstance(cached, list) else set() + return [text for text in texts if self._scanned_text_hash(text) not in seen] + + async def mark_texts_scanned( + self, + texts: list[str] | None, + request_data: dict[str, object], + cache: DualCache, + ) -> None: + """Record the hashes of all text segments present on a successful (non-blocked) scan. + + Called only after the guardrail allows the request, so a blocked segment is + never marked scanned and will be re-checked if the client retries. + """ + if not self.only_scan_new_messages or not texts: + return + if self.mask_request_content or self.mask_response_content: + return + session_id = get_session_id_from_request_data(request_data) + if not session_id: + return + + cache_key = self._scanned_texts_cache_key(session_id) + current_hashes = [self._scanned_text_hash(text) for text in texts] + try: + existing: object = await cache.async_get_cache(key=cache_key) + existing_hashes: list[str] = [str(h) for h in existing] if isinstance(existing, list) else [] + merged: list[str] = list(dict.fromkeys(existing_hashes + current_hashes)) + await cache.async_set_cache( + key=cache_key, + value=merged, + ttl=GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS, + ) + except Exception as e: # noqa: BLE001 # cache is best-effort; any failure must not block the request + verbose_logger.warning( + "Guardrail %s: failed to persist scanned-message cache (%s); next call will re-scan.", + self.guardrail_name, + e, + ) + def should_route_on_sensitive_data(self) -> bool: """ Returns True if this guardrail is configured to route requests diff --git a/litellm/litellm_core_utils/duration_parser.py b/litellm/litellm_core_utils/duration_parser.py index 438ff5600ba..b78a314dc45 100644 --- a/litellm/litellm_core_utils/duration_parser.py +++ b/litellm/litellm_core_utils/duration_parser.py @@ -7,11 +7,24 @@ duration_in_seconds is used in diff parts of the code base, example """ import re -import time -from datetime import datetime, timedelta, timezone, tzinfo -from typing import Optional, Tuple +import time as time_module +from datetime import datetime, time, timedelta, timezone, tzinfo +from typing import Final, Optional, Tuple from zoneinfo import ZoneInfo +from litellm._logging import verbose_logger + +_BUDGET_DURATION_WORD_ALIASES: Final[dict[str, str]] = { + "hourly": "1h", + "daily": "24h", + "weekly": "7d", + "monthly": "30d", +} + + +def _normalize_duration(duration: str) -> str: + return _BUDGET_DURATION_WORD_ALIASES.get(duration.strip().lower(), duration) + def _extract_from_regex(duration: str) -> Tuple[int, str]: match = re.match(r"(\d+)(mo|[smhdw]?)", duration) @@ -48,7 +61,7 @@ def duration_in_seconds(duration: str) -> int: Returns time in seconds till when budget needs to be reset """ - value, unit = _extract_from_regex(duration=duration) + value, unit = _extract_from_regex(duration=_normalize_duration(duration)) if unit == "s": return value @@ -61,7 +74,7 @@ def duration_in_seconds(duration: str) -> int: elif unit == "w": return value * 604800 elif unit == "mo": - now = time.time() + now = time_module.time() current_time = datetime.fromtimestamp(now) # Calculate target month and year, handling overflow past December @@ -94,12 +107,17 @@ def duration_in_seconds(duration: str) -> int: raise ValueError(f"Unsupported duration unit, passed duration: {duration}") -def get_next_standardized_reset_time(duration: str, current_time: datetime, timezone_str: str = "UTC") -> datetime: +def get_next_standardized_reset_time( + duration: str, + current_time: datetime, + timezone_str: str = "UTC", + reset_time_of_day: time = time(0, 0), +) -> datetime: """ Get the next standardized reset time based on the duration. All durations will reset at predictable intervals, aligned from the current time: - - Nd: If N=1, reset at next midnight; if N>1, reset every N days from now + - Nd: If N=1, reset at the next `reset_time_of_day`; if N>1, reset every N days from now - Nh: Every N hours, aligned to hour boundaries (e.g., 1:00, 2:00) - Nm: Every N minutes, aligned to minute boundaries (e.g., 1:05, 1:10) - Ns: Every N seconds, aligned to second boundaries @@ -108,17 +126,24 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time - duration: Duration string (e.g. "30s", "30m", "30h", "30d") - current_time: Current datetime - timezone_str: Timezone string (e.g. "UTC", "US/Eastern", "Asia/Kolkata") + - reset_time_of_day: Wall-clock time the reset lands on for day/week/month + durations (defaults to midnight). Ignored for sub-day durations, where a + time-of-day is meaningless. Returns: - Next reset time at a standardized interval in the specified timezone """ # Set up timezone and normalize current time - current_time, tz = _setup_timezone(current_time, timezone_str) + current_time, _ = _setup_timezone(current_time, timezone_str) # Parse duration - value, unit = _parse_duration(duration) + value, unit = _parse_duration(_normalize_duration(duration)) if value is None: - # Fall back to default if format is invalid + verbose_logger.warning( + "Unrecognized budget_duration %r; falling back to a next-midnight reset. " + "Use the format (e.g. '1h', '7d', '30d', '1mo').", + duration, + ) return current_time.replace(hour=0, minute=0, second=0, microsecond=0) + timedelta(days=1) # Midnight of the current day in the specified timezone @@ -126,9 +151,9 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time # Handle different time units if unit == "d": - return _handle_day_reset(current_time, base_midnight, value, tz) + return _handle_day_reset(current_time, base_midnight, value, reset_time_of_day) elif unit == "w": - return _handle_day_reset(current_time, base_midnight, value * 7, tz) + return _handle_day_reset(current_time, base_midnight, value * 7, reset_time_of_day) elif unit == "h": return _handle_hour_reset(current_time, base_midnight, value) elif unit == "m": @@ -136,7 +161,7 @@ def get_next_standardized_reset_time(duration: str, current_time: datetime, time elif unit == "s": return _handle_second_reset(current_time, base_midnight, value) elif unit == "mo": - return _handle_month_reset(current_time, base_midnight, value) + return _handle_month_reset(current_time, base_midnight, value, reset_time_of_day) else: # Unrecognized unit, default to next midnight return base_midnight + timedelta(days=1) @@ -175,46 +200,58 @@ def _parse_duration(duration: str) -> Tuple[Optional[int], Optional[str]]: return int(value), unit -def _handle_day_reset(current_time: datetime, base_midnight: datetime, value: int, tz: tzinfo) -> datetime: +def _apply_time_of_day(dt: datetime, reset_time_of_day: time) -> datetime: + """Set the wall-clock time of `dt` to `reset_time_of_day`, keeping its date and tzinfo.""" + return dt.replace( + hour=reset_time_of_day.hour, + minute=reset_time_of_day.minute, + second=reset_time_of_day.second, + microsecond=reset_time_of_day.microsecond, + ) + + +def _next_occurrence( + boundary_midnight: datetime, + reset_time_of_day: time, + current_time: datetime, + period: timedelta, +) -> datetime: + """Place the reset at `reset_time_of_day` on the boundary day, rolling forward one + `period` if that instant has already passed (or is exactly now).""" + candidate = _apply_time_of_day(boundary_midnight, reset_time_of_day) + if candidate <= current_time: + return candidate + period + return candidate + + +def _first_of_next_month(first_of_month: datetime) -> datetime: + """Given the 1st of some month, return the 1st of the following month.""" + if first_of_month.month == 12: + return first_of_month.replace(year=first_of_month.year + 1, month=1) + return first_of_month.replace(month=first_of_month.month + 1) + + +def _handle_day_reset( + current_time: datetime, + base_midnight: datetime, + value: int, + reset_time_of_day: time, +) -> datetime: """Handle day-based reset times.""" # Handle zero value - immediate expiration if value == 0: return current_time - if value == 1: # Daily reset at midnight - return base_midnight + timedelta(days=1) - elif value == 7: # Weekly reset on Monday at midnight + if value == 1: # Daily reset at the configured time of day + return _next_occurrence(base_midnight, reset_time_of_day, current_time, timedelta(days=1)) + elif value == 7: # Weekly reset on Monday at the configured time of day days_until_monday = (7 - current_time.weekday()) % 7 - if days_until_monday == 0: # If today is Monday - days_until_monday = 7 - return base_midnight + timedelta(days=days_until_monday) - elif value == 30: # Monthly reset on 1st at midnight - # Get 1st of next month at midnight - if current_time.month == 12: - next_reset = datetime( - year=current_time.year + 1, - month=1, - day=1, - hour=0, - minute=0, - second=0, - microsecond=0, - tzinfo=tz, - ) - else: - next_reset = datetime( - year=current_time.year, - month=current_time.month + 1, - day=1, - hour=0, - minute=0, - second=0, - microsecond=0, - tzinfo=tz, - ) - return next_reset - else: # Custom day value - next interval is value days from current - return current_time.replace(hour=0, minute=0, second=0, microsecond=0) + timedelta(days=value) + upcoming_monday = base_midnight + timedelta(days=days_until_monday) + return _next_occurrence(upcoming_monday, reset_time_of_day, current_time, timedelta(days=7)) + elif value == 30: # Monthly reset on 1st at the configured time of day + return _handle_month_reset(current_time, base_midnight, 1, reset_time_of_day) + else: # Custom day value - next interval is value days from the start of today + return _apply_time_of_day(base_midnight + timedelta(days=value), reset_time_of_day) def _handle_hour_reset(current_time: datetime, base_midnight: datetime, value: int) -> datetime: @@ -316,36 +353,30 @@ def _handle_second_reset(current_time: datetime, base_midnight: datetime, value: return current_time.replace(hour=next_hour, minute=next_minute, second=next_second, microsecond=0) -def _handle_month_reset(current_time: datetime, base_midnight: datetime, value: int) -> datetime: +def _handle_month_reset( + current_time: datetime, + base_midnight: datetime, + value: int, + reset_time_of_day: time, +) -> datetime: """ - Handle monthly reset times. For monthly resets, we always reset at the start of the next month. + Handle monthly reset times. Resets land on the 1st at `reset_time_of_day`; if the + 1st of the current month at that time has already passed, roll to the 1st of next month. Args: current_time: Current datetime base_midnight: Midnight of current day value: Number of months (currently only supports 1 month resets) + reset_time_of_day: Wall-clock time the reset lands on Returns: - datetime: First day of next month at midnight + datetime: First day of the next reset month at `reset_time_of_day` """ if value != 1: raise ValueError("Monthly resets currently only support 1 month intervals") - # Get the first day of next month - if current_time.month == 12: - next_month = 1 - next_year = current_time.year + 1 - else: - next_month = current_time.month + 1 - next_year = current_time.year - - return datetime( - year=next_year, - month=next_month, - day=1, - hour=0, - minute=0, - second=0, - microsecond=0, - tzinfo=current_time.tzinfo, - ) + first_of_this_month = base_midnight.replace(day=1) + candidate = _apply_time_of_day(first_of_this_month, reset_time_of_day) + if candidate <= current_time: + return _apply_time_of_day(_first_of_next_month(first_of_this_month), reset_time_of_day) + return candidate diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py index 1cbb1ce973f..a83cb3bc69e 100644 --- a/litellm/litellm_core_utils/url_utils.py +++ b/litellm/litellm_core_utils/url_utils.py @@ -148,6 +148,15 @@ def _parse_url_destination_allowlist_entry( return _normalize_host(parsed.hostname), scheme, port +def provider_url_destination_candidates(value: str) -> Tuple[str, ...]: + return tuple( + candidate + for part in value.split(",") + for candidate in (part.strip(), part.strip().split("/", 1)[1] if "/" in part.strip() else "") + if candidate + ) + + def is_url_destination_allowed_by_host(url: str, allowed_hosts: List[str]) -> bool: """Return True when a credential-bearing provider URL is admin-allowlisted. diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 5a0f274e3ca..e99f356f8f2 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -480,10 +480,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): """ Filter out unsupported fields from JSON schema for Anthropic's output_format API. - Anthropic's output_format doesn't support certain JSON schema properties: - - maxItems/minItems: Not supported for array types - - minimum/maximum: Not supported for numeric types - - minLength/maxLength: Not supported for string types + Anthropic's output_format doesn't support certain JSON schema properties. + These are constraints that cannot be enforced by the constrained-decoding + grammar Anthropic compiles the schema into, so the API rejects them with a + 400 ``invalid_request_error`` (e.g. "output_format.schema: For 'array' type, + property 'uniqueItems' is not supported"): + - maxItems/minItems/uniqueItems/contains/minContains/maxContains/prefixItems: array constraints + - minimum/maximum/exclusiveMinimum/exclusiveMaximum/multipleOf: numeric constraints + - minLength/maxLength: string constraints + - minProperties/maxProperties/patternProperties/propertyNames: object constraints + - dependentRequired/dependentSchemas/unevaluatedProperties: object constraints + - if/then/else/not: conditional and negation keywords + + ``oneOf`` is also rejected ("Schema type 'oneOf' is not supported") and is + rewritten to ``anyOf``, matching the Anthropic SDK. Unknown keywords are + ignored by the API, so anything not listed here passes through untouched. This mirrors the transformation done by the Anthropic Python SDK. See: https://platform.claude.com/docs/en/build-with-claude/structured-outputs#how-sdk-transformation-works @@ -504,33 +515,53 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): if not isinstance(schema, dict): return schema - # All numeric/string/array constraints not supported by Anthropic - unsupported_fields = { - "maxItems", - "minItems", # array constraints - "minimum", - "maximum", # numeric constraints - "exclusiveMinimum", - "exclusiveMaximum", # numeric constraints - "minLength", - "maxLength", # string constraints - } - - # Build description additions from removed constraints - constraint_descriptions: list = [] constraint_labels = { "minItems": "minimum number of items: {}", "maxItems": "maximum number of items: {}", + "uniqueItems": "all array items must be unique", + "contains": "array must contain an item matching: {}", + "minContains": "minimum number of matching items: {}", + "maxContains": "maximum number of matching items: {}", + "prefixItems": "leading items must match, in order: {}", "minimum": "minimum value: {}", "maximum": "maximum value: {}", "exclusiveMinimum": "exclusive minimum value: {}", "exclusiveMaximum": "exclusive maximum value: {}", + "multipleOf": "must be a multiple of {}", "minLength": "minimum length: {}", "maxLength": "maximum length: {}", + "minProperties": "minimum number of properties: {}", + "maxProperties": "maximum number of properties: {}", + "patternProperties": "properties whose names match each pattern must satisfy: {}", + "propertyNames": "property names must satisfy: {}", + "dependentRequired": "dependent required properties: {}", + "dependentSchemas": "dependent schemas: {}", + "unevaluatedProperties": "unevaluated properties must satisfy: {}", + "if": "conditional (if): {}", + "then": "conditional (then): {}", + "else": "conditional (else): {}", + "not": "must not match: {}", } - for field in unsupported_fields: - if field in schema: - constraint_descriptions.append(constraint_labels[field].format(schema[field])) + unsupported_fields = set(constraint_labels) + + # Build description additions from removed constraints. Iterating + # constraint_labels (not the set) keeps the note order deterministic across + # processes, so identical requests serialize identically regardless of + # PYTHONHASHSEED and stay cache-friendly. + constraint_descriptions: list = [] + for field, label in constraint_labels.items(): + if field not in schema: + continue + value = schema[field] + # A falsy boolean constraint (e.g. ``uniqueItems: false``) imposes no + # real requirement, so don't add a misleading advisory note for it. + if isinstance(value, bool) and not value: + continue + # Sub-schema constraints (e.g. ``contains``) are serialized as JSON so + # the advisory note preserves what the constraint actually required, + # instead of just noting that it existed. + note_value = json.dumps(value) if isinstance(value, (dict, list)) else value + constraint_descriptions.append(label.format(note_value)) result: Dict[str, Any] = {} @@ -557,11 +588,17 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): elif key == "$defs" and isinstance(value, dict): result[key] = {k: AnthropicConfig.filter_anthropic_output_schema(v) for k, v in value.items()} elif key == "anyOf" and isinstance(value, list): - result[key] = [AnthropicConfig.filter_anthropic_output_schema(item) for item in value] + result["anyOf"] = result.get("anyOf", []) + [ + AnthropicConfig.filter_anthropic_output_schema(item) for item in value + ] elif key == "allOf" and isinstance(value, list): result[key] = [AnthropicConfig.filter_anthropic_output_schema(item) for item in value] elif key == "oneOf" and isinstance(value, list): - result[key] = [AnthropicConfig.filter_anthropic_output_schema(item) for item in value] + # Anthropic rejects oneOf ("Schema type 'oneOf' is not supported"); + # the Anthropic SDK rewrites it to anyOf, so do the same. + result["anyOf"] = result.get("anyOf", []) + [ + AnthropicConfig.filter_anthropic_output_schema(item) for item in value + ] else: result[key] = value diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index c38b3593465..8ce2b982955 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -895,9 +895,8 @@ class AmazonConverseConfig(BaseConfig): if _tool_choice_value is not None: optional_params["tool_choice"] = _tool_choice_value if param == "parallel_tool_calls": - disable_parallel = not value optional_params["_parallel_tool_use_config"] = { - "tool_choice": {"disable_parallel_tool_use": disable_parallel} + "tool_choice": {"type": "auto", "disable_parallel_tool_use": not value} } if param == "thinking": if ( @@ -1208,6 +1207,22 @@ class AmazonConverseConfig(BaseConfig): return {} + @staticmethod + def _merge_parallel_tool_use_config(additional_request_params: dict, parallel_tool_use_config: dict) -> dict: + merged_entries = { + key: ( + { + **value, + **additional_request_params[key], + **{k: v for k, v in value.items() if k != "type"}, + } + if isinstance(additional_request_params.get(key), dict) and isinstance(value, dict) + else value + ) + for key, value in parallel_tool_use_config.items() + } + return {**additional_request_params, **merged_entries} + def _prepare_request_params( self, optional_params: dict, model: str, drop_params: bool = False ) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]: @@ -1276,15 +1291,9 @@ class AmazonConverseConfig(BaseConfig): # Handle parallel_tool_calls configuration parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None) if parallel_tool_use_config is not None and bedrock_converse_supports_parallel_tool_use_config(model): - for key, value in parallel_tool_use_config.items(): - if ( - key in additional_request_params - and isinstance(additional_request_params[key], dict) - and isinstance(value, dict) - ): - additional_request_params[key].update(value) - else: - additional_request_params[key] = value + additional_request_params = self._merge_parallel_tool_use_config( + additional_request_params, parallel_tool_use_config + ) additional_request_params.pop("parallel_tool_calls", None) diff --git a/litellm/llms/bedrock/realtime/handler.py b/litellm/llms/bedrock/realtime/handler.py index b48c37791c4..b7237d288ec 100644 --- a/litellm/llms/bedrock/realtime/handler.py +++ b/litellm/llms/bedrock/realtime/handler.py @@ -9,6 +9,8 @@ import contextlib import json from typing import Any, Optional +from pydantic import TypeAdapter + from litellm._logging import _redact_string, verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging @@ -16,6 +18,8 @@ from ..base_aws_llm import BaseAWSLLM from ..common_utils import BedrockError from .transformation import BedrockRealtimeConfig +_CLIENT_MODALITIES_ADAPTER: TypeAdapter["list[str] | None"] = TypeAdapter(list[str] | None) + class BedrockRealtime(BaseAWSLLM): """Handler for Bedrock Nova Sonic realtime speech-to-speech API.""" @@ -124,6 +128,9 @@ class BedrockRealtime(BaseAWSLLM): verbose_proxy_logger.debug("Bedrock Realtime: Bidirectional stream established") + await websocket.send_text(json.dumps(transformation_config.session_created_event(model, logging_obj))) + verbose_proxy_logger.debug("Bedrock Realtime: sent session.created to client on connect") + # Track state for transformation session_state = { "current_output_item_id": None, @@ -143,6 +150,7 @@ class BedrockRealtime(BaseAWSLLM): transformation_config, model, session_state, + logging_obj, ) ) @@ -179,6 +187,7 @@ class BedrockRealtime(BaseAWSLLM): transformation_config: BedrockRealtimeConfig, model: str, session_state: dict, + logging_obj: LiteLLMLogging | None = None, ): """Forward messages from client WebSocket to Bedrock stream.""" from aws_sdk_bedrock_runtime.models import ( @@ -210,6 +219,23 @@ class BedrockRealtime(BaseAWSLLM): for bedrock_message in transformed_messages: await send_to_bedrock(bedrock_message) + if logging_obj is not None: + client_message_type: str | None = None + requested_modalities: list[str] | None = None + with contextlib.suppress(Exception): + parsed_client_message = json.loads(message) + client_message_type = parsed_client_message.get("type") + if client_message_type == "session.update": + requested_modalities = _CLIENT_MODALITIES_ADAPTER.validate_python( + parsed_client_message.get("session", {}).get("modalities") + ) + if client_message_type == "session.update": + await client_ws.send_text( + json.dumps( + transformation_config.session_updated_event(model, logging_obj, requested_modalities) + ) + ) + except Exception as e: verbose_proxy_logger.debug(f"Client to Bedrock forwarding ended: {e}", exc_info=True) for close_message in transformation_config.session_close_messages(): diff --git a/litellm/llms/bedrock/realtime/transformation.py b/litellm/llms/bedrock/realtime/transformation.py index fe5f0584e03..24a40ebea1b 100644 --- a/litellm/llms/bedrock/realtime/transformation.py +++ b/litellm/llms/bedrock/realtime/transformation.py @@ -623,35 +623,42 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): verbose_logger.warning(f"Unknown message type: {message_type}") return [] - def transform_session_start_event( + def _session_object( self, - event: dict, model: str, logging_obj: LiteLLMLoggingObj, - ) -> OpenAIRealtimeStreamSessionEvents: - """ - Transform Bedrock sessionStart event to OpenAI session.created. - - Args: - event: Bedrock sessionStart event - model: Model ID - logging_obj: Logging object - - Returns: - OpenAI session.created event - """ - verbose_logger.debug("Handling sessionStart") - + modalities: list[str] | None = None, + ) -> OpenAIRealtimeStreamSession: session = OpenAIRealtimeStreamSession( id=logging_obj.litellm_trace_id, - modalities=["text", "audio"], + modalities=modalities if modalities is not None else ["text", "audio"], ) if model is not None and isinstance(model, str): session["model"] = model + return session + def session_created_event( + self, + model: str, + logging_obj: LiteLLMLoggingObj, + ) -> OpenAIRealtimeStreamSessionEvents: + """Build the OpenAI session.created event for this realtime session.""" return OpenAIRealtimeStreamSessionEvents( type="session.created", - session=session, + session=self._session_object(model, logging_obj), + event_id=str(uuid.uuid4()), + ) + + def session_updated_event( + self, + model: str, + logging_obj: LiteLLMLoggingObj, + modalities: list[str] | None = None, + ) -> OpenAIRealtimeStreamSessionEvents: + """Build the OpenAI session.updated ack reflecting the client's requested modalities.""" + return OpenAIRealtimeStreamSessionEvents( + type="session.updated", + session=self._session_object(model, logging_obj, modalities), event_id=str(uuid.uuid4()), ) @@ -1169,8 +1176,6 @@ class BedrockRealtimeConfig(BaseRealtimeConfig): # Route to appropriate transformation method if "sessionStart" in event: - session_created = self.transform_session_start_event(event, model, logging_obj) - returned_messages.append(session_created) session_configuration_request = json.dumps({"configured": True}) elif "contentStart" in event: diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index c48d75439a7..ec1301e5923 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2107,8 +2107,7 @@ class BaseLLMHTTPHandler: rust_messages_response = await self._maybe_rust_anthropic_messages( custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, - stream=stream or False, - rust_stream_eligible=bool(stream) and not self._has_agentic_completion_hook(logging_obj), + has_agentic_hook=self._has_agentic_completion_hook(logging_obj), model=model, api_key=api_key, api_base=api_base, @@ -2266,8 +2265,7 @@ class BaseLLMHTTPHandler: *, custom_llm_provider: str, litellm_params: GenericLiteLLMParams, - stream: bool, - rust_stream_eligible: bool, + has_agentic_hook: bool, model: str, api_key: str | None, api_base: str | None, @@ -2279,7 +2277,7 @@ class BaseLLMHTTPHandler: return None if litellm_params.get("rust") is not True and not BaseLLMHTTPHandler._rust_env_enabled(): return None - if stream and not rust_stream_eligible: + if has_agentic_hook: return None from litellm.rust_bridge import messages as rust_messages_bridge diff --git a/litellm/llms/huggingface/embedding/handler.py b/litellm/llms/huggingface/embedding/handler.py index 39eb430db74..f72a79e084d 100644 --- a/litellm/llms/huggingface/embedding/handler.py +++ b/litellm/llms/huggingface/embedding/handler.py @@ -322,7 +322,7 @@ class HuggingFaceEmbedding(BaseLLM): task = get_hf_task_embedding_for_model(model=model, task_type=task_type, api_base=HF_HUB_URL) # print_verbose(f"{model}, {task}") embed_url = "" - if "https" in model: + if model.startswith(("http://", "https://")): embed_url = model elif api_base: embed_url = api_base diff --git a/litellm/llms/huggingface/embedding/transformation.py b/litellm/llms/huggingface/embedding/transformation.py index 13e38ab5560..6f27e3115eb 100644 --- a/litellm/llms/huggingface/embedding/transformation.py +++ b/litellm/llms/huggingface/embedding/transformation.py @@ -316,25 +316,6 @@ class HuggingFaceEmbeddingConfig(BaseConfig): return data - def get_api_base(self, api_base: Optional[str], model: str) -> str: - """ - Get the API base for the Huggingface API. - - Do not add the chat/embedding/rerank extension here. Let the handler do this. - """ - if "https" in model: - completion_url = model - elif api_base is not None: - completion_url = api_base - elif "HF_API_BASE" in os.environ: - completion_url = os.getenv("HF_API_BASE", "") - elif "HUGGINGFACE_API_BASE" in os.environ: - completion_url = os.getenv("HUGGINGFACE_API_BASE", "") - else: - completion_url = f"https://api-inference.huggingface.co/models/{model}" - - return completion_url - def validate_environment( self, headers: Dict, diff --git a/litellm/llms/oobabooga/chat/oobabooga.py b/litellm/llms/oobabooga/chat/oobabooga.py index fe2bb9dc6d1..40d88e8e125 100644 --- a/litellm/llms/oobabooga/chat/oobabooga.py +++ b/litellm/llms/oobabooga/chat/oobabooga.py @@ -34,7 +34,7 @@ def completion( optional_params=optional_params, litellm_params=litellm_params, ) - if "https" in model: + if model.startswith(("http://", "https://")): completion_url = model elif api_base: completion_url = api_base @@ -96,7 +96,7 @@ def embedding( encoding=None, ): # Create completion URL - if "https" in model: + if model.startswith(("http://", "https://")): embeddings_url = model elif api_base: embeddings_url = f"{api_base}/v1/embeddings" diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index 4e4e088f491..4447d63e5a1 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -143,7 +143,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): raise SagemakerError(status_code=response.status_code, message=response.text) custom_stream_decoder = AWSEventStreamDecoder(model="", is_messages_api=True) - completion_stream = custom_stream_decoder.iter_bytes(response.iter_bytes(chunk_size=1024)) + completion_stream = custom_stream_decoder.iter_bytes(response.iter_bytes()) streaming_response = CustomStreamWrapper( completion_stream=completion_stream, @@ -189,7 +189,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): raise SagemakerError(status_code=response.status_code, message=response.text) custom_stream_decoder = AWSEventStreamDecoder(model="", is_messages_api=True) - completion_stream = custom_stream_decoder.aiter_bytes(response.aiter_bytes(chunk_size=1024)) + completion_stream = custom_stream_decoder.aiter_bytes(response.aiter_bytes()) streaming_response = CustomStreamWrapper( completion_stream=completion_stream, diff --git a/litellm/llms/sagemaker/completion/handler.py b/litellm/llms/sagemaker/completion/handler.py index 4b87271fd44..c27a3c3528c 100644 --- a/litellm/llms/sagemaker/completion/handler.py +++ b/litellm/llms/sagemaker/completion/handler.py @@ -200,23 +200,12 @@ class SagemakerLLM(BaseAWSLLM): # Add model_id as InferenceComponentName header # boto3 doc: https://docs.aws.amazon.com/sagemaker/latest/APIReference/API_runtime_InvokeEndpoint.html prepared_request.headers.update({"X-Amzn-SageMaker-Inference-Component": model_id}) - sync_handler = _get_httpx_client() - sync_response = sync_handler.post( - url=prepared_request.url, + completion_stream = self.make_sync_call( + api_base=prepared_request.url, headers=prepared_request.headers, # type: ignore - data=prepared_request.body, - stream=stream, + data=cast(str, prepared_request.body), # cast-ok: signed body is a JSON str, mirrors async path + logging_obj=logging_obj, ) - - if sync_response.status_code != 200: - raise SagemakerError( - status_code=sync_response.status_code, - message=str(sync_response.read()), - ) - - decoder = AWSEventStreamDecoder(model="") - - completion_stream = decoder.iter_bytes(sync_response.iter_bytes(chunk_size=1024)) streaming_response = CustomStreamWrapper( completion_stream=completion_stream, model=model, @@ -334,6 +323,29 @@ class SagemakerLLM(BaseAWSLLM): litellm_params=litellm_params, ) + def make_sync_call( + self, + api_base: str, + headers: dict, + data: str, + logging_obj, + client=None, + ): + if client is None: + client = _get_httpx_client() + sync_response = client.post( + api_base, + headers=headers, + data=data, + stream=True, + ) + + if sync_response.status_code != 200: + raise SagemakerError(status_code=sync_response.status_code, message=str(sync_response.read())) + + decoder = AWSEventStreamDecoder(model="") + return decoder.iter_bytes(sync_response.iter_bytes()) + async def make_async_call( self, api_base: str, @@ -358,7 +370,7 @@ class SagemakerLLM(BaseAWSLLM): raise SagemakerError(status_code=response.status_code, message=response.text) decoder = AWSEventStreamDecoder(model="") - completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=1024)) + completion_stream = decoder.aiter_bytes(response.aiter_bytes()) return completion_stream diff --git a/litellm/main.py b/litellm/main.py index fb05a375111..dc3ec469a1b 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -5111,7 +5111,10 @@ def completion( # type: ignore try: if base_url is not None: api_base = base_url - if num_retries is not None: + is_router_call = any("model_group" in (kwargs.get(k) or ()) for k in ("metadata", "litellm_metadata")) + if is_router_call: + max_retries = 0 + elif num_retries is not None: max_retries = num_retries logging: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, litellm_logging_obj) fallbacks = fallbacks or litellm.model_fallbacks diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index bb6243e50ed..d3917886060 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -17566,6 +17566,61 @@ }, "web_search_billing_unit": "per_query" }, + "gemini-3.5-flash-lite": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 2e-08, + "cache_read_input_token_cost_priority": 5e-08, + "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, @@ -18233,6 +18288,60 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.6-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_flex": 7.5e-08, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_batches": 7.5e-07, + "input_cost_per_token_flex": 7.5e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 7.5e-06, + "output_cost_per_token": 7.5e-06, + "output_cost_per_token_batches": 3.75e-06, + "output_cost_per_token_flex": 3.75e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "output_cost_per_token_priority": 1.35e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "vertex_ai/gemini-3.1-pro-preview": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, @@ -19585,6 +19694,63 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.5-flash-lite": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 2e-08, + "cache_read_input_token_cost_priority": 5e-08, + "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "litellm_provider": "gemini", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, + "rpm": 15, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 250000, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-3-flash-preview": { "cache_read_input_token_cost": 5e-08, "input_cost_per_audio_token": 1e-06, @@ -19691,6 +19857,63 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.6-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_flex": 7.5e-08, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_batches": 7.5e-07, + "input_cost_per_token_flex": 7.5e-07, + "litellm_provider": "gemini", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 7.5e-06, + "output_cost_per_token": 7.5e-06, + "output_cost_per_token_batches": 3.75e-06, + "output_cost_per_token_flex": 3.75e-06, + "rpm": 2000, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 800000, + "input_cost_per_token_priority": 2.7e-06, + "output_cost_per_token_priority": 1.35e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-omni-flash-preview": { "input_cost_per_audio_token": 1.5e-06, "input_cost_per_token": 1.5e-06, @@ -19971,6 +20194,61 @@ }, "web_search_billing_unit": "per_query" }, + "gemini-3.6-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_flex": 7.5e-08, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_batches": 7.5e-07, + "input_cost_per_token_flex": 7.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 7.5e-06, + "output_cost_per_token": 7.5e-06, + "output_cost_per_token_batches": 3.75e-06, + "output_cost_per_token_flex": 3.75e-06, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "output_cost_per_token_priority": 1.35e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, @@ -37232,6 +37510,61 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.5-flash-lite": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 2e-08, + "cache_read_input_token_cost_priority": 5e-08, + "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "vertex_ai/deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 882c34dbd6a..ee3196b539c 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -3,6 +3,7 @@ import html as _html import json import secrets import time +from collections.abc import Mapping from datetime import datetime, timezone from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse @@ -13,6 +14,7 @@ from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Resp from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError from litellm._logging import verbose_logger +from litellm.caching.in_memory_cache import InMemoryCache from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -20,7 +22,9 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( TokenEndpointAuthConfigError, build_token_endpoint_client_auth, + normalize_token_endpoint_auth_method, ) +from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod from litellm.proxy._experimental.mcp_server.bridge_token_flow import ( _bridge_mint_error_response, _BridgeMintReady, @@ -111,6 +115,9 @@ def encode_state_with_base_url( client_redirect_uri: Optional[str] = None, litellm_user_id: str | None = None, mcp_server_id: str | None = None, + dcr_client_id: str | None = None, + dcr_client_secret: str | None = None, + dcr_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, ) -> str: """ Encode the base_url, original state, and PKCE parameters using encryption. @@ -124,8 +131,18 @@ def encode_state_with_base_url( litellm_user_id: The SSO-authenticated litellm user captured at the bridge authorize (interactive dcr_bridge oauth_delegate only); the callback seals it into the gateway authorization code so the token mint can bind the envelope to this user - mcp_server_id: The bridge server the interactive flow targets, sealed alongside - litellm_user_id so the gateway code cannot be replayed against another server + mcp_server_id: The server the flow targets, sealed alongside litellm_user_id (bridge) or + dcr_client_id (ephemeral mint) so the gateway code cannot be replayed against another + server + dcr_client_id: The ephemeral DCR client the gateway minted at authorize for a + client-forwarded-token server with no caller-supplied client; the callback seals it + into the forwarded authorization code so the token exchange can authenticate with it + while the gateway stores nothing + dcr_client_secret: The minted client's secret, when the upstream issued one + dcr_token_endpoint_auth_method: The token-endpoint auth method the upstream's registration + response granted the minted client, sealed alongside the credentials so the exchange + authenticates the way the upstream expects instead of falling back to the server row's + configured method Returns: An encrypted string that encodes all values @@ -138,6 +155,9 @@ def encode_state_with_base_url( "client_redirect_uri": client_redirect_uri, "litellm_user_id": litellm_user_id, "mcp_server_id": mcp_server_id, + "dcr_client_id": dcr_client_id, + "dcr_client_secret": dcr_client_secret, + "dcr_token_endpoint_auth_method": dcr_token_endpoint_auth_method, } state_json = json.dumps(state_data, sort_keys=True) encrypted_state = encrypt_value_helper(state_json) @@ -217,6 +237,93 @@ def open_bridge_authorization_code(code: str) -> _BridgeAuthorizationCode | None return None +_PASSTHROUGH_AUTH_CODE_PREFIX = "llm_ptcode_" + + +class PassthroughAuthorizationCode(BaseModel): + """The ephemeral DCR client and upstream code the gateway seals into the authorization code it + forwards for a client-forwarded-token server (``true_passthrough`` / ``oauth_delegate``) whose + authorize fell through to gateway-side registration. These modes forbid the gateway from storing + an OAuth client identity, so the minted client survives only inside this sealed value: the + client echoes it back at the token endpoint, where the gateway recovers the client to + authenticate the upstream exchange. ``mcp_server_id`` binds the code to the server it was minted + for so it cannot be spent at another server's token endpoint.""" + + model_config = ConfigDict(frozen=True) + upstream_code: str = Field(min_length=1) + client_id: str = Field(min_length=1) + client_secret: str | None = None + token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None + mcp_server_id: str = Field(min_length=1) + + +def seal_passthrough_authorization_code( + upstream_code: str, + client_id: str, + client_secret: str | None, + mcp_server_id: str, + token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, +) -> str: + """Seal the upstream authorization code together with the ephemeral DCR client that authorized + it. Encrypted with the same authenticated symmetric helper as the OAuth state and bridge codes, + so the client can neither read the (possibly confidential) client credentials nor forge a + code.""" + payload = json.dumps( + { + "upstream_code": upstream_code, + "client_id": client_id, + "client_secret": client_secret, + "token_endpoint_auth_method": token_endpoint_auth_method, + "mcp_server_id": mcp_server_id, + }, + sort_keys=True, + ) + return _PASSTHROUGH_AUTH_CODE_PREFIX + encrypt_value_helper(payload) + + +def open_passthrough_authorization_code(code: str) -> PassthroughAuthorizationCode | None: + """Recover the sealed ephemeral client and upstream code, or ``None`` when ``code`` is not a + gateway passthrough code or does not decrypt / validate, so a raw upstream code falls through to + the existing caller-supplied-client behavior.""" + if not code.startswith(_PASSTHROUGH_AUTH_CODE_PREFIX): + return None + decrypted = decrypt_value_helper( + code[len(_PASSTHROUGH_AUTH_CODE_PREFIX) :], "passthrough_authorization_code", return_original_value=False + ) + if not isinstance(decrypted, str): + return None + try: + return PassthroughAuthorizationCode.model_validate_json(decrypted) + except ValidationError: + return None + + +def redeem_passthrough_authorization_code( + code: str | None, mcp_server: MCPServer, code_verifier: str | None +) -> PassthroughAuthorizationCode | None: + """The single redemption gate for sealed passthrough codes: a raw or foreign code returns + ``None`` so the caller keeps its existing behavior, while a genuine sealed code must be spent + at the server it was minted for and must carry the PKCE verifier of the S256 flow that minted + it (the mint refuses downgraded flows, so a verifier-less redemption is an interception + attempt, not a legitimate client).""" + if not code: + return None + sealed = open_passthrough_authorization_code(code) + if sealed is None: + return None + if sealed.mcp_server_id != mcp_server.server_id: + raise HTTPException( + status_code=400, + detail="Authorization code was issued for a different MCP server", + ) + if not code_verifier: + raise HTTPException( + status_code=400, + detail="code_verifier is required to redeem this authorization code", + ) + return sealed + + def _redirect_to_litellm_login(request: Request) -> RedirectResponse: """Send an unauthenticated browser through litellm login before the interactive bridge authorize can capture its identity. The bridge oauth_delegate flow seals the SSO user into the gateway code, @@ -594,10 +701,18 @@ async def authorize_with_server( code_challenge_method: Optional[str] = None, response_type: Optional[str] = None, scope: Optional[str] = None, + ephemeral_dcr_client: "EphemeralDcrClient | None" = None, ): _raise_if_not_oauth2(mcp_server) if mcp_server.authorization_url is None: - raise HTTPException(status_code=400, detail="MCP server authorization url is not set") + raise HTTPException( + status_code=400, + detail=( + "MCP server authorization url is not configured. Servers with no url (OpenAPI " + "spec or stdio) run no resource discovery, so set Authorization URL and Token URL " + "manually, or set Issuer to discover them from the identity provider (RFC 8414)." + ), + ) if mcp_server.is_dcr_bridge: # Enforce S256 PKCE on both bridge arms. The relay arm forwards the validated, @@ -605,7 +720,10 @@ async def authorize_with_server( # calling this for its enforcement side effect, then falls through to the gateway # /callback flow below, which reads the original code_challenge names. bridge_challenge, bridge_method = _require_s256_pkce(code_challenge, code_challenge_method) - if _dcr_bridge_relays_client_registration(mcp_server): + # A gateway-minted ephemeral client is registered against {base}/callback, so its + # flow must run the short-circuit arm; the relay arm is only for clients that + # registered themselves through the front door and hold their own redirect binding. + if _dcr_bridge_relays_client_registration(mcp_server) and ephemeral_dcr_client is None: return _redirect_to_upstream_authorize( mcp_server=mcp_server, client_id=client_id, @@ -649,7 +767,12 @@ async def authorize_with_server( code_challenge_method=code_challenge_method, client_redirect_uri=redirect_uri, litellm_user_id=litellm_user_id, - mcp_server_id=mcp_server.server_id if litellm_user_id else None, + mcp_server_id=mcp_server.server_id if (litellm_user_id or ephemeral_dcr_client) else None, + dcr_client_id=ephemeral_dcr_client.client_id if ephemeral_dcr_client else None, + dcr_client_secret=ephemeral_dcr_client.client_secret if ephemeral_dcr_client else None, + dcr_token_endpoint_auth_method=ephemeral_dcr_client.token_endpoint_auth_method + if ephemeral_dcr_client + else None, ) relay_state = secrets.token_urlsafe(_OAUTH_STATE_HANDLE_BYTES) @@ -696,23 +819,40 @@ async def exchange_token_with_server( code_verifier: Optional[str], refresh_token: Optional[str] = None, scope: Optional[str] = None, + client_token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None, ): _raise_if_not_oauth2(mcp_server) if grant_type not in ("authorization_code", "refresh_token"): raise HTTPException(status_code=400, detail="Unsupported grant_type") if mcp_server.token_url is None: - raise HTTPException(status_code=400, detail="MCP server token url is not set") + raise HTTPException( + status_code=400, + detail=( + "MCP server token url is not configured. Servers with no url (OpenAPI spec or " + "stdio) run no resource discovery, so set Token URL manually, or set Issuer to " + "discover it from the identity provider (RFC 8414)." + ), + ) - # The id and secret must come from the same source. When the server-side client_id wins, - # falling back to the caller's secret pairs the persisted client with a foreign secret; the - # register short-circuit hands clients a placeholder secret ("dummy"), so a re-auth against a - # persisted public PKCE client (no stored secret) would send that placeholder and the IdP 401s. + # The id, secret, and token-endpoint auth method must come from the same source. When the + # server-side client_id wins, falling back to the caller's secret pairs the persisted client + # with a foreign secret; the register short-circuit hands clients a placeholder secret + # ("dummy"), so a re-auth against a persisted public PKCE client (no stored secret) would send + # that placeholder and the IdP 401s. Symmetrically, a caller-side client (an ephemeral mint + # recovered from a sealed code) must authenticate the way its own registration was granted, + # not the way the server row is configured; callers that carry no method keep the row's method + # as before. resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id resolved_client_secret = mcp_server.client_secret if mcp_server.client_id else client_secret + resolved_auth_method = ( + mcp_server.token_endpoint_auth_method + if mcp_server.client_id + else (client_token_endpoint_auth_method or mcp_server.token_endpoint_auth_method) + ) try: client_auth = build_token_endpoint_client_auth( - auth_method=mcp_server.token_endpoint_auth_method, + auth_method=resolved_auth_method, client_id=resolved_client_id, client_secret=resolved_client_secret, ) @@ -1215,7 +1355,7 @@ async def _persist_dcr_client_registration( return "failed" -def _client_supplied_redirect_uris(value: object) -> list[str] | None: +def client_supplied_redirect_uris(value: object) -> list[str] | None: """RFC 7591 redirect_uris must be a non-empty array of URI strings. Any other shape (not a list, an empty list, or a list holding a non-string or empty-string element) yields None so every register arm falls back to the gateway callback instead of echoing a malformed value back to the @@ -1227,6 +1367,142 @@ def _client_supplied_redirect_uris(value: object) -> list[str] | None: return uris if len(uris) == len(value) else None +async def _post_dcr_registration( + registration_url: str, + register_data: Mapping[str, object], + server_id: str, +) -> httpx.Response: + """POST an RFC 7591 registration to the upstream and return its response, relaying a classified + upstream rejection instead of a generic 500 and failing loud on an absent response.""" + headers = { + "Content-Type": "application/json", + "Accept": "application/json", + } + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register) + try: + response = await async_client.post( + registration_url, + headers=headers, + json=register_data, + ) + if response is not None: + response.raise_for_status() + except httpx.HTTPStatusError as exc: + status_code, detail = dcr_fault_detail(classify_upstream_dcr_rejection(exc.response, log_context=server_id)) + raise HTTPException(status_code=status_code, detail=detail) from exc + if response is None: + raise HTTPException( + status_code=502, + detail="MCP upstream registration endpoint returned no response", + ) + return response + + +class EphemeralDcrClient(BaseModel): + """A DCR client minted for a single authorize round trip and never stored by the gateway.""" + + model_config = ConfigDict(frozen=True) + client_id: str = Field(min_length=1) + client_secret: str | None = None + token_endpoint_auth_method: MCPTokenEndpointAuthMethod | None = None + + +_EPHEMERAL_DCR_CLIENT_CACHE = InMemoryCache(default_ttl=_OAUTH_STATE_COOKIE_TTL_SECONDS) +_EPHEMERAL_DCR_MINT_LOCKS: dict[str, asyncio.Lock] = {} + + +async def mint_ephemeral_dcr_client(request: Request, mcp_server: MCPServer) -> EphemeralDcrClient | None: + """Mint a throwaway OAuth client via the upstream's RFC 7591 registration endpoint for a + client-forwarded-token server whose authorize arrived with no client_id. Returns ``None`` when + the upstream exposes no registration endpoint, so the caller keeps its existing failure path. + The minted client is deliberately not persisted anywhere: ``true_passthrough`` / + ``oauth_delegate`` require the gateway to hold no OAuth client identity, so it survives only in + the encrypted OAuth state and the sealed authorization code the callback forwards. + + Reloading the authorize page or retrying a flow must not register a fresh upstream client every + time (an OAuth client identifies the application, not the user, so reuse is semantically + correct). A per-process TTL cache bounded to the OAuth state cookie's lifetime dedupes the mint + per (server, gateway origin), and a per-server lock single-flights concurrent mints (the + ``_OAUTH_METADATA_FETCH_LOCKS`` pattern; keyed by server_id alone so the lock registry stays + bounded by the server count even when the request origin varies) so parallel authorize requests + cannot each register an upstream client; the cache stamps nothing onto the server record and + correctness never depends on it because the sealed state carries the client through the flow.""" + if mcp_server.registration_url is None: + return None + request_base_url = get_request_base_url(request) + cache_key = f"mcp_ephemeral_dcr_client:{mcp_server.server_id}:{request_base_url}" + cached = _EPHEMERAL_DCR_CLIENT_CACHE.get_cache(cache_key) + if isinstance(cached, EphemeralDcrClient): + return cached + lock = _EPHEMERAL_DCR_MINT_LOCKS.setdefault(mcp_server.server_id, asyncio.Lock()) + async with lock: + cached_after_wait = _EPHEMERAL_DCR_CLIENT_CACHE.get_cache(cache_key) + if isinstance(cached_after_wait, EphemeralDcrClient): + return cached_after_wait + register_data: dict[str, object] = { + "client_name": mcp_server.server_name or mcp_server.server_id, + "redirect_uris": [f"{request_base_url}/callback"], + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "token_endpoint_auth_method": "none", + } + response = await _post_dcr_registration( + registration_url=mcp_server.registration_url, + register_data=register_data, + server_id=mcp_server.server_id, + ) + try: + registration = _DcrClientRegistration.model_validate_json(response.text) + except ValidationError as exc: + raise HTTPException( + status_code=502, + detail="MCP upstream registration endpoint returned no usable client_id", + ) from exc + if not registration.client_id: + raise HTTPException( + status_code=502, + detail="MCP upstream registration endpoint returned no usable client_id", + ) + minted = EphemeralDcrClient( + client_id=registration.client_id, + client_secret=registration.client_secret, + token_endpoint_auth_method=normalize_token_endpoint_auth_method(registration.token_endpoint_auth_method), + ) + _EPHEMERAL_DCR_CLIENT_CACHE.set_cache(cache_key, minted) + return minted + + +async def resolve_ephemeral_dcr_client( + request: Request, + mcp_server: MCPServer, + code_challenge: str | None, + code_challenge_method: str | None, + redirect_uri: str, +) -> EphemeralDcrClient | None: + """The single owner of the gateway-side mint policy for a clientless authorize. Returns + ``None`` for servers whose mode does not permit gateway minting and for upstreams without a + registration endpoint, so those callers keep their existing failure paths: plain ``oauth2`` + keeps its persisted-client contract, and the interactive ``oauth_delegate`` dcr_bridge + sign-in has its own sealed-identity flow. ``true_passthrough`` mints regardless of the + ``dcr_bridge`` flag (the UI creates passthrough servers with the flag on by default): a + minted flow runs the bridge short-circuit arm, while the relay front door remains for + external clients that registered themselves. Flows that could never succeed fail loud + before any upstream registration: a missing ``authorization_url``, a downgraded PKCE pair + (without S256 the sealed code would be bearer-redeemable by any authenticated caller who + intercepts the redirect), or an untrusted ``redirect_uri`` (a rejected redirect must not be + usable to generate orphan IdP clients).""" + if not (mcp_server.is_true_passthrough or (mcp_server.is_oauth_delegate and not mcp_server.is_dcr_bridge)): + return None + if mcp_server.authorization_url is None: + raise HTTPException( + status_code=400, + detail="MCP server authorization url is not set", + ) + _require_s256_pkce(code_challenge, code_challenge_method) + validate_trusted_redirect_uri(request, redirect_uri) + return await mint_ephemeral_dcr_client(request, mcp_server) + + async def register_client_with_server( request: Request, mcp_server: MCPServer, @@ -1262,7 +1538,14 @@ async def register_client_with_server( return dummy_return if mcp_server.authorization_url is None: - raise HTTPException(status_code=400, detail="MCP server authorization url is not set") + raise HTTPException( + status_code=400, + detail=( + "MCP server authorization url is not configured. Servers with no url (OpenAPI " + "spec or stdio) run no resource discovery, so set Authorization URL and Token URL " + "manually, or set Issuer to discover them from the identity provider (RFC 8414)." + ), + ) if mcp_server.registration_url is None: return dummy_return @@ -1281,30 +1564,11 @@ async def register_client_with_server( "response_types": response_types or (["code"] if bridge_relay else []), "token_endpoint_auth_method": token_endpoint_auth_method or ("none" if bridge_relay else ""), } - headers = { - "Content-Type": "application/json", - "Accept": "application/json", - } - - async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Register) - try: - response = await async_client.post( - mcp_server.registration_url, - headers=headers, - json=register_data, - ) - if response is not None: - response.raise_for_status() - except httpx.HTTPStatusError as exc: - status_code, detail = dcr_fault_detail( - classify_upstream_dcr_rejection(exc.response, log_context=mcp_server.server_id) - ) - raise HTTPException(status_code=status_code, detail=detail) from exc - if response is None: - raise HTTPException( - status_code=502, - detail="MCP upstream registration endpoint returned no response", - ) + response = await _post_dcr_registration( + registration_url=mcp_server.registration_url, + register_data=register_data, + server_id=mcp_server.server_id, + ) token_response = response.json() @@ -1542,11 +1806,23 @@ async def callback( # envelope to this user. Every other flow forwards the raw code unchanged. litellm_user_id = state_data.get("litellm_user_id") mcp_server_id = state_data.get("mcp_server_id") + dcr_client_id = state_data.get("dcr_client_id") + dcr_client_secret = state_data.get("dcr_client_secret") forwarded_code = code if isinstance(litellm_user_id, str) and litellm_user_id and isinstance(mcp_server_id, str) and mcp_server_id: forwarded_code = seal_bridge_authorization_code( upstream_code=code, litellm_user_id=litellm_user_id, mcp_server_id=mcp_server_id ) + elif isinstance(dcr_client_id, str) and dcr_client_id and isinstance(mcp_server_id, str) and mcp_server_id: + forwarded_code = seal_passthrough_authorization_code( + upstream_code=code, + client_id=dcr_client_id, + client_secret=dcr_client_secret if isinstance(dcr_client_secret, str) and dcr_client_secret else None, + mcp_server_id=mcp_server_id, + token_endpoint_auth_method=normalize_token_endpoint_auth_method( + state_data.get("dcr_token_endpoint_auth_method") + ), + ) params = {"code": forwarded_code, "state": original_state} complete_returned_url = _append_query_params(redirect_uri, params) @@ -2137,7 +2413,7 @@ async def register_client(request: Request, mcp_server_name: Optional[str] = Non request_data = await _read_request_body(request=request) data: dict = {**request_data} - client_redirect_uris = _client_supplied_redirect_uris(data.get("redirect_uris")) + client_redirect_uris = client_supplied_redirect_uris(data.get("redirect_uris")) dummy_return = { "client_id": mcp_server_name or "dummy_client", diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8f30071eb5d..90b70dd01f2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -224,6 +224,20 @@ def _uses_issuer_anchor(manual_issuer: str | None, is_discovery_auth_type: bool) return _blank_to_none(manual_issuer) is not None and is_discovery_auth_type +def _has_oauth_discovery_source(server_url: str | None, use_issuer_anchor: bool) -> bool: + """Whether the server has any source OAuth discovery can fetch metadata from. + + Resource-rooted discovery (RFC 9728) is fetched from the server ``url``, so spec-only + (OpenAPI) and stdio servers, which have none, could never discover: their OAuth endpoints + stayed unset unless entered manually and ``/authorize`` served its 400 with no hint of why. + An admin-pinned issuer is a trust anchor in its own right (RFC 8414 section 3.3) whose + metadata fetch does not touch the resource at all, so an anchored server can discover with + no ``url``. Called by both build paths (config and DB) so the two cannot disagree on when + discovery is reachable. + """ + return bool(server_url) or use_issuer_anchor + + def _endpoints_yield_to_issuer( issuer: str | None, is_discovery_auth_type: bool, @@ -610,6 +624,34 @@ def _passthrough_token_from_mcp_auth_header( return None +async def _materialize_auth_headers(auth: httpx.Auth | None) -> dict[str, str] | None: + """Extract the header a resolved ``httpx.Auth`` would set, as a plain dict, or None. + + OpenAPI tool closures egress through ``AsyncHTTPHandler`` methods that accept headers but no + ``auth``, so a resolved credential must be materialized into a header value. Driving one step + of the auth's own flow (against a throwaway request that is never sent) keeps this generic + across every auth shape without per-class branching; ``header_name`` is the resolver-arm + convention for "this auth sets a header" (``NoOpAuth`` has none and yields nothing to apply). + The materialized value is point-in-time: flow behaviors past the first request, like the M2M + one-shot 401 refetch, do not apply on this arm. + """ + if auth is None: + return None + header_name = getattr(auth, "header_name", None) + if not isinstance(header_name, str) or not header_name: + return None + probe = httpx.Request("GET", "http://localhost/") + flow = auth.async_auth_flow(probe) + try: + first_request = await flow.__anext__() + except StopAsyncIteration: + return None + finally: + await flow.aclose() + header_value = first_request.headers.get(header_name) + return {header_name: header_value} if header_value else None + + def _consumes_caller_authorization(server: MCPServer) -> bool: """True when this server's egress forwards the caller's request-wide ``Authorization`` upstream: the client-forwarded token modes, legacy OAuth pass-through, and legacy upstream-delegated @@ -1226,7 +1268,12 @@ class MCPServerManager: manual_token_url = _blank_to_none(server_config.get("token_url")) manual_registration_url = _blank_to_none(server_config.get("registration_url")) is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES - use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type) + obo_needs_discovery = self._obo_needs_endpoint_discovery( + auth_type, + server_config.get("token_exchange_endpoint"), + manual_token_url, + ) + use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type or obo_needs_discovery) manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer( manual_issuer, is_discovery_auth_type, @@ -1234,17 +1281,12 @@ class MCPServerManager: manual_token_url, manual_registration_url, ) - should_discover = bool(server_url) and ( - is_discovery_auth_type - or self._obo_needs_endpoint_discovery( - auth_type, - server_config.get("token_exchange_endpoint"), - manual_token_url, - ) + should_discover = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( + is_discovery_auth_type or obo_needs_discovery ) if not should_discover: mcp_oauth_metadata = None - elif manual_issuer is not None and is_discovery_auth_type: + elif use_issuer_anchor and manual_issuer is not None: mcp_oauth_metadata = await self._fetch_issuer_anchored_oauth_metadata(manual_issuer, server_url) else: mcp_oauth_metadata = await self._descovery_metadata( @@ -1640,7 +1682,7 @@ class MCPServerManager: token_exchange_endpoint: Optional[str], ) -> Optional[MCPOAuthMetadata]: has_all_upstream_oauth_fields = bool(manual_authorization_url and manual_token_url and scopes) - needs_discovery = bool(server_url) and ( + needs_discovery = _has_oauth_discovery_source(server_url, use_issuer_anchor) and ( (is_discovery_auth_type and not has_all_upstream_oauth_fields) or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url) ) @@ -1759,13 +1801,17 @@ class MCPServerManager: manual_token_url = _blank_to_none(mcp_server.token_url) manual_registration_url = _blank_to_none(mcp_server.registration_url) is_discovery_auth_type = auth_type in _UPSTREAM_OAUTH_DISCOVERY_AUTH_TYPES - use_issuer_anchor = _uses_issuer_anchor(manual_issuer, is_discovery_auth_type) - manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer( - manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url - ) token_exchange_endpoint = mcp_server.token_exchange_endpoint or ( credentials_dict.get("token_exchange_endpoint") if credentials_dict else None ) + use_issuer_anchor = _uses_issuer_anchor( + manual_issuer, + is_discovery_auth_type + or self._obo_needs_endpoint_discovery(auth_type, token_exchange_endpoint, manual_token_url), + ) + manual_authorization_url, manual_token_url, manual_registration_url = _endpoints_yield_to_issuer( + manual_issuer, is_discovery_auth_type, manual_authorization_url, manual_token_url, manual_registration_url + ) gated_oauth_metadata = await self._resolve_table_oauth_metadata( mcp_server=mcp_server, auth_type=auth_type, @@ -1943,7 +1989,7 @@ class MCPServerManager: family: discovered ``authorization_url``/``token_url``/``scopes`` otherwise live only on the in-memory registry entry, which is rebuilt on every client connect (the DCR reuse path calls ``update_server``) and on every post-write DB reload, so one failed re-discovery - serves 400 "authorization url is not set" from /authorize until a later rebuild succeeds. + serves the 400 "authorization url is not configured" from /authorize until a later rebuild succeeds. Only fills row fields that are currently empty, never persists origin-fallback guesses (RFC 9728/8414-advertised metadata only), and deliberately skips ``registration_url`` because ``_dcr_bridge_relays_client_registration`` keys off that column. Best-effort: a @@ -4705,6 +4751,61 @@ class MCPServerManager: ) return oauth2_headers + async def resolve_openapi_upstream_auth( + self, + *, + mcp_server: MCPServer, + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + mcp_auth_header: str | dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None, + forwarded_headers: dict[str, str] | None, + ) -> tuple[dict[str, str] | None, dict[str, str] | None]: + """Resolve the gateway-owned upstream credential for a spec_path (OpenAPI) tool call. + + OpenAPI tools egress through a plain httpx call assembled from ContextVars, never through + ``_create_mcp_client``, so the v2 resolver graft there does not run for them and a resolved + credential (authorization_code's stored per-user token, client_credentials' minted M2M + token, token_exchange's exchanged token, passthrough's forwarded caller token) must be + materialized into headers here. Returns ``(resolved_auth_headers, forwarded_headers)``: + the resolved headers are authoritative over every other Authorization source (the same + rule ``_resolve_v2_auth`` applies on the MCPClient path) and ``forwarded_headers`` comes + back with any header the resolver claimed already dropped. Unmigrated (v1) servers resolve + through the stored-token lookup instead, and a missing per-user credential raises the same + discovery challenge the MCPClient path serves, rather than egressing unauthenticated. + + The resolved headers carry only credentials the gateway itself resolved (a stored per-user + token, a minted or exchanged token). Caller-supplied ``oauth2_headers`` are never promoted + into them: on the v2 arm they feed only subject-token extraction (the designed RFC 8693 + input), and on the v1 arm their presence disables the stored lookup entirely, so a + caller's gateway credential can never displace a per-server BYOK header or leak upstream + as the resolved credential. + """ + spec = to_server_spec(mcp_server) + if spec is None: + if oauth2_headers: + return None, forwarded_headers + stored_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, None, user_api_key_auth) + return stored_headers, forwarded_headers + + subject_token: str | None = None + if isinstance(spec.config, (TokenExchangeConfig, IdJagConfig)): + subject_token = self._extract_bearer_token(oauth2_headers, raw_headers) + elif isinstance(spec.config, PassthroughConfig): + inbound_token, forwarded_headers = _take_forwarded_authorization(forwarded_headers) + per_server_token = _passthrough_token_from_mcp_auth_header(mcp_auth_header) + subject_token = per_server_token if per_server_token is not None else inbound_token + + resolved_auth, forwarded_headers = await self._resolve_v2_auth( + server=mcp_server, + spec=spec, + provider=self._cred_provider, + subject_token=subject_token, + user_api_key_auth=user_api_key_auth, + extra_headers=forwarded_headers, + ) + return await _materialize_auth_headers(resolved_auth), forwarded_headers + async def _gather_openapi_tool_tasks( self, tasks: list[Any], @@ -4796,6 +4897,7 @@ class MCPServerManager: ) tasks.append(during_hook_task) + caller_oauth2_headers = oauth2_headers oauth2_headers = await self._resolve_oauth2_headers_for_tool_call(mcp_server, oauth2_headers, user_api_key_auth) # For OpenAPI servers, call the tool handler directly instead of via MCP client @@ -4813,22 +4915,32 @@ class MCPServerManager: auth_header_value = ( _format_byok_openapi_auth_header(mcp_server, mcp_auth_header) if mcp_auth_header else None ) - forwarded_headers = _openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth) + resolved_auth_headers, forwarded_headers = await self.resolve_openapi_upstream_auth( + mcp_server=mcp_server, + oauth2_headers=caller_oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + user_api_key_auth=user_api_key_auth, + forwarded_headers=_openapi_forwarded_extra_headers(mcp_server, raw_headers, user_api_key_auth), + ) async def _call_openapi_via_handler(): from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_auth_header, _request_extra_headers, + _request_resolved_auth_headers, ) auth_token = _request_auth_header.set(auth_header_value) extra_token = _request_extra_headers.set(forwarded_headers) + resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers) try: async with self._limit_outbound_concurrency(mcp_server): return await self._call_openapi_tool_handler(mcp_server, name, arguments) finally: _request_auth_header.reset(auth_token) _request_extra_headers.reset(extra_token) + _request_resolved_auth_headers.reset(resolved_token) tasks.append(asyncio.create_task(_call_openapi_via_handler())) else: diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 1ee300be718..0b795057837 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -62,6 +62,14 @@ _request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = conte "_request_extra_headers", default=None ) +# Per-request headers carrying the gateway-resolved upstream credential +# (stored per-user OAuth token, minted M2M token, exchanged OBO token). +# Set from MCPServerManager.resolve_openapi_upstream_auth; authoritative +# over every other Authorization source in _merge_openapi_tool_request_headers. +_request_resolved_auth_headers: contextvars.ContextVar[dict[str, str] | None] = contextvars.ContextVar( + "_request_resolved_auth_headers", default=None +) + def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str: """Ensure path params cannot introduce directory traversal.""" @@ -294,10 +302,15 @@ def _merge_openapi_tool_request_headers( """Merge static closure headers with per-request ContextVar overrides. Precedence (highest to lowest): - 1. ``_request_auth_header`` — BYOK override of ``Authorization`` - 2. ``static_headers`` — operator-configured headers baked into the + 1. ``_request_resolved_auth_headers`` — the gateway-resolved upstream + credential (stored per-user OAuth token, minted M2M token, + exchanged OBO token). The resolver is authoritative: a BYOK or + forwarded ``Authorization`` must not shadow it, mirroring + ``_resolve_v2_auth`` on the MCPClient path + 2. ``_request_auth_header`` — BYOK override of ``Authorization`` + 3. ``static_headers`` — operator-configured headers baked into the tool closure at registration time - 3. ``_request_extra_headers`` — per-request headers forwarded from + 4. ``_request_extra_headers`` — per-request headers forwarded from the MCP caller (allowlisted by ``MCPServer.extra_headers``) This matches the existing MCP invariant in @@ -323,6 +336,12 @@ def _merge_openapi_tool_request_headers( del effective_headers[existing] effective_headers["Authorization"] = override_auth + resolved_auth_headers = _request_resolved_auth_headers.get() or {} + for name, value in resolved_auth_headers.items(): + for existing in [k for k in effective_headers if k.lower() == name.lower()]: + del effective_headers[existing] + effective_headers[name] = value + return effective_headers diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py new file mode 100644 index 00000000000..e0927cc4f64 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -0,0 +1,214 @@ +"""Store for the enterprise IdP identity assertion captured at SSO login (EMA). + +The ``oauth2_id_jag`` egress arm needs the user's IdP ``id_token`` as its RFC 8693 +``subject_token``. A front-door client holds an identity-only ``llm_session_`` bearer, not an +IdP assertion, so the assertion captured at the one SSO login is the only usable subject +source for it. This module owns both sides of that state: the SSO callback persists here +(write-through to the DB so a login on one pod is visible to every pod) and the resolver +seam reads back by ``user_id``. Retention is gated on an ``oauth2_id_jag`` server actually +being registered, so a gateway with no EMA upstream never stores bearer material. + +The row is one encrypted payload per user, latest login wins. ``expires_at`` mirrors the +id_token ``exp`` claim and is judged by the reader, never enforced by deletion here: an +expired assertion with a refresh token is still renewable, and the DB row is the source of +truth, the same contract as the per-user OAuth credential store. +""" + +from __future__ import annotations + +import json +from datetime import datetime, timezone +from typing import TYPE_CHECKING + +import jwt +from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError + +from litellm._logging import verbose_proxy_logger + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient + +_ASSERTION_DECRYPT_LOG_KEY = "sso_identity_assertion" +_STR_ADAPTER: TypeAdapter[str] = TypeAdapter(str) +_MAYBE_STR_ADAPTER: TypeAdapter[str | None] = TypeAdapter(str | None) + + +class SSOIdentityAssertion(BaseModel): + """The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token, + ``expires_at`` bounds its usefulness, and the refresh token renews it without re-login.""" + + model_config = ConfigDict(frozen=True) + + id_token: SecretStr + refresh_token: SecretStr | None = None + issuer: str | None = None + expires_at: datetime | None = None + + +class _IdTokenClaims(BaseModel): + exp: float | None = None + iss: str | None = None + + +class _StoredAssertionPayload(BaseModel): + id_token: str + refresh_token: str | None = None + issuer: str | None = None + expires_at: datetime | None = None + + +def assertion_from_sso_login(id_token: object, refresh_token: object) -> SSOIdentityAssertion | None: + """The typed carrier built where the raw token response exists; ``None`` when the provider + sent no id_token or sent one that is not a decodable JWT, since neither is exchangeable + under EMA. Inputs are ``object`` because they come straight from the provider's untyped + token response; this is the one boundary that validates them. The token arrived over TLS + from the IdP's own token endpoint, so claims are read without signature verification, + matching how the SSO callback already decodes it for identity.""" + raw_id_token = id_token if isinstance(id_token, str) and id_token else None + if raw_id_token is None: + return None + raw_refresh_token = refresh_token if isinstance(refresh_token, str) and refresh_token else None + try: + claims = _IdTokenClaims.model_validate(jwt.decode(raw_id_token, options={"verify_signature": False})) + expires_at = datetime.fromtimestamp(claims.exp, tz=timezone.utc) if claims.exp is not None else None + except Exception: # noqa: BLE001 # decode failure = not retainable; never raise into login + verbose_proxy_logger.warning( + "SSO id_token could not be decoded or its claims were unusable; not retaining it for EMA egress." + ) + return None + return SSOIdentityAssertion( + id_token=SecretStr(raw_id_token), + refresh_token=SecretStr(raw_refresh_token) if raw_refresh_token else None, + issuer=claims.iss, + expires_at=expires_at, + ) + + +async def ema_assertion_retention_enabled() -> bool: + """Whether any MCP server uses ``oauth2_id_jag``, evaluated per login so the gateway only + retains bearer material while an EMA upstream exists to spend it on. Judged against the two + configuration authorities: the pod-local config declaration and the shared DB row. The + in-memory registry is deliberately not consulted in either direction; it is a per-process + snapshot of the DB state that can be stale both ways (a server added on another pod would + silently drop the write, one removed on another pod would keep retaining bearer material), + and a gate guarding a shared-DB write must judge against that storage's authority.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids import cycle + global_mcp_server_manager, + ) + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global + from litellm.types.mcp import MCPAuth # noqa: PLC0415 # runtime global + + config_servers = global_mcp_server_manager.config_mcp_servers.values() + if any(server.auth_type == MCPAuth.oauth2_id_jag for server in config_servers): + return True + if prisma_client is None: + return False + row = await prisma_client.db.litellm_mcpservertable.find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value}) + return row is not None + + +async def persist_sso_identity_assertion(user_id: str, assertion: SSOIdentityAssertion) -> None: + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper # noqa: PLC0415 # runtime global + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global + + if prisma_client is None: + return + payload: dict[str, str] = { + "id_token": assertion.id_token.get_secret_value(), + **({"refresh_token": assertion.refresh_token.get_secret_value()} if assertion.refresh_token else {}), + **({"issuer": assertion.issuer} if assertion.issuer else {}), + **({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}), + } + encoded = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload))) + await prisma_client.db.litellm_ssoidentityassertion.upsert( + where={"user_id": user_id}, + data={ + "create": {"user_id": user_id, "assertion_b64": encoded}, + "update": {"assertion_b64": encoded}, + }, + ) + + +async def fetch_sso_identity_assertion(user_id: str) -> SSOIdentityAssertion | None: + """The stored assertion for ``user_id``, or ``None`` when absent, undecryptable (salt-key + rotation), or unparseable. Expiry is not judged here; the reader owns that policy.""" + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper # noqa: PLC0415 # runtime global + from litellm.proxy.proxy_server import prisma_client # noqa: PLC0415 # runtime global + + if prisma_client is None: + return None + row = await prisma_client.db.litellm_ssoidentityassertion.find_unique(where={"user_id": user_id}) + if row is None: + return None + raw = _MAYBE_STR_ADAPTER.validate_python( + decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug") + ) + if raw is None: + return None + try: + payload = _StoredAssertionPayload.model_validate_json(raw) + except ValidationError: + verbose_proxy_logger.warning( + "Stored SSO identity assertion for user_id=%s could not be parsed; treating as absent.", user_id + ) + return None + return SSOIdentityAssertion( + id_token=SecretStr(payload.id_token), + refresh_token=SecretStr(payload.refresh_token) if payload.refresh_token else None, + issuer=payload.issuer, + expires_at=payload.expires_at, + ) + + +async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, new_master_key: str) -> None: + """Re-encrypt every stored assertion under ``new_master_key`` during a salt-key rotation, + mirroring the sibling per-user credential tables; an unreadable row is skipped so one + corrupt row does not abort the rotation. Rows are decrypted one at a time inside the loop + so the whole table's plaintext is never held in memory at once.""" + from prisma.models import LiteLLM_SSOIdentityAssertion as AssertionRow # noqa: PLC0415 # generated at runtime + + from litellm.proxy.common_utils.encrypt_decrypt_utils import ( # noqa: PLC0415 # runtime global + decrypt_value_helper, + encrypt_value_helper, + ) + + async def _rotate_row(row: AssertionRow) -> bool: + plaintext = _MAYBE_STR_ADAPTER.validate_python( + decrypt_value_helper(row.assertion_b64, _ASSERTION_DECRYPT_LOG_KEY, exception_type="debug") + ) + if plaintext is None: + verbose_proxy_logger.warning( + "rotate_sso_identity_assertions_master_key: could not decrypt assertion for user_id=%s, skipping", + row.user_id, + ) + return False + re_encrypted = _STR_ADAPTER.validate_python(encrypt_value_helper(plaintext, new_encryption_key=new_master_key)) + await prisma_client.db.litellm_ssoidentityassertion.update( + where={"user_id": row.user_id}, + data={"assertion_b64": re_encrypted}, + ) + return True + + rows = await prisma_client.db.litellm_ssoidentityassertion.find_many() + outcomes = [await _rotate_row(row) for row in rows] + verbose_proxy_logger.info( + "rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d", + sum(outcomes), + len(outcomes) - sum(outcomes), + ) + + +async def retain_sso_identity_assertion_for_ema(user_id: str, assertion: SSOIdentityAssertion | None) -> None: + """The SSO-callback hook: a no-op unless there is material AND an EMA server is registered. + A store failure is logged and swallowed because the login itself must not fail on an + egress-side write; the cost of a miss is a 401 challenge at the EMA upstream, not a lockout.""" + if assertion is None: + return + try: + if not await ema_assertion_retention_enabled(): + return + await persist_sso_identity_assertion(user_id, assertion) + except Exception as exc: # noqa: BLE001 # the login itself must not fail on an egress-side write + verbose_proxy_logger.warning( + "Failed to persist the SSO identity assertion for EMA egress (user_id=%s): %s", user_id, exc + ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index a8ab0937124..396dd6c7dc7 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -376,6 +376,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_auth_header, _request_extra_headers, + _request_resolved_auth_headers, ) from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport from litellm.proxy._experimental.mcp_server.tool_registry import ( @@ -2785,13 +2786,29 @@ if MCP_AVAILABLE: forwarded_headers = {} forwarded_headers[header_name] = value + resolved_auth_headers: dict[str, str] | None = None + if mcp_server: + ( + resolved_auth_headers, + forwarded_headers, + ) = await global_mcp_server_manager.resolve_openapi_upstream_auth( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=mcp_auth_header, + user_api_key_auth=user_api_key_auth, + forwarded_headers=forwarded_headers, + ) + _auth_token = _request_auth_header.set(auth_header_value) _extra_token = _request_extra_headers.set(forwarded_headers) + _resolved_token = _request_resolved_auth_headers.set(resolved_auth_headers) try: local_content = await _handle_local_mcp_tool(name, arguments) finally: _request_auth_header.reset(_auth_token) _request_extra_headers.reset(_extra_token) + _request_resolved_auth_headers.reset(_resolved_token) response = CallToolResult(content=cast(Any, local_content), isError=False) # Try managed MCP server tool (pass the full prefixed name) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8dd9dbc0afa..b73841c4793 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1856,6 +1856,17 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase): default_team_member_models: Optional[List[str]] = None # default allowed_models seeded onto new team members +class PatchTeamRequest(UpdateTeamRequest): + """ + Body of PATCH /team/{team_id}. + + Identical to UpdateTeamRequest except team_id is optional, because PATCH takes it + from the path. A team_id in the body is still accepted when it matches the path. + """ + + team_id: str | None = None + + class ResetTeamBudgetRequest(LiteLLMPydanticObjectBase): """ internal type used to reset the budget on a team @@ -2297,6 +2308,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="max response size in MB, if a response is larger than this size it will be rejected", ) + proxy_config_reload_interval_seconds: int = Field( + 30, + gt=0, + description="how often (in seconds) each pod reloads config-in-DB objects (models, credentials, guardrails, etc.) when store_model_in_db is enabled; lower values speed up multi-pod convergence at the cost of more DB load. Applied on proxy startup", + ) cancel_on_disconnect: Optional[bool] = Field( None, description="cancel the in-flight upstream LLM request (non-streaming) when the client disconnects, freeing backend capacity (e.g. a vLLM GPU slot); the request is logged as a 499 failure", @@ -2762,6 +2778,30 @@ class LiteLLM_OrganizationTableUpdate(LiteLLM_BudgetTable): return values +class OrganizationUpdateRequestV2(LiteLLMPydanticObjectBase): + """ + Typed PATCH body for ``/v2/organization/{organization_id}`` (RFC 7396 merge-patch). + + Presence is read from ``model_fields_set``, so a sent field is written and an omitted one is + left untouched. ``extra="forbid"`` makes an unknown key a 422 rather than a silent no-op, since + the contract hinges on which keys are present. See the endpoint for the per-field clear tokens. + """ + + model_config = ConfigDict(extra="forbid") + + organization_alias: str | None = None + models: list[str] | None = None + metadata: dict | None = None + tpm_limit: int | None = None + rpm_limit: int | None = None + max_budget: float | None = None + soft_budget: float | None = None + max_parallel_requests: int | None = None + model_max_budget: dict | None = None + budget_duration: str | None = None + object_permission: LiteLLM_ObjectPermissionBase | None = None + + from litellm.models.organization import ( # noqa: E402 LiteLLM_OrganizationTable as LiteLLM_OrganizationTable, ) @@ -4047,6 +4087,7 @@ class JWTAuthBuilderResult(TypedDict): token: str team_id: Optional[str] user_id: Optional[str] + user_email: str | None end_user_id: Optional[str] org_id: Optional[str] team_membership: Optional[LiteLLM_TeamMembership] diff --git a/litellm/proxy/a2a/agent_card.py b/litellm/proxy/a2a/agent_card.py index e97ab4a01ae..29a689a32de 100644 --- a/litellm/proxy/a2a/agent_card.py +++ b/litellm/proxy/a2a/agent_card.py @@ -7,23 +7,46 @@ the base; specific fields are replaced so all traffic flows through the proxy and uses LiteLLM auth. """ +import re from copy import deepcopy -from typing import Any, Dict, List, Mapping +from typing import Any, Dict, List, Literal, Mapping + +SupportedA2AVersion = Literal["0.3", "1.0"] # Protocol versions LiteLLM can serve to A2A clients. The admin pins one per agent; # responses are normalized to it regardless of the upstream agent's own version. -SUPPORTED_A2A_PROTOCOL_VERSIONS = ("0.3", "1.0") +SUPPORTED_A2A_PROTOCOL_VERSIONS: tuple[SupportedA2AVersion, ...] = ("0.3", "1.0") # Default served version when the agent card does not pin one. LITELLM_A2A_PROTOCOL_VERSION = "1.0" +_PROTOCOL_VERSION_PATTERN = re.compile( + r"^(\d+\.\d+)(?:\.\d+(?:-[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?(?:\+[0-9A-Za-z-]+(?:\.[0-9A-Za-z-]+)*)?)?$" +) + + +def normalize_protocol_version(version: object) -> SupportedA2AVersion | None: + """Map a raw ``protocolVersion`` value to the supported canonical major.minor version. + + Accepts the bare major.minor convention of the 1.0 spec (``"0.3"``, ``"1.0"``) and the + full semver forms older SDKs emit (``"0.3.0"``, ``"1.0.1"``, including prerelease and + build suffixes like ``"0.3.0-rc1"``). Malformed strings, versions outside the + supported set, and non-strings yield ``None``. + """ + if not isinstance(version, str): + return None + match = _PROTOCOL_VERSION_PATTERN.match(version) + if match is None: + return None + major_minor = match.group(1) + return next((supported for supported in SUPPORTED_A2A_PROTOCOL_VERSIONS if supported == major_minor), None) + + def resolve_served_protocol_version(card: Mapping[str, Any] | None) -> str: """Return the validated protocol version an agent card pins, else the default.""" - version = card.get("protocolVersion") if card else None - if version in SUPPORTED_A2A_PROTOCOL_VERSIONS: - return version - return LITELLM_A2A_PROTOCOL_VERSION + normalized = normalize_protocol_version(card.get("protocolVersion") if card else None) + return normalized if normalized is not None else LITELLM_A2A_PROTOCOL_VERSION # Security scheme exposed by the LiteLLM-fronted agent card. Always replaces diff --git a/litellm/proxy/a2a/version_convert.py b/litellm/proxy/a2a/version_convert.py index e8f49e6f6a9..9de33a0966a 100644 --- a/litellm/proxy/a2a/version_convert.py +++ b/litellm/proxy/a2a/version_convert.py @@ -30,6 +30,7 @@ from typing import Callable, Literal, Union from pydantic import BaseModel from litellm._logging import verbose_proxy_logger +from litellm.proxy.a2a.agent_card import normalize_protocol_version A2AVersion = Literal["0.3", "1.0"] RequestId = Union[str, int, None] @@ -103,16 +104,14 @@ def normalize_request_params(params: JsonDict, served: A2AVersion, *, method: st def _detect_card_version(card: JsonDict) -> A2AVersion: """Infer the wire version of an agent card dict. - ``protocolVersion`` is the authoritative indicator; fall back to presence of - ``supportedInterfaces`` (a 1.0-only field) only when the explicit field is absent. - Cards that set ``protocolVersion: "0.3"`` or carry neither signal are treated as 0.3. + ``protocolVersion`` is the authoritative indicator; semver values normalize to + their major.minor (``"0.3.0"`` -> ``"0.3"``). Fall back to presence of + ``supportedInterfaces`` (a 1.0-only field) only when the explicit field is + absent or unrecognized; cards carrying neither signal are treated as 0.3. """ - pv = card.get("protocolVersion") - if pv == "1.0": - return "1.0" - if pv == "0.3": - return "0.3" - # No protocolVersion field: use structural heuristic. + normalized = normalize_protocol_version(card.get("protocolVersion")) + if normalized is not None: + return normalized return "1.0" if "supportedInterfaces" in card else "0.3" diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index a7ceffed97b..2421f270974 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -23,6 +23,7 @@ from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKey from litellm.proxy.a2a.agent_card import ( SUPPORTED_A2A_PROTOCOL_VERSIONS, merge_agent_card, + normalize_protocol_version, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user @@ -51,7 +52,7 @@ def _proxy_base_url(http_request: Request) -> str: def _validate_protocol_version(upstream_card: Mapping[str, Any] | None) -> None: """Reject an agent card pinning an unsupported A2A protocol version.""" version = upstream_card.get("protocolVersion") if upstream_card else None - if version is not None and version not in SUPPORTED_A2A_PROTOCOL_VERSIONS: + if version is not None and normalize_protocol_version(version) is None: raise HTTPException( status_code=400, detail=( diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 293bb74e211..ecb37e67c14 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -3,7 +3,7 @@ import re import sys from functools import lru_cache from logging import Logger -from typing import Any, Dict, FrozenSet, List, Mapping, Optional, Tuple, Union +from typing import Any, Dict, FrozenSet, Iterator, List, Mapping, Optional, Tuple, Union from fastapi import HTTPException, Request, status @@ -12,7 +12,12 @@ from litellm import Router, provider_list from litellm._logging import verbose_proxy_logger from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HEADERS from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.litellm_core_utils.url_utils import SSRFError, validate_url +from litellm.litellm_core_utils.url_utils import ( + SSRFError, + is_url_destination_allowed_by_host, + provider_url_destination_candidates, + validate_url, +) from litellm.proxy._types import * from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_ENDPOINT_MARKER, @@ -290,6 +295,7 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = ( "use_ssl", # SDK-only field; also rejected outright in is_request_body_safe. "model_list", + "vertex_ai_credentials", # Observability credentials, hosts, and project identifiers: derived # from the canonical ``_supported_callback_params`` allowlist so new # integrations are covered automatically. Sorted for stable iteration @@ -342,6 +348,60 @@ def _check_banned_params( ) +_FALLBACK_FIELDS: tuple[str, ...] = ( + "fallbacks", + "context_window_fallbacks", + "content_policy_fallbacks", +) + + +def _iter_fallback_field_values(request_body: Mapping[str, object]) -> Iterator[object]: + override = request_body.get("router_settings_override") + for source in (request_body, override): + if isinstance(source, Mapping): + for field in _FALLBACK_FIELDS: + yield source.get(field) + + +def _iter_fallback_targets(value: object, depth: int) -> Iterator[str | Mapping[str, object]]: + if depth > 2 * litellm.ROUTER_MAX_FALLBACKS: + raise ValueError("Rejected Request: fallback nesting exceeds the allowed validation depth.") + if not isinstance(value, list): + return + for item in value: + if isinstance(item, str): + yield item + elif isinstance(item, Mapping): + values = tuple(item.values()) + if not (values and all(isinstance(v, list) for v in values)): + yield item + if isinstance(item.get("model"), str): + for field in _FALLBACK_FIELDS: + yield from _iter_fallback_targets(item.get(field), depth + 1) + else: + for target_list in values: + yield from _iter_fallback_targets(target_list, depth + 1) + + +def iter_request_fallback_targets(request_body: Mapping[str, object]) -> Iterator[str | Mapping[str, object]]: + for value in _iter_fallback_field_values(request_body): + yield from _iter_fallback_targets(value, 0) + + +def _reject_url_valued_fallback_target(value: str) -> None: + allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or [] + for candidate in provider_url_destination_candidates(value): + if not candidate.lower().startswith(("http://", "https://")): + continue + if is_url_destination_allowed_by_host(candidate, allowed_hosts): + continue + raise ValueError( + f"Rejected Request: URL-valued fallback destination '{value}' is not allowed. " + "Configure custom endpoints with api_base instead, or add the destination host to " + "`provider_url_destination_allowed_hosts` in litellm_settings." + ) + + def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Optional[Router], model: str) -> bool: """ Check if the request body is safe. @@ -379,6 +439,14 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: metadata = _coerce_metadata_to_dict(request_body.get(metadata_key)) if metadata is not None: _check_banned_params(metadata, general_settings, llm_router, model) + for target in iter_request_fallback_targets(request_body): + if isinstance(target, dict): + _check_banned_params(target, general_settings, llm_router, model) + target_model = target.get("model") + if isinstance(target_model, str): + _reject_url_valued_fallback_target(target_model) + elif isinstance(target, str): + _reject_url_valued_fallback_target(target) litellm_params = _coerce_metadata_to_dict(request_body.get("litellm_params")) if litellm_params is not None: litellm_params_metadata = _coerce_metadata_to_dict(litellm_params.get("metadata")) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index a44318c072c..ff87d0e70da 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1155,6 +1155,7 @@ class JWTAuthManager: org_id: Optional[str], api_key: str, jwt_valid_token: Optional[dict] = None, + user_email: str | None = None, ) -> Optional[JWTAuthBuilderResult]: """Check admin status and route access permissions""" if not jwt_handler.is_admin(scopes=scopes): @@ -1179,6 +1180,7 @@ class JWTAuthManager: token=api_key, team_id=None, user_id=user_id, + user_email=user_email, end_user_id=None, org_id=org_id, team_membership=None, @@ -2068,7 +2070,7 @@ class JWTAuthManager: # Check admin access admin_result = await JWTAuthManager.check_admin_access( - jwt_handler, scopes, route, user_id, org_id, api_key, jwt_valid_token + jwt_handler, scopes, route, user_id, org_id, api_key, jwt_valid_token, user_email=user_email ) if admin_result: await JWTAuthManager._attach_team_from_header_for_admin( @@ -2303,6 +2305,7 @@ class JWTAuthManager: team_id=team_id, team_object=team_object, user_id=user_id, + user_email=(user_object.user_email if user_object is not None and user_object.user_email else user_email), user_object=user_object, org_id=resolved_org_id, # Use resolved org_id (from alias lookup if applicable) org_object=org_object, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 5b21a7265a0..83a8a69511b 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -14,7 +14,7 @@ import secrets import orjson from datetime import datetime, timezone -from typing import Any, Dict, Iterator, NamedTuple, List, Optional, Protocol, Tuple, Union, cast +from typing import Any, Dict, NamedTuple, List, Optional, Protocol, Tuple, Union, cast import fastapi from fastapi import HTTPException, Request, WebSocket, status @@ -58,6 +58,7 @@ from litellm.proxy.auth.auth_utils import ( get_model_from_request, get_request_route, get_request_route_template, + iter_request_fallback_targets, normalize_request_route, pre_db_read_auth_checks, route_in_additonal_public_routes, @@ -1011,7 +1012,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None: return if getattr(request.state, "parent_otel_span", None) is not None: return - start_time = datetime.now() + start_time = datetime.now(timezone.utc) try: request.state.litellm_received_at = start_time except Exception: @@ -1061,7 +1062,7 @@ async def _user_api_key_auth_builder( # Prefer the receive-instant stamped by the early helper in # user_api_key_auth (before body parse) — overwriting it would shorten # the preprocessing-duration measurement by the body-parse window. - start_time = getattr(request.state, "litellm_received_at", None) or datetime.now() + start_time = getattr(request.state, "litellm_received_at", None) or datetime.now(timezone.utc) try: request.state.litellm_received_at = start_time except Exception: @@ -1255,6 +1256,7 @@ async def _user_api_key_auth_builder( team_id = result["team_id"] team_object = result["team_object"] user_id = result["user_id"] + user_email = result["user_email"] user_object = result["user_object"] end_user_id = result["end_user_id"] org_id = result["org_id"] @@ -1279,6 +1281,7 @@ async def _user_api_key_auth_builder( api_key=None, user_role=LitellmUserRoles.PROXY_ADMIN, user_id=user_id, + user_email=user_email, team_id=team_id, team_alias=(team_object.team_alias if team_object is not None else None), team_tpm_limit=(team_object.tpm_limit if team_object is not None else None), @@ -1304,6 +1307,7 @@ async def _user_api_key_auth_builder( else LitellmUserRoles.INTERNAL_USER ), user_id=user_id, + user_email=user_email, org_id=org_id, parent_otel_span=parent_otel_span, end_user_id=end_user_id, @@ -1345,6 +1349,7 @@ async def _user_api_key_auth_builder( ) if auto_registered is not None: auto_registered.jwt_claims = jwt_claims + auto_registered.user_email = user_email valid_token = auto_registered api_key = valid_token.token or "" @@ -2607,7 +2612,7 @@ async def _return_user_api_key_auth_obj( start_time: datetime, user_role: Optional[LitellmUserRoles] = None, ) -> UserAPIKeyAuth: - end_time = datetime.now() + end_time = datetime.now(timezone.utc) asyncio.create_task( user_api_key_service_logger_obj.async_service_success_hook( @@ -2696,9 +2701,10 @@ def _update_key_budget_with_temp_budget_increase( ) -> UserAPIKeyAuth: if valid_token.max_budget is None: return valid_token - temp_budget_increase = _get_temp_budget_increase(valid_token) or 0.0 - valid_token.max_budget = valid_token.max_budget + temp_budget_increase - return valid_token + temp_budget_increase = _get_temp_budget_increase(valid_token) + if not temp_budget_increase: + return valid_token + return valid_token.model_copy(update={"max_budget": valid_token.max_budget + temp_budget_increase}) async def _lookup_end_user_and_apply_budget( @@ -2796,19 +2802,11 @@ async def _enforce_key_and_fallback_model_access( llm_router=llm_router, ) - # Validate every fallback model name reachable by this request. - # All three fields (``fallbacks``, ``context_window_fallbacks``, - # ``content_policy_fallbacks``) are forwarded to the router as - # per-request kwargs whether they appear at the top level of - # ``request_data`` or nested under ``router_settings_override``. - # Both surfaces must be validated against the API key's model - # allowlist or a caller can smuggle a restricted model. VERIA-44. - fallback_names: List[str] = [] - override_settings = request_data.get("router_settings_override") - for _fb_key in ROUTER_FALLBACK_FIELDS: - fallback_names.extend(iter_router_fallback_model_names(request_data.get(_fb_key))) - if isinstance(override_settings, dict): - fallback_names.extend(iter_router_fallback_model_names(override_settings.get(_fb_key))) + fallback_names = tuple( + name + for target in iter_request_fallback_targets(request_data) + if (name := _fallback_target_model_name(target)) is not None + ) for _name in dict.fromkeys(fallback_names): # dedupe, preserve order await can_key_call_model( @@ -2824,36 +2822,14 @@ async def _enforce_key_and_fallback_model_access( ) -ROUTER_FALLBACK_FIELDS: Tuple[str, ...] = ( - "fallbacks", - "context_window_fallbacks", - "content_policy_fallbacks", -) - - -def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]: - """Yield leaf model names from any of the supported fallbacks shapes. - - Handles the simple top-level shape (``str`` or ``{"model": str}``) and - the nested router-config shape (``[{primary: [fallback_list]}]``). - """ - if not isinstance(fallbacks, list): - return - for entry in fallbacks: - if isinstance(entry, str): - yield entry - elif isinstance(entry, dict): - if isinstance(entry.get("model"), str): - yield entry["model"] - continue - for fallback_list in entry.values(): - if not isinstance(fallback_list, list): - continue - for m in fallback_list: - if isinstance(m, str): - yield m - elif isinstance(m, dict) and isinstance(m.get("model"), str): - yield m["model"] +def _fallback_target_model_name(target: object) -> str | None: + if isinstance(target, str): + return target + if isinstance(target, dict): + model = target.get("model") + if isinstance(model, str): + return model + return None async def _run_post_custom_auth_checks( diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index 84bc27ef0d4..2ad8a08b8c3 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -523,7 +523,7 @@ export LITELLM_PROXY_API_KEY=sk-... lite model-groups list [--format table|json] ``` -Lists the model groups your key can reach on the proxy, via `/model_group/info`, along with each group's mode (`chat`, `embedding`, etc.) and per-token pricing. This is also what `lite autoroute configure` uses internally to discover what it can offer you. +Lists the model groups your key can reach on the proxy, via `/model_group/info`, along with each group's mode (`chat`, `embedding`, etc.) and per-token pricing. Note this route needs management access; `lite autoroute configure` instead discovers models through `/v1/models`, so it works with a key scoped to just the AI API routes #### Configure the Auto-Router diff --git a/litellm/proxy/client/cli/commands/autoroute/config.py b/litellm/proxy/client/cli/commands/autoroute/config.py index 603cea38f6f..237705564ff 100644 --- a/litellm/proxy/client/cli/commands/autoroute/config.py +++ b/litellm/proxy/client/cli/commands/autoroute/config.py @@ -15,41 +15,25 @@ class DiscoveredModel(BaseModel): name: str mode: str = "chat" - input_cost_per_token: float | None = None - output_cost_per_token: float | None = None -class _RawModelGroup(BaseModel): +class _RawModelListing(BaseModel): model_config = ConfigDict(extra="ignore") - model_group: str - # Optional: some real deployments return an explicit `"mode": null` for models that - # were registered without a mode (seen for embedding models like voyage-4-large). - # ModelGroupInfo's own "chat" default (litellm/types/router.py) only applies when the - # key is missing entirely, not when it's present as null, so this must tolerate None. - mode: str | None = "chat" - input_cost_per_token: float | None = None - output_cost_per_token: float | None = None + id: str + # /v1/models attaches "mode" (sourced from the cost map) only for models it can resolve; + # a model whose mode is unknown arrives without the field, so default it to chat rather + # than dropping it, which keeps it selectable as a routing target in the wizard. + mode: str = "chat" -_RAW_MODEL_GROUPS_ADAPTER = TypeAdapter(list[_RawModelGroup]) +_RAW_MODEL_LISTING_ADAPTER = TypeAdapter(list[_RawModelListing]) def parse_discovered_models(raw: list[JsonValue]) -> tuple[DiscoveredModel, ...]: - """Validate a raw `/model_group/info` response into typed models.""" - parsed = _RAW_MODEL_GROUPS_ADAPTER.validate_python(raw) - return tuple( - DiscoveredModel( - name=group.model_group, - # A null mode means the server genuinely doesn't know what this model does; - # "unknown" (rather than guessing "chat") keeps it out of both chat_models() - # and embedding_models() instead of risking a wrong-mode deployment. - mode=group.mode or "unknown", - input_cost_per_token=group.input_cost_per_token, - output_cost_per_token=group.output_cost_per_token, - ) - for group in parsed - ) + """Validate a raw `/v1/models` response into typed models.""" + parsed = _RAW_MODEL_LISTING_ADAPTER.validate_python(raw) + return tuple(DiscoveredModel(name=item.id, mode=item.mode) for item in parsed) def chat_models(models: tuple[DiscoveredModel, ...]) -> tuple[DiscoveredModel, ...]: diff --git a/litellm/proxy/client/cli/commands/autoroute/wizard.py b/litellm/proxy/client/cli/commands/autoroute/wizard.py index 8ad87315fb9..d3fe458d9f0 100644 --- a/litellm/proxy/client/cli/commands/autoroute/wizard.py +++ b/litellm/proxy/client/cli/commands/autoroute/wizard.py @@ -111,12 +111,12 @@ def run_configure_wizard(ctx: click.Context) -> Path: api_key = ctx.obj["api_key"] client = Client(base_url=base_url, api_key=api_key) - raw_groups = client.model_groups.info() - if not isinstance(raw_groups, list): + raw_models = client.models.list() + if not isinstance(raw_models, list): raise click.ClickException( - f"Unexpected response from /model_group/info: expected a list, got {type(raw_groups).__name__}" + f"Unexpected response from /v1/models: expected a list, got {type(raw_models).__name__}" ) - discovered = parse_discovered_models(raw_groups) + discovered = parse_discovered_models(raw_models) chat_pool = chat_models(discovered) embedding_pool = embedding_models(discovered) diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index c644ecc3dae..a9c2a12aff7 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -1,4 +1,5 @@ import copy +import os from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Literal, Optional import litellm @@ -564,11 +565,8 @@ def process_callback(_callback: str, callback_type: str, environment_variables: env_vars_dict: dict[str, str | None] = {} for _var in env_vars: - env_variable = environment_variables.get(_var, None) - if env_variable is None: - env_vars_dict[_var] = None - else: - env_vars_dict[_var] = env_variable + stored_value = environment_variables.get(_var, None) + env_vars_dict[_var] = stored_value if stored_value is not None else os.getenv(_var) return {"name": _callback, "variables": env_vars_dict, "type": callback_type} diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index e758420ee37..23a5b8f9c53 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -13,6 +13,11 @@ from litellm.proxy._types import ( LiteLLM_UserTable, LiteLLM_VerificationToken, ) +from litellm.proxy.common_utils.timezone_utils import ( + BudgetResetSettings, + compute_budget_reset_at, + get_budget_reset_settings, +) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.table_repositories import ( @@ -32,9 +37,15 @@ class ResetBudgetJob: Resets the budget for all the keys, users, and teams that need it """ - def __init__(self, proxy_logging_obj: ProxyLogging, prisma_client: PrismaClient): + def __init__( + self, + proxy_logging_obj: ProxyLogging, + prisma_client: PrismaClient, + reset_settings: BudgetResetSettings | None = None, + ): self.proxy_logging_obj: ProxyLogging = proxy_logging_obj self.prisma_client: PrismaClient = prisma_client + self.reset_settings: BudgetResetSettings = reset_settings or get_budget_reset_settings() async def reset_budget( self, @@ -237,7 +248,7 @@ class ResetBudgetJob: if budgets_to_reset is not None and len(budgets_to_reset) > 0: for budget in budgets_to_reset: - budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now) + budget = await ResetBudgetJob._reset_budget_reset_at_date(budget, now, self.reset_settings) await self.prisma_client.update_data( query_type="update_many", @@ -442,7 +453,11 @@ class ResetBudgetJob: if keys_to_reset is not None and len(keys_to_reset) > 0: for key in keys_to_reset: try: - updated_key = await ResetBudgetJob._reset_budget_for_key(key=key, current_time=now) + updated_key = await ResetBudgetJob._reset_budget_for_key( + key=key, + current_time=now, + reset_settings=self.reset_settings, + ) if updated_key is not None: updated_keys.append(updated_key) else: @@ -513,7 +528,11 @@ class ResetBudgetJob: if users_to_reset is not None and len(users_to_reset) > 0: for user in users_to_reset: try: - updated_user = await ResetBudgetJob._reset_budget_for_user(user=user, current_time=now) + updated_user = await ResetBudgetJob._reset_budget_for_user( + user=user, + current_time=now, + reset_settings=self.reset_settings, + ) if updated_user is not None: updated_users.append(updated_user) else: @@ -588,7 +607,11 @@ class ResetBudgetJob: if teams_to_reset is not None and len(teams_to_reset) > 0: for team in teams_to_reset: try: - updated_team = await ResetBudgetJob._reset_budget_for_team(team=team, current_time=now) + updated_team = await ResetBudgetJob._reset_budget_for_team( + team=team, + current_time=now, + reset_settings=self.reset_settings, + ) if updated_team is not None: updated_teams.append(updated_team) else: @@ -655,10 +678,9 @@ class ResetBudgetJob: counter_key: str, spend_counter_cache: Any, now: datetime, + reset_settings: BudgetResetSettings, ) -> bool: """Reset a single budget window if expired. Returns True if the window was reset.""" - from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time - reset_at_str = window.get("reset_at") if not reset_at_str: return False @@ -671,7 +693,9 @@ class ResetBudgetJob: await spend_counter_cache.redis_cache.async_set_cache(key=counter_key, value=0.0) except Exception as redis_err: verbose_proxy_logger.warning("Failed to reset Redis counter %s: %s", counter_key, redis_err) - window["reset_at"] = get_budget_reset_time(budget_duration=window["budget_duration"]).isoformat() + window["reset_at"] = compute_budget_reset_at( + budget_duration=window["budget_duration"], settings=reset_settings + ).isoformat() return True async def reset_budget_windows(self) -> None: @@ -703,7 +727,13 @@ class ResetBudgetJob: changed = False for window in windows: counter_key = f"spend:key:{row['token']}:window:{window['budget_duration']}" - if await ResetBudgetJob._reset_expired_window(window, counter_key, spend_counter_cache, now): + if await ResetBudgetJob._reset_expired_window( + window, + counter_key, + spend_counter_cache, + now, + self.reset_settings, + ): changed = True if changed: await VerificationTokenRepository(self.prisma_client).table.update( @@ -726,7 +756,13 @@ class ResetBudgetJob: changed = False for window in windows: counter_key = f"spend:team:{row['team_id']}:window:{window['budget_duration']}" - if await ResetBudgetJob._reset_expired_window(window, counter_key, spend_counter_cache, now): + if await ResetBudgetJob._reset_expired_window( + window, + counter_key, + spend_counter_cache, + now, + self.reset_settings, + ): changed = True if changed: await TeamRepository(self.prisma_client).table.update( @@ -741,6 +777,7 @@ class ResetBudgetJob: item: Union[LiteLLM_TeamTable, LiteLLM_UserTable, LiteLLM_VerificationToken], current_time: datetime, item_type: Literal["key", "team", "user"], + reset_settings: BudgetResetSettings, ): """ In-place, updates spend=0, and sets budget_reset_at to current_time + budget_duration @@ -755,24 +792,40 @@ class ResetBudgetJob: try: item.spend = 0.0 if hasattr(item, "budget_duration") and item.budget_duration is not None: - from litellm.proxy.common_utils.timezone_utils import ( - get_budget_reset_time, + item.budget_reset_at = compute_budget_reset_at( + budget_duration=item.budget_duration, settings=reset_settings ) - - item.budget_reset_at = get_budget_reset_time(budget_duration=item.budget_duration) return item except Exception as e: verbose_proxy_logger.exception("Error resetting budget for %s: %s. Item: %s", item_type, e, item) raise e @staticmethod - async def _reset_budget_for_team(team: LiteLLM_TeamTable, current_time: datetime) -> Optional[LiteLLM_TeamTable]: - await ResetBudgetJob._reset_budget_common(item=team, current_time=current_time, item_type="team") + async def _reset_budget_for_team( + team: LiteLLM_TeamTable, + current_time: datetime, + reset_settings: BudgetResetSettings, + ) -> LiteLLM_TeamTable | None: + await ResetBudgetJob._reset_budget_common( + item=team, + current_time=current_time, + item_type="team", + reset_settings=reset_settings, + ) return team @staticmethod - async def _reset_budget_for_user(user: LiteLLM_UserTable, current_time: datetime) -> Optional[LiteLLM_UserTable]: - await ResetBudgetJob._reset_budget_common(item=user, current_time=current_time, item_type="user") + async def _reset_budget_for_user( + user: LiteLLM_UserTable, + current_time: datetime, + reset_settings: BudgetResetSettings, + ) -> LiteLLM_UserTable | None: + await ResetBudgetJob._reset_budget_common( + item=user, + current_time=current_time, + item_type="user", + reset_settings=reset_settings, + ) return user @staticmethod @@ -788,15 +841,15 @@ class ResetBudgetJob: @staticmethod async def _reset_budget_reset_at_date( - budget: LiteLLM_BudgetTableFull, current_time: datetime + budget: LiteLLM_BudgetTableFull, + current_time: datetime, + reset_settings: BudgetResetSettings, ) -> LiteLLM_BudgetTableFull: try: if budget.budget_duration is not None: - from litellm.proxy.common_utils.timezone_utils import ( - get_budget_reset_time, + budget.budget_reset_at = compute_budget_reset_at( + budget_duration=budget.budget_duration, settings=reset_settings ) - - budget.budget_reset_at = get_budget_reset_time(budget_duration=budget.budget_duration) except Exception as e: verbose_proxy_logger.exception("Error resetting budget_reset_at for budget: %s. Item: %s", e, budget) raise e @@ -804,7 +857,14 @@ class ResetBudgetJob: @staticmethod async def _reset_budget_for_key( - key: LiteLLM_VerificationToken, current_time: datetime - ) -> Optional[LiteLLM_VerificationToken]: - await ResetBudgetJob._reset_budget_common(item=key, current_time=current_time, item_type="key") + key: LiteLLM_VerificationToken, + current_time: datetime, + reset_settings: BudgetResetSettings, + ) -> LiteLLM_VerificationToken | None: + await ResetBudgetJob._reset_budget_common( + item=key, + current_time=current_time, + item_type="key", + reset_settings=reset_settings, + ) return key diff --git a/litellm/proxy/common_utils/timezone_utils.py b/litellm/proxy/common_utils/timezone_utils.py index 32f9f47d519..a50daf40144 100644 --- a/litellm/proxy/common_utils/timezone_utils.py +++ b/litellm/proxy/common_utils/timezone_utils.py @@ -1,10 +1,47 @@ -from datetime import datetime, timezone +from datetime import datetime, time, timezone + +from pydantic import BaseModel, ConfigDict import litellm from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time -def get_budget_reset_timezone(): +class BudgetResetSettings(BaseModel): + """Immutable, validated settings that govern when budgets reset. + + Parsed once from `litellm_settings` and injected into consumers (the reset + job, management endpoints) so reset times never depend on reaching into + module-level globals at call time. + """ + + model_config = ConfigDict(frozen=True) + + timezone: str = "UTC" + reset_time_of_day: time = time(0, 0) + + +def parse_budget_reset_time(raw: object) -> time: + """Parse a `budget_reset_time` config value (e.g. "12:00") into a `time`. + + Falls back to midnight when unset; raises a clear error on a malformed value + so a bad config fails loudly at startup instead of silently resetting at midnight. + """ + if raw is None or raw == "": + return time(0, 0) + if not isinstance(raw, str): + raise ValueError(f"Invalid budget_reset_time {raw!r}; must be a quoted 24-hour 'HH:MM' string, e.g. \"12:00\"") + for fmt in ("%H:%M", "%H:%M:%S"): + try: + parsed = datetime.strptime(raw, fmt) + return time(hour=parsed.hour, minute=parsed.minute, second=parsed.second) + except ValueError: + continue + raise ValueError( + f"Invalid budget_reset_time {raw!r}; expected a 24-hour 'HH:MM' or 'HH:MM:SS' string, e.g. \"12:00\"" + ) + + +def get_budget_reset_timezone() -> str: """ Get the budget reset timezone from litellm_settings. Falls back to UTC if not specified. @@ -15,15 +52,29 @@ def get_budget_reset_timezone(): return getattr(litellm, "timezone", None) or "UTC" -def get_budget_reset_time(budget_duration: str) -> datetime: - """ - Get the budget reset time based on the configured timezone. - Falls back to UTC if not specified. - """ +def get_budget_reset_settings() -> BudgetResetSettings: + """Build validated reset settings from litellm_settings. Raises on a malformed + `budget_reset_time`, which lets the proxy fail fast at startup.""" + return BudgetResetSettings( + timezone=get_budget_reset_timezone(), + reset_time_of_day=parse_budget_reset_time(getattr(litellm, "budget_reset_time", None)), + ) - reset_at = get_next_standardized_reset_time( + +def compute_budget_reset_at(budget_duration: str, settings: BudgetResetSettings) -> datetime: + """Compute the next reset time for a budget duration using injected settings.""" + return get_next_standardized_reset_time( duration=budget_duration, current_time=datetime.now(timezone.utc), - timezone_str=get_budget_reset_timezone(), + timezone_str=settings.timezone, + reset_time_of_day=settings.reset_time_of_day, ) - return reset_at + + +def get_budget_reset_time(budget_duration: str) -> datetime: + """Get the budget reset time using the globally-configured timezone and reset time. + + Thin wrapper over `compute_budget_reset_at` for callers that don't yet receive + `BudgetResetSettings` by injection (creation/update endpoints, startup backfill). + """ + return compute_budget_reset_at(budget_duration, get_budget_reset_settings()) diff --git a/litellm/proxy/config_resolvers/__init__.py b/litellm/proxy/config_resolvers/__init__.py new file mode 100644 index 00000000000..88b4c3961f0 --- /dev/null +++ b/litellm/proxy/config_resolvers/__init__.py @@ -0,0 +1,9 @@ +"""Typed, provenance-aware resolution of proxy settings from DB then env.""" + +from litellm.proxy.config_resolvers._descriptors import ( + FieldDescriptor, + FieldSource, + resolve_fields, +) + +__all__ = ["FieldDescriptor", "FieldSource", "resolve_fields"] diff --git a/litellm/proxy/config_resolvers/_descriptors.py b/litellm/proxy/config_resolvers/_descriptors.py new file mode 100644 index 00000000000..f67a690f92f --- /dev/null +++ b/litellm/proxy/config_resolvers/_descriptors.py @@ -0,0 +1,73 @@ +"""Shared primitive for resolving a settings value from its sources. + +A ``FieldDescriptor`` names, for one setting, where it lives in the stored DB +row (``db_key``), which process env var carries it (``env_var``), whether it is +a secret, and its effective default. ``resolve_fields`` reconciles a set of +descriptors against a decrypted DB row and the process environment with a fixed +precedence, returning the resolved values plus per-field provenance so a caller +can tell whether a value came from the database, the environment, a default, or +is unset. +""" + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Literal + +FieldSource = Literal["db", "env", "default", "unset"] + + +@dataclass(frozen=True, slots=True) +class FieldDescriptor: + field_name: str + db_key: str + env_var: str + is_secret: bool = False + default: str | None = None + + +def _db_is_set(db_value: object, empty_db_is_set: bool) -> bool: + if empty_db_is_set: + # A stored key that is present, even as "", is an explicit admin choice + # (e.g. clearing an alerting webhook) and must win over a stale env var. + return db_value is not None + # A blank stored value is treated as absent, so it falls through to env. This + # fits settings whose clear path also unsets the env var (e.g. SSO). + return isinstance(db_value, str) and bool(db_value.strip()) + + +def _resolve_one( + descriptor: FieldDescriptor, + db_values: Mapping[str, object], + env: Mapping[str, str], + empty_db_is_set: bool, +) -> tuple[str, str | None, FieldSource]: + db_value = db_values.get(descriptor.db_key) + if _db_is_set(db_value, empty_db_is_set): + return descriptor.field_name, db_value if isinstance(db_value, str) else str(db_value), "db" + env_value = env.get(descriptor.env_var) + if isinstance(env_value, str) and env_value.strip(): + return descriptor.field_name, env_value, "env" + if descriptor.default is not None: + return descriptor.field_name, descriptor.default, "default" + return descriptor.field_name, None, "unset" + + +def resolve_fields( + descriptors: Sequence[FieldDescriptor], + db_values: Mapping[str, object], + env: Mapping[str, str], + empty_db_is_set: bool = False, +) -> tuple[dict[str, str | None], dict[str, FieldSource]]: + """Resolve every descriptor to (values, provenance). + + Precedence per field: a set stored value wins, else a non-blank process env + var, else the descriptor default, else unset. ``empty_db_is_set`` selects + how a present-but-empty stored value is read: ``False`` treats it as absent + so it falls back to env (SSO, whose clear path also unsets the env var); + ``True`` treats it as an explicit clear that wins over env (alerting, whose + clear path stores "" without unsetting the env var). + """ + resolved = tuple(_resolve_one(descriptor, db_values, env, empty_db_is_set) for descriptor in descriptors) + values = {field_name: value for field_name, value, _ in resolved} + provenance = {field_name: source for field_name, _, source in resolved} + return values, provenance diff --git a/litellm/proxy/config_resolvers/alerting.py b/litellm/proxy/config_resolvers/alerting.py new file mode 100644 index 00000000000..3704ec09355 --- /dev/null +++ b/litellm/proxy/config_resolvers/alerting.py @@ -0,0 +1,25 @@ +"""Descriptor tables for the alerting settings surfaced by /get/config/callbacks. + +These reconcile the stored ``environment_variables`` blob (keyed by the +uppercase env-var names) with the process environment. SMTP_PORT and SMTP_TLS +carry the same effective defaults the mail-send path applies, so the settings +page shows the config that mail would actually use rather than a blank. +""" + +from litellm.proxy.config_resolvers._descriptors import FieldDescriptor + +EMAIL_DESCRIPTORS: tuple[FieldDescriptor, ...] = ( + FieldDescriptor("SMTP_HOST", "SMTP_HOST", "SMTP_HOST"), + FieldDescriptor("SMTP_PORT", "SMTP_PORT", "SMTP_PORT", default="587"), + FieldDescriptor("SMTP_TLS", "SMTP_TLS", "SMTP_TLS", default="True"), + FieldDescriptor("SMTP_USERNAME", "SMTP_USERNAME", "SMTP_USERNAME", is_secret=True), + FieldDescriptor("SMTP_PASSWORD", "SMTP_PASSWORD", "SMTP_PASSWORD", is_secret=True), + FieldDescriptor("SMTP_SENDER_EMAIL", "SMTP_SENDER_EMAIL", "SMTP_SENDER_EMAIL"), + FieldDescriptor("TEST_EMAIL_ADDRESS", "TEST_EMAIL_ADDRESS", "TEST_EMAIL_ADDRESS"), + FieldDescriptor("EMAIL_LOGO_URL", "EMAIL_LOGO_URL", "EMAIL_LOGO_URL"), + FieldDescriptor("EMAIL_SUPPORT_CONTACT", "EMAIL_SUPPORT_CONTACT", "EMAIL_SUPPORT_CONTACT"), +) + +SLACK_DESCRIPTORS: tuple[FieldDescriptor, ...] = ( + FieldDescriptor("SLACK_WEBHOOK_URL", "SLACK_WEBHOOK_URL", "SLACK_WEBHOOK_URL", is_secret=True), +) diff --git a/litellm/proxy/config_resolvers/sso.py b/litellm/proxy/config_resolvers/sso.py new file mode 100644 index 00000000000..3d83c06dd62 --- /dev/null +++ b/litellm/proxy/config_resolvers/sso.py @@ -0,0 +1,94 @@ +"""Resolved SSO config object. + +Reconciles the dedicated ``sso_config`` DB row (lowercase, per-value encrypted +keys) with the process environment (uppercase env vars) into a typed +``SSOConfig`` plus per-field provenance. This is the single source of truth for +the SSO field -> env-var mapping, used by both the read-back endpoint and the +save endpoint so the two can never drift. +""" + +from collections.abc import Mapping +from dataclasses import dataclass + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper +from litellm.proxy.config_resolvers._descriptors import ( + FieldDescriptor, + FieldSource, + resolve_fields, +) +from litellm.types.proxy.management_endpoints.ui_sso import ( + RoleMappings, + SSOConfig, + TeamMappings, +) + +SSO_DESCRIPTORS: tuple[FieldDescriptor, ...] = ( + FieldDescriptor("google_client_id", "google_client_id", "GOOGLE_CLIENT_ID"), + FieldDescriptor("google_client_secret", "google_client_secret", "GOOGLE_CLIENT_SECRET", is_secret=True), + FieldDescriptor("microsoft_client_id", "microsoft_client_id", "MICROSOFT_CLIENT_ID"), + FieldDescriptor("microsoft_client_secret", "microsoft_client_secret", "MICROSOFT_CLIENT_SECRET", is_secret=True), + FieldDescriptor("microsoft_tenant", "microsoft_tenant", "MICROSOFT_TENANT"), + FieldDescriptor("generic_client_id", "generic_client_id", "GENERIC_CLIENT_ID"), + FieldDescriptor("generic_client_secret", "generic_client_secret", "GENERIC_CLIENT_SECRET", is_secret=True), + FieldDescriptor( + "generic_authorization_endpoint", "generic_authorization_endpoint", "GENERIC_AUTHORIZATION_ENDPOINT" + ), + FieldDescriptor("generic_token_endpoint", "generic_token_endpoint", "GENERIC_TOKEN_ENDPOINT"), + FieldDescriptor("generic_userinfo_endpoint", "generic_userinfo_endpoint", "GENERIC_USERINFO_ENDPOINT"), + FieldDescriptor("generic_scope", "generic_scope", "GENERIC_SCOPE", default="openid email profile"), + FieldDescriptor("proxy_base_url", "proxy_base_url", "PROXY_BASE_URL"), +) + +# Derived from the descriptor table so read (masking) and the field->env mapping +# never diverge from the resolver. +SSO_SECRET_FIELDS: frozenset[str] = frozenset(d.field_name for d in SSO_DESCRIPTORS if d.is_secret) +SSO_FIELD_ENV_VARS: dict[str, str] = {d.field_name: d.env_var for d in SSO_DESCRIPTORS} + +# Structured sub-objects stored on the SSO row that are not simple env-backed +# scalars; handled outside the descriptor resolution. +_STRUCTURED_KEYS = ("role_mappings", "team_mappings") + + +@dataclass(frozen=True, slots=True) +class ResolvedSSOConfig: + config: SSOConfig + provenance: dict[str, FieldSource] + + +def _decrypt(raw: Mapping[str, object]) -> dict[str, object]: + return { + key: ( + decrypt_value_helper(value=value, key=key, return_original_value=True) if isinstance(value, str) else value + ) + for key, value in raw.items() + } + + +def _parse_role_mappings(data: object) -> RoleMappings | None: + # The stored row is JSON, so mappings arrive as a dict (or are absent). + return RoleMappings(**data) if isinstance(data, dict) else None + + +def _parse_team_mappings(data: object) -> TeamMappings | None: + return TeamMappings(**data) if isinstance(data, dict) else None + + +def resolve_sso_config(sso_db_settings: Mapping[str, object] | None, env: Mapping[str, str]) -> ResolvedSSOConfig: + """Resolve the effective SSO config: stored row first, then process env. + + Decryption happens here, once, via the pure ``decrypt_value_helper``; this + function never writes ``os.environ`` (unlike the legacy read path). Values + are returned unmasked so the login path could consume them; the read-back + endpoint is responsible for masking secrets before responding to the UI. + """ + raw = dict(sso_db_settings) if sso_db_settings else {} + decrypted = _decrypt({key: value for key, value in raw.items() if key not in _STRUCTURED_KEYS}) + values, provenance = resolve_fields(SSO_DESCRIPTORS, decrypted, env) + structured = { + "user_email": decrypted.get("user_email"), + "ui_access_mode": decrypted.get("ui_access_mode"), + "role_mappings": _parse_role_mappings(raw.get("role_mappings")), + "team_mappings": _parse_team_mappings(raw.get("team_mappings")), + } + config = SSOConfig(**{**values, **structured}) + return ResolvedSSOConfig(config=config, provenance=provenance) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 54156715da8..cec682d772a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -2046,6 +2046,90 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): masking_index += 1 verbose_proxy_logger.debug("Applied masking to choice text content") + @staticmethod + def _incremental_scan_cache() -> DualCache: + """Resolve the cache used to remember which segments a session already scanned. + + Prefers the proxy's shared cache (``internal_usage_cache.dual_cache``), which is + backed by Redis when the deployment configures it, so incremental state is shared + across proxy instances. Falls back to a process-local ``DualCache`` singleton when + the proxy is not running (e.g. unit tests), where sharing does not apply. + """ + from litellm.integrations.custom_guardrail import dc as fallback_cache + + try: + from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging + except Exception: # noqa: BLE001 # proxy not importable outside the server; use local fallback + return fallback_cache + if _proxy_logging is not None: + return _proxy_logging.internal_usage_cache.dual_cache + return fallback_cache + + def _bedrock_response_has_masked_output(self, response: BedrockGuardrailResponse) -> bool: + """Return True if the guardrail rewrote (masked/anonymized) any scanned text. + + Bedrock returns non-empty ``output``/``outputs`` text only when it changed the + content; an ``action == "NONE"`` response leaves both empty. + """ + for field in ("output", "outputs"): + items = response.get(field) or [] + if any(isinstance(item, dict) and item.get("text") for item in items): + return True + return False + + async def _apply_incremental_request_scan( + self, + texts: list[str], + inputs: "GenericGuardrailAPIInputs", + request_data: dict, + ) -> Optional["GenericGuardrailAPIInputs"]: + """Scan only the text segments not already seen earlier in this session. + + Returns ``None`` when incremental scanning is inactive (feature off, no + session id, masking enabled, or cache unavailable) or when the guardrail + turns out to mask content, telling the caller to run the normal full scan. + Otherwise scans only the new segments and skips the Bedrock call entirely + when nothing is new. Incremental mode is for blocking/detection guardrails + only: if the guardrail returns masked output it cannot be applied to the + skipped context, so the scan falls back to the full path and no session + state is recorded. + """ + cache = self._incremental_scan_cache() + + new_texts = await self.filter_new_texts_for_session( + texts=texts, + request_data=request_data, + cache=cache, + ) + if new_texts is None: + return None + + if not new_texts: + verbose_proxy_logger.debug("Bedrock Guardrail: no new messages to scan for this session, skipping API call") + return inputs + + bedrock_response = await self.make_bedrock_api_request( + source="INPUT", + messages=[ChatCompletionUserMessage(role="user", content=text) for text in new_texts], + request_data=request_data, + logging_event_type=GuardrailEventHooks.pre_call, + ) + + if self._bedrock_response_has_masked_output(bedrock_response): + verbose_proxy_logger.warning( + "Bedrock Guardrail %s: guardrail returned masked/anonymized content; " + "only_scan_new_messages cannot apply masking to skipped context, falling back to a full-context scan", + self.guardrail_name, + ) + return None + + await self.mark_texts_scanned( + texts=texts, + request_data=request_data, + cache=cache, + ) + return inputs + async def apply_guardrail( self, inputs: "GenericGuardrailAPIInputs", @@ -2077,6 +2161,15 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): try: verbose_proxy_logger.debug(f"Bedrock Guardrail: Applying guardrail to {len(texts)} text(s)") + if input_type == "request": + incremental_result = await self._apply_incremental_request_scan( + texts=texts, + inputs=inputs, + request_data=request_data, + ) + if incremental_result is not None: + return incremental_result + masked_texts = [] selection = self._select_messages_for_apply_guardrail( diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py index 5e62ab96f0c..d91ddffa0c1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/__init__.py @@ -27,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" mask_response_content=litellm_params.mask_response_content, fail_on_error=litellm_params.fail_on_error, skip_unscannable_attachments=litellm_params.skip_unscannable_attachments, + sanitize_error_detail=litellm_params.sanitize_error_detail, ) litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback) diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index 32a3cebfca0..31535a5b569 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -11,6 +11,7 @@ from typing import ( Union, ) +import httpx from fastapi import HTTPException if TYPE_CHECKING: @@ -35,7 +36,8 @@ from litellm.proxy.guardrails.guardrail_hooks.model_armor.file_scanning import ( MODEL_ARMOR_MAX_FILE_SIZE_BYTES, plan_file_scans, ) -from litellm.types.guardrails import GuardrailEventHooks +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH +from litellm.types.guardrails import GuardrailEventHooks, LitellmParams from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ( CallTypes, @@ -50,6 +52,33 @@ from litellm.types.utils import ( GUARDRAIL_NAME = "model_armor" +class ModelArmorAPIError(Exception): + """Model Armor API failure (non-2xx), distinct from a content-block decision so + hooks can honor fail_on_error. The detail is already sanitized per configuration.""" + + def __init__(self, detail: str): + super().__init__(detail) + self.detail = detail + + +_SCANNED_CONTENT_KEYS = frozenset({"text", "sanitizedText", "findings", "maliciousUriMatchedItems"}) + +RedactablePayload = Union[dict, list, str, int, float, bool, None] + + +def _redact_scanned_content(payload: RedactablePayload, depth: int = 0) -> RedactablePayload: + if depth >= DEFAULT_MAX_RECURSE_DEPTH: + return "[REDACTED]" + if isinstance(payload, dict): + return { + key: "[REDACTED]" if key in _SCANNED_CONTENT_KEYS else _redact_scanned_content(value, depth + 1) + for key, value in payload.items() + } + if isinstance(payload, list): + return [_redact_scanned_content(item, depth + 1) for item in payload] + return payload + + class ModelArmorGuardrail(CustomGuardrail, VertexBase): """ Google Cloud Model Armor Guardrail integration for LiteLLM. @@ -76,6 +105,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): location: Optional[str] = None, credentials: Optional[Any] = None, api_endpoint: Optional[str] = None, + sanitize_error_detail: "bool | None" = True, **kwargs, ): # Set supported event hooks if not already provided @@ -98,6 +128,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): self.location = location or "us-central1" self.credentials = credentials self.api_endpoint = api_endpoint + self.sanitize_error_detail = sanitize_error_detail is not False # Store optional params self.optional_params = kwargs @@ -141,6 +172,67 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): verbose_proxy_logger.debug("Model Armor: Skipping non-ModelResponse type: %s", type(response).__name__) return "" + def _build_api_error_detail(self, status_code: int, response_text: str) -> str: + if self.sanitize_error_detail: + return f"Model Armor API error (upstream {status_code})" + return f"Model Armor API error (upstream {status_code}): {response_text}" + + def _build_block_error_detail(self, message: str, armor_response: RedactablePayload) -> dict: + if self.sanitize_error_detail: + return {"error": message} + return {"error": message, "model_armor_response": armor_response} + + def _build_logging_response(self, armor_response: RedactablePayload) -> RedactablePayload: + if self.sanitize_error_detail: + return _redact_scanned_content(armor_response) + return armor_response + + def _raise_if_fail_closed(self, e: ModelArmorAPIError) -> None: + if self.optional_params.get("fail_on_error", True): + raise e from None + + def update_in_memory_litellm_params(self, litellm_params: LitellmParams) -> None: + super().update_in_memory_litellm_params(litellm_params) + self.sanitize_error_detail = self.sanitize_error_detail is not False + + def _log_request_debug( + self, + url: str, + body: dict, + file_bytes: "bytes | None", + file_type: "str | None", + ) -> None: + # Never log byteData: it is the full base64 of the scanned document. Log only its + # type and size so debug deployments cannot leak the contents the guardrail inspects. + if file_bytes is not None and file_type is not None: + verbose_proxy_logger.debug( + "Model Armor file request - URL: %s, byteDataType: %s, bytes: %d", + url, + file_type, + len(file_bytes), + ) + elif self.sanitize_error_detail: + verbose_proxy_logger.debug("Model Armor request - URL: %s", url) + else: + verbose_proxy_logger.debug( + "Model Armor request - URL: %s, Body: %s", + url, + body, + ) + + def _log_response_debug(self, status_code: int, response_text: str) -> None: + if self.sanitize_error_detail: + verbose_proxy_logger.debug( + "Model Armor response - Status: %s", + status_code, + ) + else: + verbose_proxy_logger.debug( + "Model Armor response - Status: %s, Body: %s", + status_code, + response_text, + ) + async def make_model_armor_request( self, content: Optional[str] = None, @@ -185,48 +277,37 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): "Authorization": f"Bearer {access_token}", } - # Never log byteData: it is the full base64 of the scanned document. Log only its - # type and size so debug deployments cannot leak the contents the guardrail inspects. - if file_bytes is not None and file_type is not None: - verbose_proxy_logger.debug( - "Model Armor file request - URL: %s, byteDataType: %s, bytes: %d", - url, - file_type, - len(file_bytes), - ) - else: - verbose_proxy_logger.debug( - "Model Armor request - URL: %s, Body: %s", - url, - body, - ) + self._log_request_debug(url=url, body=body, file_bytes=file_bytes, file_type=file_type) # Make request if self.async_handler is None: raise ValueError("Async handler not initialized") - response = await self.async_handler.post( - url=url, - json=body, - headers=headers, - ) + try: + response = await self.async_handler.post( + url=url, + json=body, + headers=headers, + ) + except httpx.HTTPStatusError as e: + detail = self._build_api_error_detail(e.response.status_code, e.response.text) + verbose_proxy_logger.error( + "Model Armor API error - Status: %s, Detail: %s", + e.response.status_code, + detail, + ) + raise ModelArmorAPIError(detail) from None - verbose_proxy_logger.debug( - "Model Armor response - Status: %s, Body: %s", - response.status_code, - response.text, - ) + self._log_response_debug(status_code=response.status_code, response_text=response.text) if response.status_code != 200: + detail = self._build_api_error_detail(response.status_code, response.text) verbose_proxy_logger.error( - "Model Armor API error - Status: %s, Response: %s", + "Model Armor API error - Status: %s, Detail: %s", response.status_code, - response.text, - ) - raise HTTPException( - status_code=400, - detail=f"Model Armor API error (upstream {response.status_code}): {response.text}", + detail, ) + raise ModelArmorAPIError(detail) json_response = response.json() if hasattr(json_response, "__await__"): @@ -351,9 +432,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): Override to store only the Model Armor API response, not the entire data dict. This prevents circular references in logging. """ - # Retrieve the Model Armor response & status stored on the per-request `metadata` object. metadata = request_data.get("metadata", {}) if isinstance(request_data, dict) else {} - guardrail_response = metadata.get("_model_armor_response", {}) # Determine status – default to "success" but prefer the explicit value if present. @@ -444,6 +523,9 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): file_bytes=attachment.file_bytes, file_type=attachment.byte_data_type, ) + except ModelArmorAPIError as e: + self._raise_if_fail_closed(e) + continue except HTTPException: raise except Exception as e: @@ -459,7 +541,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # otherwise a PII-only (SDP deidentify) document would pass through unscrubbed. blocked = self._should_block_content(armor_response, allow_sanitization=False) metadata["_model_armor_response"] = self._append_armor_response( - metadata.get("_model_armor_response"), armor_response + metadata.get("_model_armor_response"), + self._build_logging_response(armor_response), ) if blocked or metadata.get("_model_armor_status") == "blocked": metadata["_model_armor_status"] = "blocked" @@ -469,10 +552,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if blocked: raise HTTPException( status_code=400, - detail={ - "error": "Content blocked by Model Armor", - "model_armor_response": armor_response, - }, + detail=self._build_block_error_detail("Content blocked by Model Armor", armor_response), ) @log_guardrail_information @@ -530,7 +610,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): metadata = data.setdefault("metadata", {}) # ensures metadata exists and is unique per request # Accumulate so a prior file scan on the same request is not overwritten by this text scan. metadata["_model_armor_response"] = self._append_armor_response( - metadata.get("_model_armor_response"), armor_response + metadata.get("_model_armor_response"), + self._build_logging_response(armor_response), ) # Pre-compute guardrail status for downstream logging. A blocked response will eventually raise # an HTTPException, however in scenarios where the caller decides to ignore the exception (e.g. @@ -548,10 +629,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if blocked: raise HTTPException( status_code=400, - detail={ - "error": "Content blocked by Model Armor", - "model_armor_response": armor_response, - }, + detail=self._build_block_error_detail("Content blocked by Model Armor", armor_response), ) # If mask_request_content is enabled, update messages with sanitized content @@ -565,6 +643,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): data["messages"] = set_last_user_message(messages, sanitized_content) + except ModelArmorAPIError as e: + self._raise_if_fail_closed(e) except HTTPException: raise except Exception as e: @@ -625,7 +705,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): metadata = data.setdefault("metadata", {}) # Accumulate so a prior file scan on the same request is not overwritten by this text scan. metadata["_model_armor_response"] = self._append_armor_response( - metadata.get("_model_armor_response"), armor_response + metadata.get("_model_armor_response"), + self._build_logging_response(armor_response), ) if blocked or metadata.get("_model_armor_status") == "blocked": metadata["_model_armor_status"] = "blocked" @@ -640,10 +721,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if blocked: raise HTTPException( status_code=400, - detail={ - "error": "Content blocked by Model Armor", - "model_armor_response": armor_response, - }, + detail=self._build_block_error_detail("Content blocked by Model Armor", armor_response), ) # If mask_request_content is enabled, update messages with sanitized content @@ -656,6 +734,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): data["messages"] = set_last_user_message(messages, sanitized_content) + except ModelArmorAPIError as e: + self._raise_if_fail_closed(e) except HTTPException: raise except Exception as e: @@ -698,7 +778,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Attach Model Armor response & status to this request's metadata to prevent race conditions if isinstance(armor_response, dict): model_armor_logged_object = { - "model_armor_response": armor_response, + "model_armor_response": self._build_logging_response(armor_response), "model_armor_status": ( "blocked" if self._should_block_content( @@ -729,10 +809,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if self._should_block_content(armor_response, allow_sanitization=self.mask_response_content): raise HTTPException( status_code=400, - detail={ - "error": "Response blocked by Model Armor", - "model_armor_response": armor_response, - }, + detail=self._build_block_error_detail("Response blocked by Model Armor", armor_response), ) # If mask_response_content is enabled, update response with sanitized content @@ -746,6 +823,8 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if choice.message.content: choice.message.content = sanitized_content + except ModelArmorAPIError as e: + self._raise_if_fail_closed(e) except HTTPException: raise except Exception as e: @@ -790,7 +869,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): # Attach Model Armor response & status to this request's metadata to avoid race conditions if isinstance(request_data, dict): metadata = request_data.setdefault("metadata", {}) - metadata["_model_armor_response"] = armor_response + metadata["_model_armor_response"] = self._build_logging_response(armor_response) metadata["_model_armor_status"] = ( "blocked" if self._should_block_content(armor_response) else "success" ) @@ -809,10 +888,10 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): if self._should_block_content(armor_response): raise HTTPException( status_code=400, - detail={ - "error": "Streaming response blocked by Model Armor", - "model_armor_response": armor_response, - }, + detail=self._build_block_error_detail( + "Streaming response blocked by Model Armor", + armor_response, + ), ) # Apply sanitization if enabled @@ -831,6 +910,11 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): yield chunk return + except ModelArmorAPIError as e: + if self.optional_params.get("fail_on_error", True): + error_obj = {"message": e.detail, "code": "500"} + yield f"data: {json.dumps({'error': error_obj})}\n\n" + return except HTTPException as e: # Yield error as SSE event so create_response() detects it and # returns a proper JSON error response with the correct status code. diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 14e76a21093..e909c15382b 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -35,6 +35,7 @@ def initialize_bedrock(litellm_params: LitellmParams, guardrail: Guardrail): aws_sts_endpoint=litellm_params.aws_sts_endpoint, aws_bedrock_runtime_endpoint=litellm_params.aws_bedrock_runtime_endpoint, experimental_use_latest_role_message_only=litellm_params.experimental_use_latest_role_message_only, + only_scan_new_messages=litellm_params.only_scan_new_messages or False, ) litellm.logging_callback_manager.add_litellm_callback(_bedrock_callback) return _bedrock_callback diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 6514d4e1e8c..9d9ef28ec9b 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -19,7 +19,10 @@ from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( iter_client_callback_metadata_dicts, ) from litellm.litellm_core_utils.safe_json_loads import safe_json_loads -from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host +from litellm.litellm_core_utils.url_utils import ( + is_url_destination_allowed_by_host, + provider_url_destination_candidates, +) from litellm.proxy._types import ( AddTeamCallback, CommonProxyErrors, @@ -227,23 +230,26 @@ def _reject_url_valued_destinations(data: Dict[str, Any]) -> None: allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or [] for field in _URL_DESTINATION_REQUEST_FIELDS: value = data.get(field) - if not isinstance(value, str) or not value.startswith(("http://", "https://")): + if not isinstance(value, str): continue - if is_url_destination_allowed_by_host(value, allowed_hosts): - continue - raise HTTPException( - status_code=400, - detail={ - "error": "invalid_request", - "param": field, - "message": ( - f"URL-valued '{field}' is not allowed. Configure custom " - "endpoints with api_base instead, or add the destination " - "host to `provider_url_destination_allowed_hosts` in " - "litellm_settings." - ), - }, - ) + for candidate in provider_url_destination_candidates(value): + if not candidate.lower().startswith(("http://", "https://")): + continue + if is_url_destination_allowed_by_host(candidate, allowed_hosts): + continue + raise HTTPException( + status_code=400, + detail={ + "error": "invalid_request", + "param": field, + "message": ( + f"URL-valued '{field}' is not allowed. Configure custom " + "endpoints with api_base instead, or add the destination " + "host to `provider_url_destination_allowed_hosts` in " + "litellm_settings." + ), + }, + ) def _strip_untrusted_request_header_controls( diff --git a/litellm/proxy/management_endpoints/cache_settings_endpoints.py b/litellm/proxy/management_endpoints/cache_settings_endpoints.py index 9f45cb619aa..7c0d8958a28 100644 --- a/litellm/proxy/management_endpoints/cache_settings_endpoints.py +++ b/litellm/proxy/management_endpoints/cache_settings_endpoints.py @@ -18,8 +18,9 @@ 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 -from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys +from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import ( AUDIT_ACTIONS, LiteLLM_AuditLogs, @@ -43,6 +44,17 @@ router = APIRouter() # (e.g. redis://:secret@host:6379/1). _CACHE_SENSITIVE_FIELDS: set = {"password", "sentinel_password", "url"} +# The env fallback resolves the full set of redis.Redis kwargs, which includes +# credential-bearing params (azure_client_secret, ssl_password, ...) that are +# not cache UI fields. Only overlay fields the settings page actually renders, +# so the read never surfaces a credential the UI does not manage. +_CACHE_SETTINGS_FIELD_NAMES: frozenset = frozenset(field.field_name for field in CACHE_SETTINGS_FIELDS) + +# Classifier used, alongside _CACHE_SENSITIVE_FIELDS, to redact any +# credential-bearing key before it leaves the server (`url` is kept in the +# explicit set because its name carries no sensitive segment). +_CREDENTIAL_CLASSIFIER = SensitiveDataMasker() + _REDACTED_VALUE = "***REDACTED***" @@ -67,6 +79,165 @@ def _resolve_cache_url_precedence(settings: Mapping[str, Any]) -> dict[str, Any] return {k: v for k, v in settings.items() if k not in _URL_OVERRIDDEN_CONNECTION_FIELDS} +def _parse_stored_settings(cache_settings_value: object) -> dict[str, Any]: + """Normalize a stored cache_settings blob to a dict. + + The prisma column comes back as either a JSON string or an already-parsed + dict depending on the client, so callers that json.loads unconditionally + silently drop the whole (still-encrypted) row on the dict path. + """ + parsed = json.loads(cache_settings_value) if isinstance(cache_settings_value, str) else cache_settings_value + return parsed if isinstance(parsed, dict) else {} + + +def _overlay_environment(stored: Mapping[str, Any]) -> dict[str, Any]: + """Fill connection fields from the REDIS_* environment the cache actually reads. + + A response cache pointed at Redis resolves host/port/password/etc. from the + REDIS_* env vars when the stored config leaves them unset, so a cache + configured purely through the environment works while its settings page, + which reads only the database row, shows blank. Overlaying the same env + kwargs the runtime uses makes the page reflect the effective connection. + Stored values win; the environment only fills what the stored config omits. + """ + env_kwargs = { + key: value for key, value in _redis_kwargs_from_environment().items() if key in _CACHE_SETTINGS_FIELD_NAMES + } + if not env_kwargs: + return dict(stored) + effective = {**env_kwargs, **stored} + # the env fallback is a Redis connection, so name the type when the stored + # config did not, letting the UI render the Redis fields it just populated + effective.setdefault("type", "redis") + return effective + + +def _redact_credentials(settings: Mapping[str, Any]) -> dict[str, Any]: + """Replace credential-bearing values with a fixed marker, keeping the rest. + + The marker is unambiguous on the way back in: an admin who edits an + unrelated field and re-submits sends the marker for the untouched secret, + which the update path maps back to the stored value rather than persisting + the marker over a working password. + """ + return { + key: (_REDACTED_VALUE if value is not None and _is_credential_field(key) else value) + for key, value in settings.items() + } + + +def _is_credential_field(key: str) -> bool: + """Whether a cache setting carries a credential and must be redacted on read.""" + return key in _CACHE_SENSITIVE_FIELDS or _CREDENTIAL_CLASSIFIER.is_sensitive_key(key) + + +def _has_connection_target(value: object) -> bool: + """Whether a payload value names a live discrete connection target.""" + if isinstance(value, str): + return value.strip() != "" and value != _REDACTED_VALUE + return value not in (None, [], {}) + + +# Every field that identifies which Redis a credential belongs to, across node +# (host/port/url), cluster (redis_startup_nodes), and sentinel +# (sentinel_nodes/service_name) modes. A stored secret is bound to these. +_CONNECTION_TARGET_FIELDS: tuple = ( + "host", + "port", + "url", + "redis_startup_nodes", + "sentinel_nodes", + "service_name", +) + + +def _target_repr(value: object) -> str: + """Canonical string form of a connection-target value for equality checks. + + The client may serialize the same target differently from storage (a port as + "6379" vs 6379, node lists round-tripped through JSON), so compare normalized + forms rather than raw values to avoid treating an unchanged target as a change. + """ + if isinstance(value, (list, dict)): + return json.dumps(value, sort_keys=True, default=str) + return str(value) + + +def _saved_secret_is_reusable(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> bool: + """Whether a stored credential may be restored for this request. + + A stored secret belongs to the stored connection target, so it is reused only + when the request describes that same target on every dimension the stored + config pins (host/port, url, cluster nodes, sentinel nodes/service). This + prevents credential replay: a caller cannot omit the credential, point at a + different (or incomplete) target, and have the proxy send the stored secret + to a Redis of their choosing. + + Non-secret target fields (host/port/nodes/service) must be supplied and match + in normalized form, so equivalent representations (port "6379" vs 6379) are + not seen as a change while an omitted or different value is. ``url`` is the + exception: it is itself the secret and the form never re-prefills it, so a + redacted or omitted url means "keep the stored url" (same target) and only a + different supplied url blocks reuse. + """ + for field in _CONNECTION_TARGET_FIELDS: + saved_value = saved.get(field) + if saved_value in (None, "", [], {}): + continue # the stored config does not pin this dimension + incoming_value = incoming.get(field) + if field == "url": + if incoming_value in (None, "", _REDACTED_VALUE): + continue # url kept as-is (same target) + if _target_repr(incoming_value) != _target_repr(saved_value): + return False + continue + if _target_repr(incoming_value) != _target_repr(saved_value): + return False # a pinned target field is missing or different + return True + + +def _merge_over_saved(incoming: Mapping[str, Any], saved: Mapping[str, Any]) -> dict[str, Any]: + """Keep the stored secret behind any credential the caller echoed back redacted or omitted. + + GET returns credentials as the marker and the form never re-prefills a + secret, so a save that does not touch a credential arrives with the marker + or with the field absent. Either way the real secret must survive: it is + restored from the stored row, or dropped when there is no stored row (the + value is env-sourced and the marker must never be persisted). Non-secret + fields are taken from the incoming payload as-is, so clearing one still works. + + ``url`` is the exception: it is credential-bearing (redacted) yet also a + connection-mode selector that url-precedence resolves against host/port. If + the caller supplies a discrete target (host, cluster, or sentinel nodes), a + stored url is a stale mode the caller is leaving, so it is dropped rather + than restored, otherwise url-precedence would resurrect it and discard the + submitted host/port. + """ + switching_to_discrete_target = ( + _has_connection_target(incoming.get("host")) + or _has_connection_target(incoming.get("redis_startup_nodes")) + or _has_connection_target(incoming.get("sentinel_nodes")) + ) + reuse_saved_secret = _saved_secret_is_reusable(incoming, saved) + merged = dict(incoming) + for field in _CACHE_SENSITIVE_FIELDS: + # A value the caller explicitly supplied is honored verbatim: a new + # secret, or an empty string / null to clear the stored one. Only an + # omitted field or the echoed-back marker triggers preserve-or-drop. + if field in incoming and incoming[field] != _REDACTED_VALUE: + continue + if field == "url" and switching_to_discrete_target: + merged.pop(field, None) + continue + if field in saved and reuse_saved_secret: + merged[field] = saved[field] + else: + # nothing stored to reuse, or the caller is pointing at a different + # target: never persist/replay the marker or the stored secret + merged.pop(field, None) + return merged + + def _redact_settings(settings: Optional[Mapping[str, Any]]) -> Dict[str, Any]: """Replace every value in a settings map with a fixed marker. @@ -270,34 +441,34 @@ async def get_cache_settings( # Get cache settings fields from types file cache_fields = [field.model_copy(deep=True) for field in CACHE_SETTINGS_FIELDS] - # Try to get cache settings from database - current_values = {} + # Read the stored settings (decrypted); an env-only cache has none. + stored: dict[str, Any] = {} if prisma_client is not None: cache_config = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"}) if cache_config is not None and cache_config.cache_settings: - # Decrypt cache settings - cache_settings_json = cache_config.cache_settings - if isinstance(cache_settings_json, str): - cache_settings_dict = json.loads(cache_settings_json) - else: - cache_settings_dict = cache_settings_json + stored = proxy_config._decrypt_db_variables( + variables_dict=_parse_stored_settings(cache_config.cache_settings) + ) - # Decrypt environment variables - decrypted_settings = proxy_config._decrypt_db_variables(variables_dict=cache_settings_dict) + # Fill connection fields from the REDIS_* environment the cache resolves + # from when the stored config leaves them unset, then apply url precedence + # so a url-mode config does not surface conflicting discrete fields (which + # would otherwise let a no-op save silently switch it to host/port). + effective = _resolve_cache_url_precedence(_overlay_environment(stored)) - # Derive redis_type for UI based on settings - # UI uses redis_type to show/hide fields, backend only stores 'type' - if decrypted_settings.get("type") == "redis": - if decrypted_settings.get("redis_startup_nodes"): - decrypted_settings["redis_type"] = "cluster" - elif decrypted_settings.get("sentinel_nodes"): - decrypted_settings["redis_type"] = "sentinel" - else: - decrypted_settings["redis_type"] = "node" + # Derive redis_type for UI based on settings + # UI uses redis_type to show/hide fields, backend only stores 'type' + if effective.get("type") == "redis": + if effective.get("redis_startup_nodes"): + effective["redis_type"] = "cluster" + elif effective.get("sentinel_nodes"): + effective["redis_type"] = "sentinel" + else: + effective["redis_type"] = "node" - # Mask credential fields so the GET response never carries - # plaintext Redis / Sentinel passwords off the server. - current_values = mask_sensitive_keys(decrypted_settings, _CACHE_SENSITIVE_FIELDS) + # Redact credential fields so the GET response never carries a plaintext + # Redis / Sentinel password off the server. + current_values = _redact_credentials(effective) # Update field values with current values for field in cache_fields: @@ -331,10 +502,27 @@ async def test_cache_connection( to verify the credentials work without affecting global state. """ from litellm import Cache + from litellm.proxy.proxy_server import prisma_client, proxy_config try: - cache_settings = _resolve_cache_url_precedence(request.cache_settings) - verbose_proxy_logger.debug("Testing cache connection with settings: %s", cache_settings) + # A credential the form left untouched arrives redacted; resolve it back + # to the stored secret so the test connects with the real password. A + # lookup failure must not block the test, so fall back to no stored row. + saved_settings: dict[str, Any] = {} + if prisma_client is not None: + try: + existing_row = await CacheConfigRepository(prisma_client).table.find_unique( + where={"id": "cache_config"} + ) + if existing_row is not None and existing_row.cache_settings: + saved_settings = proxy_config._decrypt_db_variables( + variables_dict=_parse_stored_settings(existing_row.cache_settings) + ) + except Exception: # noqa: BLE001 - a saved-settings lookup failure must not block a connection test + saved_settings = {} + cache_settings = _resolve_cache_url_precedence(_merge_over_saved(request.cache_settings, saved_settings)) + # cache_settings now carries the resolved plaintext credential; never log it raw + verbose_proxy_logger.debug("Testing cache connection with settings: %s", _redact_credentials(cache_settings)) # Only support Redis for now if cache_settings.get("type") != "redis": @@ -400,19 +588,20 @@ async def update_cache_settings( ) try: - cache_settings = _resolve_cache_url_precedence(request.cache_settings) - - # Snapshot the prior settings (key set only — values get redacted in - # the audit row) so the audit-log entry shows which fields changed. + # Read the stored row first: its decrypted values back any credential the + # caller echoed back redacted, and its key set drives the audit diff. existing_row = await CacheConfigRepository(prisma_client).table.find_unique(where={"id": "cache_config"}) before_settings: Optional[Dict[str, Any]] = None + saved_settings: dict[str, Any] = {} if existing_row is not None and existing_row.cache_settings: - try: - before_settings = json.loads(existing_row.cache_settings) - except (TypeError, ValueError): - before_settings = None + before_settings = _parse_stored_settings(existing_row.cache_settings) + saved_settings = proxy_config._decrypt_db_variables(variables_dict=before_settings) action: AUDIT_ACTIONS = "updated" if existing_row is not None else "created" + # Preserve stored secrets behind any redacted or omitted credential, then + # resolve the url-vs-discrete-fields precedence. + cache_settings = _resolve_cache_url_precedence(_merge_over_saved(request.cache_settings, saved_settings)) + # Encrypt sensitive fields (keep redis_type for storage) encrypted_settings = proxy_config._encrypt_env_variables(environment_variables=cache_settings) @@ -461,7 +650,7 @@ async def update_cache_settings( return { "message": "Cache settings updated successfully", "status": "success", - "settings": cache_settings, + "settings": _redact_credentials(cache_settings), } except Exception as e: verbose_proxy_logger.error(f"Error updating cache settings: {str(e)}") diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 01f4e040e58..ac6a2a4a7db 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -42,6 +42,9 @@ from litellm.proxy._experimental.mcp_server.db import ( rotate_mcp_user_credentials_master_key, rotate_mcp_user_env_vars_master_key, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import ( + rotate_sso_identity_assertions_master_key, +) from litellm.proxy._types import * from litellm.proxy._types import LiteLLM_VerificationToken, hash_token from litellm.proxy.auth.auth_checks import ( @@ -4242,6 +4245,15 @@ async def _rotate_master_key( except Exception as e: verbose_proxy_logger.warning("Failed to rotate MCP user env vars: %s", str(e)) + # 4d. process SSO identity assertion table (EMA subject tokens) + try: + await rotate_sso_identity_assertions_master_key( + prisma_client=prisma_client, + new_master_key=new_master_key, + ) + except Exception as e: # noqa: BLE001 # one store's failure must not abort the master-key rotation + verbose_proxy_logger.warning("Failed to rotate SSO identity assertions: %s", str(e)) + # 5. process credentials table try: credentials = await CredentialsRepository(prisma_client).table.find_many() diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index d50db8324ef..89f28a30a84 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -136,9 +136,12 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( _raise_if_not_oauth2, authorize_with_server, + client_supplied_redirect_uris, exchange_token_with_server, get_request_base_url, + redeem_passthrough_authorization_code, register_client_with_server, + resolve_ephemeral_dcr_client, ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( global_mcp_server_manager, @@ -1661,7 +1664,21 @@ if MCP_AVAILABLE: mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request) _raise_if_not_oauth2(mcp_server) # Use the server's stored client_id when the caller doesn't supply one - resolved_client_id = mcp_server.client_id or client_id or "" + stored_or_supplied_client_id = mcp_server.client_id or client_id or "" + ephemeral_dcr_client = ( + await resolve_ephemeral_dcr_client( + request=request, + mcp_server=mcp_server, + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + redirect_uri=redirect_uri, + ) + if not stored_or_supplied_client_id + else None + ) + resolved_client_id = stored_or_supplied_client_id or ( + ephemeral_dcr_client.client_id if ephemeral_dcr_client else "" + ) if not resolved_client_id: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -1683,6 +1700,7 @@ if MCP_AVAILABLE: code_challenge_method=code_challenge_method, response_type=response_type, scope=scope, + ephemeral_dcr_client=ephemeral_dcr_client, ) @router.post( @@ -1705,7 +1723,21 @@ if MCP_AVAILABLE: ): mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request) _raise_if_not_oauth2(mcp_server) - resolved_client_id = mcp_server.client_id or client_id or "" + # Sealed passthrough codes exist only for the authorization_code grant. A refresh_token + # grant must never open one: the minted client is unrecoverable after the single flow by + # contract, so an expired browser-held token re-runs authorize instead. + sealed_code = ( + redeem_passthrough_authorization_code(code=code, mcp_server=mcp_server, code_verifier=code_verifier) + if grant_type == "authorization_code" + else None + ) + resolved_code = sealed_code.upstream_code if sealed_code else code + # A sealed flow ran the gateway /callback as its upstream redirect (bridge short-circuit + # or plain flow alike), so the exchange must present that binding, not the browser page. + resolved_redirect_uri = f"{get_request_base_url(request)}/callback" if sealed_code else redirect_uri + caller_client_id = sealed_code.client_id if sealed_code else client_id + caller_client_secret = sealed_code.client_secret if sealed_code else client_secret + resolved_client_id = mcp_server.client_id or caller_client_id or "" if not resolved_client_id: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -1721,13 +1753,14 @@ if MCP_AVAILABLE: request=request, mcp_server=mcp_server, grant_type=grant_type, - code=code, - redirect_uri=redirect_uri, + code=resolved_code, + redirect_uri=resolved_redirect_uri, client_id=resolved_client_id, - client_secret=client_secret, + client_secret=caller_client_secret, code_verifier=code_verifier, refresh_token=refresh_token, scope=scope, + client_token_endpoint_auth_method=sealed_code.token_endpoint_auth_method if sealed_code else None, ) @router.post( @@ -1743,6 +1776,7 @@ if MCP_AVAILABLE: mcp_server = await _get_cached_temporary_mcp_server_or_404(server_id, user_api_key_dict, request=request) request_data = await _read_request_body(request=request) data: dict = {**request_data} + client_redirect_uris = client_supplied_redirect_uris(data.get("redirect_uris")) return await register_client_with_server( request=request, @@ -1753,6 +1787,7 @@ if MCP_AVAILABLE: token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""), fallback_client_id=server_id, persist_credentials=_user_is_full_admin(user_api_key_dict), + client_redirect_uris=client_redirect_uris, ) @router.delete( diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 138a55d9227..5a289d22f99 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -13,16 +13,18 @@ Endpoints for /organization operations #### ORGANIZATION MANAGEMENT #### -from typing import Any, Dict, List, Optional, Tuple +from typing import Annotated, Any, Dict, List, Mapping, Optional, Tuple import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, status +from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.proxy._types import * from litellm.proxy.auth.auth_checks import can_user_call_model, get_user_object from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.management_endpoints.budget_management_endpoints import ( new_budget, update_budget, @@ -34,6 +36,7 @@ from litellm.proxy.management_endpoints.common_utils import ( ) from litellm.proxy.management_helpers.object_permission_utils import ( handle_update_object_permission_common, + prepare_object_permission_upsert, ) from litellm.proxy.management_helpers.utils import ( get_new_internal_user_defaults, @@ -101,6 +104,30 @@ async def _verify_org_access( ) +_STR_OBJECT_DICT_ADAPTER = TypeAdapter(dict[str, object]) +_BUDGET_SETTABLE_FIELDS = frozenset(LiteLLM_BudgetTable.model_fields.keys()) - {"budget_id"} +_ORG_COLUMN_FIELDS = frozenset({"organization_alias", "models"}) + + +def build_budget_write_data(budget_updates: Mapping[str, object], updated_by: str) -> Mapping[str, object]: + """ + Budget-row columns to write. ``budget_reset_at`` tracks any sent ``budget_duration``: + recomputed for a new duration, cleared alongside a ``None`` duration so no stale reset + timestamp survives. Other sent fields (including a ``None`` clear) are written as-is. + """ + budget_duration = budget_updates.get("budget_duration") + recomputed_reset_at: Mapping[str, object] = ( + { + "budget_reset_at": ( + get_budget_reset_time(budget_duration=budget_duration) if isinstance(budget_duration, str) else None + ) + } + if "budget_duration" in budget_updates + else {} + ) + return {**budget_updates, **recomputed_reset_at, "updated_by": updated_by} + + def handle_nested_budget_structure_in_organization_update_request( raw_data: dict, ) -> dict: @@ -556,6 +583,154 @@ async def handle_update_object_permission( return data_json +@router.patch( + "/v2/organization/{organization_id}", + tags=["organization management"], + dependencies=[Depends(user_api_key_auth)], + response_model=LiteLLM_OrganizationTableWithMembers, + include_in_schema=False, +) +async def update_organization_v2( + organization_id: str, + data: OrganizationUpdateRequestV2, + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +): + """ + Partial update of an organization (RESTful PATCH, RFC 7396 merge-patch semantics). + + A sent field is written and an omitted one is left untouched (presence is read from + ``model_fields_set``). Clear tokens are per field: budget limits and ``metadata`` clear with + ``null``, ``models`` with ``[]``, and ``object_permission`` with ``null`` (it merges when sent, + so an empty ``{}`` is rejected). ``organization_alias`` is required and cannot be cleared. + Validation failures return 422; the object-permission upsert, budget-row write, and + org-row write are one transaction. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + + if user_api_key_dict.user_id is None: + raise HTTPException( + status_code=400, + detail={ + "error": "Cannot associate a user_id to this action. Check `/key/info` to validate if 'user_id' is set." + }, + ) + + if data.max_budget is not None and (not math.isfinite(data.max_budget) or data.max_budget < 0): + raise HTTPException( + status_code=422, + detail={"error": f"max_budget must be a non-negative finite number. Received: {data.max_budget}"}, + ) + if data.soft_budget is not None and (not math.isfinite(data.soft_budget) or data.soft_budget < 0): + raise HTTPException( + status_code=422, + detail={"error": f"soft_budget must be a non-negative finite number. Received: {data.soft_budget}"}, + ) + if data.model_max_budget: + from litellm.proxy.management_endpoints.key_management_endpoints import ( + validate_model_max_budget, + ) + + try: + validate_model_max_budget(data.model_max_budget) + except ValueError as e: + raise HTTPException(status_code=422, detail={"error": str(e)}) + + if "organization_alias" in data.model_fields_set and data.organization_alias is None: + raise HTTPException( + status_code=422, + detail={"error": "organization_alias cannot be cleared; it is required"}, + ) + if "models" in data.model_fields_set and data.models is None: + raise HTTPException( + status_code=422, + detail={"error": "models cannot be set to null; send [] to clear it"}, + ) + if data.object_permission is not None and not data.object_permission.model_dump(exclude_none=True): + raise HTTPException( + status_code=422, + detail={ + "error": "object_permission cannot be an empty object; send null to clear it, or a non-empty object to set grants" + }, + ) + + await _verify_org_access( + organization_id=organization_id, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + + existing_organization_row = await OrganizationRepository(prisma_client).table.find_unique( + where={"organization_id": organization_id}, + ) + if existing_organization_row is None: + raise HTTPException( + status_code=404, + detail={"error": f"Organization not found for organization_id={organization_id}"}, + ) + + field_values = _STR_OBJECT_DICT_ADAPTER.validate_python(data.model_dump()) + present_fields = data.model_fields_set + budget_updates = {field: field_values[field] for field in present_fields if field in _BUDGET_SETTABLE_FIELDS} + org_column_updates: Mapping[str, object] = { + **{field: field_values[field] for field in present_fields if field in _ORG_COLUMN_FIELDS}, + **({"metadata": data.metadata or {}} if "metadata" in present_fields else {}), + } + + object_permission_cleared = "object_permission" in present_fields and data.object_permission is None + object_permission_upsert = ( + await prepare_object_permission_upsert( + new_object_permission=data.object_permission.model_dump(exclude_none=True), + existing_object_permission_id=existing_organization_row.object_permission_id, + prisma_client=prisma_client, + ) + if data.object_permission is not None + else None + ) + object_permission_write: Mapping[str, object] = ( + {"object_permission_id": object_permission_upsert.object_permission_id} + if object_permission_upsert is not None + else ({"object_permission_id": None} if object_permission_cleared else {}) + ) + + organization_write_data = prisma_client.jsonify_object( + { + **org_column_updates, + **object_permission_write, + "updated_by": user_api_key_dict.user_id, + } + ) + + async with prisma_client.db.tx() as tx: + if object_permission_upsert is not None: + await tx.litellm_objectpermissiontable.upsert( + where={"object_permission_id": object_permission_upsert.object_permission_id}, + data={ + "create": object_permission_upsert.record, + "update": object_permission_upsert.record, + }, + ) + if budget_updates: + await tx.litellm_budgettable.update( + where={"budget_id": existing_organization_row.budget_id}, + data=prisma_client.jsonify_object( + dict(build_budget_write_data(budget_updates, user_api_key_dict.user_id)) + ), + ) + response = await tx.litellm_organizationtable.update( + where={"organization_id": organization_id}, + data=organization_write_data, + include={"members": True, "teams": True, "litellm_budget_table": True}, + ) + + return response + + @router.delete( "/organization/delete", tags=["organization management"], diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index fa123b7d76c..90eae5bbb21 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -98,14 +98,27 @@ class UserProvisionerHelpers: if not existing_user: return None - # Update the user + new_teams = list(dict.fromkeys(new_user_request.teams or [])) + + if new_user_request.user_id != existing_user.user_id: + await UserRepository(prisma_client).table.update( + where={"user_id": existing_user.user_id}, + data={"user_id": new_user_request.user_id}, + ) + + await _handle_team_membership_changes( + user_id=new_user_request.user_id, + existing_teams=existing_user.teams or [], + new_teams=new_teams, + raise_on_error=True, + ) + updated_user = await UserRepository(prisma_client).table.update( - where={"user_id": existing_user.user_id}, + where={"user_id": new_user_request.user_id}, data={ - "user_id": new_user_request.user_id, "user_email": new_user_request.user_email, "user_alias": new_user_request.user_alias, - "teams": new_user_request.teams, + "teams": new_teams, "metadata": safe_dumps(new_user_request.metadata), **({"user_role": new_user_request.user_role} if admin_group is not None else {}), }, @@ -440,7 +453,12 @@ async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]: return members -async def _handle_team_membership_changes(user_id: str, existing_teams: List[str], new_teams: List[str]) -> None: +async def _handle_team_membership_changes( + user_id: str, + existing_teams: List[str], + new_teams: List[str], + raise_on_error: bool = False, +) -> None: """Handle adding/removing user from teams based on changes.""" existing_teams_set = set(existing_teams) new_teams_set = set(new_teams) @@ -453,6 +471,7 @@ async def _handle_team_membership_changes(user_id: str, existing_teams: List[str user_id=user_id, teams_ids_to_add_user_to=list(teams_to_add), teams_ids_to_remove_user_from=list(teams_to_remove), + raise_on_error=raise_on_error, ) @@ -1298,6 +1317,13 @@ async def delete_user( where={"team_id": team.team_id}, data={"members": new_members} ) + team_row = LiteLLM_TeamTable(**team.model_dump()) + if any(member.user_id == user_id for member in team_row.members_with_roles or []): + await team_member_delete( + data=TeamMemberDeleteRequest(team_id=team_row.team_id, user_id=user_id), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + await _set_user_keys_blocked(user_id=user_id, blocked=True) await _delete_rows_referencing_user(prisma_client, user_id=user_id) @@ -1327,6 +1353,31 @@ def _extract_group_values(value: Any) -> List[str]: return group_values +def _extract_ids_from_path_filter(path: str | None, attribute: str) -> List[str]: + """Return ids from a SCIM filtered path like ``members[value eq "id"]``. + + Okta commonly sends membership removals as a filtered path and omits the + request body ``value``, so the id lives only inside the ``[value eq "..."]`` + filter. The ``eq`` operator is matched case-insensitively per the SCIM + spec; the id keeps its original case. Per the SCIM filter grammar the + compared value must be quoted (single or double), so malformed unquoted + filters yield no id. A quoted id may contain escaped quotes and + backslashes (``\\"`` and ``\\\\``), which are unescaped before use. + ``path`` must be the raw, case-preserving path from the patch op. + """ + if not path: + return [] + match = re.match( + rf"""\s*{re.escape(attribute)}\s*\[\s*value\s+eq\s+(['"])((?:\\.|[^\\])*?)\1\s*\]\s*$""", + path, + flags=re.IGNORECASE, + ) + if not match: + return [] + extracted = re.sub(r"\\(.)", r"\1", match.group(2)) + return [extracted] if extracted else [] + + def _handle_displayname_update(op_type: str, value: Any, update_data: Dict[str, Any]) -> None: """Handle displayname updates.""" if op_type == "remove": @@ -1370,9 +1421,11 @@ def _handle_name_update(path: str, op_type: str, value: Any, scim_metadata: Dict scim_metadata["familyName"] = str(value) -def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str]) -> Optional[Set[str]]: +def _handle_group_operations(op_type: str, value: Any, teams_set: Set[str], path: str | None) -> Set[str] | None: """Handle group/team membership operations.""" group_values = _extract_group_values(value) + if not group_values and value is None: + group_values = _extract_ids_from_path_filter(path, "groups") if op_type == "replace": return set(group_values) elif op_type == "add": @@ -1485,7 +1538,7 @@ def _apply_patch_ops( elif _multi_valued_attribute_base(path) in SCIM_MULTI_VALUED_ATTRIBUTE_METADATA_KEYS: _handle_multi_valued_attribute_update(path, op_type, value, metadata) elif path.startswith("groups"): - new_replace_set = _handle_group_operations(op_type, value, teams_set) + new_replace_set = _handle_group_operations(op_type, value, teams_set, op.path) if new_replace_set is not None: replace_team_set = new_replace_set else: @@ -1497,16 +1550,29 @@ def _apply_patch_ops( return update_data, final_team_set +def _is_user_not_in_team_error(exc: HTTPException) -> bool: + """True when team_member_delete reports the user was already absent from the + team, which is the idempotent no-op case for a removal.""" + detail = exc.detail + return isinstance(detail, dict) and detail.get("error") == "User not found in team" + + async def patch_team_membership( user_id: str, teams_ids_to_add_user_to: List[str], teams_ids_to_remove_user_from: List[str], + raise_on_error: bool = False, ) -> bool: """ Add or remove user from teams Handles duplicate membership gracefully (idempotent operation). - If a user is already in a team, that's fine - we don't treat it as an error. + A user already being in a team (on add) or already absent from it (on + remove) is treated as a no-op, not an error. + + When ``raise_on_error`` is True a genuine add or remove failure (anything + other than those idempotent no-ops) propagates instead of being swallowed, + so a caller can avoid persisting a teams array the roster never received. """ for _team_id in teams_ids_to_add_user_to: try: @@ -1521,9 +1587,13 @@ async def patch_team_membership( # Handle duplicate membership gracefully - this is idempotent if e.type == ProxyErrorTypes.team_member_already_in_team: verbose_proxy_logger.debug(f"User {user_id} is already in team {_team_id}, skipping add") + elif raise_on_error: + raise else: verbose_proxy_logger.exception(f"Error adding user to team {_team_id}: {e}") except Exception as e: + if raise_on_error: + raise verbose_proxy_logger.exception(f"Error adding user to team {_team_id}: {e}") for _team_id in teams_ids_to_remove_user_from: @@ -1532,7 +1602,16 @@ async def patch_team_membership( data=TeamMemberDeleteRequest(team_id=_team_id, user_id=user_id), user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), ) + except HTTPException as e: + if _is_user_not_in_team_error(e): + verbose_proxy_logger.debug(f"User {user_id} is not in team {_team_id}, skipping remove") + elif raise_on_error: + raise + else: + verbose_proxy_logger.exception(f"Error removing user from team {_team_id}: {e}") except Exception as e: + if raise_on_error: + raise verbose_proxy_logger.exception(f"Error removing user from team {_team_id}: {e}") return True @@ -1654,8 +1733,11 @@ async def get_groups( # Convert to SCIM format scim_groups = [] for team in teams: - # Get team members with display names - members = await _get_team_members_display(team.members or []) + # Get team members with display names. members_with_roles is the + # source of truth; the legacy `members` column is not populated by + # team creation, so reading it here would report an empty member + # list to the IdP and trigger repeated re-provisioning. + members = await _get_team_members_display(await _get_team_member_user_ids_from_team(team)) verbose_proxy_logger.debug(f"SCIM GET GROUPS members: {members}") team_alias = getattr(team, "team_alias", team.team_id) team_created_at = team.created_at.isoformat() if team.created_at else None @@ -1877,16 +1959,28 @@ async def delete_group( async def _process_group_patch_operations( patch_ops: SCIMPatchOp, existing_team, prisma_client -) -> Tuple[Dict[str, Any], Set[str]]: - """Process patch operations for a group and return update data and final members.""" +) -> Tuple[Dict[str, Any], Set[str], Set[str] | None]: + """Process patch operations for a group and return update data, final members + and, when the request contained a member ``replace`` op, the absolute target + roster it declared (``None`` otherwise). + + ``add``/``remove`` are deltas relative to the current roster, but ``replace`` + is absolute: it declares the roster is exactly this set, so the caller must + reconcile against it as a set-to-target rather than rebasing it onto a + concurrently-mutated roster. + """ update_data: Dict[str, Any] = {} # Create a fresh copy of existing metadata to avoid Prisma issues existing_metadata = existing_team.metadata or {} metadata = dict(existing_metadata) if existing_metadata else {} - # Track member changes - current_members = set(existing_team.members or []) + # Track member changes. members_with_roles is the source of truth for team + # membership; the legacy `members` column is not populated by team creation + # or the real team endpoints, so seeding from it would make an `add`/`remove` + # operation recompute the member set from an empty base and silently drop + # everyone already in the team. + current_members = set(await _get_team_member_user_ids_from_team(existing_team)) final_members = current_members.copy() # Process each patch operation @@ -1908,6 +2002,8 @@ async def _process_group_patch_operations( elif path.startswith("members"): # Handle member operations member_values = _extract_group_values(value) + if not member_values and value is None: + member_values = _extract_ids_from_path_filter(op.path, "members") # Check the feature flag scim_upsert_user = await _get_scim_upsert_user_setting() # Validate all users exist or create them based on feature flag @@ -1960,27 +2056,32 @@ async def _process_group_patch_operations( if metadata: update_data["metadata"] = metadata - return update_data, final_members + member_replace_present = any( + op.op == "replace" and (op.path or "").lower().startswith("members") for op in patch_ops.Operations + ) + replace_target = set(final_members) if member_replace_present else None + + return update_data, final_members, replace_target -async def _apply_group_patch_updates( - group_id: str, update_data: Dict[str, Any], final_members: Set[str], prisma_client -): - """Apply patch updates to the group in the database.""" - # Serialize metadata if present +async def _apply_group_patch_updates(group_id: str, update_data: Dict[str, Any], prisma_client): + """Apply the group's metadata/displayName patch updates to the database. + + Membership itself is not written here; it is reconciled onto the source of + truth (members_with_roles and each member's user.teams) by + _handle_group_membership_changes via team_member_add/team_member_delete. + Writing the legacy `members` column here too would create a second, unread + copy of membership that could drift from the source of truth. + """ if "metadata" in update_data and isinstance(update_data["metadata"], dict): update_data["metadata"] = safe_dumps(update_data["metadata"]) - # Update members list - update_data["members"] = list(final_members) - - # Update team in database - updated_team = await TeamRepository(prisma_client).table.update( - where={"team_id": group_id}, - data=update_data, - ) - - return updated_team + if update_data: + return await TeamRepository(prisma_client).table.update( + where={"team_id": group_id}, + data=update_data, + ) + return await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id}) async def _handle_group_membership_changes(group_id: str, current_members: Set[str], final_members: Set[str]): @@ -2031,27 +2132,29 @@ async def patch_group( existing_team = await _check_team_exists(group_id) # Process patch operations - update_data, final_members = await _process_group_patch_operations(patch_ops, existing_team, prisma_client) + update_data, final_members, replace_target = await _process_group_patch_operations( + patch_ops, existing_team, prisma_client + ) - # Track current members BEFORE update for comparison - current_members = set(await _get_team_member_user_ids_from_team(existing_team)) + snapshot_members = set(await _get_team_member_user_ids_from_team(existing_team)) + intended_add = final_members - snapshot_members + intended_remove = snapshot_members - final_members - # Apply updates to the database - updated_team = await _apply_group_patch_updates(group_id, update_data, final_members, prisma_client) + # Apply the metadata/displayName updates to the database + updated_team = await _apply_group_patch_updates(group_id, update_data, prisma_client) - # Refresh team data from database to get the latest state after concurrent updates - # This prevents race conditions when multiple PATCH requests come in simultaneously refreshed_team = await TeamRepository(prisma_client).table.find_unique(where={"team_id": group_id}) - if refreshed_team: - # Re-read current members from refreshed team to account for concurrent updates - refreshed_current_members = set( - await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump())) - ) - # Use the refreshed members for comparison - current_members = refreshed_current_members + refreshed_current = ( + set(await _get_team_member_user_ids_from_team(LiteLLM_TeamTable(**refreshed_team.model_dump()))) + if refreshed_team + else snapshot_members + ) - # Handle user-team relationship changes - await _handle_group_membership_changes(group_id, current_members, final_members) + effective_final = ( + replace_target if replace_target is not None else (refreshed_current | intended_add) - intended_remove + ) + + await _handle_group_membership_changes(group_id, refreshed_current, effective_final) # A rename can flip whether this group matches scim_admin_group by display # name, so retained members must be re-resolved too, not just the ones whose @@ -2060,7 +2163,7 @@ async def patch_group( alias_changed = new_alias != existing_team.team_alias await _recompute_scim_member_roles( prisma_client, - (current_members | final_members if alias_changed else current_members ^ final_members), + (refreshed_current | effective_final if alias_changed else refreshed_current ^ effective_final), ) # Refresh team one more time to get final state after membership changes diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 70c002d2d2d..59b0cbc4ae7 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -47,6 +47,7 @@ from litellm.proxy._types import ( Member, NewTeamRequest, OrgMember, + PatchTeamRequest, ProxyErrorTypes, ProxyException, SpecialManagementEndpointEnums, @@ -1956,6 +1957,7 @@ async def update_team( ) async def patch_team( team_id: str, + data: PatchTeamRequest, http_request: Request, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], litellm_changed_by: Annotated[ @@ -1968,11 +1970,12 @@ async def patch_team( """ Partially update a team using RFC 7386 JSON Merge Patch semantics. - `team_id` is taken from the path. `metadata` is merged with the team's stored - metadata rather than replacing it: an omitted key is preserved, `key: null` - deletes it, and any other value overwrites (recursing into nested objects). - Every other field behaves exactly like `POST /team/update` (omitted preserves, - a value overwrites). Returns the full updated team. + `team_id` is taken from the path; a `team_id` in the body is accepted only when it + matches. `metadata` is merged with the team's stored metadata rather than replacing + it: an omitted key is preserved, `key: null` deletes it, and any other value + overwrites (recursing into nested objects). Every other field behaves exactly like + `POST /team/update` (omitted preserves, a value overwrites). Returns the full + updated team. ``` curl --location --request PATCH 'http://0.0.0.0:4000/team/8d916b1c-510d-4894-a334-1c16a93344f5' \ @@ -1992,21 +1995,15 @@ async def patch_team( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - try: - body = await http_request.json() - except (json.JSONDecodeError, ValueError): - raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"}) - if not isinstance(body, dict): - raise HTTPException(status_code=400, detail={"error": "Request body must be a JSON object"}) - - body_team_id = body.pop("team_id", None) - if body_team_id is not None and body_team_id != team_id: + if data.team_id is not None and data.team_id != team_id: raise HTTPException( status_code=400, - detail={"error": f"team_id in body ({body_team_id}) does not match team_id in path ({team_id})"}, + detail={"error": f"team_id in body ({data.team_id}) does not match team_id in path ({team_id})"}, ) - if "metadata" in body: + patch_fields = data.model_dump(exclude_unset=True, exclude={"team_id"}) + + if "metadata" in patch_fields: existing_team_row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) if existing_team_row is None: raise HTTPException( @@ -2014,9 +2011,9 @@ async def patch_team( detail={"error": f"Team not found, passed team_id={team_id}"}, ) existing_metadata = existing_team_row.metadata if isinstance(existing_team_row.metadata, dict) else {} - body["metadata"] = apply_json_merge_patch(existing_metadata, body["metadata"]) + patch_fields["metadata"] = apply_json_merge_patch(existing_metadata, patch_fields["metadata"]) - update_request = UpdateTeamRequest(team_id=team_id, **body) + update_request = UpdateTeamRequest(team_id=team_id, **patch_fields) result = await update_team( data=update_request, @@ -2375,7 +2372,15 @@ async def _add_team_members_to_team( user_api_key_dict: UserAPIKeyAuth, litellm_proxy_admin_name: str, ) -> Tuple[LiteLLM_TeamTable, List[LiteLLM_UserTable], List[LiteLLM_TeamMembership]]: - """Add team members to the team.""" + """Add team members to the team. + + The members_with_roles reconciliation runs inside a transaction that locks + the team row with ``SELECT ... FOR UPDATE`` before reading the current + membership. Concurrent /team/member_add calls for the same team therefore + serialize on the row lock and each appends onto the other's committed + result, instead of both rewriting the whole JSON array from a stale + snapshot (which silently drops one member on the losing write). + """ # Process and add new members updated_users, updated_team_memberships = await _process_team_members( data=data, @@ -2385,19 +2390,22 @@ async def _add_team_members_to_team( litellm_proxy_admin_name=litellm_proxy_admin_name, ) - # Update team members list - await _update_team_members_list( - data=data, - complete_team_data=complete_team_data, - updated_users=updated_users, - ) + async with prisma_client.tx() as tx: + complete_team_data.members_with_roles = await TeamRepository(prisma_client).get_members_with_roles_locked( + tx, data.team_id + ) - # ADD MEMBER TO TEAM - _db_team_members = [m.model_dump() for m in complete_team_data.members_with_roles] - updated_team = await TeamRepository(prisma_client).table.update( - where={"team_id": data.team_id}, - data={"members_with_roles": json.dumps(_db_team_members)}, # type: ignore - ) + await _update_team_members_list( + data=data, + complete_team_data=complete_team_data, + updated_users=updated_users, + ) + + _db_team_members = [m.model_dump() for m in complete_team_data.members_with_roles] + updated_team = await tx.litellm_teamtable.update( + where={"team_id": data.team_id}, + data={"members_with_roles": json.dumps(_db_team_members)}, + ) return updated_team, updated_users, updated_team_memberships diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 6c2e06a418c..de988a0140f 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -12,6 +12,7 @@ import asyncio import base64 import hashlib import inspect +import json import os import re import secrets @@ -62,6 +63,11 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) +from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import ( + SSOIdentityAssertion, + assertion_from_sso_login, + retain_sso_identity_assertion_for_ema, +) from litellm.proxy._types import ( CommonProxyErrors, LiteLLM_UserTable, @@ -253,11 +259,20 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic raise HTTPException(status_code=400, detail="Invalid CLI login session id") cache_key = _get_cli_sso_flow_cache_key(cast(str, login_id)) - flow = cache.get_cache(key=cache_key) + redis_cache = cache.redis_cache + if redis_cache is not None: + flow = redis_cache.get_cache(key=cache_key) + else: + flow = cache.get_cache(key=cache_key) + if isinstance(flow, str): + try: + flow = json.loads(flow) + except ValueError: + flow = None if not isinstance(flow, dict) or "poll_secret_hash" not in flow: verbose_proxy_logger.warning( "CLI SSO login session not found in cache for login_id=%s. If the proxy runs multiple replicas, " - "a shared Redis cache (enable_redis_auth_cache: true) is required for CLI login to work.", + "a shared Redis cache is required for CLI login to work.", login_id, ) raise HTTPException( @@ -265,7 +280,7 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic detail=( "CLI login session not found or expired. Run `litellm-proxy login` again. " "If this happens immediately after starting a login, the proxy is likely running multiple " - "replicas without a shared cache; configure Redis with `enable_redis_auth_cache: true` " + "replicas without a shared cache; configure a Redis cache " "so every replica can see the login session." ), ) @@ -273,11 +288,12 @@ def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dic def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None: - cache.set_cache( - key=_get_cli_sso_flow_cache_key(login_id), - value=flow, - ttl=CLI_SSO_SESSION_TTL_SECONDS, - ) + cache_key = _get_cli_sso_flow_cache_key(login_id) + redis_cache = cache.redis_cache + if redis_cache is not None: + redis_cache.set_cache(key=cache_key, value=json.dumps(flow), ttl=CLI_SSO_SESSION_TTL_SECONDS) + else: + cache.set_cache(key=cache_key, value=flow, ttl=CLI_SSO_SESSION_TTL_SECONDS) def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool: @@ -588,11 +604,11 @@ def _render_cli_sso_verification_page( @router.post("/sso/cli/start", tags=["experimental"], include_in_schema=False) async def cli_sso_start(request: Request): - from litellm.proxy.proxy_server import general_settings, user_api_key_cache + from litellm.proxy.proxy_server import cli_sso_session_cache, general_settings _check_cli_sso_start_rate_limit( request=request, - cache=user_api_key_cache, + cache=cli_sso_session_cache, use_x_forwarded_for=bool((general_settings or {}).get("use_x_forwarded_for", False)), ) @@ -607,7 +623,7 @@ async def cli_sso_start(request: Request): "user_code_verified": False, "session_data": None, } - _set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow) + _set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow) verification_uri_complete: str | None = ( ( @@ -639,9 +655,9 @@ async def cli_sso_complete(request: Request, login_id: str): from litellm.proxy.common_utils.html_forms.cli_sso_success import ( render_cli_sso_success_page, ) - from litellm.proxy.proxy_server import user_api_key_cache + from litellm.proxy.proxy_server import cli_sso_session_cache - flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=user_api_key_cache) + flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cli_sso_session_cache) if not flow.get("sso_complete") or not flow.get("session_data"): raise HTTPException(status_code=400, detail="CLI login is not ready") @@ -665,7 +681,7 @@ async def cli_sso_complete(request: Request, login_id: str): raise HTTPException(status_code=400, detail="Invalid verification code") flow["user_code_verified"] = True - _set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow) + _set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow) html_content = render_cli_sso_success_page() return HTMLResponse(content=html_content, status_code=200) @@ -856,10 +872,10 @@ async def google_login( Example: """ from litellm.proxy.proxy_server import ( + cli_sso_session_cache, general_settings, premium_user, prisma_client, - user_api_key_cache, user_custom_ui_sso_sign_in_handler, ) @@ -907,7 +923,7 @@ async def google_login( ) if source == LITELLM_CLI_SOURCE_IDENTIFIER: - _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache) + _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache) # Store CLI login handle in state for OAuth flow cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state( @@ -1311,12 +1327,15 @@ async def get_generic_sso_response( sso_jwt_handler: Optional[JWTHandler], # sso specific jwt handler - used for restricted sso group access control generic_client_id: str, redirect_url: str, -) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload) +) -> tuple[ + Union[OpenID, dict], dict | None, dict | None, SSOIdentityAssertion | None +]: # (result, received_response, access_token_payload, sso_assertion) # make generic sso provider from fastapi_sso.sso.base import DiscoveryDocument from fastapi_sso.sso.generic import create_provider received_response: Optional[dict] = None + sso_assertion: SSOIdentityAssertion | None = None # Setup environment variables ( @@ -1450,6 +1469,9 @@ async def get_generic_sso_response( # Assign directly rather than relying on nonlocal mutation so that Pyright # can track that received_response is non-None from this point on. received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS} + sso_assertion = assertion_from_sso_login( + combined_response.get("id_token"), combined_response.get("refresh_token") + ) # In the PKCE path verify_and_process is skipped, so generic_sso.access_token # is never set. Read the token directly from the exchange response instead so # process_sso_jwt_access_token can extract JWT-embedded roles/teams. @@ -1461,6 +1483,7 @@ async def get_generic_sso_response( headers=additional_generic_sso_headers_dict, ) access_token_str = generic_sso.access_token + sso_assertion = assertion_from_sso_login(generic_sso.id_token, generic_sso.refresh_token) access_token_payload = process_sso_jwt_access_token( access_token_str, sso_jwt_handler, result, role_mappings=role_mappings @@ -1480,7 +1503,7 @@ async def get_generic_sso_response( additional_generic_sso_headers_dict, ) verbose_proxy_logger.debug("generic result: %s", result) - return result or {}, received_response, access_token_payload + return result or {}, received_response, access_token_payload, sso_assertion async def create_team_member_add_task(team_id, user_info): @@ -1812,6 +1835,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): generic_client_id = os.getenv("GENERIC_CLIENT_ID", None) received_response: Optional[dict] = None access_token_payload: Optional[dict] = None + sso_assertion: SSOIdentityAssertion | None = None # get url from request if master_key is None: raise ProxyException( @@ -1842,6 +1866,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): result, received_response, access_token_payload, + sso_assertion, ) = await get_generic_sso_response( request=request, jwt_handler=jwt_handler, @@ -1869,6 +1894,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): prefill_user_code=prefill_user_code, result=result, received_response=received_response, + sso_assertion=sso_assertion, ) # Control-plane cross-origin: read return_to from cookie. @@ -1884,6 +1910,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): access_token_payload=access_token_payload, jwt_handler=jwt_handler, return_to=cp_return_to, + sso_assertion=sso_assertion, ) @@ -1941,8 +1968,10 @@ async def _complete_cli_sso_callback_session( user_defined_values: Optional[SSOUserDefinedValues], prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, + cli_sso_session_cache: DualCache, proxy_logging_obj: ProxyLogging, prefill_user_code: str | None = None, + sso_assertion: SSOIdentityAssertion | None = None, ): from fastapi.responses import HTMLResponse @@ -1962,6 +1991,8 @@ async def _complete_cli_sso_callback_session( if not user_info.user_id: raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO") + await retain_sso_identity_assertion_for_ema(user_id=user_info.user_id, assertion=sso_assertion) + teams: List[str] = [] if hasattr(user_info, "teams") and user_info.teams: teams = user_info.teams if isinstance(user_info.teams, list) else [] @@ -1987,7 +2018,7 @@ async def _complete_cli_sso_callback_session( flow["sso_complete"] = True browser_complete_token = secrets.token_urlsafe(32) flow["browser_complete_token_hash"] = _hash_cli_sso_secret(browser_complete_token) - _set_cli_sso_flow(login_id=key, cache=user_api_key_cache, flow=flow) + _set_cli_sso_flow(login_id=key, cache=cli_sso_session_cache, flow=flow) verbose_proxy_logger.info( f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}" @@ -2012,18 +2043,20 @@ async def cli_sso_callback( result: Optional[Union[OpenID, dict]] = None, received_response: Optional[dict] = None, prefill_user_code: str | None = None, + sso_assertion: SSOIdentityAssertion | None = None, ): """CLI SSO callback - stores session info for JWT generation on polling""" verbose_proxy_logger.info("CLI SSO callback") from litellm.proxy.proxy_server import ( + cli_sso_session_cache, general_settings, prisma_client, proxy_logging_obj, user_api_key_cache, ) - flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache) + flow = _get_cli_sso_flow_or_raise(login_id=key, cache=cli_sso_session_cache) if prisma_client is None: raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) @@ -2063,8 +2096,10 @@ async def cli_sso_callback( user_defined_values=user_defined_values, prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, + cli_sso_session_cache=cli_sso_session_cache, proxy_logging_obj=proxy_logging_obj, prefill_user_code=prefill_user_code, + sso_assertion=sso_assertion, ) except ProxyException: raise @@ -2093,10 +2128,10 @@ async def cli_poll_key( team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams. """ from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken - from litellm.proxy.proxy_server import user_api_key_cache + from litellm.proxy.proxy_server import cli_sso_session_cache try: - flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache) + flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=cli_sso_session_cache) if not _verify_cli_sso_poll_secret(flow=flow, poll_secret=x_litellm_cli_poll_secret): raise HTTPException(status_code=403, detail="Invalid CLI polling secret") @@ -2171,7 +2206,7 @@ async def cli_poll_key( ) # Delete cache entry (single-use) - user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id)) + cli_sso_session_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id)) verbose_proxy_logger.info(f"CLI JWT generated for user: {user_id}, team: {team_id}") poll_response = { @@ -3018,6 +3053,7 @@ class SSOAuthenticationHandler: access_token_payload: Optional[dict] = None, jwt_handler: Optional[JWTHandler] = None, return_to: Optional[str] = None, + sso_assertion: SSOIdentityAssertion | None = None, ) -> RedirectResponse: import jwt @@ -3148,6 +3184,9 @@ class SSOAuthenticationHandler: }, ) + if isinstance(user_id, str) and user_id: + await retain_sso_identity_assertion_for_ema(user_id=user_id, assertion=sso_assertion) + disabled_non_admin_personal_key_creation = get_disabled_non_admin_personal_key_creation() litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/") @@ -4241,6 +4280,7 @@ async def debug_sso_callback(request: Request): result, received_response, access_token_payload, + _sso_assertion, ) = await get_generic_sso_response( request=request, jwt_handler=jwt_handler, diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 6bbd41b93ed..015ab1d6df6 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -4,7 +4,8 @@ organizations, teams, and keys. """ import json -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Union +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Dict, List, Mapping, Optional, Set, Union from fastapi import HTTPException, status @@ -64,6 +65,57 @@ async def attach_object_permission_to_dict( return data_dict +@dataclass(frozen=True, slots=True) +class ObjectPermissionUpsert: + object_permission_id: str + record: dict[str, object] + + +async def prepare_object_permission_upsert( + new_object_permission: Mapping[str, object], + existing_object_permission_id: str | None, + prisma_client: PrismaClient, +) -> ObjectPermissionUpsert: + """ + Read-and-merge half of an object permission upsert; performs no writes. + + Merges the sent grants over the existing row (looked up by + ``existing_object_permission_id``, or a fresh uuid when the entity has none) and + returns the id plus the full record to upsert. The id is pinned inside the record + because the column has ``@default(uuid())``, so a create without it would mint a + different id than the one the caller links. ``mcp_tool_permissions`` is serialized + to a JSON string to avoid GraphQL parsing issues (e.g. server IDs starting with + "3e64" being interpreted as floats). + + Keeping this separate from the write lets callers run the upsert inside the same + transaction as the row that links ``object_permission_id``, so a rolled-back + update cannot leave permission changes live. + """ + object_permission_id = existing_object_permission_id or str(uuid.uuid4()) + existing_object_permission = await ObjectPermissionRepository(prisma_client).table.find_unique( + where={"object_permission_id": object_permission_id}, + ) + existing_fields: dict[str, object] = ( + existing_object_permission.model_dump(exclude_unset=True, exclude_none=True) + if existing_object_permission is not None + else {} + ) + merged: dict[str, object] = { + **existing_fields, + **new_object_permission, + "object_permission_id": object_permission_id, + } + record: dict[str, object] = { + **merged, + **( + {"mcp_tool_permissions": safe_dumps(merged["mcp_tool_permissions"])} + if "mcp_tool_permissions" in merged + else {} + ), + } + return ObjectPermissionUpsert(object_permission_id=object_permission_id, record=record) + + async def handle_update_object_permission_common( data_json: Dict, existing_object_permission_id: Optional[str], @@ -93,50 +145,23 @@ async def handle_update_object_permission_common( if prisma_client is None: raise ValueError("Prisma client not found") - ######################################################### - # Ensure `object_permission` is not added to the data_json - # We need to update the entity at the object_permission_id level in the LiteLLM_ObjectPermissionTable - ######################################################### - new_object_permission: Union[dict, str] = data_json.pop("object_permission", None) + new_object_permission: Union[dict, str, None] = data_json.pop("object_permission", None) if new_object_permission is None: return None - # Lookup existing object permission ID and update that entry - object_permission_id_to_use: str = existing_object_permission_id or str(uuid.uuid4()) - existing_object_permissions_dict: Dict = {} - - existing_object_permission = await ObjectPermissionRepository(prisma_client).table.find_unique( - where={"object_permission_id": object_permission_id_to_use}, - ) - - # Update the object permission - if existing_object_permission is not None: - existing_object_permissions_dict = existing_object_permission.model_dump(exclude_unset=True, exclude_none=True) - - # Handle string JSON object permission if isinstance(new_object_permission, str): new_object_permission = json.loads(new_object_permission) - if isinstance(new_object_permission, dict): - existing_object_permissions_dict.update(new_object_permission) - - ######################################################### - # Serialize mcp_tool_permissions JSON field to avoid GraphQL parsing issues - # (e.g., server IDs starting with "3e64" being interpreted as floats) - ######################################################### - if "mcp_tool_permissions" in existing_object_permissions_dict: - existing_object_permissions_dict["mcp_tool_permissions"] = safe_dumps( - existing_object_permissions_dict["mcp_tool_permissions"] - ) - - ######################################################### - # Commit the update to the LiteLLM_ObjectPermissionTable - ######################################################### + upsert = await prepare_object_permission_upsert( + new_object_permission=new_object_permission if isinstance(new_object_permission, dict) else {}, + existing_object_permission_id=existing_object_permission_id, + prisma_client=prisma_client, + ) created_object_permission_row = await ObjectPermissionRepository(prisma_client).table.upsert( - where={"object_permission_id": object_permission_id_to_use}, + where={"object_permission_id": upsert.object_permission_id}, data={ - "create": existing_object_permissions_dict, - "update": existing_object_permissions_dict, + "create": upsert.record, + "update": upsert.record, }, ) diff --git a/litellm/proxy/management_helpers/utils.py b/litellm/proxy/management_helpers/utils.py index 11a99caebf5..86beb063667 100644 --- a/litellm/proxy/management_helpers/utils.py +++ b/litellm/proxy/management_helpers/utils.py @@ -252,6 +252,21 @@ async def _resolve_member_budget_id( return response.budget_id +async def _append_team_id_if_absent(prisma_client: PrismaClient, user_id: str, team_id: str) -> None: + """Append team_id to a user's teams array, only if it is not already present. + + The row-level filter makes the append a no-op once the team is present, so + repeated or concurrent adds of the same team cannot accumulate duplicate + team ids in user.teams (a duplicate also breaks auth logic that keys off the + number of teams a user belongs to). Teams added concurrently for a different + team id are unaffected, since each update filters on its own team id. + """ + await UserRepository(prisma_client).table.update_many( + where={"user_id": user_id, "NOT": {"teams": {"has": team_id}}}, + data={"teams": {"push": [team_id]}}, + ) + + async def add_new_member( new_member: Member, max_budget_in_team: Optional[float], @@ -276,13 +291,16 @@ async def add_new_member( ## ADD TEAM ID, to USER TABLE IF NEW ## if new_member.user_id is not None: new_user_defaults = get_new_internal_user_defaults(user_id=new_member.user_id) + # Upsert ensures the user row exists atomically (no create race when the + # same new user is provisioned concurrently), seeding teams on create. + # The teams append lives in the filtered update below rather than the + # upsert's update branch so an already-existing user does not get a + # duplicate team id. _returned_user = await UserRepository(prisma_client).table.upsert( where={"user_id": new_member.user_id}, - data={ - "update": {"teams": {"push": [team_id]}}, - "create": {"teams": [team_id], **new_user_defaults}, # type: ignore - }, + data={"create": {"teams": [team_id], **new_user_defaults}, "update": {}}, ) + await _append_team_id_if_absent(prisma_client, new_member.user_id, team_id) if _returned_user is not None: returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) elif new_member.user_email is not None: @@ -302,12 +320,8 @@ async def add_new_member( returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) elif len(existing_user_row) == 1: user_info = existing_user_row[0] - _returned_user = await UserRepository(prisma_client).table.update( - where={"user_id": user_info.user_id}, # type: ignore - data={"teams": {"push": [team_id]}}, - ) - if _returned_user is not None: - returned_user = LiteLLM_UserTable(**_returned_user.model_dump()) + await _append_team_id_if_absent(prisma_client, user_info.user_id, team_id) + returned_user = LiteLLM_UserTable(**user_info.model_dump()) elif len(existing_user_row) > 1: raise HTTPException( status_code=400, diff --git a/litellm/proxy/mcp_registry.json b/litellm/proxy/mcp_registry.json index 84431634e24..f37fc39813e 100644 --- a/litellm/proxy/mcp_registry.json +++ b/litellm/proxy/mcp_registry.json @@ -198,13 +198,9 @@ "icon_url": "https://cdn.simpleicons.org/googledrive", "category": "Productivity", "registry_url": null, - "transport": "stdio", - "command": "npx", - "args": ["-y", "@modelcontextprotocol/server-gdrive"], - "env_vars": [ - {"name": "GOOGLE_CLIENT_ID", "description": "Google OAuth Client ID", "secret": false}, - {"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true} - ] + "transport": "http", + "url": "https://drivemcp.googleapis.com/mcp/v1", + "env_vars": [] }, { "name": "google_calendar", diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3b40abed19e..32845763f22 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -226,6 +226,7 @@ from litellm.constants import ( APSCHEDULER_MAX_INSTANCES, APSCHEDULER_MISFIRE_GRACE_TIME, APSCHEDULER_REPLACE_EXISTING, + CLI_SSO_SESSION_TTL_SECONDS, DAYS_IN_A_MONTH, DEFAULT_HEALTH_CHECK_INTERVAL, DEFAULT_MODEL_CREATED_AT_TIME, @@ -236,6 +237,7 @@ from litellm.constants import ( PROXY_BATCH_WRITE_AT, PROXY_BUDGET_RESCHEDULER_MAX_TIME, PROXY_BUDGET_RESCHEDULER_MIN_TIME, + PROXY_CONFIG_RELOAD_INTERVAL_SECONDS, ) from litellm.exceptions import RejectedRequestError from litellm.integrations.custom_guardrail import ModifyResponseException @@ -302,6 +304,11 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form +from litellm.proxy.config_resolvers import resolve_fields +from litellm.proxy.config_resolvers.alerting import ( + EMAIL_DESCRIPTORS, + SLACK_DESCRIPTORS, +) from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, @@ -319,7 +326,10 @@ from litellm.proxy.common_utils.openai_endpoint_utils import ( from litellm.proxy.common_utils.proxy_state import ProxyState from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob from litellm.proxy.common_utils.swagger_utils import ERROR_RESPONSES -from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time +from litellm.proxy.common_utils.timezone_utils import ( + get_budget_reset_settings, + get_budget_reset_time, +) from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, get_management_object_ttl, @@ -1154,9 +1164,9 @@ _OPENAPI_HTTP_METHODS = { # Credentials surfaced by `/get/config/callbacks` in the alerting block: the # full Slack incoming-webhook URL is itself a credential, and the SMTP # password is a service password. Masked on read so plaintext never reaches -# the UI. Kept here at module scope to match the analogous -# `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO -# and cache endpoint files. +# the UI. Kept here at module scope to match the analogous descriptor +# `is_secret` flags in litellm.proxy.config_resolvers and the +# `_CACHE_SENSITIVE_FIELDS` constant in the cache endpoint file. _ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"} @@ -1966,6 +1976,7 @@ user_api_key_cache: UserApiKeyCache = UserApiKeyCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) spend_counter_cache = DualCache(default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value) +cli_sso_session_cache = DualCache(default_in_memory_ttl=CLI_SSO_SESSION_TTL_SECONDS) model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter(dual_cache=user_api_key_cache) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) redis_usage_cache: Optional[RedisCache] = None # redis cache used for tracking spend, tpm/rpm limits @@ -1998,6 +2009,7 @@ proxy_budget_rescheduler_min_time = PROXY_BUDGET_RESCHEDULER_MIN_TIME proxy_budget_rescheduler_max_time = PROXY_BUDGET_RESCHEDULER_MAX_TIME proxy_batch_polling_interval = PROXY_BATCH_POLLING_INTERVAL proxy_batch_write_at = PROXY_BATCH_WRITE_AT +proxy_config_reload_interval_seconds = PROXY_CONFIG_RELOAD_INTERVAL_SECONDS litellm_master_key_hash = None disable_spend_logs = False jwt_handler = JWTHandler() @@ -3691,13 +3703,22 @@ def _build_redis_usage_cache_from_environment() -> RedisCache | None: def _attach_redis_usage_cache(redis_cache: RedisCache, enable_redis_auth_cache: bool) -> None: """ Wires an established coordination Redis into the proxy-level caches that - consume it directly: the spend counter cache, the cluster-wide config - cache, and (only when opted in) the virtual-key auth cache. + consume it directly: the spend counter cache, the CLI SSO login-session + cache, the cluster-wide config cache, and (only when opted in) the + virtual-key auth cache. + + The CLI SSO login-session cache is always backed by Redis when available so + that the browser SSO flow behind `lite login` survives landing on different + workers; it must not be gated behind enable_redis_auth_cache. """ spend_counter_cache.attach_redis_cache( redis_cache, default_redis_ttl=litellm.default_redis_ttl, ) + cli_sso_session_cache.attach_redis_cache( + redis_cache, + default_redis_ttl=CLI_SSO_SESSION_TTL_SECONDS, + ) if enable_redis_auth_cache is True: user_api_key_cache.attach_redis_cache( redis_cache, @@ -3879,7 +3900,7 @@ class ProxyConfig: del config["include"] return config - async def save_config(self, new_config: dict): + async def save_config(self, new_config: dict, include_env_vars: bool = False): global prisma_client, general_settings, user_config_file_path, store_model_in_db # Load existing config ## DB - writes valid config to db @@ -3896,6 +3917,17 @@ class ProxyConfig: # Make a copy to avoid mutating the original config config_to_save = new_config.copy() + # environment_variables are persisted to the DB only when a caller + # explicitly opts in. Most callers reach save_config after + # get_config() merged YAML + OS env into new_config (with + # os.environ/ placeholders already resolved to plaintext), so + # persisting them here would snapshot file/container env vars into + # a config row that then shadows those sources on every restart. + # The dedicated /config/update path writes env vars directly, so + # no current caller needs include_env_vars=True. + if not include_env_vars: + config_to_save.pop("environment_variables", None) + # SECURITY: Always encrypt environment_variables before DB write. # _encrypt_env_variables_for_db is idempotent — a caller that # already encrypted the values (or re-submitted ciphertext read @@ -3913,6 +3945,38 @@ class ProxyConfig: with open(f"{user_config_file_path}", "w") as config_file: yaml.dump(new_config, config_file, default_flow_style=False) + async def save_environment_variables(self, updates: dict[str, str | None]) -> None: + """Persist specific environment variables to the DB config row. + + Each key in ``updates`` is written to the ``environment_variables`` + config row; a ``None`` value deletes that key. Env vars the caller does + not name are preserved, so a caller that owns a couple of keys can + update just those without snapshotting unrelated (YAML/OS-sourced) + values the way a full ``save_config`` write would. No-op when config is + not DB-backed. + """ + global prisma_client, general_settings, store_model_in_db + if prisma_client is None or not (general_settings.get("store_model_in_db", False) is True or store_model_in_db): + return + + row = await ConfigRepository(prisma_client).table.find_first(where={"param_name": "environment_variables"}) + existing: dict = dict(row.param_value) if row is not None and row.param_value is not None else {} + + to_set = {k: v for k, v in updates.items() if v is not None} + encrypted = self._encrypt_env_variables_for_db(environment_variables=to_set) if to_set else {} + deleted_keys = {k for k, v in updates.items() if v is None} + merged = {**{k: v for k, v in existing.items() if k not in deleted_keys}, **encrypted} + + serialized = json.dumps(merged) + await ConfigRepository(prisma_client).table.upsert( + where={"param_name": "environment_variables"}, + data={ + "create": {"param_name": "environment_variables", "param_value": serialized}, + "update": {"param_value": serialized}, + }, + ) + await invalidate_config_param("environment_variables") + def _check_for_os_environ_vars( self, config: dict, depth: int = 0, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH ) -> dict: @@ -4291,6 +4355,7 @@ class ProxyConfig: open_telemetry_logger, \ health_check_details, \ proxy_batch_polling_interval, \ + proxy_config_reload_interval_seconds, \ config_passthrough_endpoints config: dict = await self.get_config(config_file_path=config_file_path) @@ -4569,6 +4634,11 @@ class ProxyConfig: verbose_proxy_logger.debug( f"{blue_color_code} Initialized polling via cache: enabled={polling_via_cache_enabled}, native_background_mode={native_background_mode}, ttl={polling_cache_ttl}{reset_color_code}" ) + elif key == "max_ui_session_budget": + litellm.max_ui_session_budget = float(value) if value is not None else None + verbose_proxy_logger.debug( + f"{blue_color_code} setting litellm.max_ui_session_budget={litellm.max_ui_session_budget}{reset_color_code}" + ) elif key == "default_team_settings": for idx, team_setting in enumerate(value): # run through pydantic validation try: @@ -4597,6 +4667,13 @@ class ProxyConfig: litellm.json_logs = True litellm._turn_on_json() verbose_proxy_logger.debug(f"{blue_color_code} Enabled JSON logging via config{reset_color_code}") + elif key == "budget_reset_time": + from litellm.proxy.common_utils.timezone_utils import ( + parse_budget_reset_time, + ) + + parse_budget_reset_time(value) + setattr(litellm, key, value) else: verbose_proxy_logger.debug( f"{blue_color_code} setting litellm.{key}={_redact_general_setting_value(key, value, is_full_admin=False)}{reset_color_code}" @@ -4773,6 +4850,10 @@ class ProxyConfig: ) ## BATCH WRITER ## proxy_batch_write_at = general_settings.get("proxy_batch_write_at", proxy_batch_write_at) + ## DB CONFIG RELOAD INTERVAL ## + proxy_config_reload_interval_seconds = general_settings.get( + "proxy_config_reload_interval_seconds", proxy_config_reload_interval_seconds + ) ## DISABLE SPEND LOGS ## - gives a perf improvement disable_spend_logs = general_settings.get("disable_spend_logs", disable_spend_logs) ### BACKGROUND HEALTH CHECKS ### @@ -7868,6 +7949,7 @@ class ProxyStartupEvent: budget_reset_job = ResetBudgetJob( proxy_logging_obj=proxy_logging_obj, prisma_client=prisma_client, + reset_settings=get_budget_reset_settings(), ) scheduler.add_job( @@ -7944,12 +8026,20 @@ class ProxyStartupEvent: verbose_proxy_logger.debug("Failed to check DB for store_model_in_db: %s", str(e)) if store_model_in_db is True: + config_reload_interval_seconds = proxy_config_reload_interval_seconds + if not isinstance(config_reload_interval_seconds, int) or config_reload_interval_seconds <= 0: + verbose_proxy_logger.warning( + "proxy_config_reload_interval_seconds=%s must be a positive integer; falling back to 30s", + config_reload_interval_seconds, + ) + config_reload_interval_seconds = 30 + # MEMORY LEAK FIX: Increase interval from 10s to 30s minimum # Frequent polling was causing excessive memory allocations scheduler.add_job( proxy_config.add_deployment, "interval", - seconds=30, # increased from 10s to reduce memory pressure + seconds=config_reload_interval_seconds, # REMOVED jitter parameter - major cause of memory leak args=[prisma_client, proxy_logging_obj], id="add_deployment_job", @@ -7964,7 +8054,7 @@ class ProxyStartupEvent: scheduler.add_job( proxy_config.get_credentials, "interval", - seconds=30, # increased from 10s to reduce memory pressure + seconds=config_reload_interval_seconds, # REMOVED jitter parameter - major cause of memory leak args=[prisma_client], id="get_credentials_job", @@ -14845,10 +14935,11 @@ GeneralSettingsUILiteLLMValue = Union[float, bool, str, None] class GeneralSettingsUILiteLLMFieldSpec(TypedDict): - type: Literal["Float", "Boolean", "Select"] + type: Literal["Float", "Dollar", "Boolean", "Select"] description: str options: NotRequired[tuple[str, ...]] tab: NotRequired[str] # Admin UI sub-tab this field renders under; None groups it with the rest + default: NotRequired[float] # reset/clear restores this instead of None; fields whose None means fail-open set it _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec] = { @@ -14874,21 +14965,32 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec "tab": "prompt_caching", "description": "Empty uses Anthropic's 5m default. 1h suits long sessions but doubles the cache write cost.", }, + "max_ui_session_budget": { + "type": "Dollar", + "default": 1.0, + "description": ( + "USD spend cap for each dashboard login session; covers LLM calls made from the dashboard " + "such as the playground and auto router Test Connection. Each login starts a fresh session " + "with this budget. Clearing restores the $1 default." + ), + }, } def _general_settings_ui_litellm_default( - field_type: Literal["Float", "Boolean", "Select"], + spec: GeneralSettingsUILiteLLMFieldSpec, ) -> GeneralSettingsUILiteLLMValue: """The value a field falls back to when it is cleared or reset.""" - return False if field_type == "Boolean" else None + if "default" in spec: + return spec["default"] + return False if spec["type"] == "Boolean" else None def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> GeneralSettingsUILiteLLMValue: spec = _GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name] field_type = spec["type"] if value is None or value == "": - return _general_settings_ui_litellm_default(field_type) + return _general_settings_ui_litellm_default(spec) match field_type: case "Boolean": if not isinstance(value, bool): @@ -14912,6 +15014,13 @@ def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> detail={"error": f"{field_name} must be a number in (0, 1] or empty"}, ) return float(value) + case "Dollar": + if isinstance(value, bool) or not isinstance(value, (int, float)) or float(value) <= 0: + raise HTTPException( + status_code=400, + detail={"error": f"{field_name} must be a positive dollar amount or empty"}, + ) + return float(value) case _: assert_never(field_type) @@ -14934,7 +15043,7 @@ async def _persist_general_settings_ui_litellm_field( async def _reset_general_settings_ui_litellm_field(field_name: str, user_api_key_dict: UserAPIKeyAuth) -> dict: config = await proxy_config.get_config() before_value = config.get("litellm_settings", {}).get(field_name) - default_value = _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]["type"]) + default_value = _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]) setattr(litellm, field_name, default_value) if "litellm_settings" in config: config["litellm_settings"].pop(field_name, None) @@ -14998,6 +15107,7 @@ async def get_config_list( "global_max_parallel_requests": {"type": "Integer"}, "max_request_size_mb": {"type": "Integer"}, "max_response_size_mb": {"type": "Integer"}, + "proxy_config_reload_interval_seconds": {"type": "Integer"}, "pass_through_endpoints": {"type": "PydanticModel"}, "store_model_in_db": {"type": "Boolean"}, "store_prompts_in_spend_logs": {"type": "Boolean"}, @@ -15108,7 +15218,7 @@ async def get_config_list( ) for litellm_field_name, spec in _GENERAL_SETTINGS_UI_LITELLM_FIELDS.items(): current_value: GeneralSettingsUILiteLLMValue = getattr(litellm, litellm_field_name, None) - default_value = _general_settings_ui_litellm_default(spec["type"]) + default_value = _general_settings_ui_litellm_default(spec) stored_in_db_litellm: Optional[bool] if litellm_field_name in db_litellm_settings: stored_in_db_litellm = True @@ -15386,14 +15496,10 @@ async def get_config( _alerting = _general_settings.get("alerting", []) alerting_data = [] if "slack" in _alerting: - _slack_vars = [ - "SLACK_WEBHOOK_URL", - ] - _slack_env_vars = { - _var: (value if (value := environment_variables.get(_var)) is not None else os.getenv(_var)) - for _var in _slack_vars - } - _slack_env_vars = _apply_alerting_env_role_gate(_slack_env_vars, is_full_admin) + _slack_values, _ = resolve_fields( + SLACK_DESCRIPTORS, environment_variables, os.environ, empty_db_is_set=True + ) + _slack_env_vars = _apply_alerting_env_role_gate(_slack_values, is_full_admin) _alerting_types = proxy_logging_obj.slack_alerting_instance.alert_types _all_alert_types = proxy_logging_obj.slack_alerting_instance._all_possible_alert_types() @@ -15409,19 +15515,8 @@ async def get_config( } ) # pass email alerting vars - _email_vars = [ - "SMTP_HOST", - "SMTP_PORT", - "SMTP_USERNAME", - "SMTP_PASSWORD", - "SMTP_SENDER_EMAIL", - "TEST_EMAIL_ADDRESS", - "EMAIL_LOGO_URL", - "EMAIL_SUPPORT_CONTACT", - ] - _email_env_vars = _apply_alerting_env_role_gate( - {_var: environment_variables.get(_var) for _var in _email_vars}, is_full_admin - ) + _email_values, _ = resolve_fields(EMAIL_DESCRIPTORS, environment_variables, os.environ, empty_db_is_set=True) + _email_env_vars = _apply_alerting_env_role_gate(_email_values, is_full_admin) alerting_data.append( { diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index b27ddea010b..23a9c086c73 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient { updated_at DateTime @default(now()) @updatedAt @map("updated_at") } +// The enterprise IdP identity assertion captured at SSO login, one row per user. +// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}. +model LiteLLM_SSOIdentityAssertion { + user_id String @id + assertion_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") +} + // Generate Tokens for Proxy model LiteLLM_VerificationToken { token String @id diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index a8926d26047..10c71c00110 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -1,6 +1,8 @@ #### CRUD ENDPOINTS for UI Settings ##### import asyncio import json +import os +from collections.abc import Mapping from typing import Any, Dict, List, Optional, Set, Tuple, Type, Union from urllib.parse import urlparse @@ -13,6 +15,11 @@ from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.config_resolvers.sso import ( + SSO_FIELD_ENV_VARS, + SSO_SECRET_FIELDS, + resolve_sso_config, +) from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.table_repositories import ( SSOConfigRepository, @@ -25,17 +32,45 @@ from litellm.types.proxy.management_endpoints.ui_sso import ( router = APIRouter() -# SSO secret fields returned by /get/sso_settings. These are masked on read so -# the UI can show "(set)" without ever transporting the plaintext OAuth secret -# off the server, matching the write-once + masked-on-read contract used for -# the HashiCorp Vault config override. -_SSO_SENSITIVE_FIELDS: Set[str] = { - "google_client_secret", - "microsoft_client_secret", - "generic_client_secret", +# Maps each UIThemeConfig field to the env var the UI branding path reads it +# from. /update/ui_theme_settings writes both the stored ui_theme_config and +# these env vars, so /get/ui_theme_settings resolves the same env vars to +# reflect a deployment branded purely through process env. +_UI_THEME_FIELD_ENV_VARS: dict[str, str] = { + "logo_url": "UI_LOGO_PATH", + "favicon_url": "LITELLM_FAVICON_URL", } +def _is_public_http_url(value: str | None) -> bool: + """Whether a value is a plain http(s) URL with a host, safe to disclose publicly.""" + if not isinstance(value, str) or not value.strip(): + return False + parsed = urlparse(value.strip()) + return parsed.scheme in ("http", "https") and bool(parsed.netloc) + + +def _resolve_ui_theme_field(stored_values: Mapping[str, Any], field_name: str) -> str | None: + """Resolve one UI theme field to the value the branding path actually uses. + + The stored ui_theme_config wins; a field absent or blank there falls back to + the process environment. The branding path reads the env var, and stored + settings reach it by being pushed into the environment on save, so a value + supplied only as a process env var is live even though no stored entry exists. + + This endpoint is unauthenticated, so the env fallback only surfaces a public + http(s) URL: an operator can point UI_LOGO_PATH at a local filesystem path + (the branding path serves it server-side), and that path must not be + disclosed to anonymous callers. A stored value is already validated as a + public URL on write, so it passes through. + """ + stored = stored_values.get(field_name) + if isinstance(stored, str) and stored.strip(): + return stored + env_value = os.environ.get(_UI_THEME_FIELD_ENV_VARS[field_name]) + return env_value if _is_public_http_url(env_value) else None + + class IPAddress(BaseModel): ip: str @@ -69,7 +104,8 @@ class SettingsResponse(BaseModel): class SSOSettingsResponse(SettingsResponse): """Response model for SSO settings""" - pass + provenance: Dict[str, str] = Field(default_factory=dict) + """Per-field source of each value: 'db', 'env', 'default', or 'unset'.""" class InternalUserSettingsResponse(SettingsResponse): @@ -717,7 +753,7 @@ async def get_sso_settings(): Returns a structured object with values and descriptions for UI display. """ - from litellm.proxy.proxy_server import prisma_client, proxy_config + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: raise HTTPException( @@ -725,59 +761,12 @@ async def get_sso_settings(): detail={"error": "Database not connected. Please connect a database."}, ) - # Get SSO config from dedicated table + # Resolve the effective SSO config: the stored row wins, else the process + # environment, else each field's default. Unlike the legacy read path this + # does not write os.environ; a GET has no business mutating the environment. sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}) - - # Initialize with defaults - sso_settings_dict = {} - - if sso_db_record and sso_db_record.sso_settings: - # Load settings from database - sso_settings_dict = dict(sso_db_record.sso_settings) - - role_mappings_data = sso_settings_dict.pop("role_mappings", None) - role_mappings = None - if role_mappings_data: - from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings - - if isinstance(role_mappings_data, dict): - role_mappings = RoleMappings(**role_mappings_data) - elif isinstance(role_mappings_data, RoleMappings): - role_mappings = role_mappings_data - - team_mappings_data = sso_settings_dict.pop("team_mappings", None) - team_mappings = None - if team_mappings_data: - from litellm.types.proxy.management_endpoints.ui_sso import TeamMappings - - if isinstance(team_mappings_data, dict): - team_mappings = TeamMappings(**team_mappings_data) - elif isinstance(team_mappings_data, TeamMappings): - team_mappings = team_mappings_data - - decrypted_sso_settings_dict = proxy_config._decrypt_and_set_db_env_variables( - environment_variables=sso_settings_dict - ) - - # Build SSO config with database values or environment fallback - - sso_config = SSOConfig( - google_client_id=decrypted_sso_settings_dict.get("google_client_id", None), - google_client_secret=decrypted_sso_settings_dict.get("google_client_secret", None), - microsoft_client_id=decrypted_sso_settings_dict.get("microsoft_client_id", None), - microsoft_client_secret=decrypted_sso_settings_dict.get("microsoft_client_secret", None), - microsoft_tenant=decrypted_sso_settings_dict.get("microsoft_tenant", None), - generic_client_id=decrypted_sso_settings_dict.get("generic_client_id", None), - generic_client_secret=decrypted_sso_settings_dict.get("generic_client_secret", None), - generic_authorization_endpoint=decrypted_sso_settings_dict.get("generic_authorization_endpoint", None), - generic_token_endpoint=decrypted_sso_settings_dict.get("generic_token_endpoint", None), - generic_userinfo_endpoint=decrypted_sso_settings_dict.get("generic_userinfo_endpoint", None), - proxy_base_url=decrypted_sso_settings_dict.get("proxy_base_url", None), - user_email=decrypted_sso_settings_dict.get("user_email"), - ui_access_mode=decrypted_sso_settings_dict.get("ui_access_mode"), - role_mappings=role_mappings, - team_mappings=team_mappings, - ) + sso_db_settings = dict(sso_db_record.sso_settings) if sso_db_record and sso_db_record.sso_settings else None + resolved = resolve_sso_config(sso_db_settings, os.environ) # Get the schema for UI display from pydantic import TypeAdapter @@ -786,11 +775,12 @@ async def get_sso_settings(): # Convert to dict for response, masking OAuth client secrets so plaintext # is never sent to the UI. - sso_dict = mask_sensitive_keys(sso_config.model_dump(), _SSO_SENSITIVE_FIELDS) + sso_dict = mask_sensitive_keys(resolved.config.model_dump(), set(SSO_SECRET_FIELDS)) # Add descriptions to the response result = { "values": sso_dict, + "provenance": resolved.provenance, "field_schema": { "description": schema.get("description", ""), "properties": {}, @@ -841,21 +831,6 @@ async def update_sso_settings( detail={"error": "Set `'STORE_MODEL_IN_DB='True'` in your env to enable this feature."}, ) - # Update environment variables - env_var_mapping = { - "google_client_id": "GOOGLE_CLIENT_ID", - "google_client_secret": "GOOGLE_CLIENT_SECRET", - "microsoft_client_id": "MICROSOFT_CLIENT_ID", - "microsoft_client_secret": "MICROSOFT_CLIENT_SECRET", - "microsoft_tenant": "MICROSOFT_TENANT", - "generic_client_id": "GENERIC_CLIENT_ID", - "generic_client_secret": "GENERIC_CLIENT_SECRET", - "generic_authorization_endpoint": "GENERIC_AUTHORIZATION_ENDPOINT", - "generic_token_endpoint": "GENERIC_TOKEN_ENDPOINT", - "generic_userinfo_endpoint": "GENERIC_USERINFO_ENDPOINT", - "proxy_base_url": "PROXY_BASE_URL", - } - # Read the existing SSO row first so the audit log captures a real # before/after diff. Stored values are encrypted; decrypt them so the # before-snapshot has the same shape as after_value, and rely on @@ -884,8 +859,8 @@ async def update_sso_settings( # Update environment variables in config and in memory sso_data = sso_config.model_dump() for field_name, value in sso_data.items(): - if field_name in env_var_mapping: - env_var_name = env_var_mapping[field_name] + if field_name in SSO_FIELD_ENV_VARS: + env_var_name = SSO_FIELD_ENV_VARS[field_name] if value: os.environ[env_var_name] = value else: @@ -935,7 +910,7 @@ async def update_sso_settings( else: environment_variables = {} - env_vars_to_remove = set(env_var_mapping.values()) + env_vars_to_remove = set(SSO_FIELD_ENV_VARS.values()) filtered_env_vars = { key: value for key, value in environment_variables.items() if key not in env_vars_to_remove } @@ -977,12 +952,19 @@ async def get_ui_theme_settings(): # Load existing config config = await proxy_config.get_config() - return await _get_settings_with_schema( + result = await _get_settings_with_schema( settings_key="ui_theme_config", settings_class=UIThemeConfig, config=config, ) + stored_values = result.get("values", {}) + result["values"] = { + **stored_values, + **{field: _resolve_ui_theme_field(stored_values, field) for field in _UI_THEME_FIELD_ENV_VARS}, + } + return result + def _validate_public_image_url(value: Optional[str], field_name: str) -> None: """ @@ -1041,13 +1023,6 @@ async def update_ui_theme_settings( config = await proxy_config.get_config() before_theme = config.get("litellm_settings", {}).get("ui_theme_config") - # Update config with UI theme settings - if "general_settings" not in config: - config["general_settings"] = {} - - if "environment_variables" not in config: - config["environment_variables"] = {} - # Convert theme config to dict theme_data = theme_config.model_dump(exclude_none=True) @@ -1056,55 +1031,29 @@ async def update_ui_theme_settings( config["litellm_settings"] = {} config["litellm_settings"]["ui_theme_config"] = theme_data - # Update UI_LOGO_PATH environment variable if logo_url is provided - # If logo_url is empty string, None, or null, remove the environment variable to use default - logo_url = theme_data.get("logo_url") - verbose_proxy_logger.debug(f"Updating logo_url: {logo_url}") + # UI_LOGO_PATH and LITELLM_FAVICON_URL are the only environment variables + # this endpoint owns. A non-empty value sets the var; an empty or missing + # one clears it back to the default. Apply to the live process immediately, + # then persist only these two keys so an unrelated env var (a YAML/OS value + # merged in by get_config) is never snapshotted into the DB. + def _clean(url: str | None) -> str | None: + return url if url is not None and url.strip() else None - if ( - logo_url and isinstance(logo_url, str) and logo_url.strip() - ): # Check if logo_url exists and is not empty/whitespace - config["environment_variables"]["UI_LOGO_PATH"] = logo_url - os.environ["UI_LOGO_PATH"] = logo_url - verbose_proxy_logger.debug(f"Set UI_LOGO_PATH to: {logo_url}") - else: - # Remove the environment variable to restore default logo - if "UI_LOGO_PATH" in config.get("environment_variables", {}): - del config["environment_variables"]["UI_LOGO_PATH"] - verbose_proxy_logger.debug("Removed UI_LOGO_PATH from config") - if "UI_LOGO_PATH" in os.environ: - del os.environ["UI_LOGO_PATH"] - verbose_proxy_logger.debug("Removed UI_LOGO_PATH from environment") + env_updates: dict[str, str | None] = { + "UI_LOGO_PATH": _clean(theme_config.logo_url), + "LITELLM_FAVICON_URL": _clean(theme_config.favicon_url), + } + for env_key, env_value in env_updates.items(): + if env_value is not None: + os.environ[env_key] = env_value + else: + os.environ.pop(env_key, None) - # Update LITELLM_FAVICON_URL environment variable if favicon_url is provided - favicon_url = theme_data.get("favicon_url") - verbose_proxy_logger.debug(f"Updating favicon_url: {favicon_url}") - - if ( - favicon_url and isinstance(favicon_url, str) and favicon_url.strip() - ): # Check if favicon_url exists and is not empty/whitespace - config["environment_variables"]["LITELLM_FAVICON_URL"] = favicon_url - os.environ["LITELLM_FAVICON_URL"] = favicon_url - verbose_proxy_logger.debug(f"Set LITELLM_FAVICON_URL to: {favicon_url}") - else: - # Remove the environment variable to restore default favicon - if "LITELLM_FAVICON_URL" in config.get("environment_variables", {}): - del config["environment_variables"]["LITELLM_FAVICON_URL"] - verbose_proxy_logger.debug("Removed LITELLM_FAVICON_URL from config") - if "LITELLM_FAVICON_URL" in os.environ: - del os.environ["LITELLM_FAVICON_URL"] - verbose_proxy_logger.debug("Removed LITELLM_FAVICON_URL from environment") - - # Handle environment variable encryption if needed - stored_config = config.copy() - if "environment_variables" in stored_config and len(stored_config["environment_variables"]) > 0: - # Only encrypt if there are environment variables to encrypt - stored_config["environment_variables"] = proxy_config._encrypt_env_variables( - environment_variables=stored_config["environment_variables"] - ) - - # Save the updated config - await proxy_config.save_config(new_config=stored_config) + # Persist the theme config (litellm_settings). save_config defaults to + # include_env_vars=False, so it does not snapshot environment_variables. + await proxy_config.save_config(new_config=config) + # Persist only the two owned env vars, merged against the existing DB row. + await proxy_config.save_environment_variables(env_updates) asyncio.create_task( create_config_audit_log( diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 43921a847a9..5b81d1f2da3 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -172,6 +172,7 @@ from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams if TYPE_CHECKING: from opentelemetry.trace import Span as _Span + from prisma.client import TransactionManager from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -2922,6 +2923,14 @@ class PrismaClient: return self.db.writer return self.db + def tx(self) -> "TransactionManager": + """Open an interactive transaction on the writer. + + Callers go through this instead of reaching into ``self.db`` so writer + selection and read-replica routing stay encapsulated in the wrapper. + """ + return cast("TransactionManager", self.db.tx()) # cast-ok: wrappers delegate tx via __getattr__ (untyped) + def get_request_status(self, payload: Union[dict, SpendLogsPayload]) -> Literal["success", "failure"]: """ Determine if a request was successful or failed based on payload metadata. @@ -6159,6 +6168,9 @@ def create_model_info_response( if model_cost_info is not None: max_input_tokens = coerce_token_limit(model_cost_info.get("max_input_tokens")) max_output_tokens = coerce_token_limit(model_cost_info.get("max_output_tokens")) + mode = model_cost_info.get("mode") + if isinstance(mode, str): + base["mode"] = mode if llm_router is not None: configured_input, configured_output = llm_router.get_configured_token_limits(model_id) diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 3227aa812ca..68875bd7972 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -4,11 +4,18 @@ Team repository for database operations on LiteLLM_TeamTable. import json from datetime import datetime -from typing import Any, Dict, List, Optional, Type +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type -from litellm.models.team import LiteLLM_TeamTable +from pydantic import TypeAdapter + +from litellm.models.team import LiteLLM_TeamTable, Member from litellm.repositories.base_repository import BaseRepository +if TYPE_CHECKING: + from prisma import Prisma + +_MEMBERS_WITH_ROLES_ADAPTER = TypeAdapter(list[Member]) + class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" @@ -46,6 +53,24 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): return LiteLLM_TeamTable(**data) + async def get_members_with_roles_locked(self, tx: "Prisma", team_id: str) -> List[Member]: + """Return the team's members_with_roles, locking the row FOR UPDATE. + + Must be called inside a transaction so the row lock is held until + commit. This serializes concurrent membership writers on the team row + so the losing writer appends onto the winner's committed result instead + of overwriting it from a stale snapshot. + """ + rows = await tx.query_raw( + 'SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = $1 FOR UPDATE', + team_id, + ) + raw_value = rows[0]["members_with_roles"] if rows else None + parsed = json.loads(raw_value) if isinstance(raw_value, str) else raw_value + if not parsed: + return [] + return _MEMBERS_WITH_ROLES_ADAPTER.validate_python(parsed) + async def find_by_id(self, team_id: str, id_field: str = "team_id") -> Optional[LiteLLM_TeamTable]: return await super().find_by_id(team_id, id_field) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 206736f501a..944cf58df1c 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -494,7 +494,14 @@ async def aresponses( prompt_label=kwargs.get("prompt_label", None), prompt_version=kwargs.get("prompt_version", None), ) - input = cast(Union[str, ResponseInputParam], merged_input) + input = cast( + Union[str, ResponseInputParam], + ResponsesAPIRequestUtils.merge_prompt_management_input( + original_input=input, + client_input=client_input, + merged_input=merged_input, + ), + ) if model != original_model: _, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model) kwargs.pop("prompt_id", None) @@ -609,7 +616,14 @@ def _apply_prompt_management_to_responses_call( prompt_label=kwargs.get("prompt_label", None), prompt_version=kwargs.get("prompt_version", None), ) - input = cast(Union[str, ResponseInputParam], merged_input) + input = cast( + Union[str, ResponseInputParam], + ResponsesAPIRequestUtils.merge_prompt_management_input( + original_input=input, + client_input=client_input, + merged_input=merged_input, + ), + ) local_vars["input"] = input local_vars["model"] = model if model != original_model: diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py index 7a42cb96566..ac92e5d6dcc 100644 --- a/litellm/responses/utils.py +++ b/litellm/responses/utils.py @@ -19,7 +19,9 @@ import litellm from litellm._logging import verbose_logger from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.types.llms.openai import ( + AllMessageValues, ResponseAPIUsage, + ResponseInputParam, ResponsesAPIOptionalRequestParams, ResponsesAPIResponse, ResponseText, @@ -36,6 +38,57 @@ from litellm.types.utils import ( class ResponsesAPIRequestUtils: """Helper utils for constructing ResponseAPI requests""" + @staticmethod + def merge_prompt_management_input( + original_input: str | ResponseInputParam, + client_input: list[AllMessageValues], + merged_input: list[AllMessageValues], + ) -> list[object]: + if isinstance(original_input, str): + return [*merged_input] + + original_items = tuple(original_input) + client_item_ids = frozenset(id(item) for item in client_input) + message_positions = tuple(index for index, item in enumerate(original_items) if id(item) in client_item_ids) + + if len(message_positions) == len(original_items): + return [*merged_input] + if not message_positions: + verbose_logger.warning( + "Prompt management hook returned messages without Responses API input messages; merged messages were ignored" + ) + return [*original_items] + + corresponding_messages = len(client_input) == len(merged_input) and all( + original.get("role") == merged.get("role") + and (not isinstance(original.get("id"), str) or original.get("id") == merged.get("id")) + for original, merged in zip(client_input, merged_input) + ) + if corresponding_messages: + merged_by_position = dict(zip(message_positions, merged_input)) + return [ + merged_by_position[index] if index in merged_by_position else item + for index, item in enumerate(original_items) + ] + + all_messages_preserved = all(any(original is merged for merged in merged_input) for original in client_input) + if all_messages_preserved: + prefixes = { + id(original_items[position]): original_items[ + message_positions[index - 1] + 1 if index else 0 : position + ] + for index, position in enumerate(message_positions) + } + trailing_items = original_items[message_positions[-1] + 1 :] + return [item for merged in merged_input for item in (*prefixes.get(id(merged), ()), merged)] + list( + trailing_items + ) + + verbose_logger.warning( + "Prompt management hook replaced Responses API messages; non-message input items were dropped" + ) + return [*merged_input] + @staticmethod def _check_valid_arg( supported_params: Optional[List[str]], diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 3605ab95d1b..c86794b90f8 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -725,6 +725,18 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up description="When True, guardrails only receive the latest message for the relevant role (e.g., newest user input pre-call, newest assistant output post-call)", ) + only_scan_new_messages: Optional[bool] = Field( + default=False, + description=( + "When True, the guardrail only scans messages that have not already been scanned " + "earlier in the same session (identified by litellm_session_id / session_id). " + "Message content is hashed per session and cached; only the diff (new or edited " + "messages) is sent to the guardrail provider on follow-up calls. Falls back to a " + "full scan when the request has no session id or the cache is unavailable. Intended " + "for blocking/detection guardrails; not applied when mask_request_content is set." + ), + ) + skip_system_message_in_guardrail: Optional[bool] = Field( default=None, description=( @@ -826,6 +838,13 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up "while fail_on_error still governs real Model Armor API errors. Default False blocks them." ), ) + sanitize_error_detail: Optional[bool] = Field( + default=True, + description=( + "For guardrail='model_armor': omit the raw Model Armor response from " + "caller-facing errors and logs by default. Set False to restore verbose output." + ), + ) additional_provider_specific_params: Optional[Dict[str, Any]] = Field( default=None, diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index 793cc02ff17..4a1ef5ed696 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -173,6 +173,7 @@ class Status1(Enum): cancelled = "cancelled" incomplete = "incomplete" budget_exceeded = "budget_exceeded" + queued = "queued" class InteractionStatusUpdate(BaseModel): @@ -341,6 +342,7 @@ class Status3(Enum): CANCELLED = "cancelled" INCOMPLETE = "incomplete" BUDGET_EXCEEDED = "budget_exceeded" + QUEUED = "queued" class ModelOption(RootModel[str]): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py b/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py index 628ac0442de..d5e601ce8ea 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/model_armor.py @@ -20,6 +20,13 @@ class ModelArmorGuardrailConfigModel(GuardrailConfigModel): default=True, description="Whether to fail the request if Model Armor encounters an error", ) + sanitize_error_detail: Optional[bool] = Field( + default=True, + description=( + "Omit the raw Model Armor response from caller-facing errors and logs " + "by default. Set False to restore verbose output." + ), + ) @staticmethod def ui_friendly_name() -> str: diff --git a/litellm/types/proxy/management_endpoints/ui_sso.py b/litellm/types/proxy/management_endpoints/ui_sso.py index 7234cc2650f..742e0f7818f 100644 --- a/litellm/types/proxy/management_endpoints/ui_sso.py +++ b/litellm/types/proxy/management_endpoints/ui_sso.py @@ -148,6 +148,10 @@ class SSOConfig(LiteLLMPydanticObjectBase): default=None, description="User info endpoint URL for generic OAuth provider", ) + generic_scope: Optional[str] = Field( + default=None, + description="Space-separated OAuth scopes requested from the generic provider, e.g. 'openid email profile'", + ) # Common settings proxy_base_url: Optional[str] = Field( diff --git a/litellm/types/proxy/model_listing.py b/litellm/types/proxy/model_listing.py index c3330da0d66..b59c0f2cf19 100644 --- a/litellm/types/proxy/model_listing.py +++ b/litellm/types/proxy/model_listing.py @@ -10,12 +10,16 @@ class ModelInfoMetadata(TypedDict): class ModelInfoResponse(TypedDict): - """OpenAI-compatible model object. `metadata` is present only when the - endpoint is called with include_metadata=true. + """OpenAI-compatible model object. `mode`, `max_input_tokens`, and + `max_output_tokens` are attached when the cost map knows them; `metadata` + is present only when the endpoint is called with include_metadata=true. """ id: str object: Literal["model"] created: int owned_by: str + mode: NotRequired[str] + max_input_tokens: NotRequired[int] + max_output_tokens: NotRequired[int] metadata: NotRequired[ModelInfoMetadata] diff --git a/litellm/utils.py b/litellm/utils.py index 174bed09396..a11c5500503 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1864,7 +1864,9 @@ def client(original_function): except Exception: pass - setattr(e, "num_retries", num_retries) ## IMPORTANT: returns the deployment's num_retries to the router + deployment_num_retries = kwargs.get("num_retries") + if deployment_num_retries is not None: + setattr(e, "num_retries", deployment_num_retries) timeout = _get_wrapper_timeout(kwargs=kwargs, exception=e) setattr(e, "timeout", timeout) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 5efd61f9747..c9d871fc41d 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -17644,6 +17644,61 @@ }, "web_search_billing_unit": "per_query" }, + "gemini-3.5-flash-lite": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 2e-08, + "cache_read_input_token_cost_priority": 5e-08, + "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, @@ -18311,6 +18366,60 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.6-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_flex": 7.5e-08, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_batches": 7.5e-07, + "input_cost_per_token_flex": 7.5e-07, + "litellm_provider": "vertex_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 7.5e-06, + "output_cost_per_token": 7.5e-06, + "output_cost_per_token_batches": 3.75e-06, + "output_cost_per_token_flex": 3.75e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "output_cost_per_token_priority": 1.35e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "vertex_ai/gemini-3.1-pro-preview": { "cache_read_input_token_cost": 2e-07, "cache_read_input_token_cost_above_200k_tokens": 4e-07, @@ -19663,6 +19772,63 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.5-flash-lite": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 2e-08, + "cache_read_input_token_cost_priority": 5e-08, + "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "litellm_provider": "gemini", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, + "rpm": 15, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 250000, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-3-flash-preview": { "cache_read_input_token_cost": 5e-08, "input_cost_per_audio_token": 1e-06, @@ -19769,6 +19935,63 @@ }, "web_search_billing_unit": "per_query" }, + "gemini/gemini-3.6-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_flex": 7.5e-08, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_batches": 7.5e-07, + "input_cost_per_token_flex": 7.5e-07, + "litellm_provider": "gemini", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 7.5e-06, + "output_cost_per_token": 7.5e-06, + "output_cost_per_token_batches": 3.75e-06, + "output_cost_per_token_flex": 3.75e-06, + "rpm": 2000, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "tpm": 800000, + "input_cost_per_token_priority": 2.7e-06, + "output_cost_per_token_priority": 1.35e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-omni-flash-preview": { "input_cost_per_audio_token": 1.5e-06, "input_cost_per_token": 1.5e-06, @@ -20049,6 +20272,61 @@ }, "web_search_billing_unit": "per_query" }, + "gemini-3.6-flash": { + "cache_read_input_token_cost": 1.5e-07, + "cache_read_input_token_cost_flex": 7.5e-08, + "input_cost_per_token": 1.5e-06, + "input_cost_per_token_batches": 7.5e-07, + "input_cost_per_token_flex": 7.5e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 7.5e-06, + "output_cost_per_token": 7.5e-06, + "output_cost_per_token_batches": 3.75e-06, + "output_cost_per_token_flex": 3.75e-06, + "source": "https://ai.google.dev/pricing/gemini-3", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_output": false, + "supports_audio_input": true, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "input_cost_per_token_priority": 2.7e-06, + "output_cost_per_token_priority": 1.35e-05, + "cache_read_input_token_cost_priority": 2.7e-07, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "gemini/gemini-2.5-pro-preview-tts": { "cache_read_input_token_cost": 1.25e-07, "cache_read_input_token_cost_above_200k_tokens": 2.5e-07, @@ -37323,6 +37601,61 @@ }, "web_search_billing_unit": "per_query" }, + "vertex_ai/gemini-3.5-flash-lite": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_flex": 2e-08, + "cache_read_input_token_cost_priority": 5e-08, + "input_cost_per_token": 3e-07, + "input_cost_per_token_batches": 1.5e-07, + "input_cost_per_token_flex": 1.5e-07, + "input_cost_per_token_priority": 5.4e-07, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 1048576, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_reasoning_token": 2.5e-06, + "output_cost_per_token": 2.5e-06, + "output_cost_per_token_batches": 1.25e-06, + "output_cost_per_token_flex": 1.25e-06, + "output_cost_per_token_priority": 4.5e-06, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#gemini-models", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image", + "audio", + "video" + ], + "supported_output_modalities": [ + "text" + ], + "supports_audio_input": true, + "supports_audio_output": false, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_url_context": true, + "supports_video_input": true, + "supports_vision": true, + "supports_web_search": true, + "supports_native_streaming": true, + "search_context_cost_per_query": { + "search_context_size_low": 0.014, + "search_context_size_medium": 0.014, + "search_context_size_high": 0.014 + }, + "web_search_billing_unit": "per_query" + }, "vertex_ai/deep-research-pro-preview-12-2025": { "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, diff --git a/pyproject.toml b/pyproject.toml index 9e2f5c4e3ac..62bd37c3db6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm" -version = "1.94.0" +version = "1.95.0" description = "Library to easily interface with LLM API providers" readme = "README.md" requires-python = ">=3.10, <3.15" @@ -62,7 +62,7 @@ proxy = [ "azure-identity>=1.25.2,<2.0", "azure-storage-blob>=12.28.0,<13.0", "mcp>=1.28.1,<2.0", - "litellm-proxy-extras==0.4.79", + "litellm-proxy-extras==0.4.80", "litellm-enterprise==0.1.51", "RestrictedPython>=8.1,<9.0", "rich>=13.9.4,<14.0", @@ -117,6 +117,14 @@ stt-nvidia-riva = [ "numpy>=1.26.0", ] google = ["google-cloud-aiplatform>=1.133.0,<2.0"] +bedrock-realtime = [ + # Bedrock Nova Sonic realtime (speech-to-speech) uses the + # InvokeModelWithBidirectionalStream API, which boto3 cannot do. This + # experimental AWS SDK (with its smithy-* deps, pulled transitively) + # provides the bidirectional stream; imported lazily in the realtime + # handler so litellm core stays usable without it. + "aws-sdk-bedrock-runtime>=0.7.0,<0.8.0; python_version >= '3.12'", +] proxy-runtime = [ # Historically bundled in the proxy Docker images via requirements.txt. # Keep these in a dedicated extra so uv-based images preserve the same @@ -289,7 +297,7 @@ members = ["enterprise", "litellm-proxy-extras"] profile = "black" [tool.commitizen] -version = "1.94.0" +version = "1.95.0" version_files = [ "pyproject.toml:^version", ] diff --git a/pyrightconfig.json b/pyrightconfig.json index eabfbf515c4..2686ccd73d9 100644 --- a/pyrightconfig.json +++ b/pyrightconfig.json @@ -1,7 +1,7 @@ { "include": ["litellm"], "ignore": [], - "exclude": ["**/node_modules", "**/__pycache__", "tests/e2e/claude_code", "litellm/types/utils.py", "litellm/proxy/_types.py"], + "exclude": ["**/node_modules", "**/__pycache__", "tests/e2e/claude_code", "tests/e2e/ui", "litellm/types/utils.py", "litellm/proxy/_types.py"], "pythonVersion": "3.12", "typeCheckingMode": "strict", "enableTypeIgnoreComments": false, diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 448a0079674..d3d70ff5ff4 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -93,7 +93,7 @@ "limit": 33 }, "DTZ005": { - "limit": 244 + "limit": 241 }, "DTZ006": { "limit": 13 diff --git a/schema.prisma b/schema.prisma index b27ddea010b..23a9c086c73 100644 --- a/schema.prisma +++ b/schema.prisma @@ -403,6 +403,15 @@ model LiteLLM_MCPServerOAuthClient { updated_at DateTime @default(now()) @updatedAt @map("updated_at") } +// The enterprise IdP identity assertion captured at SSO login, one row per user. +// assertion_b64 is an encrypted JSON payload: {id_token, refresh_token?, issuer?, expires_at?}. +model LiteLLM_SSOIdentityAssertion { + user_id String @id + assertion_b64 String + created_at DateTime @default(now()) @map("created_at") + updated_at DateTime @default(now()) @updatedAt @map("updated_at") +} + // Generate Tokens for Proxy model LiteLLM_VerificationToken { token String @id diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index cce0cb61c1e..150a4bbf9de 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -5,6 +5,7 @@ # gating CI checks, so a clean run means a green CI lint: # - litellm/ Python staged -> `make lint` (test-linting.yml's lint job) # - tests/e2e Python staged -> `make lint-e2e-basedpyright` (test-linting.yml's e2e type-check step) +# + raw HTTP client ban (test-code-quality.yml's check_e2e_no_raw_requests) # - dashboard staged -> prettier + eslint + lint budgets (test-litellm-ui-build.yml's frontend-lint) # - proxy/types staged -> regenerate dashboard API types and fail on drift (check-ui-api-types.yml) # @@ -112,6 +113,12 @@ if [ -n "$e2e_py_files" ] && [ -z "$litellm_py_files" ]; then make lint-e2e-basedpyright || { echo "✗ tests/e2e basedpyright failed. Fix the errors above, then re-run make pre-commit." >&2; status=1; } fi +if [ -n "$e2e_py_files" ]; then + echo "pre-commit: checking tests/e2e raw HTTP client ban (check_e2e_no_raw_requests)" + uv run --no-sync python tests/code_coverage_tests/check_e2e_no_raw_requests.py \ + || { echo "✗ Raw HTTP client import in tests/e2e. Route the call through tests/e2e/e2e_http.py, then re-run make pre-commit." >&2; status=1; } +fi + if [ -n "$ui_prettier_files" ] || [ -n "$ui_eslint_files" ]; then echo "pre-commit: linting dashboard (prettier + eslint + lint budgets)" if [ ! -d ui/litellm-dashboard/node_modules ]; then diff --git a/tests/code_coverage_tests/check_e2e_no_raw_requests.py b/tests/code_coverage_tests/check_e2e_no_raw_requests.py new file mode 100644 index 00000000000..e70e83652d1 --- /dev/null +++ b/tests/code_coverage_tests/check_e2e_no_raw_requests.py @@ -0,0 +1,81 @@ +"""tests/e2e routes every HTTP call through the typed transport (e2e_http.py), so +raw HTTP client imports (requests, urllib.request, httpx, aiohttp, http.client) are +banned in suite code. Importing requests' exception types for catching is fine +anywhere; a small allowlist grandfathers the files that legitimately make raw calls +(the transport itself, the root conftest liveness probe, and the claude_code version +resolver's constant registry URL fetch). Referenced by tests/e2e/CLAUDE.md.""" + +from __future__ import annotations + +import ast +import sys +from pathlib import Path + +E2E_DIR = Path(__file__).resolve().parents[1] / "e2e" + +BANNED_MODULES = ("requests", "urllib.request", "http.client", "httpx", "aiohttp") + +ALLOWED_RAW_CLIENT_FILES = { + "e2e_http.py": ("requests",), + "conftest.py": ("requests",), + "claude_code/pr_gate_version_resolver.py": ("urllib.request",), +} + +EXCEPTION_ONLY_NAMES = frozenset({"RequestException", "ConnectionError", "Timeout", "HTTPError"}) + + +def _is_banned(module: str) -> bool: + return any(module == banned or module.startswith(banned + ".") for banned in BANNED_MODULES) + + +def _banned_imports(tree: ast.Module) -> tuple[tuple[str, int], ...]: + plain = tuple( + (alias.name, node.lineno) + for node in ast.walk(tree) + if isinstance(node, ast.Import) + for alias in node.names + if _is_banned(alias.name) + ) + from_imports = tuple( + (node.module, node.lineno) + for node in ast.walk(tree) + if isinstance(node, ast.ImportFrom) + and node.module is not None + and _is_banned(node.module) + and not all(alias.name in EXCEPTION_ONLY_NAMES for alias in node.names) + ) + return plain + from_imports + + +def _violations_in(path: Path) -> tuple[str, ...]: + relative = path.relative_to(E2E_DIR).as_posix() + allowed = ALLOWED_RAW_CLIENT_FILES.get(relative, ()) + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + return tuple( + f"tests/e2e/{relative}:{lineno}: raw HTTP client import '{module}'" + for module, lineno in _banned_imports(tree) + if module not in allowed + ) + + +def main() -> int: + violations = tuple( + violation + for path in sorted(E2E_DIR.rglob("*.py")) + for violation in _violations_in(path) + ) + for violation in violations: + print(violation) + if violations: + print( + f"\n{len(violations)} raw HTTP client import(s) in tests/e2e. " + "Route the call through tests/e2e/e2e_http.py (get_external for absolute " + "third-party URLs) so it gets the typed Result handling." + ) + return 1 + print("tests/e2e raw HTTP client check passed") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index e08d703d21f..0bc3cebdd5a 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -55,6 +55,8 @@ IGNORE_FUNCTIONS = [ "_freeze_for_dedupe", # OTEL: max depth set (default 16, _FREEZE_MAX_DEPTH); fails closed by returning repr(value) at the cap. "apply_json_merge_patch", # max depth set (_MAX_MERGE_DEPTH=64); fails closed by raising ValueError at the cap. "_filter_mcp_argument_value", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by blocking the MCP call at the cap. + "_redact_scanned_content", # max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by returning "[REDACTED]" at the cap. + "_iter_fallback_targets", # max depth set (2 * ROUTER_MAX_FALLBACKS); fails closed by raising ValueError at the cap. ] diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index 47f3c74d7f1..0e39664e358 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -13,13 +13,16 @@ Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family - `realtime/` - realtime websocket sessions, including the pipecat audio path - `quota_management/` - quota enforcement and accounting, one subfolder per behavior: `ratelimit/` (rpm/tpm blocks, window reset, pacing headers on live traffic), `budgets/` (budget definition, enforcement, and reset windows: key, team, tag, soft, multi-window), and `spend_tracking/` (spend logging and cost attribution on `/spend/*`) - `management/` - key/team/user/organization management routes: create/update/delete persistence via the info routes, team membership, and llm-only-key route denials (API surface; not Playwright) +- `a2a/` - the A2A (agent-to-agent) surface: admin registration via `/v1/agents`, proxy-fronted card discovery at `/.well-known/agent-card.json`, and JSON-RPC `message/send` invocation, driving agents backed by the litellm completion bridge (a real provider) and asserting protocol-version normalization (0.3 vs 1.0) - `mcp/` - the MCP server surface over api_key auth against the real Datadog remote MCP server only (see "MCP suite: real Datadog only" below) - `logging/` - logging-integration delivery (datadog and friends) - `security/` - secret handling and log-leak protection - `router/` - routing and reliability behavior (fallbacks, cooldowns) -- `load/` - throughput/performance under concurrency: drives real concurrent traffic through the whole stack with Locust and asserts a throughput SLO; marked `load` so the parent conftest collects it last and it never perturbs latency-sensitive suites +- `load/` - throughput/performance under concurrency: drives real concurrent traffic through the whole stack with Locust and asserts a throughput SLO; marked `load` so the parent conftest collects it last and it never perturbs latency-sensitive suites. Also home of the weekly session-anomaly test (`test_weekly_session_anomaly_e2e.py`): Claude Code-shaped multi-turn sessions against real providers with ceilings on error rate, cache read/write, turn time, and spend; additionally marked `weekly` and deselected unless `E2E_WEEKLY_ANOMALY` is set, because it spends real provider money (driven by `.github/workflows/weekly_load_anomaly.yml`) +- `other/` - the holding-pen suite for the `other.*` registry cluster with no home of its own yet: the master-key auth gate and the process-lifecycle health probes (liveness, public readiness, authenticated readiness diagnostics). Promote a cluster out once it is large/stable enough for its own suite - `gateway/` - proxy configuration only (`litellm-config.yml`); no tests - `claude_code/` - the Claude Code compatibility matrix: drives the real `claude` CLI (and HTTP probes) against a proxy for each feature x provider cell, reporting tagged-union outcomes via the `compat_result` fixture; ships its own driver/builder/publisher plus `_*_unit_tests/` trees. The HTTP probes ride the shared transport (`ProxyClient.count_tokens` / `ProxyClient.messages`); the CLI-driving path stays bespoke +- `ui/` - the Admin UI browser suite: Playwright in TypeScript, driving the dashboard served by a live proxy on port 4000 (seeded postgres + mock LLM upstream; see its `run_e2e.sh`). It is a self-contained npm package with its own lockfile and does not use the Python harness, pytest markers, or the shared transport; the Python rules in this file (typed models, `Result` unions, basedpyright zero-error gate) do not apply inside it. Its only Python file, `fixtures/mock_llm_server/server.py`, is excluded from the e2e basedpyright gate via the root `pyrightconfig.json` ## MCP suite: real Datadog only @@ -130,7 +133,7 @@ reliability... behavior : fallback | retry | cooldown | timeout | routing | cache | circuit_breaker | perf variant : 5xx | context_window | content_policy | 429 | timeout simple_shuffle | usage_based | latency_based | cost_based | least_busy - latency | throughput (perf only; SLO/threshold assertion, not binary) + latency | throughput | session_anomaly (perf only; SLO/threshold assertion, not binary) assertion : routes_to_fallback | succeeds_within_retries | picks_under_tpm | returns_cached | trips_then_recovers | under_slo e.g. reliability.fallback.context_window.routes_to_fallback exercised_on=[chat_completions] diff --git a/tests/e2e/a2a/a2a_client.py b/tests/e2e/a2a/a2a_client.py new file mode 100644 index 00000000000..97ffa8c34a3 --- /dev/null +++ b/tests/e2e/a2a/a2a_client.py @@ -0,0 +1,292 @@ +"""Client for the proxy's A2A (agent-to-agent) surface. + +An A2A agent is registered admin-side via POST /v1/agents with an agent card and +litellm_params; the proxy fronts it at /a2a/{id}, serving a proxy-owned agent card +at /.well-known/agent-card.json and accepting A2A JSON-RPC calls at /a2a/{id}. This +suite registers agents backed by the litellm_completion_bridge (custom_llm_provider ++ model), so message/send runs a real provider completion and comes back in the +agent's pinned A2A protocol version. The A2A request/response models are co-located +here because only this suite uses them. +""" + +from __future__ import annotations + +import warnings +from dataclasses import dataclass + +from pydantic import BaseModel, ConfigDict, Field + +from e2e_http import NoBody, Result, get_external, is_ok +from proxy_client import ProxyClient + + +class A2ACapabilities(BaseModel): + streaming: bool | None = None + push_notifications: bool | None = Field(default=None, serialization_alias="pushNotifications") + + +class A2ASkill(BaseModel): + id: str + name: str + description: str + tags: list[str] + examples: list[str] | None = None + + +class A2AProvider(BaseModel): + organization: str + url: str + + +class AgentCardParams(BaseModel): + """The upstream agent card an admin registers. `protocolVersion` is the field the + proxy validates against SUPPORTED_A2A_PROTOCOL_VERSIONS on registration.""" + + protocol_version: str = Field(serialization_alias="protocolVersion") + name: str + description: str + version: str + url: str | None = None + capabilities: A2ACapabilities = A2ACapabilities() + skills: list[A2ASkill] + default_input_modes: list[str] = Field(default=["text"], serialization_alias="defaultInputModes") + default_output_modes: list[str] = Field(default=["text"], serialization_alias="defaultOutputModes") + preferred_transport: str | None = Field(default=None, serialization_alias="preferredTransport") + + +class UpstreamAgentCard(BaseModel): + """A real published agent card parsed from a public /.well-known endpoint. Keys on + the A2A wire aliases so `model_validate_json` reads the served JSON and + `model_dump(by_alias=True)` re-emits it unchanged for verbatim registration; it is + only ever fetched-and-validated, never hand-constructed, so aliasing on the wire + names does not affect any call site.""" + + model_config = ConfigDict(populate_by_name=True) + + protocol_version: str = Field(alias="protocolVersion") + name: str + description: str + version: str + url: str + provider: A2AProvider | None = None + documentation_url: str | None = Field(default=None, alias="documentationUrl") + capabilities: A2ACapabilities = A2ACapabilities() + skills: list[A2ASkill] + default_input_modes: list[str] = Field(default=["text"], alias="defaultInputModes") + default_output_modes: list[str] = Field(default=["text"], alias="defaultOutputModes") + preferred_transport: str | None = Field(default=None, alias="preferredTransport") + + +class A2ABridgeParams(BaseModel): + """litellm_params that route the agent through the completion bridge: an A2A + message/send is transformed into a litellm.acompletion against this provider.""" + + model_config = ConfigDict(protected_namespaces=()) + + custom_llm_provider: str + model: str + + +class AgentRegisterBody(BaseModel): + agent_name: str + agent_card_params: AgentCardParams | UpstreamAgentCard + litellm_params: A2ABridgeParams | None = None + + +class A2ASecurityScheme(BaseModel): + type: str + scheme: str + + +class A2AInterface(BaseModel): + model_config = ConfigDict(populate_by_name=True) + + url: str + protocol_version: str | None = Field(default=None, alias="protocolVersion") + + +class ServedAgentCard(BaseModel): + """The proxy-owned card, either nested under a registration response's + `agent_card_params` or served raw at /.well-known/agent-card.json. The proxy + rewrites `url`/`supportedInterfaces` to itself and replaces the security scheme + with its own virtual-key bearer scheme.""" + + model_config = ConfigDict(populate_by_name=True) + + protocol_version: str = Field(alias="protocolVersion") + name: str + url: str | None = None + security_schemes: dict[str, A2ASecurityScheme] | None = Field(default=None, alias="securitySchemes") + security: list[dict[str, list[str]]] | None = None + supported_interfaces: list[A2AInterface] | None = Field(default=None, alias="supportedInterfaces") + + +class AgentResponse(BaseModel): + agent_id: str + agent_name: str + agent_card_params: ServedAgentCard + + +class A2ATextPart(BaseModel): + kind: str = "text" + text: str + + +class A2ASearchPropertiesParams(BaseModel): + """The strict param schema of the published property agent's `search_properties` + skill (unknown keys are rejected upstream), so a natural-language query like + "properties for sale in SF under $2M" is expressed as typed fields.""" + + un_locode: str | None = None + service_type: str | None = None + property_type: str | None = None + bedrooms_min: int | None = None + asking_price_max: float | None = None + limit: int | None = None + + +class A2ASkillInvocation(BaseModel): + skill: str + params: A2ASearchPropertiesParams + + +class A2ADataPart(BaseModel): + kind: str = "data" + data: A2ASkillInvocation + + +class A2AOutboundMessage(BaseModel): + role: str = "user" + parts: list[A2ATextPart | A2ADataPart] + message_id: str = Field(serialization_alias="messageId") + + +class A2AMessageSendParams(BaseModel): + message: A2AOutboundMessage + + +class A2AJsonRpcRequest(BaseModel): + jsonrpc: str = "2.0" + id: str + method: str = "message/send" + params: A2AMessageSendParams + + +class A2AResponsePart(BaseModel): + kind: str | None = None + text: str | None = None + + +class A2AResponseMessage(BaseModel): + model_config = ConfigDict(populate_by_name=True) + + message_id: str | None = Field(default=None, alias="messageId") + role: str | None = None + parts: list[A2AResponsePart] = [] + + +class A2ATaskStatus(BaseModel): + state: str | None = None + message: A2AResponseMessage | None = None + + +class A2AResult(BaseModel): + """A message/send result. In 0.3 the message fields sit directly on the result + (`kind`/`role`/`parts`); in 1.0 they are nested under `message`; a real agent that + runs a task replies with a `task` whose agent text lives on `status.message`. + `text` reads the agent's reply from whichever shape the served version produced.""" + + model_config = ConfigDict(populate_by_name=True) + + kind: str | None = None + role: str | None = None + message_id: str | None = Field(default=None, alias="messageId") + parts: list[A2AResponsePart] = [] + message: A2AResponseMessage | None = None + status: A2ATaskStatus | None = None + + @property + def text(self) -> str: + if self.message is not None: + parts = self.message.parts + elif self.parts: + parts = self.parts + elif self.status is not None and self.status.message is not None: + parts = self.status.message.parts + else: + parts = [] + return "".join(part.text or "" for part in parts) + + @property + def is_nested_v1_shape(self) -> bool: + return self.message is not None + + +class A2AError(BaseModel): + code: int + message: str + + +class A2AResponse(BaseModel): + jsonrpc: str + id: str | None = None + result: A2AResult | None = None + error: A2AError | None = None + + +@dataclass(frozen=True, slots=True) +class A2AClient: + proxy: ProxyClient + + def register_agent(self, body: AgentRegisterBody) -> Result[AgentResponse]: + return self.proxy.transport.post( + "/v1/agents", + headers=self.proxy.transport.master, + json=body, + response_type=AgentResponse, + ) + + def get_agent(self, agent_id: str) -> Result[AgentResponse]: + return self.proxy.transport.get( + f"/v1/agents/{agent_id}", + headers=self.proxy.transport.master, + params=NoBody(), + response_type=AgentResponse, + ) + + def delete_agent(self, agent_id: str) -> None: + result = self.proxy.transport.delete( + f"/v1/agents/{agent_id}", + headers=self.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) + if not is_ok(result): + warnings.warn(f"delete_agent({agent_id!r}) failed: {result}", stacklevel=2) + + def agent_card(self, agent_id: str, key: str) -> Result[ServedAgentCard]: + return self.proxy.transport.get( + f"/a2a/{agent_id}/.well-known/agent-card.json", + headers=self.proxy.transport.bearer(key), + params=NoBody(), + response_type=ServedAgentCard, + ) + + def send_message(self, agent_id: str, key: str, body: A2AJsonRpcRequest) -> Result[A2AResponse]: + return self.proxy.transport.post( + f"/a2a/{agent_id}", + headers=self.proxy.transport.bearer(key), + json=body, + response_type=A2AResponse, + ) + + +def build_a2a_client(proxy: ProxyClient) -> A2AClient: + return A2AClient(proxy=proxy) + + +def fetch_agent_card(url: str, *, timeout: float = 20.0) -> Result[UpstreamAgentCard]: + """Fetch a live A2A agent card from its /.well-known endpoint and parse it into the + registration model, so a test can register a real published card verbatim rather + than a hand-rolled one.""" + return get_external(url, response_type=UpstreamAgentCard, timeout=timeout) diff --git a/tests/e2e/a2a/conftest.py b/tests/e2e/a2a/conftest.py new file mode 100644 index 00000000000..93f3b56c8f7 --- /dev/null +++ b/tests/e2e/a2a/conftest.py @@ -0,0 +1,17 @@ +"""A2A suite's `client` fixture. + +The shared lifecycle (resources/scoped_key), proxy liveness gate, and e2e marker +live in the parent tests/e2e/conftest.py. A2AClient holds the shared ProxyClient, +so the `resources` fixture cleans up keys this suite creates; agents are torn down +via `resources.defer(...)` in each test. +""" + +import pytest + +from a2a_client import A2AClient, build_a2a_client +from proxy_client import ProxyClient + + +@pytest.fixture(scope="session") +def client(proxy: ProxyClient) -> A2AClient: + return build_a2a_client(proxy) diff --git a/tests/e2e/a2a/test_a2a_agent_e2e.py b/tests/e2e/a2a/test_a2a_agent_e2e.py new file mode 100644 index 00000000000..aa60b57f99b --- /dev/null +++ b/tests/e2e/a2a/test_a2a_agent_e2e.py @@ -0,0 +1,202 @@ +"""A2A agents end to end, against a live proxy. + +An admin registers an agent whose card pins an A2A protocol version and whose +litellm_params route it through the completion bridge; a caller then discovers the +proxy-owned card and drives it over A2A JSON-RPC. These tests assert the recorded +state (the agent persists, a spend row lands) and the enforced behavior (the served +card points back at the proxy, message/send returns a real completion in the pinned +protocol version, and an unsupported version is refused at registration). +""" + +from __future__ import annotations + +import pytest + +from a2a_client import ( + A2ABridgeParams, + A2AClient, + A2ADataPart, + A2AJsonRpcRequest, + A2AMessageSendParams, + A2AOutboundMessage, + A2ASearchPropertiesParams, + A2ASkill, + A2ASkillInvocation, + A2ATextPart, + AgentCardParams, + AgentRegisterBody, + AgentResponse, + fetch_agent_card, +) +from e2e_config import unique_marker +from e2e_http import Result, UnknownApiError, unwrap +from lifecycle import ResourceManager + +BRIDGE = A2ABridgeParams(custom_llm_provider="anthropic", model="claude-haiku-4-5") + +MOVEHOME_AGENT_CARD_URL = "https://movehome.org/.well-known/agent.json" +MOVEHOME_ORIGIN = "https://movehome.org" + +pytestmark = pytest.mark.e2e + + +def _register(client: A2AClient, resources: ResourceManager, protocol_version: str) -> AgentResponse: + marker = unique_marker() + body = AgentRegisterBody( + agent_name=f"e2e-a2a-{marker}", + agent_card_params=AgentCardParams( + protocol_version=protocol_version, + name=f"E2E A2A {marker}", + description="e2e agent backed by the litellm completion bridge", + version="1.0.0", + skills=[A2ASkill(id="chat", name="Chat", description="general chat", tags=["chat"])], + ), + litellm_params=BRIDGE, + ) + agent = unwrap(client.register_agent(body)) + resources.defer(lambda: client.delete_agent(agent.agent_id)) + return agent + + +def _register_rejection(client: A2AClient, protocol_version: str) -> Result[AgentResponse]: + marker = unique_marker() + body = AgentRegisterBody( + agent_name=f"e2e-a2a-bad-{marker}", + agent_card_params=AgentCardParams( + protocol_version=protocol_version, + name=f"E2E A2A bad {marker}", + description="rejected at registration", + version="1.0.0", + skills=[A2ASkill(id="chat", name="Chat", description="c", tags=["chat"])], + ), + litellm_params=BRIDGE, + ) + return client.register_agent(body) + + +def _ask(text: str) -> A2AJsonRpcRequest: + return A2AJsonRpcRequest( + id=f"e2e-{unique_marker()}", + params=A2AMessageSendParams( + message=A2AOutboundMessage(parts=[A2ATextPart(text=text)], message_id=unique_marker()) + ), + ) + + +class TestA2AAgentLifecycle: + @pytest.mark.covers("other.a2a.register.persists") + def test_register_persists(self, client: A2AClient, resources: ResourceManager) -> None: + agent = _register(client, resources, "0.3") + fetched = unwrap(client.get_agent(agent.agent_id)) + assert fetched.agent_id == agent.agent_id + assert fetched.agent_name == agent.agent_name + assert fetched.agent_card_params.protocol_version == "0.3" + + @pytest.mark.covers("other.a2a.register.semver_version_accepted") + def test_semver_protocol_version_registers_and_serves(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: + agent = _register(client, resources, "0.3.0") + assert agent.agent_card_params.protocol_version == "0.3" + card = unwrap(client.agent_card(agent.agent_id, scoped_key)) + assert card.protocol_version == "0.3" + assert card.supported_interfaces is not None + assert card.supported_interfaces[0].protocol_version == "0.3" + result = unwrap(client.send_message(agent.agent_id, scoped_key, _ask("Say hi in one word"))).result + assert result is not None + assert result.text != "" + + @pytest.mark.covers("other.a2a.message_send.real_world_agent_replies") + def test_real_world_agent_replies_to_property_query(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: + upstream = unwrap(fetch_agent_card(MOVEHOME_AGENT_CARD_URL)).model_copy(update={"url": MOVEHOME_ORIGIN}) + assert upstream.protocol_version == "0.3.0" + marker = unique_marker() + body = AgentRegisterBody(agent_name=f"e2e-a2a-real-{marker}", agent_card_params=upstream) + agent = unwrap(client.register_agent(body)) + resources.defer(lambda: client.delete_agent(agent.agent_id)) + assert agent.agent_card_params.protocol_version == "0.3" + request = A2AJsonRpcRequest( + id=f"e2e-{unique_marker()}", + params=A2AMessageSendParams( + message=A2AOutboundMessage( + parts=[ + A2ADataPart( + data=A2ASkillInvocation( + skill="search_properties", + params=A2ASearchPropertiesParams(un_locode="USSFO", service_type="sale", asking_price_max=2_000_000, limit=3), + ) + ) + ], + message_id=unique_marker(), + ) + ), + ) + response = unwrap(client.send_message(agent.agent_id, scoped_key, request)) + assert response.error is None + assert response.result is not None + assert response.result.text.strip() != "" + + @pytest.mark.covers("other.a2a.discovery.proxy_fronted_card") + def test_discovery_card_is_proxy_fronted(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: + agent = _register(client, resources, "0.3") + card = unwrap(client.agent_card(agent.agent_id, scoped_key)) + assert card.url is not None and card.url.endswith(f"/a2a/{agent.agent_id}") + assert card.security_schemes is not None + scheme = next(iter(card.security_schemes.values())) + assert scheme.scheme == "bearer" + assert card.supported_interfaces is not None + assert card.supported_interfaces[0].url == card.url + + @pytest.mark.covers("other.a2a.message_send.bridge_invokes") + def test_message_send_runs_completion_bridge(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: + agent = _register(client, resources, "0.3") + request = _ask("Reply with exactly the word PONG and nothing else") + response = unwrap(client.send_message(agent.agent_id, scoped_key, request)) + assert response.error is None + assert response.result is not None + assert "PONG" in response.result.text.upper() + + rows = client.proxy.poll_logs_for_request_id(request.id) + assert rows, f"no spend log row landed for a2a request {request.id}" + assert rows[0].call_type == "asend_message" + assert rows[0].model == f"a2a_agent/{agent.agent_card_params.name}" + + @pytest.mark.covers("other.a2a.version.serves_pinned_0_3") + def test_pinned_v0_3_serves_flat_message_shape(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: + agent = _register(client, resources, "0.3") + request = _ask("Say hi in one word") + result = unwrap(client.send_message(agent.agent_id, scoped_key, request)).result + assert result is not None + assert not result.is_nested_v1_shape + assert result.kind == "message" + assert result.role == "agent" + assert result.text != "" + + @pytest.mark.covers("other.a2a.version.serves_pinned_1_0") + def test_pinned_v1_0_serves_nested_message_shape(self, client: A2AClient, resources: ResourceManager, scoped_key: str) -> None: + agent = _register(client, resources, "1.0") + request = _ask("Say hi in one word") + result = unwrap(client.send_message(agent.agent_id, scoped_key, request)).result + assert result is not None + assert result.is_nested_v1_shape + assert result.message is not None + assert result.message.role == "ROLE_AGENT" + assert result.text != "" + + @pytest.mark.covers("other.a2a.register.unsupported_version_rejected") + def test_unsupported_protocol_version_rejected(self, client: A2AClient) -> None: + result = _register_rejection(client, "9.9") + match result: + case UnknownApiError(status_code=status, body=detail): + assert status == 400 + assert "protocolVersion" in detail + case _: + pytest.fail(f"expected 400 for unsupported protocolVersion, got {result}") + + @pytest.mark.covers("other.a2a.register.malformed_version_rejected") + def test_malformed_protocol_version_rejected(self, client: A2AClient) -> None: + result = _register_rejection(client, "0.3.garbage") + match result: + case UnknownApiError(status_code=status, body=detail): + assert status == 400 + assert "Unsupported protocolVersion '0.3.garbage'" in detail + case _: + pytest.fail(f"expected 400 for malformed protocolVersion, got {result}") diff --git a/tests/e2e/batches/batch_client.py b/tests/e2e/batches/batch_client.py index 7db5d0b6beb..5cc5d1dae3b 100644 --- a/tests/e2e/batches/batch_client.py +++ b/tests/e2e/batches/batch_client.py @@ -26,16 +26,24 @@ from e2e_http import ( ) from models import LiteLLMParamsBody +UPLOAD_FILENAME = "batch_input.jsonl" + class FileObject(BaseModel): id: str object: str | None = None purpose: str | None = None + filename: str | None = None bytes: int | None = None status: str | None = None created_at: int | None = None +class FileList(BaseModel): + object: str | None = None + data: list[FileObject] = [] + + class BatchObject(BaseModel): id: str object: str | None = None @@ -106,12 +114,30 @@ class BatchClient: _files_path(provider), headers=self.proxy.transport.bearer(key), form=form, - filename="batch_input.jsonl", + filename=UPLOAD_FILENAME, content=content, params=ModelQuery(model=model), response_type=FileObject, ) + def retrieve_file( + self, file_id: str, *, key: str, provider: str | None = None + ) -> Result[FileObject]: + return self.proxy.transport.get( + f"{_files_path(provider)}/{file_id}", + headers=self.proxy.transport.bearer(key), + params=NoBody(), + response_type=FileObject, + ) + + def list_files(self, *, key: str, provider: str | None = None) -> Result[FileList]: + return self.proxy.transport.get( + _files_path(provider), + headers=self.proxy.transport.bearer(key), + params=NoBody(), + response_type=FileList, + ) + def create_batch( self, *, body: BatchCreateBody, key: str, provider: str | None = None ) -> StreamingResponse: diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index f886c09b705..b0c53becb6b 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -26,6 +26,7 @@ import pytest from e2e_config import require_env, unique_marker from batch_client import ( + UPLOAD_FILENAME, BatchClient, BatchCreateBody, BatchObject, @@ -511,6 +512,72 @@ class TestBatchFileContent: ) +class TestOpenAIFiles: + """GET /v1/files (list) and GET /v1/files/{id} (retrieve) over the OpenAI route. + + The proxy lists the OpenAI org's raw file ids, so the list case uploads a raw + (provider-routed) file whose id matches what list returns; retrieve re-encodes + the id it was called with, so the model-encoded upload round-trips unchanged. + """ + + @pytest.mark.covers( + "llm.files.openai.list.nonstream.works", + exercised_on=["files"], + ) + def test_uploaded_file_appears_in_list( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + key = resources.key() + file = unwrap( + client.upload_file( + content=render_jsonl(OPENAI_BATCH_MODEL), + form=FileUploadForm(purpose="batch"), + key=key, + provider="openai", + ) + ) + resources.defer( + quietly(lambda: client.delete_file(file.id, key=key, provider="openai")) + ) + + listed = unwrap(client.list_files(key=key)) + assert listed.object is None or listed.object == "list", ( + f"list envelope object={listed.object!r}" + ) + match = next((entry for entry in listed.data if entry.id == file.id), None) + assert match is not None, f"uploaded file {file.id!r} absent from GET /v1/files" + assert match.purpose == "batch", ( + f"listed file must round-trip the upload purpose, got {match.purpose!r}" + ) + + @pytest.mark.covers( + "llm.files.openai.retrieve.nonstream.works", + exercised_on=["files"], + ) + def test_retrieve_round_trips_metadata( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + key = resources.key() + file = unwrap( + client.upload_file( + content=render_jsonl(OPENAI_BATCH_MODEL), + form=FileUploadForm(purpose="batch"), + model=OPENAI_BATCH_MODEL, + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + + fetched = unwrap(client.retrieve_file(file.id, key=key)) + assert fetched.id == file.id, "retrieve must echo the uploaded file id" + assert fetched.purpose == "batch", ( + f"retrieve must round-trip purpose, got {fetched.purpose!r}" + ) + assert fetched.filename == UPLOAD_FILENAME, ( + f"retrieve must round-trip filename, got {fetched.filename!r}" + ) + + BATCH_RL_REQUEST_LINES = 3 BATCH_RL_RPM_LIMIT = 2 diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 609da6a9b07..eff3b4ddf58 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -5,9 +5,9 @@ answers or when credentials/env are missing; they never skip. Pure unit coverage of the harness itself carries no `e2e` marker and runs regardless of whether a proxy is up. -Lifecycle: the `resources` fixture maps the init -> run -> teardown contract -(lifecycle.E2ECase) onto pytest - setup is init(), the test body is run(), and -teardown deletes every resource the test created on the long-lived proxy. +Lifecycle: the `resources` fixture hands each test a lifecycle.ResourceManager - +the test registers a cleanup for every resource it creates, and the fixture's +teardown deletes them all on the long-lived proxy, even when the test fails. Each suite provides its own `client` fixture (a lifecycle.ResourceClient); these shared fixtures build on it. @@ -43,6 +43,10 @@ def pytest_configure(config: pytest.Config) -> None: "markers", "load: heavy throughput/load test; collected last so it never perturbs latency-sensitive suites", ) + config.addinivalue_line( + "markers", + "weekly: real-provider anomaly load test that spends real money; deselected unless E2E_WEEKLY_ANOMALY is set", + ) def pytest_collection_modifyitems(items: list[pytest.Item]) -> None: diff --git a/tests/e2e/coverage_registry/guardrail.yaml b/tests/e2e/coverage_registry/guardrail.yaml index 68722fbbb96..d54c12ba6dc 100644 --- a/tests/e2e/coverage_registry/guardrail.yaml +++ b/tests/e2e/coverage_registry/guardrail.yaml @@ -31,3 +31,4 @@ - {id: guardrail.tool_policy.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/tool_policy/tool_policy_guardrail.py", rationale: "Tool-use policy enforcement"} - {id: guardrail.mcp_security.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/mcp_security", rationale: "MCP protocol security"} - {id: guardrail.llm_as_a_judge.pre_call.blocks, module: guardrail, tier: P2, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/llm_as_a_judge", rationale: "LLM-based judgment guardrail"} +- {id: guardrail.litellm_content_filter.pre_mcp_call.blocks, module: guardrail, tier: P1, hook_point: pre_mcp_call, assertions: [blocks], exercised_on: [mcp_operations], source: "guardrail_hooks/litellm_content_filter/content_filter.py:_scan_mcp_tool_call_arguments", rationale: "A general content-filter guardrail configured mode=pre_mcp_call blocks a banned keyword in an MCP tool call's arguments before it reaches the upstream MCP server; a clean argument passes"} diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 3b1aff80024..26280d35da0 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -46,10 +46,10 @@ - {id: llm.messages.anthropic.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Extended thinking via Messages API"} - {id: llm.messages.bedrock_invoke.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Flagged Claude 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (#32578/#32831/#32882)", fail_before_fix: proven} - {id: llm.messages.bedrock_invoke.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: bedrock_invoke, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py", rationale: "Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (#32831)", fail_before_fix: proven} -- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (Kraken Tech RCA gap)", fail_before_fix: proven} -- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (Kraken Tech RCA gap)", fail_before_fix: proven} -- {id: llm.messages.vertex.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (Kraken Tech RCA gap)", fail_before_fix: proven} -- {id: llm.messages.vertex.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (Kraken Tech RCA gap)", fail_before_fix: proven} +- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (customer RCA gap)", fail_before_fix: proven} +- {id: llm.messages.azure_foundry.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: azure_foundry, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/azure_ai/anthropic/messages_transformation.py", rationale: "Azure Foundry Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (customer RCA gap)", fail_before_fix: proven} +- {id: llm.messages.vertex.mid_conversation_system.nonstream.cache_hit, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works, cache_hit], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex serves Claude on the native Anthropic contract, so flagged 4.8+/5 must keep mid-conversation system reminders in messages; hoisting mutates the system prefix and collapses the prompt cache (customer RCA gap)", fail_before_fix: proven} +- {id: llm.messages.vertex.mid_conversation_system.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: vertex, capability: mid_conversation_system, streaming: nonstream, assertions: [works], source: "llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py", rationale: "Vertex Claude <= 4.7 rejects role system inside messages; unflagged models must hoist reminders into top-level system or every Claude Code session 400s (customer RCA gap)", fail_before_fix: proven} - {id: llm.responses.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Core endpoint; OpenAI Responses native"} - {id: llm.responses.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: stream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Streaming via /v1/responses"} - {id: llm.responses.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "response_api_endpoints/endpoints.py:26", rationale: "Cost logged on responses"} diff --git a/tests/e2e/coverage_registry/llm_nonconversational.yaml b/tests/e2e/coverage_registry/llm_nonconversational.yaml index 2b456aacefc..63e6fde14a3 100644 --- a/tests/e2e/coverage_registry/llm_nonconversational.yaml +++ b/tests/e2e/coverage_registry/llm_nonconversational.yaml @@ -32,7 +32,7 @@ - {id: llm.files.hosted_vllm.upload.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible file upload"} - {id: llm.rerank.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_rerank_e2e.py:29", rationale: "Cohere rerank, top_n + relevance_score"} - {id: llm.files.openai.content.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "GET /v1/files/{id}/content returns uploaded batch JSONL bytes"} -- {id: llm.realtime.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "test_nova_sonic_realtime_e2e.py", rationale: "Nova Sonic realtime session emits response.done (LIT-2239)"} +- {id: llm.realtime.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "test_realtime_bedrock_e2e.py", rationale: "Nova Sonic realtime session emits response.done (LIT-2239)"} - {id: llm.rerank.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/rerank/handler.py", rationale: "Bedrock rerank"} - {id: llm.rerank.together_ai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: together_ai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/together_ai/rerank/handler.py", rationale: "Together rerank"} - {id: llm.images_generations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_image_generation_e2e.py:22", rationale: "OpenAI image gen, b64/url"} diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index 6b183cbf9f3..ace4f8bcdc9 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -28,3 +28,12 @@ - {id: other.config.overrides.audit_logged, module: other, tier: P1, area: config, assertions: [audit_logged], source: "config_override_endpoints.py:67-100", rationale: "Config override mutations audit-logged, values redacted"} - {id: other.key_mgmt.regenerate.grace_period_honored, module: other, tier: P1, area: auth, assertions: [grace_period_honored], source: "key_management_endpoints.py:4503-4560", rationale: "Old key valid during grace_period then revoked"} - {id: other.key_mgmt.spend_reset.resets_to_value, module: other, tier: P1, area: auth, assertions: [resets_to_value], source: "key_management_endpoints.py:4841", rationale: "reset_spend resets accumulated spend"} +- {id: other.a2a.register.persists, module: other, tier: P1, area: a2a, assertions: [persists], source: "agent_endpoints/endpoints.py:325-443", rationale: "POST /v1/agents registers an agent card; GET /v1/agents/{id} reads it back"} +- {id: other.a2a.register.unsupported_version_rejected, module: other, tier: P1, area: a2a, assertions: [unsupported_version_rejected], source: "agent_endpoints/endpoints.py _validate_protocol_version", rationale: "A card pinning a protocolVersion outside SUPPORTED_A2A_PROTOCOL_VERSIONS is refused with 400"} +- {id: other.a2a.register.semver_version_accepted, module: other, tier: P1, area: a2a, assertions: [semver_version_accepted], source: "agent_endpoints/endpoints.py _validate_protocol_version", rationale: "A card pinning a patch-level semver like 0.3.0 (what the Google A2A SDK emits) registers, stores and serves the canonical 0.3 rather than 400ing; regression guard for the v1.92 report"} +- {id: other.a2a.message_send.real_world_agent_replies, module: other, tier: P1, area: a2a, assertions: [real_world_agent_replies], source: "agent_endpoints/a2a_endpoints.py asend_message", rationale: "A real published a2a agent fetched live from a public /.well-known endpoint (pinning the full semver 0.3.0 the a2a-sdk emits) registers, serves the canonical 0.3, and a message/send skill invocation proxies to the live upstream and returns the agent's reply"} +- {id: other.a2a.register.malformed_version_rejected, module: other, tier: P1, area: a2a, assertions: [malformed_version_rejected], source: "a2a/agent_card.py normalize_protocol_version", rationale: "A malformed protocolVersion like 0.3.garbage fails full-string semver validation and is refused with 400 instead of truncating to a supported family"} +- {id: other.a2a.discovery.proxy_fronted_card, module: other, tier: P1, area: a2a, assertions: [proxy_fronted_card], source: "agent_endpoints/a2a_endpoints.py get_agent_card", rationale: "/.well-known/agent-card.json serves the proxy url + supportedInterfaces and the LiteLLM virtual-key bearer scheme, not the upstream"} +- {id: other.a2a.message_send.bridge_invokes, module: other, tier: P1, area: a2a, assertions: [bridge_invokes], source: "a2a_protocol/litellm_completion_bridge/handler.py", rationale: "A2A message/send routes through the completion bridge to a real provider and logs an asend_message spend row"} +- {id: other.a2a.version.serves_pinned_0_3, module: other, tier: P1, area: a2a, assertions: [serves_pinned_0_3], source: "agent_endpoints/a2a_endpoints.py _served_version", rationale: "An agent pinning 0.3 returns the flat 0.3 message shape (parts on the result)"} +- {id: other.a2a.version.serves_pinned_1_0, module: other, tier: P1, area: a2a, assertions: [serves_pinned_1_0], source: "agent_endpoints/a2a_endpoints.py _served_version", rationale: "An agent pinning 1.0 returns the nested 1.0 message shape (result.message with ROLE_AGENT)"} diff --git a/tests/e2e/coverage_registry/reliability.yaml b/tests/e2e/coverage_registry/reliability.yaml index 1538d3f3cda..ebbfd3415a5 100644 --- a/tests/e2e/coverage_registry/reliability.yaml +++ b/tests/e2e/coverage_registry/reliability.yaml @@ -24,3 +24,4 @@ - {id: reliability.timeout.request_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: request_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions, messages], source: "litellm/router.py:545-551", rationale: "Per-request timeout raises Timeout"} - {id: reliability.timeout.stream_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: stream_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions], source: "litellm/router.py:551", rationale: "Streaming chunk-delivery timeout"} - {id: reliability.perf.throughput.under_slo, module: reliability, tier: P1, behavior: perf, variant: throughput, assertions: [under_slo], exercised_on: [chat_completions, messages], source: grammar, rationale: "Throughput SLO under load"} +- {id: reliability.perf.session_anomaly.under_slo, module: reliability, tier: P1, behavior: perf, variant: session_anomaly, assertions: [under_slo], exercised_on: [messages], source: grammar, rationale: "Weekly Claude Code-shaped multi-turn session load against real providers; ceilings on error rate, warm-turn cache read/write, p95 turn time, and gateway-recorded spend (LIT-4562)"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 3be339d28a0..e7c48690c0a 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -72,12 +72,32 @@ POLL_TIMEOUT = float(os.environ.get("E2E_POLL_TIMEOUT", "120")) POLL_INTERVAL = float(os.environ.get("E2E_POLL_INTERVAL", "5")) REQUEST_TIMEOUT = float(os.environ.get("E2E_REQUEST_TIMEOUT", "60")) +EXPECT_RUST = os.environ.get("E2E_EXPECT_RUST", "").strip().lower() in ("1", "true", "yes") + LOAD_USERS = int(os.environ.get("E2E_LOAD_USERS", "750")) LOAD_SPAWN_RATE = float(os.environ.get("E2E_LOAD_SPAWN_RATE", "50")) LOAD_DURATION_SECONDS = float(os.environ.get("E2E_LOAD_DURATION_SECONDS", "60")) LOAD_MIN_RPS = float(os.environ.get("E2E_LOAD_MIN_RPS", "355")) LOAD_MAX_FAILURE_RATIO = float(os.environ.get("E2E_LOAD_MAX_FAILURE_RATIO", "0.01")) +WEEKLY_ANOMALY_OPT_IN_ENV = "E2E_WEEKLY_ANOMALY" +ANOMALY_SESSIONS = int(os.environ.get("E2E_ANOMALY_SESSIONS", "6")) +ANOMALY_TURNS_PER_SESSION = int(os.environ.get("E2E_ANOMALY_TURNS_PER_SESSION", "6")) +ANOMALY_TURN_ATTEMPTS = int(os.environ.get("E2E_ANOMALY_TURN_ATTEMPTS", "3")) +ANOMALY_MAX_ERROR_RATIO = float(os.environ.get("E2E_ANOMALY_MAX_ERROR_RATIO", "0.05")) +ANOMALY_MIN_WARM_CACHE_READ_SHARE = float( + os.environ.get("E2E_ANOMALY_MIN_WARM_CACHE_READ_SHARE", "0.65") +) +ANOMALY_MAX_P95_TURN_SECONDS = float( + os.environ.get("E2E_ANOMALY_MAX_P95_TURN_SECONDS", "30") +) +ANOMALY_MAX_KEY_SPEND_USD = float( + os.environ.get("E2E_ANOMALY_MAX_KEY_SPEND_USD", "0.60") +) +ANOMALY_SPEND_SETTLE_SECONDS = float( + os.environ.get("E2E_ANOMALY_SPEND_SETTLE_SECONDS", "75") +) + def require_env(*names: str) -> tuple[str, ...]: """Return the non-empty values for each env name, or hard-fail naming which are missing. diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 1c6048a3688..03d7b5d051a 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -62,6 +62,7 @@ R = TypeVar("R", bound=BaseModel) class Success(BaseModel, Generic[R]): kind: Literal["success"] = "success" + status_code: int data: R @@ -146,6 +147,32 @@ class StreamingResponse(BaseModel): return "text/event-stream" in (self.content_type or "") +class BinaryStream(BaseModel): + """Outcome of consuming a binary chunked response (e.g. TTS audio) as a stream. + + Unlike StreamingResponse, which line-splits an SSE text body, this iterates the + raw bytes with iter_content and reports how many non-empty chunks arrived and + the total byte count, so a caller can assert customer-observable streaming + (multiple chunks, real bytes) without decoding the payload.""" + + status_code: int + content_type: str | None = None + call_id: str | None = None + transfer_encoding: str | None = None + content_length: str | None = None + error_body: str | None = None + chunk_count: int = 0 + total_bytes: int = 0 + + @property + def ok(self) -> bool: + return 200 <= self.status_code < 300 + + @property + def chunked(self) -> bool: + return "chunked" in (self.transfer_encoding or "") + + def _hdr(resp: requests.Response, name: str) -> str | None: value = resp.headers.get(name) return value if isinstance(value, str) else None @@ -159,6 +186,18 @@ def unwrap[R: BaseModel](result: Result[R]) -> R: raise AssertionError(result) +def unwrap_status[R: BaseModel](result: Result[R], expected_status: int) -> R: + """Like unwrap, but also pins the exact HTTP status the success came back on, + for routes whose contract is a specific 2xx (e.g. 201 Created on a submission).""" + match result: + case Success(status_code=status_code, data=data) if status_code == expected_status: + return data + case Success(status_code=status_code): + raise AssertionError(f"expected HTTP {expected_status}, got {status_code}") + case _: + raise AssertionError(result) + + def is_ok[R: BaseModel](result: Result[R]) -> bool: match result: case Success(): @@ -199,7 +238,7 @@ def _classify[R: BaseModel]( if not resp.ok: return UnknownApiError(status_code=resp.status_code, body=resp.text) try: - return Success(data=response_type.model_validate(resp.json())) + return Success(status_code=resp.status_code, data=response_type.model_validate(resp.json())) except Exception as exc: # noqa: BLE001 - any parse/validation failure is a value return ValidationError(message=str(exc)) @@ -244,6 +283,26 @@ def get[R: BaseModel]( return _classify(resp, response_type) +def get_external[R: BaseModel]( + url: str, + *, + response_type: type[R], + timeout: float = 30.0, +) -> Result[R]: + """GET an absolute URL outside the proxy (e.g. a public /.well-known document). + Unlike the transport wrappers there is no proxy base url and no proxy auth; the + response still gets the same tagged-union classification as every other call.""" + try: + resp = requests.get( + url, + headers={"Accept": "application/json"}, + timeout=timeout, + ) + except requests.RequestException as exc: + return NetworkError(message=str(exc)) + return _classify(resp, response_type) + + def delete[R: BaseModel]( url: URL, *, @@ -286,6 +345,26 @@ def patch[R: BaseModel]( return _classify(resp, response_type) +def put[R: BaseModel]( + url: URL, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + timeout: float = 30.0, +) -> Result[R]: + try: + resp = requests.put( + str(url), + headers=_headers(headers), + json=json.model_dump(by_alias=True, exclude_none=True), + timeout=timeout, + ) + except requests.RequestException as exc: + return NetworkError(message=str(exc)) + return _classify(resp, response_type) + + def probe( url: URL, *, headers: BaseModel, params: BaseModel, timeout: float = 30.0 ) -> ProbeResult: @@ -397,16 +476,18 @@ def upload[R: BaseModel]( url: URL, *, headers: BaseModel, - form: FileUploadForm, + form: BaseModel, filename: str, content: bytes, + file_content_type: str = "application/jsonl", params: BaseModel | None = None, response_type: type[R], timeout: float = 60.0, ) -> Result[R]: - """Multipart POST for file uploads (/v1/files). Form fields come from `form`, - the file bytes are sent as the `file` part, and `params` carries any query - routing (e.g. ?model=). requests sets the multipart Content-Type itself.""" + """Multipart POST for file-bearing routes (/v1/files, /v1/audio/transcriptions). + Form fields come from `form`, the file bytes are sent as the `file` part with + `file_content_type`, and `params` carries any query routing (e.g. ?model=). + requests sets the multipart Content-Type itself.""" dumped: dict[str, object] = form.model_dump(by_alias=True, exclude_none=True) data = {key: str(value) for key, value in dumped.items()} try: @@ -415,7 +496,7 @@ def upload[R: BaseModel]( headers=_headers(headers), params=_params(params), data=data, - files={"file": (filename, content, "application/jsonl")}, + files={"file": (filename, content, file_content_type)}, timeout=timeout, ) except requests.RequestException as exc: @@ -423,6 +504,54 @@ def upload[R: BaseModel]( return _classify(resp, response_type) +def stream_binary( + url: URL, + *, + headers: BaseModel, + json: BaseModel, + chunk_size: int = 8192, + timeout: float = 60.0, +) -> BinaryStream: + """POST that consumes a binary chunked response (e.g. TTS audio) as a stream, + counting non-empty chunks and total bytes with iter_content. A non-2xx status + short-circuits with the counts left at zero so the caller can fail loudly.""" + try: + resp = requests.post( + str(url), + headers=_headers(headers), + json=json.model_dump(by_alias=True, exclude_none=True), + stream=True, + timeout=timeout, + ) + except requests.RequestException as exc: + return BinaryStream(status_code=-1, error_body=str(exc)[:300]) + with resp: + content_type = _hdr(resp, "content-type") + call_id = _hdr(resp, "x-litellm-call-id") + transfer_encoding = _hdr(resp, "transfer-encoding") + content_length = _hdr(resp, "content-length") + if not (200 <= resp.status_code < 300): + return BinaryStream( + status_code=resp.status_code, + content_type=content_type, + call_id=call_id, + transfer_encoding=transfer_encoding, + content_length=content_length, + error_body=resp.text[:300], + ) + raw_chunks = cast("Iterator[bytes]", resp.iter_content(chunk_size=chunk_size)) + chunks = tuple(chunk for chunk in raw_chunks if chunk) + return BinaryStream( + status_code=resp.status_code, + content_type=content_type, + call_id=call_id, + transfer_encoding=transfer_encoding, + content_length=content_length, + chunk_count=len(chunks), + total_bytes=sum(len(chunk) for chunk in chunks), + ) + + def download( url: URL, *, headers: BaseModel, timeout: float = 60.0 ) -> StreamingResponse: diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py index d24cd36c2fd..53f2e4480df 100644 --- a/tests/e2e/guardrails/guardrails_client.py +++ b/tests/e2e/guardrails/guardrails_client.py @@ -10,13 +10,15 @@ from typing import Literal from pydantic import BaseModel -from e2e_config import POLL_INTERVAL, POLL_TIMEOUT +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, unique_marker from e2e_http import NoBody, Result, Success, unwrap +from lifecycle import ResourceManager from models import ( ChatBody, ChatMessage, ChatResponse, KeyGenerateBody, + LiteLLMParamsBody, TeamDeleteBody, TeamInfoParams, TeamInfoResponse, @@ -54,7 +56,33 @@ class BedrockGuardrailParamsBody(GuardrailParamsBase): aws_region_name: str | None = None -GuardrailParamsBody = ContentFilterParamsBody | BedrockGuardrailParamsBody +class OpenAIModerationParamsBody(GuardrailParamsBase): + guardrail: Literal["openai_moderation"] = "openai_moderation" + api_key: str | None = None + model: str | None = None + + +class PresidioParamsBody(GuardrailParamsBase): + guardrail: Literal["presidio"] = "presidio" + presidio_analyzer_api_base: str | None = None + presidio_anonymizer_api_base: str | None = None + # apply_to_output masks PII the model itself emitted, which also makes the + # guardrail run post_call. logging_only masks what the proxy logs. + apply_to_output: bool | None = None + logging_only: bool | None = None + + +class BlockCodeExecutionParamsBody(GuardrailParamsBase): + guardrail: Literal["block_code_execution"] = "block_code_execution" + + +GuardrailParamsBody = ( + ContentFilterParamsBody + | BedrockGuardrailParamsBody + | OpenAIModerationParamsBody + | PresidioParamsBody + | BlockCodeExecutionParamsBody +) class GuardrailSpecBody(BaseModel): @@ -135,6 +163,35 @@ class GuardrailsClient: ) ).guardrail_id + def create_backend_model(self, resources: ResourceManager, prefix: str = "e2e-guard-backend") -> str: + """Register a gemini chat deployment for a guardrail test to run against + (deleted on teardown). The guardrails under test here gate on prompt/output + content, not the backend, so a single cheap deployment stands in for the + model the customer would call.""" + model_name = f"{prefix}-{unique_marker()}" + model_id = self.proxy.create_model( + model_name, + LiteLLMParamsBody(model="gemini/gemini-2.5-flash", api_key="os.environ/GEMINI_API_KEY"), + ) + resources.defer(lambda: self.proxy.delete_model(model_id)) + return model_name + + def register(self, name: str, params: GuardrailParamsBody) -> str: + """Register any guardrail via POST /guardrails and return its id. New + built-ins register with default_on=False and are opted into per request + via the chat body's `guardrails` list, so one guardrail under test never + intercepts unrelated traffic on the shared proxy.""" + return unwrap( + self.proxy.transport.post( + "/guardrails", + headers=self.proxy.transport.master, + json=GuardrailCreateBody( + guardrail=GuardrailSpecBody(guardrail_name=name, litellm_params=params) + ), + response_type=GuardrailCreateResponse, + ) + ).guardrail_id + def delete_guardrail(self, guardrail_id: str) -> None: _ = self.proxy.transport.delete( f"/guardrails/{guardrail_id}", @@ -171,13 +228,27 @@ class GuardrailsClient: KeyGenerateBody(team_id=team_id, user_id="e2e-guardrails-user") ) - def chat(self, key: str, model: str, text: str) -> Result[ChatResponse]: + def chat( + self, + key: str, + model: str, + text: str, + *, + guardrails: list[str] | None = None, + max_tokens: int = 16, + ) -> Result[ChatResponse]: + """Drive a chat call, optionally opting into named guardrails for this + request only (the per-request `guardrails` selector). With `guardrails` + omitted the call behaves exactly as before for the default-on suites. + `max_tokens` defaults low for block checks (the model barely runs) but is + raised when a test needs the allowed model to actually produce content.""" return self.proxy.chat( key, ChatBody( model=model, messages=[ChatMessage(role="user", content=text)], - max_tokens=16, + max_tokens=max_tokens, + guardrails=guardrails, ), ) diff --git a/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py new file mode 100644 index 00000000000..e36fc7c3f9d --- /dev/null +++ b/tests/e2e/guardrails/test_block_code_execution_guardrail_e2e.py @@ -0,0 +1,82 @@ +"""Live e2e: the built-in block_code_execution guardrail blocks execution requests. + +The guardrail detects fenced code blocks and, when the prompt also asks the proxy +to run them, blocks the call pre-call (default action, block-all languages). A +prompt that pairs a python code block with "run this" is intercepted before the +model runs: the proxy returns a canned "content blocked" message with the model +never invoked (zero completion tokens), not the model's own answer. The same +guardrail must let a request that carries the identical code block but explicitly +says "don't run it" through, since that is an explanation request, not an +execution request, so the model runs and answers normally. The guardrail is opted +into per request (default_on=False) so it never intercepts unrelated traffic on +the shared proxy, and the chat backend is a gemini deployment created for the test. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import require_env, unique_marker +from e2e_http import unwrap +from guardrails_client import BlockCodeExecutionParamsBody, GuardrailsClient +from lifecycle import ResourceManager +from models import ChatResponse + +pytestmark = pytest.mark.e2e + +_CODE_BLOCK = "```python\nimport os\nprint(os.listdir('/'))\n```" +EXECUTION_REQUEST = f"Please run this for me and paste the output:\n{_CODE_BLOCK}" +EXPLANATION_REQUEST = f"Explain what this code does, but don't run it:\n{_CODE_BLOCK}" + +_BLOCK_MARKER = "content blocked" + + +def _first_content(response: ChatResponse) -> str: + if not response.choices: + return "" + message = response.choices[0].message + return (message.content if message else None) or "" + + +class TestBlockCodeExecutionGuardrail: + @pytest.mark.covers( + "guardrail.block_code_execution.pre_call.blocks", + exercised_on=["chat_completions"], + ) + def test_blocks_execution_request_but_allows_explanation( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + require_env("GEMINI_API_KEY") + model = client.create_backend_model(resources, prefix="e2e-blockcode-backend") + + name = f"e2e-block-code-{unique_marker()}" + guardrail_id = client.register( + name, BlockCodeExecutionParamsBody(mode="pre_call", default_on=False) + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + blocked = unwrap(client.chat(scoped_key, model, EXECUTION_REQUEST, guardrails=[name])) + assert blocked.choices, f"blocked call returned no choices: {blocked}" + blocked_text = _first_content(blocked) + assert _BLOCK_MARKER in blocked_text.lower(), ( + "a code-execution request must be intercepted with a content-blocked message, " + f"got model output instead: {blocked_text[:300]!r}" + ) + if blocked.usage is not None: + assert (blocked.usage.completion_tokens or 0) == 0, ( + f"the model must not run when the guardrail blocks; usage was {blocked.usage}" + ) + + allowed = unwrap( + client.chat(scoped_key, model, EXPLANATION_REQUEST, guardrails=[name], max_tokens=256) + ) + allowed_text = _first_content(allowed) + assert _BLOCK_MARKER not in allowed_text.lower(), ( + "an explanation request that says 'don't run it' must not be blocked, but got the " + f"content-blocked message: {allowed_text[:300]!r}" + ) + ran = allowed.usage is not None and (allowed.usage.prompt_tokens or 0) > 0 + assert ran, ( + "the explanation request must reach the model (the guardrail lets it through), but " + f"the model was never invoked; usage was {allowed.usage}" + ) diff --git a/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py new file mode 100644 index 00000000000..4e2fcbf8fba --- /dev/null +++ b/tests/e2e/guardrails/test_openai_moderation_guardrail_e2e.py @@ -0,0 +1,66 @@ +"""Live e2e: the built-in OpenAI moderation guardrail blocks flagged input pre-call. + +Registers an `openai_moderation` guardrail (default_on=False) and opts into it for +one request via the chat body's `guardrails` selector, so it never intercepts +unrelated traffic on the shared proxy. The guardrail calls OpenAI's moderation API +before the upstream model runs; a prompt that trips the policy must be rejected +with HTTP 400 naming the moderation policy, and the same guardrail must let a +benign prompt through. The chat backend is a gemini deployment created for the +test (and torn down); moderation runs independently of it, so the block is +attributable to the guardrail, not the model. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import require_env, unique_marker +from e2e_http import UnknownApiError, unwrap +from guardrails_client import GuardrailsClient, OpenAIModerationParamsBody +from lifecycle import ResourceManager + +pytestmark = pytest.mark.e2e + +# A direct threat of violence reliably trips OpenAI's harassment/violence categories. +FLAGGED_PROMPT = "I am going to find you and violently kill you and your entire family tonight." +BENIGN_PROMPT = "Reply with the single word: hello" + + +class TestOpenAIModerationGuardrail: + @pytest.mark.covers( + "guardrail.openai_moderations.pre_call.blocks", + exercised_on=["chat_completions"], + ) + def test_moderation_blocks_flagged_input( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + require_env("OPENAI_API_KEY", "GEMINI_API_KEY") + model = client.create_backend_model(resources, prefix="e2e-moderation-backend") + + name = f"e2e-openai-moderation-{unique_marker()}" + guardrail_id = client.register( + name, + OpenAIModerationParamsBody( + mode="pre_call", default_on=False, api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + blocked = client.chat(scoped_key, model, FLAGGED_PROMPT, guardrails=[name]) + match blocked: + case UnknownApiError(status_code=400, body=body): + assert "moderation" in body.lower(), ( + f"the block body must name the moderation policy, got: {body[:400]}" + ) + case UnknownApiError(status_code=status, body=body): + pytest.fail(f"expected a 400 moderation block, got {status}: {body[:400]}") + case _: + pytest.fail( + f"openai moderation did not block a flagged prompt; got {blocked}" + ) + + allowed = unwrap(client.chat(scoped_key, model, BENIGN_PROMPT, guardrails=[name])) + assert allowed.choices, ( + "the same moderation guardrail must let a benign prompt through, but the " + f"call returned no choices: {allowed}" + ) diff --git a/tests/e2e/guardrails/test_presidio_guardrail_e2e.py b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py new file mode 100644 index 00000000000..a911f387382 --- /dev/null +++ b/tests/e2e/guardrails/test_presidio_guardrail_e2e.py @@ -0,0 +1,211 @@ +"""Live e2e: the built-in Presidio PII guardrail masks PII on the request, on the +model output, and in what the proxy logs. + +Presidio replaces detected PII with `` placeholders (e.g. +``) via a real analyzer + anonymizer. Three modes are checked +independently, each opted into per request (default_on=False) so it never touches +unrelated traffic: + +- pre_call: the prompt is anonymized before it reaches the model, so a + repeat-verbatim request comes back with the placeholder, never the raw email +- post_call (apply_to_output): PII the model itself emits is masked on the way + out, so the caller never receives the raw value the model produced +- logging_only: the call is not blocked, and the request the proxy records is + masked. That is read back from the real OTEL destination (Jaeger): the gen-AI + span's `gen_ai.input.messages` attribute carries the masked placeholder, never + the raw email + +Analyzer/anonymizer bases come from PRESIDIO_ANALYZER_API_BASE / +PRESIDIO_ANONYMIZER_API_BASE (compose provides the in-network hosts; point them at +locally published container ports for a host run). The logging_only check needs +the OTEL v2 logger active and its destination readable at OTEL_QUERY_URL, with +message-content capture on (OTEL_INSTRUMENTATION_GENAI_CAPTURE_MESSAGE_CONTENT). +The chat backend is a gemini deployment created for the test. +""" + +from __future__ import annotations + +import time + +import pytest + +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT, require_env, unique_marker +from e2e_http import NoBody, require_successful_call, unwrap +from guardrails_client import GuardrailMode, GuardrailsClient, PresidioParamsBody +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, ChatResponse, ReadinessDetailsResponse +from otel_client import JaegerSpan, OtelReader, build_otel_reader + +pytestmark = pytest.mark.e2e + +RAW_EMAIL = "alice.example.person@example.com" +PLACEHOLDER = "" + +ECHO_REQUEST = f"Repeat the following text back exactly, verbatim, with no changes: My email is {RAW_EMAIL}" +EMIT_REQUEST = f"Output exactly this one line and nothing else: Please contact {RAW_EMAIL} today" +LOG_REQUEST = f"Say hello and include this email once verbatim: {RAW_EMAIL}" + +OTEL_V2_LOGGER = "OpenTelemetryV2" +INPUT_MESSAGES_TAG = "gen_ai.input.messages" + + +def _content(response: ChatResponse) -> str: + if not response.choices: + return "" + message = response.choices[0].message + return (message.content if message else None) or "" + + +def _span_tag(span: JaegerSpan, key: str) -> str | None: + for tag in span.tags: + if tag.key == key and isinstance(tag.value, str): + return tag.value + return None + + +def _poll_logged_prompt(reader: OtelReader, *, call_id: str, genai_span: str) -> str | None: + """Poll the OTEL destination until the call's gen-AI span carries a masked + logged prompt, and return it. logging_only masks the payload asynchronously, + so the span can briefly export before the mask lands; polling to a deadline + waits that out and returns the last value seen so the caller's assertions + report the real final state if it never masks.""" + deadline = time.monotonic() + POLL_TIMEOUT + last: str | None = None + while time.monotonic() < deadline: + for trace in reader.traces_for_call(call_id): + for span in trace.spans: + if span.operation_name != genai_span: + continue + value = _span_tag(span, INPUT_MESSAGES_TAG) + if value is not None: + last = value + if PLACEHOLDER in value and RAW_EMAIL not in value: + return value + time.sleep(POLL_INTERVAL) + return last + + +def _presidio_params( + mode: GuardrailMode, *, apply_to_output: bool = False, logging_only: bool = False +) -> PresidioParamsBody: + analyzer, anonymizer = require_env( + "PRESIDIO_ANALYZER_API_BASE", "PRESIDIO_ANONYMIZER_API_BASE" + ) + return PresidioParamsBody( + mode=mode, + default_on=False, + presidio_analyzer_api_base=analyzer, + presidio_anonymizer_api_base=anonymizer, + apply_to_output=apply_to_output, + logging_only=logging_only, + ) + + +def _require_otel_v2_active(client: GuardrailsClient) -> None: + details = unwrap( + client.proxy.transport.get( + "/health/readiness/details", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=ReadinessDetailsResponse, + ) + ) + assert OTEL_V2_LOGGER in details.success_callbacks, ( + f"the logging_only check reads the masked prompt back from OTEL, so the proxy must have " + f"the {OTEL_V2_LOGGER} logger active; got callbacks: {details.success_callbacks}" + ) + + +class TestPresidioGuardrail: + @pytest.mark.covers( + "guardrail.presidio.pre_call.masks", + exercised_on=["chat_completions"], + ) + def test_pre_call_masks_pii_before_the_model_sees_it( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + require_env("GEMINI_API_KEY") + model = client.create_backend_model(resources, prefix="e2e-presidio-pre") + name = f"e2e-presidio-pre-{unique_marker()}" + guardrail_id = client.register(name, _presidio_params("pre_call")) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + echoed = _content( + unwrap(client.chat(scoped_key, model, ECHO_REQUEST, guardrails=[name], max_tokens=128)) + ) + assert RAW_EMAIL not in echoed, ( + "pre_call masking must strip the raw email before the model sees it, but the " + f"model echoed it back: {echoed[:300]!r}" + ) + assert PLACEHOLDER in echoed, ( + "the model should have echoed the masked placeholder the guardrail substituted, " + f"got: {echoed[:300]!r}" + ) + + @pytest.mark.covers( + "guardrail.presidio.post_call.masks", + exercised_on=["chat_completions"], + ) + def test_post_call_masks_pii_in_model_output( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + require_env("GEMINI_API_KEY") + model = client.create_backend_model(resources, prefix="e2e-presidio-post") + name = f"e2e-presidio-post-{unique_marker()}" + guardrail_id = client.register(name, _presidio_params("post_call", apply_to_output=True)) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + out = _content( + unwrap(client.chat(scoped_key, model, EMIT_REQUEST, guardrails=[name], max_tokens=128)) + ) + assert RAW_EMAIL not in out, ( + "post_call masking must strip PII the model emitted, but the raw email reached the " + f"caller: {out[:300]!r}" + ) + assert PLACEHOLDER in out, ( + f"the masked placeholder should replace the model's PII output, got: {out[:300]!r}" + ) + + @pytest.mark.covers( + "guardrail.presidio.logging_only.masks", + exercised_on=["chat_completions"], + ) + def test_logging_only_masks_the_logged_prompt( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + require_env("GEMINI_API_KEY") + _require_otel_v2_active(client) + reader = build_otel_reader() + + model = client.create_backend_model(resources, prefix="e2e-presidio-log") + name = f"e2e-presidio-log-{unique_marker()}" + guardrail_id = client.register(name, _presidio_params("logging_only", logging_only=True)) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + outcome = client.proxy.transport.send( + "/chat/completions", + headers=client.proxy.transport.bearer(scoped_key), + json=ChatBody( + model=model, + messages=[ChatMessage(role="user", content=LOG_REQUEST)], + max_tokens=64, + guardrails=[name], + ), + ) + require_successful_call(outcome) # logging_only must not block + assert outcome.call_id is not None, "the response must carry x-litellm-call-id to find its trace" + + genai_span = f"chat {model}" + logged_prompt = _poll_logged_prompt(reader, call_id=outcome.call_id, genai_span=genai_span) + assert logged_prompt is not None, ( + f"the gen-AI span {genai_span!r} never recorded {INPUT_MESSAGES_TAG} at the OTEL " + "destination within the deadline (message-content capture must be on, and the trace " + "must reach the destination)" + ) + assert RAW_EMAIL not in logged_prompt, ( + "logging_only must mask the PII the proxy records for the request, but the raw email " + f"is present in the logged prompt: {logged_prompt[:400]!r}" + ) + assert PLACEHOLDER in logged_prompt, ( + f"the logged prompt must carry the masked placeholder, got: {logged_prompt[:400]!r}" + ) diff --git a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py index cd32a19d54a..db917d6ede9 100644 --- a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py +++ b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py @@ -8,6 +8,8 @@ suite was removed. from __future__ import annotations +import time + import pytest from e2e_config import unique_marker @@ -19,11 +21,39 @@ pytestmark = pytest.mark.e2e MODEL = "gemini-2.5-flash" +# A guardrail created via POST /guardrails is registered in-process immediately +# on the worker that served the create call, but the proxy runs multiple +# pods/workers behind the shared key, and every other one only picks up the new +# guardrail on its next periodic DB sync (every 30s), so the very next request +# can race a worker that has not synced yet. +GUARDRAIL_PROPAGATION_DEADLINE_SECONDS = 40.0 +GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS = 5.0 + def _prompt_with(banned_keyword: str) -> str: return f"Reply with the single word OK. {banned_keyword}" +def _assert_eventually_blocked(client: GuardrailsClient, key: str, banned: str) -> None: + deadline = time.monotonic() + GUARDRAIL_PROPAGATION_DEADLINE_SECONDS + while True: + result = client.chat(key, MODEL, _prompt_with(banned)) + match result: + case UnknownApiError(status_code=status, body=body): + assert status == 400, f"expected a 400 guardrail block, got {status}: {body[:300]}" + assert "content blocked" in body.lower() or banned in body, ( + f"block response missing content-filter reason: {body[:300]}" + ) + return + case _ if time.monotonic() < deadline: + time.sleep(GUARDRAIL_PROPAGATION_POLL_INTERVAL_SECONDS) + case _: + pytest.fail( + f"default-on guardrail never blocked the banned keyword within " + f"{GUARDRAIL_PROPAGATION_DEADLINE_SECONDS}s; got {result}" + ) + + class TestTeamDisableGlobalGuardrail: @pytest.mark.covers( "guardrail.litellm_content_filter.pre_call.blocks", @@ -33,25 +63,10 @@ class TestTeamDisableGlobalGuardrail: self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str ) -> None: banned = unique_marker() - guardrail_id = client.create_content_filter_guardrail( - f"e2e-content-filter-{banned}", banned - ) + guardrail_id = client.create_content_filter_guardrail(f"e2e-content-filter-{banned}", banned) resources.defer(lambda: client.delete_guardrail(guardrail_id)) - result = client.chat(scoped_key, MODEL, _prompt_with(banned)) - - match result: - case UnknownApiError(status_code=status, body=body): - assert status == 400, ( - f"expected a 400 guardrail block, got {status}: {body[:300]}" - ) - assert "content blocked" in body.lower() or banned in body, ( - f"block response missing content-filter reason: {body[:300]}" - ) - case _: - pytest.fail( - f"default-on guardrail did not block the banned keyword; got {result}" - ) + _assert_eventually_blocked(client, scoped_key, banned) @pytest.mark.covers( "guardrail.litellm_content_filter.pre_call.allows", @@ -61,14 +76,10 @@ class TestTeamDisableGlobalGuardrail: self, client: GuardrailsClient, resources: ResourceManager ) -> None: banned = unique_marker() - guardrail_id = client.create_content_filter_guardrail( - f"e2e-content-filter-{banned}", banned - ) + guardrail_id = client.create_content_filter_guardrail(f"e2e-content-filter-{banned}", banned) resources.defer(lambda: client.delete_guardrail(guardrail_id)) - team_id = client.create_team_opted_out_of_global_guardrails( - f"e2e-guardrail-optout-{banned}" - ) + team_id = client.create_team_opted_out_of_global_guardrails(f"e2e-guardrail-optout-{banned}") resources.defer(lambda: client.delete_team(team_id)) key = client.create_key_in_team(team_id) resources.defer(lambda: client.proxy.delete_key(key)) diff --git a/tests/e2e/lifecycle.py b/tests/e2e/lifecycle.py index c0de074f0d2..4ef25509905 100644 --- a/tests/e2e/lifecycle.py +++ b/tests/e2e/lifecycle.py @@ -1,13 +1,11 @@ -"""Lifecycle contract and resource cleanup for stateful e2e tests. +"""Resource cleanup for stateful e2e tests. Shared by every e2e suite under tests/e2e/. The proxy under test is long-lived and never reset between tests, so anything a test creates (keys, customers, teams, orgs, users, guardrails, budgets, ...) persists unless -explicitly deleted. Every check follows an init -> run -> teardown lifecycle; -teardown releases each resource init() created, even when run() raises. - -In pytest terms (see conftest.py): the `resources` fixture's setup is init(), -the test body is run(), and the fixture's teardown is teardown(). +explicitly deleted. The `resources` fixture (see conftest.py) hands each test a +ResourceManager; the test registers a cleanup for every resource it creates, and +the fixture's teardown releases them all even when the test body raises. """ from dataclasses import dataclass, field @@ -17,38 +15,6 @@ from proxy_client import ProxyClient from models import KeyGenerateBody -@runtime_checkable -class E2ECase(Protocol): - """A stateful e2e check run against a long-lived proxy. - - init() acquires resources, run() exercises behaviour and asserts, teardown() - releases everything init() created. teardown() must run even if init() fails - partway or run() raises. - """ - - def init(self) -> None: ... - - def run(self) -> None: ... - - def teardown(self) -> None: ... - - -def run_case(case: E2ECase) -> None: - """Drive a case through its lifecycle: init -> run -> teardown. - - teardown always runs - even when init() fails partway or run() raises (or - skips) - so resources the case already registered on the long-lived proxy are - released. init() is inside the try because cases register cleanups - progressively (e.g. create team, then user, then key), and a failure after - the first creation must still release what came before. - """ - try: - case.init() - case.run() - finally: - case.teardown() - - @runtime_checkable class ResourceClient(Protocol): """Proxy operations the convenience creators use. Resource types without a diff --git a/tests/e2e/llm_translation/endpoints_client.py b/tests/e2e/llm_translation/endpoints_client.py index 32d9922c775..ace621d03b3 100644 --- a/tests/e2e/llm_translation/endpoints_client.py +++ b/tests/e2e/llm_translation/endpoints_client.py @@ -15,7 +15,7 @@ from typing import Literal from pydantic import BaseModel from proxy_client import ProxyClient -from e2e_http import StreamingResponse +from e2e_http import BinaryStream, Result, StreamingResponse from models import CacheControl, ChatMessage, LiteLLMParamsBody, RichMessage, TextBlock __all__ = [ @@ -110,6 +110,16 @@ class ImageRequest(BaseModel): size: str = "1024x1024" +class TranscriptionForm(BaseModel): + model: str + response_format: str = "json" + + +class ModerationRequest(BaseModel): + model: str + input: str + + class ResponsesOutputContent(BaseModel): type: str | None = None text: str | None = None @@ -213,6 +223,27 @@ class ImagesResult(BaseModel): data: list[ImageItem] = [] +class TranscriptionResult(BaseModel): + text: str = "" + + +class ModerationResultItem(BaseModel): + flagged: bool + categories: dict[str, bool] = {} + + @property + def flagged_categories(self) -> tuple[str, ...]: + return tuple(name for name, hit in self.categories.items() if hit) + + +class ModerationResult(BaseModel): + results: list[ModerationResultItem] = [] + + @property + def first(self) -> ModerationResultItem | None: + return self.results[0] if self.results else None + + @dataclass(frozen=True, slots=True) class EndpointsClient: proxy: ProxyClient @@ -314,6 +345,36 @@ class EndpointsClient: "/v1/audio/speech", key, SpeechRequest(model=model, input=text, voice=voice) ) + def audio_speech_stream( + self, key: str, model: str, text: str, *, voice: str = "alloy" + ) -> BinaryStream: + return self.proxy.transport.stream_binary( + "/v1/audio/speech", + headers=self.proxy.transport.bearer(key), + json=SpeechRequest(model=model, input=text, voice=voice), + ) + + def transcribe( + self, key: str, model: str, *, filename: str, content: bytes + ) -> Result[TranscriptionResult]: + return self.proxy.transport.upload( + "/v1/audio/transcriptions", + headers=self.proxy.transport.bearer(key), + form=TranscriptionForm(model=model), + filename=filename, + content=content, + file_content_type="audio/wav", + response_type=TranscriptionResult, + ) + + def moderations(self, key: str, model: str, text: str) -> Result[ModerationResult]: + return self.proxy.transport.post( + "/v1/moderations", + headers=self.proxy.transport.bearer(key), + json=ModerationRequest(model=model, input=text), + response_type=ModerationResult, + ) + def images(self, key: str, model: str, prompt: str) -> StreamingResponse: return self._send( "/v1/images/generations", key, ImageRequest(model=model, prompt=prompt) diff --git a/tests/e2e/llm_translation/realtime/test_nova_sonic_realtime_e2e.py b/tests/e2e/llm_translation/realtime/test_realtime_bedrock_e2e.py similarity index 100% rename from tests/e2e/llm_translation/realtime/test_nova_sonic_realtime_e2e.py rename to tests/e2e/llm_translation/realtime/test_realtime_bedrock_e2e.py diff --git a/tests/e2e/llm_translation/test_audio_speech_e2e.py b/tests/e2e/llm_translation/test_audio_speech_e2e.py index f7a04d94cb3..b95cef8db4d 100644 --- a/tests/e2e/llm_translation/test_audio_speech_e2e.py +++ b/tests/e2e/llm_translation/test_audio_speech_e2e.py @@ -1,8 +1,9 @@ -"""Live e2e: POST /v1/audio/speech returns audio. +"""Live e2e: POST /v1/audio/speech returns audio, non-streamed and streamed. -Registers an OpenAI text-to-speech deployment at runtime and asserts the response -is an audio body (binary, not JSON). Migrated from -litellm-regression-tests/tests/test_inference_endpoints.py. +The non-streamed call asserts an audio (not JSON) body. The streamed call consumes +the response the way a player would and asserts customer-observable streaming: +chunked transfer encoding (a buffered body would carry a content-length) with +non-zero audio bytes. """ from __future__ import annotations @@ -19,6 +20,7 @@ pytestmark = pytest.mark.e2e class TestAudioSpeech: + @pytest.mark.covers("llm.audio_speech.openai.basic.nonstream.works") def test_audio_speech_returns_audio( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: @@ -38,3 +40,39 @@ class TestAudioSpeech: f"/audio/speech content-type is not audio: {result.content_type!r}" ) assert result.body, "/audio/speech returned an empty body" + + @pytest.mark.covers("llm.audio_speech.openai.basic.stream.works") + def test_audio_speech_streams_audio_chunks( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-speech-stream-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/gpt-4o-mini-tts", api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.audio_speech_stream( + key, + model, + "Streaming speech should arrive in several audio chunks so a client can " + "begin playback well before the whole clip has finished generating.", + ) + assert result.ok, ( + f"/audio/speech stream failed (status {result.status_code}); body={result.error_body}" + ) + assert "audio" in (result.content_type or ""), ( + f"/audio/speech content-type is not audio: {result.content_type!r}" + ) + assert result.chunked, ( + f"/audio/speech did not stream: transfer-encoding={result.transfer_encoding!r}, " + f"content-length={result.content_length!r} (a buffered body is not a stream)" + ) + assert result.content_length is None, ( + f"/audio/speech advertised content-length={result.content_length!r} on a " + f"streamed response (a buffered body is not a stream)" + ) + assert result.total_bytes > 0, "/audio/speech stream returned no audio bytes" diff --git a/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py b/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py new file mode 100644 index 00000000000..af6123dc46a --- /dev/null +++ b/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py @@ -0,0 +1,51 @@ +"""Live e2e: POST /v1/audio/transcriptions turns speech into text. + +Registers an OpenAI speech-to-text deployment at runtime and uploads a spoken +weather question (the realtime suite's 24kHz WAV fixture) as multipart, asserting +the returned transcript is non-empty and mentions the word it was asked about. +""" + +from __future__ import annotations + +from pathlib import Path + +import pytest + +from e2e_config import unique_marker +from e2e_http import unwrap +from endpoints_client import EndpointsClient +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + +WEATHER_WAV = ( + Path(__file__).resolve().parent / "realtime" / "fixtures" / "weather_question_24k.wav" +) + + +class TestAudioTranscriptions: + @pytest.mark.covers("llm.audio_transcriptions.openai.basic.nonstream.works") + def test_audio_transcriptions_returns_text( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-transcribe-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/gpt-4o-mini-transcribe", api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = unwrap( + endpoints_client.transcribe( + key, model, filename=WEATHER_WAV.name, content=WEATHER_WAV.read_bytes() + ) + ) + text = result.text.strip() + assert text, "/audio/transcriptions returned an empty transcript" + assert "weather" in text.lower(), ( + f"transcript of a spoken weather question does not mention weather: {text!r}" + ) diff --git a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py index 13992744f42..8d3622e441a 100644 --- a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py +++ b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py @@ -19,17 +19,176 @@ from __future__ import annotations import os import pytest +from pydantic import BaseModel from e2e_config import require_env, unique_marker -from e2e_http import unwrap +from e2e_http import StreamingResponse, unwrap from lifecycle import ResourceManager -from models import ChatBody, ChatMessage, LiteLLMParamsBody +from models import ( + ChatBody, + ChatMessage, + ChatResponse, + ChatTool, + ChatToolFunction, + ImageContentPart, + ImageUrl, + LiteLLMParamsBody, + TextContentPart, + ThinkingParam, +) from passthrough_client import PassthroughClient pytestmark = pytest.mark.e2e COHERE_BACKEND = "cohere/command-r-08-2024" GEMINI_BACKEND = "gemini/gemini-2.5-flash" +OPENAI_BACKEND = "openai/gpt-5.6" +BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" + + +class _StreamToolCallFunction(BaseModel): + name: str | None = None + arguments: str | None = None + + +class _StreamToolCall(BaseModel): + function: _StreamToolCallFunction = _StreamToolCallFunction() + + +class _StreamDelta(BaseModel): + content: str | None = None + tool_calls: list[_StreamToolCall] | None = None + + +class _StreamChoice(BaseModel): + delta: _StreamDelta = _StreamDelta() + + +class _StreamChunk(BaseModel): + choices: list[_StreamChoice] = [] + + +def _streamed_tool_call(events: list[str]) -> tuple[str, str]: + """Reassemble the tool call streamed across chunks: the name arrives once and the + arguments arrive as fragments, so concatenating both and parsing the arguments as + JSON catches a stream that never completes the call or splits its argument JSON.""" + chunks = [_StreamChunk.model_validate_json(event) for event in events] + calls = [call for chunk in chunks for choice in chunk.choices for call in (choice.delta.tool_calls or [])] + name = "".join(call.function.name or "" for call in calls) + arguments = "".join(call.function.arguments or "" for call in calls) + return name, arguments + + +CAT_IMAGE_URL = "https://upload.wikimedia.org/wikipedia/commons/3/3a/Cat03.jpg" +OPENAI_VISION_BACKEND = "openai/gpt-4o" + +# OpenAI caches a shared prompt prefix once it exceeds ~1024 tokens; this is well +# past that, so a repeat call reports cached prompt tokens. +CACHE_PREFIX = ( + "You are a meticulous assistant. Follow these standing instructions exactly. " + * 300 +) + + +def _vision_messages() -> list[ChatMessage]: + return [ + ChatMessage( + role="user", + content=[ + TextContentPart(text="What animal is in this image? Answer in one word."), + ImageContentPart(image_url=ImageUrl(url=CAT_IMAGE_URL)), + ], + ) + ] + + +def _assert_describes_cat(response: ChatResponse) -> None: + assert response.choices, f"vision returned no choices: {response}" + message = response.choices[0].message + content = (message.content if message else None) or "" + assert "cat" in content.lower() or "feline" in content.lower(), ( + f"vision response did not describe the image: {content[:200]}" + ) + + +def _streamed_text(events: list[str]) -> str: + """Concatenate the delta content across streamed chunks. Parsing every event as + JSON also fails loudly on a truncated or garbled chunk (the vertex/gemini image + streaming regression class), so an incomplete stream cannot pass as content.""" + chunks = [_StreamChunk.model_validate_json(event) for event in events] + return "".join(choice.delta.content or "" for chunk in chunks for choice in chunk.choices) + + +def _assert_streamed_completion(result: StreamingResponse) -> None: + """A streamed /chat/completions must deliver real content, not a clean-but-empty + stream (the #28991 class on the streaming path).""" + assert result.ok and result.is_streaming, f"stream was not established: {result}" + assert result.stream_error is None, f"stream carried an error event: {result.stream_error}" + assert len(result.stream_events) > 1, f"stream did not deliver multiple data events: {result}" + assert _streamed_text(result.stream_events).strip(), ( + f"stream completed with no content deltas: {result.stream_events[:3]}" + ) + + +def _bedrock_params() -> LiteLLMParamsBody: + return LiteLLMParamsBody( + model=BEDROCK_CONVERSE_BACKEND, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ) + + +class _WeatherArgs(BaseModel): + location: str + + +_WEATHER_TOOL = ChatTool( + function=ChatToolFunction( + name="get_weather", + description="Get the current weather for a location", + parameters={ + "type": "object", + "properties": {"location": {"type": "string"}}, + "required": ["location"], + }, + ) +) + + +def _assert_weather_tool_call(response: ChatResponse) -> None: + """The model, forced to call the tool, must return a get_weather call whose + arguments parse as JSON and carry a location. A regression that drops tool_calls + or emits malformed argument JSON fails here rather than passing on a 200.""" + assert response.choices, f"chat returned no choices: {response}" + message = response.choices[0].message + calls = message.tool_calls if message else None + assert calls, f"model returned no tool call for a tool-forced prompt: {response}" + weather = next((call for call in calls if call.function.name == "get_weather"), None) + assert weather is not None, f"expected a get_weather call, got {[c.function.name for c in calls]}" + assert weather.function.arguments, f"get_weather call carried no arguments: {weather}" + args = _WeatherArgs.model_validate_json(weather.function.arguments) + assert args.location.strip(), f"get_weather arguments missing location: {weather.function.arguments}" + + +class _Person(BaseModel): + name: str + age: int + + +_PERSON_SCHEMA: dict[str, object] = { + "type": "json_schema", + "json_schema": { + "name": "person", + "strict": True, + "schema": { + "type": "object", + "properties": {"name": {"type": "string"}, "age": {"type": "integer"}}, + "required": ["name", "age"], + "additionalProperties": False, + }, + }, +} CHAT_MODELS: tuple[tuple[str, str], ...] = ( ("gpt-5.5", "openai"), @@ -219,3 +378,391 @@ class TestHostedVllmChat: assert response.choices, f"hosted_vllm chat returned no choices: {response}" content = response.choices[0].message.content if response.choices[0].message else None assert content and content.strip(), f"hosted_vllm empty content: {response}" + + +class TestOpenAIChatCompletions: + """OpenAI /chat/completions, the SDK path the customer runs against the proxy. + + The streamed call must deliver real content deltas (a clean-but-empty stream is + the regression), and a non-streamed call must be costed so per-request spend and + the response-cost header stay accurate. + """ + + @pytest.mark.covers( + "llm.chat_completions.openai.basic.stream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_streams_real_content( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-chat-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + result = client.proxy.chat_stream( + key, + ChatBody( + model=model, + messages=[ + ChatMessage(role="user", content=f"Count from 1 to 5, one number per line. {unique_marker()}") + ], + max_tokens=64, + stream=True, + ), + ) + _assert_streamed_completion(result) + + @pytest.mark.covers( + "llm.chat_completions.openai.basic.nonstream.cost_logged", + exercised_on=["chat_completions"], + ) + def test_openai_chat_logs_cost( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-cost-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=f"Reply with the single word pong. {unique_marker()}")], + max_tokens=16, + ), + ) + ) + assert response.choices, f"openai chat returned no choices: {response}" + + rows = client.proxy.poll_logs_for_key( + key, min_rows=1, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs) + ) + priced = [r for r in rows if (r.spend or 0) > 0] + assert priced, f"openai chat was not costed on key ...{key[-6:]}: {rows}" + assert priced[0].status == "success", f"openai chat spend status={priced[0].status!r}" + + @pytest.mark.covers( + "llm.chat_completions.openai.tool_use.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_returns_tool_call( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-tool-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage(role="user", content="What is the weather in San Francisco? Use the get_weather tool.") + ], + tools=[_WEATHER_TOOL], + tool_choice="required", + max_tokens=128, + ), + ) + ) + _assert_weather_tool_call(response) + + @pytest.mark.covers( + "llm.chat_completions.openai.structured_output.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_structured_output_conforms_to_schema( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-schema-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content="Extract the person. John Doe is 42 years old.")], + response_format=_PERSON_SCHEMA, + max_tokens=128, + ), + ) + ) + assert response.choices, f"structured output returned no choices: {response}" + content = response.choices[0].message.content if response.choices[0].message else None + assert content, f"structured output returned empty content: {response}" + person = _Person.model_validate_json(content) + assert person.name.strip() and person.age == 42, ( + f"schema-constrained extraction was wrong: {person}" + ) + + @pytest.mark.covers( + "llm.chat_completions.openai.thinking.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_reasoning_reports_reasoning_tokens( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-reasoning-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", + content="A train travels 60 miles in 1.5 hours. What is its average speed in mph?", + ) + ], + reasoning_effort="low", + max_tokens=2048, + ), + ) + ) + assert response.choices, f"reasoning call returned no choices: {response}" + message = response.choices[0].message + assert message and message.content and message.content.strip(), f"reasoning call had no answer: {response}" + details = response.usage.completion_tokens_details if response.usage else None + assert details and details.reasoning_tokens and details.reasoning_tokens > 0, ( + f"a reasoning model must report reasoning tokens, got usage={response.usage}" + ) + + @pytest.mark.covers( + "llm.chat_completions.openai.vision.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_vision_describes_image( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-vision-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_VISION_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_vision_messages(), max_tokens=32))) + _assert_describes_cat(response) + + @pytest.mark.covers( + "llm.chat_completions.openai.prompt_cache_5m.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_prompt_cache_hits_on_repeat( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-cache-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + body = ChatBody( + model=model, + messages=[ + ChatMessage(role="system", content=CACHE_PREFIX), + ChatMessage(role="user", content="Reply with the single word pong."), + ], + max_tokens=16, + ) + unwrap(client.proxy.chat(key, body)) + second = unwrap(client.proxy.chat(key, body)) + + details = second.usage.prompt_tokens_details if second.usage else None + assert details and details.cached_tokens and details.cached_tokens > 0, ( + f"a repeated large-prefix prompt must report cached prompt tokens, got usage={second.usage}" + ) + + @pytest.mark.covers( + "llm.chat_completions.openai.tool_use.stream.works", + exercised_on=["chat_completions"], + ) + def test_openai_chat_streams_tool_call( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + require_env("OPENAI_API_KEY") + model = f"e2e-openai-tool-stream-{unique_marker()}" + model_id = client.proxy.create_model( + model, LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY") + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + result = client.proxy.chat_stream( + key, + ChatBody( + model=model, + messages=[ + ChatMessage(role="user", content="What is the weather in San Francisco? Use the get_weather tool.") + ], + tools=[_WEATHER_TOOL], + tool_choice="required", + max_tokens=128, + stream=True, + ), + ) + assert result.ok and result.is_streaming, f"tool stream was not established: {result}" + assert result.stream_error is None, f"tool stream carried an error event: {result.stream_error}" + name, arguments = _streamed_tool_call(result.stream_events) + assert name == "get_weather", f"streamed tool call named {name!r}: {result.stream_events[:5]}" + args = _WeatherArgs.model_validate_json(arguments) + assert args.location.strip(), f"streamed tool call arguments missing location: {arguments!r}" + + +class TestBedrockConverseChatCompletions: + """Bedrock Converse via /chat/completions, the customer's AWS stack. A non-OpenAI + provider must return real content on both the non-streamed and streamed paths. + """ + + def _register(self, client: PassthroughClient, resources: ResourceManager, prefix: str) -> str: + require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + model = f"{prefix}-{unique_marker()}" + model_id = client.proxy.create_model(model, _bedrock_params()) + resources.defer(lambda: client.proxy.delete_model(model_id)) + return model + + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.basic.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_bedrock_converse_chat_returns_content( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = self._register(client, resources, "e2e-bedrock-chat") + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=f"Reply with the single word pong. {unique_marker()}")], + max_tokens=32, + ), + ) + ) + assert response.choices, f"bedrock converse chat returned no choices: {response}" + content = response.choices[0].message.content if response.choices[0].message else None + assert content and content.strip(), f"bedrock converse returned empty content: {response}" + + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.basic.stream.works", + exercised_on=["chat_completions"], + ) + def test_bedrock_converse_chat_streams_real_content( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = self._register(client, resources, "e2e-bedrock-stream") + key = resources.key() + + result = client.proxy.chat_stream( + key, + ChatBody( + model=model, + messages=[ + ChatMessage(role="user", content=f"Count from 1 to 5, one number per line. {unique_marker()}") + ], + max_tokens=64, + stream=True, + ), + ) + _assert_streamed_completion(result) + + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.tool_use.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_bedrock_converse_chat_returns_tool_call( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = self._register(client, resources, "e2e-bedrock-tool") + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage(role="user", content="What is the weather in San Francisco? Use the get_weather tool.") + ], + tools=[_WEATHER_TOOL], + tool_choice="required", + max_tokens=128, + ), + ) + ) + _assert_weather_tool_call(response) + + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.thinking.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_bedrock_converse_chat_returns_reasoning( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = self._register(client, resources, "e2e-bedrock-thinking") + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content="What is 17 times 23? Think it through step by step.")], + thinking=ThinkingParam(type="enabled", budget_tokens=1024), + max_tokens=2048, + ), + ) + ) + assert response.choices, f"bedrock thinking returned no choices: {response}" + message = response.choices[0].message + assert message and message.content and message.content.strip(), ( + f"bedrock thinking returned no answer content: {response}" + ) + assert message.reasoning_content and message.reasoning_content.strip(), ( + "thinking was enabled but no reasoning_content came back on the Bedrock Converse path" + ) + + @pytest.mark.covers( + "llm.chat_completions.bedrock_converse.vision.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_bedrock_converse_chat_vision_describes_image( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = self._register(client, resources, "e2e-bedrock-vision") + key = resources.key() + + response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_vision_messages(), max_tokens=32))) + _assert_describes_cat(response) diff --git a/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py b/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py index 56f2de8bd4f..157caedd561 100644 --- a/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py +++ b/tests/e2e/llm_translation/test_embeddings_endpoint_e2e.py @@ -1,9 +1,9 @@ -"""Live e2e: POST /embeddings returns a real vector. +"""Live e2e: POST /embeddings returns a real vector across OpenAI, Bedrock, Vertex. -Registers an OpenAI embedding deployment at runtime and asserts a non-empty, -non-zero vector came back. Migrated from -litellm-regression-tests/tests/test_inference_endpoints.py; the LIT-3167 guard in -tests/e2e/embeddings/ covers the Gemini embedding path. +Each test registers the deployment it needs at runtime (deleted on teardown) and +asserts a non-empty, non-zero vector came back. The LIT-3167 guard in +tests/e2e/embeddings/ covers the Gemini embedding path; embeddings cost tracking is +covered by tests/e2e/quota_management/spend_tracking/. """ from __future__ import annotations @@ -20,6 +20,7 @@ pytestmark = pytest.mark.e2e class TestEmbeddingsEndpoint: + @pytest.mark.covers("llm.embeddings.openai.basic.nonstream.works") def test_embeddings_returns_vector( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: @@ -40,3 +41,49 @@ class TestEmbeddingsEndpoint: assert any(component != 0.0 for component in parsed.first_vector), ( f"embedding vector is all zeros: {result.body[:300]}" ) + + @pytest.mark.covers("llm.embeddings.bedrock.basic.nonstream.works") + def test_bedrock_embeddings_returns_vector( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-embeddings-bedrock-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="bedrock/amazon.titan-embed-text-v2:0", aws_region_name="us-west-2" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.embeddings(key, model, "Say this is a test!") + require_successful_call(result) + parsed = EmbeddingsResult.model_validate_json(result.body) + assert parsed.first_vector, f"/embeddings returned no vector: {result.body[:300]}" + assert any(component != 0.0 for component in parsed.first_vector), ( + f"embedding vector is all zeros: {result.body[:300]}" + ) + + @pytest.mark.covers("llm.embeddings.vertex.basic.nonstream.works") + def test_vertex_embeddings_returns_vector( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-embeddings-vertex-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="vertex_ai/gemini-embedding-2", + vertex_project="os.environ/VERTEXAI_PROJECT", + vertex_location="us-central1", + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.embeddings(key, model, "Say this is a test!") + require_successful_call(result) + parsed = EmbeddingsResult.model_validate_json(result.body) + assert parsed.first_vector, f"/embeddings returned no vector: {result.body[:300]}" + assert any(component != 0.0 for component in parsed.first_vector), ( + f"embedding vector is all zeros: {result.body[:300]}" + ) diff --git a/tests/e2e/llm_translation/test_image_generation_e2e.py b/tests/e2e/llm_translation/test_image_generation_e2e.py index 4d2211f3be4..1ba78a7e083 100644 --- a/tests/e2e/llm_translation/test_image_generation_e2e.py +++ b/tests/e2e/llm_translation/test_image_generation_e2e.py @@ -9,7 +9,7 @@ from __future__ import annotations import pytest -from e2e_config import unique_marker +from e2e_config import require_env, unique_marker from e2e_http import require_successful_call from endpoints_client import EndpointsClient, ImagesResult from lifecycle import ResourceManager @@ -18,7 +18,17 @@ from models import LiteLLMParamsBody pytestmark = pytest.mark.e2e +def _assert_image_returned(body: str) -> None: + parsed = ImagesResult.model_validate_json(body) + assert parsed.data, f"/images/generations returned no data: {body[:300]}" + first = parsed.data[0] + assert first.b64_json or first.url, ( + f"generated image has neither b64_json nor url: {body[:300]}" + ) + + class TestImageGeneration: + @pytest.mark.covers("llm.images_generations.openai.basic.nonstream.works") def test_image_generation_returns_image( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: @@ -34,9 +44,26 @@ class TestImageGeneration: result = endpoints_client.images(key, model, "Draw a cute cat") require_successful_call(result) - parsed = ImagesResult.model_validate_json(result.body) - assert parsed.data, f"/images/generations returned no data: {result.body[:300]}" - first = parsed.data[0] - assert first.b64_json or first.url, ( - f"generated image has neither b64_json nor url: {result.body[:300]}" + _assert_image_returned(result.body) + + @pytest.mark.covers("llm.images_generations.bedrock.basic.nonstream.works", exercised_on=["images_generations"]) + def test_bedrock_image_generation_returns_image( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + model = f"e2e-bedrock-image-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="bedrock/amazon.titan-image-generator-v2:0", + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ), ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.images(key, model, "Draw a cute cat") + require_successful_call(result) + _assert_image_returned(result.body) diff --git a/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py b/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py new file mode 100644 index 00000000000..d8d44820e80 --- /dev/null +++ b/tests/e2e/llm_translation/test_messages_azure_foundry_e2e.py @@ -0,0 +1,162 @@ +"""Live e2e: POST /v1/messages routed to Azure AI Foundry Anthropic deployments. + +Registers `azure_ai/` deployments at runtime and drives the Messages +endpoint through the gateway across the behaviors an Anthropic client relies on: +a basic completion, a streamed completion, and tool use (non-streaming and +streaming). Auth is the Azure API key (`x-api-key`); the deployment reads +`AZURE_AI_API_BASE` / `AZURE_AI_API_KEY` from the proxy env, so no secret is +sent in the request. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import EXPECT_RUST, unique_marker +from e2e_http import StreamingResponse, require_successful_call, unwrap +from endpoints_client import EndpointsClient +from lifecycle import ResourceManager +from models import ( + AnthropicCustomTool, + AnthropicMessagesBody, + ChatMessage, + JsonSchemaProperty, + LiteLLMParamsBody, + ToolInputSchema, +) + +pytestmark = pytest.mark.e2e + +AZURE_FOUNDRY_MODEL = "azure_ai/claude-haiku-4-5" + +WEATHER_TOOL = AnthropicCustomTool( + name="get_weather", + description="Get the current weather for a city.", + input_schema=ToolInputSchema( + properties={"city": JsonSchemaProperty(type="string")}, + required=["city"], + ), +) + + +def _assert_streamed_ok(result: StreamingResponse) -> None: + require_successful_call(result) + assert result.is_streaming, f"response was not streamed: {result.headers}" + assert not result.stream_error, f"stream errored: {result.stream_error}" + assert result.stream_events, "stream produced no SSE events" + assert any("content_block_delta" in event for event in result.stream_events), ( + "stream carried no content deltas" + ) + assert any("message_stop" in event for event in result.stream_events), ( + "stream never reached message_stop" + ) + if EXPECT_RUST: + assert result.headers.get("x-litellm-rust") == "true", ( + "E2E_EXPECT_RUST is set, so this gateway must serve /v1/messages through the " + "Rust path, but the response carried no x-litellm-rust marker. The request " + "still succeeded, which is exactly the failure mode: a gateway whose native " + f"extension is unavailable falls back to Python silently. headers={result.headers}" + ) + + +class TestAzureFoundryMessages: + def _register( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> tuple[str, str]: + model = f"e2e-azure-foundry-messages-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model=AZURE_FOUNDRY_MODEL, + api_base="os.environ/AZURE_AI_API_BASE", + api_key="os.environ/AZURE_AI_API_KEY", + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + return model, resources.key(models=[model]) + + @pytest.mark.covers("llm.messages.azure_foundry.basic.nonstream.works") + def test_basic_nonstream( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) + response = unwrap( + endpoints_client.proxy.messages( + key, + AnthropicMessagesBody( + model=model, + max_tokens=64, + messages=[ChatMessage(role="user", content="Reply with one word.")], + ), + ) + ) + assert response.content, f"no content blocks in response: {response}" + text = "".join(block.text or "" for block in response.content if block.type == "text") + assert text.strip(), f"/v1/messages returned no text: {response}" + + @pytest.mark.covers("llm.messages.azure_foundry.basic.stream.works") + def test_basic_stream( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) + result = endpoints_client.proxy.messages_stream( + key, + AnthropicMessagesBody( + model=model, + max_tokens=64, + stream=True, + messages=[ChatMessage(role="user", content="Count from one to three.")], + ), + ) + _assert_streamed_ok(result) + + @pytest.mark.covers("llm.messages.azure_foundry.tool_use.nonstream.works") + def test_tool_use_nonstream( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) + response = unwrap( + endpoints_client.proxy.messages( + key, + AnthropicMessagesBody( + model=model, + max_tokens=256, + tools=[WEATHER_TOOL], + messages=[ + ChatMessage(role="user", content="What is the weather in Paris? Use the tool.") + ], + ), + ) + ) + assert response.content, f"no content blocks in response: {response}" + assert any(block.type == "tool_use" for block in response.content), ( + f"model did not call the tool: {response}" + ) + + @pytest.mark.covers("llm.messages.azure_foundry.tool_use.stream.works") + def test_tool_use_stream( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) + result = endpoints_client.proxy.messages_stream( + key, + AnthropicMessagesBody( + model=model, + max_tokens=256, + stream=True, + tools=[WEATHER_TOOL], + messages=[ + ChatMessage(role="user", content="What is the weather in Paris? Use the tool.") + ], + ), + ) + require_successful_call(result) + assert result.is_streaming, f"response was not streamed: {result.headers}" + assert not result.stream_error, f"stream errored: {result.stream_error}" + assert result.stream_events, "stream produced no SSE events" + assert any("tool_use" in event for event in result.stream_events), ( + "stream carried no tool_use block" + ) + assert any("message_stop" in event for event in result.stream_events), ( + "stream never reached message_stop" + ) diff --git a/tests/e2e/llm_translation/test_messages_e2e.py b/tests/e2e/llm_translation/test_messages_e2e.py index b0a48f22118..44376218c6b 100644 --- a/tests/e2e/llm_translation/test_messages_e2e.py +++ b/tests/e2e/llm_translation/test_messages_e2e.py @@ -1,7 +1,8 @@ """Live e2e: POST /v1/messages (Anthropic Messages API) returns a real completion. Registers an Anthropic deployment at runtime, drives the Messages endpoint through -the gateway, and asserts an assistant message with text came back. Migrated from +the gateway, and asserts an assistant message with text came back, both +non-streaming and streamed. Migrated from litellm-regression-tests/tests/test_inference_endpoints.py. """ @@ -9,31 +10,163 @@ from __future__ import annotations import pytest -from e2e_config import unique_marker -from e2e_http import require_successful_call +from e2e_config import require_env, unique_marker +from e2e_http import require_successful_call, unwrap from endpoints_client import EndpointsClient, MessagesResult from lifecycle import ResourceManager -from models import LiteLLMParamsBody +from models import ( + AnthropicCustomTool, + AnthropicMessagesBody, + ChatMessage, + JsonSchemaProperty, + LiteLLMParamsBody, + SpendLogRow, + ToolInputSchema, +) pytestmark = pytest.mark.e2e +ANTHROPIC_BACKEND = "anthropic/claude-haiku-4-5" + +WEATHER_TOOL = AnthropicCustomTool( + name="get_weather", + description="Get the current weather for a city.", + input_schema=ToolInputSchema( + properties={"city": JsonSchemaProperty(type="string")}, + required=["city"], + ), +) + + +def _approx_equal(actual: float, expected: float) -> bool: + """Within 1% or 1e-9 absolute - spend math, not exact float identity.""" + return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2) + class TestAnthropicMessages: - def test_messages_returns_completion( + def _register( self, endpoints_client: EndpointsClient, resources: ResourceManager - ) -> None: + ) -> tuple[str, str]: model = f"e2e-messages-{unique_marker()}" model_id = endpoints_client.create_model( model, LiteLLMParamsBody( - model="anthropic/claude-haiku-4-5", api_key="os.environ/ANTHROPIC_API_KEY" + model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY" ), ) resources.defer(lambda: endpoints_client.delete_model(model_id)) - key = resources.key() + return model, resources.key() + + @pytest.mark.covers("llm.messages.anthropic.basic.nonstream.works") + def test_messages_returns_completion( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) result = endpoints_client.messages(key, model, "reply with one word") require_successful_call(result) parsed = MessagesResult.model_validate_json(result.body) assert parsed.role == "assistant", f"unexpected role: {result.body[:300]}" assert parsed.text.strip(), f"/v1/messages returned no text: {result.body[:300]}" + + @pytest.mark.covers("llm.messages.anthropic.basic.nonstream.cost_logged") + def test_messages_logs_cost_matching_the_response_header( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + require_env("ANTHROPIC_API_KEY") + model = f"e2e-messages-cost-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.messages(key, model, f"reply with one word {unique_marker()}") + require_successful_call(result) + parsed = MessagesResult.model_validate_json(result.body) + assert parsed.role == "assistant" and parsed.text.strip(), ( + f"/v1/messages returned no assistant text: {result.body[:300]}" + ) + + # The customer reads per-request cost off the response header (LIT-4076), so + # it must be present and positive on /v1/messages, not only /chat/completions. + header_cost = result.response_cost + assert header_cost is not None and header_cost > 0, ( + "x-litellm-response-cost header missing or non-positive on /v1/messages; " + f"headers={result.headers}" + ) + + # Correlate the spend row by the unique scoped key, not the Anthropic response + # id: on /v1/messages the spend-log request_id is the proxy's own call id, which + # need not equal the message body id, so an id-based poll can miss a correctly + # logged row and time out. The key is fresh per test, so its only priced row is + # this call. + def _priced(rows: list[SpendLogRow]) -> bool: + return any(r.spend is not None and r.spend > 0 for r in rows) + + rows = endpoints_client.proxy.poll_logs_for_key(key, predicate=_priced) + priced = [r for r in rows if r.spend is not None and r.spend > 0] + assert priced, ( + f"no priced /spend/logs row landed for key {key} within the poll window; got {rows}" + ) + row = priced[0] + assert (row.prompt_tokens or 0) > 0 and (row.completion_tokens or 0) > 0, ( + f"messages spend row missing token counts, so the cost is not real usage: {row}" + ) + assert row.spend is not None and _approx_equal(row.spend, header_cost), ( + f"logged spend {row.spend} disagrees with the x-litellm-response-cost header {header_cost}; " + "the customer bills against the header, so the two must match" + ) + + @pytest.mark.covers("llm.messages.anthropic.basic.stream.works") + def test_messages_streams_completion( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) + + result = endpoints_client.proxy.messages_stream( + key, + AnthropicMessagesBody( + model=model, + max_tokens=64, + stream=True, + messages=[ChatMessage(role="user", content="Count from one to three.")], + ), + ) + require_successful_call(result) + assert result.is_streaming, f"response was not streamed: {result.headers}" + assert not result.stream_error, f"stream errored: {result.stream_error}" + assert result.stream_events, "stream produced no SSE events" + assert any("content_block_delta" in event for event in result.stream_events), ( + "stream carried no content deltas" + ) + assert any("message_stop" in event for event in result.stream_events), ( + "stream never reached message_stop" + ) + + @pytest.mark.covers("llm.messages.anthropic.tool_use.nonstream.works") + def test_messages_tool_use( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model, key = self._register(endpoints_client, resources) + + response = unwrap( + endpoints_client.proxy.messages( + key, + AnthropicMessagesBody( + model=model, + max_tokens=256, + tools=[WEATHER_TOOL], + messages=[ + ChatMessage(role="user", content="What is the weather in Paris? Use the tool.") + ], + ), + ) + ) + assert response.content, f"no content blocks in response: {response}" + assert any(block.type == "tool_use" for block in response.content), ( + f"model did not call the tool: {response}" + ) diff --git a/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py b/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py index 7a04b044634..97d24e0564b 100644 --- a/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py +++ b/tests/e2e/llm_translation/test_messages_mid_conversation_system_native_providers_e2e.py @@ -7,7 +7,7 @@ accepted in place on Claude 4.8+/5 (200) but rejected on Claude 4.7 and older ("role 'system' is not supported on this model", 400), and a *leading* system entry is rejected on every model ("messages.0: use the top-level 'system' parameter"). This mirrors Bedrock Invoke (PRs #32578/#32831/#32882); the same -model-gated hoist now runs for these two providers (Kraken Tech RCA gap #3). +model-gated hoist now runs for these two providers (customer RCA gap #3). Flagged models (``supports_mid_conversation_system`` in the cost map: Claude 4.8+ and the 5 family) must keep the reminder in ``messages`` so the top-level @@ -88,9 +88,9 @@ def _system_reminder_turn() -> RichMessage: def _post_messages(client: EndpointsClient, key: str, body: RichMessagesRequest) -> Result[MessagesResult]: - return client.gateway.transport.post( + return client.proxy.transport.post( "/v1/messages", - headers=client.gateway.transport.bearer(key), + headers=client.proxy.transport.bearer(key), json=body, response_type=MessagesResult, ) diff --git a/tests/e2e/llm_translation/test_moderations_e2e.py b/tests/e2e/llm_translation/test_moderations_e2e.py new file mode 100644 index 00000000000..69cf4414a48 --- /dev/null +++ b/tests/e2e/llm_translation/test_moderations_e2e.py @@ -0,0 +1,65 @@ +"""Live e2e: POST /v1/moderations classifies content against the provider policy. + +Registers OpenAI's omni moderation model at runtime and asserts the product +promise on both sides of the decision: clearly violent text comes back flagged +with at least one policy category tripped, and benign text comes back not flagged. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import unwrap +from endpoints_client import EndpointsClient +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + +VIOLENT_TEXT = "I am going to find you and kill you, and I will hurt everyone you love." +BENIGN_TEXT = "I enjoyed the sunny afternoon and a relaxing walk in the park today." + + +def _register_moderation_model( + endpoints_client: EndpointsClient, resources: ResourceManager +) -> str: + model = f"e2e-moderation-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="openai/omni-moderation-latest", api_key="os.environ/OPENAI_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + return model + + +class TestModerations: + @pytest.mark.covers("llm.moderations.openai.basic.nonstream.works") + def test_moderations_flags_violent_content( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = _register_moderation_model(endpoints_client, resources) + key = resources.key() + + result = unwrap(endpoints_client.moderations(key, model, VIOLENT_TEXT)) + item = result.first + assert item is not None, f"/moderations returned no results: {result}" + assert item.flagged, f"violent text was not flagged: {item}" + assert item.flagged_categories, ( + f"flagged result reported no true category: {item}" + ) + + def test_moderations_passes_benign_content( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = _register_moderation_model(endpoints_client, resources) + key = resources.key() + + result = unwrap(endpoints_client.moderations(key, model, BENIGN_TEXT)) + item = result.first + assert item is not None, f"/moderations returned no results: {result}" + assert not item.flagged, ( + f"benign text was flagged as {item.flagged_categories}: {item}" + ) diff --git a/tests/e2e/llm_translation/test_passthrough_e2e.py b/tests/e2e/llm_translation/test_passthrough_e2e.py index c8806faf3ea..ed5c657d23e 100644 --- a/tests/e2e/llm_translation/test_passthrough_e2e.py +++ b/tests/e2e/llm_translation/test_passthrough_e2e.py @@ -15,7 +15,8 @@ import pytest from e2e_config import unique_marker from e2e_http import StreamingResponse, require_successful_call -from models import SpendLogRow +from lifecycle import ResourceManager +from models import KeyGenerateBody, SpendLogRow from passthrough_client import ( AnthropicTool, GeminiFunctionDeclaration, @@ -157,3 +158,25 @@ def test_anthropic_passthrough_tool_call_logs_cost( row = _fetch_cost_breakdown(client, result) assert row.custom_llm_provider == "anthropic" + + +class TestPassthroughModelAllowlist: + """A passthrough route must honor the calling key's model allow-list. + + The customer fronts native provider calls through the proxy with custom auth, + so a key scoped to one model must not reach a different model just because the + request goes through the passthrough route rather than /chat/completions. + """ + + @pytest.mark.covers("other.auth.passthrough.model_allowlist_enforced") + def test_passthrough_denies_model_outside_key_allowlist( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + key = client.proxy.generate_key(KeyGenerateBody(models=["gemini-2.5-flash"])) + resources.defer(lambda: client.proxy.delete_key(key)) + + result = client.anthropic_message(key, "claude-haiku-4-5", f"say hi {unique_marker()}") + assert result.status_code == 403, ( + "a key restricted to gemini-2.5-flash must be denied a claude passthrough call, " + f"got {result.status_code}: {result.body[:300]}" + ) diff --git a/tests/e2e/llm_translation/test_passthrough_headers_e2e.py b/tests/e2e/llm_translation/test_passthrough_headers_e2e.py index d9fe37c79ed..045988334d5 100644 --- a/tests/e2e/llm_translation/test_passthrough_headers_e2e.py +++ b/tests/e2e/llm_translation/test_passthrough_headers_e2e.py @@ -1,29 +1,32 @@ """Live e2e: custom pass-through endpoints inject configured headers and honor x-pass-* client headers (prefix stripped) on the way to the upstream. -The upstream is a real public echo service (httpbin.org/anything). Creating the -route via POST /config/pass_through_endpoint, calling it with a virtual key, and -asserting the echo body is the product path operators use; a mock would not -prove the proxy actually rewrote the outbound request. +The upstream is the real Anthropic Messages API rather than an echo service: +Anthropic doesn't echo request headers back, but it does gate real behavior on +two of them, which is enough to prove forwarding without a mock. A static +x-api-key configured on the pass-through endpoint (the caller never supplies +one) must reach upstream, or every call 401s; an invalid x-pass-anthropic-version +sent by the caller must reach upstream with the prefix stripped, and Anthropic +echoes the exact value back in its 400 body, so a unique-per-run marker proves +this specific request's header - not a stale or cached one - got there. """ from __future__ import annotations import pytest -from pydantic import BaseModel, Field, ValidationError +from pydantic import BaseModel, Field from e2e_config import unique_marker -from e2e_http import AuthHeaders, NoBody, StreamingResponse, require_successful_call, unwrap +from e2e_http import AuthHeaders, NoBody, require_successful_call, unwrap +from endpoints_client import MessagesResult from lifecycle import ResourceManager -from models import KeyGenerateBody +from models import ChatMessage, KeyGenerateBody from passthrough_client import PassthroughClient pytestmark = pytest.mark.e2e -ECHO_TARGET = "https://httpbin.org/anything" -STATIC_HEADER_NAME = "x-e2e-static-header" -PASS_HEADER_STEM = "e2e-client-marker" -PASS_HEADER_NAME = f"x-pass-{PASS_HEADER_STEM}" +ANTHROPIC_MESSAGES_TARGET = "https://api.anthropic.com/v1/messages" +MODEL = "claude-haiku-4-5-20251001" class PassThroughCreateBody(BaseModel): @@ -48,30 +51,26 @@ class PassThroughDeleteParams(BaseModel): endpoint_id: str -class EchoCallHeaders(AuthHeaders): +class AnthropicPassThroughHeaders(AuthHeaders): content_type: str = Field(default="application/json", serialization_alias="Content-Type") - x_pass_e2e_client_marker: str = Field(serialization_alias="x-pass-e2e-client-marker") + x_pass_anthropic_version: str = Field(serialization_alias="x-pass-anthropic-version") -class EchoBody(BaseModel): - ping: str +class AnthropicMessagesBody(BaseModel): + model: str + max_tokens: int = 8 + messages: list[ChatMessage] -class EchoResponse(BaseModel): - headers: dict[str, str] - - -def _create_passthrough( - client: PassthroughClient, *, path: str, static_value: str -) -> PassThroughEndpoint: +def _create_passthrough(client: PassthroughClient, *, path: str) -> PassThroughEndpoint: created = unwrap( client.proxy.transport.post( "/config/pass_through_endpoint", headers=client.proxy.transport.master, json=PassThroughCreateBody( path=path, - target=ECHO_TARGET, - headers={STATIC_HEADER_NAME: static_value}, + target=ANTHROPIC_MESSAGES_TARGET, + headers={"x-api-key": "os.environ/ANTHROPIC_API_KEY"}, ), response_type=PassThroughCreateResponse, ) @@ -92,12 +91,8 @@ def _delete_passthrough(client: PassthroughClient, endpoint_id: str) -> None: ) -def _echo_headers(resp: StreamingResponse) -> dict[str, str]: - try: - echo = EchoResponse.model_validate_json(resp.body) - except ValidationError as exc: - pytest.fail(f"echo upstream did not return a headers map: {exc}; body={resp.body[:300]}") - return {k.lower(): v for k, v in echo.headers.items()} +def _messages_body() -> AnthropicMessagesBody: + return AnthropicMessagesBody(model=MODEL, messages=[ChatMessage(role="user", content="Say hi.")]) class TestPassthroughHeaders: @@ -110,10 +105,8 @@ class TestPassthroughHeaders: ) -> None: marker = unique_marker() path = f"/e2e-passthrough-headers-{marker}" - static_value = f"static-{marker}" - client_value = f"client-{marker}" - endpoint = _create_passthrough(client, path=path, static_value=static_value) + endpoint = _create_passthrough(client, path=path) assert endpoint.id is not None resources.defer(lambda: _delete_passthrough(client, endpoint.id or "")) @@ -128,23 +121,32 @@ class TestPassthroughHeaders: result = client.proxy.transport.send( path, - headers=EchoCallHeaders( + headers=AnthropicPassThroughHeaders( authorization=f"Bearer {key}", - x_pass_e2e_client_marker=client_value, + x_pass_anthropic_version="2023-06-01", ), - json=EchoBody(ping=marker), + json=_messages_body(), ) require_successful_call(result) + completion = MessagesResult.model_validate_json(result.body) + assert completion.text.strip(), ( + f"static x-api-key must reach Anthropic for the call to succeed at all; got {result.body[:300]}" + ) - upstream = _echo_headers(result) - assert upstream.get(STATIC_HEADER_NAME) == static_value, ( - f"configured pass-through header {STATIC_HEADER_NAME!r} not on upstream " - f"request; got {upstream}" + invalid_version = f"e2e-passhdr-{unique_marker()}" + blocked = client.proxy.transport.send( + path, + headers=AnthropicPassThroughHeaders( + authorization=f"Bearer {key}", + x_pass_anthropic_version=invalid_version, + ), + json=_messages_body(), ) - assert upstream.get(PASS_HEADER_STEM) == client_value, ( - f"x-pass-* header should strip the prefix and forward as {PASS_HEADER_STEM!r}; " - f"got {upstream}" + assert blocked.status_code == 400, ( + f"expected Anthropic to reject the invalid anthropic-version, got " + f"{blocked.status_code}: {blocked.body[:300]}" ) - assert PASS_HEADER_NAME not in upstream, ( - "upstream must not see the x-pass- prefix; proxy should strip it" + assert invalid_version in blocked.body, ( + f"x-pass-anthropic-version must reach upstream with the prefix stripped; " + f"marker missing from Anthropic's error body: {blocked.body[:300]}" ) diff --git a/tests/e2e/llm_translation/test_rerank_e2e.py b/tests/e2e/llm_translation/test_rerank_e2e.py index 4b30ac1ea5c..0857ff65a52 100644 --- a/tests/e2e/llm_translation/test_rerank_e2e.py +++ b/tests/e2e/llm_translation/test_rerank_e2e.py @@ -9,7 +9,7 @@ from __future__ import annotations import pytest -from e2e_config import unique_marker +from e2e_config import require_env, unique_marker from e2e_http import require_successful_call from endpoints_client import EndpointsClient, RerankResult from lifecycle import ResourceManager @@ -23,9 +23,20 @@ DOCUMENTS = [ "Washington, D.C. is the capital of the United States.", "Capital punishment has existed in the United States since before it was a country.", ] +QUERY = "What is the capital of the United States?" + + +def _assert_top_n_scored(body: str) -> None: + parsed = RerankResult.model_validate_json(body) + assert parsed.results, f"/rerank returned no results: {body[:300]}" + assert len(parsed.results) <= 3, f"top_n=3 not honored: {body[:300]}" + assert parsed.results[0].relevance_score is not None, ( + f"top rerank result has no relevance_score: {body[:300]}" + ) class TestRerank: + @pytest.mark.covers("llm.rerank.cohere.basic.nonstream.works") def test_rerank_scores_top_n( self, endpoints_client: EndpointsClient, resources: ResourceManager ) -> None: @@ -37,13 +48,28 @@ class TestRerank: resources.defer(lambda: endpoints_client.delete_model(model_id)) key = resources.key() - result = endpoints_client.rerank( - key, model, "What is the capital of the United States?", DOCUMENTS, top_n=3 - ) + result = endpoints_client.rerank(key, model, QUERY, DOCUMENTS, top_n=3) require_successful_call(result) - parsed = RerankResult.model_validate_json(result.body) - assert parsed.results, f"/rerank returned no results: {result.body[:300]}" - assert len(parsed.results) <= 3, f"top_n=3 not honored: {result.body[:300]}" - assert parsed.results[0].relevance_score is not None, ( - f"top rerank result has no relevance_score: {result.body[:300]}" + _assert_top_n_scored(result.body) + + @pytest.mark.covers("llm.rerank.bedrock.basic.nonstream.works", exercised_on=["rerank"]) + def test_bedrock_rerank_scores_top_n( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + model = f"e2e-bedrock-rerank-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="bedrock/amazon.rerank-v1:0", + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ), ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.rerank(key, model, QUERY, DOCUMENTS, top_n=3) + require_successful_call(result) + _assert_top_n_scored(result.body) diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py index bd98f11c045..d24d2b53b71 100644 --- a/tests/e2e/llm_translation/test_responses_e2e.py +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -13,7 +13,7 @@ from typing import cast import pytest from pydantic import BaseModel, ValidationError -from e2e_config import unique_marker +from e2e_config import require_env, unique_marker from e2e_http import require_successful_call from endpoints_client import ( EndpointsClient, @@ -29,6 +29,26 @@ from models import LiteLLMParamsBody pytestmark = pytest.mark.e2e +BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" + +WEATHER_TOOL = ResponsesFunctionTool( + name="get_weather", + description="Get the weather for a location", + parameters=FunctionParameters( + properties={"location": FunctionParameterProperty(type="string")}, + required=["location"], + ), +) + + +def _bedrock_params() -> LiteLLMParamsBody: + return LiteLLMParamsBody( + model=BEDROCK_CONVERSE_BACKEND, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ) + class WeatherArguments(BaseModel): location: str @@ -190,6 +210,84 @@ class TestResponses: parsed = ResponsesResult.model_validate_json(result.body) assert parsed.text.strip(), f"/responses returned no output text: {result.body[:300]}" + @pytest.mark.covers("llm.responses.anthropic.tool_use.nonstream.works") + def test_responses_anthropic_returns_function_call( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + model = f"e2e-responses-{unique_marker()}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="anthropic/claude-haiku-4-5", api_key="os.environ/ANTHROPIC_API_KEY" + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.responses_with_tools( + key, + model, + "What is the weather in San Francisco? Use the get_weather tool.", + [ + ResponsesFunctionTool( + name="get_weather", + description="Get the weather for a location", + parameters=FunctionParameters( + properties={"location": FunctionParameterProperty(type="string")}, + required=["location"], + ), + ) + ], + ) + require_successful_call(result) + parsed = ResponsesResult.model_validate_json(result.body) + function_call = next( + (call for call in parsed.function_calls if call.name == "get_weather"), + None, + ) + assert function_call is not None, f"no get_weather function call: {result.body[:500]}" + assert function_call.arguments is not None + raw_arguments = cast(object, json.loads(function_call.arguments)) + arguments = WeatherArguments.model_validate(raw_arguments) + assert arguments.location, f"function call arguments missing location: {function_call.arguments}" + + @pytest.mark.covers("llm.responses.bedrock_converse.basic.nonstream.works") + def test_responses_bedrock_returns_completion( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + model = f"e2e-responses-{unique_marker()}" + model_id = endpoints_client.create_model(model, _bedrock_params()) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.responses(key, model, "reply with one word") + require_successful_call(result) + parsed = ResponsesResult.model_validate_json(result.body) + assert parsed.text.strip(), f"/responses over bedrock returned no output text: {result.body[:300]}" + + @pytest.mark.covers("llm.responses.bedrock_converse.tool_use.nonstream.works") + def test_responses_bedrock_returns_function_call( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + model = f"e2e-responses-{unique_marker()}" + model_id = endpoints_client.create_model(model, _bedrock_params()) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + result = endpoints_client.responses_with_tools( + key, model, "What is the weather in San Francisco? Use the get_weather tool.", [WEATHER_TOOL] + ) + require_successful_call(result) + parsed = ResponsesResult.model_validate_json(result.body) + function_call = next((call for call in parsed.function_calls if call.name == "get_weather"), None) + assert function_call is not None, f"no get_weather function call over bedrock: {result.body[:500]}" + assert function_call.arguments is not None + raw_arguments = cast(object, json.loads(function_call.arguments)) + arguments = WeatherArguments.model_validate(raw_arguments) + assert arguments.location, f"function call arguments missing location: {function_call.arguments}" + def _parse_stream_event( event: str, diff --git a/tests/e2e/load/conftest.py b/tests/e2e/load/conftest.py index e9fba02680d..89a571af83e 100644 --- a/tests/e2e/load/conftest.py +++ b/tests/e2e/load/conftest.py @@ -1,10 +1,12 @@ from __future__ import annotations +import os from collections.abc import Iterator import pytest from requests import RequestException +from e2e_config import WEEKLY_ANOMALY_OPT_IN_ENV from e2e_http import NoBody, Success from load_client import LoadClient, build_client from load_constants import LOAD_MODEL @@ -18,6 +20,22 @@ LOAD_MODEL_PARAMS = LiteLLMParamsBody( ) +def pytest_collection_modifyitems( + config: pytest.Config, items: list[pytest.Item] +) -> None: + if os.environ.get(WEEKLY_ANOMALY_OPT_IN_ENV): + return + deselected = [ + item for item in items if item.get_closest_marker("weekly") is not None + ] + if not deselected: + return + config.hook.pytest_deselected(items=deselected) + items[:] = [ + item for item in items if item.get_closest_marker("weekly") is None + ] + + @pytest.fixture(scope="session") def client(proxy: ProxyClient) -> LoadClient: return build_client(proxy) @@ -33,10 +51,8 @@ def _model_is_servable(proxy: ProxyClient, model_name: str) -> bool: return isinstance(result, Success) and any(entry.id == model_name for entry in result.data.data) -@pytest.fixture(scope="session", autouse=True) -def _ensure_load_model( # pyright: ignore[reportUnusedFunction] # pytest autouse session fixture, wired by name - client: LoadClient, -) -> Iterator[None]: +@pytest.fixture(scope="session") +def ensure_load_model(client: LoadClient) -> Iterator[None]: proxy = client.proxy if _model_is_servable(proxy, LOAD_MODEL): yield @@ -60,7 +76,9 @@ def _ensure_load_model( # pyright: ignore[reportUnusedFunction] # pytest autou @pytest.fixture -def load_key(resources: ResourceManager, client: LoadClient) -> str: +def load_key( + resources: ResourceManager, client: LoadClient, ensure_load_model: None +) -> str: key = client.proxy.generate_key(KeyGenerateBody(models=[LOAD_MODEL], user_id="e2e-load")) resources.defer(lambda: client.proxy.delete_key(key)) return key diff --git a/tests/e2e/load/session_anomaly.py b/tests/e2e/load/session_anomaly.py new file mode 100644 index 00000000000..c29b833635d --- /dev/null +++ b/tests/e2e/load/session_anomaly.py @@ -0,0 +1,299 @@ +from __future__ import annotations + +import time +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from dataclasses import dataclass + +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import Result, Success +from models import CacheControl, RichMessage, TextBlock +from transport import Transport + + +class SessionMessagesRequest(BaseModel): + model: str + max_tokens: int = 128 + system: list[TextBlock] + messages: list[RichMessage] + + +class SessionUsage(BaseModel): + input_tokens: int = 0 + output_tokens: int = 0 + cache_creation_input_tokens: int = 0 + cache_read_input_tokens: int = 0 + + +class SessionContentBlock(BaseModel): + type: str | None = None + text: str | None = None + + +class SessionMessagesResponse(BaseModel): + content: list[SessionContentBlock] = [] + usage: SessionUsage = SessionUsage() + + @property + def text(self) -> str: + return "".join(block.text or "" for block in self.content) + + +@dataclass(frozen=True, slots=True) +class TurnMetric: + turn_index: int + ok: bool + latency_seconds: float + uncached_input_tokens: int + cache_read_tokens: int + cache_creation_tokens: int + failure: str | None + + +@dataclass(frozen=True, slots=True) +class AnomalyReport: + planned_turns: int + attempted_turns: int + failed_turns: int + warm_turns: int + warm_uncached_input_tokens: int + warm_cache_read_tokens: int + warm_cache_creation_tokens: int + p95_turn_seconds: float + + @property + def error_ratio(self) -> float: + return self.failed_turns / self.planned_turns if self.planned_turns else 1.0 + + @property + def warm_cache_read_share(self) -> float: + billed = ( + self.warm_uncached_input_tokens + + self.warm_cache_read_tokens + + self.warm_cache_creation_tokens + ) + return self.warm_cache_read_tokens / billed if billed else 0.0 + + +def _system_prefix_block(marker: str) -> TextBlock: + text = " ".join( + f"Project context paragraph {index} for session {marker}." for index in range(300) + ) + return TextBlock(text=text, cache_control=CacheControl()) + + +def _user_turn_text(marker: str, turn_index: int) -> str: + notes = " ".join( + f"Working note {index} of turn {turn_index} in session {marker}." + for index in range(80) + ) + return f"Reply with one short sentence.\n{notes}" + + +def _reminder_turn() -> RichMessage: + return RichMessage( + role="system", + content=[ + TextBlock( + text="Keep the answer to one short sentence." + ) + ], + ) + + +def _without_cache_control(message: RichMessage) -> RichMessage: + return RichMessage( + role=message.role, + content=[TextBlock(text=block.text) for block in message.content], + ) + + +RETRY_BACKOFF_SECONDS = 2.0 + + +def retried( + call: Callable[[], Result[SessionMessagesResponse]], + attempts: int, + backoff_seconds: float = RETRY_BACKOFF_SECONDS, + sleep: Callable[[float], None] = time.sleep, +) -> Result[SessionMessagesResponse]: + result = call() + if isinstance(result, Success) or attempts <= 1: + return result + sleep(backoff_seconds) + return retried(call, attempts - 1, backoff_seconds, sleep) + + +def _metric( + result: Result[SessionMessagesResponse], turn_index: int, latency_seconds: float +) -> TurnMetric: + if isinstance(result, Success): + usage = result.data.usage + return TurnMetric( + turn_index=turn_index, + ok=True, + latency_seconds=latency_seconds, + uncached_input_tokens=usage.input_tokens, + cache_read_tokens=usage.cache_read_input_tokens, + cache_creation_tokens=usage.cache_creation_input_tokens, + failure=None, + ) + return TurnMetric( + turn_index=turn_index, + ok=False, + latency_seconds=latency_seconds, + uncached_input_tokens=0, + cache_read_tokens=0, + cache_creation_tokens=0, + failure=repr(result), + ) + + +def _drive_turns( + transport: Transport, + key: str, + model: str, + marker: str, + system_block: TextBlock, + history: tuple[RichMessage, ...], + turn_index: int, + remaining_turns: int, + attempts_per_turn: int, +) -> tuple[TurnMetric, ...]: + if remaining_turns == 0: + return () + user_turn = RichMessage( + role="user", + content=[ + TextBlock( + text=_user_turn_text(marker, turn_index), cache_control=CacheControl() + ) + ], + ) + started = time.monotonic() + result = retried( + lambda: transport.post( + "/v1/messages", + headers=transport.bearer(key), + json=SessionMessagesRequest( + model=model, + system=[system_block], + messages=[*history, user_turn], + ), + response_type=SessionMessagesResponse, + ), + attempts_per_turn, + ) + turn = _metric(result, turn_index, time.monotonic() - started) + if not isinstance(result, Success): + return (turn,) + assistant_turn = RichMessage( + role="assistant", content=[TextBlock(text=result.data.text or "Understood.")] + ) + return ( + turn, + *_drive_turns( + transport, + key, + model, + marker, + system_block, + ( + *history, + _without_cache_control(user_turn), + _reminder_turn(), + assistant_turn, + ), + turn_index + 1, + remaining_turns - 1, + attempts_per_turn, + ), + ) + + +def run_session( + transport: Transport, key: str, model: str, turns: int, attempts_per_turn: int +) -> tuple[TurnMetric, ...]: + marker = unique_marker() + return _drive_turns( + transport, + key, + model, + marker, + _system_prefix_block(marker), + (), + 1, + turns, + attempts_per_turn, + ) + + +def run_concurrent_sessions( + transport: Transport, + key: str, + model: str, + sessions: int, + turns_per_session: int, + attempts_per_turn: int, +) -> tuple[TurnMetric, ...]: + with ThreadPoolExecutor(max_workers=sessions) as pool: + futures = [ + pool.submit( + run_session, transport, key, model, turns_per_session, attempts_per_turn + ) + for _ in range(sessions) + ] + return tuple(turn for future in futures for turn in future.result()) + + +def settled_spend( + read_spend: Callable[[], float], + poll_interval: float, + settle_seconds: float, + timeout_seconds: float, + now: Callable[[], float] = time.monotonic, + sleep: Callable[[float], None] = time.sleep, +) -> float: + deadline = now() + timeout_seconds + settle_seconds + + def settle(previous: float, stable_since: float) -> float: + current = read_spend() + observed = now() + since = stable_since if current == previous else observed + if current > 0 and observed - since >= settle_seconds: + return current + if observed >= deadline: + raise AssertionError( + f"key spend never held a stable non-zero value for {settle_seconds}s " + f"within {timeout_seconds + settle_seconds}s (last read {current}); " + f"spend stopped being recorded, which is itself a spend anomaly" + ) + sleep(poll_interval) + return settle(current, since) + + return settle(-1.0, now()) + + +def _p95(latencies: tuple[float, ...]) -> float: + if not latencies: + return 0.0 + ranked = sorted(latencies) + return ranked[max(0, -(-len(ranked) * 95 // 100) - 1)] + + +def summarize(turns: tuple[TurnMetric, ...], planned_turns: int) -> AnomalyReport: + warm = tuple(turn for turn in turns if turn.ok and turn.turn_index >= 2) + return AnomalyReport( + planned_turns=planned_turns, + attempted_turns=len(turns), + failed_turns=planned_turns - sum(1 for turn in turns if turn.ok), + warm_turns=len(warm), + warm_uncached_input_tokens=sum(turn.uncached_input_tokens for turn in warm), + warm_cache_read_tokens=sum(turn.cache_read_tokens for turn in warm), + warm_cache_creation_tokens=sum(turn.cache_creation_tokens for turn in warm), + p95_turn_seconds=_p95( + tuple(turn.latency_seconds for turn in turns if turn.ok) + ), + ) diff --git a/tests/e2e/load/test_session_anomaly.py b/tests/e2e/load/test_session_anomaly.py new file mode 100644 index 00000000000..7062587352b --- /dev/null +++ b/tests/e2e/load/test_session_anomaly.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +from itertools import count, repeat + +import pytest + +from e2e_http import NetworkError, Success +from session_anomaly import ( + SessionMessagesResponse, + TurnMetric, + retried, + settled_spend, + summarize, +) + + +def _ok_turn(turn_index: int) -> TurnMetric: + return TurnMetric( + turn_index=turn_index, + ok=True, + latency_seconds=1.0, + uncached_input_tokens=10, + cache_read_tokens=100, + cache_creation_tokens=5, + failure=None, + ) + + +def _failed_turn(turn_index: int) -> TurnMetric: + return TurnMetric( + turn_index=turn_index, + ok=False, + latency_seconds=1.0, + uncached_input_tokens=0, + cache_read_tokens=0, + cache_creation_tokens=0, + failure="NetworkError()", + ) + + +class TestSummarizePlannedTurns: + def test_session_aborted_on_first_turn_counts_all_its_planned_turns_as_failed( + self, + ) -> None: + completed_session = tuple(_ok_turn(index) for index in range(1, 7)) + aborted_session = (_failed_turn(1),) + + report = summarize((*completed_session, *aborted_session), planned_turns=12) + + assert report.attempted_turns == 7 + assert report.failed_turns == 6 + assert report.error_ratio == 0.5 + + def test_all_planned_turns_completing_reports_zero_failures(self) -> None: + report = summarize( + tuple(_ok_turn(index) for index in range(1, 7)), planned_turns=6 + ) + + assert report.failed_turns == 0 + assert report.error_ratio == 0.0 + + +class TestRetried: + def test_transient_failures_then_success_returns_the_success(self) -> None: + outcome = Success[SessionMessagesResponse](status_code=200, data=SessionMessagesResponse()) + calls = iter( + (NetworkError(message="overloaded"), NetworkError(message="overloaded"), outcome) + ) + + result = retried(lambda: next(calls), attempts=3, sleep=lambda _: None) + + assert result is outcome + + def test_exhausted_attempts_return_the_last_failure(self) -> None: + last_attempt = NetworkError(message="still overloaded") + never_reached = NetworkError(message="a fourth attempt would break the budget") + calls = iter( + (NetworkError(message="overloaded"), last_attempt, never_reached) + ) + + result = retried(lambda: next(calls), attempts=2, sleep=lambda _: None) + + assert result is last_attempt + assert next(calls) is never_reached + + def test_first_try_success_never_sleeps(self) -> None: + def sleep_means_retry(_: float) -> None: + raise AssertionError("slept after a successful attempt") + + result = retried( + lambda: Success[SessionMessagesResponse](status_code=200, data=SessionMessagesResponse()), + attempts=3, + sleep=sleep_means_retry, + ) + + assert isinstance(result, Success) + + +class TestSettledSpend: + def test_partial_total_between_batch_flushes_is_not_accepted_as_final(self) -> None: + reads = iter((0.1, 0.1, 0.1, 0.35, 0.35, 0.35, 0.35, 0.35)) + ticks = count(0.0, 2.5) + + spend = settled_spend( + lambda: next(reads), + poll_interval=5.0, + settle_seconds=10.0, + timeout_seconds=100.0, + now=lambda: next(ticks), + sleep=lambda _: None, + ) + + assert spend == 0.35 + + def test_spend_that_never_stabilizes_raises(self) -> None: + reads = (0.1 * step for step in count(1)) + ticks = count(0.0, 2.5) + + with pytest.raises(AssertionError, match="spend anomaly"): + settled_spend( + lambda: next(reads), + poll_interval=5.0, + settle_seconds=5.0, + timeout_seconds=10.0, + now=lambda: next(ticks), + sleep=lambda _: None, + ) + + def test_spend_that_never_becomes_nonzero_raises(self) -> None: + reads = repeat(0.0) + ticks = count(0.0, 2.5) + + with pytest.raises(AssertionError, match="spend anomaly"): + settled_spend( + lambda: next(reads), + poll_interval=5.0, + settle_seconds=5.0, + timeout_seconds=10.0, + now=lambda: next(ticks), + sleep=lambda _: None, + ) diff --git a/tests/e2e/load/test_weekly_session_anomaly_e2e.py b/tests/e2e/load/test_weekly_session_anomaly_e2e.py new file mode 100644 index 00000000000..d4ef883702e --- /dev/null +++ b/tests/e2e/load/test_weekly_session_anomaly_e2e.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import pytest + +from e2e_config import ( + ANOMALY_MAX_ERROR_RATIO, + ANOMALY_MAX_KEY_SPEND_USD, + ANOMALY_MAX_P95_TURN_SECONDS, + ANOMALY_MIN_WARM_CACHE_READ_SHARE, + ANOMALY_SESSIONS, + ANOMALY_SPEND_SETTLE_SECONDS, + ANOMALY_TURN_ATTEMPTS, + ANOMALY_TURNS_PER_SESSION, + unique_marker, +) +from lifecycle import ResourceManager +from load_client import LoadClient +from models import KeyGenerateBody, LiteLLMParamsBody +from proxy_client import ProxyClient +from session_anomaly import run_concurrent_sessions, settled_spend, summarize + +pytestmark = [pytest.mark.e2e, pytest.mark.load, pytest.mark.weekly] + + +@dataclass(frozen=True, slots=True) +class AnomalyRoute: + route_id: str + params: LiteLLMParamsBody + + +ANOMALY_ROUTES = ( + AnomalyRoute( + route_id="anthropic", + params=LiteLLMParamsBody(model="anthropic/claude-sonnet-5"), + ), + AnomalyRoute( + route_id="bedrock_invoke", + params=LiteLLMParamsBody( + model="bedrock/invoke/us.anthropic.claude-sonnet-5", + aws_region_name="us-east-1", + ), + ), +) + + +def _route_id(route: AnomalyRoute) -> str: + return route.route_id + + +def _settled_key_spend(proxy: ProxyClient, key: str) -> float: + return settled_spend( + lambda: proxy.key_info(key).spend or 0.0, + proxy.poll_interval, + ANOMALY_SPEND_SETTLE_SECONDS, + proxy.poll_timeout, + ) + + +class TestWeeklySessionAnomaly: + @pytest.mark.covers("reliability.perf.session_anomaly.under_slo") + @pytest.mark.parametrize("route", ANOMALY_ROUTES, ids=_route_id) + def test_session_load_stays_within_baselines( + self, client: LoadClient, resources: ResourceManager, route: AnomalyRoute + ) -> None: + model_name = f"weekly-anomaly-{route.route_id}-{unique_marker()}" + model_id = client.proxy.create_model(model_name, route.params) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = client.proxy.generate_key( + KeyGenerateBody(models=[model_name], key_alias=model_name) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + + turns = run_concurrent_sessions( + client.proxy.transport, + key, + model_name, + ANOMALY_SESSIONS, + ANOMALY_TURNS_PER_SESSION, + ANOMALY_TURN_ATTEMPTS, + ) + report = summarize(turns, ANOMALY_SESSIONS * ANOMALY_TURNS_PER_SESSION) + failures = tuple(turn.failure for turn in turns if turn.failure) + print(f"{route.route_id} anomaly report: {report}") + + assert report.error_ratio <= ANOMALY_MAX_ERROR_RATIO, ( + f"{route.route_id}: {report.failed_turns}/{report.planned_turns} planned " + f"turns failed or never ran because their session aborted " + f"({report.error_ratio:.1%} > {ANOMALY_MAX_ERROR_RATIO:.1%} allowed); " + f"error rate is anomalously high. Failures: {failures}" + ) + assert report.warm_turns > 0, ( + f"{route.route_id}: no session got past its first turn, so cache and " + f"latency baselines have nothing to read. Failures: {failures}" + ) + assert report.warm_cache_read_share >= ANOMALY_MIN_WARM_CACHE_READ_SHARE, ( + f"{route.route_id}: warm turns read only {report.warm_cache_read_share:.1%} " + f"of billed input tokens from the prompt cache " + f"(read={report.warm_cache_read_tokens}, " + f"creation={report.warm_cache_creation_tokens}, " + f"uncached={report.warm_uncached_input_tokens}), below the " + f"{ANOMALY_MIN_WARM_CACHE_READ_SHARE:.0%} floor; the cached prefix is " + f"being invalidated between turns (the mid-conversation-system cache " + f"collapse signature) or caching stopped working" + ) + assert report.warm_cache_creation_tokens > 0, ( + f"{route.route_id}: warm turns wrote 0 cache-creation tokens across " + f"{report.warm_turns} turns; the moving cache breakpoint stopped writing " + f"new prefix increments" + ) + assert report.p95_turn_seconds <= ANOMALY_MAX_P95_TURN_SECONDS, ( + f"{route.route_id}: p95 turn time {report.p95_turn_seconds:.1f}s exceeds " + f"the {ANOMALY_MAX_P95_TURN_SECONDS:.0f}s ceiling under " + f"{ANOMALY_SESSIONS} concurrent sessions; turn times are anomalously slow" + ) + + spend = _settled_key_spend(client.proxy, key) + assert spend <= ANOMALY_MAX_KEY_SPEND_USD, ( + f"{route.route_id}: gateway recorded ${spend:.4f} for " + f"{report.attempted_turns} turns, above the " + f"${ANOMALY_MAX_KEY_SPEND_USD} ceiling; spend per session is " + f"anomalously high (cache regressions surface here as 2-3x spend)" + ) diff --git a/tests/e2e/load/weekly_anomaly_config.yml b/tests/e2e/load/weekly_anomaly_config.yml new file mode 100644 index 00000000000..08972969cf0 --- /dev/null +++ b/tests/e2e/load/weekly_anomaly_config.yml @@ -0,0 +1,3 @@ +general_settings: + master_key: os.environ/LITELLM_MASTER_KEY + store_model_in_db: true diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index a9dedac8e61..cdc31aeea79 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -14,6 +14,10 @@ from e2e_http import NoBody, ProbeResult, Result, StreamingResponse, Success, Un from models import ( ChatBody, ChatMessage, + CustomerDeleteBody, + CustomerInfoParams, + CustomerNewBody, + CustomerResponse, KeyBlockBody, KeyDeleteBody, KeyGenerateBody, @@ -270,6 +274,35 @@ class ManagementClient: ) ).user_id + def create_customer(self, user_id: str) -> str: + _ = unwrap( + self.proxy.transport.post( + "/customer/new", + headers=self.proxy.transport.master, + json=CustomerNewBody(user_id=user_id), + response_type=CustomerResponse, + ) + ) + return user_id + + def customer_info(self, end_user_id: str) -> CustomerResponse: + return unwrap( + self.proxy.transport.get( + "/customer/info", + headers=self.proxy.transport.master, + params=CustomerInfoParams(end_user_id=end_user_id), + response_type=CustomerResponse, + ) + ) + + def delete_customer(self, user_id: str) -> None: + _ = self.proxy.transport.post( + "/customer/delete", + headers=self.proxy.transport.master, + json=CustomerDeleteBody(user_ids=[user_id]), + response_type=NoBody, + ) + def update_user(self, body: UserUpdateBody) -> None: _ = unwrap( self.proxy.transport.post( diff --git a/tests/e2e/management/test_budget_customer_user_org_e2e.py b/tests/e2e/management/test_budget_customer_user_org_e2e.py new file mode 100644 index 00000000000..54cc18b228b --- /dev/null +++ b/tests/e2e/management/test_budget_customer_user_org_e2e.py @@ -0,0 +1,415 @@ +"""Live e2e coverage for the budget, customer/end-user, user-info and +organization-membership management routes. + +Each test creates its resources under unique ids (deleted on teardown) and +asserts the recorded state the route promises: the budget table reflects a +create/update, a customer round-trips through the info route and disappears after +delete, /user/info echoes what /user/new stored, and an added org member shows up +both in the add response and in /organization/info. The budget/new admin gate is +proven by driving the route under a non-admin key and asserting it is refused. + +Response bodies validate into local pydantic models (only the fields asserted are +modelled) so a shape change fails here instead of passing vacuously. +""" + +from __future__ import annotations + +import math +import time +from collections.abc import Callable + +import pytest +from pydantic import BaseModel, RootModel + +from e2e_config import unique_marker +from e2e_http import NoBody, unwrap +from lifecycle import ResourceManager +from management_client import ManagementClient +from models import KeyGenerateBody, OrgInfoParams, OrgNewBody, UserNewBody + +pytestmark = pytest.mark.e2e + + +def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T: + deadline = time.monotonic() + client.proxy.poll_timeout + while time.monotonic() < deadline: + found = attempt() + if found is not None: + return found + time.sleep(client.proxy.poll_interval) + pytest.fail(failure) + + +# ---------- budget ---------- + + +class BudgetNewBody(BaseModel): + max_budget: float + soft_budget: float | None = None + budget_duration: str | None = None + + +class BudgetNewResponse(BaseModel): + budget_id: str + + +class BudgetUpdateBody(BaseModel): + budget_id: str + max_budget: float + + +class BudgetInfoBody(BaseModel): + budgets: list[str] + + +class BudgetRow(BaseModel): + budget_id: str | None = None + max_budget: float | None = None + soft_budget: float | None = None + + +class BudgetInfoResponse(RootModel[list[BudgetRow]]): + pass + + +class BudgetListResponse(RootModel[list[BudgetRow]]): + """GET /budget/list answers with a bare array of budget rows, not an object + wrapping them. Read the rows off .root.""" + + +class BudgetDeleteBody(BaseModel): + id: str + + +def _delete_budget(client: ManagementClient, budget_id: str) -> None: + _ = client.proxy.transport.post( + "/budget/delete", + headers=client.proxy.transport.master, + json=BudgetDeleteBody(id=budget_id), + response_type=NoBody, + ) + + +def _create_budget(client: ManagementClient, resources: ResourceManager, body: BudgetNewBody) -> str: + budget_id = unwrap( + client.proxy.transport.post( + "/budget/new", + headers=client.proxy.transport.master, + json=body, + response_type=BudgetNewResponse, + ) + ).budget_id + resources.defer(lambda: _delete_budget(client, budget_id)) + return budget_id + + +def _budget_rows(client: ManagementClient, budget_id: str) -> tuple[BudgetRow, ...]: + return tuple( + unwrap( + client.proxy.transport.post( + "/budget/info", + headers=client.proxy.transport.master, + json=BudgetInfoBody(budgets=[budget_id]), + response_type=BudgetInfoResponse, + ) + ).root + ) + + +def _budget_list_ids(client: ManagementClient) -> tuple[str, ...]: + return tuple( + row.budget_id + for row in unwrap( + client.proxy.transport.get( + "/budget/list", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=BudgetListResponse, + ) + ).root + if row.budget_id is not None + ) + + +_INITIAL_MAX_BUDGET = 5.5 +_UPDATED_MAX_BUDGET = 91.25 + + +class TestBudgetManagement: + @pytest.mark.covers("mgmt.budget.list.happy_path") + def test_created_budget_appears_in_budget_list( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + budget_id = _create_budget(client, resources, BudgetNewBody(max_budget=_INITIAL_MAX_BUDGET)) + + _ = _poll( + client, + lambda: budget_id if budget_id in _budget_list_ids(client) else None, + f"/budget/list never included the created budget {budget_id}", + ) + + @pytest.mark.covers("mgmt.budget.update.persists") + def test_update_max_budget_persists_to_budget_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + budget_id = _create_budget(client, resources, BudgetNewBody(max_budget=_INITIAL_MAX_BUDGET)) + + rows = _budget_rows(client, budget_id) + assert rows, f"/budget/info returned nothing for the freshly created budget {budget_id}" + initial = rows[0].max_budget + assert initial is not None and math.isclose(initial, _INITIAL_MAX_BUDGET, rel_tol=1e-9), ( + f"/budget/info reports max_budget {initial}, created with {_INITIAL_MAX_BUDGET}" + ) + + _ = unwrap( + client.proxy.transport.post( + "/budget/update", + headers=client.proxy.transport.master, + json=BudgetUpdateBody(budget_id=budget_id, max_budget=_UPDATED_MAX_BUDGET), + response_type=NoBody, + ) + ) + + def updated() -> BudgetRow | None: + row = next((r for r in _budget_rows(client, budget_id) if r.budget_id == budget_id), None) + if row is None or row.max_budget is None: + return None + return row if math.isclose(row.max_budget, _UPDATED_MAX_BUDGET, rel_tol=1e-9) else None + + _ = _poll( + client, + updated, + f"/budget/info never reported max_budget {_UPDATED_MAX_BUDGET} for {budget_id} after /budget/update", + ) + + @pytest.mark.covers("mgmt.budget.new.admin_only") + def test_new_is_refused_for_a_non_admin_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + key = client.proxy.generate_key(KeyGenerateBody()) + resources.defer(lambda: client.proxy.delete_key(key)) + + outcome = client.proxy.transport.send( + "/budget/new", + headers=client.proxy.transport.bearer(key), + json=BudgetNewBody(max_budget=1.0), + ) + + assert outcome.status_code in (401, 403), ( + f"non-admin key POSTing /budget/new must be refused 401/403, got " + f"{outcome.status_code}: {outcome.body[:300]}" + ) + assert "proxy admin" in outcome.body.lower() or "not allowed" in outcome.body.lower(), ( + f"/budget/new denial body must name the admin-only gate, got: {outcome.body[:300]}" + ) + + +# ---------- customer / end-user ---------- + + +class CustomerNewBody(BaseModel): + user_id: str + max_budget: float | None = None + + +class CustomerNewResponse(BaseModel): + user_id: str + + +class CustomerInfoParams(BaseModel): + end_user_id: str + + +class CustomerInfoResponse(BaseModel): + user_id: str + + +class CustomerDeleteBody(BaseModel): + user_ids: list[str] + + +class CustomerDeleteResponse(BaseModel): + deleted_customers: int + + +def _create_customer( + client: ManagementClient, resources: ResourceManager, route: str, body: CustomerNewBody +) -> str: + user_id = unwrap( + client.proxy.transport.post( + route, + headers=client.proxy.transport.master, + json=body, + response_type=CustomerNewResponse, + ) + ).user_id + resources.defer(lambda: client.proxy.delete_customers([user_id])) + return user_id + + +def _customer_info(client: ManagementClient, route: str, user_id: str) -> CustomerInfoResponse: + return unwrap( + client.proxy.transport.get( + route, + headers=client.proxy.transport.master, + params=CustomerInfoParams(end_user_id=user_id), + response_type=CustomerInfoResponse, + ) + ) + + +class TestCustomerManagement: + @pytest.mark.covers("mgmt.customer.new.happy_path") + def test_new_persists_to_customer_info(self, client: ManagementClient, resources: ResourceManager) -> None: + customer_id = f"e2e-mgmt-cust-{unique_marker()}" + created = _create_customer( + client, resources, "/customer/new", CustomerNewBody(user_id=customer_id, max_budget=7.0) + ) + assert created == customer_id, f"/customer/new echoed user_id {created!r}, created {customer_id!r}" + + info = _customer_info(client, "/customer/info", customer_id) + assert info.user_id == customer_id, ( + f"/customer/info reports user_id {info.user_id!r} for the created customer {customer_id!r}" + ) + + @pytest.mark.covers("mgmt.customer.delete.persists") + def test_delete_removes_the_customer(self, client: ManagementClient, resources: ResourceManager) -> None: + """The teardown's deferred delete fires again on the already-deleted customer + by design: it is the safety net if this test fails before the in-body delete, + and a repeat /customer/delete is absorbed by the warn-only teardown.""" + customer_id = f"e2e-mgmt-cust-{unique_marker()}" + _ = _create_customer(client, resources, "/customer/new", CustomerNewBody(user_id=customer_id, max_budget=3.0)) + + assert _customer_info(client, "/customer/info", customer_id).user_id == customer_id, ( + f"customer {customer_id} was not readable before deletion" + ) + + deleted = unwrap( + client.proxy.transport.post( + "/customer/delete", + headers=client.proxy.transport.master, + json=CustomerDeleteBody(user_ids=[customer_id]), + response_type=CustomerDeleteResponse, + ) + ).deleted_customers + assert deleted == 1, f"/customer/delete reported {deleted} rows removed for one customer" + + def gone() -> bool | None: + return True if client.proxy.transport.probe( + "/customer/info", params=CustomerInfoParams(end_user_id=customer_id) + ).status_code == 404 else None + + _ = _poll(client, gone, f"customer {customer_id} still resolved on /customer/info after /customer/delete") + + @pytest.mark.covers("mgmt.end_user.new.happy_path") + def test_end_user_new_persists_to_end_user_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + end_user_id = f"e2e-mgmt-euser-{unique_marker()}" + created = _create_customer(client, resources, "/end_user/new", CustomerNewBody(user_id=end_user_id)) + assert created == end_user_id, f"/end_user/new echoed user_id {created!r}, created {end_user_id!r}" + + info = _customer_info(client, "/end_user/info", end_user_id) + assert info.user_id == end_user_id, ( + f"/end_user/info reports user_id {info.user_id!r} for the created end user {end_user_id!r}" + ) + + +# ---------- user info ---------- + + +class TestUserManagement: + @pytest.mark.covers("mgmt.user.info.happy_path") + def test_new_user_is_readable_via_user_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + email = f"e2e-mgmt-{unique_marker()}@example.com" + user_id = client.create_user(UserNewBody(user_email=email, user_role="internal_user")) + resources.defer(lambda: client.delete_user(user_id)) + + info = client.user_info(user_id).user_info + assert info.user_id == user_id, f"/user/info reports user_id {info.user_id!r}, created {user_id!r}" + assert info.user_email == email, f"/user/info reports user_email {info.user_email!r}, configured {email!r}" + assert info.user_role == "internal_user", ( + f"/user/info reports user_role {info.user_role!r}, configured 'internal_user'" + ) + + +# ---------- organization membership ---------- + + +class OrgMemberEntry(BaseModel): + role: str + user_id: str + + +class OrgMemberAddBody(BaseModel): + organization_id: str + member: OrgMemberEntry + + +class OrgMembershipRow(BaseModel): + user_id: str + organization_id: str | None = None + + +class OrgMemberAddResponse(BaseModel): + organization_id: str + updated_organization_memberships: list[OrgMembershipRow] + + +class OrgInfoMembersResponse(BaseModel): + members: list[OrgMembershipRow] = [] + + +class TestOrganizationMembership: + @pytest.mark.covers("mgmt.organization.member_add.happy_path") + def test_member_add_records_membership( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + org_id = client.create_org(OrgNewBody(organization_alias=f"e2e-mgmt-org-{unique_marker()}")) + resources.defer(lambda: client.delete_org(org_id)) + + user_id = client.create_user( + UserNewBody(user_email=f"e2e-mgmt-{unique_marker()}@example.com", user_role="internal_user") + ) + resources.defer(lambda: client.delete_user(user_id)) + + added = unwrap( + client.proxy.transport.post( + "/organization/member_add", + headers=client.proxy.transport.master, + json=OrgMemberAddBody( + organization_id=org_id, + member=OrgMemberEntry(role="internal_user", user_id=user_id), + ), + response_type=OrgMemberAddResponse, + ) + ) + assert added.organization_id == org_id, ( + f"/organization/member_add echoed organization_id {added.organization_id!r}, added to {org_id!r}" + ) + assert any( + row.user_id == user_id and row.organization_id == org_id + for row in added.updated_organization_memberships + ), ( + f"/organization/member_add response does not record {user_id} in org {org_id}: " + f"{added.updated_organization_memberships}" + ) + + def listed() -> bool | None: + members = unwrap( + client.proxy.transport.get( + "/organization/info", + headers=client.proxy.transport.master, + params=OrgInfoParams(organization_id=org_id), + response_type=OrgInfoMembersResponse, + ) + ).members + return True if any(member.user_id == user_id for member in members) else None + + _ = _poll( + client, + listed, + f"/organization/info never listed member {user_id} in org {org_id} after /organization/member_add", + ) diff --git a/tests/e2e/management/test_config_misc_endpoints_e2e.py b/tests/e2e/management/test_config_misc_endpoints_e2e.py new file mode 100644 index 00000000000..6c4de621271 --- /dev/null +++ b/tests/e2e/management/test_config_misc_endpoints_e2e.py @@ -0,0 +1,698 @@ +"""Live e2e: the config and miscellaneous Management/UI routes. + +One method per registry cell, each asserting the real contract against a live +proxy: read-only inventory routes return their documented shape, stateless +validators compute their verdict from the request, and the write routes persist +so a read-back reflects the change. The two routes that mutate global proxy state +(cache settings and router settings, both driven from the admin UI) are exercised +with a benign, self-restoring change so a shared proxy is left as it was found. +""" + +from __future__ import annotations + +import math +import time +from collections.abc import Callable + +import pytest +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import NoBody, Success, unwrap, unwrap_status +from lifecycle import ResourceManager +from management_client import ManagementClient +from models import KeyGenerateBody, LiteLLMParamsBody, TeamNewBody + +pytestmark = pytest.mark.e2e + + +def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T: + deadline = time.monotonic() + client.proxy.poll_timeout + while time.monotonic() < deadline: + found = attempt() + if found is not None: + return found + time.sleep(client.proxy.poll_interval) + pytest.fail(failure) + + +# ---- callbacks ------------------------------------------------------------- + + +class CallbacksListResponse(BaseModel): + success: list[str] + failure: list[str] + success_and_failure: list[str] + + +# ---- cost estimate --------------------------------------------------------- + + +class CostEstimateBody(BaseModel): + model: str + input_tokens: int + output_tokens: int + num_requests_per_day: int | None = None + + +class CostEstimateResponse(BaseModel): + model: str + input_tokens: int + output_tokens: int + cost_per_request: float + input_cost_per_request: float + output_cost_per_request: float + margin_cost_per_request: float + daily_cost: float | None = None + provider: str | None = None + + +# ---- credential migration check -------------------------------------------- + + +class MigrationReport(BaseModel): + residual_legacy: int + total_undecryptable: int + + +class MigrationCheckResponse(BaseModel): + status: str + report: MigrationReport + + +# ---- tool + workflow inventories ------------------------------------------- + + +class ToolListEntry(BaseModel): + name: str | None = None + + +class ToolListResponse(BaseModel): + tools: list[ToolListEntry] + total: int + + +class WorkflowRunEntry(BaseModel): + workflow_id: str | None = None + + +class WorkflowRunsResponse(BaseModel): + runs: list[WorkflowRunEntry] + count: int + + +# ---- compliance ------------------------------------------------------------ + + +class ComplianceGdprBody(BaseModel): + request_id: str + user_id: str + model: str + timestamp: str + + +class ComplianceCheck(BaseModel): + check_name: str + article: str + passed: bool + detail: str + + +class ComplianceResponse(BaseModel): + compliant: bool + regulation: str + checks: list[ComplianceCheck] + + +# ---- cache settings -------------------------------------------------------- + + +class CacheSettingsValue(BaseModel): + type: str + host: str = "" + port: str = "" + + +class CacheSettingsUpdateBody(BaseModel): + cache_settings: CacheSettingsValue + + +class CacheCurrentValues(BaseModel): + type: str | None = None + host: str | None = None + port: str | None = None + + +class CacheGetResponse(BaseModel): + current_values: CacheCurrentValues + + +class CacheUpdateResponse(BaseModel): + status: str + settings: CacheSettingsValue + + +# ---- fallback management --------------------------------------------------- + + +class FallbackShape(BaseModel): + model: str + fallback_models: list[str] + fallback_type: str + + +class FallbackCreateBody(FallbackShape): + pass + + +class FallbackResponse(FallbackShape): + message: str + + +class FallbackGetParams(BaseModel): + fallback_type: str + + +class FallbackGetResponse(FallbackShape): + pass + + +# ---- jwt key mapping ------------------------------------------------------- + + +class JwtKeyMappingNewBody(BaseModel): + jwt_claim_name: str + jwt_claim_value: str + key: str + description: str + + +class JwtInfoParams(BaseModel): + id: str + + +class JwtDeleteBody(BaseModel): + id: str + + +class JwtKeyMappingResponse(BaseModel): + id: str + jwt_claim_name: str + jwt_claim_value: str + is_active: bool + description: str | None = None + + +# ---- router settings via /config/update ------------------------------------ + + +class RouterSettingsPatch(BaseModel): + num_retries: int + + +class ConfigUpdateBody(BaseModel): + router_settings: RouterSettingsPatch + + +class ConfigUpdateResponse(BaseModel): + message: str + + +class RouterCurrentValues(BaseModel): + num_retries: int | None = None + + +class RouterSettingsResponse(BaseModel): + current_values: RouterCurrentValues + + +# ---- mcp server submission ------------------------------------------------- + + +class McpRegisterBody(BaseModel): + server_name: str + url: str + transport: str + description: str + + +class McpServerResponse(BaseModel): + server_id: str + server_name: str | None = None + approval_status: str + transport: str + url: str | None = None + + +class TestInventoryRoutes: + @pytest.mark.covers("mgmt.callback.list.happy_path") + def test_callbacks_list_reports_active_logging_callbacks(self, client: ManagementClient) -> None: + listing = unwrap( + client.proxy.transport.get( + "/callbacks/list", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=CallbacksListResponse, + ) + ) + every = [*listing.success, *listing.failure, *listing.success_and_failure] + assert every, "/callbacks/list reported no active logging callbacks; the proxy always runs the db logger" + assert "_ProxyDBLogger" in every, ( + f"/callbacks/list omitted the always-on _ProxyDBLogger spend logger; got {every}" + ) + + @pytest.mark.covers("mgmt.tool_management.list.happy_path") + def test_tool_list_returns_catalog_with_consistent_total(self, client: ManagementClient) -> None: + listing = unwrap( + client.proxy.transport.get( + "/v1/tool/list", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=ToolListResponse, + ) + ) + assert listing.total == len(listing.tools), ( + f"/v1/tool/list total {listing.total} disagrees with the {len(listing.tools)} tools returned" + ) + + @pytest.mark.covers("mgmt.workflow.list.happy_path") + def test_workflow_runs_list_returns_consistent_count(self, client: ManagementClient) -> None: + listing = unwrap( + client.proxy.transport.get( + "/v1/workflows/runs", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=WorkflowRunsResponse, + ) + ) + assert listing.count == len(listing.runs), ( + f"/v1/workflows/runs count {listing.count} disagrees with the {len(listing.runs)} runs returned" + ) + + @pytest.mark.covers("mgmt.credential_migration.check.happy_path") + def test_credential_migration_check_reports_residual_scan(self, client: ManagementClient) -> None: + report = unwrap( + client.proxy.transport.get( + "/credentials/migrate-encryption/check", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=MigrationCheckResponse, + ) + ) + assert report.status == "success", f"migrate-encryption/check status {report.status!r}, expected 'success'" + assert report.report.residual_legacy >= 0, ( + f"residual_legacy count is negative ({report.report.residual_legacy}); the scan is broken" + ) + assert report.report.total_undecryptable >= 0, ( + f"total_undecryptable count is negative ({report.report.total_undecryptable}); the scan is broken" + ) + + +class TestCostEstimate: + @pytest.mark.covers("mgmt.cost_tracking.estimate.happy_path") + def test_estimate_computes_cost_from_token_counts(self, client: ManagementClient) -> None: + estimate = unwrap( + client.proxy.transport.post( + "/cost/estimate", + headers=client.proxy.transport.master, + json=CostEstimateBody( + model="gpt-4o-mini", input_tokens=1000, output_tokens=500, num_requests_per_day=100 + ), + response_type=CostEstimateResponse, + ) + ) + assert estimate.input_cost_per_request > 0, ( + f"input cost per request is {estimate.input_cost_per_request}; a priced model must cost more than zero" + ) + assert estimate.output_cost_per_request > 0, ( + f"output cost per request is {estimate.output_cost_per_request}; a priced model must cost more than zero" + ) + expected_per_request = ( + estimate.input_cost_per_request + estimate.output_cost_per_request + estimate.margin_cost_per_request + ) + assert math.isclose(estimate.cost_per_request, expected_per_request, rel_tol=1e-9), ( + f"cost_per_request {estimate.cost_per_request} != input+output+margin {expected_per_request}" + ) + assert estimate.daily_cost is not None and math.isclose( + estimate.daily_cost, estimate.cost_per_request * 100, rel_tol=1e-9 + ), f"daily_cost {estimate.daily_cost} != cost_per_request * 100 requests {estimate.cost_per_request * 100}" + + +class TestComplianceRoutes: + @pytest.mark.covers("mgmt.compliance.gdpr.happy_path") + def test_gdpr_check_derives_verdict_from_the_request(self, client: ManagementClient) -> None: + result = unwrap( + client.proxy.transport.post( + "/compliance/gdpr", + headers=client.proxy.transport.master, + json=ComplianceGdprBody( + request_id=f"e2e-gdpr-{unique_marker()}", + user_id=f"e2e-user-{unique_marker()}", + model="gpt-4o-mini", + timestamp="2026-07-21T00:00:00Z", + ), + response_type=ComplianceResponse, + ) + ) + assert result.regulation == "GDPR", ( + f"/compliance/gdpr reported regulation {result.regulation!r}, expected 'GDPR'" + ) + articles = {check.article for check in result.checks} + assert articles == {"Art. 32", "Art. 5(1)(c)", "Art. 30"}, ( + f"/compliance/gdpr returned articles {articles}, expected the three GDPR articles" + ) + assert result.compliant == all(check.passed for check in result.checks), ( + "the overall compliant verdict must be the conjunction of the individual checks" + ) + assert all(check.check_name and check.detail for check in result.checks), ( + "every compliance check must carry a name and a human-readable detail" + ) + + +class TestCacheSettings: + @pytest.mark.covers("mgmt.cache_settings.update.happy_path") + def test_update_persists_cache_backend_to_get( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """Exercise the update route without changing global state: capture the live + cache backend and write exactly that back, so the config the proxy ends on is + byte-for-byte the one it started with. A teardown restore of the same captured + settings is the safety net if the body fails partway. The update route is only + meaningful against a configured cache, so an unconfigured proxy fails loudly + here rather than being silently switched to redis.""" + before = self._read_settings(client) + assert before.type is not None, ( + "GET /cache/settings reported no cache type; refusing to invent one and mutate the shared proxy" + ) + captured = CacheSettingsValue(type=before.type, host=before.host or "", port=before.port or "") + resources.defer(lambda: self._write_settings(client, captured)) + + updated = unwrap( + client.proxy.transport.post( + "/cache/settings", + headers=client.proxy.transport.master, + json=CacheSettingsUpdateBody(cache_settings=captured), + response_type=CacheUpdateResponse, + ) + ) + assert updated.status == "success", f"/cache/settings update status {updated.status!r}, expected 'success'" + assert updated.settings.type == captured.type, ( + f"/cache/settings echoed type {updated.settings.type!r}, wrote {captured.type!r}" + ) + + def reflected() -> CacheCurrentValues | None: + current = self._read_settings(client) + return current if current.type == captured.type else None + + after = _poll(client, reflected, f"/cache/settings never reported type {captured.type!r} after the update") + assert after.host == captured.host and after.port == captured.port, ( + f"/cache/settings persisted host/port {after.host!r}/{after.port!r}, " + f"wrote {captured.host!r}/{captured.port!r}" + ) + + @staticmethod + def _read_settings(client: ManagementClient) -> CacheCurrentValues: + return unwrap( + client.proxy.transport.get( + "/cache/settings", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=CacheGetResponse, + ) + ).current_values + + @staticmethod + def _write_settings(client: ManagementClient, settings: CacheSettingsValue) -> None: + _ = unwrap( + client.proxy.transport.post( + "/cache/settings", + headers=client.proxy.transport.master, + json=CacheSettingsUpdateBody(cache_settings=settings), + response_type=CacheUpdateResponse, + ) + ) + + +class TestFallbackManagement: + @pytest.mark.covers("mgmt.fallback_management.update.happy_path") + def test_create_persists_and_is_read_back(self, client: ManagementClient, resources: ResourceManager) -> None: + primary = f"e2e-fallback-primary-{unique_marker()}" + secondary = f"e2e-fallback-secondary-{unique_marker()}" + params = LiteLLMParamsBody(model="openai/gpt-5.5", api_key="e2e-dummy-key") + primary_id = client.proxy.create_model(primary, params) + resources.defer(lambda: client.proxy.delete_model(primary_id)) + secondary_id = client.proxy.create_model(secondary, params) + resources.defer(lambda: client.proxy.delete_model(secondary_id)) + resources.defer(lambda: self._delete_fallback(client, primary)) + + created = unwrap( + client.proxy.transport.post( + "/fallback", + headers=client.proxy.transport.master, + json=FallbackCreateBody(model=primary, fallback_models=[secondary], fallback_type="general"), + response_type=FallbackResponse, + ) + ) + assert created.model == primary and created.fallback_models == [secondary], ( + f"/fallback echoed model={created.model!r} fallbacks={created.fallback_models}, " + f"configured {primary!r} -> [{secondary!r}]" + ) + + def read_back() -> FallbackGetResponse | None: + result = client.proxy.transport.get( + f"/fallback/{primary}", + headers=client.proxy.transport.master, + params=FallbackGetParams(fallback_type="general"), + response_type=FallbackGetResponse, + ) + match result: + case Success(data=data) if secondary in data.fallback_models: + return data + case _: + return None + + got = _poll(client, read_back, f"GET /fallback/{primary} never reported {secondary} after /fallback") + assert got.fallback_models == [secondary], ( + f"GET /fallback/{primary} reports fallbacks {got.fallback_models}, configured [{secondary!r}]" + ) + + @staticmethod + def _delete_fallback(client: ManagementClient, model: str) -> None: + _ = client.proxy.transport.delete( + f"/fallback/{model}", + headers=client.proxy.transport.master, + json=NoBody(), + params=FallbackGetParams(fallback_type="general"), + response_type=NoBody, + ) + + +class TestJwtKeyMapping: + @pytest.mark.covers("mgmt.jwt_key_mapping.new.happy_path") + def test_new_persists_mapping_and_is_read_back( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + key = client.proxy.generate_key(KeyGenerateBody()) + resources.defer(lambda: client.proxy.delete_key(key)) + claim_value = f"e2e_jwt_{unique_marker()}" + + created = unwrap( + client.proxy.transport.post( + "/jwt/key/mapping/new", + headers=client.proxy.transport.master, + json=JwtKeyMappingNewBody( + jwt_claim_name="team_id", + jwt_claim_value=claim_value, + key=key, + description="e2e coverage mapping", + ), + response_type=JwtKeyMappingResponse, + ) + ) + resources.defer(lambda: self._delete_mapping(client, created.id)) + assert created.jwt_claim_value == claim_value and created.is_active, ( + f"/jwt/key/mapping/new returned claim_value={created.jwt_claim_value!r} active={created.is_active}, " + f"configured {claim_value!r} active=True" + ) + + info = unwrap( + client.proxy.transport.get( + "/jwt/key/mapping/info", + headers=client.proxy.transport.master, + params=JwtInfoParams(id=created.id), + response_type=JwtKeyMappingResponse, + ) + ) + assert info.id == created.id and info.jwt_claim_name == "team_id" and info.jwt_claim_value == claim_value, ( + f"/jwt/key/mapping/info reports {info.jwt_claim_name!r}={info.jwt_claim_value!r} for id {info.id}, " + f"created team_id={claim_value!r}" + ) + + @staticmethod + def _delete_mapping(client: ManagementClient, mapping_id: str) -> None: + _ = client.proxy.transport.post( + "/jwt/key/mapping/delete", + headers=client.proxy.transport.master, + json=JwtDeleteBody(id=mapping_id), + response_type=NoBody, + ) + + +class TestRouterSettings: + @pytest.mark.covers("mgmt.router_settings.update.happy_path") + def test_config_update_persists_router_setting_to_get( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """/config/update is the only write path for router_settings (there is no + dedicated router-settings write route). The change is restored on teardown so + the shared proxy keeps its original retry policy.""" + original = self._read_num_retries(client) + assert original is not None, "GET /router/settings did not report num_retries; cannot prove a change" + resources.defer(lambda: self._write_num_retries(client, original)) + + target = original + 5 + response = unwrap( + client.proxy.transport.post( + "/config/update", + headers=client.proxy.transport.master, + json=ConfigUpdateBody(router_settings=RouterSettingsPatch(num_retries=target)), + response_type=ConfigUpdateResponse, + ) + ) + assert "success" in response.message.lower(), ( + f"/config/update reported {response.message!r}, expected a success message" + ) + + _ = _poll( + client, + lambda: True if self._read_num_retries(client) == target else None, + f"GET /router/settings never reported num_retries {target} after /config/update", + ) + + self._write_num_retries(client, original) + restored = _poll( + client, + lambda: original if self._read_num_retries(client) == original else None, + f"GET /router/settings never returned to the original num_retries {original} after the restore", + ) + assert restored == original, f"router num_retries left at {restored}, expected the original {original}" + + @staticmethod + def _read_num_retries(client: ManagementClient) -> int | None: + return unwrap( + client.proxy.transport.get( + "/router/settings", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=RouterSettingsResponse, + ) + ).current_values.num_retries + + @staticmethod + def _write_num_retries(client: ManagementClient, value: int) -> None: + _ = unwrap( + client.proxy.transport.post( + "/config/update", + headers=client.proxy.transport.master, + json=ConfigUpdateBody(router_settings=RouterSettingsPatch(num_retries=value)), + response_type=ConfigUpdateResponse, + ) + ) + + +class TestMcpServerSubmission: + @pytest.mark.covers("mgmt.mcp_server.register.happy_path") + def test_register_submits_pending_server(self, client: ManagementClient, resources: ResourceManager) -> None: + """A non-admin, team-scoped key submits an MCP server for review; the proxy + stores it as pending_review without loading it into the runtime registry.""" + team_id = client.create_team(TeamNewBody(team_alias=f"e2e-mcp-team-{unique_marker()}")) + resources.defer(lambda: client.delete_team(team_id)) + team_key = client.proxy.generate_key(KeyGenerateBody(team_id=team_id)) + resources.defer(lambda: client.proxy.delete_key(team_key)) + + server_name = f"e2e_mcp_{unique_marker()}" + submitted = unwrap_status( + client.proxy.transport.post( + "/v1/mcp/server/register", + headers=client.proxy.transport.bearer(team_key), + json=McpRegisterBody( + server_name=server_name, + url="https://example.com/mcp", + transport="sse", + description="e2e coverage submission", + ), + response_type=McpServerResponse, + ), + 201, + ) + resources.defer(lambda: self._delete_server(client, submitted.server_id)) + assert submitted.approval_status == "pending_review", ( + f"a user submission must be pending_review, got {submitted.approval_status!r}" + ) + assert submitted.server_name == server_name and submitted.transport == "sse", ( + f"/v1/mcp/server/register echoed name={submitted.server_name!r} transport={submitted.transport!r}, " + f"configured {server_name!r}/sse" + ) + + @pytest.mark.covers("mgmt.mcp_server.approve.persists") + def test_approve_activates_submission_and_persists( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """An admin approving a pending submission flips it to active, and the change + persists to a fresh read of the server.""" + team_id = client.create_team(TeamNewBody(team_alias=f"e2e-mcp-team-{unique_marker()}")) + resources.defer(lambda: client.delete_team(team_id)) + team_key = client.proxy.generate_key(KeyGenerateBody(team_id=team_id)) + resources.defer(lambda: client.proxy.delete_key(team_key)) + + submitted = unwrap( + client.proxy.transport.post( + "/v1/mcp/server/register", + headers=client.proxy.transport.bearer(team_key), + json=McpRegisterBody( + server_name=f"e2e_mcp_{unique_marker()}", + url="https://example.com/mcp", + transport="sse", + description="e2e coverage submission", + ), + response_type=McpServerResponse, + ) + ) + resources.defer(lambda: self._delete_server(client, submitted.server_id)) + assert submitted.approval_status == "pending_review", ( + f"a fresh submission must be pending_review before approval, got {submitted.approval_status!r}" + ) + + approved = unwrap( + client.proxy.transport.put( + f"/v1/mcp/server/{submitted.server_id}/approve", + headers=client.proxy.transport.master, + json=NoBody(), + response_type=McpServerResponse, + ) + ) + assert approved.approval_status == "active", ( + f"approve must flip the submission to active, got {approved.approval_status!r}" + ) + + fetched = unwrap( + client.proxy.transport.get( + f"/v1/mcp/server/{submitted.server_id}", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=McpServerResponse, + ) + ) + assert fetched.server_id == submitted.server_id and fetched.approval_status == "active", ( + f"GET /v1/mcp/server/{submitted.server_id} reports approval_status {fetched.approval_status!r} " + "after approve, expected 'active'" + ) + + @staticmethod + def _delete_server(client: ManagementClient, server_id: str) -> None: + _ = client.proxy.transport.delete( + f"/v1/mcp/server/{server_id}", + headers=client.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) diff --git a/tests/e2e/management/test_key_management_e2e.py b/tests/e2e/management/test_key_management_e2e.py new file mode 100644 index 00000000000..711175abb0d --- /dev/null +++ b/tests/e2e/management/test_key_management_e2e.py @@ -0,0 +1,251 @@ +"""Live e2e: the /key management routes' persistence, health, bulk-update, and +admin-only contracts. + +Each test creates its keys under the master key with unique aliases (deleted on +teardown) and asserts the real contract: the info route reflects the write +(persistence), the health route reports the calling key, bulk_update applies to +the target key, and the write routes refuse a non-admin caller. Key writes reach +the auth cache eventually, so the read-backs poll to a deadline instead of +asserting once. +""" + +from __future__ import annotations + +import time +from collections.abc import Callable +from typing import Literal + +import pytest + +from e2e_config import unique_marker +from e2e_http import NoBody, unwrap +from lifecycle import ResourceManager +from management_client import ManagementClient +from models import KeyDeleteBody, KeyGenerateBody, KeyUpdateBody +from pydantic import BaseModel + +pytestmark = pytest.mark.e2e + + +class KeyToggleBlockBody(BaseModel): + key: str + + +class LoggingCallbackStatus(BaseModel): + callbacks: list[str] | None = None + status: str | None = None + details: str | None = None + + +class KeyHealthResponse(BaseModel): + key: Literal["healthy", "unhealthy"] + logging_callbacks: LoggingCallbackStatus | None = None + + +class BulkKeyUpdateItem(BaseModel): + key: str + max_budget: float | None = None + + +class BulkKeyUpdateBody(BaseModel): + keys: list[BulkKeyUpdateItem] + + +class BulkKeyUpdateSuccess(BaseModel): + key: str + + +class BulkKeyUpdateFailure(BaseModel): + key: str + failed_reason: str + + +class BulkKeyUpdateResponse(BaseModel): + total_requested: int + successful_updates: list[BulkKeyUpdateSuccess] + failed_updates: list[BulkKeyUpdateFailure] + + +def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T: + deadline = time.monotonic() + client.proxy.poll_timeout + while time.monotonic() < deadline: + found = attempt() + if found is not None: + return found + time.sleep(client.proxy.poll_interval) + pytest.fail(failure) + + +def _generate_key(client: ManagementClient, resources: ResourceManager, body: KeyGenerateBody) -> str: + key = client.proxy.generate_key(body) + resources.defer(lambda: client.proxy.delete_key(key)) + return key + + +def _block(client: ManagementClient, key: str) -> None: + _ = unwrap( + client.proxy.transport.post( + "/key/block", + headers=client.proxy.transport.master, + json=KeyToggleBlockBody(key=key), + response_type=NoBody, + ) + ) + + +def _unblock(client: ManagementClient, key: str) -> None: + _ = unwrap( + client.proxy.transport.post( + "/key/unblock", + headers=client.proxy.transport.master, + json=KeyToggleBlockBody(key=key), + response_type=NoBody, + ) + ) + + +class TestKeyManagementRoutes: + @pytest.mark.covers("mgmt.key.info.persists") + def test_info_reflects_the_fields_the_key_was_created_with( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + alias = f"e2e-mgmt-keyinfo-{unique_marker()}" + key = _generate_key( + client, + resources, + KeyGenerateBody( + models=["gpt-5.5", "gemini-2.5-flash"], + key_alias=alias, + tpm_limit=131313, + rpm_limit=141414, + ), + ) + + info = client.proxy.key_info(key) + assert info.key_alias == alias, f"/key/info reports key_alias {info.key_alias!r}, configured {alias!r}" + assert info.models == ["gpt-5.5", "gemini-2.5-flash"], ( + f"/key/info reports models {info.models}, configured ['gpt-5.5', 'gemini-2.5-flash']" + ) + assert info.tpm_limit == 131313, f"/key/info reports tpm_limit {info.tpm_limit}, configured 131313" + assert info.rpm_limit == 141414, f"/key/info reports rpm_limit {info.rpm_limit}, configured 141414" + + @pytest.mark.covers("mgmt.key.unblock.persists") + def test_unblock_flips_key_info_blocked_back( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + + _block(client, key) + _ = _poll( + client, + lambda: True if client.proxy.key_info(key).blocked else None, + "/key/info never reported the key blocked after /key/block before the deadline", + ) + + _unblock(client, key) + _ = _poll( + client, + lambda: True if client.proxy.key_info(key).blocked is False else None, + "/key/info never reported the key unblocked after /key/unblock before the deadline", + ) + + @pytest.mark.covers("mgmt.key.health.happy_path") + def test_health_reports_the_calling_key_healthy( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + + health = unwrap( + client.proxy.transport.post( + "/key/health", + headers=client.proxy.transport.bearer(key), + json=NoBody(), + response_type=KeyHealthResponse, + ) + ) + assert health.key == "healthy", f"/key/health reports {health.key!r} for a key with no logging configured" + assert health.logging_callbacks is None, ( + f"/key/health reports logging_callbacks {health.logging_callbacks!r} for a key with no logging configured" + ) + + @pytest.mark.covers("mgmt.key.bulk_update.happy_path") + def test_bulk_update_applies_max_budget_to_target_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"], max_budget=5.0)) + assert client.proxy.key_info(key).max_budget == 5.0, ( + f"/key/info reports max_budget {client.proxy.key_info(key).max_budget}, configured 5.0" + ) + + result = unwrap( + client.proxy.transport.post( + "/key/bulk_update", + headers=client.proxy.transport.master, + json=BulkKeyUpdateBody(keys=[BulkKeyUpdateItem(key=key, max_budget=42.0)]), + response_type=BulkKeyUpdateResponse, + ) + ) + assert result.total_requested == 1, f"/key/bulk_update reports total_requested {result.total_requested}, sent 1" + assert result.failed_updates == [], f"/key/bulk_update reported failed updates: {result.failed_updates}" + assert [entry.key for entry in result.successful_updates] == [key], ( + f"/key/bulk_update successful_updates {[entry.key for entry in result.successful_updates]} did not target {key}" + ) + + _ = _poll( + client, + lambda: True if client.proxy.key_info(key).max_budget == 42.0 else None, + "/key/info never reported max_budget 42.0 after /key/bulk_update before the deadline", + ) + + @pytest.mark.covers("mgmt.key.generate.admin_only") + def test_generate_forbidden_for_non_admin_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + nonadmin = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + + outcome = client.proxy.transport.send( + "/key/generate", + headers=client.proxy.transport.bearer(nonadmin), + json=KeyGenerateBody(models=["gpt-5.5"], key_alias=f"e2e-mgmt-forbidden-{unique_marker()}"), + ) + assert outcome.status_code in (401, 403), ( + f"non-admin key POSTing /key/generate must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}" + ) + + @pytest.mark.covers("mgmt.key.delete.admin_only") + def test_delete_forbidden_for_non_admin_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + nonadmin = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + victim = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + + outcome = client.proxy.transport.send( + "/key/delete", + headers=client.proxy.transport.bearer(nonadmin), + json=KeyDeleteBody(keys=[victim]), + ) + assert outcome.status_code in (401, 403), ( + f"non-admin key POSTing /key/delete must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}" + ) + assert client.proxy.key_info(victim).blocked in (None, False), ( + "victim key should be unaffected by the denied /key/delete" + ) + + @pytest.mark.covers("mgmt.key.update.admin_only") + def test_update_forbidden_for_non_admin_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + nonadmin = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + target = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + + outcome = client.proxy.transport.send( + "/key/update", + headers=client.proxy.transport.bearer(nonadmin), + json=KeyUpdateBody(key=target, models=["gemini-2.5-flash"]), + ) + assert outcome.status_code in (401, 403), ( + f"non-admin key POSTing /key/update must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}" + ) + assert client.proxy.key_info(target).models == ["gpt-5.5"], ( + f"target key models changed to {client.proxy.key_info(target).models} despite the denied /key/update" + ) diff --git a/tests/e2e/management/test_management_e2e.py b/tests/e2e/management/test_management_e2e.py index 18bc384a879..9b398963ac9 100644 --- a/tests/e2e/management/test_management_e2e.py +++ b/tests/e2e/management/test_management_e2e.py @@ -610,3 +610,18 @@ class TestManagementRoutePermissions: f"/team/info returned {team_probe.status_code}: {team_probe.body[:300]}" ) assert client.user_count(user_id) == 0, f"user {user_id} was created despite the 403 route denial" + + +class TestCustomer: + @pytest.mark.covers("mgmt.end_user.new.happy_path") + def test_customer_create_persists_to_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + customer = f"e2e-customer-{unique_marker()}" + client.create_customer(customer) + resources.defer(lambda: client.delete_customer(customer)) + + info = client.customer_info(customer) + assert info.user_id == customer, ( + f"/customer/info did not report the created end-user; got {info.user_id!r}" + ) diff --git a/tests/e2e/management/test_model_tag_accessgroup_e2e.py b/tests/e2e/management/test_model_tag_accessgroup_e2e.py new file mode 100644 index 00000000000..e6a187ae105 --- /dev/null +++ b/tests/e2e/management/test_model_tag_accessgroup_e2e.py @@ -0,0 +1,385 @@ +"""Live e2e: the model, tag, and model-access-group management routes. + +Each test creates its resources under unique names (deleted on teardown) and +asserts the route's contract against a live proxy: the admin-only guard on +adding a global model, the tag inventory round-trip through /tag/list and +/tag/delete, and creating a model access group then reading it back through +/access_group/{name}/info. Reads that lag a write poll to a deadline instead of +asserting once. + +Request bodies for /model/new are the shared pydantic models; every response +this suite reads is modelled locally so the file is self-contained and no +untyped dict crosses the boundary. +""" + +from __future__ import annotations + +import time +from collections.abc import Callable + +import pytest +from pydantic import BaseModel, ConfigDict, RootModel + +from e2e_config import unique_marker +from e2e_http import NoBody, unwrap +from lifecycle import ResourceManager +from management_client import ManagementClient +from models import KeyGenerateBody, LiteLLMParamsBody, ModelInfoBody, ModelNewBody +from proxy_client import ProxyClient + +pytestmark = pytest.mark.e2e + +_MODEL_PERMISSION_DENIED_MARKER = "does not have permission to make this model call" +_DUMMY_MODEL = "openai/gpt-5.5" +_DUMMY_API_KEY = "e2e-dummy-key" + + +def _poll[T](proxy: ProxyClient, attempt: Callable[[], T | None], failure: str) -> T: + deadline = time.monotonic() + proxy.poll_timeout + while time.monotonic() < deadline: + found = attempt() + if found is not None: + return found + time.sleep(proxy.poll_interval) + pytest.fail(failure) + + +# ---------- tag route models / helpers ---------- + + +class TagCreateBody(BaseModel): + name: str + description: str | None = None + + +class TagDeleteBody(BaseModel): + name: str + + +class TagEntry(BaseModel): + name: str + description: str | None = None + + +class TagCatalog(RootModel[list[TagEntry]]): + """GET /tag/list answers with a bare array of tag configs, not an object + wrapping them; read the rows off .root.""" + + +def _tag_list(client: ManagementClient) -> tuple[TagEntry, ...]: + return tuple( + unwrap( + client.proxy.transport.get( + "/tag/list", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=TagCatalog, + ) + ).root + ) + + +def _create_tag(client: ManagementClient, body: TagCreateBody) -> None: + _ = unwrap( + client.proxy.transport.post( + "/tag/new", + headers=client.proxy.transport.master, + json=body, + response_type=NoBody, + ) + ) + + +def _delete_tag(client: ManagementClient, name: str) -> None: + """Best-effort delete for teardown: a repeat /tag/delete on an already-deleted + tag is a no-op the warn-only teardown absorbs.""" + _ = client.proxy.transport.post( + "/tag/delete", + headers=client.proxy.transport.master, + json=TagDeleteBody(name=name), + response_type=NoBody, + ) + + +def _delete_tag_strict(client: ManagementClient, name: str) -> None: + """Strict delete for the act phase: a failed /tag/delete is a hard failure.""" + _ = unwrap( + client.proxy.transport.post( + "/tag/delete", + headers=client.proxy.transport.master, + json=TagDeleteBody(name=name), + response_type=NoBody, + ) + ) + + +# ---------- access group route models / helpers ---------- + + +class AccessGroupNewBody(BaseModel): + access_group: str + model_names: list[str] + + +class AccessGroupNewResponse(BaseModel): + access_group: str + models_updated: int + + +class AccessGroupInfoResponse(BaseModel): + access_group: str + model_names: list[str] + deployment_count: int + + +def _create_access_group(client: ManagementClient, body: AccessGroupNewBody) -> AccessGroupNewResponse: + return unwrap( + client.proxy.transport.post( + "/access_group/new", + headers=client.proxy.transport.master, + json=body, + response_type=AccessGroupNewResponse, + ) + ) + + +def _access_group_info(client: ManagementClient, access_group: str) -> AccessGroupInfoResponse | None: + result = client.proxy.transport.get( + f"/access_group/{access_group}/info", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=AccessGroupInfoResponse, + ) + return unwrap(result) if result.kind == "success" else None + + +def _delete_access_group(client: ManagementClient, access_group: str) -> None: + """Best-effort delete for teardown; deleting the model behind it removes the + access group too, so a repeat delete is a no-op the teardown absorbs.""" + _ = client.proxy.transport.delete( + f"/access_group/{access_group}/delete", + headers=client.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) + + +def _create_db_model(client: ManagementClient, resources: ResourceManager, model_name: str) -> str: + model_id = client.proxy.create_model( + model_name, LiteLLMParamsBody(model=_DUMMY_MODEL, api_key=_DUMMY_API_KEY) + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + return model_id + + +# ---------- model block route models / helpers ---------- + + +class ModelBlockBody(BaseModel): + model_config = ConfigDict(protected_namespaces=()) + model_id: str + + +class ModelInfoBlockDetail(BaseModel): + id: str | None = None + blocked: bool | None = None + + +class ModelInfoBlockEntry(BaseModel): + model_config = ConfigDict(protected_namespaces=()) + model_name: str + model_info: ModelInfoBlockDetail = ModelInfoBlockDetail() + + +class ModelInfoCatalog(BaseModel): + data: list[ModelInfoBlockEntry] = [] + + +def _model_blocked_flag(client: ManagementClient, model_id: str) -> bool | None: + catalog = unwrap( + client.proxy.transport.get( + "/model/info", + headers=client.proxy.transport.master, + params=NoBody(), + response_type=ModelInfoCatalog, + ) + ) + entry = next((row for row in catalog.data if row.model_info.id == model_id), None) + return entry.model_info.blocked if entry is not None else None + + +class TestModelRoutes: + @pytest.mark.covers("mgmt.model.add.admin_only") + def test_non_admin_key_cannot_add_global_model( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + key = client.proxy.generate_key(KeyGenerateBody(models=[])) + resources.defer(lambda: client.proxy.delete_key(key)) + + model_name = f"e2e-mgmt-model-forbidden-{unique_marker()}" + outcome = client.proxy.transport.send( + "/model/new", + headers=client.proxy.transport.bearer(key), + json=ModelNewBody( + model_name=model_name, + litellm_params=LiteLLMParamsBody(model=_DUMMY_MODEL, api_key=_DUMMY_API_KEY), + model_info=ModelInfoBody(), + ), + ) + + assert outcome.status_code == 403, ( + f"non-admin key adding a global model (no team_id) must be denied 403, got " + f"{outcome.status_code}: {outcome.body[:300]}" + ) + assert _MODEL_PERMISSION_DENIED_MARKER in outcome.body, ( + f"403 body must be the model-permission denial, got: {outcome.body[:300]}" + ) + + cataloged = [entry.model_name for entry in client.proxy.model_info()] + assert model_name not in cataloged, ( + f"{model_name!r} was registered in /model/info despite the 403; the admin-only " + f"guard did not block the write" + ) + + @pytest.mark.covers("mgmt.model.block.persists") + def test_block_then_unblock_persists_to_model_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """The blocked flag's persistence is read back from /model/info, not from the + /model/block response: that route currently returns a non-2xx serialization + envelope even though the DB write lands, so the /model/info read-back is the + authoritative persistence contract and keeps this test valid once the + response shape is fixed.""" + model_name = f"e2e-mgmt-model-block-{unique_marker()}" + model_id = _create_db_model(client, resources, model_name) + + assert _model_blocked_flag(client, model_id) is not True, ( + f"{model_name!r} already reports blocked in /model/info before /model/block ran" + ) + + _ = client.proxy.transport.send( + "/model/block", + headers=client.proxy.transport.master, + json=ModelBlockBody(model_id=model_id), + ) + _ = _poll( + client.proxy, + lambda: True if _model_blocked_flag(client, model_id) is True else None, + f"/model/info never reported {model_name!r} blocked after /model/block", + ) + + _ = client.proxy.transport.send( + "/model/unblock", + headers=client.proxy.transport.master, + json=ModelBlockBody(model_id=model_id), + ) + _ = _poll( + client.proxy, + lambda: True if _model_blocked_flag(client, model_id) is not True else None, + f"/model/info never cleared blocked for {model_name!r} after /model/unblock", + ) + + +class TestTagRoutes: + @pytest.mark.covers("mgmt.tag.list.happy_path") + def test_tag_list_reports_created_tag(self, client: ManagementClient, resources: ResourceManager) -> None: + name = f"e2e-mgmt-tag-{unique_marker()}" + description = "coverage: tag inventory" + assert all(entry.name != name for entry in _tag_list(client)), ( + f"tag {name!r} was already listed by /tag/list before /tag/new created it" + ) + + _create_tag(client, TagCreateBody(name=name, description=description)) + resources.defer(lambda: _delete_tag(client, name)) + + entry = _poll( + client.proxy, + lambda: next((entry for entry in _tag_list(client) if entry.name == name), None), + f"/tag/list never listed {name!r} after /tag/new", + ) + assert entry.description == description, ( + f"/tag/list reports description {entry.description!r} for {name!r}, configured {description!r}" + ) + + @pytest.mark.covers("mgmt.tag.delete.persists") + def test_tag_delete_removes_from_list(self, client: ManagementClient, resources: ResourceManager) -> None: + """The teardown's deferred delete fires again on the already-deleted tag by + design: it is the safety net if this test fails before the in-body delete, + and a repeat /tag/delete is a warn-only no-op the teardown absorbs.""" + name = f"e2e-mgmt-tag-{unique_marker()}" + _create_tag(client, TagCreateBody(name=name)) + resources.defer(lambda: _delete_tag(client, name)) + + _ = _poll( + client.proxy, + lambda: True if any(entry.name == name for entry in _tag_list(client)) else None, + f"/tag/list never listed {name!r} after /tag/new; cannot prove deletion removes it", + ) + + _delete_tag_strict(client, name) + + _ = _poll( + client.proxy, + lambda: True if all(entry.name != name for entry in _tag_list(client)) else None, + f"{name!r} still present in /tag/list after /tag/delete at the deadline", + ) + + +class TestModelAccessGroupRoutes: + @pytest.mark.covers("mgmt.access_group.new.happy_path") + def test_new_access_group_tags_the_deployment( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + model_name = f"e2e-mgmt-agmodel-{unique_marker()}" + _ = _create_db_model(client, resources, model_name) + + access_group = f"e2e-mgmt-ag-{unique_marker()}" + created = _create_access_group( + client, AccessGroupNewBody(access_group=access_group, model_names=[model_name]) + ) + resources.defer(lambda: _delete_access_group(client, access_group)) + + assert created.access_group == access_group, ( + f"/access_group/new echoed access_group {created.access_group!r}, requested {access_group!r}" + ) + assert created.models_updated >= 1, ( + f"/access_group/new tagged {created.models_updated} deployments for {model_name!r}, expected >= 1" + ) + + info = _poll( + client.proxy, + lambda: _access_group_info(client, access_group), + f"/access_group/{access_group}/info never resolved the group created by /access_group/new", + ) + assert model_name in info.model_names, ( + f"the group created by /access_group/new does not list {model_name!r} on read-back; " + f"/access_group/info reports members {info.model_names}" + ) + + @pytest.mark.covers("mgmt.access_group.info.happy_path") + def test_access_group_info_reports_membership( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + model_name = f"e2e-mgmt-agmodel-{unique_marker()}" + _ = _create_db_model(client, resources, model_name) + + access_group = f"e2e-mgmt-ag-{unique_marker()}" + _ = _create_access_group( + client, AccessGroupNewBody(access_group=access_group, model_names=[model_name]) + ) + resources.defer(lambda: _delete_access_group(client, access_group)) + + info = _poll( + client.proxy, + lambda: _access_group_info(client, access_group), + f"/access_group/{access_group}/info never resolved the created access group", + ) + assert info.access_group == access_group, ( + f"/access_group/info reports access_group {info.access_group!r}, created {access_group!r}" + ) + assert model_name in info.model_names, ( + f"/access_group/info reports members {info.model_names}, expected to include {model_name!r}" + ) + assert info.deployment_count >= 1, ( + f"/access_group/info reports deployment_count {info.deployment_count}, expected >= 1" + ) diff --git a/tests/e2e/management/test_team_management_e2e.py b/tests/e2e/management/test_team_management_e2e.py new file mode 100644 index 00000000000..108aeaad21b --- /dev/null +++ b/tests/e2e/management/test_team_management_e2e.py @@ -0,0 +1,303 @@ +"""Live e2e: the /team/* management routes' block, membership, and admin-only +contract. + +Each test creates its team/user/key resources under unique names (deleted on +teardown) and asserts both halves of the contract: the recorded state (the info +route reflects the write) and the enforced behavior (a non-admin key is refused). +Team writes reach the read path once their db/cache entry propagates, so the +read-backs poll to a deadline instead of asserting once. + +Everything the shared harness does not already model lives here: the local +request/response models for /team/block, /team/member_update, and the +/team/info fields (blocked flag and per-member budget) these tests assert on. +""" + +from __future__ import annotations + +import time +from collections.abc import Callable +from typing import Literal + +import pytest +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import NoBody, StreamingResponse, unwrap +from lifecycle import ResourceManager +from management_client import ManagementClient +from models import ( + KeyGenerateBody, + TeamInfoParams, + TeamMemberAddBody, + TeamMemberDeleteBody, + TeamMemberEntry, + TeamNewBody, + UserNewBody, +) + +pytestmark = pytest.mark.e2e + +TeamRole = Literal["admin", "user"] + + +class TeamBlockBody(BaseModel): + team_id: str + + +class MemberUpdateBody(BaseModel): + team_id: str + user_id: str + role: TeamRole | None = None + max_budget_in_team: float | None = None + + +class MemberRoleEntry(BaseModel): + user_id: str | None = None + user_email: str | None = None + role: TeamRole + + +class MemberBudgetTable(BaseModel): + max_budget: float | None = None + + +class TeamMembership(BaseModel): + user_id: str + litellm_budget_table: MemberBudgetTable | None = None + + +class TeamInfoData(BaseModel): + team_alias: str | None = None + models: list[str] = [] + blocked: bool | None = None + members_with_roles: list[MemberRoleEntry] = [] + + +class TeamInfoRead(BaseModel): + team_id: str + team_info: TeamInfoData + team_memberships: list[TeamMembership] = [] + + +def _poll[T](client: ManagementClient, attempt: Callable[[], T | None], failure: str) -> T: + deadline = time.monotonic() + client.proxy.poll_timeout + while time.monotonic() < deadline: + found = attempt() + if found is not None: + return found + time.sleep(client.proxy.poll_interval) + pytest.fail(failure) + + +def _create_team(client: ManagementClient, resources: ResourceManager, alias: str, models: list[str]) -> str: + team_id = client.create_team(TeamNewBody(team_alias=alias, models=models)) + resources.defer(lambda: client.delete_team(team_id)) + return team_id + + +def _create_user(client: ManagementClient, resources: ResourceManager, email: str) -> str: + user_id = client.create_user(UserNewBody(user_email=email, user_role="internal_user")) + resources.defer(lambda: client.delete_user(user_id)) + return user_id + + +def _generate_key(client: ManagementClient, resources: ResourceManager, body: KeyGenerateBody) -> str: + key = client.proxy.generate_key(body) + resources.defer(lambda: client.proxy.delete_key(key)) + return key + + +def _read_team(client: ManagementClient, team_id: str) -> TeamInfoRead: + return unwrap( + client.proxy.transport.get( + "/team/info", + headers=client.proxy.transport.master, + params=TeamInfoParams(team_id=team_id), + response_type=TeamInfoRead, + ) + ) + + +def _set_blocked(client: ManagementClient, team_id: str, *, blocked: bool) -> None: + _ = unwrap( + client.proxy.transport.post( + "/team/unblock" if not blocked else "/team/block", + headers=client.proxy.transport.master, + json=TeamBlockBody(team_id=team_id), + response_type=NoBody, + ) + ) + + +def _member_update(client: ManagementClient, body: MemberUpdateBody) -> None: + _ = unwrap( + client.proxy.transport.post( + "/team/member_update", + headers=client.proxy.transport.master, + json=body, + response_type=NoBody, + ) + ) + + +def _member_role(info: TeamInfoRead, user_id: str) -> TeamRole | None: + return next((m.role for m in info.team_info.members_with_roles if m.user_id == user_id), None) + + +def _member_max_budget(info: TeamInfoRead, user_id: str) -> float | None: + membership = next((tm for tm in info.team_memberships if tm.user_id == user_id), None) + if membership is None or membership.litellm_budget_table is None: + return None + return membership.litellm_budget_table.max_budget + + +def _member_add_status(client: ManagementClient, key: str, team_id: str, user_id: str) -> StreamingResponse: + return client.proxy.transport.send( + "/team/member_add", + headers=client.proxy.transport.bearer(key), + json=TeamMemberAddBody(team_id=team_id, member=TeamMemberEntry(role="user", user_id=user_id)), + ) + + +def _member_delete_status(client: ManagementClient, key: str, team_id: str, user_id: str) -> StreamingResponse: + return client.proxy.transport.send( + "/team/member_delete", + headers=client.proxy.transport.bearer(key), + json=TeamMemberDeleteBody(team_id=team_id, user_id=user_id), + ) + + +class TestTeamManagementRoutes: + @pytest.mark.covers("mgmt.team.info.happy_path") + def test_info_returns_created_team_fields( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + alias = f"e2e-team-info-{unique_marker()}" + team_id = _create_team(client, resources, alias, ["gemini-2.5-flash"]) + + info = _read_team(client, team_id) + assert info.team_id == team_id, f"/team/info echoed team_id {info.team_id!r}, requested {team_id!r}" + assert info.team_info.team_alias == alias, ( + f"/team/info reports team_alias {info.team_info.team_alias!r}, configured {alias!r}" + ) + assert info.team_info.models == ["gemini-2.5-flash"], ( + f"/team/info reports models {info.team_info.models}, configured ['gemini-2.5-flash']" + ) + + @pytest.mark.covers("mgmt.team.block.persists") + def test_block_then_unblock_persists_to_team_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + team_id = _create_team(client, resources, f"e2e-team-block-{unique_marker()}", ["gemini-2.5-flash"]) + assert not _read_team(client, team_id).team_info.blocked, "/team/info reports the team blocked before /team/block" + + _set_blocked(client, team_id, blocked=True) + _ = _poll( + client, + lambda: True if _read_team(client, team_id).team_info.blocked else None, + "/team/info never reflected blocked=True after /team/block", + ) + + _set_blocked(client, team_id, blocked=False) + _ = _poll( + client, + lambda: True if _read_team(client, team_id).team_info.blocked is False else None, + "/team/info never reflected blocked=False after /team/unblock", + ) + + @pytest.mark.covers("mgmt.team.member_update.persists") + def test_member_update_persists_role_and_budget( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + user_id = _create_user(client, resources, f"e2e-team-mu-{unique_marker()}@example.com") + team_id = _create_team(client, resources, f"e2e-team-mu-{unique_marker()}", ["gemini-2.5-flash"]) + client.add_team_member(team_id, user_id) + assert _member_role(_read_team(client, team_id), user_id) == "user", ( + f"member {user_id} should start as role 'user' after /team/member_add" + ) + + budget = 4242.0 + _member_update(client, MemberUpdateBody(team_id=team_id, user_id=user_id, role="admin", max_budget_in_team=budget)) + + def updated() -> bool | None: + info = _read_team(client, team_id) + return True if _member_role(info, user_id) == "admin" and _member_max_budget(info, user_id) == budget else None + + _ = _poll( + client, + updated, + f"/team/info never reflected role=admin and max_budget={budget} for {user_id} after /team/member_update", + ) + + @pytest.mark.covers("mgmt.team.member_delete.persists") + def test_member_delete_persists_to_team_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + user_id = _create_user(client, resources, f"e2e-team-md-{unique_marker()}@example.com") + team_id = _create_team(client, resources, f"e2e-team-md-{unique_marker()}", ["gemini-2.5-flash"]) + client.add_team_member(team_id, user_id) + assert _member_role(_read_team(client, team_id), user_id) == "user", ( + f"/team/info does not list {user_id} as a member after /team/member_add" + ) + + client.delete_team_member(team_id, user_id) + _ = _poll( + client, + lambda: True if _member_role(_read_team(client, team_id), user_id) is None else None, + f"/team/info still lists {user_id} after /team/member_delete", + ) + + @pytest.mark.covers("mgmt.team.new.admin_only") + def test_new_is_denied_to_non_admin_keys( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + no_role_key = _generate_key(client, resources, KeyGenerateBody(models=[])) + internal_user_id = _create_user(client, resources, f"e2e-team-adm-{unique_marker()}@example.com") + internal_user_key = _generate_key(client, resources, KeyGenerateBody(user_id=internal_user_id)) + + for key, label in ((no_role_key, "role=None"), (internal_user_key, "internal_user")): + outcome = client.team_new_status(key, TeamNewBody(team_alias=f"e2e-team-adm-{unique_marker()}")) + assert outcome.status_code in (401, 403), ( + f"/team/new by a {label} key must be denied 401/403, got {outcome.status_code}: {outcome.body[:300]}" + ) + + @pytest.mark.covers("mgmt.team.member_add.member_forbidden") + def test_member_add_forbidden_to_plain_member( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + _member_id, other_id, member_key, team_id = self._team_with_member_key(client, resources) + + outcome = _member_add_status(client, member_key, team_id, other_id) + assert outcome.status_code == 403, ( + f"/team/member_add by a plain team member must be 403, got {outcome.status_code}: {outcome.body[:300]}" + ) + assert "not allowed" in outcome.body.lower(), ( + f"403 body should say the call is not allowed, got: {outcome.body[:300]}" + ) + + @pytest.mark.covers("mgmt.team.member_delete.member_forbidden") + def test_member_delete_forbidden_to_plain_member( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + member_id, _other_id, member_key, team_id = self._team_with_member_key(client, resources) + + outcome = _member_delete_status(client, member_key, team_id, member_id) + assert outcome.status_code == 403, ( + f"/team/member_delete by a plain team member must be 403, got {outcome.status_code}: {outcome.body[:300]}" + ) + assert "not allowed" in outcome.body.lower(), ( + f"403 body should say the call is not allowed, got: {outcome.body[:300]}" + ) + + @staticmethod + def _team_with_member_key( + client: ManagementClient, resources: ResourceManager + ) -> tuple[str, str, str, str]: + """A team with a plain member (role user) whose key is scoped to that + user + team, plus a second user id the member could try to add.""" + member_id = _create_user(client, resources, f"e2e-team-fb-{unique_marker()}@example.com") + other_id = _create_user(client, resources, f"e2e-team-fb-{unique_marker()}@example.com") + team_id = _create_team(client, resources, f"e2e-team-fb-{unique_marker()}", ["gemini-2.5-flash"]) + client.add_team_member(team_id, member_id) + member_key = _generate_key(client, resources, KeyGenerateBody(user_id=member_id, team_id=team_id)) + return member_id, other_id, member_key, team_id diff --git a/tests/e2e/mcp/mcp_client.py b/tests/e2e/mcp/mcp_client.py index f68fdf63b3f..b0aa4c68e3a 100644 --- a/tests/e2e/mcp/mcp_client.py +++ b/tests/e2e/mcp/mcp_client.py @@ -85,6 +85,37 @@ class McpToolsListResponse(BaseModel): return None +class BlockedWordSpec(BaseModel): + keyword: str + action: str = "BLOCK" + + +class ContentFilterMcpParams(BaseModel): + """litellm_content_filter params scoped to the MCP tool-call hook. mode is + pre_mcp_call because a pre_call config silently no-ops on the tools/call path + (the event type is rewritten to pre_mcp_call for call_mcp_tool), and default_on + is required there because per-key/request guardrail selection is dropped from + the synthetic MCP request the hook sees.""" + + guardrail: str = "litellm_content_filter" + mode: str = "pre_mcp_call" + default_on: bool = True + blocked_words: list[BlockedWordSpec] + + +class GuardrailSpecBody(BaseModel): + guardrail_name: str + litellm_params: ContentFilterMcpParams + + +class GuardrailCreateBody(BaseModel): + guardrail: GuardrailSpecBody + + +class GuardrailCreateResponse(BaseModel): + guardrail_id: str + + class McpCallToolBody(BaseModel): name: str arguments: dict[str, McpToolArg] @@ -186,6 +217,35 @@ class McpClient: response_type=McpToolsListResponse, ) + def register_mcp_content_filter(self, *, name: str, blocked_keyword: str) -> str: + """Register a default-on content-filter guardrail that runs on the MCP + tool-call hook (pre_mcp_call) and blocks a single keyword. The keyword is + unique per test, so default_on only ever intercepts this test's own + banned tool call on the shared proxy.""" + return unwrap( + self.proxy.transport.post( + "/guardrails", + headers=self.proxy.transport.master, + json=GuardrailCreateBody( + guardrail=GuardrailSpecBody( + guardrail_name=name, + litellm_params=ContentFilterMcpParams( + blocked_words=[BlockedWordSpec(keyword=blocked_keyword)], + ), + ) + ), + response_type=GuardrailCreateResponse, + ) + ).guardrail_id + + def delete_guardrail(self, guardrail_id: str) -> None: + _ = self.proxy.transport.delete( + f"/guardrails/{guardrail_id}", + headers=self.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) + def call_tool( self, key: str, diff --git a/tests/e2e/mcp/test_mcp_guardrail_e2e.py b/tests/e2e/mcp/test_mcp_guardrail_e2e.py new file mode 100644 index 00000000000..63239444454 --- /dev/null +++ b/tests/e2e/mcp/test_mcp_guardrail_e2e.py @@ -0,0 +1,146 @@ +"""Live e2e: a guardrail on the MCP tool-call path blocks banned content in the +tool arguments before the call reaches the upstream MCP server. + +A general litellm_content_filter guardrail is configured with mode=pre_mcp_call +(the event type the proxy rewrites pre_call to for a call_mcp_tool) and default_on +(per-key/request guardrail selection is dropped from the synthetic MCP request the +hook sees, so default_on is how it attaches to tools/call). The banned keyword is +unique per run, so default_on only ever intercepts this test's own banned call. + +Against the real Datadog MCP server, calling search_datadog_logs with the banned +keyword in the query is blocked with HTTP 400 attributed to the pre_mcp_call hook, +and the tool never runs; the same guardrail lets a clean query through to Datadog. +This is the enforced half (the block) plus the pass-through half in one spec. +""" + +from __future__ import annotations + +import time +from collections.abc import Callable + +import pytest + +from datadog_mcp import SEARCH_LOGS_TOOL, assert_dd_mcp_creds, register_datadog_mcp +from e2e_config import DD_SEARCH_FROM, unique_marker +from e2e_http import Result, Success, UnknownApiError, unwrap +from lifecycle import ResourceManager +from mcp_client import McpCallToolResponse, McpClient, McpToolArguments + +pytestmark = pytest.mark.e2e + +# Stage runs several data-plane pods behind the shared key, and each picks up a +# newly registered guardrail only on its next periodic DB sync (~30s in +# proxy_server.py). Every pod is guaranteed to have refreshed only once a full sync +# interval has elapsed since the create; before then a banned call routed to a +# lagging pod passes through as legitimate in-flight propagation, not a leak. +GUARDRAIL_FULL_SYNC_SECONDS = 40.0 +POST_SYNC_VERIFICATION_CALLS = 4 + + +def _poll_until_blocked( + search: Callable[[str], Result[McpCallToolResponse]], banned_keyword: str, client: McpClient +) -> Result[McpCallToolResponse]: + """Retry a banned tool call until the guardrail blocks it (400) or the deadline + passes, returning the last result. Absorbs the control-plane -> data-plane + guardrail-sync delay so the check waits for enforcement instead of racing it.""" + deadline = time.monotonic() + client.proxy.poll_timeout + last: Result[McpCallToolResponse] = search(f"tell me about {banned_keyword}") + while time.monotonic() < deadline: + if isinstance(last, UnknownApiError) and last.status_code == 400: + return last + time.sleep(client.proxy.poll_interval) + last = search(f"tell me about {banned_keyword}") + return last + + +class TestMcpToolCallGuardrail: + @pytest.mark.covers( + "guardrail.litellm_content_filter.pre_mcp_call.blocks", + exercised_on=["mcp_operations"], + ) + def test_content_filter_blocks_banned_keyword_in_tool_args( + self, client: McpClient, resources: ResourceManager + ) -> None: + assert_dd_mcp_creds() + marker = unique_marker() + banned_keyword = f"e2eblocked{marker}" + + guardrail_id = client.register_mcp_content_filter( + name=f"e2e-mcp-cf-{marker}", blocked_keyword=banned_keyword + ) + guardrail_created_at = time.monotonic() + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + server_id = register_datadog_mcp(client, resources) + key = client.generate_key(user_id=f"e2e-mcp-guard-{marker}", mcp_servers=[server_id]) + resources.defer(lambda: client.proxy.delete_key(key)) + + tools = unwrap(client.list_tools(key)) + tool_name = tools.tool_name_containing(server_id, SEARCH_LOGS_TOOL) + assert tool_name is not None, ( + f"granted key never saw {SEARCH_LOGS_TOOL} on server {server_id}; " + f"tools={tools.tool_names_for_server(server_id)}" + ) + + def search(query: str) -> Result[McpCallToolResponse]: + arguments: McpToolArguments = { + "query": query, + "from": DD_SEARCH_FROM, + "to": "now", + "max_tokens": 500, + "telemetry": {"intent": "e2e mcp guardrail check"}, + } + return client.call_tool(key, server_id=server_id, name=tool_name, arguments=arguments) + + # Registering the guardrail is a control-plane write; the data-plane worker + # that serves tools/call picks it up on its next guardrail sync, so an + # immediate call can race the propagation and slip through. Poll the banned + # call to the deadline and require a block, so the check proves enforcement + # rather than catching a pre-sync pass-through. The keyword is unique per + # run, so this only ever intercepts this test's own call. + blocked = _poll_until_blocked(search, banned_keyword, client) + match blocked: + case UnknownApiError(status_code=400, body=body): + assert banned_keyword in body or "content blocked" in body.lower(), ( + f"the block must name the content-filter reason, got: {body[:300]}" + ) + assert "pre_mcp_call" in body, ( + f"the block must be attributed to the MCP tool-call hook (pre_mcp_call), got: {body[:300]}" + ) + case _: + pytest.fail( + "content_filter never blocked the banned keyword on the MCP tool call within " + f"{client.proxy.poll_timeout}s (guardrail sync to the data plane never landed); " + f"last result: {blocked}" + ) + + # The block above only proves the one pod that served it has synced; another + # pod could still lack the guardrail and let the banned call reach Datadog. + # Wait out the full sync interval from the create so every pod has refreshed + # from the DB, then require the banned call to stay blocked across several + # attempts. A pass-through now is a genuine partial-propagation leak, not a + # race. Client load balancing still can't guarantee every pod is hit, so this + # samples several worker selections rather than proving all pods synced. + sync_remaining = guardrail_created_at + GUARDRAIL_FULL_SYNC_SECONDS - time.monotonic() + if sync_remaining > 0: + time.sleep(sync_remaining) + for attempt in range(1, POST_SYNC_VERIFICATION_CALLS + 1): + reblocked = search(f"still about {banned_keyword} #{attempt}") + assert isinstance(reblocked, UnknownApiError) and reblocked.status_code == 400, ( + "after the guardrail sync interval every data-plane pod must block the banned " + f"keyword, but attempt {attempt} of {POST_SYNC_VERIFICATION_CALLS} was allowed " + f"through (a pod still lacks the guardrail): {reblocked}" + ) + if attempt < POST_SYNC_VERIFICATION_CALLS: + time.sleep(client.proxy.poll_interval) + + allowed = search(f"e2e-clean-{marker}") + match allowed: + case Success(data=result): + assert result.is_error is not True, ( + f"a clean MCP tool call must reach the server and not error, got: {result}" + ) + case _: + pytest.fail( + f"a clean MCP tool call must pass the guardrail and reach the server; got {allowed}" + ) diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 28cc7984598..b3ea9346180 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -115,6 +115,18 @@ class KeyInfoResponse(BaseModel): # ---------- customers ---------- +class CustomerNewBody(BaseModel): + user_id: str + + +class CustomerResponse(BaseModel): + user_id: str | None = None + + +class CustomerInfoParams(BaseModel): + end_user_id: str + + class CustomerDeleteBody(BaseModel): user_ids: list[str] @@ -126,9 +138,26 @@ class ChatMetadata(BaseModel): tags: list[str] | None = None +class ImageUrl(BaseModel): + url: str + + +class TextContentPart(BaseModel): + type: str = "text" + text: str + + +class ImageContentPart(BaseModel): + type: str = "image_url" + image_url: ImageUrl + + +ContentPart = TextContentPart | ImageContentPart + + class ChatMessage(BaseModel): role: str - content: str + content: str | list[ContentPart] class CacheControl(BaseModel): @@ -180,6 +209,7 @@ class ChatBody(BaseModel): tools: list[ChatTool] | None = None tool_choice: str | None = None guardrails: list[str] | None = None + response_format: dict[str, object] | None = None class RouterSettingsOverride(BaseModel): @@ -203,9 +233,19 @@ class ReliabilityChatBody(ChatBody): router_settings_override: RouterSettingsOverride | None = None +class ToolCallFunction(BaseModel): + name: str | None = None + arguments: str | None = None + + +class ToolCall(BaseModel): + function: ToolCallFunction = ToolCallFunction() + + class OutMessage(BaseModel): content: str | None = None reasoning_content: str | None = None + tool_calls: list[ToolCall] | None = None class ChatChoice(BaseModel): @@ -216,6 +256,10 @@ class PromptTokensDetails(BaseModel): cached_tokens: int | None = None +class CompletionTokensDetails(BaseModel): + reasoning_tokens: int | None = None + + class Usage(BaseModel): prompt_tokens: int | None = None completion_tokens: int | None = None @@ -223,6 +267,7 @@ class Usage(BaseModel): cache_read_input_tokens: int | None = None cache_creation_input_tokens: int | None = None prompt_tokens_details: PromptTokensDetails | None = None + completion_tokens_details: CompletionTokensDetails | None = None class ChatResponse(BaseModel): @@ -286,6 +331,7 @@ class CountTokensBody(BaseModel): class AnthropicContentBlock(BaseModel): type: str | None = None + text: str | None = None class AnthropicMessagesResponse(BaseModel): @@ -817,3 +863,23 @@ class TagListResponse(RootModel[list[TagListEntry]]): """GET /tag/list answers with a bare array of tag configs (the stored tags plus any dynamically-seen spend tags), not an object wrapping them. Read the rows off .root.""" + + +# ---------- health / lifecycle ---------- + + +class ReadinessResponse(BaseModel): + """GET /health/readiness (public probe). The low-detail payload a load + balancer sees: `status` plus the resolved DB state (`connected`, + `disconnected`, or `Not connected`).""" + + status: str + db: str | None = None + + +class ReadinessDetailsResponse(ReadinessResponse): + """GET /health/readiness/details (authenticated). Extends the public payload + with the diagnostics only an authenticated caller may read.""" + + litellm_version: str | None = None + success_callbacks: list[str] = [] diff --git a/tests/e2e/logging/otel_client.py b/tests/e2e/otel_client.py similarity index 100% rename from tests/e2e/logging/otel_client.py rename to tests/e2e/otel_client.py diff --git a/tests/e2e/other/conftest.py b/tests/e2e/other/conftest.py new file mode 100644 index 00000000000..9141b6e364e --- /dev/null +++ b/tests/e2e/other/conftest.py @@ -0,0 +1,18 @@ +"""`other` suite's `client` fixture. + +Lifecycle (resources/scoped_key), proxy liveness gate, and the e2e/covers +markers all live in the parent tests/e2e/conftest.py. OtherClient holds the +shared ProxyClient so anything these tests create tears down through it. +""" + +from __future__ import annotations + +import pytest + +from other_client import OtherClient, build_client +from proxy_client import ProxyClient + + +@pytest.fixture(scope="session") +def client(proxy: ProxyClient) -> OtherClient: + return build_client(proxy) diff --git a/tests/e2e/other/other_client.py b/tests/e2e/other/other_client.py new file mode 100644 index 00000000000..1aa83ac42c7 --- /dev/null +++ b/tests/e2e/other/other_client.py @@ -0,0 +1,73 @@ +"""Client for the `other` holding-pen suite: the auth gate (master key vs an +invalid key on an admin route) and the process-lifecycle health probes +(liveness, public readiness, authenticated readiness diagnostics). + +Holds the shared ProxyClient so `resources` / `scoped_key` still clean up, and +adds only the routes these behaviors need. The health probes deliberately send +no auth header (public routes), so they go through the transport with an empty +headers model rather than a bearer. +""" + +from __future__ import annotations + +from dataclasses import dataclass + +from e2e_http import NoBody, ProbeResult, Result +from models import ( + ReadinessDetailsResponse, + ReadinessResponse, + UserListParams, + UserListResponse, +) +from proxy_client import ProxyClient + + +@dataclass(frozen=True, slots=True) +class OtherClient: + proxy: ProxyClient + + def liveness(self) -> ProbeResult: + """GET /health/liveliness. Unauthenticated; the probe returns status + + raw body so the test can assert the worker reports itself alive.""" + return self.proxy.transport.probe("/health/liveliness", params=NoBody()) + + def readiness_public(self) -> Result[ReadinessResponse]: + """GET /health/readiness with no credential at all, proving the probe is + safe to expose to an unauthenticated load balancer.""" + return self.proxy.transport.get( + "/health/readiness", + headers=NoBody(), + params=NoBody(), + response_type=ReadinessResponse, + ) + + def readiness_details(self, key: str) -> Result[ReadinessDetailsResponse]: + return self.proxy.transport.get( + "/health/readiness/details", + headers=self.proxy.transport.bearer(key), + params=NoBody(), + response_type=ReadinessDetailsResponse, + ) + + def readiness_details_unauthenticated(self) -> Result[ReadinessDetailsResponse]: + return self.proxy.transport.get( + "/health/readiness/details", + headers=NoBody(), + params=NoBody(), + response_type=ReadinessDetailsResponse, + ) + + def list_users_as(self, key: str) -> Result[UserListResponse]: + """GET /user/list under `key`. Admin-only, so it doubles as the master + key's authorization proof: the master key (proxy admin) reads it, a + non-matching key is rejected before it ever reaches the handler.""" + return self.proxy.transport.get( + "/user/list", + headers=self.proxy.transport.bearer(key), + params=UserListParams(user_ids="e2e-test-user"), + response_type=UserListResponse, + ) + + +def build_client(proxy: ProxyClient) -> OtherClient: + return OtherClient(proxy=proxy) diff --git a/tests/e2e/other/test_health_lifecycle_e2e.py b/tests/e2e/other/test_health_lifecycle_e2e.py new file mode 100644 index 00000000000..2551352e8fa --- /dev/null +++ b/tests/e2e/other/test_health_lifecycle_e2e.py @@ -0,0 +1,65 @@ +"""Live e2e: the process-lifecycle probes Kubernetes and load balancers depend on. + +Liveness and public readiness must answer without a credential (a load balancer +has none), and public readiness must distinguish a healthy worker from one whose +DB is unreachable by reporting the resolved DB state. The detailed readiness +route, by contrast, is authenticated: it exposes diagnostics (version, callbacks, +DB) and must reject an anonymous caller. The suite runs against a proxy configured +with a real database, so a healthy readiness payload reports the DB as connected; +a regression that stopped checking the DB, or dropped the public exposure, fails +here. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import MASTER_KEY +from e2e_http import UnauthorizedError, unwrap +from other_client import OtherClient + +pytestmark = pytest.mark.e2e + + +class TestHealthLifecycle: + @pytest.mark.covers("other.lifecycle.liveness.ping") + def test_liveness_reports_alive_without_auth(self, client: OtherClient) -> None: + probe = client.liveness() + assert probe.status_code == 200, ( + f"liveness must answer 200 for an unauthenticated probe, got " + f"{probe.status_code}: {probe.body[:200]}" + ) + assert "alive" in probe.body.lower(), ( + f"liveness body must confirm the worker is alive, got {probe.body[:200]}" + ) + + @pytest.mark.covers("other.lifecycle.readiness.public_probe") + def test_readiness_is_reachable_without_credentials(self, client: OtherClient) -> None: + readiness = unwrap(client.readiness_public()) + assert readiness.status == "healthy", ( + f"public readiness must report a healthy worker, got status {readiness.status!r}" + ) + + @pytest.mark.covers("other.lifecycle.readiness.reports_db_status") + def test_readiness_reports_connected_db(self, client: OtherClient) -> None: + readiness = unwrap(client.readiness_public()) + assert readiness.db == "connected", ( + "readiness must report the configured database as connected so an " + f"orchestrator can tell a healthy worker from a DB-unreachable one, got {readiness.db!r}" + ) + + @pytest.mark.covers("other.lifecycle.readiness_details.authenticated_diagnostics") + def test_readiness_details_require_auth_and_expose_diagnostics(self, client: OtherClient) -> None: + anonymous = client.readiness_details_unauthenticated() + assert isinstance(anonymous, UnauthorizedError), ( + f"/health/readiness/details must reject an unauthenticated caller, got {anonymous}" + ) + + details = unwrap(client.readiness_details(MASTER_KEY)) + assert details.status == "healthy", f"authenticated readiness status must be healthy, got {details.status!r}" + assert details.litellm_version is not None, ( + "authenticated diagnostics must expose the litellm version" + ) + assert details.db == "connected", ( + f"authenticated diagnostics must report the DB as connected, got {details.db!r}" + ) diff --git a/tests/e2e/other/test_master_key_auth_e2e.py b/tests/e2e/other/test_master_key_auth_e2e.py new file mode 100644 index 00000000000..6ab33c9b62a --- /dev/null +++ b/tests/e2e/other/test_master_key_auth_e2e.py @@ -0,0 +1,37 @@ +"""Live e2e: the master key authenticates and is treated as a proxy admin, and a +key that is not the master key is rejected before reaching the handler. + +/user/list is admin-only, so it proves both halves of the master-key contract in +one route: the master key reads it (authenticated + authorized as admin), while a +freshly minted, never-provisioned token is denied 401 by the auth layer. The +invalid case uses a unique, master-key-shaped token so the check exercises the +credential comparison rather than a value that could collide with a real key. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import MASTER_KEY, unique_marker +from e2e_http import UnauthorizedError, unwrap +from other_client import OtherClient + +pytestmark = pytest.mark.e2e + + +class TestMasterKeyAuth: + @pytest.mark.covers("other.auth.master_key.valid_allows") + def test_master_key_authenticates_and_grants_admin_route(self, client: OtherClient) -> None: + listing = unwrap(client.list_users_as(MASTER_KEY)) + assert listing.total >= 0, ( + "master key reached the admin /user/list handler but the response did not " + f"carry a user count: {listing}" + ) + + @pytest.mark.covers("other.auth.master_key.invalid_denied") + def test_non_matching_master_key_is_denied(self, client: OtherClient) -> None: + bogus = f"sk-{unique_marker()}" + result = client.list_users_as(bogus) + assert isinstance(result, UnauthorizedError), ( + f"a token that is not the master key must be rejected with 401, got {result}" + ) diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index e9611df139b..2998a4b83c6 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -6,3 +6,4 @@ addopts = --strict-markers --strict-config markers = e2e: live test that requires a running proxy and real provider keys load: heavy throughput/load test; collected last so it never perturbs latency-sensitive suites + weekly: real-provider anomaly load test that spends real money; deselected unless E2E_WEEKLY_ANOMALY is set diff --git a/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py b/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py index c4ad0c38f31..8b93afb4752 100644 --- a/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py +++ b/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py @@ -1,28 +1,36 @@ """Live e2e: a tiny max_budget on an entity actually blocks requests. -Each entity is an E2ECase (lifecycle.E2ECase) driven by run_case: init() creates -the budgeted entity + a key, run() drives spend until a `budget_exceeded` block, -teardown() deletes everything init() created (always runs, even on failure/skip). -Covers the entities with no prior live coverage - internal user, end-user, -organization, team member - plus key and team. See BUDGET_TEST_COVERAGE_MATRIX.md. +One test per budget level (key, team, internal user, end-user, organization, +team member): put the tiny cap on that level, drive spend until a +`budget_exceeded` block, and where a cap could be confused with a neighbor, +prove isolation with an uncapped control key that must keep serving. The +capped-key sweep proves the key's own max_budget blocks across mint shapes +(personal, team, team-member) with roomy surroundings, so the key-level cap is +provably the blocker no matter who the key was minted to. A non-budget error fails hard (never a skip); if calls never get blocked, budget enforcement is broken -> fail. """ import time -from dataclasses import dataclass, field -from typing import Callable, List, Type import pytest from budget_client import BudgetClient, is_budget_block from e2e_config import unique_marker from e2e_http import StreamingResponse, require_successful_call -from lifecycle import run_case +from lifecycle import ResourceManager pytestmark = pytest.mark.e2e +TINY_CAP = 3e-6 +ROOMY_CAP = 100.0 + + +def _chat(client: BudgetClient, key: str, *, user: str | None = None) -> StreamingResponse: + return client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16, user=user) + + def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> StreamingResponse: """Send paid calls until the entity's budget blocks one; return the blocked response so callers can assert on its shape. Key/user/org/member block within @@ -30,13 +38,7 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> enforces off table spend that lands on the batch write, so it takes a few more. A non-budget error fails hard (never a skip).""" for _ in range(40): - result = client.chat( - key, - "claude-haiku-4-5", - f"spend {unique_marker()}", - max_tokens=16, - user=user or None, - ) + result = _chat(client, key, user=user or None) if is_budget_block(result): return result require_successful_call(result) @@ -44,225 +46,154 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> pytest.fail("budget never enforced within the call budget") -@dataclass -class _BudgetCase: - """Base E2ECase: a key under some budgeted entity must get blocked. - - Subclasses set up the budgeted entity in init() and register every created id - in `_undo` (run LIFO in teardown so a key is deleted before its team/org). - """ - - client: BudgetClient - key: str = "" - _undo: List[Callable[[], None]] = field( - default_factory=list - ) # mutable-ok: per-case teardown registry - - def init(self) -> None: - raise NotImplementedError - - def run(self) -> None: - _assert_budget_blocks(self.client, self.key) - - def teardown(self) -> None: - for undo in reversed(self._undo): - undo() +def _assert_blocked_429(client: BudgetClient, key: str) -> StreamingResponse: + blocked = _assert_budget_blocks(client, key) + assert blocked.status_code == 429, ( + f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" + ) + return blocked -class KeyBudgetCase(_BudgetCase): - """A bare key (no team_id / user_id) carrying its own max_budget, so only the - key-level budget can be the thing that blocks. The refusal must be a 429 - budget_exceeded; any other error already fails via _assert_budget_blocks.""" +class TestBudgetBlocksPerLevel: + @pytest.mark.covers("quota_management.budget.key.blocks_over_limit") + def test_bare_key_blocks_over_its_own_budget(self, client: BudgetClient, resources: ResourceManager) -> None: + key = client.generate_key(max_budget=TINY_CAP) + resources.defer(lambda: client.delete_key(key)) - def init(self) -> None: - self.key = self.client.generate_key(max_budget=3e-6) - self._undo.append(lambda: self.client.delete_key(self.key)) + _assert_blocked_429(client, key) - def run(self) -> None: - blocked = _assert_budget_blocks(self.client, self.key) - assert blocked.status_code == 429, ( - f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" - ) + @pytest.mark.covers("quota_management.budget.team.blocks_over_limit") + def test_team_budget_blocks_every_team_key(self, client: BudgetClient, resources: ResourceManager) -> None: + team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", max_budget=TINY_CAP) + resources.defer(lambda: client.delete_team(team_id)) + spender_key = client.generate_key(team_id=team_id) + resources.defer(lambda: client.delete_key(spender_key)) + sibling_key = client.generate_key(team_id=team_id) + resources.defer(lambda: client.delete_key(sibling_key)) - -class TeamBudgetCase(_BudgetCase): - """An admin caps a whole team: two keys under a tiny-budget team, neither with - a key-level budget. Key A is driven until the team cap blocks it; key B's very - first call must then be refused too, proving the cap sits on the team, not the - key that spent. Both refusals must be 429 budget_exceeded.""" - - def init(self) -> None: - team_id = self.client.create_team( - alias=f"e2e-budget-team-{unique_marker()}", max_budget=3e-6 - ) - self._undo.append(lambda: self.client.delete_team(team_id)) - self.key = self.client.generate_key(team_id=team_id) - self._undo.append(lambda: self.client.delete_key(self.key)) - self._sibling_key = self.client.generate_key(team_id=team_id) - self._undo.append(lambda: self.client.delete_key(self._sibling_key)) - - def run(self) -> None: - blocked = _assert_budget_blocks(self.client, self.key) - assert blocked.status_code == 429, ( - f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" - ) - sibling = self.client.chat( - self._sibling_key, - "claude-haiku-4-5", - f"spend {unique_marker()}", - max_tokens=16, - ) + _assert_blocked_429(client, spender_key) + sibling = _chat(client, sibling_key) assert is_budget_block(sibling) and sibling.status_code == 429, ( f"a sibling key on the capped team must get the same 429 budget_exceeded, " f"got {sibling.status_code}: {sibling.body[:200]}" ) + @pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit") + def test_user_budget_enforced_across_all_their_keys( + self, client: BudgetClient, resources: ResourceManager + ) -> None: + user_id = client.create_user(max_budget=TINY_CAP) + resources.defer(lambda: client.delete_user(user_id)) + first_key = client.generate_key(user_id=user_id) + resources.defer(lambda: client.delete_key(first_key)) + second_key = client.generate_key(user_id=user_id) + resources.defer(lambda: client.delete_key(second_key)) + team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}") + resources.defer(lambda: client.delete_team(team_id)) + client.add_team_member(team_id, user_id) + team_key = client.generate_key(team_id=team_id, user_id=user_id) + resources.defer(lambda: client.delete_key(team_key)) -class InternalUserBudgetCase(_BudgetCase): - """A user's max_budget follows the person, not the key. The capped user holds - two personal keys (no team, no key budgets) plus a team-member key on an - uncapped team; once the first personal key is refused, the other two must be - refused as well - a second key is not a fresh allowance, and since #32005 the - user budget draws down team keys too. All refusals must be 429 budget_exceeded.""" - - def init(self) -> None: - user_id = self.client.create_user(max_budget=3e-6) - self._undo.append(lambda: self.client.delete_user(user_id)) - self.key = self.client.generate_key(user_id=user_id) - self._undo.append(lambda: self.client.delete_key(self.key)) - self._second_key = self.client.generate_key(user_id=user_id) - self._undo.append(lambda: self.client.delete_key(self._second_key)) - team_id = self.client.create_team(alias=f"e2e-budget-team-{unique_marker()}") - self._undo.append(lambda: self.client.delete_team(team_id)) - self.client.add_team_member(team_id, user_id) - self._team_key = self.client.generate_key(team_id=team_id, user_id=user_id) - self._undo.append(lambda: self.client.delete_key(self._team_key)) - - def run(self) -> None: - blocked = _assert_budget_blocks(self.client, self.key) - assert blocked.status_code == 429, ( - f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" - ) - for label, key in (("second personal key", self._second_key), ("team-member key", self._team_key)): - result = self.client.chat(key, "claude-haiku-4-5", f"spend {unique_marker()}", max_tokens=16) + _assert_blocked_429(client, first_key) + for label, key in (("second personal key", second_key), ("team-member key", team_key)): + result = _chat(client, key) assert is_budget_block(result) and result.status_code == 429, ( f"the {label} of a user over budget must get the same 429 budget_exceeded, " f"got {result.status_code}: {result.body[:200]}" ) - -class EndUserBudgetCase(_BudgetCase): - def init(self) -> None: + @pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit") + def test_end_user_budget_blocks_attributed_calls( + self, client: BudgetClient, resources: ResourceManager + ) -> None: customer = f"e2e-budget-cust-{unique_marker()}" - self.client.create_customer(customer, max_budget=3e-6) - self._undo.append(lambda: self.client.delete_customers([customer])) - self.key = self.client.generate_key(models=["claude-haiku-4-5"]) - self._undo.append(lambda: self.client.delete_key(self.key)) - self._customer = customer + client.create_customer(customer, max_budget=TINY_CAP) + resources.defer(lambda: client.delete_customers([customer])) + key = client.generate_key(models=["claude-haiku-4-5"]) + resources.defer(lambda: client.delete_key(key)) - def run(self) -> None: - _assert_budget_blocks(self.client, self.key, user=self._customer) + _assert_budget_blocks(client, key, user=customer) + @pytest.mark.covers("quota_management.budget.organization.blocks_over_limit") + def test_org_budget_blocks_keys_under_it(self, client: BudgetClient, resources: ResourceManager) -> None: + org_id = client.create_org(max_budget=TINY_CAP, alias=f"e2e-budget-org-{unique_marker()}") + resources.defer(lambda: client.delete_org(org_id)) + team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", organization_id=org_id) + resources.defer(lambda: client.delete_team(team_id)) + key = client.generate_key(team_id=team_id) + resources.defer(lambda: client.delete_key(key)) -class OrganizationBudgetCase(_BudgetCase): - """Org carries the tiny budget; the team under it and the key carry none, so - the org is the only entity that can block (the historically weak link). The - refusal must be a 429 budget_exceeded that names the org as the blocker.""" - - def init(self) -> None: - self._org_id = self.client.create_org( - max_budget=3e-6, alias=f"e2e-budget-org-{unique_marker()}" - ) - self._undo.append(lambda: self.client.delete_org(self._org_id)) - team_id = self.client.create_team( - alias=f"e2e-budget-team-{unique_marker()}", organization_id=self._org_id - ) - self._undo.append(lambda: self.client.delete_team(team_id)) - self.key = self.client.generate_key(team_id=team_id) - self._undo.append(lambda: self.client.delete_key(self.key)) - - def run(self) -> None: - blocked = _assert_budget_blocks(self.client, self.key) - assert blocked.status_code == 429, ( - f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" - ) - assert f"Organization={self._org_id}" in blocked.body, ( + blocked = _assert_blocked_429(client, key) + assert f"Organization={org_id}" in blocked.body, ( f"refusal must name the org as the blocker, got: {blocked.body[:200]}" ) + @pytest.mark.covers("quota_management.budget.team_member.blocks_over_limit") + def test_member_budget_blocks_without_touching_teammates( + self, client: BudgetClient, resources: ResourceManager + ) -> None: + team_id = client.create_team(alias=f"e2e-budget-team-{unique_marker()}", max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_team(team_id)) + member_id = client.create_user(max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_user(member_id)) + client.add_team_member(team_id, member_id, max_budget_in_team=TINY_CAP) + member_key = client.generate_key(team_id=team_id, user_id=member_id) + resources.defer(lambda: client.delete_key(member_key)) + teammate_id = client.create_user(max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_user(teammate_id)) + client.add_team_member(team_id, teammate_id) + teammate_key = client.generate_key(team_id=team_id, user_id=teammate_id) + resources.defer(lambda: client.delete_key(teammate_key)) -class TeamMemberBudgetCase(_BudgetCase): - """Member A's per-team budget is tiny while the team and both members' user - budgets are roomy (100.0), so the only cap that can trip is A's: a block - proves member-level enforcement and must be a 429 budget_exceeded. Teammate - B, uncapped on the same team, must keep serving after A is cut off, proving - the member cap does not leak onto the team or its members.""" - - def init(self) -> None: - self._team_id = self.client.create_team( - alias=f"e2e-budget-team-{unique_marker()}", max_budget=100.0 - ) - self._undo.append(lambda: self.client.delete_team(self._team_id)) - self._member_id = self.client.create_user(max_budget=100.0) - self._undo.append(lambda: self.client.delete_user(self._member_id)) - self.client.add_team_member(self._team_id, self._member_id, max_budget_in_team=3e-6) - self.key = self.client.generate_key(team_id=self._team_id, user_id=self._member_id) - self._undo.append(lambda: self.client.delete_key(self.key)) - teammate_id = self.client.create_user(max_budget=100.0) - self._undo.append(lambda: self.client.delete_user(teammate_id)) - self.client.add_team_member(self._team_id, teammate_id) - self._teammate_key = self.client.generate_key(team_id=self._team_id, user_id=teammate_id) - self._undo.append(lambda: self.client.delete_key(self._teammate_key)) - - def run(self) -> None: - blocked = _assert_budget_blocks(self.client, self.key) - assert blocked.status_code == 429, ( - f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" - ) - teammate = self.client.chat( - self._teammate_key, - "claude-haiku-4-5", - f"spend {unique_marker()}", - max_tokens=16, - ) - require_successful_call(teammate) + _assert_blocked_429(client, member_key) + require_successful_call(_chat(client, teammate_key)) -def _case_id(case_cls: Type[_BudgetCase]) -> str: - return case_cls.__name__ +class TestKeyBudgetBlocksAcrossKeyKinds: + """The tiny max_budget sits on the key itself while every budget around it + (user / team / membership) is roomy, so only the key-level cap can block; the + uncapped control key minted to the same surroundings must keep serving after + the capped key is refused, proving nothing around the key was the blocker.""" + @pytest.mark.covers("quota_management.budget.key.blocks_over_limit") + def test_personal_key_blocks_over_its_own_budget( + self, client: BudgetClient, resources: ResourceManager + ) -> None: + user_id = client.create_user(max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_user(user_id)) + capped_key = client.generate_key(user_id=user_id, max_budget=TINY_CAP) + resources.defer(lambda: client.delete_key(capped_key)) + control_key = client.generate_key(user_id=user_id) + resources.defer(lambda: client.delete_key(control_key)) -@pytest.mark.parametrize( - "case_cls", - [ - pytest.param( - KeyBudgetCase, - marks=pytest.mark.covers("quota_management.budget.key.blocks_over_limit"), - ), - pytest.param( - TeamBudgetCase, - marks=pytest.mark.covers("quota_management.budget.team.blocks_over_limit"), - ), - pytest.param( - InternalUserBudgetCase, - marks=pytest.mark.covers("quota_management.budget.internal_user.blocks_over_limit"), - ), - pytest.param( - EndUserBudgetCase, - marks=pytest.mark.covers("quota_management.budget.end_user.blocks_over_limit"), - ), - pytest.param( - OrganizationBudgetCase, - marks=pytest.mark.covers("quota_management.budget.organization.blocks_over_limit"), - ), - pytest.param( - TeamMemberBudgetCase, - marks=pytest.mark.covers("quota_management.budget.team_member.blocks_over_limit"), - ), - ], - ids=_case_id, -) -def test_budget_enforcement( - client: BudgetClient, case_cls: Type[_BudgetCase] -) -> None: - run_case(case_cls(client)) + _assert_blocked_429(client, capped_key) + require_successful_call(_chat(client, control_key)) + + @pytest.mark.covers("quota_management.budget.key.blocks_over_limit") + def test_team_key_blocks_over_its_own_budget(self, client: BudgetClient, resources: ResourceManager) -> None: + team_id = client.create_team(alias=f"e2e-key-cap-team-{unique_marker()}", max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_team(team_id)) + capped_key = client.generate_key(team_id=team_id, max_budget=TINY_CAP) + resources.defer(lambda: client.delete_key(capped_key)) + control_key = client.generate_key(team_id=team_id) + resources.defer(lambda: client.delete_key(control_key)) + + _assert_blocked_429(client, capped_key) + require_successful_call(_chat(client, control_key)) + + @pytest.mark.covers("quota_management.budget.key.blocks_over_limit") + def test_team_member_key_blocks_over_its_own_budget( + self, client: BudgetClient, resources: ResourceManager + ) -> None: + team_id = client.create_team(alias=f"e2e-key-cap-team-{unique_marker()}", max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_team(team_id)) + member_id = client.create_user(max_budget=ROOMY_CAP) + resources.defer(lambda: client.delete_user(member_id)) + client.add_team_member(team_id, member_id, max_budget_in_team=ROOMY_CAP) + capped_key = client.generate_key(team_id=team_id, user_id=member_id, max_budget=TINY_CAP) + resources.defer(lambda: client.delete_key(capped_key)) + control_key = client.generate_key(team_id=team_id, user_id=member_id) + resources.defer(lambda: client.delete_key(control_key)) + + _assert_blocked_429(client, capped_key) + require_successful_call(_chat(client, control_key)) diff --git a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py index d43d8e94898..b6e33627a5e 100644 --- a/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py +++ b/tests/e2e/quota_management/spend_tracking/test_spend_tracking_e2e.py @@ -23,7 +23,7 @@ import pytest from e2e_http import Result, Success from lifecycle import ResourceManager -from models import ChatResponse, SpendLogs, SpendLogsParams +from models import ChatResponse, LiteLLMParamsBody, SpendLogs, SpendLogsParams from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap pytestmark = pytest.mark.e2e @@ -194,6 +194,7 @@ def test_streaming_messages_via_responses_bridge_tracks_spend( @pytest.mark.covers("quota_management.spend_tracking.embeddings.logs_cost") +@pytest.mark.covers("llm.embeddings.openai.basic.nonstream.cost_logged") def test_embedding_writes_nonzero_spend_row( client: SpendClient, scoped_key: str ) -> None: @@ -231,14 +232,12 @@ def test_cache_hit_is_zero_cost_and_suffixed( rows = client.poll_logs_for_key( scoped_key, predicate=lambda rs: any(r.cache_hit == "True" for r in rs) ) - cache_rows = [r for r in rows if r.cache_hit == "True"] - if not cache_rows: - pytest.skip( - "no cache-hit row observed; caching may be disabled on this proxy. " - f"rows seen: {_summarize(rows)}" - ) - - cache_row = cache_rows[0] + cache_row = _require_row( + rows, + lambda r: r.cache_hit == "True", + "with cache_hit=True (caching is enabled on the e2e proxy, so an identical " + "repeat call must hit the cache)", + ) assert ( cache_row.spend or 0 ) == 0.0, f"cache hit was charged (double-charge regression): {_summarize(rows)}" @@ -503,22 +502,27 @@ def test_each_model_on_a_shared_key_gets_its_own_row( @pytest.mark.covers("quota_management.spend_tracking.failure.writes_failure_row") def test_failure_call_writes_failure_status_row( - client: SpendClient, scoped_key: str + client: SpendClient, resources: ResourceManager, scoped_key: str ) -> None: - result = client.chat(scoped_key, "gemini-2.5-flash", "", max_tokens=1) - if is_ok(result): - pytest.skip("call unexpectedly succeeded; could not induce a failure row") + model = f"e2e-spend-failure-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody(model="openai/gpt-5.5", api_key="sk-invalid-e2e-failure-row"), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + result = client.chat(scoped_key, model, f"trigger failure {unique_marker()}", max_tokens=1) + assert not is_ok(result), ( + f"a call to a deployment with an invalid upstream key must fail, not succeed: {result}" + ) rows = client.poll_logs_for_key( scoped_key, predicate=lambda rs: any(r.status == "failure" for r in rs) ) - failure_rows = [r for r in rows if r.status == "failure"] - if not failure_rows: - pytest.skip( - "no failure-status row was logged for the rejected call; " - "failure logging is environment-specific" - ) - assert (failure_rows[0].spend or 0) == 0.0, "failed call must not be charged" + failure_row = _require_row( + rows, lambda r: r.status == "failure", "with status=failure for the rejected call" + ) + assert (failure_row.spend or 0) == 0.0, "failed call must not be charged" @pytest.mark.covers("quota_management.spend_tracking.spend_calculate.returns_cost") diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py index 005b49272e8..da4252e550e 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -16,7 +16,7 @@ import e2e_http from e2e_http import ( URL, AuthHeaders, - FileUploadForm, + BinaryStream, ProbeResult, Result, StreamingResponse, @@ -32,6 +32,15 @@ class Transport(Protocol): self, path: str, *, headers: BaseModel, json: BaseModel ) -> StreamingResponse: ... + def stream_binary( + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + chunk_size: int = 8192, + ) -> BinaryStream: ... + def send( self, path: str, @@ -65,6 +74,10 @@ class Transport(Protocol): self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] ) -> Result[R]: ... + def put[R: BaseModel]( + self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + ) -> Result[R]: ... + def probe(self, path: str, *, params: BaseModel) -> ProbeResult: ... def upload[R: BaseModel]( @@ -72,9 +85,10 @@ class Transport(Protocol): path: str, *, headers: BaseModel, - form: FileUploadForm, + form: BaseModel, filename: str, content: bytes, + file_content_type: str = "application/jsonl", params: BaseModel | None = None, response_type: type[R], ) -> Result[R]: ... @@ -159,6 +173,17 @@ class HttpTransport: timeout=self.request_timeout, ) + def put[R: BaseModel]( + self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + ) -> Result[R]: + return e2e_http.put( + self._url(path), + headers=headers, + json=json, + response_type=response_type, + timeout=self.request_timeout, + ) + def stream( self, path: str, *, headers: BaseModel, json: BaseModel ) -> StreamingResponse: @@ -166,6 +191,22 @@ class HttpTransport: self._url(path), headers=headers, json=json, timeout=self.request_timeout ) + def stream_binary( + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + chunk_size: int = 8192, + ) -> BinaryStream: + return e2e_http.stream_binary( + self._url(path), + headers=headers, + json=json, + chunk_size=chunk_size, + timeout=self.request_timeout, + ) + def send( self, path: str, @@ -197,9 +238,10 @@ class HttpTransport: path: str, *, headers: BaseModel, - form: FileUploadForm, + form: BaseModel, filename: str, content: bytes, + file_content_type: str = "application/jsonl", params: BaseModel | None = None, response_type: type[R], ) -> Result[R]: @@ -209,6 +251,7 @@ class HttpTransport: form=form, filename=filename, content=content, + file_content_type=file_content_type, params=params, response_type=response_type, timeout=self.request_timeout, @@ -234,6 +277,7 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = ( "/tag", "/budget", "/model/", + "/access_group", "/spend", "/global", "/config", @@ -318,11 +362,30 @@ class SplitTransport: path, headers=headers, json=json, response_type=response_type ) + def put[R: BaseModel]( + self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + ) -> Result[R]: + return self._route(path).put( + path, headers=headers, json=json, response_type=response_type + ) + def stream( self, path: str, *, headers: BaseModel, json: BaseModel ) -> StreamingResponse: return self._route(path).stream(path, headers=headers, json=json) + def stream_binary( + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + chunk_size: int = 8192, + ) -> BinaryStream: + return self._route(path).stream_binary( + path, headers=headers, json=json, chunk_size=chunk_size + ) + def send( self, path: str, @@ -344,9 +407,10 @@ class SplitTransport: path: str, *, headers: BaseModel, - form: FileUploadForm, + form: BaseModel, filename: str, content: bytes, + file_content_type: str = "application/jsonl", params: BaseModel | None = None, response_type: type[R], ) -> Result[R]: @@ -356,6 +420,7 @@ class SplitTransport: form=form, filename=filename, content=content, + file_content_type=file_content_type, params=params, response_type=response_type, ) diff --git a/ui/litellm-dashboard/e2e_tests/constants.ts b/tests/e2e/ui/constants.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/constants.ts rename to tests/e2e/ui/constants.ts diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/config.yml b/tests/e2e/ui/fixtures/config.yml similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/config.yml rename to tests/e2e/ui/fixtures/config.yml diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts b/tests/e2e/ui/fixtures/menuMappings.ts similarity index 93% rename from ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts rename to tests/e2e/ui/fixtures/menuMappings.ts index 4a4bb64c8ed..d6e7ea86982 100644 --- a/ui/litellm-dashboard/e2e_tests/fixtures/menuMappings.ts +++ b/tests/e2e/ui/fixtures/menuMappings.ts @@ -26,7 +26,8 @@ export const menuLabelToPage: Record = { "Cost Tracking": Page.CostTracking, "UI Theme": Page.UiTheme, // Experimental submenu items - Caching: Page.Caching, + "Response Cache": Page.Caching, + Caching: Page.Caching, // Legacy label support Prompts: Page.Prompts, Budgets: Page.Budgets, "API Playground": Page.TransformRequest, diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/migratedPages.ts b/tests/e2e/ui/fixtures/migratedPages.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/migratedPages.ts rename to tests/e2e/ui/fixtures/migratedPages.ts diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py b/tests/e2e/ui/fixtures/mock_llm_server/server.py similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/mock_llm_server/server.py rename to tests/e2e/ui/fixtures/mock_llm_server/server.py diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/pages.ts b/tests/e2e/ui/fixtures/pages.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/pages.ts rename to tests/e2e/ui/fixtures/pages.ts diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/roles.ts b/tests/e2e/ui/fixtures/roles.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/roles.ts rename to tests/e2e/ui/fixtures/roles.ts diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/seed.sql b/tests/e2e/ui/fixtures/seed.sql similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/seed.sql rename to tests/e2e/ui/fixtures/seed.sql diff --git a/ui/litellm-dashboard/e2e_tests/fixtures/users.ts b/tests/e2e/ui/fixtures/users.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/fixtures/users.ts rename to tests/e2e/ui/fixtures/users.ts diff --git a/ui/litellm-dashboard/e2e_tests/globalSetup.ts b/tests/e2e/ui/globalSetup.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/globalSetup.ts rename to tests/e2e/ui/globalSetup.ts diff --git a/ui/litellm-dashboard/e2e_tests/helpers/navigation.ts b/tests/e2e/ui/helpers/navigation.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/helpers/navigation.ts rename to tests/e2e/ui/helpers/navigation.ts diff --git a/ui/litellm-dashboard/e2e_tests/migration.serverRootPath.config.ts b/tests/e2e/ui/migration.serverRootPath.config.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/migration.serverRootPath.config.ts rename to tests/e2e/ui/migration.serverRootPath.config.ts diff --git a/ui/litellm-dashboard/e2e_tests/migration.serverRootPath.globalSetup.ts b/tests/e2e/ui/migration.serverRootPath.globalSetup.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/migration.serverRootPath.globalSetup.ts rename to tests/e2e/ui/migration.serverRootPath.globalSetup.ts diff --git a/tests/e2e/ui/package-lock.json b/tests/e2e/ui/package-lock.json new file mode 100644 index 00000000000..b22673a3535 --- /dev/null +++ b/tests/e2e/ui/package-lock.json @@ -0,0 +1,111 @@ +{ + "name": "litellm-ui-e2e", + "version": "0.0.0", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "litellm-ui-e2e", + "version": "0.0.0", + "devDependencies": { + "@playwright/test": "1.58.1", + "@types/node": "20.19.37", + "typescript": "5.9.3" + } + }, + "node_modules/@playwright/test": { + "version": "1.58.1", + "resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.58.1.tgz", + "integrity": "sha512-6LdVIUERWxQMmUSSQi0I53GgCBYgM2RpGngCPY7hSeju+VrKjq3lvs7HpJoPbDiY5QM5EYRtRX5fvrinnMAz3w==", + "dev": true, + "license": "Apache-2.0", + "dependencies": { + "playwright": "1.58.1" + }, + "bin": { + "playwright": "cli.js" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/@types/node": { + "version": "20.19.37", + "resolved": "https://registry.npmjs.org/@types/node/-/node-20.19.37.tgz", + "integrity": "sha512-8kzdPJ3FsNsVIurqBs7oodNnCEVbni9yUEkaHbgptDACOPW04jimGagZ51E6+lXUwJjgnBw+hyko/lkFWCldqw==", + "dev": true, + "license": "MIT", + "dependencies": { + "undici-types": "~6.21.0" + } + }, + "node_modules/fsevents": { + "version": "2.3.2", + "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz", + "integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==", + "dev": true, + "hasInstallScript": true, + "license": "MIT", + "optional": true, + "os": [ + "darwin" + ], + "engines": { + "node": "^8.16.0 || ^10.6.0 || >=11.0.0" + } + }, + "node_modules/playwright": { + "version": "1.58.1", + "resolved": "https://registry.npmjs.org/playwright/-/playwright-1.58.1.tgz", + "integrity": "sha512-+2uTZHxSCcxjvGc5C891LrS1/NlxglGxzrC4seZiVjcYVQfUa87wBL6rTDqzGjuoWNjnBzRqKmF6zRYGMvQUaQ==", + "dev": true, + "license": "Apache-2.0", + "dependencies": { + "playwright-core": "1.58.1" + }, + "bin": { + "playwright": "cli.js" + }, + "engines": { + "node": ">=18" + }, + "optionalDependencies": { + "fsevents": "2.3.2" + } + }, + "node_modules/playwright-core": { + "version": "1.58.1", + "resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.58.1.tgz", + "integrity": "sha512-bcWzOaTxcW+VOOGBCQgnaKToLJ65d6AqfLVKEWvexyS3AS6rbXl+xdpYRMGSRBClPvyj44njOWoxjNdL/H9UNg==", + "dev": true, + "license": "Apache-2.0", + "bin": { + "playwright-core": "cli.js" + }, + "engines": { + "node": ">=18" + } + }, + "node_modules/typescript": { + "version": "5.9.3", + "resolved": "https://registry.npmjs.org/typescript/-/typescript-5.9.3.tgz", + "integrity": "sha512-jl1vZzPDinLr9eUt3J/t7V6FgNEw9QjvBPdysz9KfQDD41fQrC2Y4vKQdiaUpFT4bXlb1RHhLpp8wtm6M5TgSw==", + "dev": true, + "license": "Apache-2.0", + "bin": { + "tsc": "bin/tsc", + "tsserver": "bin/tsserver" + }, + "engines": { + "node": ">=14.17" + } + }, + "node_modules/undici-types": { + "version": "6.21.0", + "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz", + "integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==", + "dev": true, + "license": "MIT" + } + } +} diff --git a/tests/e2e/ui/package.json b/tests/e2e/ui/package.json new file mode 100644 index 00000000000..ede759d97cb --- /dev/null +++ b/tests/e2e/ui/package.json @@ -0,0 +1,16 @@ +{ + "name": "litellm-ui-e2e", + "version": "0.0.0", + "private": true, + "scripts": { + "e2e": "playwright test --config playwright.config.ts", + "e2e:ui": "playwright test --ui --config playwright.config.ts", + "e2e:migration": "playwright test tests/migration/migratedPages.spec.ts --config playwright.config.ts", + "e2e:migration:root": "playwright test --config migration.serverRootPath.config.ts" + }, + "devDependencies": { + "@playwright/test": "1.58.1", + "@types/node": "20.19.37", + "typescript": "5.9.3" + } +} diff --git a/ui/litellm-dashboard/e2e_tests/playwright.config.ts b/tests/e2e/ui/playwright.config.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/playwright.config.ts rename to tests/e2e/ui/playwright.config.ts diff --git a/ui/litellm-dashboard/e2e_tests/run_e2e.sh b/tests/e2e/ui/run_e2e.sh similarity index 91% rename from ui/litellm-dashboard/e2e_tests/run_e2e.sh rename to tests/e2e/ui/run_e2e.sh index ed0641d04e6..858eb401c8e 100755 --- a/ui/litellm-dashboard/e2e_tests/run_e2e.sh +++ b/tests/e2e/ui/run_e2e.sh @@ -20,12 +20,13 @@ set -euo pipefail # ================================================================ SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" -DASHBOARD_DIR="$(cd "$SCRIPT_DIR/.." && pwd)" REPO_ROOT="$(cd "$SCRIPT_DIR/../../.." && pwd)" +DASHBOARD_DIR="$REPO_ROOT/ui/litellm-dashboard" IS_CI="${CI:-false}" CONTAINER_NAME="litellm-e2e-postgres-$$" MOCK_PID="" PROXY_PID="" +PROXY_LOG="" # --- Ensure common tool paths are available (local dev only) --- if [ "$IS_CI" = "false" ]; then @@ -40,6 +41,7 @@ cleanup() { echo "Cleaning up..." [ -n "$MOCK_PID" ] && kill "$MOCK_PID" 2>/dev/null || true [ -n "$PROXY_PID" ] && kill "$PROXY_PID" 2>/dev/null || true + [ -n "$PROXY_LOG" ] && rm -f "$PROXY_LOG" || true if [ "$IS_CI" = "false" ]; then docker stop "$CONTAINER_NAME" 2>/dev/null || true fi @@ -124,6 +126,7 @@ echo "UI build copied and restructured" # --- Python environment --- echo "=== Setting up Python environment ===" cd "$REPO_ROOT" +export UV_PYTHON="${UV_PYTHON:-3.13}" uv sync --group dev --group proxy-dev --extra proxy --frozen --quiet uv run --no-sync python -m prisma generate --schema litellm/proxy/schema.prisma @@ -143,16 +146,18 @@ done # --- LiteLLM proxy --- echo "=== Starting LiteLLM proxy ===" cd "$REPO_ROOT" +PROXY_LOG="${TMPDIR:-/tmp}/litellm-e2e-proxy-$$.log" uv run --no-sync python -m litellm.proxy.proxy_cli \ --config "$SCRIPT_DIR/fixtures/config.yml" \ - --port 4000 & + --port 4000 >"$PROXY_LOG" 2>&1 & PROXY_PID=$! -echo "Waiting for proxy..." +echo "Waiting for proxy (logs: $PROXY_LOG)..." PROXY_READY=0 for i in $(seq 1 180); do if ! kill -0 "$PROXY_PID" 2>/dev/null; then - echo "Error: proxy process exited unexpectedly" + echo "Error: proxy process exited unexpectedly. Proxy output:" + tail -n 100 "$PROXY_LOG" exit 1 fi HTTP_CODE=$(curl -s -o /dev/null -w "%{http_code}" http://127.0.0.1:4000/health -H "Authorization: Bearer $LITELLM_MASTER_KEY" 2>/dev/null || true) @@ -163,7 +168,8 @@ for i in $(seq 1 180); do sleep 1 done if [ "$PROXY_READY" -ne 1 ]; then - echo "Error: proxy did not become healthy within 180 seconds" + echo "Error: proxy did not become healthy within 180 seconds. Proxy output:" + tail -n 100 "$PROXY_LOG" exit 1 fi echo "Proxy is ready." @@ -181,12 +187,12 @@ PGPASSWORD="$DB_PASS" psql -h "$DB_HOST" -p "$DB_PORT" -U "$DB_USER" -d "$DB_NAM # --- Playwright --- echo "=== Installing Playwright dependencies ===" -cd "$DASHBOARD_DIR" +cd "$SCRIPT_DIR" npm install --silent 2>/dev/null || true npx playwright install chromium --with-deps 2>/dev/null || npx playwright install chromium echo "=== Running Playwright tests ===" -npx playwright test --config e2e_tests/playwright.config.ts "$@" +npx playwright test --config playwright.config.ts "$@" EXIT_CODE=$? exit $EXIT_CODE diff --git a/ui/litellm-dashboard/e2e_tests/serverRootPath.config.ts b/tests/e2e/ui/serverRootPath.config.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/serverRootPath.config.ts rename to tests/e2e/ui/serverRootPath.config.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/auth/logout.spec.ts b/tests/e2e/ui/tests/auth/logout.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/auth/logout.spec.ts rename to tests/e2e/ui/tests/auth/logout.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/auth/proxyLogoutUrl.spec.ts b/tests/e2e/ui/tests/auth/proxyLogoutUrl.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/auth/proxyLogoutUrl.spec.ts rename to tests/e2e/ui/tests/auth/proxyLogoutUrl.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/auth/unauthenticatedRedirect.spec.ts b/tests/e2e/ui/tests/auth/unauthenticatedRedirect.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/auth/unauthenticatedRedirect.spec.ts rename to tests/e2e/ui/tests/auth/unauthenticatedRedirect.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/internal-user/internalUser.spec.ts b/tests/e2e/ui/tests/internal-user/internalUser.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/internal-user/internalUser.spec.ts rename to tests/e2e/ui/tests/internal-user/internalUser.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/internal-user/internalUserNoTeam.spec.ts b/tests/e2e/ui/tests/internal-user/internalUserNoTeam.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/internal-user/internalUserNoTeam.spec.ts rename to tests/e2e/ui/tests/internal-user/internalUserNoTeam.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/internal-user/internalUserWithTeams.spec.ts b/tests/e2e/ui/tests/internal-user/internalUserWithTeams.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/internal-user/internalUserWithTeams.spec.ts rename to tests/e2e/ui/tests/internal-user/internalUserWithTeams.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/internal-viewer/internalViewer.spec.ts b/tests/e2e/ui/tests/internal-viewer/internalViewer.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/internal-viewer/internalViewer.spec.ts rename to tests/e2e/ui/tests/internal-viewer/internalViewer.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/login/internalUserIdentity.spec.ts b/tests/e2e/ui/tests/login/internalUserIdentity.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/login/internalUserIdentity.spec.ts rename to tests/e2e/ui/tests/login/internalUserIdentity.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/login/login.spec.ts b/tests/e2e/ui/tests/login/login.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/login/login.spec.ts rename to tests/e2e/ui/tests/login/login.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/login/serverRootPathRedirect.spec.ts b/tests/e2e/ui/tests/login/serverRootPathRedirect.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/login/serverRootPathRedirect.spec.ts rename to tests/e2e/ui/tests/login/serverRootPathRedirect.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/mcp/mcpServers.spec.ts b/tests/e2e/ui/tests/mcp/mcpServers.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/mcp/mcpServers.spec.ts rename to tests/e2e/ui/tests/mcp/mcpServers.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/migration/README.md b/tests/e2e/ui/tests/migration/README.md similarity index 84% rename from ui/litellm-dashboard/e2e_tests/tests/migration/README.md rename to tests/e2e/ui/tests/migration/README.md index 4b3a391d421..d6b33598ec4 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/migration/README.md +++ b/tests/e2e/ui/tests/migration/README.md @@ -9,8 +9,9 @@ the default mount and a non-root `SERVER_ROOT_PATH` mount. ## Adding a page When a page's migration merges, add its route segment to -`e2e_tests/fixtures/migratedPages.ts` (keep it in lockstep with `MIGRATED_PAGES` -in `src/utils/migratedPages.ts`). Both suites pick it up automatically. +`tests/e2e/ui/fixtures/migratedPages.ts` (keep it in lockstep with `MIGRATED_PAGES` +in `ui/litellm-dashboard/src/utils/migratedPages.ts`). Both suites pick it up +automatically. ## Running diff --git a/ui/litellm-dashboard/e2e_tests/tests/migration/migratedPages.spec.ts b/tests/e2e/ui/tests/migration/migratedPages.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/migration/migratedPages.spec.ts rename to tests/e2e/ui/tests/migration/migratedPages.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelHub/modelHub.spec.ts b/tests/e2e/ui/tests/modelHub/modelHub.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/modelHub/modelHub.spec.ts rename to tests/e2e/ui/tests/modelHub/modelHub.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts b/tests/e2e/ui/tests/modelsPage/addModel.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/modelsPage/addModel.spec.ts rename to tests/e2e/ui/tests/modelsPage/addModel.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/clearCustomPricing.spec.ts b/tests/e2e/ui/tests/modelsPage/clearCustomPricing.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/modelsPage/clearCustomPricing.spec.ts rename to tests/e2e/ui/tests/modelsPage/clearCustomPricing.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/credentials.spec.ts b/tests/e2e/ui/tests/modelsPage/credentials.spec.ts similarity index 95% rename from ui/litellm-dashboard/e2e_tests/tests/modelsPage/credentials.spec.ts rename to tests/e2e/ui/tests/modelsPage/credentials.spec.ts index 8b7824813a4..7c836068567 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/modelsPage/credentials.spec.ts +++ b/tests/e2e/ui/tests/modelsPage/credentials.spec.ts @@ -38,7 +38,8 @@ test.describe("Edit LLM credential", () => { const row = page.locator("tr", { hasText: credentialName }); await expect(row).toBeVisible({ timeout: 15_000 }); - await row.getByRole("button").first().click(); + await row.getByTestId(`credential-actions-${credentialName}`).click(); + await page.getByTestId("credential-action-edit").click(); const modal = page.locator(".ant-modal-content").filter({ hasText: "Edit Credential" }); await expect(modal).toBeVisible({ timeout: 10_000 }); diff --git a/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts b/tests/e2e/ui/tests/navigation/sidebar.spec.ts similarity index 94% rename from ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts rename to tests/e2e/ui/tests/navigation/sidebar.spec.ts index 7e42d07ae7c..b220dc09ae2 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/navigation/sidebar.spec.ts +++ b/tests/e2e/ui/tests/navigation/sidebar.spec.ts @@ -8,7 +8,16 @@ import { MIGRATED_E2E_PAGES } from "../../fixtures/migratedPages"; import type { Page as PlaywrightPage } from "@playwright/test"; const sidebarButtons = { - [Role.ProxyAdmin]: ["Virtual Keys", "Playground", "Models", "Usage", "Teams", "Internal Users", "AI Hub"], + [Role.ProxyAdmin]: [ + "Virtual Keys", + "Playground", + "Models", + "Usage", + "Teams", + "Internal Users", + "AI Hub", + "Response Cache", + ], }; /** Migrated pages live at a path route; legacy pages keep the ?page= query param. */ diff --git a/ui/litellm-dashboard/e2e_tests/tests/proxy-admin/keys.spec.ts b/tests/e2e/ui/tests/proxy-admin/keys.spec.ts similarity index 98% rename from ui/litellm-dashboard/e2e_tests/tests/proxy-admin/keys.spec.ts rename to tests/e2e/ui/tests/proxy-admin/keys.spec.ts index a55c19a53de..c44957ea737 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/proxy-admin/keys.spec.ts +++ b/tests/e2e/ui/tests/proxy-admin/keys.spec.ts @@ -103,7 +103,8 @@ test.describe("Proxy Admin - Keys", () => { await expect(page.getByText("Back to Keys")).toBeVisible({ timeout: 10_000 }); - await page.getByRole("button", { name: "Delete Key" }).click(); + await page.getByRole("button", { name: "More key actions" }).click(); + await page.getByRole("menuitem", { name: "Delete Key" }).click(); const modal = page.locator(".ant-modal:visible"); await expect(modal).toBeVisible({ timeout: 5_000 }); diff --git a/ui/litellm-dashboard/e2e_tests/tests/proxy-admin/license.spec.ts b/tests/e2e/ui/tests/proxy-admin/license.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/proxy-admin/license.spec.ts rename to tests/e2e/ui/tests/proxy-admin/license.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/proxy-admin/teams.spec.ts b/tests/e2e/ui/tests/proxy-admin/teams.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/proxy-admin/teams.spec.ts rename to tests/e2e/ui/tests/proxy-admin/teams.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/settings/adminSettings.spec.ts b/tests/e2e/ui/tests/settings/adminSettings.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/settings/adminSettings.spec.ts rename to tests/e2e/ui/tests/settings/adminSettings.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/settings/routerSettings.spec.ts b/tests/e2e/ui/tests/settings/routerSettings.spec.ts similarity index 98% rename from ui/litellm-dashboard/e2e_tests/tests/settings/routerSettings.spec.ts rename to tests/e2e/ui/tests/settings/routerSettings.spec.ts index 3e140b9ab56..ffa5f2c2ae2 100644 --- a/ui/litellm-dashboard/e2e_tests/tests/settings/routerSettings.spec.ts +++ b/tests/e2e/ui/tests/settings/routerSettings.spec.ts @@ -6,7 +6,7 @@ import { Role, users } from "../../fixtures/users"; // Type-only import of the OpenAPI-generated backend schema, erased at runtime by // esbuild. It types the round-trips below so mistakes surface in the editor; the live // test against the real proxy is what actually enforces the contract. -import type { components } from "../../../src/lib/http/schema"; +import type { components } from "../../../../../ui/litellm-dashboard/src/lib/http/schema"; // These tests mutate the proxy's shared router_settings, and the Loadbalancing save // echoes the whole settings object, so they must not run concurrently. diff --git a/ui/litellm-dashboard/e2e_tests/tests/team-admin/teamAdmin.spec.ts b/tests/e2e/ui/tests/team-admin/teamAdmin.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/team-admin/teamAdmin.spec.ts rename to tests/e2e/ui/tests/team-admin/teamAdmin.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/users/searchUsers.spec.ts b/tests/e2e/ui/tests/users/searchUsers.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/users/searchUsers.spec.ts rename to tests/e2e/ui/tests/users/searchUsers.spec.ts diff --git a/ui/litellm-dashboard/e2e_tests/tests/users/viewInternalUsers.spec.ts b/tests/e2e/ui/tests/users/viewInternalUsers.spec.ts similarity index 100% rename from ui/litellm-dashboard/e2e_tests/tests/users/viewInternalUsers.spec.ts rename to tests/e2e/ui/tests/users/viewInternalUsers.spec.ts diff --git a/tests/e2e/ui/tsconfig.json b/tests/e2e/ui/tsconfig.json new file mode 100644 index 00000000000..f9290fe7b49 --- /dev/null +++ b/tests/e2e/ui/tsconfig.json @@ -0,0 +1,16 @@ +{ + "compilerOptions": { + "target": "ES2022", + "module": "commonjs", + "moduleResolution": "node", + "lib": ["ES2022", "DOM"], + "strict": true, + "esModuleInterop": true, + "skipLibCheck": true, + "resolveJsonModule": true, + "noEmit": true, + "types": ["node"] + }, + "include": ["**/*.ts"], + "exclude": ["node_modules"] +} diff --git a/tests/litellm/llms/anthropic/test_anthropic_schema_filter.py b/tests/litellm/llms/anthropic/test_anthropic_schema_filter.py index bd3b4198e9e..c10ac5532a0 100644 --- a/tests/litellm/llms/anthropic/test_anthropic_schema_filter.py +++ b/tests/litellm/llms/anthropic/test_anthropic_schema_filter.py @@ -45,21 +45,14 @@ class TestFilterAnthropicOutputSchema: assert "minimum value: 0" in result["properties"]["age"]["description"] assert "maximum value: 150" in result["properties"]["age"]["description"] # Score had no description, should get one from constraints - assert ( - "exclusive minimum value: 0" in result["properties"]["score"]["description"] - ) - assert ( - "exclusive maximum value: 100" - in result["properties"]["score"]["description"] - ) + assert "exclusive minimum value: 0" in result["properties"]["score"]["description"] + assert "exclusive maximum value: 100" in result["properties"]["score"]["description"] def test_removes_string_constraints(self): """Test that minLength/maxLength are removed from string schemas.""" schema = { "type": "object", - "properties": { - "name": {"type": "string", "minLength": 1, "maxLength": 100} - }, + "properties": {"name": {"type": "string", "minLength": 1, "maxLength": 100}}, } result = AnthropicConfig.filter_anthropic_output_schema(schema) @@ -154,3 +147,203 @@ class TestFilterAnthropicOutputSchema: result = AnthropicConfig.filter_anthropic_output_schema(schema) assert result == schema # Should be unchanged + + def test_removes_uniqueitems(self): + """Test that uniqueItems is removed from array schemas. + + Reproduces the 400 ``invalid_request_error``: + "output_format.schema: For 'array' type, property 'uniqueItems' is not + supported". + """ + schema = { + "type": "object", + "properties": { + "tags": { + "type": "array", + "items": {"type": "string"}, + "uniqueItems": True, + } + }, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "uniqueItems" not in result["properties"]["tags"] + assert result["properties"]["tags"]["items"] == {"type": "string"} + # Constraint intent preserved in the description + assert "all array items must be unique" in result["properties"]["tags"]["description"] + + def test_removes_contains_constraints(self): + """Test that contains/minContains/maxContains are removed from arrays.""" + schema = { + "type": "array", + "items": {"type": "integer"}, + "contains": {"type": "integer", "const": 1}, + "minContains": 1, + "maxContains": 3, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "contains" not in result + assert "minContains" not in result + assert "maxContains" not in result + assert result["items"] == {"type": "integer"} + # The contains sub-schema is serialized into the advisory note so the model + # knows what item the array must contain. + assert "array must contain an item matching:" in result["description"] + assert '"const": 1' in result["description"] + assert "minimum number of matching items: 1" in result["description"] + assert "maximum number of matching items: 3" in result["description"] + + def test_removes_object_property_constraints(self): + """Test that minProperties/maxProperties are removed from object schemas.""" + schema = { + "type": "object", + "properties": {"a": {"type": "string"}}, + "minProperties": 1, + "maxProperties": 5, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "minProperties" not in result + assert "maxProperties" not in result + assert "minimum number of properties: 1" in result["description"] + assert "maximum number of properties: 5" in result["description"] + + def test_uniqueitems_false_skips_misleading_note(self): + """``uniqueItems: false`` is stripped but must not add a 'unique' note.""" + schema = { + "type": "array", + "items": {"type": "string"}, + "uniqueItems": False, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "uniqueItems" not in result + # A disabled constraint imposes no requirement -> no advisory note + assert "unique" not in result.get("description", "") + + def test_removes_multipleof(self): + """multipleOf is rejected by Anthropic for integer and number types.""" + schema = { + "type": "object", + "properties": {"n": {"type": "integer", "multipleOf": 5}}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "multipleOf" not in result["properties"]["n"] + assert "must be a multiple of 5" in result["properties"]["n"]["description"] + + def test_removes_conditional_and_negation_keywords(self): + """if/then/else and not are rejected by Anthropic and stripped into notes.""" + schema = { + "type": "object", + "properties": {"kind": {"type": "string"}, "sound": {"type": "string", "not": {"const": "moo"}}}, + "if": {"properties": {"kind": {"const": "dog"}}}, + "then": {"required": ["sound"]}, + "else": {"required": ["kind"]}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "if" not in result + assert "then" not in result + assert "else" not in result + assert "not" not in result["properties"]["sound"] + assert 'conditional (if): {"properties": {"kind": {"const": "dog"}}}' in result["description"] + assert 'conditional (then): {"required": ["sound"]}' in result["description"] + assert 'conditional (else): {"required": ["kind"]}' in result["description"] + assert 'must not match: {"const": "moo"}' in result["properties"]["sound"]["description"] + + def test_removes_object_shape_keywords(self): + """patternProperties/propertyNames/dependent*/unevaluatedProperties are stripped.""" + schema = { + "type": "object", + "properties": {"first": {"type": "string"}}, + "patternProperties": {"^x": {"type": "string"}}, + "propertyNames": {"pattern": "^[a-z]+$"}, + "dependentRequired": {"first": ["last"]}, + "dependentSchemas": {"first": {"required": ["last"]}}, + "unevaluatedProperties": {"type": "string"}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + for field in ( + "patternProperties", + "propertyNames", + "dependentRequired", + "dependentSchemas", + "unevaluatedProperties", + ): + assert field not in result + assert 'properties whose names match each pattern must satisfy: {"^x": {"type": "string"}}' in result["description"] + assert 'property names must satisfy: {"pattern": "^[a-z]+$"}' in result["description"] + assert 'dependent required properties: {"first": ["last"]}' in result["description"] + assert 'dependent schemas: {"first": {"required": ["last"]}}' in result["description"] + assert 'unevaluated properties must satisfy: {"type": "string"}' in result["description"] + + def test_removes_prefixitems(self): + """prefixItems is rejected by Anthropic for array types.""" + schema = { + "type": "array", + "prefixItems": [{"type": "number"}, {"type": "string"}], + "items": {"type": "number"}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "prefixItems" not in result + assert result["items"] == {"type": "number"} + assert 'leading items must match, in order: [{"type": "number"}, {"type": "string"}]' in result["description"] + + def test_oneof_rewritten_to_anyof(self): + """oneOf 400s ("Schema type 'oneOf' is not supported") and becomes anyOf, like the SDK.""" + schema = { + "type": "object", + "properties": {"id": {"oneOf": [{"type": "string", "minLength": 1}, {"type": "integer"}]}}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + id_schema = result["properties"]["id"] + assert "oneOf" not in id_schema + assert [v["type"] for v in id_schema["anyOf"]] == ["string", "integer"] + assert "minLength" not in id_schema["anyOf"][0] + assert "minimum length: 1" in id_schema["anyOf"][0]["description"] + + def test_oneof_merges_into_existing_anyof(self): + schema = { + "anyOf": [{"type": "string"}], + "oneOf": [{"type": "integer"}], + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "oneOf" not in result + assert [v["type"] for v in result["anyOf"]] == ["string", "integer"] + + def test_constraint_note_order_is_deterministic(self): + """Note order must not depend on set iteration order (PYTHONHASHSEED), or the + serialized request differs across proxy workers and breaks caching.""" + schema = { + "type": "array", + "items": {"type": "string"}, + "minItems": 1, + "maxItems": 10, + "uniqueItems": True, + "minContains": 2, + "maxContains": 3, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert result["description"] == ( + "Note: minimum number of items: 1, maximum number of items: 10, " + "all array items must be unique, minimum number of matching items: 2, " + "maximum number of matching items: 3." + ) diff --git a/tests/litellm_utils_tests/test_proxy_budget_reset.py b/tests/litellm_utils_tests/test_proxy_budget_reset.py index 5c96eb619bf..44da3ea06a0 100644 --- a/tests/litellm_utils_tests/test_proxy_budget_reset.py +++ b/tests/litellm_utils_tests/test_proxy_budget_reset.py @@ -30,6 +30,7 @@ def _attrify(d: dict): None)` (et al), which returns None for plain dicts — that would silently skip the row. """ + class _AttrDict(dict): def __getattr__(self, k): try: @@ -120,9 +121,11 @@ async def test_reset_budget_keys_partial_failure(): key1, key2, key3, key4, key5, key6 = ( _attrify(k) for k in [key1, key2, key3, key4, key5, key6] ) - prisma_client.get_data = AsyncMock(return_value=[key1, key2, key3, key4, key5, key6]) + prisma_client.get_data = AsyncMock( + return_value=[key1, key2, key3, key4, key5, key6] + ) - async def fake_reset_key(key, current_time): + async def fake_reset_key(key, current_time, reset_settings=None): if key["id"] == "key1": # Simulate a failure on key1 (for example, this might be due to an invariant check) raise Exception("Simulated failure for key1") @@ -207,9 +210,11 @@ async def test_reset_budget_users_partial_failure(): user1, user2, user3, user4, user5, user6 = ( _attrify(u) for u in [user1, user2, user3, user4, user5, user6] ) - prisma_client.get_data = AsyncMock(return_value=[user1, user2, user3, user4, user5, user6]) + prisma_client.get_data = AsyncMock( + return_value=[user1, user2, user3, user4, user5, user6] + ) - async def fake_reset_user(user, current_time): + async def fake_reset_user(user, current_time, reset_settings=None): if user["id"] == "user1": raise Exception("Simulated failure for user1") else: @@ -397,7 +402,7 @@ async def test_reset_budget_teams_partial_failure(): team1, team2 = _attrify(team1), _attrify(team2) prisma_client.get_data = AsyncMock(return_value=[team1, team2]) - async def fake_reset_team(team, current_time): + async def fake_reset_team(team, current_time, reset_settings=None): if team["id"] == "team1": raise Exception("Simulated failure for team1") else: @@ -513,14 +518,14 @@ async def test_reset_budget_continues_other_categories_on_failure(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_key(key, current_time): + async def fake_reset_key(key, current_time, reset_settings=None): key["spend"] = 0.0 key["budget_reset_at"] = ( current_time + timedelta(seconds=key["budget_duration"]) ).isoformat() return key - async def fake_reset_user(user, current_time): + async def fake_reset_user(user, current_time, reset_settings=None): if user["id"] == "user1": raise Exception("Simulated failure for user1") user["spend"] = 0.0 @@ -529,7 +534,7 @@ async def test_reset_budget_continues_other_categories_on_failure(): ).isoformat() return user - async def fake_reset_team(team, current_time): + async def fake_reset_team(team, current_time, reset_settings=None): team["spend"] = 0.0 team["budget_reset_at"] = ( current_time + timedelta(seconds=team["budget_duration"]) @@ -632,7 +637,7 @@ async def test_service_logger_keys_success(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_key(key, current_time): + async def fake_reset_key(key, current_time, reset_settings=None): key["spend"] = 0.0 key["budget_reset_at"] = ( current_time + timedelta(seconds=key["budget_duration"]) @@ -688,7 +693,7 @@ async def test_service_logger_keys_failure(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_key(key, current_time): + async def fake_reset_key(key, current_time, reset_settings=None): if key["id"] == "key1": raise Exception("Simulated failure for key1") key["spend"] = 0.0 @@ -750,7 +755,7 @@ async def test_service_logger_users_success(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_user(user, current_time): + async def fake_reset_user(user, current_time, reset_settings=None): user["spend"] = 0.0 user["budget_reset_at"] = ( current_time + timedelta(seconds=user["budget_duration"]) @@ -802,7 +807,7 @@ async def test_service_logger_users_failure(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_user(user, current_time): + async def fake_reset_user(user, current_time, reset_settings=None): if user["id"] == "user1": raise Exception("Simulated failure for user1") user["spend"] = 0.0 @@ -863,7 +868,7 @@ async def test_service_logger_teams_success(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_team(team, current_time): + async def fake_reset_team(team, current_time, reset_settings=None): team["spend"] = 0.0 team["budget_reset_at"] = ( current_time + timedelta(seconds=team["budget_duration"]) @@ -915,7 +920,7 @@ async def test_service_logger_teams_failure(): job = ResetBudgetJob(proxy_logging_obj, prisma_client) - async def fake_reset_team(team, current_time): + async def fake_reset_team(team, current_time, reset_settings=None): if team["id"] == "team1": raise Exception("Simulated failure for team1") team["spend"] = 0.0 diff --git a/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json b/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json index a4c50d3c575..41b4c3efb63 100644 --- a/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json +++ b/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json @@ -11,7 +11,7 @@ "user": "", "team_id": "", "organization_id": "", - "metadata": "{\"applied_guardrails\": [], \"batch_models\": null, \"mcp_tool_call_metadata\": null, \"vector_store_request_metadata\": null, \"guardrail_information\": null, \"usage_object\": {\"completion_tokens\": 20, \"prompt_tokens\": 10, \"total_tokens\": 30, \"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"model_map_information\": {\"model_map_key\": \"gpt-4o\", \"model_map_value\": {\"key\": \"gpt-4o\", \"max_tokens\": 16384, \"max_input_tokens\": 128000, \"max_output_tokens\": 16384, \"input_cost_per_token\": 2.5e-06, \"cache_creation_input_token_cost\": null, \"cache_read_input_token_cost\": 1.25e-06, \"input_cost_per_character\": null, \"input_cost_per_token_above_128k_tokens\": null, \"input_cost_per_token_above_200k_tokens\": null, \"input_cost_per_query\": null, \"input_cost_per_second\": null, \"input_cost_per_audio_token\": null, \"input_cost_per_token_batches\": 1.25e-06, \"output_cost_per_token_batches\": 5e-06, \"output_cost_per_token\": 1e-05, \"output_cost_per_audio_token\": null, \"output_cost_per_character\": null, \"output_cost_per_token_above_128k_tokens\": null, \"output_cost_per_character_above_128k_tokens\": null, \"output_cost_per_token_above_200k_tokens\": null, \"output_cost_per_second\": null, \"output_cost_per_image\": null, \"output_vector_size\": null, \"litellm_provider\": \"openai\", \"mode\": \"chat\", \"supports_system_messages\": true, \"supports_response_schema\": true, \"supports_vision\": true, \"supports_function_calling\": true, \"supports_tool_choice\": true, \"supports_assistant_prefill\": false, \"supports_prompt_caching\": true, \"supports_audio_input\": false, \"supports_audio_output\": false, \"supports_pdf_input\": false, \"supports_embedding_image_input\": false, \"supports_native_streaming\": null, \"supports_web_search\": true, \"supports_reasoning\": false, \"search_context_cost_per_query\": {\"search_context_size_low\": 0.03, \"search_context_size_medium\": 0.035, \"search_context_size_high\": 0.05}, \"tpm\": null, \"rpm\": null, \"supported_openai_params\": [\"frequency_penalty\", \"logit_bias\", \"logprobs\", \"top_logprobs\", \"max_tokens\", \"max_completion_tokens\", \"modalities\", \"prediction\", \"n\", \"presence_penalty\", \"seed\", \"stop\", \"stream\", \"stream_options\", \"temperature\", \"top_p\", \"tools\", \"tool_choice\", \"function_call\", \"functions\", \"max_retries\", \"extra_headers\", \"parallel_tool_calls\", \"audio\", \"response_format\", \"user\"]}}, \"additional_usage_values\": {\"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"user_api_key\": null, \"user_api_key_alias\": null, \"user_api_key_team_id\": null, \"user_api_key_project_id\": null, \"user_api_key_project_alias\": null, \"user_api_key_org_id\": null, \"user_api_key_user_id\": null, \"user_api_key_team_alias\": null, \"spend_logs_metadata\": null, \"requester_ip_address\": null, \"status\": null, \"proxy_server_request\": null, \"error_information\": null, \"attempted_retries\": null, \"max_retries\": null}", + "metadata": "{\"applied_guardrails\": [], \"batch_models\": null, \"mcp_tool_call_metadata\": null, \"vector_store_request_metadata\": null, \"guardrail_information\": null, \"compression_savings\": null, \"usage_object\": {\"completion_tokens\": 20, \"prompt_tokens\": 10, \"total_tokens\": 30, \"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"model_map_information\": {\"model_map_key\": \"gpt-4o\", \"model_map_value\": {\"key\": \"gpt-4o\", \"max_tokens\": 16384, \"max_input_tokens\": 128000, \"max_output_tokens\": 16384, \"input_cost_per_token\": 2.5e-06, \"cache_creation_input_token_cost\": null, \"cache_read_input_token_cost\": 1.25e-06, \"input_cost_per_character\": null, \"input_cost_per_token_above_128k_tokens\": null, \"input_cost_per_token_above_200k_tokens\": null, \"input_cost_per_query\": null, \"input_cost_per_second\": null, \"input_cost_per_audio_token\": null, \"input_cost_per_token_batches\": 1.25e-06, \"output_cost_per_token_batches\": 5e-06, \"output_cost_per_token\": 1e-05, \"output_cost_per_audio_token\": null, \"output_cost_per_character\": null, \"output_cost_per_token_above_128k_tokens\": null, \"output_cost_per_character_above_128k_tokens\": null, \"output_cost_per_token_above_200k_tokens\": null, \"output_cost_per_second\": null, \"output_cost_per_image\": null, \"output_vector_size\": null, \"litellm_provider\": \"openai\", \"mode\": \"chat\", \"supports_system_messages\": true, \"supports_response_schema\": true, \"supports_vision\": true, \"supports_function_calling\": true, \"supports_tool_choice\": true, \"supports_assistant_prefill\": false, \"supports_prompt_caching\": true, \"supports_audio_input\": false, \"supports_audio_output\": false, \"supports_pdf_input\": false, \"supports_embedding_image_input\": false, \"supports_native_streaming\": null, \"supports_web_search\": true, \"supports_reasoning\": false, \"search_context_cost_per_query\": {\"search_context_size_low\": 0.03, \"search_context_size_medium\": 0.035, \"search_context_size_high\": 0.05}, \"tpm\": null, \"rpm\": null, \"supported_openai_params\": [\"frequency_penalty\", \"logit_bias\", \"logprobs\", \"top_logprobs\", \"max_tokens\", \"max_completion_tokens\", \"modalities\", \"prediction\", \"n\", \"presence_penalty\", \"seed\", \"stop\", \"stream\", \"stream_options\", \"temperature\", \"top_p\", \"tools\", \"tool_choice\", \"function_call\", \"functions\", \"max_retries\", \"extra_headers\", \"parallel_tool_calls\", \"audio\", \"response_format\", \"user\"]}}, \"additional_usage_values\": {\"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"user_api_key\": null, \"user_api_key_alias\": null, \"user_api_key_team_id\": null, \"user_api_key_project_id\": null, \"user_api_key_project_alias\": null, \"user_api_key_org_id\": null, \"user_api_key_user_id\": null, \"user_api_key_team_alias\": null, \"spend_logs_metadata\": null, \"requester_ip_address\": null, \"status\": null, \"proxy_server_request\": null, \"error_information\": null, \"attempted_retries\": null, \"max_retries\": null}", "cache_key": "Cache OFF", "spend": 0.00022500000000000002, "total_tokens": 30, diff --git a/tests/proxy_migration_tests/test_offline_image_migration.py b/tests/proxy_migration_tests/test_offline_image_migration.py new file mode 100644 index 00000000000..8ed88278e69 --- /dev/null +++ b/tests/proxy_migration_tests/test_offline_image_migration.py @@ -0,0 +1,143 @@ +"""Image-level regression net for the prisma bake in the shipped runtime image. + +Boots a built image's migration entrypoint the way an OpenShift / air-gapped +deployment does (an internal-only network with no egress, an arbitrary non-root +uid in GID 0) against a brand-new Postgres, and asserts the schema was created. + +This catches the whole failure class, not one symptom: a bake that only works +under `docker run` as the default uid with network still passes every existing +check, because the migration entrypoint exits 0 even when it applied nothing. +Asserting the table count is what turns that silent success into a hard fail. + +Gated on LITELLM_IMAGE (the tag of the image to exercise) so it is skipped in +the normal unit-test run and exercised only where an image has been built (the +image-scan workflow). Requires a working docker CLI. +""" + +import shutil +import subprocess +import uuid + +import os +import pytest + +IMAGE = os.getenv("LITELLM_IMAGE") +POSTGRES_IMAGE = os.getenv("LITELLM_TEST_POSTGRES_IMAGE", "postgres:16-alpine") +MIN_TABLES = int(os.getenv("LITELLM_TEST_MIN_TABLES", "20")) +NON_ROOT_UID = "12345:0" # arbitrary uid in GID 0, as OpenShift restricted-v2 assigns + +pytestmark = [ + pytest.mark.skipif(IMAGE is None, reason="requires a built image (set LITELLM_IMAGE)"), + pytest.mark.skipif(shutil.which("docker") is None, reason="requires the docker CLI"), +] + + +def _docker(*args: str, check: bool = True) -> subprocess.CompletedProcess: + return subprocess.run( + ["docker", *args], capture_output=True, text=True, check=check + ) + + +@pytest.fixture() +def offline_postgres(): + """A fresh Postgres reachable only over an internal-only (no egress) network. + + Yields (network_name, postgres_host). Both are torn down afterwards. + """ + run_id = f"offlinemig-{uuid.uuid4().hex[:8]}" + network = f"{run_id}-net" + pg = f"{run_id}-pg" + + # Pull Postgres while egress still exists; the internal network below has none. + _docker("pull", "--quiet", POSTGRES_IMAGE) + # --internal => containers on this network cannot reach the internet, so a + # prisma engine download (binaries.prisma.sh / npm) fails instead of masking + # a non-self-contained bake. + _docker("network", "create", "--internal", network) + try: + _docker( + "run", "-d", "--name", pg, "--network", network, + "-e", "POSTGRES_PASSWORD=pw", "-e", "POSTGRES_DB=litellm", + POSTGRES_IMAGE, + ) + _wait_until_ready(pg) + yield network, pg + finally: + _docker("rm", "-f", pg, check=False) + _docker("network", "rm", network, check=False) + + +def _wait_until_ready(pg: str, attempts: int = 60) -> None: + for _ in range(attempts): + running = _docker( + "ps", "--filter", f"name={pg}", "--filter", "status=running", + "--format", "{{.Names}}", check=False, + ).stdout + if pg not in running: + logs = _docker("logs", pg, check=False).stdout + _docker("logs", pg, check=False).stderr + pytest.fail(f"postgres container is not running:\n{logs}") + ready = _docker( + "exec", pg, "pg_isready", "-U", "postgres", "-d", "litellm", check=False + ) + if ready.returncode == 0: + return + subprocess.run(["sleep", "1"]) + pytest.fail(f"postgres never became ready after {attempts}s") + + +def _table_count(pg: str) -> int: + result = _docker( + "exec", pg, "psql", "-U", "postgres", "-d", "litellm", "-tAc", + "SELECT count(*) FROM information_schema.tables WHERE table_schema='public';", + ) + return int(result.stdout.strip() or "0") + + +def test_migration_offline_as_non_root_uid(offline_postgres): + """The migration entrypoint creates the full schema offline as an arbitrary uid. + + Reproduces the OpenShift / air-gapped failure: on the pre-fix image the + migration exits 0 having created 0 tables (every DB endpoint then 500s on + missing columns); a self-contained bake creates the full schema. + """ + network, pg = offline_postgres + assert IMAGE is not None + + migrate = _docker( + "run", "--rm", "--network", network, "--user", NON_ROOT_UID, + "-e", f"DATABASE_URL=postgresql://postgres:pw@{pg}:5432/litellm", + "-e", "LITELLM_MASTER_KEY=sk-offline-migration-test", + "-e", "DISABLE_SCHEMA_UPDATE=false", + "-w", "/app", "--entrypoint", "python", + IMAGE, "litellm/proxy/prisma_migration.py", + check=False, + ) + tables = _table_count(pg) + + assert migrate.returncode == 0, ( + f"migration entrypoint exited {migrate.returncode} offline as uid {NON_ROOT_UID}\n" + f"stdout:\n{migrate.stdout}\nstderr:\n{migrate.stderr}" + ) + assert tables >= MIN_TABLES, ( + f"only {tables} tables created (need >= {MIN_TABLES}) offline as uid {NON_ROOT_UID}. " + "The prisma bake is not self-contained: it needs a runtime download or a " + "writable HOME/cache, so OpenShift and air-gapped deployments start on an " + f"empty database.\nstdout:\n{migrate.stdout}\nstderr:\n{migrate.stderr}" + ) + + +def test_runtime_cache_env_not_read_only(): + """No runtime cache env var may point at the world-read-only /opt/prisma bake. + + /opt/prisma is baked `a+rX` (no write). Pointing XDG_CACHE_HOME (or any cache + var an XDG-aware library honours) there would deny writes for every uid, so + guard against a future edit reintroducing that. + """ + assert IMAGE is not None + env = _docker("run", "--rm", "--entrypoint", "env", IMAGE).stdout + offenders = [ + line for line in env.splitlines() + if line.startswith(("XDG_CACHE_HOME=", "XDG_DATA_HOME=", "HOME=")) + and line.split("=", 1)[1].startswith("/opt/prisma") + ] + assert not offenders, f"cache/home env points at the read-only bake: {offenders}" diff --git a/tests/proxy_unit_tests/test_proxy_server.py b/tests/proxy_unit_tests/test_proxy_server.py index 212f7772cad..bedd4dd1838 100644 --- a/tests/proxy_unit_tests/test_proxy_server.py +++ b/tests/proxy_unit_tests/test_proxy_server.py @@ -1252,6 +1252,17 @@ async def test_create_team_member_add(prisma_client, new_member_method): return_value=LiteLLM_TeamTableCachedObj(team_id="1234") ) + tx_mock = AsyncMock() + tx_mock.query_raw = AsyncMock(return_value=[{"members_with_roles": []}]) + tx_mock.litellm_teamtable = team_mock_client + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=tx_mock) + tx_cm.__aexit__ = AsyncMock(return_value=None) + original_tx = litellm.proxy.proxy_server.prisma_client.tx + litellm.proxy.proxy_server.prisma_client.tx = MagicMock( + return_value=tx_cm + ) + print(f"team_member_add_request={team_member_add_request}") await team_member_add( data=team_member_add_request, @@ -1273,6 +1284,7 @@ async def test_create_team_member_add(prisma_client, new_member_method): ) litellm.proxy.proxy_server.prisma_client.db.litellm_teamtable = original_val + litellm.proxy.proxy_server.prisma_client.tx = original_tx @pytest.mark.parametrize("team_member_role", ["admin", "user"]) @@ -1434,42 +1446,51 @@ async def test_create_team_member_add_team_admin( mock_litellm_usertable.find_unique = AsyncMock(return_value=None) team_mock_client = AsyncMock() - original_val = getattr( - litellm.proxy.proxy_server.prisma_client.db, "litellm_teamtable" - ) - litellm.proxy.proxy_server.prisma_client.db.litellm_teamtable = team_mock_client - team_mock_client.update = AsyncMock( return_value=LiteLLM_TeamTableCachedObj(team_id="1234") ) - try: - await team_member_add( - data=team_member_add_request, - user_api_key_dict=valid_token, + tx_mock = AsyncMock() + tx_mock.query_raw = AsyncMock(return_value=[{"members_with_roles": []}]) + tx_mock.litellm_teamtable = team_mock_client + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=tx_mock) + tx_cm.__aexit__ = AsyncMock(return_value=None) + + with ( + patch.object( + litellm.proxy.proxy_server.prisma_client.db, + "litellm_teamtable", + team_mock_client, + ), + patch.object( + litellm.proxy.proxy_server.prisma_client, + "tx", + MagicMock(return_value=tx_cm), + ), + ): + try: + await team_member_add( + data=team_member_add_request, + user_api_key_dict=valid_token, + ) + except HTTPException as e: + if user_role == "user": + assert e.status_code == 403 + return + else: + raise e + + mock_client.assert_called() + + assert ( + mock_client.call_args.kwargs["data"]["create"]["max_budget"] + == litellm.max_internal_user_budget + ) + assert ( + mock_client.call_args.kwargs["data"]["create"]["budget_duration"] + == litellm.internal_user_budget_duration ) - except HTTPException as e: - if user_role == "user": - assert e.status_code == 403 - return - else: - raise e - - mock_client.assert_called() - - print(f"mock_client.call_args: {mock_client.call_args}") - print("mock_client.call_args.kwargs: {}".format(mock_client.call_args.kwargs)) - - assert ( - mock_client.call_args.kwargs["data"]["create"]["max_budget"] - == litellm.max_internal_user_budget - ) - assert ( - mock_client.call_args.kwargs["data"]["create"]["budget_duration"] - == litellm.internal_user_budget_duration - ) - - litellm.proxy.proxy_server.prisma_client.db.litellm_teamtable = original_val @pytest.mark.asyncio diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index d22b343d843..ee18c96c393 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1780,7 +1780,10 @@ def test_update_key_budget_with_temp_budget_increase(): "temp_budget_expiry": expiry_in_isoformat, }, ) - assert _update_key_budget_with_temp_budget_increase(valid_token).max_budget == 200 + result = _update_key_budget_with_temp_budget_increase(valid_token) + assert result.max_budget == 200 + assert result is not valid_token + assert valid_token.max_budget == 100 @pytest.mark.asyncio diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index 5471d2668e4..59c4caefa33 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -1115,6 +1115,7 @@ async def test_jwt_non_admin_team_route_access(monkeypatch): "team_id": None, "team_object": None, "user_id": None, + "user_email": None, "user_object": None, "org_id": None, "org_object": None, diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index b745ca8eadf..fbd7e36e298 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -219,8 +219,7 @@ def _gate(**overrides): kwargs = { "custom_llm_provider": "azure_ai", "litellm_params": GenericLiteLLMParams(api_key="sk-azure", rust=True), - "stream": False, - "rust_stream_eligible": False, + "has_agentic_hook": False, "model": "claude-sonnet-4-5", "api_key": "sk-azure", "api_base": "https://resource.services.ai.azure.com/anthropic", @@ -345,11 +344,11 @@ async def test_gate_skips_rust_for_unsupported_provider(): @pytest.mark.asyncio -async def test_gate_skips_rust_when_streaming_but_not_eligible(): +async def test_gate_skips_rust_for_agentic_hook(): bridge = ExplodingAsyncMessages() litellm.use_litellm_rust(True, amessages=bridge) - response = await _gate(stream=True, rust_stream_eligible=False) + response = await _gate(has_agentic_hook=True) assert response is None assert bridge.calls == 0 @@ -362,8 +361,7 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): streaming_body = {**REQUEST_BODY, "stream": True} response = await _gate( - stream=True, - rust_stream_eligible=True, + has_agentic_hook=False, request_body=streaming_body, ) diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 9289dece83f..64813c1eda7 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -1716,3 +1716,201 @@ class TestApplyGuardrailStyleDeploymentDispatch: await guardrail.async_pre_call_deployment_hook(kwargs, CallTypes.acompletion) assert guardrail.apply_called is False + + +class TestOnlyScanNewMessages: + """Incremental guardrail scanning: only send text segments not already scanned this session.""" + + def _guardrail(self, **overrides): + params = dict(guardrail_name="test-guard", only_scan_new_messages=True) + params.update(overrides) + return CustomGuardrail(**params) + + def _cache(self): + from litellm.caching import DualCache + + return DualCache() + + @pytest.mark.asyncio + async def test_disabled_returns_none(self): + guardrail = self._guardrail(only_scan_new_messages=False) + result = await guardrail.filter_new_texts_for_session( + texts=["hi"], + request_data={"litellm_session_id": "s1"}, + cache=self._cache(), + ) + assert result is None + + @pytest.mark.asyncio + async def test_no_session_id_fails_safe_to_full_scan(self): + guardrail = self._guardrail() + result = await guardrail.filter_new_texts_for_session( + texts=["hi"], + request_data={"metadata": {}}, + cache=self._cache(), + ) + assert result is None + + @pytest.mark.asyncio + async def test_masking_guardrail_not_supported(self): + guardrail = self._guardrail(mask_request_content=True) + result = await guardrail.filter_new_texts_for_session( + texts=["hi"], + request_data={"litellm_session_id": "s1"}, + cache=self._cache(), + ) + assert result is None + + @pytest.mark.asyncio + async def test_cache_read_failure_fails_safe_to_full_scan(self): + from unittest.mock import AsyncMock + + guardrail = self._guardrail() + cache = self._cache() + cache.async_get_cache = AsyncMock(side_effect=RuntimeError("redis down")) + result = await guardrail.filter_new_texts_for_session( + texts=["hi"], + request_data={"litellm_session_id": "s1"}, + cache=cache, + ) + assert result is None + + @pytest.mark.asyncio + async def test_dedupes_previously_scanned_texts(self): + guardrail = self._guardrail() + cache = self._cache() + request = {"litellm_session_id": "sess-dedupe"} + turn1 = ["you are helpful", "first question"] + + first = await guardrail.filter_new_texts_for_session(texts=turn1, request_data=request, cache=cache) + assert first == turn1 + await guardrail.mark_texts_scanned(texts=turn1, request_data=request, cache=cache) + + turn2 = turn1 + ["an answer", "second question"] + second = await guardrail.filter_new_texts_for_session(texts=turn2, request_data=request, cache=cache) + assert second == ["an answer", "second question"] + + @pytest.mark.asyncio + async def test_no_new_texts_returns_empty(self): + guardrail = self._guardrail() + cache = self._cache() + request = {"litellm_session_id": "sess-empty"} + texts = ["only message"] + + await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + await guardrail.mark_texts_scanned(texts=texts, request_data=request, cache=cache) + + again = await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + assert again == [] + + @pytest.mark.asyncio + async def test_modified_earlier_text_is_rescanned(self): + guardrail = self._guardrail() + cache = self._cache() + request = {"litellm_session_id": "sess-edit"} + original = ["original"] + + await guardrail.filter_new_texts_for_session(texts=original, request_data=request, cache=cache) + await guardrail.mark_texts_scanned(texts=original, request_data=request, cache=cache) + + edited = ["original EDITED"] + result = await guardrail.filter_new_texts_for_session(texts=edited, request_data=request, cache=cache) + assert result == edited + + @pytest.mark.asyncio + async def test_blocked_scan_does_not_persist_hashes(self): + guardrail = self._guardrail() + cache = self._cache() + request = {"litellm_session_id": "sess-blocked"} + texts = ["please block me"] + + filtered = await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + assert filtered == texts + + again = await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + assert again == texts + + @pytest.mark.asyncio + async def test_scanned_hashes_written_with_fixed_ttl(self): + from unittest.mock import AsyncMock + + from litellm.constants import GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS + + guardrail = self._guardrail() + cache = self._cache() + cache.async_set_cache = AsyncMock() + request = {"litellm_session_id": "sess-ttl"} + + await guardrail.mark_texts_scanned(texts=["a", "b"], request_data=request, cache=cache) + + cache.async_set_cache.assert_awaited_once() + assert cache.async_set_cache.await_args.kwargs["ttl"] == GUARDRAIL_SCANNED_MESSAGES_CACHE_TTL_SECONDS + + @pytest.mark.asyncio + async def test_session_id_from_metadata_is_used_for_dedupe(self): + guardrail = self._guardrail() + cache = self._cache() + request = {"metadata": {"session_id": "sess-meta"}} + texts = ["shared message"] + + await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + await guardrail.mark_texts_scanned(texts=texts, request_data=request, cache=cache) + + again = await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + assert again == [] + + @pytest.mark.asyncio + async def test_session_id_from_litellm_metadata_is_used_for_dedupe(self): + guardrail = self._guardrail() + cache = self._cache() + request = {"litellm_metadata": {"session_id": "sess-lmeta"}} + texts = ["shared message"] + + await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + await guardrail.mark_texts_scanned(texts=texts, request_data=request, cache=cache) + + again = await guardrail.filter_new_texts_for_session(texts=texts, request_data=request, cache=cache) + assert again == [] + + @pytest.mark.asyncio + async def test_mark_texts_scanned_disabled_does_not_persist(self): + from unittest.mock import AsyncMock + + guardrail = self._guardrail(only_scan_new_messages=False) + cache = self._cache() + cache.async_set_cache = AsyncMock() + + await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache) + cache.async_set_cache.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mark_texts_scanned_masking_does_not_persist(self): + from unittest.mock import AsyncMock + + guardrail = self._guardrail(mask_request_content=True) + cache = self._cache() + cache.async_set_cache = AsyncMock() + + await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache) + cache.async_set_cache.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mark_texts_scanned_without_session_does_not_persist(self): + from unittest.mock import AsyncMock + + guardrail = self._guardrail() + cache = self._cache() + cache.async_set_cache = AsyncMock() + + await guardrail.mark_texts_scanned(texts=["a"], request_data={"metadata": {}}, cache=cache) + cache.async_set_cache.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mark_texts_scanned_survives_cache_write_failure(self): + from unittest.mock import AsyncMock + + guardrail = self._guardrail() + cache = self._cache() + cache.async_set_cache = AsyncMock(side_effect=RuntimeError("redis down")) + + await guardrail.mark_texts_scanned(texts=["a"], request_data={"litellm_session_id": "s1"}, cache=cache) diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 209e99895db..11b08fa45a8 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -194,6 +194,7 @@ class TestResponseCompliance: "cancelled", "incomplete", "budget_exceeded", + "queued", ] assert status_prop["enum"] == expected_statuses print(f"✓ Status enum values: {expected_statuses}") 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 b156faf3ea6..9ff67a82f40 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 @@ -2237,3 +2237,75 @@ def test_token_type_cost_breakdown_applies_regional_uplift(): text_input_cost = 600 * model_info["input_cost_per_token"] * uplift assert text_output_cost + eu.reasoning_cost == pytest.approx(completion_cost) assert text_input_cost + eu.cache_read_cost == pytest.approx(prompt_cost) + + +GEMINI_DAY0_LAUNCH_PRICING = [ + ("gemini-3.6-flash", 1.5e-06, 7.5e-06, 1.5e-07), + ("gemini/gemini-3.6-flash", 1.5e-06, 7.5e-06, 1.5e-07), + ("vertex_ai/gemini-3.6-flash", 1.5e-06, 7.5e-06, 1.5e-07), + ("gemini-3.5-flash-lite", 3e-07, 2.5e-06, 3e-08), + ("gemini/gemini-3.5-flash-lite", 3e-07, 2.5e-06, 3e-08), + ("vertex_ai/gemini-3.5-flash-lite", 3e-07, 2.5e-06, 3e-08), +] + + +@pytest.mark.parametrize("model,input_cost,output_cost,cache_read_cost", GEMINI_DAY0_LAUNCH_PRICING) +def test_gemini_36_flash_and_35_flash_lite_launch_pricing(model, input_cost, output_cost, cache_read_cost): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model_cost_map = litellm.model_cost[model] + assert model_cost_map["input_cost_per_token"] == input_cost + assert model_cost_map["output_cost_per_token"] == output_cost + assert model_cost_map["output_cost_per_reasoning_token"] == output_cost + assert model_cost_map["cache_read_input_token_cost"] == cache_read_cost + assert model_cost_map["mode"] == "chat" + assert model_cost_map["supports_reasoning"] is True + assert model_cost_map["supports_function_calling"] is True + assert model_cost_map["max_input_tokens"] == 1048576 + + +def test_generic_cost_per_token_gemini_36_flash(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + usage = Usage( + prompt_tokens=1000, + completion_tokens=500, + total_tokens=1500, + completion_tokens_details=CompletionTokensDetailsWrapper( + reasoning_tokens=200, + text_tokens=300, + ), + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=1000), + ) + prompt_cost, completion_cost = generic_cost_per_token( + model="gemini-3.6-flash", + usage=usage, + custom_llm_provider="gemini", + ) + assert prompt_cost == pytest.approx(0.0015) + assert completion_cost == pytest.approx(0.00375) + + +def test_generic_cost_per_token_gemini_35_flash_lite(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + usage = Usage( + prompt_tokens=1000, + completion_tokens=500, + total_tokens=1500, + completion_tokens_details=CompletionTokensDetailsWrapper( + reasoning_tokens=200, + text_tokens=300, + ), + prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=1000), + ) + prompt_cost, completion_cost = generic_cost_per_token( + model="gemini-3.5-flash-lite", + usage=usage, + custom_llm_provider="gemini", + ) + assert prompt_cost == pytest.approx(0.0003) + assert completion_cost == pytest.approx(0.00125) diff --git a/tests/test_litellm/litellm_core_utils/test_duration_parser.py b/tests/test_litellm/litellm_core_utils/test_duration_parser.py index 3e4446c6672..cb9f273a0a7 100644 --- a/tests/test_litellm/litellm_core_utils/test_duration_parser.py +++ b/tests/test_litellm/litellm_core_utils/test_duration_parser.py @@ -1,8 +1,13 @@ import unittest -from datetime import datetime, timezone +from datetime import datetime, time, timezone +from unittest.mock import patch from zoneinfo import ZoneInfo -from litellm.litellm_core_utils.duration_parser import get_next_standardized_reset_time +import litellm.litellm_core_utils.duration_parser as duration_parser +from litellm.litellm_core_utils.duration_parser import ( + duration_in_seconds, + get_next_standardized_reset_time, +) class TestStandardizedResetTime(unittest.TestCase): @@ -199,5 +204,186 @@ class TestStandardizedResetTime(unittest.TestCase): self.assertEqual(result, expected) +class TestResetTimeOfDay(unittest.TestCase): + """A configurable reset_time_of_day shifts day/week/month resets off midnight.""" + + def test_daily_reset_before_offset_is_today(self): + now = datetime(2023, 5, 15, 8, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1d", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 15, 12, 0, 0, tzinfo=timezone.utc)) + + def test_daily_reset_after_offset_is_tomorrow(self): + now = datetime(2023, 5, 15, 14, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1d", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 16, 12, 0, 0, tzinfo=timezone.utc)) + + def test_daily_reset_exactly_at_offset_rolls_forward(self): + now = datetime(2023, 5, 15, 12, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1d", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 16, 12, 0, 0, tzinfo=timezone.utc)) + + def test_daily_reset_with_seconds_offset(self): + now = datetime(2023, 5, 15, 8, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1d", now, "UTC", reset_time_of_day=time(9, 30, 15) + ) + self.assertEqual(result, datetime(2023, 5, 15, 9, 30, 15, tzinfo=timezone.utc)) + + def test_offset_applies_in_configured_timezone(self): + # 2023-05-15 22:30 UTC == 2023-05-16 01:30 in Jerusalem (IDT, UTC+3), + # so the next noon-Jerusalem reset is 2023-05-16 12:00 IDT. + now = datetime(2023, 5, 15, 22, 30, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1d", now, "Asia/Jerusalem", reset_time_of_day=time(12, 0) + ) + jerusalem = result.astimezone(ZoneInfo("Asia/Jerusalem")) + self.assertEqual( + (jerusalem.year, jerusalem.month, jerusalem.day), (2023, 5, 16) + ) + self.assertEqual(jerusalem.hour, 12) + self.assertEqual(jerusalem.minute, 0) + + def test_weekly_reset_lands_on_monday_at_offset(self): + wednesday = datetime(2023, 5, 17, 15, 45, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "7d", wednesday, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 22, 12, 0, 0, tzinfo=timezone.utc)) + + def test_weekly_reset_today_is_monday_before_offset_is_today(self): + monday_morning = datetime(2023, 5, 22, 9, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "7d", monday_morning, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 22, 12, 0, 0, tzinfo=timezone.utc)) + + def test_weekly_reset_today_is_monday_after_offset_is_next_week(self): + monday_afternoon = datetime(2023, 5, 22, 15, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "7d", monday_afternoon, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 29, 12, 0, 0, tzinfo=timezone.utc)) + + def test_monthly_30d_lands_on_first_at_offset(self): + now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "30d", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 6, 1, 12, 0, 0, tzinfo=timezone.utc)) + + def test_monthly_1mo_today_is_first_before_offset_is_today(self): + now = datetime(2023, 5, 1, 9, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1mo", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 1, 12, 0, 0, tzinfo=timezone.utc)) + + def test_monthly_year_rollover_at_offset(self): + now = datetime(2023, 12, 15, 9, 0, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "1mo", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2024, 1, 1, 12, 0, 0, tzinfo=timezone.utc)) + + def test_custom_day_reset_applies_offset(self): + now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) + result = get_next_standardized_reset_time( + "3d", now, "UTC", reset_time_of_day=time(12, 0) + ) + self.assertEqual(result, datetime(2023, 5, 18, 12, 0, 0, tzinfo=timezone.utc)) + + def test_sub_day_durations_ignore_offset(self): + base = datetime(2023, 5, 15, 15, 20, 30, tzinfo=timezone.utc) + self.assertEqual( + get_next_standardized_reset_time( + "2h", base, "UTC", reset_time_of_day=time(12, 0) + ), + datetime(2023, 5, 15, 16, 0, 0, tzinfo=timezone.utc), + ) + self.assertEqual( + get_next_standardized_reset_time( + "30m", base, "UTC", reset_time_of_day=time(12, 0) + ), + datetime(2023, 5, 15, 15, 30, 0, tzinfo=timezone.utc), + ) + + def test_default_offset_is_midnight(self): + now = datetime(2023, 5, 15, 10, 30, 0, tzinfo=timezone.utc) + self.assertEqual( + get_next_standardized_reset_time("1d", now, "UTC"), + datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc), + ) + + +class TestWordFormBudgetDurations(unittest.TestCase): + """The Admin UI historically persisted word-form budget durations + (hourly/daily/weekly/monthly). They must resolve to their real interval + instead of silently collapsing to a next-midnight (daily) reset. + """ + + def test_word_forms_map_to_correct_reset_times(self): + base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc) + + self.assertEqual( + get_next_standardized_reset_time("hourly", base_time, "UTC"), + datetime(2023, 5, 17, 16, 0, 0, tzinfo=timezone.utc), + ) + self.assertEqual( + get_next_standardized_reset_time("daily", base_time, "UTC"), + datetime(2023, 5, 18, 0, 0, 0, tzinfo=timezone.utc), + ) + self.assertEqual( + get_next_standardized_reset_time("weekly", base_time, "UTC"), + datetime(2023, 5, 22, 0, 0, 0, tzinfo=timezone.utc), + ) + self.assertEqual( + get_next_standardized_reset_time("monthly", base_time, "UTC"), + datetime(2023, 6, 1, 0, 0, 0, tzinfo=timezone.utc), + ) + + def test_word_forms_are_not_all_collapsed_to_daily(self): + base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc) + results = { + word: get_next_standardized_reset_time(word, base_time, "UTC") + for word in ("hourly", "daily", "weekly", "monthly") + } + self.assertEqual(len(set(results.values())), len(results)) + + def test_word_forms_match_canonical_int_unit_forms(self): + base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc) + for word, canonical in (("hourly", "1h"), ("daily", "24h"), ("weekly", "7d"), ("monthly", "30d")): + self.assertEqual( + get_next_standardized_reset_time(word, base_time, "UTC"), + get_next_standardized_reset_time(canonical, base_time, "UTC"), + ) + + def test_word_forms_are_case_and_whitespace_insensitive(self): + base_time = datetime(2023, 5, 17, 15, 20, 30, tzinfo=timezone.utc) + self.assertEqual( + get_next_standardized_reset_time(" Monthly ", base_time, "UTC"), + datetime(2023, 6, 1, 0, 0, 0, tzinfo=timezone.utc), + ) + + def test_duration_in_seconds_accepts_word_forms(self): + self.assertEqual(duration_in_seconds("hourly"), 3600) + self.assertEqual(duration_in_seconds("daily"), 86400) + self.assertEqual(duration_in_seconds("weekly"), 604800) + self.assertEqual(duration_in_seconds("monthly"), 2592000) + + def test_invalid_duration_logs_warning_and_falls_back(self): + base_time = datetime(2023, 5, 15, 15, 0, 0, tzinfo=timezone.utc) + with patch.object(duration_parser.verbose_logger, "warning") as mock_warning: + result = get_next_standardized_reset_time("garbage", base_time, "UTC") + self.assertEqual(result, datetime(2023, 5, 16, 0, 0, 0, tzinfo=timezone.utc)) + mock_warning.assert_called_once() + self.assertIn("garbage", mock_warning.call_args.args) + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index c5422e0d70f..9cd1fbb59a6 100644 --- a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -373,3 +373,132 @@ class TestAnthropicMessagesHandlerToolInjection: if __name__ == "__main__": # Run the tests pytest.main([__file__, "-v"]) + + +class TestAnthropicMessagesIncrementalScan: + """PR #33278: only_scan_new_messages through the real /v1/messages translation + handler (the path Claude Code uses). Encodes the wire payloads observed in the + live validation against a real Bedrock guardrail. + """ + + def _bedrock_guardrail(self): + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + return BedrockGuardrail( + guardrail_name="bedrock-incremental-anthropic", + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + default_on=True, + only_scan_new_messages=True, + ) + + def _data(self, messages, session_id): + return { + "model": "claude-sonnet-4-5", + "messages": messages, + "system": "You are a helpful geography assistant.", + "litellm_session_id": session_id, + } + + @pytest.mark.asyncio + async def test_first_turn_scans_all_eligible_then_second_turn_scans_only_diff(self): + from unittest.mock import AsyncMock, patch + + handler = AnthropicMessagesHandler() + guardrail = self._bedrock_guardrail() + sid = "anth-sess-diff" + turn1 = [{"role": "user", "content": "What is the capital of France?"}] + turn2 = turn1 + [ + {"role": "assistant", "content": "Paris."}, + {"role": "user", "content": "What is the capital of Germany?"}, + ] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await handler.process_input_messages( + data=self._data(turn1, sid), guardrail_to_apply=guardrail + ) + assert mock_api.call_count == 1 + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ + "What is the capital of France?" + ] + mock_api.reset_mock() + await handler.process_input_messages( + data=self._data(turn2, sid), guardrail_to_apply=guardrail + ) + assert mock_api.call_count == 1 + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ + "Paris.", + "What is the capital of Germany?", + ] + + @pytest.mark.asyncio + async def test_identical_resend_makes_no_guardrail_call(self): + from unittest.mock import AsyncMock, patch + + handler = AnthropicMessagesHandler() + guardrail = self._bedrock_guardrail() + sid = "anth-sess-resend" + msgs = [ + {"role": "user", "content": "What is the capital of France?"}, + {"role": "assistant", "content": "Paris."}, + {"role": "user", "content": "What is the capital of Germany?"}, + ] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await handler.process_input_messages(data=self._data(msgs, sid), guardrail_to_apply=guardrail) + assert mock_api.call_count == 1 + mock_api.reset_mock() + await handler.process_input_messages(data=self._data(msgs, sid), guardrail_to_apply=guardrail) + mock_api.assert_not_called() + + @pytest.mark.asyncio + async def test_edited_history_message_is_rescanned(self): + from unittest.mock import AsyncMock, patch + + handler = AnthropicMessagesHandler() + guardrail = self._bedrock_guardrail() + sid = "anth-sess-edit" + msgs = [{"role": "user", "content": "What is the capital of France?"}] + edited = [{"role": "user", "content": "What is the capital and population of France?"}] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await handler.process_input_messages(data=self._data(msgs, sid), guardrail_to_apply=guardrail) + mock_api.reset_mock() + await handler.process_input_messages(data=self._data(edited, sid), guardrail_to_apply=guardrail) + assert mock_api.call_count == 1 + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ + "What is the capital and population of France?" + ] + + @pytest.mark.asyncio + async def test_mixed_text_and_tool_use_keeps_text_segments(self): + """A message carrying both text and a tool_use block must not lose its text. + (tool_use inputs and tool_result content are dropped from texts on the + anthropic input path today; that is pre-existing baseline behavior.)""" + from unittest.mock import AsyncMock, patch + + handler = AnthropicMessagesHandler() + guardrail = self._bedrock_guardrail() + sid = "anth-sess-tools" + msgs = [ + {"role": "user", "content": "Search for the weather in Paris"}, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me look that up for you."}, + {"type": "tool_use", "id": "toolu_1", "name": "search", "input": {"query": "canary-args"}}, + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": "canary-result"}], + }, + {"role": "user", "content": "Thanks, summarize the result."}, + ] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await handler.process_input_messages(data=self._data(msgs, sid), guardrail_to_apply=guardrail) + scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]] + assert "Let me look that up for you." in scanned, "text beside a tool_use must be scanned" + assert "Search for the weather in Paris" in scanned + assert "Thanks, summarize the result." in scanned diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_output_format_filter.py b/tests/test_litellm/llms/anthropic/test_anthropic_output_format_filter.py new file mode 100644 index 00000000000..90cba035760 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/test_anthropic_output_format_filter.py @@ -0,0 +1,147 @@ +""" +Coverage for filter_anthropic_output_schema's array/object constraint stripping. + +Mirrors tests/litellm/llms/anthropic/test_anthropic_schema_filter.py, but lives +under tests/test_litellm/ so the coverage-uploading CI job exercises the stripped +keyword handling (uniqueItems / contains / minProperties / maxProperties plus +multipleOf / patternProperties / propertyNames / dependentRequired / +dependentSchemas / unevaluatedProperties / if / then / else / not / prefixItems), +the ``uniqueItems: false`` branch, the oneOf to anyOf rewrite, and the +deterministic note ordering. +""" + +from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + +class TestOutputFormatArrayObjectConstraints: + def test_removes_uniqueitems(self): + schema = { + "type": "object", + "properties": { + "tags": { + "type": "array", + "items": {"type": "string"}, + "uniqueItems": True, + } + }, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "uniqueItems" not in result["properties"]["tags"] + assert "all array items must be unique" in result["properties"]["tags"]["description"] + + def test_uniqueitems_false_skips_misleading_note(self): + schema = { + "type": "array", + "items": {"type": "string"}, + "uniqueItems": False, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "uniqueItems" not in result + assert "unique" not in result.get("description", "") + + def test_removes_contains_constraints(self): + schema = { + "type": "array", + "items": {"type": "integer"}, + "contains": {"type": "integer", "const": 1}, + "minContains": 1, + "maxContains": 3, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "contains" not in result + assert "minContains" not in result + assert "maxContains" not in result + assert "array must contain an item matching:" in result["description"] + assert '"const": 1' in result["description"] + + def test_removes_object_property_constraints(self): + schema = { + "type": "object", + "properties": {"a": {"type": "string"}}, + "minProperties": 1, + "maxProperties": 5, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert "minProperties" not in result + assert "maxProperties" not in result + assert "minimum number of properties: 1" in result["description"] + assert "maximum number of properties: 5" in result["description"] + + def test_removes_remaining_rejected_keywords(self): + schema = { + "type": "object", + "properties": { + "n": {"type": "integer", "multipleOf": 5}, + "pair": {"type": "array", "prefixItems": [{"type": "number"}], "items": {"type": "number"}}, + "color": {"type": "string", "not": {"const": "red"}}, + }, + "patternProperties": {"^x": {"type": "string"}}, + "propertyNames": {"pattern": "^[a-z]+$"}, + "dependentRequired": {"n": ["pair"]}, + "dependentSchemas": {"n": {"required": ["pair"]}}, + "unevaluatedProperties": {"type": "string"}, + "if": {"properties": {"n": {"const": 5}}}, + "then": {"required": ["pair"]}, + "else": {"required": ["color"]}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + for field in ( + "patternProperties", + "propertyNames", + "dependentRequired", + "dependentSchemas", + "unevaluatedProperties", + "if", + "then", + "else", + ): + assert field not in result + assert "multipleOf" not in result["properties"]["n"] + assert "must be a multiple of 5" in result["properties"]["n"]["description"] + assert "prefixItems" not in result["properties"]["pair"] + assert 'leading items must match, in order: [{"type": "number"}]' in result["properties"]["pair"]["description"] + assert "not" not in result["properties"]["color"] + assert 'must not match: {"const": "red"}' in result["properties"]["color"]["description"] + assert 'conditional (if): {"properties": {"n": {"const": 5}}}' in result["description"] + + def test_oneof_rewritten_to_anyof(self): + schema = { + "type": "object", + "properties": {"id": {"oneOf": [{"type": "string", "minLength": 1}, {"type": "integer"}]}}, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + id_schema = result["properties"]["id"] + assert "oneOf" not in id_schema + assert [v["type"] for v in id_schema["anyOf"]] == ["string", "integer"] + assert "minLength" not in id_schema["anyOf"][0] + + def test_constraint_note_order_is_deterministic(self): + schema = { + "type": "array", + "items": {"type": "string"}, + "minItems": 1, + "maxItems": 10, + "uniqueItems": True, + "minContains": 2, + "maxContains": 3, + } + + result = AnthropicConfig.filter_anthropic_output_schema(schema) + + assert result["description"] == ( + "Note: minimum number of items: 1, maximum number of items: 10, " + "all array items must be unique, minimum number of matching items: 2, " + "maximum number of matching items: 3." + ) diff --git a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py index 25d24cfc3ac..1e1b98861b4 100644 --- a/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py +++ b/tests/test_litellm/llms/azure_ai/claude/test_azure_anthropic_messages_transformation.py @@ -412,7 +412,7 @@ class TestAzureAnthropicMidConversationSystem: older Claude, and a *leading* system entry 400s on every model ("messages.0: use the top-level 'system' parameter"). These tests pin the model-aware hoist the config applies so Claude Code sessions neither collapse the prompt cache - on 4.8+ nor hard-fail on 4.7 and older (RCA: Kraken Tech high-spend).""" + on 4.8+ nor hard-fail on 4.7 and older (RCA: customer high-spend).""" def test_supported_model_keeps_mid_conversation_system_in_place(self, local_model_cost_map): messages = [ diff --git a/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py b/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py index 448afd5f3a5..e5a2ea9b28f 100644 --- a/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py @@ -70,25 +70,33 @@ class TestAgentCoreAcceptHeader: """ End-to-end test: verify Accept header appears in the final HTTP request when using JWT auth through litellm.completion(). + + No exception swallowing: if completion() raises (for example because the + injected client was silently ignored and a real network call was made), + the test must fail with that error, not a misleading mock assertion. """ from litellm.llms.custom_httpx.http_handler import HTTPHandler client = HTTPHandler() - with patch.object(client, "post", return_value=MagicMock()) as mock_post: - try: - litellm.completion( - model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime", - messages=[{"role": "user", "content": "test"}], - api_key="test-jwt-token", - client=client, - ) - except Exception: - pass + mock_response = Mock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = { + "result": {"role": "assistant", "content": [{"text": "agent reply"}]} + } - mock_post.assert_called_once() - headers = mock_post.call_args.kwargs["headers"] - assert "Accept" in headers - assert headers["Accept"] == "application/json, text/event-stream" + with patch.object(client, "post", return_value=mock_response) as mock_post: + response = litellm.completion( + model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime", + messages=[{"role": "user", "content": "test"}], + api_key="test-jwt-token", + client=client, + ) + + mock_post.assert_called_once() + headers = mock_post.call_args.kwargs["headers"] + assert headers["Accept"] == "application/json, text/event-stream" + assert response.choices[0].message.content == "agent reply" class TestAgentCoreJsonResponseParsing: diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index fc12ead36a1..f832a4087ec 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -633,7 +633,8 @@ def test_parallel_tool_calls_config_kept_for_sonnet_5(): ) assert data["additionalModelRequestFields"]["tool_choice"] == { - "disable_parallel_tool_use": True + "type": "auto", + "disable_parallel_tool_use": True, } finally: litellm.model_cost = old_cost @@ -4251,6 +4252,49 @@ def test_parallel_tool_calls_older_model_drops_disable_flag(): assert "parallel_tool_calls" not in additional +@pytest.mark.parametrize( + "parallel_tool_calls, expected_disable", + [(True, False), (False, True)], +) +def test_parallel_tool_calls_emits_typed_auto_tool_choice(parallel_tool_calls, expected_disable): + config = AmazonConverseConfig() + model = "us.anthropic.claude-opus-4-8" + messages = [{"role": "user", "content": "What's the weather in SF and NYC?"}] + + optional_params = config.map_openai_params( + non_default_params={"parallel_tool_calls": parallel_tool_calls, "tools": _TOOL_PARAM}, + optional_params={}, + model=model, + drop_params=False, + ) + + request_data = config.transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert request_data["additionalModelRequestFields"]["tool_choice"] == { + "type": "auto", + "disable_parallel_tool_use": expected_disable, + } + + +def test_parallel_tool_use_merge_preserves_user_tool_choice_type(): + merged = AmazonConverseConfig._merge_parallel_tool_use_config( + {"tool_choice": {"type": "tool", "name": "get_weather", "disable_parallel_tool_use": False}}, + {"tool_choice": {"type": "auto", "disable_parallel_tool_use": True}}, + ) + + assert merged["tool_choice"] == { + "type": "tool", + "name": "get_weather", + "disable_parallel_tool_use": True, + } + + class TestBedrockMinThinkingBudgetTokens: """Test that thinking.budget_tokens is clamped to the Bedrock minimum (1024).""" diff --git a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py index ddc2e026e83..ffe21b91ab2 100644 --- a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py +++ b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_handler.py @@ -54,15 +54,24 @@ class FakeBedrockStream: self.input_stream = input_stream if input_stream is not None else FakeInputStream() +class FakeLogging: + def __init__(self, trace_id="trace-nova-sonic"): + self.litellm_trace_id = trace_id + + class DisconnectingClientWS: def __init__(self, messages): self._messages = list(messages) + self.sent_to_client = [] async def receive_text(self): if self._messages: return self._messages.pop(0) raise RuntimeError("client disconnected") + async def send_text(self, message): + self.sent_to_client.append(message) + class ClosableClientWS: def __init__(self): @@ -85,10 +94,14 @@ class EndedBedrockStream: class RealtimeClientWS: def __init__(self): self.closed = False + self.sent_to_client = [] async def receive_text(self): raise RuntimeError("client disconnected") + async def send_text(self, message): + self.sent_to_client.append(message) + async def close(self, code=None, reason=None): self.closed = True @@ -277,6 +290,61 @@ class TestBedrockRealtimeHandler: assert client_ws.closed +class TestBedrockRealtimeSessionLifecycle: + """Server must emit session.created on connect and session.updated on session.update (LIT-4655 regression)""" + + @pytest.mark.asyncio + async def test_session_created_sent_on_connect_before_any_client_input(self, stub_aws_sdk_client): + handler = BedrockRealtime() + websocket = RealtimeClientWS() + + await handler.async_realtime( + model="amazon.nova-sonic-v1:0", + websocket=websocket, + logging_obj=FakeLogging(), + aws_region_name="us-east-1", + aws_access_key_id="k", + aws_secret_access_key="s", + ) + + assert websocket.sent_to_client, "server sent nothing on connect: spec-conformant clients deadlock" + first_event = json.loads(websocket.sent_to_client[0]) + assert first_event["type"] == "session.created" + assert first_event["session"]["id"] == "trace-nova-sonic" + assert first_event["session"]["model"] == "amazon.nova-sonic-v1:0" + + @pytest.mark.asyncio + async def test_session_update_is_acked_with_session_updated(self, stub_aws_models): + handler = BedrockRealtime() + config = BedrockRealtimeConfig() + stream = FakeBedrockStream() + client_ws = DisconnectingClientWS( + [json.dumps({"type": "session.update", "session": {"instructions": "hi", "modalities": ["text"]}})] + ) + + await handler._forward_client_to_bedrock( + client_ws, stream, config, "amazon.nova-sonic-v1:0", {}, FakeLogging() + ) + + acked = [json.loads(message) for message in client_ws.sent_to_client] + updated = [event for event in acked if event["type"] == "session.updated"] + assert updated, "session.update was not acked" + assert updated[0]["session"]["modalities"] == ["text"], "ack must reflect the requested modalities" + + @pytest.mark.asyncio + async def test_no_session_updated_without_logging_obj(self, stub_aws_models): + handler = BedrockRealtime() + config = BedrockRealtimeConfig() + stream = FakeBedrockStream() + client_ws = DisconnectingClientWS( + [json.dumps({"type": "session.update", "session": {"instructions": "hi"}})] + ) + + await handler._forward_client_to_bedrock(client_ws, stream, config, "amazon.nova-sonic-v1:0", {}) + + assert client_ws.sent_to_client == [] + + class TestBedrockRealtimeAwsAuth: """AWS auth params passed via litellm_params must reach the Smithy client config (LIT-3923 regression)""" @@ -288,7 +356,7 @@ class TestBedrockRealtimeAwsAuth: await handler.async_realtime( model="amazon.nova-sonic-v1:0", websocket=websocket, - logging_obj=MagicMock(), + logging_obj=FakeLogging(), aws_region_name="us-east-1", aws_access_key_id="litellm-params-access-key", aws_secret_access_key="litellm-params-secret-key", @@ -318,7 +386,7 @@ class TestBedrockRealtimeAwsAuth: await handler.async_realtime( model="amazon.nova-sonic-v1:0", websocket=RealtimeClientWS(), - logging_obj=MagicMock(), + logging_obj=FakeLogging(), aws_region_name="eu-west-1", aws_role_name="arn:aws:iam::123456789012:role/nova-sonic", aws_session_name="realtime-session", diff --git a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py index a68aa603b26..aa002b6e302 100644 --- a/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py +++ b/tests/test_litellm/llms/bedrock/realtime/test_bedrock_realtime_transformation.py @@ -403,8 +403,9 @@ class TestBedrockRealtimeResponseCreate: class TestBedrockRealtimeResponseTransformation: """Test suite for response transformation""" - def test_transform_session_start_response(self): - """Test sessionStart response transformation""" + def test_bedrock_session_start_does_not_emit_duplicate_session_created(self): + """A Bedrock output sessionStart must not forward a second session.created to the + client; session.created is sent exactly once on connect (LIT-4655)""" config = BedrockRealtimeConfig() logging_obj = MagicMock() logging_obj.litellm_trace_id = "trace_123" @@ -428,10 +429,8 @@ class TestBedrockRealtimeResponseTransformation: }, ) - assert len(result["response"]) == 1 - assert result["response"][0]["type"] == "session.created" - assert result["response"][0]["session"]["id"] == "trace_123" - assert "model" in result["response"][0]["session"] + assert result["response"] == [] + assert result["session_configuration_request"] == json.dumps({"configured": True}) def test_transform_text_output_response(self): """Test textOutput response transformation""" @@ -789,5 +788,47 @@ class TestBedrockRealtimeResponseTransformation: assert len(set(response_ids)) == 1, "Response IDs should be consistent" +class TestBedrockRealtimeSessionEvents: + """session.created / session.updated builders produce spec-shaped events (LIT-4655)""" + + @staticmethod + def _logging(): + from types import SimpleNamespace + + return SimpleNamespace(litellm_trace_id="trace_123") + + def test_session_created_event_shape(self): + event = BedrockRealtimeConfig().session_created_event("amazon.nova-sonic-v1:0", self._logging()) + assert event["type"] == "session.created" + assert event["session"]["id"] == "trace_123" + assert event["session"]["model"] == "amazon.nova-sonic-v1:0" + assert event["session"]["modalities"] == ["text", "audio"] + assert event["event_id"] + + def test_session_updated_event_shape(self): + event = BedrockRealtimeConfig().session_updated_event("amazon.nova-sonic-v1:0", self._logging()) + assert event["type"] == "session.updated" + assert event["session"]["id"] == "trace_123" + assert event["session"]["model"] == "amazon.nova-sonic-v1:0" + assert event["event_id"] + + def test_created_and_updated_have_distinct_event_ids(self): + config = BedrockRealtimeConfig() + logging_obj = self._logging() + created = config.session_created_event("amazon.nova-sonic-v1:0", logging_obj) + updated = config.session_updated_event("amazon.nova-sonic-v1:0", logging_obj) + assert created["event_id"] != updated["event_id"] + + def test_session_updated_reflects_requested_modalities(self): + event = BedrockRealtimeConfig().session_updated_event( + "amazon.nova-sonic-v1:0", self._logging(), modalities=["text"] + ) + assert event["session"]["modalities"] == ["text"] + + def test_session_updated_defaults_modalities_when_unspecified(self): + event = BedrockRealtimeConfig().session_updated_event("amazon.nova-sonic-v1:0", self._logging()) + assert event["session"]["modalities"] == ["text", "audio"] + + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py b/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py index 8a072fa5097..c907e3249d1 100644 --- a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py +++ b/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py @@ -1,4 +1,3 @@ -import importlib import json import os import sys @@ -16,22 +15,7 @@ MOCK_EMBEDDING_RESPONSE = [[0.1, 0.2, 0.3, 0.4, 0.5]] @pytest.fixture -def reload_huggingface_modules(): - """ - Reload modules to ensure fresh references after conftest reloads litellm. - This ensures the HTTPHandler class being patched is the same one used by - the embedding handler during parallel test execution. - """ - import litellm.llms.custom_httpx.http_handler as http_handler_module - import litellm.llms.huggingface.embedding.handler as hf_embedding_handler_module - - importlib.reload(http_handler_module) - importlib.reload(hf_embedding_handler_module) - yield - - -@pytest.fixture -def mock_embedding_http_handler(reload_huggingface_modules): +def mock_embedding_http_handler(): """Fixture to mock the HTTP handler for embedding tests""" with patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post") as mock_post: mock_response = MagicMock() @@ -43,7 +27,7 @@ def mock_embedding_http_handler(reload_huggingface_modules): @pytest.fixture -def mock_embedding_async_http_handler(reload_huggingface_modules): +def mock_embedding_async_http_handler(): """Fixture to mock the async HTTP handler for embedding tests""" with patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", @@ -121,6 +105,20 @@ class TestHuggingFaceEmbedding: assert response.usage.prompt_tokens > 0 assert response.usage.total_tokens == response.usage.prompt_tokens + def test_model_name_with_https_substring_uses_api_base(self): + api_base = "https://legit.example/embed" + + litellm.embedding( + model="huggingface/my-https-endpoint", + input=["hello world"], + input_type="embed", + api_base=api_base, + ) + + self.mock_http.assert_called_once() + called_url = self.mock_http.call_args[0][0] + assert called_url == api_base + def test_embedding_with_sentence_similarity_task(self): """Test embedding when task type is sentence-similarity (requires 2+ sentences)""" diff --git a/tests/test_litellm/llms/oobabooga/chat/test_oobabooga.py b/tests/test_litellm/llms/oobabooga/chat/test_oobabooga.py new file mode 100644 index 00000000000..91ebb2bd9d4 --- /dev/null +++ b/tests/test_litellm/llms/oobabooga/chat/test_oobabooga.py @@ -0,0 +1,55 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm + +MOCK_COMPLETION_RESPONSE = { + "choices": [{"message": {"role": "assistant", "content": "hi there"}}], + "usage": {"prompt_tokens": 3, "completion_tokens": 2, "total_tokens": 5}, +} + + +def _mock_post_response(): + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.text = "ok" + mock_response.json.return_value = MOCK_COMPLETION_RESPONSE + return mock_response + + +def test_model_name_with_https_substring_uses_api_base(): + api_base = "https://legit.example" + + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post" + ) as mock_post: + mock_post.return_value = _mock_post_response() + + litellm.completion( + model="oobabooga/my-https-model", + messages=[{"role": "user", "content": "hello"}], + api_base=api_base, + ) + + mock_post.assert_called_once() + called_url = mock_post.call_args[0][0] + assert called_url == f"{api_base}/v1/chat/completions" + + +def test_url_valued_model_still_targets_that_url(): + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post" + ) as mock_post: + mock_post.return_value = _mock_post_response() + + litellm.completion( + model="oobabooga/https://sdk-user.example", + messages=[{"role": "user", "content": "hello"}], + ) + + mock_post.assert_called_once() + called_url = mock_post.call_args[0][0] + assert called_url == "https://sdk-user.example/v1/chat/completions" diff --git a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 4c268d9dfc9..7730b664c5e 100644 --- a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -1137,3 +1137,95 @@ class TestGetStructuredMessages: if __name__ == "__main__": # Run the tests pytest.main([__file__, "-v"]) + + +class TestIncrementalScanRespectsSkipFlags: + """PR #33278: skip_system_message_in_guardrail and skip_tool_message_in_guardrail + are enforced while this handler builds inputs["texts"] (_extract_inputs early + returns for system/tool roles), upstream of BedrockGuardrail's incremental path. + Bypassing _select_messages_for_apply_guardrail therefore cannot resurrect skipped + content on any turn, including a session's first turn where every segment is new. + Verified live against a real Bedrock ApplyGuardrail before being encoded here. + The flags are set as instance attributes, mirroring how guardrail_registry + applies litellm_params to the callback (they are not constructor kwargs). + """ + + def _bedrock_guardrail(self): + from litellm.proxy.guardrails.guardrail_hooks.bedrock_guardrails import BedrockGuardrail + + guardrail = BedrockGuardrail( + guardrail_name="bedrock-incremental-skip-flags", + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + default_on=True, + only_scan_new_messages=True, + ) + guardrail.skip_system_message_in_guardrail = True + guardrail.skip_tool_message_in_guardrail = True + return guardrail + + def _messages(self, followup=None): + base = [ + {"role": "system", "content": "SYSTEM-PROMPT-must-not-be-scanned"}, + {"role": "user", "content": "Search for the weather in Paris"}, + { + "role": "assistant", + "content": "Let me look that up.", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "search", "arguments": '{"query": "weather"}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "call_1", "content": "TOOL-RESULT-must-not-be-scanned"}, + {"role": "user", "content": "Thanks, summarize."}, + ] + return base + (followup or []) + + @pytest.mark.asyncio + async def test_first_turn_scans_no_system_or_tool_content(self): + from unittest.mock import AsyncMock, patch + + handler = OpenAIChatCompletionsHandler() + guardrail = self._bedrock_guardrail() + data = {"messages": self._messages(), "litellm_session_id": "skip-flags-turn1"} + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + assert mock_api.call_count == 1 + scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]] + assert scanned == [ + "Search for the weather in Paris", + "Let me look that up.", + "Thanks, summarize.", + ] + assert not any("SYSTEM-PROMPT" in text for text in scanned) + assert not any("TOOL-RESULT" in text for text in scanned) + + @pytest.mark.asyncio + async def test_second_turn_scans_only_new_eligible_content(self): + from unittest.mock import AsyncMock, patch + + handler = OpenAIChatCompletionsHandler() + guardrail = self._bedrock_guardrail() + session = "skip-flags-turn2" + followup = [ + {"role": "assistant", "content": "It is sunny in Paris."}, + {"role": "user", "content": "And tomorrow?"}, + ] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await handler.process_input_messages( + data={"messages": self._messages(), "litellm_session_id": session}, + guardrail_to_apply=guardrail, + ) + mock_api.reset_mock() + await handler.process_input_messages( + data={"messages": self._messages(followup), "litellm_session_id": session}, + guardrail_to_apply=guardrail, + ) + assert mock_api.call_count == 1 + scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]] + assert scanned == ["It is sunny in Paris.", "And tomorrow?"] diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_chat_transformation.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_chat_transformation.py new file mode 100644 index 00000000000..54e6f95c795 --- /dev/null +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_chat_transformation.py @@ -0,0 +1,235 @@ +""" +Regression tests for LIT-4313: sagemaker_chat streaming must forward each AWS +event-stream frame as it arrives instead of buffering to a fixed 1024-byte +threshold and then draining a burst of deltas. + +The buffering came from `response.iter_bytes(chunk_size=1024)` / +`response.aiter_bytes(chunk_size=1024)`: httpx's ByteChunker withholds bytes until +`chunk_size` accumulates, so the first client delta could not be produced until +enough later frames had arrived to cross 1024 bytes, inflating TTFT and turning a +steady provider stream into gap-then-burst delivery. +""" + +import binascii +import json +import struct +from typing import AsyncIterator, Iterator +from unittest.mock import MagicMock + +import httpx +import pytest + +from litellm.llms.sagemaker.chat.transformation import SagemakerChatConfig + + +def _encode_header(name: str, value: str) -> bytes: + name_b = name.encode("utf-8") + value_b = value.encode("utf-8") + return struct.pack("B", len(name_b)) + name_b + struct.pack("B", 7) + struct.pack(">H", len(value_b)) + value_b + + +def _encode_event_frame(payload: bytes) -> bytes: + """Encode one AWS event-stream message that botocore's EventStreamBuffer decodes.""" + headers = { + ":event-type": "PayloadPart", + ":content-type": "application/json", + ":message-type": "event", + } + headers_b = b"".join(_encode_header(k, v) for k, v in headers.items()) + total_len = 16 + len(headers_b) + len(payload) + prelude = struct.pack(">I", total_len) + struct.pack(">I", len(headers_b)) + prelude_crc = struct.pack(">I", binascii.crc32(prelude) & 0xFFFFFFFF) + message = prelude + prelude_crc + headers_b + payload + message_crc = struct.pack(">I", binascii.crc32(message) & 0xFFFFFFFF) + return message + message_crc + + +def _delta_frame(index: int, content: str) -> bytes: + sse = ( + "data: " + + json.dumps( + { + "id": "chatcmpl-test", + "object": "chat.completion.chunk", + "created": 1700000000, + "choices": [{"index": 0, "delta": {"content": content}, "finish_reason": None}], + } + ) + + "\n\n" + ) + return _encode_event_frame(sse.encode("utf-8")) + + +def _make_frames(n: int) -> list[bytes]: + # Small single-token frames (< 1024 bytes each) so a fixed 1024-byte chunker + # would have to swallow several frames before releasing the first delta. + frames = [_delta_frame(i, f"token{i} ") for i in range(n)] + assert all(len(f) < 1024 for f in frames) + return frames + + +class _CountingSyncStream(httpx.SyncByteStream): + """Yields provider frames one at a time and records how many have been pulled.""" + + def __init__(self, frames: list[bytes]) -> None: + self._frames = frames + self.consumed = 0 + + def __iter__(self) -> Iterator[bytes]: + for frame in self._frames: + self.consumed += 1 + yield frame + + +class _CountingAsyncStream(httpx.AsyncByteStream): + def __init__(self, frames: list[bytes]) -> None: + self._frames = frames + self.consumed = 0 + + async def __aiter__(self) -> AsyncIterator[bytes]: + for frame in self._frames: + self.consumed += 1 + yield frame + + +class _FakeSyncClient: + def __init__(self, response: httpx.Response) -> None: + self._response = response + + def post(self, *args, **kwargs) -> httpx.Response: + return self._response + + +class _FakeAsyncClient: + def __init__(self, response: httpx.Response) -> None: + self._response = response + + async def post(self, *args, **kwargs) -> httpx.Response: + return self._response + + +def _content_of(chunk) -> str | None: + return chunk.choices[0].delta.content + + +def test_sync_first_event_emitted_after_a_single_frame(): + """The first delta must be available after exactly one source frame is pulled. + + With the old chunk_size=1024 the httpx chunker would consume several small + frames before yielding, so `consumed` would be > 1 at the first delta. + """ + frames = _make_frames(24) + stream = _CountingSyncStream(frames) + response = httpx.Response(200, stream=stream) + + wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper( + model="phi-4", + custom_llm_provider="sagemaker_chat", + logging_obj=MagicMock(), + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + headers={}, + data={}, + messages=[], + client=_FakeSyncClient(response), + ) + + first = next(c for c in wrapper.completion_stream if c is not None and _content_of(c) is not None) + assert _content_of(first) == "token0 " + assert stream.consumed == 1 + + +def test_sync_events_emitted_incrementally_without_bursting(): + """Each successive delta must correspond to exactly one newly-pulled frame.""" + frames = _make_frames(24) + stream = _CountingSyncStream(frames) + response = httpx.Response(200, stream=stream) + + wrapper = SagemakerChatConfig().get_sync_custom_stream_wrapper( + model="phi-4", + custom_llm_provider="sagemaker_chat", + logging_obj=MagicMock(), + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + headers={}, + data={}, + messages=[], + client=_FakeSyncClient(response), + ) + + consumed_at_delta = [ + stream.consumed for chunk in wrapper.completion_stream if chunk is not None and _content_of(chunk) is not None + ] + + assert consumed_at_delta == list(range(1, len(frames) + 1)) + + +@pytest.mark.asyncio +async def test_async_first_event_emitted_after_a_single_frame(): + frames = _make_frames(24) + stream = _CountingAsyncStream(frames) + response = httpx.Response(200, stream=stream) + + wrapper = await SagemakerChatConfig().get_async_custom_stream_wrapper( + model="phi-4", + custom_llm_provider="sagemaker_chat", + logging_obj=MagicMock(), + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + headers={}, + data={}, + messages=[], + client=_FakeAsyncClient(response), + ) + + consumed_at_delta = [] + async for chunk in wrapper.completion_stream: + if chunk is not None and _content_of(chunk) is not None: + consumed_at_delta.append(stream.consumed) + + assert consumed_at_delta == list(range(1, len(frames) + 1)) + + +def test_signed_body_includes_stream_flag(): + """A streaming request must carry `stream: true` in the signed body sent to SageMaker. + + `stream` flows into the request body through the transformed request (`{**optional_params}`) + and must survive SigV4 signing so the endpoint enables token-level streaming. + """ + headers, signed_body = SagemakerChatConfig().sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIATESTTESTTESTTEST", + "aws_secret_access_key": "test-secret-key", + "aws_region_name": "us-east-1", + }, + request_data={"model": "phi-4", "messages": [{"role": "user", "content": "hi"}], "stream": True}, + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + model="phi-4", + stream=True, + ) + assert signed_body is not None + assert json.loads(signed_body)["stream"] is True + + +@pytest.mark.parametrize("split_size", [1, 3, 7, 64, 4096]) +def test_decoder_reassembles_frames_across_arbitrary_byte_boundaries(split_size): + """Correctness must not depend on chunk boundaries falling on frame edges. + + Removing `chunk_size=1024` lets httpx yield raw transport reads, so in + production a single read can straddle several frames or split one frame in + half. This re-chunks the concatenated stream at boundaries that deliberately + ignore frame edges and asserts every delta still decodes, in order, exactly + once - the guarantee botocore's EventStreamBuffer provides. + """ + from litellm.llms.sagemaker.chat.transformation import AWSEventStreamDecoder + + frames = _make_frames(24) + blob = b"".join(frames) + chunks = [blob[i : i + split_size] for i in range(0, len(blob), split_size)] + + decoder = AWSEventStreamDecoder(model="phi-4", is_messages_api=True) + texts = [ + _content_of(chunk) + for chunk in decoder.iter_bytes(iter(chunks)) + if chunk is not None and _content_of(chunk) is not None + ] + + assert texts == [f"token{i} " for i in range(len(frames))] diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_completion_handler.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_completion_handler.py new file mode 100644 index 00000000000..1cb27b7cf5f --- /dev/null +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_completion_handler.py @@ -0,0 +1,174 @@ +""" +Regression tests for LIT-4313: the native `sagemaker/` streaming path must +forward each AWS event-stream frame as it arrives instead of buffering to a +fixed 1024-byte threshold and then draining a burst of tokens. + +The buffering came from `response.aiter_bytes(chunk_size=1024)`: httpx's +ByteChunker withholds bytes until `chunk_size` accumulates, so the first token +could not be produced until enough later frames had arrived to cross 1024 bytes, +inflating TTFT and turning a steady provider stream into gap-then-burst delivery. +""" + +import binascii +import json +import struct +from typing import AsyncIterator, Iterator +from unittest.mock import MagicMock + +import httpx +import pytest + +from litellm.llms.sagemaker.common_utils import SagemakerError +from litellm.llms.sagemaker.completion.handler import SagemakerLLM + + +def _encode_header(name: str, value: str) -> bytes: + name_b = name.encode("utf-8") + value_b = value.encode("utf-8") + return struct.pack("B", len(name_b)) + name_b + struct.pack("B", 7) + struct.pack(">H", len(value_b)) + value_b + + +def _encode_event_frame(payload: bytes) -> bytes: + """Encode one AWS event-stream message that botocore's EventStreamBuffer decodes.""" + headers = { + ":event-type": "PayloadPart", + ":content-type": "application/json", + ":message-type": "event", + } + headers_b = b"".join(_encode_header(k, v) for k, v in headers.items()) + total_len = 16 + len(headers_b) + len(payload) + prelude = struct.pack(">I", total_len) + struct.pack(">I", len(headers_b)) + prelude_crc = struct.pack(">I", binascii.crc32(prelude) & 0xFFFFFFFF) + message = prelude + prelude_crc + headers_b + payload + message_crc = struct.pack(">I", binascii.crc32(message) & 0xFFFFFFFF) + return message + message_crc + + +def _token_frame(text: str) -> bytes: + # SageMaker HF TGI streaming payloads are `{"token": {"text": ...}}` blobs. + sse = "data: " + json.dumps({"token": {"text": text}}) + "\n\n" + return _encode_event_frame(sse.encode("utf-8")) + + +def _make_frames(n: int) -> list[bytes]: + frames = [_token_frame(f"token{i} ") for i in range(n)] + assert all(len(f) < 1024 for f in frames) + return frames + + +class _CountingSyncStream(httpx.SyncByteStream): + """Yields provider frames one at a time and records how many have been pulled.""" + + def __init__(self, frames: list[bytes]) -> None: + self._frames = frames + self.consumed = 0 + + def __iter__(self) -> Iterator[bytes]: + for frame in self._frames: + self.consumed += 1 + yield frame + + +class _CountingAsyncStream(httpx.AsyncByteStream): + """Yields provider frames one at a time and records how many have been pulled.""" + + def __init__(self, frames: list[bytes]) -> None: + self._frames = frames + self.consumed = 0 + + async def __aiter__(self) -> AsyncIterator[bytes]: + for frame in self._frames: + self.consumed += 1 + yield frame + + +class _FakeSyncClient: + def __init__(self, response: httpx.Response) -> None: + self._response = response + + def post(self, *args, **kwargs) -> httpx.Response: + return self._response + + +class _FakeAsyncClient: + def __init__(self, response: httpx.Response) -> None: + self._response = response + + async def post(self, *args, **kwargs) -> httpx.Response: + return self._response + + +def test_sync_native_streaming_forwards_each_frame_incrementally(): + """Each token must be emitted after exactly one newly-pulled source frame. + + With the old `chunk_size=1024` the httpx chunker would swallow several small + frames before yielding, so the first token would arrive only after `consumed` + had already crossed multiple frames, and tokens would then replay in a burst. + """ + frames = _make_frames(24) + stream = _CountingSyncStream(frames) + response = httpx.Response(200, stream=stream) + + completion_stream = SagemakerLLM().make_sync_call( + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + headers={}, + data="", + logging_obj=MagicMock(), + client=_FakeSyncClient(response), + ) + + consumed_at_token = [] + texts = [] + for chunk in completion_stream: + if chunk is not None and chunk["text"]: + consumed_at_token.append(stream.consumed) + texts.append(chunk["text"]) + + assert texts == [f"token{i} " for i in range(len(frames))] + assert consumed_at_token == list(range(1, len(frames) + 1)) + + +def test_sync_native_streaming_raises_sagemaker_error_on_non_200(): + response = httpx.Response(500, text="boom") + + with pytest.raises(SagemakerError) as exc_info: + SagemakerLLM().make_sync_call( + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + headers={}, + data="", + logging_obj=MagicMock(), + client=_FakeSyncClient(response), + ) + + assert exc_info.value.status_code == 500 + + +@pytest.mark.asyncio +async def test_async_native_streaming_forwards_each_frame_incrementally(): + """Each token must be emitted after exactly one newly-pulled source frame. + + With the old `chunk_size=1024` the httpx chunker would swallow several small + frames before yielding, so the first token would arrive only after `consumed` + had already crossed multiple frames, and tokens would then replay in a burst. + """ + frames = _make_frames(24) + stream = _CountingAsyncStream(frames) + response = httpx.Response(200, stream=stream) + + completion_stream = await SagemakerLLM().make_async_call( + api_base="https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/phi-4/invocations-response-stream", + headers={}, + data="", + logging_obj=MagicMock(), + client=_FakeAsyncClient(response), + ) + + consumed_at_token = [] + texts = [] + async for chunk in completion_stream: + if chunk is not None and chunk["text"]: + consumed_at_token.append(stream.consumed) + texts.append(chunk["text"]) + + assert texts == [f"token{i} " for i in range(len(frames))] + assert consumed_at_token == list(range(1, len(frames) + 1)) diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py index 6af4cf698e2..7fea5ac0965 100644 --- a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py @@ -3,26 +3,16 @@ Integration tests for Vertex AI rerank functionality. These tests demonstrate end-to-end usage of the Vertex AI rerank feature. """ -import importlib from unittest.mock import MagicMock import httpx +from litellm.llms.vertex_ai.rerank.transformation import VertexAIRerankConfig + class TestVertexAIRerankIntegration: def setup_method(self): - # Reload modules to ensure fresh references after conftest reloads litellm. - # This ensures the class being patched is the same one used by the tests. - import litellm.llms.vertex_ai.rerank.transformation as rerank_transformation_module - - importlib.reload(rerank_transformation_module) - - # Re-import after reload to get the fresh class - from litellm.llms.vertex_ai.rerank.transformation import ( - VertexAIRerankConfig as FreshConfig, - ) - - self.config = FreshConfig() + self.config = VertexAIRerankConfig() self.model = "semantic-ranker-default@latest" def test_end_to_end_rerank_flow(self): diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py index 2d09cc0ed32..292bddf1274 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_messages_config.py @@ -591,7 +591,7 @@ class TestVertexAnthropicMidConversationSystem: Claude, and a *leading* system entry 400s on every model ("messages.0: use the top-level 'system' parameter"). These tests pin the model-aware hoist so Claude Code sessions neither collapse the prompt cache on 4.8+ nor hard-fail - on 4.7 and older (RCA: Kraken Tech high-spend).""" + on 4.7 and older (RCA: customer high-spend).""" def test_supported_model_keeps_mid_conversation_system_in_place(self, local_model_cost_map): messages = [ diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py new file mode 100644 index 00000000000..a3f46a49ba9 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -0,0 +1,343 @@ +"""Tests for the SSO identity assertion store (EMA subject-token capture). + +Pins the contract of the store that PR 2's ``_id_jag`` subject-sourcing seam will read: +the carrier validates untyped IdP token-response values at the boundary, retention is +gated on an ``oauth2_id_jag`` server being registered, the row is encrypted at rest and +round-trips exactly, a store failure never escapes into the login path, and a salt-key +rotation re-encrypts stored rows like the sibling per-user credential tables. +""" + +import json +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import jwt as pyjwt +import pytest + +from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import ( + assertion_from_sso_login, + ema_assertion_retention_enabled, + fetch_sso_identity_assertion, + persist_sso_identity_assertion, + retain_sso_identity_assertion_for_ema, + rotate_sso_identity_assertions_master_key, +) +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper +from litellm.types.mcp import MCPAuth + +SALT_KEY = "test-salt-key-for-sso-assertion-tests-1234" +SIGNING_KEY = "test-idp-signing-key-32-bytes-long-xxxx" +ISSUER = "https://idp.example.com" + + +@pytest.fixture(autouse=True) +def _set_salt_key(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY) + + +def _make_id_token(exp_offset: int = 3600, iss: str = ISSUER) -> str: + return pyjwt.encode( + {"iss": iss, "sub": "u1", "exp": int(time.time()) + exp_offset}, + SIGNING_KEY, + algorithm="HS256", + ) + + +def _make_prisma(stored: dict, db_has_id_jag_server: bool = False): + """A fake prisma client whose sso-assertion table reads and writes ``stored`` + (user_id -> assertion_b64), covering upsert, find_unique, find_many, and update. + ``db_has_id_jag_server`` drives the retention gate's authoritative DB fallback; + it is wired explicitly so the gate never reads a truthy bare MagicMock.""" + prisma = MagicMock() + prisma.db.litellm_mcpservertable.find_first = AsyncMock( + return_value=MagicMock() if db_has_id_jag_server else None + ) + + async def _upsert(where, data): + stored[where["user_id"]] = data["update"]["assertion_b64"] + + async def _find_unique(where): + blob = stored.get(where["user_id"]) + if blob is None: + return None + row = MagicMock() + row.user_id = where["user_id"] + row.assertion_b64 = blob + return row + + async def _find_many(): + rows = [] + for user_id, blob in stored.items(): + row = MagicMock() + row.user_id = user_id + row.assertion_b64 = blob + rows.append(row) + return rows + + async def _update(where, data): + stored[where["user_id"]] = data["assertion_b64"] + + prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=_upsert) + prisma.db.litellm_ssoidentityassertion.find_unique = AsyncMock(side_effect=_find_unique) + prisma.db.litellm_ssoidentityassertion.find_many = AsyncMock(side_effect=_find_many) + prisma.db.litellm_ssoidentityassertion.update = AsyncMock(side_effect=_update) + return prisma + + +def _server_with_auth(auth_type): + server = MagicMock() + server.auth_type = auth_type + return server + + +def test_assertion_from_sso_login_happy_path(): + token = _make_id_token() + assertion = assertion_from_sso_login(token, "rt_1") + assert assertion is not None + assert assertion.id_token.get_secret_value() == token + assert assertion.refresh_token is not None + assert assertion.refresh_token.get_secret_value() == "rt_1" + assert assertion.issuer == ISSUER + assert assertion.expires_at is not None + assert assertion.expires_at.timestamp() == pytest.approx(time.time() + 3600, abs=5) + + +def test_assertion_repr_never_leaks_token_material(): + token = _make_id_token() + assertion = assertion_from_sso_login(token, "rt_secret_value") + rendered = repr(assertion) + str(assertion) + assert token not in rendered + assert "rt_secret_value" not in rendered + + +@pytest.mark.parametrize("id_token", [None, "", "not-a-jwt", 12345, ["x"], {"a": 1}]) +def test_assertion_from_sso_login_rejects_unusable_id_token(id_token): + assert assertion_from_sso_login(id_token, "rt") is None + + +@pytest.mark.parametrize("refresh_token", [None, "", 123, ["rt"], {"rt": 1}]) +def test_assertion_from_sso_login_drops_malformed_refresh_token(refresh_token): + assertion = assertion_from_sso_login(_make_id_token(), refresh_token) + assert assertion is not None + assert assertion.refresh_token is None + + +def test_assertion_without_exp_or_iss_still_retained(): + token = pyjwt.encode({"sub": "u1"}, SIGNING_KEY, algorithm="HS256") + assertion = assertion_from_sso_login(token, None) + assert assertion is not None + assert assertion.expires_at is None + assert assertion.issuer is None + + +@pytest.mark.asyncio +async def test_retention_gate_requires_an_id_jag_server(): + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)), + ): + manager.config_mcp_servers = { + "s1": _server_with_auth(MCPAuth.oauth2), + "s2": _server_with_auth(None), + } + assert await ema_assertion_retention_enabled() is False + manager.config_mcp_servers = { + "s1": _server_with_auth(MCPAuth.oauth2), + "s2": _server_with_auth(MCPAuth.oauth2_id_jag), + } + assert await ema_assertion_retention_enabled() is True + + +@pytest.mark.asyncio +async def test_retention_gate_reads_the_db_when_config_declares_no_id_jag_server(): + """A DB-backed server added on another pod (or before this pod's DB load) must still enable + retention off the authoritative DB row; False only when neither authority knows one.""" + with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager: + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)} + db_backed = _make_prisma({}, db_has_id_jag_server=True) + with patch("litellm.proxy.proxy_server.prisma_client", db_backed): + assert await ema_assertion_retention_enabled() is True + db_backed.db.litellm_mcpservertable.find_first.assert_awaited_once_with( + where={"auth_type": MCPAuth.oauth2_id_jag.value} + ) + with patch("litellm.proxy.proxy_server.prisma_client", None): + assert await ema_assertion_retention_enabled() is False + + +@pytest.mark.asyncio +async def test_retention_gate_never_consults_the_registry_snapshot(): + """The registry is a per-process snapshot of DB state, stale in either direction: trusting + it positively would keep retaining bearer material after the last EMA server was removed on + another pod, trusting it negatively would drop writes for one added elsewhere. The gate must + judge only the config declaration and the DB row, so a stale snapshot listing an id_jag + server changes nothing.""" + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", _make_prisma({}, db_has_id_jag_server=False)), + ): + manager.config_mcp_servers = {} + manager.get_registry.return_value = {"stale": _server_with_auth(MCPAuth.oauth2_id_jag)} + assert await ema_assertion_retention_enabled() is False + manager.get_registry.assert_not_called() + + +@pytest.mark.asyncio +async def test_retain_persists_when_only_the_db_knows_the_id_jag_server(): + stored = {} + prisma = _make_prisma(stored, db_has_id_jag_server=True) + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", prisma), + ): + manager.config_mcp_servers = {} + await retain_sso_identity_assertion_for_ema( + user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) + ) + assert "user-a" in stored + + +@pytest.mark.asyncio +async def test_persist_and_fetch_round_trip_encrypted_at_rest(): + stored = {} + prisma = _make_prisma(stored) + token = _make_id_token() + assertion = assertion_from_sso_login(token, "rt_1") + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await persist_sso_identity_assertion("user-a", assertion) + fetched = await fetch_sso_identity_assertion("user-a") + assert fetched is not None + assert fetched.id_token.get_secret_value() == token + assert fetched.refresh_token is not None + assert fetched.refresh_token.get_secret_value() == "rt_1" + assert fetched.issuer == assertion.issuer + assert fetched.expires_at == assertion.expires_at + assert token not in stored["user-a"] + assert "rt_1" not in stored["user-a"] + decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug") + assert json.loads(decrypted)["id_token"] == token + + +@pytest.mark.asyncio +async def test_persist_overwrites_previous_login(): + stored = {} + prisma = _make_prisma(stored) + first = _make_id_token(exp_offset=100) + second = _make_id_token(exp_offset=7200) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(first, None)) + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(second, "rt_new")) + fetched = await fetch_sso_identity_assertion("user-a") + assert fetched is not None + assert fetched.id_token.get_secret_value() == second + assert fetched.refresh_token is not None + + +@pytest.mark.asyncio +async def test_fetch_missing_row_returns_none(): + prisma = _make_prisma({}) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + assert await fetch_sso_identity_assertion("nobody") is None + + +@pytest.mark.asyncio +async def test_fetch_undecryptable_row_returns_none(): + prisma = _make_prisma({"user-a": "not-an-encrypted-blob"}) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + assert await fetch_sso_identity_assertion("user-a") is None + + +@pytest.mark.asyncio +async def test_fetch_unparseable_payload_returns_none(): + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + prisma = _make_prisma({"user-a": encrypt_value_helper("]]not json")}) + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + assert await fetch_sso_identity_assertion("user-a") is None + + +@pytest.mark.asyncio +async def test_retain_noop_when_no_id_jag_server(): + stored = {} + prisma = _make_prisma(stored) + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", prisma), + ): + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2)} + await retain_sso_identity_assertion_for_ema( + user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) + ) + prisma.db.litellm_ssoidentityassertion.upsert.assert_not_called() + assert stored == {} + + +@pytest.mark.asyncio +async def test_retain_persists_when_id_jag_server_registered(): + stored = {} + prisma = _make_prisma(stored) + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", prisma), + ): + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)} + await retain_sso_identity_assertion_for_ema( + user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) + ) + assert "user-a" in stored + + +@pytest.mark.asyncio +async def test_retain_none_assertion_never_consults_gate_or_store(): + gate = MagicMock() + with patch( + "litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store.ema_assertion_retention_enabled", + gate, + ): + await retain_sso_identity_assertion_for_ema(user_id="user-a", assertion=None) + gate.assert_not_called() + + +@pytest.mark.asyncio +async def test_retain_swallows_store_failure(): + prisma = MagicMock() + prisma.db.litellm_ssoidentityassertion.upsert = AsyncMock(side_effect=RuntimeError("db down")) + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as manager, + patch("litellm.proxy.proxy_server.prisma_client", prisma), + ): + manager.config_mcp_servers = {"s1": _server_with_auth(MCPAuth.oauth2_id_jag)} + await retain_sso_identity_assertion_for_ema( + user_id="user-a", assertion=assertion_from_sso_login(_make_id_token(), None) + ) + + +@pytest.mark.asyncio +async def test_rotation_reencrypts_under_new_key(monkeypatch): + stored = {} + prisma = _make_prisma(stored) + token = _make_id_token() + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await persist_sso_identity_assertion("user-a", assertion_from_sso_login(token, None)) + original_blob = stored["user-a"] + + new_key = "rotated-sso-assertion-salt-key-5678" + await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key=new_key) + assert stored["user-a"] != original_blob + + monkeypatch.setenv("LITELLM_SALT_KEY", new_key) + decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug") + assert decrypted is not None + assert json.loads(decrypted)["id_token"] == token + + +@pytest.mark.asyncio +async def test_rotation_skips_unreadable_rows_but_rotates_readable_ones(): + stored = {"good": None, "bad": "garbage-blob"} + prisma = _make_prisma(stored) + token = _make_id_token() + with patch("litellm.proxy.proxy_server.prisma_client", prisma): + await persist_sso_identity_assertion("good", assertion_from_sso_login(token, None)) + good_blob_before = stored["good"] + await rotate_sso_identity_assertions_master_key(prisma_client=prisma, new_master_key="another-new-salt-key-0000") + assert stored["bad"] == "garbage-blob" + assert stored["good"] != good_blob_before diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 0489b197652..692e5340f48 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -8031,3 +8031,606 @@ async def test_bare_origin_discovery_resolves_single_server_not_aggregate(): assert resource_response["authorization_servers"] == ["https://llm.example.com/test_oauth"] finally: global_mcp_server_manager.registry.clear() + + +@pytest.mark.asyncio +async def test_authorize_wall_names_the_fix_for_urlless_servers(): + """LIT-4629: the authorize wall previously said only "authorization url is not set" with no + hint that spec-only servers never discover; the detail must now name both remedies (manual + Authorization URL + Token URL, or an Issuer for RFC 8414 discovery).""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize_with_server, + ) + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="urlless-wall", + name="sheets_wall", + server_name="sheets_wall", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + spec_path="https://example.com/openapi.yaml", + ) + mock_request = MagicMock() + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await authorize_with_server( + request=mock_request, + mcp_server=server, + client_id="client", + redirect_uri="http://localhost/callback", + ) + assert exc_info.value.status_code == 400 + detail_text = str(exc_info.value.detail) + assert "set Authorization URL and Token URL" in detail_text + assert "Issuer" in detail_text + + +@pytest.mark.asyncio +async def test_token_wall_names_the_fix_for_urlless_servers(): + """The /token wall is the second stop on the same misconfiguration (LIT-4629): after an admin + fills only the Authorization URL, the code exchange dies here; the detail must name the + remedies like the authorize wall does.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="urlless-token-wall", + name="sheets_token_wall", + server_name="sheets_token_wall", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + spec_path="https://example.com/openapi.yaml", + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + ) + mock_request = MagicMock() + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await exchange_token_with_server( + request=mock_request, + mcp_server=server, + grant_type="authorization_code", + code="auth-code", + redirect_uri="http://localhost/callback", + client_id="client", + client_secret=None, + code_verifier="verifier", + ) + assert exc_info.value.status_code == 400 + detail_text = str(exc_info.value.detail) + assert "set Token URL manually" in detail_text + assert "Issuer" in detail_text + + +@pytest.mark.asyncio +async def test_register_wall_names_the_fix_for_urlless_servers(): + """The /register wall serves the same missing-authorization-url 400 as authorize; its detail + must carry the same actionable remedies.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + register_client_with_server, + ) + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer( + server_id="urlless-register-wall", + name="sheets_register_wall", + server_name="sheets_register_wall", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + spec_path="https://example.com/openapi.yaml", + ) + mock_request = MagicMock() + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with pytest.raises(HTTPException) as exc_info: + await register_client_with_server( + request=mock_request, + mcp_server=server, + client_name="client", + grant_types=None, + response_types=None, + token_endpoint_auth_method=None, + ) + assert exc_info.value.status_code == 400 + detail_text = str(exc_info.value.detail) + assert "set Authorization URL and Token URL" in detail_text + assert "Issuer" in detail_text + + +def test_passthrough_authorization_code_round_trips_and_rejects_hostile_input(): + """The passthrough gateway code seals and recovers the ephemeral DCR client and upstream code, + and is total over hostile input: a raw upstream code opens to None, and a tampered or + non-gateway value opens to None rather than raising, so every existing caller-supplied-client + flow is untouched.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + open_passthrough_authorization_code, + seal_passthrough_authorization_code, + ) + + with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY): + sealed = seal_passthrough_authorization_code( + upstream_code="up-code", + client_id="minted-77", + client_secret="mint-secret", + mcp_server_id="srv-1", + token_endpoint_auth_method="client_secret_basic", + ) + opened = open_passthrough_authorization_code(sealed) + assert opened is not None + assert opened.upstream_code == "up-code" + assert opened.client_id == "minted-77" + assert opened.client_secret == "mint-secret" + assert opened.mcp_server_id == "srv-1" + assert opened.token_endpoint_auth_method == "client_secret_basic" + assert open_passthrough_authorization_code("raw-upstream-code") is None + assert open_passthrough_authorization_code(sealed[:-4] + "aaaa") is None + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _BRIDGE_AUTH_CODE_PREFIX, + _PASSTHROUGH_AUTH_CODE_PREFIX, + open_bridge_authorization_code, + seal_bridge_authorization_code, + ) + + bridge_sealed = seal_bridge_authorization_code( + upstream_code="up-code", litellm_user_id="sso-user-9", mcp_server_id="srv-1" + ) + reprefixed_as_passthrough = _PASSTHROUGH_AUTH_CODE_PREFIX + bridge_sealed[len(_BRIDGE_AUTH_CODE_PREFIX) :] + reprefixed_as_bridge = _BRIDGE_AUTH_CODE_PREFIX + sealed[len(_PASSTHROUGH_AUTH_CODE_PREFIX) :] + assert open_passthrough_authorization_code(reprefixed_as_passthrough) is None + assert open_bridge_authorization_code(reprefixed_as_bridge) is None + + +@pytest.mark.asyncio +async def test_authorize_with_ephemeral_dcr_client_seals_client_into_state(): + """When mcp_authorize fell through to a gateway-side DCR mint, authorize_with_server seals the + minted client and the target server into the encrypted OAuth state, so the callback can bind + them into the forwarded authorization code while the gateway stores nothing.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + EphemeralDcrClient, + authorize_with_server, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None) + captured: dict = {} + + def _capture(**kwargs): + captured.update(kwargs) + return "mocked_encrypted_state" + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encode_state_with_base_url", + side_effect=_capture, + ): + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=server, + client_id="minted-77", + redirect_uri="http://127.0.0.1:60108/callback", + state="s", + code_challenge="chal", + code_challenge_method="S256", + ephemeral_dcr_client=EphemeralDcrClient( + client_id="minted-77", client_secret="mint-secret", token_endpoint_auth_method="client_secret_basic" + ), + ) + + assert captured["dcr_client_id"] == "minted-77" + assert captured["dcr_client_secret"] == "mint-secret" + assert captured["dcr_token_endpoint_auth_method"] == "client_secret_basic" + assert captured["mcp_server_id"] == server.server_id + assert "client_id=minted-77" in response.headers["location"] + + +@pytest.mark.asyncio +async def test_callback_wraps_code_into_passthrough_code_for_ephemeral_dcr_state(): + """When the OAuth state carries an ephemeral DCR client, the callback forwards a sealed + passthrough code (binding the client and the upstream code to the server) instead of the raw + upstream code, so the client's later token call can authenticate the exchange with a client the + gateway never stored.""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + callback, + open_passthrough_authorization_code, + ) + + state_data = { + "original_state": "client-state", + "client_redirect_uri": "http://127.0.0.1:60108/cb", + "base_url": "http://127.0.0.1:60108/cb", + "mcp_server_id": "srv-1", + "dcr_client_id": "minted-77", + "dcr_client_secret": "mint-secret", + "dcr_token_endpoint_auth_method": "client_secret_basic", + } + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_encoded_oauth_state", + return_value="enc", + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash", + return_value=state_data, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._get_validated_client_redirect_uri", + return_value="http://127.0.0.1:60108/cb", + ), + patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), + ): + response = await callback(request=_bridge_mock_request(), code="REAL-UPSTREAM-CODE", state="relay") + + forwarded_code = parse_qs(urlparse(response.headers["location"]).query)["code"][0] + opened = open_passthrough_authorization_code(forwarded_code) + + assert opened is not None + assert opened.upstream_code == "REAL-UPSTREAM-CODE" + assert opened.client_id == "minted-77" + assert opened.client_secret == "mint-secret" + assert opened.mcp_server_id == "srv-1" + assert opened.token_endpoint_auth_method == "client_secret_basic" + + +@pytest.mark.asyncio +async def test_authorize_bridge_server_with_ephemeral_client_takes_short_circuit_arm(): + """A gateway-minted client is registered against {base}/callback, so a bridge server's + authorize with an ephemeral client must run the short-circuit (gateway /callback) arm with a + relay state cookie, never the verbatim relay: relaying would send the browser's redirect_uri + to an IdP that has the gateway callback registered, stranding the flow.""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + EphemeralDcrClient, + authorize_with_server, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.true_passthrough) + with patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY): + response = await authorize_with_server( + request=_bridge_mock_request(), + mcp_server=server, + client_id="minted-77", + redirect_uri="http://127.0.0.1:60108/callback", + state="client-state", + code_challenge="chal", + code_challenge_method="S256", + ephemeral_dcr_client=EphemeralDcrClient(client_id="minted-77", client_secret=None), + ) + + location = response.headers["location"] + params = parse_qs(urlparse(location).query) + assert params["redirect_uri"] == ["https://litellm.example.com/callback"] + assert params["client_id"] == ["minted-77"] + assert params["state"] != ["client-state"] + assert any(cookie.startswith("mcp_oauth_state_") for cookie in response.headers.get("set-cookie", "").split(";")) + + +@pytest.mark.asyncio +async def test_callback_forwards_raw_code_when_dcr_state_lacks_server_binding(): + """A state carrying a dcr client but no server id cannot produce a server-bound sealed code, so + the callback falls back to forwarding the raw upstream code instead of sealing an unbindable + one.""" + from urllib.parse import parse_qs, urlparse + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import callback + + state_data = { + "original_state": "client-state", + "client_redirect_uri": "http://127.0.0.1:60108/cb", + "base_url": "http://127.0.0.1:60108/cb", + "dcr_client_id": "minted-77", + } + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._resolve_encoded_oauth_state", + return_value="enc", + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.decode_state_hash", + return_value=state_data, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints._get_validated_client_redirect_uri", + return_value="http://127.0.0.1:60108/cb", + ), + patch("litellm.proxy.proxy_server.master_key", _BRIDGE_MASTER_KEY), + ): + response = await callback(request=_bridge_mock_request(), code="REAL-UPSTREAM-CODE", state="relay") + + forwarded_code = parse_qs(urlparse(response.headers["location"]).query)["code"][0] + assert forwarded_code == "REAL-UPSTREAM-CODE" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth_type_value", + [ + "none", + "api_key", + "bearer_token", + "basic", + "authorization", + "oauth2", + "aws_sigv4", + "token", + "oauth2_token_exchange", + "oauth2_id_jag", + "true_passthrough", + "oauth_delegate", + ], +) +@pytest.mark.parametrize("dcr_bridge", [True, False]) +async def test_resolve_ephemeral_dcr_client_mint_set_is_exact(auth_type_value, dcr_bridge): + """The full authorize-time mint decision matrix, one cell per (auth_type, dcr_bridge). The gateway + mints iff true_passthrough (any bridge) or oauth_delegate-and-not-dcr_bridge; every other mode + returns None so no non-OAuth mode ever registers an upstream client, and the interactive + oauth_delegate dcr_bridge sign-in is left to its own browser-front-door flow. The UI + gatewayMintsClientFor helper mirrors this exact set; ui/.../mcp_tools/types.test.tsx pins the + frontend side against the same table, so a divergence fails on one side or the other.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + EphemeralDcrClient, + resolve_ephemeral_dcr_client, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server( + auth_type=MCPAuth(auth_type_value), + dcr_bridge=dcr_bridge, + server_id=f"matrix_{auth_type_value}_{dcr_bridge}", + server_name=f"matrix_{auth_type_value}_{dcr_bridge}", + ) + expected_mint = server.is_true_passthrough or (server.is_oauth_delegate and not server.is_dcr_bridge) + mint_mock = AsyncMock(return_value=EphemeralDcrClient(client_id="minted", client_secret=None)) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client", + mint_mock, + ): + result = await resolve_ephemeral_dcr_client( + request=_bridge_mock_request(), + mcp_server=server, + code_challenge="chal", + code_challenge_method="S256", + redirect_uri="http://127.0.0.1:9/callback", + ) + + if expected_mint: + mint_mock.assert_awaited_once() + assert result is not None + else: + mint_mock.assert_not_awaited() + assert result is None + + +@pytest.mark.asyncio +async def test_mint_ephemeral_dcr_client_returns_none_without_registration_endpoint(): + """A server whose upstream exposes no RFC 7591 registration endpoint cannot mint, so the + fall-through reports None and the caller keeps its existing missing_client_id failure.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None, registration_url=None) + assert await mint_ephemeral_dcr_client(_bridge_mock_request(), server) is None + + +@pytest.mark.asyncio +async def test_mint_ephemeral_dcr_client_posts_rfc7591_and_returns_client(): + """The mint POSTs a public-client RFC 7591 registration bound to the gateway /callback and hands + back the upstream's client without persisting it anywhere.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server( + auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id="mint_posts_srv", server_name="mint_posts_srv" + ) + mock_response = MagicMock() + mock_response.text = json.dumps( + {"client_id": "minted-77", "client_secret": "mint-secret", "token_endpoint_auth_method": "client_secret_basic"} + ) + mock_response.raise_for_status = MagicMock() + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + minted = await mint_ephemeral_dcr_client(_bridge_mock_request(), server) + + assert minted is not None + assert minted.client_id == "minted-77" + assert minted.client_secret == "mint-secret" + assert minted.token_endpoint_auth_method == "client_secret_basic" + register_data = mock_async_client.post.call_args.kwargs["json"] + assert register_data["redirect_uris"] == ["https://litellm.example.com/callback"] + assert register_data["token_endpoint_auth_method"] == "none" + assert register_data["grant_types"] == ["authorization_code", "refresh_token"] + + +@pytest.mark.asyncio +async def test_mint_ephemeral_dcr_client_reuses_minted_client_within_flow_ttl(): + """Reloading the authorize page must not spam the upstream registration endpoint with orphan + clients: within the OAuth state's lifetime a second mint for the same server and gateway origin + reuses the cached client and performs no second upstream POST.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server( + auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id="mint_reuse_srv", server_name="mint_reuse_srv" + ) + mock_response = MagicMock() + mock_response.text = json.dumps({"client_id": "minted-77"}) + mock_response.raise_for_status = MagicMock() + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + first = await mint_ephemeral_dcr_client(_bridge_mock_request(), server) + second = await mint_ephemeral_dcr_client(_bridge_mock_request(), server) + + assert first is not None + assert second == first + mock_async_client.post.assert_called_once() + + +@pytest.mark.asyncio +async def test_mint_ephemeral_dcr_client_single_flights_concurrent_mints(): + """Two in-flight authorize requests for the same server must not both register an upstream + client: the per-key lock makes the second waiter reuse the first mint, so exactly one upstream + POST happens.""" + import asyncio + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server( + auth_type=MCPAuth.true_passthrough, + dcr_bridge=None, + server_id="mint_concurrent_srv", + server_name="mint_concurrent_srv", + ) + mock_response = MagicMock() + mock_response.text = json.dumps({"client_id": "minted-77"}) + mock_response.raise_for_status = MagicMock() + + async def _slow_post(*args, **kwargs): + await asyncio.sleep(0.05) + return mock_response + + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(side_effect=_slow_post) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + first, second = await asyncio.gather( + mint_ephemeral_dcr_client(_bridge_mock_request(), server), + mint_ephemeral_dcr_client(_bridge_mock_request(), server), + ) + + assert first is not None + assert second == first + mock_async_client.post.assert_called_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "payload, server_id", + [ + ({"unexpected": "shape"}, "mint_bad_shape_srv"), + ({"client_id": ""}, "mint_empty_id_srv"), + ], +) +async def test_mint_ephemeral_dcr_client_unusable_registration_response_is_502(payload, server_id): + """An upstream registration response without a usable client_id, whether the field is missing or + an empty string, surfaces as a loud 502 instead of letting the authorize proceed with an empty + client and fail opaquely at the IdP.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id=server_id, server_name=server_id) + mock_response = MagicMock() + mock_response.text = json.dumps(payload) + mock_response.raise_for_status = MagicMock() + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + with pytest.raises(HTTPException) as exc: + await mint_ephemeral_dcr_client(_bridge_mock_request(), server) + + assert exc.value.status_code == 502 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "sealed_auth_method, expects_basic_header", + [ + ("client_secret_basic", True), + (None, False), + ], +) +async def test_token_exchange_authenticates_with_the_sealed_clients_own_auth_method( + sealed_auth_method, expects_basic_header +): + """The id, secret, and token-endpoint auth method must come from the same source: a client + recovered from a sealed passthrough code authenticates the upstream exchange the way its own + registration was granted, not the way the server row is configured. A sealed + ``client_secret_basic`` grant sends the Basic header and keeps the secret out of the body; a + sealed public client (no method) keeps the body-credential path.""" + import base64 + + import httpx + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + exchange_token_with_server, + ) + from litellm.types.mcp import MCPAuth + + server = _bridge_server(auth_type=MCPAuth.true_passthrough, dcr_bridge=None, server_id="sealed_method_srv") + upstream_request = httpx.Request("POST", server.token_url) + upstream_response = httpx.Response( + 200, json={"access_token": "up-token", "token_type": "Bearer"}, request=upstream_request + ) + mock_async_client = MagicMock() + mock_async_client.post = AsyncMock(return_value=upstream_response) + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client", + return_value=mock_async_client, + ): + await exchange_token_with_server( + request=_bridge_mock_request(), + mcp_server=server, + grant_type="authorization_code", + code="up-code", + redirect_uri="https://litellm.example.com/callback", + client_id="minted-77", + client_secret="mint-secret", + code_verifier="verifier", + client_token_endpoint_auth_method=sealed_auth_method, + ) + + sent_headers = mock_async_client.post.call_args.kwargs["headers"] + sent_body = mock_async_client.post.call_args.kwargs["data"] + if expects_basic_header: + expected = base64.b64encode(b"minted-77:mint-secret").decode() + assert sent_headers["Authorization"] == f"Basic {expected}" + assert "client_secret" not in sent_body + else: + assert "Authorization" not in sent_headers + assert sent_body["client_id"] == "minted-77" + assert sent_body["client_secret"] == "mint-secret" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 73486fe0b6a..b56a12db5b1 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -1033,3 +1033,166 @@ class TestResolveByokMcpAuthHeader: check_mock.assert_awaited_once_with(server, user_auth) assert result == "caller-header" + + +class TestOpenApiResolvedUpstreamAuth: + """LIT-4629: spec_path servers egress through plain httpx, so the manager's OpenAPI arm must + materialize the v2-resolved credential into the `_request_resolved_auth_headers` ContextVar; + before the fix the resolved token never reached the upstream API.""" + + def _oauth_server(self, **overrides: Any) -> MCPServer: + fields: Dict[str, Any] = dict( + server_id="srv-sheets", + name="google_sheets", + server_name="google_sheets", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + spec_path="https://example.com/sheets-openapi.yaml", + ) + fields.update(overrides) + return MCPServer(**fields) + + @pytest.mark.asyncio + async def test_call_tool_openapi_injects_v2_resolved_token_contextvar(self): + """The managed spec_path arm resolves the v2 credential and sets the ContextVar; kills + the mutant that drops the resolve_openapi_upstream_auth call in call_tool.""" + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_resolved_auth_headers, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + manager = MCPServerManager() + server = self._oauth_server() + user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user") + captured: Dict[str, Any] = {} + + async def fake_openapi_handler(_server, _name, _arguments): + captured["resolved"] = _request_resolved_auth_headers.get() + return MagicMock() + + with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): + with patch.object( + manager._cred_provider, + "resolve_credentials", + new=AsyncMock(return_value=Ok(StaticHeaderAuth("Bearer stored-user-token"))), + ): + with patch.object(manager, "_call_openapi_tool_handler", side_effect=fake_openapi_handler): + await manager.call_tool( + server_name=server.server_name, + name="get_values", + arguments={}, + user_api_key_auth=user_auth, + ) + + assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"} + assert _request_resolved_auth_headers.get() is None + + @pytest.mark.asyncio + async def test_call_tool_openapi_m2m_missing_token_url_fails_closed(self): + """A url-less M2M spec server with no token_url must fail with a typed error instead of + egressing unauthenticated (the pre-#32259 silent failure this arm previously preserved). + Drives the real adapter/resolver chain: ClientCredentialsConfig with missing grant fields + resolves to a misconfigured CredError, raised as an HTTPException.""" + from fastapi import HTTPException + + manager = MCPServerManager() + server = self._oauth_server( + oauth2_flow="client_credentials", + client_id="m2m-client", + client_secret="m2m-secret", + token_url=None, + ) + called = AsyncMock() + + with patch.object(manager, "_resolve_mcp_server_for_tool_call", return_value=server): + with patch.object(manager, "_call_openapi_tool_handler", new=called): + with pytest.raises(HTTPException): + await manager.call_tool( + server_name=server.server_name, + name="get_values", + arguments={}, + user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="sk-user"), + ) + + called.assert_not_awaited() + + @pytest.mark.asyncio + async def test_caller_oauth2_headers_never_become_resolved_for_byok_server(self): + """Greptile P1 regression: BYOK servers defer to v1 (to_server_spec None), and the v1 arm + must never promote caller-supplied oauth2 headers into the resolved-auth slot, where they + would override the per-server BYOK credential and leak the caller's gateway Authorization + upstream.""" + manager = MCPServerManager() + server = MCPServer( + server_id="byok-spec", + name="byok_spec", + server_name="byok_spec", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.api_key, + spec_path="https://example.com/openapi.yaml", + is_byok=True, + ) + + resolved, forwarded = await manager.resolve_openapi_upstream_auth( + mcp_server=server, + oauth2_headers={"Authorization": "Bearer sk-litellm-gateway-key"}, + raw_headers=None, + mcp_auth_header="user-byok-key", + user_api_key_auth=UserAPIKeyAuth(user_id="alice", api_key="sk-user"), + forwarded_headers=None, + ) + + assert resolved is None + assert forwarded is None + + @pytest.mark.asyncio + async def test_v1_server_threads_stored_headers_only_without_caller_headers(self): + """The v1 (unmigrated) arm resolves the stored per-user token only when the caller sent no + oauth2 headers of their own; with caller headers present the stored lookup is skipped and + nothing is promoted to resolved.""" + manager = MCPServerManager() + server = MCPServer( + server_id="v1-spec", + name="v1_spec", + server_name="v1_spec", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + spec_path="https://example.com/openapi.yaml", + delegate_auth_to_upstream=True, + ) + stored = {"Authorization": "Bearer stored-v1-token"} + user_auth = UserAPIKeyAuth(user_id="alice", api_key="sk-user") + + with patch.object( + manager, "_resolve_oauth2_headers_for_tool_call", new=AsyncMock(return_value=stored) + ) as lookup: + resolved, _ = await manager.resolve_openapi_upstream_auth( + mcp_server=server, + oauth2_headers=None, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=user_auth, + forwarded_headers=None, + ) + assert resolved == stored + lookup.assert_awaited_once_with(server, None, user_auth) + + with patch.object( + manager, "_resolve_oauth2_headers_for_tool_call", new=AsyncMock(return_value=stored) + ) as lookup: + resolved, _ = await manager.resolve_openapi_upstream_auth( + mcp_server=server, + oauth2_headers={"Authorization": "Bearer caller-supplied"}, + raw_headers=None, + mcp_auth_header=None, + user_api_key_auth=user_auth, + forwarded_headers=None, + ) + assert resolved is None + lookup.assert_not_awaited() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 03f91260955..a5cb16822cf 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -5597,7 +5597,7 @@ class TestMCPServerTimestamps: async def test_build_mcp_server_from_table_persists_discovered_oauth_endpoints(self): """A DB-backed oauth2 server with no configured endpoints discovers them and must write authorization_url, token_url, and scopes back to the row; otherwise the resolved values - live only in memory and one failed re-discovery serves 400 "authorization url is not set" + live only in memory and one failed re-discovery serves the 400 "authorization url is not configured" from /authorize. registration_url must never be persisted because _dcr_bridge_relays_client_registration keys off that column.""" manager = MCPServerManager() @@ -8891,3 +8891,140 @@ async def test_resolve_toolset_tool_permissions_single_db_fetch_across_checks(): assert first == {"server-a": ["lookup_status"]} assert second == first list_toolsets_mock.assert_awaited_once() + + +class TestMaterializeAuthHeaders: + """_materialize_auth_headers drives one step of a resolved httpx.Auth's own flow to turn it + into a header dict for the OpenAPI egress arm, which sends plain headers and cannot carry an + httpx.Auth. Generic across auth shapes via the resolver-arm header_name convention.""" + + @pytest.mark.asyncio + async def test_static_header_auth_materializes_its_header(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _materialize_auth_headers, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + + headers = await _materialize_auth_headers(StaticHeaderAuth("Bearer stored-token")) + assert headers == {"Authorization": "Bearer stored-token"} + + @pytest.mark.asyncio + async def test_client_credentials_bearer_auth_materializes_bearer(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _materialize_auth_headers, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ( + ClientCredentialsBearerAuth, + ) + + async def _refetch(_stale: str): + return None + + headers = await _materialize_auth_headers(ClientCredentialsBearerAuth("m2m-token", _refetch)) + assert headers == {"Authorization": "Bearer m2m-token"} + + @pytest.mark.asyncio + async def test_noop_and_none_materialize_to_none(self): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _materialize_auth_headers, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + NoOpAuth, + ) + + assert await _materialize_auth_headers(None) is None + assert await _materialize_auth_headers(NoOpAuth()) is None + + +class TestUrllessIssuerDiscovery: + """LIT-4629: servers with no url (OpenAPI spec_path, stdio) run no resource discovery, so + their OAuth endpoints could only ever come from manual entry; an admin-pinned issuer is a + url-independent trust anchor (RFC 8414 section 3.3) and must unlock discovery for them.""" + + def _urlless_row(self, **overrides): + fields = dict( + server_id="urlless-1", + alias="sheets_urlless", + description="spec-only server", + url=None, + spec_path="https://example.com/sheets-openapi.yaml", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + created_at=datetime.now(), + updated_at=datetime.now(), + ) + fields.update(overrides) + return LiteLLM_MCPServerTable(**fields) + + @pytest.mark.asyncio + async def test_urlless_server_with_issuer_discovers_endpoints(self): + """The gate previously required bool(server_url), so a url-less server with an issuer + configured never ran the issuer-anchored fetch and /authorize 400d. Kills the mutant that + restores the bare bool(server_url) term.""" + manager = MCPServerManager() + row = self._urlless_row(issuer="https://accounts.google.com") + + resolved = MCPOAuthMetadata( + authorization_url="https://accounts.google.com/o/oauth2/v2/auth", + token_url="https://oauth2.googleapis.com/token", + ) + resource_rooted = AsyncMock(return_value=None) + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored, + patch.object(manager, "_descovery_metadata", new=resource_rooted), + ): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + anchored.assert_awaited_once_with("https://accounts.google.com", None) + resource_rooted.assert_not_awaited() + assert built.issuer_is_anchored is True + assert built.authorization_url == "https://accounts.google.com/o/oauth2/v2/auth" + assert built.token_url == "https://oauth2.googleapis.com/token" + + @pytest.mark.asyncio + async def test_urlless_server_without_issuer_stays_undiscovered(self): + """With neither a url nor an issuer there is no discovery source; the build must not + attempt any fetch and the endpoints stay unset (manual entry remains the only path).""" + manager = MCPServerManager() + row = self._urlless_row() + + anchored = AsyncMock() + resource_rooted = AsyncMock() + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=anchored), + patch.object(manager, "_descovery_metadata", new=resource_rooted), + ): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + anchored.assert_not_awaited() + resource_rooted.assert_not_awaited() + assert built.authorization_url is None + assert built.token_url is None + assert built.issuer_is_anchored is False + + @pytest.mark.asyncio + async def test_urlless_obo_with_issuer_discovers_token_url(self): + """oauth2_token_exchange is not a discovery auth type, so the plain gate relax alone + would leave a url-less OBO server undiscovered; with an issuer pinned and no configured + exchange endpoint it must resolve token_url through the issuer-anchored fetch. Kills the + mutant that drops the OBO widening from the anchor computation.""" + manager = MCPServerManager() + row = self._urlless_row( + alias="obo_urlless", + auth_type=MCPAuth.oauth2_token_exchange, + issuer="https://idp.example.com", + ) + + resolved = MCPOAuthMetadata(token_url="https://idp.example.com/token") + resource_rooted = AsyncMock(return_value=None) + with ( + patch.object(manager, "_fetch_issuer_anchored_oauth_metadata", new=AsyncMock(return_value=resolved)) as anchored, + patch.object(manager, "_descovery_metadata", new=resource_rooted), + ): + built = await manager.build_mcp_server_from_table(row, credentials_are_encrypted=False) + + anchored.assert_awaited_once_with("https://idp.example.com", None) + resource_rooted.assert_not_awaited() + assert built.token_url == "https://idp.example.com/token" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index 39f3c767220..7bcacb3ff4a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -17,6 +17,7 @@ import pytest from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _request_auth_header, _request_extra_headers, + _request_resolved_auth_headers, _resolve_param_list, _resolve_ref, build_input_schema, @@ -1207,3 +1208,61 @@ class TestRequestExtraHeaders: call_args = async_client.get.call_args headers_sent = call_args[1]["headers"] assert "X-TOKEN" not in headers_sent + + @pytest.mark.asyncio + async def test_resolved_auth_headers_win_over_every_other_authorization_source(self): + """The gateway-resolved credential (stored per-user OAuth / minted M2M token) is + authoritative: it must override the BYOK override, static headers, and forwarded caller + headers on the Authorization name, case-insensitively, mirroring _resolve_v2_auth's rule + on the MCPClient path. Without this, a spec_path oauth2 server's completed OAuth flow + stores a token that never reaches the upstream API (LIT-4629).""" + operation = {} + func = create_tool_function( + path="/secure", + method="get", + operation=operation, + base_url="https://api.example.com", + headers={"authorization": "Bearer static-operator"}, + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "secure-data") + mock_client.return_value = async_client + + extra_token = _request_extra_headers.set({"Authorization": "Bearer caller-forwarded"}) + auth_token = _request_auth_header.set("Bearer byok-credential") + resolved_token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) + try: + result = await func() + finally: + _request_auth_header.reset(auth_token) + _request_extra_headers.reset(extra_token) + _request_resolved_auth_headers.reset(resolved_token) + + assert result == "secure-data" + headers_sent = async_client.get.call_args[1]["headers"] + authorization_values = [v for k, v in headers_sent.items() if k.lower() == "authorization"] + assert authorization_values == ["Bearer resolved-oauth"] + + @pytest.mark.asyncio + async def test_resolved_auth_headers_not_leaked_between_calls(self): + """After resetting the resolved-auth ContextVar, subsequent calls send no credential.""" + operation = {} + func = create_tool_function( + path="/data", + method="get", + operation=operation, + base_url="https://api.example.com", + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "ok") + mock_client.return_value = async_client + + token = _request_resolved_auth_headers.set({"Authorization": "Bearer resolved-oauth"}) + _request_resolved_auth_headers.reset(token) + + await func() + + headers_sent = async_client.get.call_args[1]["headers"] + assert "Authorization" not in headers_sent diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index 3ad01e9c3ec..1e4349c3143 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -218,3 +218,86 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): assert exc.value.status_code == 503 pre_call.assert_not_awaited() handle_local.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_openapi_local_tool_injects_resolved_oauth_token(): + """LIT-4629: the local-registry (OpenAPI) dispatch is the primary egress for spec_path + tools, and before the fix it dropped the gateway-resolved OAuth credential entirely, so a + user's completed OAuth flow stored a token that never reached the upstream API. The resolved + credential must land in the `_request_resolved_auth_headers` ContextVar the tool closure + reads. Kills the mutant that deletes the resolve_openapi_upstream_auth call in server.py.""" + from litellm.proxy._experimental.mcp_server import server as mcp_module + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_resolved_auth_headers, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + from litellm.types.mcp import MCPAuth, MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + user = UserAPIKeyAuth( + api_key="sk-user", + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + oauth_server = MCPServer( + server_id="srv-sheets", + name="google_sheets", + server_name="google_sheets", + url=None, + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + spec_path="https://example.com/sheets-openapi.yaml", + ) + + fake_tool = MagicMock() + fake_tool.name = "get_values" + captured: dict = {} + + async def handle_local(_name, _arguments): + captured["resolved"] = _request_resolved_auth_headers.get() + return [] + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=oauth_server, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "pre_call_tool_check", + new=AsyncMock(return_value={}), + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=fake_tool, + ), + patch.object( + mcp_module.global_mcp_server_manager._cred_provider, + "resolve_credentials", + new=AsyncMock(return_value=Ok(StaticHeaderAuth("Bearer stored-user-token"))), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + new=handle_local, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), + ): + await mcp_module.execute_mcp_tool( + name="get_values", + arguments={}, + allowed_mcp_servers=[oauth_server], + start_time=datetime.now(timezone.utc), + user_api_key_auth=user, + ) + + assert captured["resolved"] == {"Authorization": "Bearer stored-user-token"} + assert _request_resolved_auth_headers.get() is None diff --git a/tests/test_litellm/proxy/a2a/test_agent_card.py b/tests/test_litellm/proxy/a2a/test_agent_card.py index d302bde7895..dfa848e335e 100644 --- a/tests/test_litellm/proxy/a2a/test_agent_card.py +++ b/tests/test_litellm/proxy/a2a/test_agent_card.py @@ -1,10 +1,14 @@ """Unit tests for the pure merge logic in litellm/proxy/a2a/agent_card.py.""" +import pytest + from litellm.proxy.a2a.agent_card import ( LITELLM_A2A_PROTOCOL_VERSION, LITELLM_SECURITY_REQUIREMENTS, LITELLM_SECURITY_SCHEMES, merge_agent_card, + normalize_protocol_version, + resolve_served_protocol_version, ) PROXY_URL = "https://proxy.example/a2a/agent-xyz" @@ -205,3 +209,54 @@ def test_strips_additional_interfaces_to_prevent_backend_url_leak(): ] merged = merge_agent_card(upstream, proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE) assert "additionalInterfaces" not in merged + + +@pytest.mark.parametrize( + ("raw", "expected"), + [ + ("0.3", "0.3"), + ("0.3.0", "0.3"), + ("1.0", "1.0"), + ("1.0.0", "1.0"), + ("1.0.1", "1.0"), + ("0.3.0-rc1", "0.3"), + ("1.0.0-rc.1+build.5", "1.0"), + ("0.2.6", None), + ("2.0", None), + ("0.30", None), + ("0.3.garbage", None), + ("0.3.", None), + ("1.0.not-semver", None), + ("0.3.0.0", None), + ("0.3-rc1", None), + ("garbage", None), + ("", None), + (None, None), + (1.0, None), + ], +) +def test_normalize_protocol_version(raw, expected): + assert normalize_protocol_version(raw) == expected + + +def test_resolve_served_protocol_version_canonicalizes_semver_pins(): + assert resolve_served_protocol_version({"protocolVersion": "0.3.0"}) == "0.3" + assert resolve_served_protocol_version({"protocolVersion": "1.0.0"}) == "1.0" + assert resolve_served_protocol_version({"protocolVersion": "0.3"}) == "0.3" + assert resolve_served_protocol_version({"protocolVersion": "1.0"}) == "1.0" + + +def test_resolve_served_protocol_version_falls_back_for_unsupported(): + assert ( + resolve_served_protocol_version({"protocolVersion": "0.2.6"}) + == LITELLM_A2A_PROTOCOL_VERSION + ) + assert resolve_served_protocol_version(None) == LITELLM_A2A_PROTOCOL_VERSION + + +def test_serves_semver_pinned_protocol_version_as_major_minor(): + card = _full_upstream_card() + card["protocolVersion"] = "0.3.0" + merged = merge_agent_card(card, proxy_url=PROXY_URL, proxy_base_url=PROXY_BASE) + assert merged["protocolVersion"] == "0.3" + assert merged["supportedInterfaces"][0]["protocolVersion"] == "0.3" diff --git a/tests/test_litellm/proxy/a2a/test_version_convert.py b/tests/test_litellm/proxy/a2a/test_version_convert.py index f3c51ca6b72..7eb5debb792 100644 --- a/tests/test_litellm/proxy/a2a/test_version_convert.py +++ b/tests/test_litellm/proxy/a2a/test_version_convert.py @@ -313,3 +313,13 @@ def test_agent_card_with_0_3_pin_and_supported_interfaces_is_lowered(): def test_agent_card_same_version_passthrough(): card = _extended_card_1_0() assert normalize_agent_card(card, "1.0") is card + + +def test_detect_card_version_normalizes_semver_protocol_version(): + from litellm.proxy.a2a.version_convert import _detect_card_version + + assert _detect_card_version({"protocolVersion": "1.0.0"}) == "1.0" + assert ( + _detect_card_version({"protocolVersion": "0.3.0", "supportedInterfaces": []}) + == "0.3" + ) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py index 3740c01b7fc..bcd3333baf9 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_endpoints.py @@ -540,6 +540,53 @@ class TestAgentRBACProxyAdmin: assert resp.status_code == 200 +class TestAgentProtocolVersionValidation: + """Registration accepts spec-default semver protocolVersion values and still + rejects genuinely unsupported versions.""" + + @pytest.fixture(autouse=True) + def _setup(self, monkeypatch): + self.admin_client = _make_app_with_role(LitellmUserRoles.PROXY_ADMIN) + self.mock_registry = MagicMock() + monkeypatch.setattr(agent_endpoints, "AGENT_REGISTRY", self.mock_registry) + + def _create_agent_with_protocol_version(self, protocol_version: str): + config = _sample_agent_config() + config["agent_card_params"]["protocolVersion"] = protocol_version + with patch("litellm.proxy.proxy_server.prisma_client"): + self.mock_registry.get_agent_by_name = MagicMock(return_value=None) + self.mock_registry.add_agent_to_db = AsyncMock( + return_value=_sample_agent_response() + ) + self.mock_registry.register_agent = MagicMock() + return self.admin_client.post( + "/v1/agents", + json=config, + headers={"Authorization": "Bearer k"}, + ) + + def test_semver_protocol_version_registers_and_stores_major_minor(self): + resp = self._create_agent_with_protocol_version("0.3.0") + assert resp.status_code == 200 + stored_card = self.mock_registry.add_agent_to_db.await_args.kwargs["agent"][ + "agent_card_params" + ] + assert stored_card["protocolVersion"] == "0.3" + assert stored_card["supportedInterfaces"][0]["protocolVersion"] == "0.3" + + def test_unsupported_protocol_version_is_rejected(self): + resp = self._create_agent_with_protocol_version("0.2.6") + assert resp.status_code == 400 + assert "Unsupported protocolVersion '0.2.6'" in resp.json()["detail"] + self.mock_registry.add_agent_to_db.assert_not_awaited() + + def test_malformed_protocol_version_is_rejected(self): + resp = self._create_agent_with_protocol_version("0.3.garbage") + assert resp.status_code == 400 + assert "Unsupported protocolVersion '0.3.garbage'" in resp.json()["detail"] + self.mock_registry.add_agent_to_db.assert_not_awaited() + + class TestCheckAgentManagementPermission: """Unit tests for the _check_agent_management_permission helper.""" diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 72bd215b9be..9f24c662581 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -1587,7 +1587,7 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: } out = get_dynamic_litellm_params( litellm_params=dict(admin_params), - request_kwargs={"base_url": "https://attacker.example"}, + request_kwargs={"base_url": "https://attacker.example", "api_key": "sk-caller"}, ) assert "aws_access_key_id" not in out assert "aws_secret_access_key" not in out @@ -1608,7 +1608,7 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: } out = get_dynamic_litellm_params( litellm_params=dict(admin_params), - request_kwargs={"api_base": "self-hosted.example.com:50051"}, + request_kwargs={"api_base": "self-hosted.example.com:50051", "api_key": "sk-caller"}, ) assert out["api_base"] == "self-hosted.example.com:50051" assert "nvcf_function_id" not in out @@ -1626,7 +1626,7 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: } out = get_dynamic_litellm_params( litellm_params=dict(admin_params), - request_kwargs={"api_base": "self-hosted.example.com:50051"}, + request_kwargs={"api_base": "self-hosted.example.com:50051", "api_key": "sk-caller"}, ) assert out["api_base"] == "self-hosted.example.com:50051" assert "use_ssl" not in out @@ -1651,6 +1651,7 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: }, request_kwargs={ "api_base": "https://attacker.example", + "api_key": "sk-caller", "organization": "org-attacker", "extra_body": {"attacker": "value"}, }, @@ -1674,6 +1675,7 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: }, request_kwargs={ "api_base": "https://attacker.example", + "api_key": "sk-caller", "organization": "", "extra_body": "", }, @@ -1701,6 +1703,310 @@ class TestGetDynamicLitellmParamsClearsAdminConfigOnBaseOverride: assert out["api_version"] == "2026-04-01" assert out["api_base"] == "https://admin.upstream/v1" + def test_client_api_key_used_when_supplied_with_base_override(self): + from litellm.router_utils.clientside_credential_handler import ( + get_dynamic_litellm_params, + ) + + out = get_dynamic_litellm_params( + litellm_params={ + "model": "gpt-4", + "api_key": "sk-admin-secret", + "api_base": "https://admin.upstream/v1", + }, + request_kwargs={ + "api_base": "https://attacker.example", + "api_key": "sk-client-byok", + }, + ) + assert out["api_key"] == "sk-client-byok" + assert "sk-admin-secret" not in str(out) + + +_OPENAI_CHAT_RESPONSE = { + "id": "chatcmpl-x", + "object": "chat.completion", + "created": 1, + "model": "gpt-4", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, +} + + +class TestClientsideBaseOverrideOutboundKey: + """Drive a completion through the router and assert on the outbound request + when the caller overrides ``api_base``.""" + + def _router(self): + from litellm import Router + + return Router( + model_list=[ + { + "model_name": "gpt-4", + "litellm_params": { + "model": "openai/gpt-4", + "api_key": "sk-SERVER-CONFIG", + "api_base": "https://admin.upstream/v1", + }, + } + ] + ) + + @pytest.fixture(autouse=True) + def _ambient_server_key(self, monkeypatch): + import litellm + + monkeypatch.setenv("OPENAI_API_KEY", "sk-SERVER-ENV") + monkeypatch.setattr(litellm, "api_key", None, raising=False) + + def test_caller_key_override_sends_caller_key_never_server_key(self): + import httpx + import respx + + with respx.mock: + route = respx.post("https://caller.example/v1/chat/completions").mock( + return_value=httpx.Response(200, json=_OPENAI_CHAT_RESPONSE) + ) + self._router().completion( + model="gpt-4", + messages=[{"role": "user", "content": "hi"}], + api_base="https://caller.example/v1", + api_key="sk-CALLER", + ) + authorization = route.calls.last.request.headers.get("authorization") + assert authorization == "Bearer sk-CALLER" + assert "SERVER" not in (authorization or "") + + +def _rounds_deep_api_base_payload(rounds, field): + """Build a fallbacks payload with ``api_base`` on a target nested ``rounds`` + fallback-rounds deep, each round wrapped in its own grouping dict.""" + node = {"model": "leaf", "api_base": "https://attacker.example"} + for i in range(rounds): + node = {"model": f"m{i}", field: [{"grp": [node]}]} + return {"model": "gpt-4", field: [{"grp": [node]}]} + + +class TestIsRequestBodySafeBlocksFallbackSmuggle: + """``is_request_body_safe`` runs the banned-param check on every dict target + inside the fallback lists.""" + + @pytest.fixture(autouse=True) + def _disable_url_validation(self, monkeypatch): + import litellm + + monkeypatch.setattr(litellm, "user_url_validation", False, raising=False) + + @pytest.mark.parametrize( + "fallback_key", + ["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"], + ) + def test_api_base_smuggled_via_nested_fallback_is_rejected(self, fallback_key): + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body={ + "model": "gpt-4", + fallback_key: [ + { + "gpt-4": [ + {"model": "evil", "api_base": "https://attacker.example"}, + ] + } + ], + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_string_only_fallbacks_are_accepted(self): + assert ( + is_request_body_safe( + request_body={ + "model": "gpt-4", + "fallbacks": [{"gpt-4": ["gpt-3.5-turbo", "claude-3-haiku"]}], + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + def test_benign_dict_fallback_entry_is_accepted(self): + assert ( + is_request_body_safe( + request_body={ + "model": "gpt-4", + "fallbacks": [{"gpt-4": [{"model": "gpt-3.5-turbo"}]}], + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + def test_smuggled_fallback_allowed_under_proxy_wide_opt_in(self): + assert ( + is_request_body_safe( + request_body={ + "model": "gpt-4", + "fallbacks": [ + {"gpt-4": [{"model": "byok", "api_base": "https://my-byok.example"}]} + ], + }, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + @pytest.mark.parametrize( + "fallback_field", + ["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"], + ) + @pytest.mark.parametrize("surface", ["top_level", "router_settings_override"]) + def test_deeply_nested_api_base_smuggle_rejected_on_both_surfaces(self, fallback_field, surface): + nested = [ + { + "always-fail": [ + { + "model": "x", + fallback_field: [ + {"x": [{"model": "deepseek-chat", "api_base": "http://attacker"}]} + ], + } + ] + } + ] + request_body = {"model": "gpt-4"} + if surface == "top_level": + request_body[fallback_field] = nested + else: + request_body["router_settings_override"] = {fallback_field: nested} + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body=request_body, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_router_settings_override_single_level_api_base_rejected(self): + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body={ + "model": "gpt-4", + "router_settings_override": { + "fallbacks": [{"gpt-4": [{"model": "x", "api_base": "http://attacker"}]}] + }, + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_model_less_config_dict_api_base_rejected(self): + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body={ + "model": "gpt-4", + "fallbacks": [{"gpt-4": [{"api_base": "http://attacker"}]}], + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_nested_api_base_caught_across_router_fallback_rounds(self): + """An ``api_base`` target nested ``ROUTER_MAX_FALLBACKS - 1`` rounds deep + is still reached and rejected.""" + import litellm + + with pytest.raises(ValueError, match="api_base"): + is_request_body_safe( + request_body=_rounds_deep_api_base_payload(litellm.ROUTER_MAX_FALLBACKS - 1, "fallbacks"), + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_grouping_only_deep_chain_is_rejected_at_depth_limit(self): + """A deep grouping-only chain (``{"g": [{"g": [...]}]}``) is rejected at the + validation-depth limit rather than accepted or raising RecursionError.""" + node: object = ["safe-model"] + for _ in range(5000): + node = [{"grp": node}] + with pytest.raises(ValueError, match="depth"): + is_request_body_safe( + request_body={"model": "gpt-4", "fallbacks": node}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_pathologically_deep_model_nesting_is_rejected(self): + with pytest.raises(ValueError, match="depth"): + is_request_body_safe( + request_body=_rounds_deep_api_base_payload(5000, "fallbacks"), + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + +class TestIsRequestBodySafeRejectsUrlValuedFallback: + @pytest.mark.parametrize("fallback_field", ["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"]) + def test_url_valued_string_fallback_is_rejected(self, fallback_field): + with pytest.raises(ValueError, match="URL-valued fallback"): + is_request_body_safe( + request_body={ + "model": "gpt-4", + fallback_field: [{"gpt-4": ["huggingface/http://attacker.example/path"]}], + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + @pytest.mark.parametrize("fallback_field", ["fallbacks", "context_window_fallbacks", "content_policy_fallbacks"]) + def test_url_valued_dict_model_fallback_is_rejected(self, fallback_field): + with pytest.raises(ValueError, match="URL-valued fallback"): + is_request_body_safe( + request_body={ + "model": "gpt-4", + fallback_field: [{"gpt-4": [{"model": "huggingface/http://attacker.example/path"}]}], + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_ordinary_string_fallback_is_allowed(self): + assert ( + is_request_body_safe( + request_body={"model": "gpt-4", "fallbacks": [{"gpt-4": ["gpt-4-backup"]}]}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + def test_ordinary_dict_model_fallback_is_allowed(self): + assert ( + is_request_body_safe( + request_body={"model": "gpt-4", "fallbacks": [{"gpt-4": [{"model": "gpt-4-backup"}]}]}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + is True + ) + class TestIsRequestBodySafeBlocksEndpointTargetingFields: """ @@ -1823,6 +2129,46 @@ class TestIsRequestBodySafeBlocksBedrockProjectOverride: ) +class TestIsRequestBodySafeBlocksVertexCredentialAlias: + @pytest.mark.parametrize("field", ["vertex_ai_credentials"]) + def test_field_in_request_body_is_rejected(self, field): + with pytest.raises(ValueError, match=field): + is_request_body_safe( + request_body={"model": "gpt-4", field: "attacker-supplied"}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + @pytest.mark.parametrize("field", ["vertex_ai_credentials"]) + def test_admin_opt_in_proxy_wide_allows(self, field): + assert ( + is_request_body_safe( + request_body={"model": "gpt-4", field: "byok-supplied"}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + def test_legitimate_request_body_param_still_allowed(self): + assert ( + is_request_body_safe( + request_body={ + "model": "gpt-4", + "temperature": 0.7, + "max_tokens": 128, + "user": "end-user-123", + }, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + class TestIsRequestBodySafeBlocksNVCFFunctionOverride: """``nvcf_function_id`` is rejected as a request-body param unless the admin opted in proxy-wide or per-deployment.""" diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 13041950f98..ffc5241d027 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -458,6 +458,118 @@ async def test_auth_builder_non_proxy_admin_user_role(): assert result["user_id"] == "test_user_1" +@pytest.mark.asyncio +@pytest.mark.parametrize( + "row_email,expected_email", + [ + ("row@example.com", "row@example.com"), + (None, "claim@example.com"), + ("", "claim@example.com"), + ], +) +async def test_auth_builder_result_includes_user_email(row_email, expected_email): + """LIT-4238: auth_builder must return user_email (user row wins, JWT claim + is the fallback) so the auth object and metrics get the email.""" + api_key = "test_jwt_token" + request_data = {"model": "gpt-4"} + general_settings = {"enforce_rbac": False} + route = "/chat/completions" + + user_object = LiteLLM_UserTable( + user_id="test_user_1", + user_email=row_email, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth() + + with ( + patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt, + patch.object(JWTAuthManager, "check_rbac_role", new_callable=AsyncMock), + patch.object(jwt_handler, "get_rbac_role", return_value=None), + patch.object(jwt_handler, "get_scopes", return_value=[]), + patch.object(jwt_handler, "get_object_id", return_value=None), + patch.object( + JWTAuthManager, + "get_user_info", + new_callable=AsyncMock, + return_value=("test_user_1", "claim@example.com", True), + ), + patch.object(jwt_handler, "get_org_id", return_value=None), + patch.object(jwt_handler, "get_end_user_id", return_value=None), + patch.object( + JWTAuthManager, + "check_admin_access", + new_callable=AsyncMock, + return_value=None, + ) as mock_check_admin, + patch.object( + JWTAuthManager, + "find_and_validate_specific_team_id", + new_callable=AsyncMock, + return_value=(None, None), + ), + patch.object(JWTAuthManager, "get_all_team_ids", return_value=set()), + patch.object( + JWTAuthManager, + "find_team_with_model_access", + new_callable=AsyncMock, + return_value=(None, None), + ), + patch.object( + JWTAuthManager, + "get_objects", + new_callable=AsyncMock, + return_value=(user_object, None, None, None, user_object.user_id), + ), + patch.object(JWTAuthManager, "map_user_to_teams", new_callable=AsyncMock), + patch.object(JWTAuthManager, "validate_object_id", return_value=True), + ): + mock_auth_jwt.return_value = {"sub": "test_user_1", "scope": ""} + + result = await JWTAuthManager.auth_builder( + api_key=api_key, + jwt_handler=jwt_handler, + request_data=request_data, + general_settings=general_settings, + route=route, + prisma_client=None, + user_api_key_cache=None, + parent_otel_span=None, + proxy_logging_obj=None, + ) + + assert result["user_email"] == expected_email + assert mock_check_admin.call_args.kwargs["user_email"] == "claim@example.com" + + +@pytest.mark.asyncio +async def test_check_admin_access_result_includes_user_email(): + """LIT-4238: the scope-based admin path has no user row, so the JWT claim + email must ride the JWTAuthBuilderResult.""" + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + admin_jwt_scope="litellm_proxy_admin", + admin_allowed_routes=["/chat/completions"], + ) + + result = await JWTAuthManager.check_admin_access( + jwt_handler=jwt_handler, + scopes=["litellm_proxy_admin"], + route="/chat/completions", + user_id="admin-user", + user_email="admin@example.com", + org_id=None, + api_key="test_jwt_token", + jwt_valid_token={"sub": "admin-user"}, + ) + + assert result is not None + assert result["is_proxy_admin"] is True + assert result["user_email"] == "admin@example.com" + + @pytest.mark.asyncio async def test_sync_user_role_and_teams(): from unittest.mock import MagicMock diff --git a/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py b/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py index fc0e9aec501..eb1135a240a 100644 --- a/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py +++ b/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py @@ -11,12 +11,22 @@ from unittest.mock import AsyncMock, patch import pytest from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.auth_utils import iter_request_fallback_targets from litellm.proxy.auth.user_api_key_auth import ( _enforce_key_and_fallback_model_access, - iter_router_fallback_model_names, + _fallback_target_model_name, ) +def _fallback_model_names(fallbacks): + """Model names the auth check validates for a top-level ``fallbacks`` value.""" + return [ + name + for target in iter_request_fallback_targets({"fallbacks": fallbacks}) + if (name := _fallback_target_model_name(target)) is not None + ] + + def _key_with_models(models: List[str]) -> UserAPIKeyAuth: return UserAPIKeyAuth( api_key="hashed", @@ -26,37 +36,40 @@ def _key_with_models(models: List[str]) -> UserAPIKeyAuth: ) -# ── iter_router_fallback_model_names ───────────────────────────────────────── +# ── fallback model-name extraction ─────────────────────────────────────────── -def testiter_router_fallback_model_names_router_config_shape(): +def test_fallback_model_names_router_config_shape(): """Router-config shape: ``[{primary: [fallback_list]}]``.""" - assert list( - iter_router_fallback_model_names( - [{"gpt-3.5-turbo": ["gpt-4", "claude-3"]}, {"gpt-4o": ["o1"]}] - ) + assert _fallback_model_names( + [{"gpt-3.5-turbo": ["gpt-4", "claude-3"]}, {"gpt-4o": ["o1"]}] ) == ["gpt-4", "claude-3", "o1"] -def testiter_router_fallback_model_names_simple_string_shape(): +def test_fallback_model_names_simple_string_shape(): """Simple top-level shape: list of strings.""" - assert list(iter_router_fallback_model_names(["gpt-4", "claude-3"])) == [ + assert _fallback_model_names(["gpt-4", "claude-3"]) == ["gpt-4", "claude-3"] + + +def test_fallback_model_names_client_side_shape(): + """ClientSideFallbackModel shape: ``[{"model": "..."}]``.""" + assert _fallback_model_names([{"model": "gpt-4"}, {"model": "claude-3"}]) == [ "gpt-4", "claude-3", ] -def testiter_router_fallback_model_names_client_side_shape(): - """ClientSideFallbackModel shape: ``[{"model": "..."}]``.""" - assert list( - iter_router_fallback_model_names([{"model": "gpt-4"}, {"model": "claude-3"}]) - ) == ["gpt-4", "claude-3"] +def test_fallback_model_names_nested_deployment_fallbacks(): + """A deployment target's own nested fallback field is unrolled too.""" + assert _fallback_model_names( + [{"primary": [{"model": "gpt-4", "fallbacks": [{"gpt-4": ["deepseek-chat"]}]}]}] + ) == ["gpt-4", "deepseek-chat"] -def testiter_router_fallback_model_names_empty_or_none(): - assert list(iter_router_fallback_model_names(None)) == [] - assert list(iter_router_fallback_model_names([])) == [] - assert list(iter_router_fallback_model_names("not a list")) == [] +def test_fallback_model_names_empty_or_none(): + assert _fallback_model_names(None) == [] + assert _fallback_model_names([]) == [] + assert _fallback_model_names("not a list") == [] # ── _enforce_key_and_fallback_model_access ──────────────────────────────────── @@ -200,6 +213,98 @@ async def test_top_level_fallback_fields_validated(fallback_field): assert "top-level-smuggled" in seen +@pytest.mark.asyncio +async def test_nested_deployment_fallback_inner_model_validated(): + """A model name nested several fallback rounds deep, inside a deployment + target's own ``fallbacks``, is extracted and passed to can_key_call_model.""" + valid_token = _key_with_models(["gpt-3.5-turbo"]) + request_data = { + "model": "gpt-3.5-turbo", + "fallbacks": [ + { + "gpt-3.5-turbo": [ + { + "model": "gpt-3.5-turbo", + "fallbacks": [{"gpt-3.5-turbo": ["deep-smuggled-model"]}], + } + ] + } + ], + } + + seen: List[str] = [] + + async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router): + seen.append(model) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + side_effect=fake_can_key_call_model, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model", + new=AsyncMock(), + ), + ): + await _enforce_key_and_fallback_model_access( + valid_token=valid_token, + request_data=request_data, + route="/v1/chat/completions", + request=None, + llm_model_list=None, + llm_router=None, + ) + + assert "deep-smuggled-model" in seen + + +@pytest.mark.asyncio +async def test_model_less_fallback_dict_is_skipped_never_passed_as_none(): + """A fallback target dict without a ``model`` key is skipped, never passed + as ``None`` into can_key_call_model / is_valid_fallback_model.""" + valid_token = _key_with_models(["gpt-3.5-turbo"]) + request_data = { + "model": "gpt-3.5-turbo", + "fallbacks": [ + { + "gpt-3.5-turbo": [ + {"model": "real-fallback"}, + {"api_base": "http://attacker"}, + "string-fallback", + ] + } + ], + } + + seen: List[str] = [] + + async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router): + seen.append(model) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + side_effect=fake_can_key_call_model, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model", + new=AsyncMock(), + ), + ): + await _enforce_key_and_fallback_model_access( + valid_token=valid_token, + request_data=request_data, + route="/v1/chat/completions", + request=None, + llm_model_list=None, + llm_router=None, + ) + + assert None not in seen + assert seen == ["gpt-3.5-turbo", "real-fallback", "string-fallback"] + + @pytest.mark.asyncio async def test_router_override_without_fallbacks_does_not_break_auth(): """``router_settings_override`` set without any fallback fields is a diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 59b8228530a..2c1948adca1 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1569,6 +1569,7 @@ class TestJWTOAuth2Coexistence: "token": jwt_token, "team_id": "jwt-team", "user_id": "jwt-human-user", + "user_email": None, "end_user_id": None, "org_id": None, "team_membership": None, @@ -1643,6 +1644,7 @@ class TestJWTOAuth2Coexistence: "token": jwt_token, "team_id": "validated-team", "user_id": "validated-user", + "user_email": "validated@example.com", "end_user_id": "validated-end-user", "org_id": "validated-org", "team_membership": None, @@ -1702,6 +1704,7 @@ class TestJWTOAuth2Coexistence: mock_auto_register.call_args.kwargs["end_user_id"] == "validated-end-user" ) assert result.org_id == "validated-org" + assert result.user_email == "validated@example.com" @pytest.mark.asyncio async def test_routing_override_routes_matching_jwt_to_oauth2(self): @@ -1788,6 +1791,7 @@ class TestJWTOAuth2Coexistence: "token": jwt_token, "team_id": "jwt-team", "user_id": "jwt-user-no-override", + "user_email": None, "end_user_id": None, "org_id": None, "team_membership": None, @@ -1988,6 +1992,7 @@ class TestJWTOAuth2Coexistence: "token": jwt_token, "team_id": "jwt-team", "user_id": "jwt-user-scope-mismatch", + "user_email": None, "end_user_id": None, "org_id": None, "team_membership": None, @@ -2296,6 +2301,7 @@ class TestJWTOAuth2Coexistence: "token": jwt_token, "team_id": None, "user_id": "jwt-admin-user", + "user_email": None, "end_user_id": None, "org_id": None, "team_membership": None, @@ -4255,6 +4261,98 @@ async def test_auth_does_not_rewrite_cached_key_object_back_into_cache(): setattr(_proxy_server_mod, k, v) +class TestJWTAuthUserEmail: + """JWT auth must populate `UserAPIKeyAuth.user_email` (LIT-4238); it feeds + the Prometheus `user_email` label and `user_api_key_user_email` in + StandardLogging/SpendLogs metadata, which were always None for JWT traffic.""" + + def _jwt_request(self, jwt_token): + mock_request = MagicMock() + mock_request.url.path = "/v1/chat/completions" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {jwt_token}"} + mock_request.query_params = {} + return mock_request + + async def _run_jwt_auth(self, mock_jwt_result, jwt_token): + with ( + patch( + "litellm.proxy.proxy_server.general_settings", + {"enable_jwt_auth": True}, + ), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.auth.user_api_key_auth.JWTAuthManager.auth_builder", + new_callable=AsyncMock, + return_value=mock_jwt_result, + ), + ): + litellm.proxy.proxy_server.jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + return await user_api_key_auth( + request=self._jwt_request(jwt_token), + api_key=f"Bearer {jwt_token}", + ) + + @pytest.mark.asyncio + async def test_jwt_auth_populates_user_email_on_valid_token(self): + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + mock_jwt_result = { + "is_proxy_admin": False, + "team_object": None, + "user_object": LiteLLM_UserTable( + user_id="jwt-human-user", + user_email="row@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ), + "end_user_object": None, + "org_object": None, + "token": jwt_token, + "team_id": None, + "user_id": "jwt-human-user", + "user_email": "resolved@example.com", + "end_user_id": None, + "org_id": None, + "team_membership": None, + "jwt_claims": {"sub": "user1"}, + } + + result = await self._run_jwt_auth(mock_jwt_result, jwt_token) + + assert result.user_id == "jwt-human-user" + assert result.user_email == "resolved@example.com" + + @pytest.mark.asyncio + async def test_jwt_auth_populates_user_email_on_proxy_admin(self): + jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + mock_jwt_result = { + "is_proxy_admin": True, + "team_object": None, + "user_object": None, + "end_user_object": None, + "org_object": None, + "token": jwt_token, + "team_id": None, + "user_id": "jwt-admin-user", + "user_email": "admin@example.com", + "end_user_id": None, + "org_id": None, + "team_membership": None, + "jwt_claims": {"sub": "user1"}, + } + + result = await self._run_jwt_auth(mock_jwt_result, jwt_token) + + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert result.user_id == "jwt-admin-user" + assert result.user_email == "admin@example.com" + + class TestCheckKeyModelBudgetWithFallback: """`_check_key_model_budget_with_fallback` must reroute a request to the first configured `budget_fallbacks` entry still within its own budget, @@ -4529,6 +4627,9 @@ async def test_temp_budget_increase_applied_for_cached_key(): Seed the auth cache with a key whose spend (5.0) exceeds its original max_budget (2.0) but is under the effective budget (2.0 + 100.0). The cache-hit request must not raise and the resolved token must carry max_budget == 102.0. + + Resolving twice must yield 102.0 both times and leave the cached object at the + original 2.0: the increase is derived per request, never compounded or persisted. """ from datetime import datetime, timedelta @@ -4574,14 +4675,22 @@ async def test_temp_budget_increase_applied_for_cached_key(): new_callable=AsyncMock, ), ): - result = await _user_api_key_auth_builder( - request=mock_request, - api_key=f"Bearer {api_key}", - azure_api_key_header="", - anthropic_api_key_header=None, - google_ai_studio_api_key_header=None, - azure_apim_header=None, - request_data={"model": "gpt-4o-mini"}, + results = tuple( + [ + await _user_api_key_auth_builder( + request=mock_request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={"model": "gpt-4o-mini"}, + ) + for _ in range(2) + ] ) - assert result.max_budget == 102.0 + assert all(result.max_budget == 102.0 for result in results) + + cached_after = await user_api_key_cache.async_get_cache(key=hashed_token) + assert cached_after.max_budget == 2.0 diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py index 6ab484d5004..f399b6f957a 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py @@ -48,27 +48,22 @@ def _base_config(**overrides: Any) -> AutorouteConfig: class TestParseDiscoveredModels: def test_parses_valid_raw_list_into_typed_tuple(self): raw = [ - { - "model_group": "gpt-4o", - "mode": "chat", - "input_cost_per_token": 0.01, - "output_cost_per_token": 0.02, - }, - {"model_group": "text-embedding-3-small", "mode": "embedding"}, + {"id": "gpt-4o", "object": "model", "mode": "chat"}, + {"id": "text-embedding-3-small", "object": "model", "mode": "embedding"}, ] result = parse_discovered_models(raw) assert result == ( - DiscoveredModel(name="gpt-4o", mode="chat", input_cost_per_token=0.01, output_cost_per_token=0.02), + DiscoveredModel(name="gpt-4o", mode="chat"), DiscoveredModel(name="text-embedding-3-small", mode="embedding"), ) def test_ignores_unknown_extra_fields(self): - raw = [{"model_group": "gpt-4o", "mode": "chat", "totally_unknown_field": "whatever"}] + raw = [{"id": "gpt-4o", "mode": "chat", "created": 123, "owned_by": "openai", "max_input_tokens": 128000}] result = parse_discovered_models(raw) assert result == (DiscoveredModel(name="gpt-4o", mode="chat"),) def test_missing_mode_defaults_to_chat(self): - raw = [{"model_group": "gpt-4o"}] + raw = [{"id": "gpt-4o", "object": "model"}] result = parse_discovered_models(raw) assert result[0].mode == "chat" diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py index 78d4bd20338..a17fed36f52 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py @@ -16,22 +16,22 @@ from litellm.proxy.client.cli.commands.autoroute.config import DiscoveredModel from litellm.proxy.client.cli.commands.autoroute.wizard import run_configure_wizard CHAT_AND_EMBEDDING_GROUPS: List[Dict[str, Any]] = [ - {"model_group": "gpt-4o-mini", "mode": "chat", "input_cost_per_token": 0.01, "output_cost_per_token": 0.02}, - {"model_group": "gpt-4o", "mode": "chat", "input_cost_per_token": 0.01, "output_cost_per_token": 0.02}, - {"model_group": "claude-opus", "mode": "chat"}, - {"model_group": "o1", "mode": "chat"}, - {"model_group": "text-embedding-3-small", "mode": "embedding"}, + {"id": "gpt-4o-mini", "object": "model", "mode": "chat", "max_input_tokens": 128000}, + {"id": "gpt-4o", "object": "model", "mode": "chat", "max_input_tokens": 128000}, + {"id": "claude-opus", "object": "model", "mode": "chat"}, + {"id": "o1", "object": "model", "mode": "chat"}, + {"id": "text-embedding-3-small", "object": "model", "mode": "embedding"}, ] CHAT_ONLY_GROUPS: List[Dict[str, Any]] = [ - {"model_group": "gpt-4o-mini", "mode": "chat"}, - {"model_group": "gpt-4o", "mode": "chat"}, - {"model_group": "claude-opus", "mode": "chat"}, - {"model_group": "o1", "mode": "chat"}, + {"id": "gpt-4o-mini", "object": "model", "mode": "chat"}, + {"id": "gpt-4o", "object": "model", "mode": "chat"}, + {"id": "claude-opus", "object": "model", "mode": "chat"}, + {"id": "o1", "object": "model", "mode": "chat"}, ] EMBEDDING_ONLY_GROUPS: List[Dict[str, Any]] = [ - {"model_group": "text-embedding-3-small", "mode": "embedding"}, + {"id": "text-embedding-3-small", "object": "model", "mode": "embedding"}, ] @@ -73,7 +73,7 @@ def _run( patch.object(wizard_module, "_render_and_prompt_for_models", side_effect=_fake_prompt_for_models), patch.object(wizard_module, "_render_and_prompt_for_model", side_effect=_fake_prompt_for_model), ): - mock_client_cls.return_value.model_groups.info.return_value = raw_groups + mock_client_cls.return_value.models.list.return_value = raw_groups result = runner.invoke( _invoke_wizard, obj={"base_url": "http://localhost:4000", "api_key": "sk-test"}, @@ -262,7 +262,7 @@ class TestRunConfigureWizardNoChatModels: assert result.exit_code != 0 assert result.exception is None or not isinstance(result.exception, AssertionError) - assert "Unexpected response from /model_group/info" in result.output + assert "Unexpected response from /v1/models" in result.output assert not config_path.exists() @@ -275,7 +275,7 @@ class TestRunConfigureWizardNotInteractive: patch.object(wizard_module, "CONFIG_PATH", config_path), patch.object(wizard_module, "_is_interactive", return_value=False), ): - mock_client_cls.return_value.model_groups.info.return_value = CHAT_AND_EMBEDDING_GROUPS + mock_client_cls.return_value.models.list.return_value = CHAT_AND_EMBEDDING_GROUPS result = runner.invoke(_invoke_wizard, obj={"base_url": "http://localhost:4000", "api_key": "sk-test"}) assert result.exit_code != 0 diff --git a/tests/test_litellm/proxy/common_utils/test_callback_utils.py b/tests/test_litellm/proxy/common_utils/test_callback_utils.py index 36ff3f3c399..8f390c096d7 100644 --- a/tests/test_litellm/proxy/common_utils/test_callback_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_callback_utils.py @@ -84,6 +84,52 @@ def test_process_callback_with_no_required_env_vars(mock_get_env_vars): assert result["variables"] == {} +@patch( + "litellm.proxy.common_utils.callback_utils.CustomLogger.get_callback_env_vars", + return_value=["LANGFUSE_PUBLIC_KEY", "LANGFUSE_SECRET_KEY"], +) +def test_process_callback_falls_back_to_process_env(mock_get_env_vars, monkeypatch): + """A callback env var set only in the process env must be surfaced. + + The logging integrations read their config from the process environment, so a + callback configured purely via env vars (IaC) is live even with no stored + entry. Reporting it as unset makes a working callback read as unconfigured. + """ + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "env-public-key") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "env-secret-key") + # stored config only carries the public key; the secret is env-only + environment_variables = {"LANGFUSE_PUBLIC_KEY": "db-public-key"} + + result = process_callback( + _callback="langfuse", + callback_type="success", + environment_variables=environment_variables, + ) + + # stored value wins; the env-only var is resolved rather than reported None + assert result["variables"] == { + "LANGFUSE_PUBLIC_KEY": "db-public-key", + "LANGFUSE_SECRET_KEY": "env-secret-key", + } + + +@patch( + "litellm.proxy.common_utils.callback_utils.CustomLogger.get_callback_env_vars", + return_value=["LANGFUSE_SECRET_KEY"], +) +def test_process_callback_reports_none_when_absent_everywhere(mock_get_env_vars, monkeypatch): + """A var set in neither the stored config nor the process env stays None.""" + monkeypatch.delenv("LANGFUSE_SECRET_KEY", raising=False) + + result = process_callback( + _callback="langfuse", + callback_type="success", + environment_variables={}, + ) + + assert result["variables"] == {"LANGFUSE_SECRET_KEY": None} + + def test_normalize_callback_names_none_returns_empty_list(): assert normalize_callback_names(None) == [] assert normalize_callback_names([]) == [] diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 5e348b1bb7e..be5bc74c385 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -5,25 +5,23 @@ import sys import time import types from datetime import datetime, timedelta, timezone +from datetime import time as dt_time from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock import pytest -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path from litellm._logging import verbose_proxy_logger from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob +from litellm.proxy.common_utils.timezone_utils import BudgetResetSettings from litellm.proxy.utils import ProxyLogging # Mock classes for testing class MockLiteLLMTeamMembership: - async def update_many( - self, where: Dict[str, Any], data: Dict[str, Any] - ) -> Dict[str, Any]: + async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: # Mock the update_many method for litellm_teammembership return {"count": 1} @@ -32,9 +30,7 @@ class MockLiteLLMVerificationToken: def __init__(self): self.update_many_calls: List[Dict[str, Any]] = [] - async def update_many( - self, where: Dict[str, Any], data: Dict[str, Any] - ) -> Dict[str, Any]: + async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: self.update_many_calls.append({"where": where, "data": data}) return {"count": 1} @@ -52,9 +48,7 @@ class MockLiteLLMOrganizationTable: self.find_many_calls.append({"where": where}) return self._find_many_results - async def update_many( - self, where: Dict[str, Any], data: Dict[str, Any] - ) -> Dict[str, Any]: + async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: self.update_many_calls.append({"where": where, "data": data}) return {"count": 1} @@ -72,9 +66,7 @@ class MockLiteLLMTagTable: self.find_many_calls.append({"where": where}) return self._find_many_results - async def update_many( - self, where: Dict[str, Any], data: Dict[str, Any] - ) -> Dict[str, Any]: + async def update_many(self, where: Dict[str, Any], data: Dict[str, Any]) -> Dict[str, Any]: self.update_many_calls.append({"where": where, "data": data}) return {"count": 1} @@ -110,9 +102,7 @@ class MockBatcher: _self._outer = outer def update(_self, where, data): - _self._outer.calls.append( - {"table": _self._table_name, "where": where, "data": data} - ) + _self._outer.calls.append({"table": _self._table_name, "where": where, "data": data}) self.litellm_verificationtoken = _Table("key", self) self.litellm_usertable = _Table("user", self) @@ -172,11 +162,7 @@ class MockPrismaClient: return [item for item in data if hasattr(item, "budget_reset_at")] # Handle specific filtering for enduser table queries - if ( - table_name == "enduser" - and query_type == "find_all" - and "budget_id_list" in kwargs - ): + if table_name == "enduser" and query_type == "find_all" and "budget_id_list" in kwargs: budget_id_list = kwargs["budget_id_list"] # Return endusers that match the budget IDs return [ @@ -188,11 +174,7 @@ class MockPrismaClient: ] # Handle key queries with expires and reset_at - if ( - table_name == "key" - and query_type == "find_all" - and ("expires" in kwargs or "reset_at" in kwargs) - ): + if table_name == "key" and query_type == "find_all" and ("expires" in kwargs or "reset_at" in kwargs): return [item for item in data if hasattr(item, "budget_reset_at")] return data @@ -227,9 +209,7 @@ def mock_proxy_logging(): @pytest.fixture def reset_budget_job(mock_prisma_client, mock_proxy_logging): - return ResetBudgetJob( - proxy_logging_obj=mock_proxy_logging, prisma_client=mock_prisma_client - ) + return ResetBudgetJob(proxy_logging_obj=mock_proxy_logging, prisma_client=mock_prisma_client) # Helper function to run async tests @@ -270,6 +250,40 @@ def test_reset_budget_for_key(reset_budget_job, mock_prisma_client): assert set(write["data"].keys()) == {"spend", "budget_reset_at"} +def test_reset_budget_for_key_honors_injected_reset_time(mock_prisma_client, mock_proxy_logging): + """Injected BudgetResetSettings drives the written reset time end to end (DI, no globals). + + Before the configurable-reset-time change this wrote a midnight reset_at (hour 0); + with noon injected it must write a noon reset_at. + """ + job = ResetBudgetJob( + proxy_logging_obj=mock_proxy_logging, + prisma_client=mock_prisma_client, + reset_settings=BudgetResetSettings(timezone="UTC", reset_time_of_day=dt_time(12, 0)), + ) + now = datetime.now(timezone.utc) + test_key = type( + "LiteLLM_VerificationToken", + (), + { + "spend": 100.0, + "budget_duration": "1d", + "budget_reset_at": now, + "id": "test-key-noon", + "token": "tok-noon", + }, + ) + mock_prisma_client.data["key"] = [test_key] + + asyncio.run(job.reset_budget_for_litellm_keys()) + + key_writes = [c for c in mock_prisma_client.db.batch_calls if c["table"] == "key"] + assert len(key_writes) == 1 + reset_at = key_writes[0]["data"]["budget_reset_at"].astimezone(timezone.utc) + assert reset_at.hour == 12 + assert reset_at.minute == 0 + + def test_reset_budget_for_user(reset_budget_job, mock_prisma_client): # Setup test data with timezone-aware datetime now = datetime.now(timezone.utc) @@ -486,11 +500,7 @@ def test_reset_budget_for_keys_linked_to_budgets(reset_budget_job, mock_prisma_c budgets_to_reset = [test_budget] # Run the method - asyncio.run( - reset_budget_job.reset_budget_for_keys_linked_to_budgets( - budgets_to_reset=budgets_to_reset - ) - ) + asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)) # Verify that update_many was called on litellm_verificationtoken calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls @@ -531,11 +541,7 @@ def test_reset_budget_for_keys_linked_to_budgets_excludes_keys_with_own_budget_d budgets_to_reset = [test_budget] - asyncio.run( - reset_budget_job.reset_budget_for_keys_linked_to_budgets( - budgets_to_reset=budgets_to_reset - ) - ) + asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=budgets_to_reset)) calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls assert len(calls) == 1 @@ -548,17 +554,13 @@ def test_reset_budget_for_keys_linked_to_budgets_excludes_keys_with_own_budget_d assert call["where"]["budget_id"] == {"in": ["7d-budget-tier"]} -def test_reset_budget_for_keys_linked_to_budgets_empty( - reset_budget_job, mock_prisma_client -): +def test_reset_budget_for_keys_linked_to_budgets_empty(reset_budget_job, mock_prisma_client): """ Test that when there are no budgets to reset, no update is performed on the verification token table. """ # Run with empty list - asyncio.run( - reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=[]) - ) + asyncio.run(reset_budget_job.reset_budget_for_keys_linked_to_budgets(budgets_to_reset=[])) # Verify no update_many calls were made calls = mock_prisma_client.db.litellm_verificationtoken.update_many_calls @@ -584,11 +586,7 @@ def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_c }, ) - asyncio.run( - reset_budget_job.reset_budget_for_orgs_linked_to_budgets( - budgets_to_reset=[test_budget] - ) - ) + asyncio.run(reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[test_budget])) calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls assert len(calls) == 1 @@ -598,16 +596,12 @@ def test_reset_budget_for_orgs_linked_to_budgets(reset_budget_job, mock_prisma_c assert call["data"]["spend"] == 0 -def test_reset_budget_for_orgs_linked_to_budgets_empty( - reset_budget_job, mock_prisma_client -): +def test_reset_budget_for_orgs_linked_to_budgets_empty(reset_budget_job, mock_prisma_client): """ Test that when there are no budgets to reset, no update is performed on the organization table. """ - asyncio.run( - reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[]) - ) + asyncio.run(reset_budget_job.reset_budget_for_orgs_linked_to_budgets(budgets_to_reset=[])) calls = mock_prisma_client.db.litellm_organizationtable.update_many_calls assert len(calls) == 0 @@ -631,11 +625,7 @@ def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_c }, ) - asyncio.run( - reset_budget_job.reset_budget_for_tags_linked_to_budgets( - budgets_to_reset=[test_budget] - ) - ) + asyncio.run(reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[test_budget])) calls = mock_prisma_client.db.litellm_tagtable.update_many_calls assert len(calls) == 1 @@ -645,16 +635,12 @@ def test_reset_budget_for_tags_linked_to_budgets(reset_budget_job, mock_prisma_c assert call["data"]["spend"] == 0 -def test_reset_budget_for_tags_linked_to_budgets_empty( - reset_budget_job, mock_prisma_client -): +def test_reset_budget_for_tags_linked_to_budgets_empty(reset_budget_job, mock_prisma_client): """ Test that when there are no budgets to reset, no update is performed on the tag table. """ - asyncio.run( - reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[]) - ) + asyncio.run(reset_budget_job.reset_budget_for_tags_linked_to_budgets(budgets_to_reset=[])) calls = mock_prisma_client.db.litellm_tagtable.update_many_calls assert len(calls) == 0 @@ -668,9 +654,7 @@ def test_reset_budget_for_tags_linked_to_budgets_empty( ], ids=["30d-calendar-month", "1mo-calendar-month", "1d-next-midnight"], ) -def test_reset_budget_reset_at_date_calendar_aligned( - budget_duration, expected_day, expected_month -): +def test_reset_budget_reset_at_date_calendar_aligned(budget_duration, expected_day, expected_month): """ Verify that _reset_budget_reset_at_date produces calendar-aligned reset times (matching get_budget_reset_time), not sliding-window offsets. @@ -694,7 +678,7 @@ def test_reset_budget_reset_at_date_calendar_aligned( with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: mock_dt.now.return_value = fixed_now mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now)) + asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings())) assert test_budget.budget_reset_at.day == expected_day assert test_budget.budget_reset_at.month == expected_month @@ -724,7 +708,7 @@ def test_reset_budget_reset_at_date_7d_next_monday(): with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: mock_dt.now.return_value = fixed_now mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now)) + asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings())) # Next Monday after Wednesday June 14 is June 19 assert test_budget.budget_reset_at.day == 19 @@ -749,7 +733,7 @@ def test_reset_budget_reset_at_date_none_duration(): }, ) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, now)) + asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, now, BudgetResetSettings())) assert test_budget.budget_reset_at == original_reset_at @@ -773,7 +757,7 @@ def test_reset_budget_reset_at_date_none_reset_at(): with patch("litellm.proxy.common_utils.timezone_utils.datetime") as mock_dt: mock_dt.now.return_value = fixed_now mock_dt.side_effect = lambda *args, **kwargs: datetime(*args, **kwargs) - asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now)) + asyncio.run(ResetBudgetJob._reset_budget_reset_at_date(test_budget, fixed_now, BudgetResetSettings())) # Should be set to 1st of next month (July 1) assert test_budget.budget_reset_at is not None @@ -781,9 +765,7 @@ def test_reset_budget_reset_at_date_none_reset_at(): assert test_budget.budget_reset_at.month == 7 -def test_budget_table_reset_also_resets_linked_keys( - reset_budget_job, mock_prisma_client -): +def test_budget_table_reset_also_resets_linked_keys(reset_budget_job, mock_prisma_client): """ Integration-style test: when reset_budget_for_litellm_budget_table runs, it should also reset spend for keys linked to the expiring budget tiers @@ -818,9 +800,7 @@ def test_budget_table_reset_also_resets_linked_keys( assert calls[0]["data"]["spend"] == 0 -def test_budget_table_reset_also_resets_linked_orgs( - reset_budget_job, mock_prisma_client -): +def test_budget_table_reset_also_resets_linked_orgs(reset_budget_job, mock_prisma_client): """ Integration-style test: when reset_budget_for_litellm_budget_table runs, it should also reset spend for orgs linked to the expiring budget tiers @@ -853,9 +833,7 @@ def test_budget_table_reset_also_resets_linked_orgs( assert calls[0]["data"]["spend"] == 0 -def test_budget_table_reset_also_resets_linked_tags( - reset_budget_job, mock_prisma_client -): +def test_budget_table_reset_also_resets_linked_tags(reset_budget_job, mock_prisma_client): """ Integration-style test: when reset_budget_for_litellm_budget_table runs, it should also reset spend for tags linked to the expiring budget tiers. @@ -887,9 +865,7 @@ def test_budget_table_reset_also_resets_linked_tags( assert calls[0]["data"]["spend"] == 0 -def test_reset_budget_resets_endusers_with_null_budget_id( - reset_budget_job, mock_prisma_client -): +def test_reset_budget_resets_endusers_with_null_budget_id(reset_budget_job, mock_prisma_client): """ When litellm.max_end_user_budget_id is configured and that budget is being reset, end users with budget_id=NULL should also have their spend @@ -959,17 +935,13 @@ def test_reset_budget_resets_endusers_with_null_budget_id( mock_prisma_client.data["enduser"] = [enduser_with_budget] # Set up the DB mock for NULL-budget-id end users - mock_prisma_client.db.litellm_endusertable.set_find_many_results( - [enduser_no_budget_row] - ) + mock_prisma_client.db.litellm_endusertable.set_find_many_results([enduser_no_budget_row]) asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table()) # Both end users should have been reset updated = mock_prisma_client.updated_data["enduser"] - assert ( - len(updated) == 2 - ), f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}" + assert len(updated) == 2, f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}" user_ids = {u.user_id for u in updated} assert "enduser-explicit" in user_ids @@ -986,9 +958,7 @@ def test_reset_budget_resets_endusers_with_null_budget_id( litellm.max_end_user_budget_id = None -def test_reset_budget_skips_null_budget_id_endusers_when_default_not_configured( - reset_budget_job, mock_prisma_client -): +def test_reset_budget_skips_null_budget_id_endusers_when_default_not_configured(reset_budget_job, mock_prisma_client): """ When litellm.max_end_user_budget_id is NOT configured, end users with budget_id=NULL should NOT be fetched or reset. @@ -1073,20 +1043,14 @@ def test_reset_budget_for_team_members_preserves_total_spend(): mock_prisma_client = MagicMock() mock_prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[]) - mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock( - return_value={"count": 1} - ) + mock_prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1}) - job = ResetBudgetJob( - proxy_logging_obj=MagicMock(), prisma_client=mock_prisma_client - ) + job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=mock_prisma_client) asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) mock_prisma_client.db.litellm_teammembership.update_many.assert_called_once() - call_kwargs = ( - mock_prisma_client.db.litellm_teammembership.update_many.call_args.kwargs - ) + call_kwargs = mock_prisma_client.db.litellm_teammembership.update_many.call_args.kwargs assert call_kwargs["where"]["budget_id"]["in"] == ["budget-1"] assert call_kwargs["data"] == {"spend": 0} assert "total_spend" not in call_kwargs["data"] @@ -1142,9 +1106,7 @@ def test_reset_budget_windows_uses_is_not_null_filter(monkeypatch): raises `MissingRequiredValueError`. We work around it by using `query_raw` with `IS NOT NULL`. If someone reverts to the ORM filter, this test fails. """ - job, prisma_client, _ = _make_reset_budget_windows_job( - monkeypatch, key_rows=[], team_rows=[] - ) + job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=[], team_rows=[]) asyncio.run(job.reset_budget_windows()) @@ -1184,15 +1146,11 @@ def test_reset_budget_windows_resets_expired_key_window(monkeypatch): # The `budget_limits` payload is re-serialized JSON with a bumped reset_at. written_windows = json.loads(call_kwargs["data"]["budget_limits"]) assert len(written_windows) == 1 - new_reset_at = datetime.fromisoformat( - written_windows[0]["reset_at"].replace("Z", "+00:00") - ).replace(tzinfo=None) + new_reset_at = datetime.fromisoformat(written_windows[0]["reset_at"].replace("Z", "+00:00")).replace(tzinfo=None) assert new_reset_at > now # The spend counter for this key+window was cleared. - spend_counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:key:sk-expired:window:1d", value=0.0 - ) + spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-expired:window:1d", value=0.0) def test_reset_budget_windows_skips_unexpired_key_window(monkeypatch): @@ -1206,9 +1164,7 @@ def test_reset_budget_windows_skips_unexpired_key_window(monkeypatch): "budget_limits": [{"budget_duration": "1d", "reset_at": future}], } ] - job, prisma_client, _ = _make_reset_budget_windows_job( - monkeypatch, key_rows=key_rows, team_rows=[] - ) + job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[]) asyncio.run(job.reset_budget_windows()) @@ -1237,9 +1193,7 @@ def test_reset_budget_windows_resets_expired_team_window(monkeypatch): assert call_kwargs["where"] == {"team_id": "team-expired"} assert "budget_limits" in call_kwargs["data"] - spend_counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:team:team-expired:window:30d", value=0.0 - ) + spend_counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team:team-expired:window:30d", value=0.0) def test_reset_budget_windows_handles_string_budget_limits(monkeypatch): @@ -1252,14 +1206,10 @@ def test_reset_budget_windows_handles_string_budget_limits(monkeypatch): key_rows = [ { "token": "sk-string-limits", - "budget_limits": json.dumps( - [{"budget_duration": "1d", "reset_at": expired}] - ), + "budget_limits": json.dumps([{"budget_duration": "1d", "reset_at": expired}]), } ] - job, prisma_client, _ = _make_reset_budget_windows_job( - monkeypatch, key_rows=key_rows, team_rows=[] - ) + job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[]) asyncio.run(job.reset_budget_windows()) @@ -1274,9 +1224,7 @@ def test_reset_budget_windows_skips_row_with_empty_budget_limits(monkeypatch): {"token": "sk-empty-list", "budget_limits": []}, {"token": "sk-empty-str", "budget_limits": ""}, ] - job, prisma_client, _ = _make_reset_budget_windows_job( - monkeypatch, key_rows=key_rows, team_rows=[] - ) + job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[]) asyncio.run(job.reset_budget_windows()) @@ -1361,27 +1309,17 @@ def test_reset_budget_for_team_members_invalidates_redis_counter(monkeypatch): ) prisma_client = MagicMock() - prisma_client.db.litellm_teammembership.find_many = AsyncMock( - return_value=[membership] - ) - prisma_client.db.litellm_teammembership.update_many = AsyncMock( - return_value={"count": 1} - ) + prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership]) + prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1}) job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:team_member:alice:team-x", value=0.0, ttl=60 - ) - counter_cache.redis_cache.async_set_cache.assert_any_await( - key="spend:team_member:alice:team-x", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team_member:alice:team-x", value=0.0, ttl=60) + counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:team_member:alice:team-x", value=0.0, ttl=60) -def test_reset_budget_for_keys_invalidates_redis_counter( - reset_budget_job, mock_prisma_client, monkeypatch -): +def test_reset_budget_for_keys_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch): """Key budget reset must clear the Redis spend counter.""" counter_cache = _make_counter_invalidation_job(monkeypatch) @@ -1402,14 +1340,10 @@ def test_reset_budget_for_keys_invalidates_redis_counter( asyncio.run(reset_budget_job.reset_budget_for_litellm_keys()) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:key:sk-abc", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-abc", value=0.0, ttl=60) -def test_reset_budget_for_users_invalidates_redis_counter( - reset_budget_job, mock_prisma_client, monkeypatch -): +def test_reset_budget_for_users_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch): """User budget reset must clear the Redis spend counter.""" counter_cache = _make_counter_invalidation_job(monkeypatch) @@ -1430,14 +1364,10 @@ def test_reset_budget_for_users_invalidates_redis_counter( asyncio.run(reset_budget_job.reset_budget_for_litellm_users()) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:user:alice", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:user:alice", value=0.0, ttl=60) -def test_reset_budget_for_teams_invalidates_redis_counter( - reset_budget_job, mock_prisma_client, monkeypatch -): +def test_reset_budget_for_teams_invalidates_redis_counter(reset_budget_job, mock_prisma_client, monkeypatch): """Team budget reset must clear the Redis spend counter.""" counter_cache = _make_counter_invalidation_job(monkeypatch) @@ -1458,9 +1388,7 @@ def test_reset_budget_for_teams_invalidates_redis_counter( asyncio.run(reset_budget_job.reset_budget_for_litellm_teams()) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:team:team-x", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:team:team-x", value=0.0, ttl=60) def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch): @@ -1511,9 +1439,7 @@ def test_reset_does_not_zero_counter_when_db_write_fails(monkeypatch): batcher.commit = failing_commit prisma_client.db.batch_ = MagicMock(return_value=batcher) - job = ResetBudgetJob( - proxy_logging_obj=MockProxyLogging(), prisma_client=prisma_client - ) + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_litellm_keys()) @@ -1543,8 +1469,8 @@ def test_reset_budget_for_keys_writes_only_spend_and_reset_at(reset_budget_job, "budget_duration": "30d", "budget_reset_at": now, "token": "sk-problematic", - "object_permission_id": "perm-abc", # would be rejected on update - "budget_limits": [{"max_budget": 5}], # would be rejected on update + "object_permission_id": "perm-abc", # would be rejected on update + "budget_limits": [{"max_budget": 5}], # would be rejected on update "metadata": {"some": "thing"}, }, ) @@ -1570,19 +1496,13 @@ def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monke linked_key = type("Key", (), {"token": "sk-linked"}) prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[linked_key] - ) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 1} - ) + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[linked_key]) + prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(return_value={"count": 1}) job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget])) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:key:sk-linked", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:key:sk-linked", value=0.0, ttl=60) def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monkeypatch): @@ -1593,22 +1513,14 @@ def test_reset_budget_for_orgs_linked_to_budgets_invalidates_redis_counter(monke linked_org = type("Org", (), {"organization_id": "org-acme"}) prisma_client = MagicMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[linked_org] - ) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 1} - ) + prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[linked_org]) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock(return_value={"count": 1}) job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget])) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:org:org-acme", value=0.0, ttl=60 - ) - counter_cache.redis_cache.async_set_cache.assert_any_await( - key="spend:org:org-acme", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:org:org-acme", value=0.0, ttl=60) + counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:org:org-acme", value=0.0, ttl=60) def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monkeypatch): @@ -1625,12 +1537,8 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_redis_counter(monke job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) - counter_cache.in_memory_cache.set_cache.assert_any_call( - key="spend:tag:tenant-42", value=0.0, ttl=60 - ) - counter_cache.redis_cache.async_set_cache.assert_any_await( - key="spend:tag:tenant-42", value=0.0, ttl=60 - ) + counter_cache.in_memory_cache.set_cache.assert_any_call(key="spend:tag:tenant-42", value=0.0, ttl=60) + counter_cache.redis_cache.async_set_cache.assert_any_await(key="spend:tag:tenant-42", value=0.0, ttl=60) def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache( @@ -1657,9 +1565,7 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_management_cache( job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) - counter_cache.user_api_key_cache.async_delete_cache.assert_any_await( - key="tag:tenant-42" - ) + counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="tag:tenant-42") def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management_cache( @@ -1684,8 +1590,7 @@ def test_reset_budget_for_tags_linked_to_budgets_invalidates_each_tag_management asyncio.run(job.reset_budget_for_tags_linked_to_budgets([expired_budget])) deleted_keys = { - call.kwargs.get("key") - for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list + call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list } assert deleted_keys == {"tag:tenant-a", "tag:tenant-b", "tag:tenant-c"} @@ -1711,19 +1616,13 @@ def test_reset_budget_for_keys_linked_to_budgets_invalidates_management_cache( linked_key = type("Key", (), {"token": "sk-linked"}) prisma_client = MagicMock() - prisma_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[linked_key] - ) - prisma_client.db.litellm_verificationtoken.update_many = AsyncMock( - return_value={"count": 1} - ) + prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[linked_key]) + prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(return_value={"count": 1}) job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget])) - counter_cache.user_api_key_cache.async_delete_cache.assert_any_await( - key="sk-linked" - ) + counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="sk-linked") def test_reset_budget_for_orgs_linked_to_budgets_invalidates_management_cache( @@ -1736,19 +1635,14 @@ def test_reset_budget_for_orgs_linked_to_budgets_invalidates_management_cache( linked_org = type("Org", (), {"organization_id": "org-acme"}) prisma_client = MagicMock() - prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[linked_org] - ) - prisma_client.db.litellm_organizationtable.update_many = AsyncMock( - return_value={"count": 1} - ) + prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[linked_org]) + prisma_client.db.litellm_organizationtable.update_many = AsyncMock(return_value={"count": 1}) job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_orgs_linked_to_budgets([expired_budget])) deleted_keys = { - call.kwargs.get("key") - for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list + call.kwargs.get("key") for call in counter_cache.user_api_key_cache.async_delete_cache.await_args_list } assert deleted_keys == { "org_id:org-acme", @@ -1768,19 +1662,13 @@ def test_reset_budget_for_team_members_invalidates_management_cache(monkeypatch) ) prisma_client = MagicMock() - prisma_client.db.litellm_teammembership.find_many = AsyncMock( - return_value=[membership] - ) - prisma_client.db.litellm_teammembership.update_many = AsyncMock( - return_value={"count": 1} - ) + prisma_client.db.litellm_teammembership.find_many = AsyncMock(return_value=[membership]) + prisma_client.db.litellm_teammembership.update_many = AsyncMock(return_value={"count": 1}) job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client) asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget])) - counter_cache.user_api_key_cache.async_delete_cache.assert_any_await( - key="team-x_alice" - ) + counter_cache.user_api_key_cache.async_delete_cache.assert_any_await(key="team-x_alice") def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure_still_resets( @@ -1788,9 +1676,7 @@ def test_reset_budget_for_tags_linked_to_budgets_management_cache_delete_failure ): """If ``async_delete_cache`` raises, the DB cascade must still complete.""" counter_cache = _make_counter_invalidation_job(monkeypatch) - counter_cache.user_api_key_cache.async_delete_cache = AsyncMock( - side_effect=RuntimeError("cache unavailable") - ) + counter_cache.user_api_key_cache.async_delete_cache = AsyncMock(side_effect=RuntimeError("cache unavailable")) expired_budget = type("B", (), {"budget_id": "budget-1"}) linked_tag = type("Tag", (), {"tag_name": "tenant-42"}) diff --git a/tests/test_litellm/proxy/common_utils/test_timezone_utils.py b/tests/test_litellm/proxy/common_utils/test_timezone_utils.py index 80b813226df..7f686c53c95 100644 --- a/tests/test_litellm/proxy/common_utils/test_timezone_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_timezone_utils.py @@ -1,19 +1,33 @@ import os import sys -from datetime import datetime, timezone +from datetime import datetime, time, timezone from zoneinfo import ZoneInfo +import pytest + sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path import litellm from litellm.proxy.common_utils.timezone_utils import ( + BudgetResetSettings, + compute_budget_reset_at, + get_budget_reset_settings, get_budget_reset_time, get_budget_reset_timezone, + parse_budget_reset_time, ) +def _restore_attr(obj, name, original): + if original is None: + if hasattr(obj, name): + delattr(obj, name) + else: + setattr(obj, name, original) + + def test_get_budget_reset_time(): """ Test that the budget reset time is set to the first of the next month @@ -100,3 +114,69 @@ def test_get_budget_reset_time_respects_timezone(): delattr(litellm, "timezone") else: litellm.timezone = original + + +def test_parse_budget_reset_time_hh_mm(): + assert parse_budget_reset_time("12:00") == time(12, 0) + + +def test_parse_budget_reset_time_hh_mm_ss(): + assert parse_budget_reset_time("09:30:15") == time(9, 30, 15) + + +def test_parse_budget_reset_time_unset_defaults_to_midnight(): + assert parse_budget_reset_time(None) == time(0, 0) + assert parse_budget_reset_time("") == time(0, 0) + + +def test_parse_budget_reset_time_invalid_string_raises(): + with pytest.raises(ValueError): + parse_budget_reset_time("25:00") + with pytest.raises(ValueError): + parse_budget_reset_time("noon") + + +def test_parse_budget_reset_time_non_string_raises(): + # Unquoted "12:00" in YAML parses to the int 720; it must fail loudly, + # not silently fall back to midnight. + with pytest.raises(ValueError): + parse_budget_reset_time(720) + + +def test_get_budget_reset_settings_reads_globals(): + orig_tz = getattr(litellm, "timezone", None) + orig_rt = getattr(litellm, "budget_reset_time", None) + try: + litellm.timezone = "Asia/Jerusalem" + litellm.budget_reset_time = "12:00" + settings = get_budget_reset_settings() + assert settings.timezone == "Asia/Jerusalem" + assert settings.reset_time_of_day == time(12, 0) + finally: + _restore_attr(litellm, "timezone", orig_tz) + _restore_attr(litellm, "budget_reset_time", orig_rt) + + +def test_compute_budget_reset_at_applies_offset(): + settings = BudgetResetSettings( + timezone="Asia/Jerusalem", reset_time_of_day=time(12, 0) + ) + reset_at = compute_budget_reset_at("1d", settings) + jerusalem = reset_at.astimezone(ZoneInfo("Asia/Jerusalem")) + assert jerusalem.hour == 12 + assert jerusalem.minute == 0 + assert reset_at > datetime.now(timezone.utc) + + +def test_get_budget_reset_time_honors_global_budget_reset_time(): + orig_tz = getattr(litellm, "timezone", None) + orig_rt = getattr(litellm, "budget_reset_time", None) + try: + litellm.timezone = "UTC" + litellm.budget_reset_time = "12:00" + reset_at = get_budget_reset_time(budget_duration="1d") + assert reset_at.astimezone(timezone.utc).hour == 12 + assert reset_at.astimezone(timezone.utc).minute == 0 + finally: + _restore_attr(litellm, "timezone", orig_tz) + _restore_attr(litellm, "budget_reset_time", orig_rt) diff --git a/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py b/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py new file mode 100644 index 00000000000..20bea98351f --- /dev/null +++ b/tests/test_litellm/proxy/config_resolvers/test_config_resolvers.py @@ -0,0 +1,105 @@ +import os + +from litellm.proxy.config_resolvers._descriptors import FieldDescriptor, resolve_fields +from litellm.proxy.config_resolvers.sso import ( + SSO_FIELD_ENV_VARS, + SSO_SECRET_FIELDS, + resolve_sso_config, +) + +_D = ( + FieldDescriptor("client_id", "client_id", "CLIENT_ID"), + FieldDescriptor("scope", "scope", "SCOPE", default="openid"), +) + + +def test_resolve_fields_db_wins_over_env(): + values, provenance = resolve_fields(_D, {"client_id": "from-db"}, {"CLIENT_ID": "from-env"}) + assert values["client_id"] == "from-db" + assert provenance["client_id"] == "db" + + +def test_resolve_fields_blank_db_falls_back_to_env(): + values, provenance = resolve_fields(_D, {"client_id": " "}, {"CLIENT_ID": "from-env"}) + assert values["client_id"] == "from-env" + assert provenance["client_id"] == "env" + + +def test_resolve_fields_blank_everywhere_falls_to_default(): + values, provenance = resolve_fields(_D, {}, {"SCOPE": ""}) + assert values["scope"] == "openid" + assert provenance["scope"] == "default" + + +def test_resolve_fields_unset_everywhere(): + values, provenance = resolve_fields(_D, {}, {}) + assert values["client_id"] is None + assert provenance["client_id"] == "unset" + + +def test_resolve_fields_empty_db_absent_by_default_falls_to_env(): + # SSO semantics: a present-but-empty stored value is absent, so env wins. + values, provenance = resolve_fields(_D, {"client_id": ""}, {"CLIENT_ID": "from-env"}) + assert values["client_id"] == "from-env" + assert provenance["client_id"] == "env" + + +def test_resolve_fields_empty_db_is_explicit_clear_when_flag_set(): + # Alerting semantics: a present-but-empty stored value is an explicit clear + # that must win over a stale env var. + values, provenance = resolve_fields( + _D, {"client_id": ""}, {"CLIENT_ID": "stale-env"}, empty_db_is_set=True + ) + assert values["client_id"] == "" + assert provenance["client_id"] == "db" + + +def test_sso_descriptor_mapping_is_single_sourced(): + # The write path and read path both consume this mapping; it must cover every + # env-backed SSO field and map to the uppercase env var. + assert SSO_FIELD_ENV_VARS["generic_client_id"] == "GENERIC_CLIENT_ID" + assert SSO_SECRET_FIELDS == frozenset( + {"google_client_secret", "microsoft_client_secret", "generic_client_secret"} + ) + + +def test_resolve_sso_config_returns_unmasked_secret_and_provenance(): + # The resolver hands back plaintext; masking is the endpoint's job. If the + # resolver masked, the login path would consume a masked secret and fail. + resolved = resolve_sso_config( + {"generic_client_secret": "super-secret-value"}, + {"GENERIC_CLIENT_ID": "env-id"}, + ) + assert resolved.config.generic_client_secret == "super-secret-value" + assert resolved.provenance["generic_client_secret"] == "db" + assert resolved.config.generic_client_id == "env-id" + assert resolved.provenance["generic_client_id"] == "env" + + +def test_resolve_sso_config_parses_structured_mappings(): + resolved = resolve_sso_config( + { + "generic_client_id": "id", + "role_mappings": { + "provider": "generic", + "group_claim": "groups", + "default_role": "internal_user", + "roles": {}, + }, + "team_mappings": {"team_ids_jwt_field": "teams"}, + }, + {}, + ) + assert resolved.config.role_mappings is not None + assert resolved.config.role_mappings.group_claim == "groups" + assert resolved.config.team_mappings is not None + assert resolved.config.team_mappings.team_ids_jwt_field == "teams" + + +def test_resolve_sso_config_does_not_mutate_os_environ(monkeypatch): + # Unlike the legacy read path, resolving must not write os.environ. + monkeypatch.delenv("GENERIC_CLIENT_ID", raising=False) + before = dict(os.environ) + resolve_sso_config({"generic_client_id": "id-from-db"}, os.environ) + assert dict(os.environ) == before + assert "GENERIC_CLIENT_ID" not in os.environ diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 15827b80bcf..53f32fd96fb 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -3274,3 +3274,343 @@ async def test_chat_completion_modify_response_exception_streaming_logging_obj_n # CustomStreamWrapper would raise AttributeError inside __init__ and this # call would never reach here. assert response is not None + + +class TestBedrockOnlyScanNewMessages: + """Bedrock apply_guardrail honors only_scan_new_messages: scans only the per-session diff. + + apply_guardrail is the path the proxy actually runs for Bedrock (via the unified + guardrail interface), so these tests exercise it directly rather than the legacy + async_pre_call_hook. Each test uses a unique session id to isolate the process-wide + incremental cache. + """ + + def _guardrail(self): + return BedrockGuardrail( + guardrail_name="bedrock-incremental", + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + default_on=True, + only_scan_new_messages=True, + ) + + @pytest.mark.asyncio + async def test_second_turn_scans_only_new_messages(self): + guardrail = self._guardrail() + session = {"litellm_session_id": "sess-bedrock-diff"} + bedrock_none = {"action": "NONE", "output": [], "outputs": []} + + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = bedrock_none + + await guardrail.apply_guardrail( + inputs={"texts": ["be helpful", "first question"]}, + request_data=session, + input_type="request", + ) + assert mock_api.call_count == 1 + first_scanned = mock_api.call_args.kwargs["messages"] + assert [m["content"] for m in first_scanned] == ["be helpful", "first question"] + + mock_api.reset_mock() + + await guardrail.apply_guardrail( + inputs={"texts": ["be helpful", "first question", "first answer", "second question"]}, + request_data=session, + input_type="request", + ) + assert mock_api.call_count == 1 + second_scanned = mock_api.call_args.kwargs["messages"] + assert [m["content"] for m in second_scanned] == ["first answer", "second question"] + + @pytest.mark.asyncio + async def test_identical_resend_skips_api_call(self): + guardrail = self._guardrail() + session = {"litellm_session_id": "sess-bedrock-resend"} + + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + + await guardrail.apply_guardrail( + inputs={"texts": ["only question"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 1 + + mock_api.reset_mock() + result = await guardrail.apply_guardrail( + inputs={"texts": ["only question"]}, request_data=session, input_type="request" + ) + mock_api.assert_not_called() + assert result["texts"] == ["only question"] + + @pytest.mark.asyncio + async def test_no_session_id_scans_full_context(self): + guardrail = self._guardrail() + + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + + await guardrail.apply_guardrail( + inputs={"texts": ["q1", "a1", "q2"]}, + request_data={"metadata": {}}, + input_type="request", + ) + assert mock_api.call_count == 1 + scanned = mock_api.call_args.kwargs["messages"] + assert [m["content"] for m in scanned] == ["q1", "a1", "q2"] + + @pytest.mark.asyncio + async def test_masking_guardrail_falls_back_and_does_not_persist(self): + """A guardrail that anonymizes content must not be short-circuited. + + Regression: the incremental fast path used to ignore the guardrail response, + so masked/anonymized output was dropped, the raw text reached the model, and + the segment was marked scanned so it was never re-checked. Detecting masked + output must force a full-context scan (which applies the masking) and must not + persist session state, so an identical resend is scanned again. + """ + guardrail = self._guardrail() + session = {"litellm_session_id": "sess-bedrock-mask"} + masked = { + "action": "GUARDRAIL_INTERVENED", + "output": [], + "outputs": [{"text": "my ssn is [REDACTED]"}], + } + + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = masked + + result = await guardrail.apply_guardrail( + inputs={"texts": ["my ssn is 123-45-6789"]}, + request_data=session, + input_type="request", + ) + assert mock_api.call_count == 2 + assert result["texts"] == ["my ssn is [REDACTED]"] + + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={"texts": ["my ssn is 123-45-6789"]}, + request_data=session, + input_type="request", + ) + assert mock_api.call_count >= 1 + first_scanned = mock_api.call_args_list[0].kwargs.get("messages") + assert first_scanned is not None + assert [m["content"] for m in first_scanned] == ["my ssn is 123-45-6789"] + + @pytest.mark.asyncio + async def test_generic_agent_multi_turn_scans_only_new_each_turn(self): + """A generic agent (not Claude Code) opts in by propagating a session id. + + Agent frameworks on the OpenAI SDK carry the session through the request + body (metadata.session_id here), not the x-claude-code-session-id header. + Across a growing multi-turn conversation every turn after the first must + send Bedrock only the newly appended segments, never the whole context. + """ + guardrail = self._guardrail() + session = {"metadata": {"session_id": "agent-multi-turn"}} + + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + + await guardrail.apply_guardrail( + inputs={"texts": ["system prompt", "turn 1 question"]}, + request_data=session, + input_type="request", + ) + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ + "system prompt", + "turn 1 question", + ] + + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={"texts": ["system prompt", "turn 1 question", "turn 1 answer", "turn 2 question"]}, + request_data=session, + input_type="request", + ) + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ + "turn 1 answer", + "turn 2 question", + ] + + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={ + "texts": [ + "system prompt", + "turn 1 question", + "turn 1 answer", + "turn 2 question", + "turn 2 answer", + "turn 3 question", + ] + }, + request_data=session, + input_type="request", + ) + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ + "turn 2 answer", + "turn 3 question", + ] + + def test_incremental_scan_cache_prefers_proxy_shared_cache(self): + guardrail = self._guardrail() + shared = DualCache() + proxy_logging = MagicMock() + proxy_logging.internal_usage_cache.dual_cache = shared + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging): + assert guardrail._incremental_scan_cache() is shared + + def test_incremental_scan_cache_falls_back_when_proxy_logging_missing(self): + from litellm.integrations.custom_guardrail import dc as fallback_cache + + guardrail = self._guardrail() + with patch("litellm.proxy.proxy_server.proxy_logging_obj", None): + assert guardrail._incremental_scan_cache() is fallback_cache + + def test_incremental_scan_cache_falls_back_when_proxy_not_importable(self): + from litellm.integrations.custom_guardrail import dc as fallback_cache + + guardrail = self._guardrail() + with patch.dict(sys.modules, {"litellm.proxy.proxy_server": None}): + assert guardrail._incremental_scan_cache() is fallback_cache + + @pytest.mark.asyncio + async def test_blocked_turn_is_rescanned_on_retry(self): + guardrail = self._guardrail() + session = {"litellm_session_id": "sess-bedrock-blocked"} + + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.side_effect = HTTPException(status_code=400, detail="blocked") + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs={"texts": ["blocked prompt"]}, request_data=session, input_type="request" + ) + + mock_api.reset_mock() + mock_api.side_effect = None + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await guardrail.apply_guardrail( + inputs={"texts": ["blocked prompt"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 1 + scanned = mock_api.call_args.kwargs["messages"] + assert [m["content"] for m in scanned] == ["blocked prompt"] + + +class TestBedrockIncrementalFlagInteractions: + """Regression coverage for only_scan_new_messages combined with the other + Bedrock guardrail flags, from the PR #33278 live validation. Live evidence: + each of these was reproduced against a real Bedrock ApplyGuardrail first; + the mocks here encode the wire payloads observed there. + """ + + def _guardrail(self, **overrides): + params = dict( + guardrail_name="bedrock-incremental-flags", + guardrailIdentifier="test-guardrail", + guardrailVersion="DRAFT", + default_on=True, + only_scan_new_messages=True, + ) + params.update(overrides) + return BedrockGuardrail(**params) + + @pytest.mark.asyncio + async def test_edited_history_segment_rescans_only_that_segment(self): + guardrail = self._guardrail() + session = {"litellm_session_id": "sess-flags-edit"} + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await guardrail.apply_guardrail( + inputs={"texts": ["q1", "a1", "q2"]}, request_data=session, input_type="request" + ) + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={"texts": ["q1 EDITED", "a1", "q2"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 1 + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == ["q1 EDITED"] + + @pytest.mark.asyncio + async def test_same_content_different_session_rescans_everything(self): + guardrail = self._guardrail() + texts = ["shared question", "shared answer"] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await guardrail.apply_guardrail( + inputs={"texts": list(texts)}, request_data={"litellm_session_id": "sess-x1"}, input_type="request" + ) + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={"texts": list(texts)}, request_data={"litellm_session_id": "sess-x2"}, input_type="request" + ) + assert mock_api.call_count == 1 + assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == texts + + @pytest.mark.asyncio + async def test_litellm_masking_flag_disables_incremental_single_full_scan(self): + """mask_request_content must fall back to exactly ONE full scan per turn + and never persist hashes (verified live: 1 call/turn, no cache writes).""" + guardrail = self._guardrail(mask_request_content=True) + session = {"litellm_session_id": "sess-flags-mask"} + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await guardrail.apply_guardrail( + inputs={"texts": ["q1"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 1 + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={"texts": ["q1"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 1, "masking mode must re-scan every turn, exactly once" + + @pytest.mark.asyncio + async def test_server_side_anonymize_falls_back_full_scan_and_never_persists(self): + """A guardrail that rewrites content (Bedrock-side ANONYMIZE) must fall back + to the full scan so masking applies, and record no session state. Live + validation showed this costs 2 provider calls per turn; the count is + asserted here as documentation of that intended-tradeoff behavior.""" + guardrail = self._guardrail() + session = {"litellm_session_id": "sess-flags-anon"} + masked = {"action": "NONE", "output": [{"text": "MASKED q1"}], "outputs": [{"text": "MASKED q1"}]} + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = masked + result = await guardrail.apply_guardrail( + inputs={"texts": ["q1"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 2, "incremental attempt + full-scan fallback" + assert result["texts"] == ["MASKED q1"], "masked content must be applied" + mock_api.reset_mock() + await guardrail.apply_guardrail( + inputs={"texts": ["q1"]}, request_data=session, input_type="request" + ) + assert mock_api.call_count == 2, "no hashes persisted, so the double scan repeats" + + @pytest.mark.asyncio + @pytest.mark.xfail( + reason="PR #33278 known gap: incremental path bypasses _select_messages_for_apply_guardrail, " + "so experimental_use_latest_role_message_only is silently ignored. Intended semantics " + "(pending DRI decision): incremental mode defers to the latest-role selection.", + strict=False, + ) + async def test_latest_role_only_is_respected_with_incremental(self): + guardrail = self._guardrail(experimental_use_latest_role_message_only=True) + session = {"litellm_session_id": "sess-flags-latestrole"} + structured = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "q1"}, + ] + with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: + mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} + await guardrail.apply_guardrail( + inputs={"texts": ["sys", "q1"], "structured_messages": structured}, + request_data=session, + input_type="request", + ) + scanned = [m["content"] for m in mock_api.call_args.kwargs["messages"]] + assert scanned == ["q1"], "latest-role selection must exclude the system prompt" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py index 07c40aa763d..4021f922877 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_model_armor.py @@ -10,14 +10,19 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) +import httpx from fastapi import HTTPException import litellm import litellm.types.utils from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache +from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.model_armor import ModelArmorGuardrail +from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( + ModelArmorAPIError, +) from litellm.types.guardrails import GuardrailEventHooks @@ -403,8 +408,9 @@ async def test_model_armor_api_error_handling(): "metadata": {"guardrails": ["model-armor-test"]}, } - # Should raise HTTPException for API error - with pytest.raises(HTTPException) as exc_info: + # An API failure propagates as ModelArmorAPIError, not a content-block + # HTTPException, so guardrail trace status stays guardrail_failed_to_respond + with pytest.raises(ModelArmorAPIError) as exc_info: await guardrail.async_pre_call_hook( user_api_key_dict=mock_user_api_key_dict, cache=mock_cache, @@ -412,9 +418,8 @@ async def test_model_armor_api_error_handling(): call_type="completion", ) - assert exc_info.value.status_code == 400 - assert "Model Armor API error" in str(exc_info.value.detail) - assert "upstream 500" in str(exc_info.value.detail) + assert exc_info.value.detail == "Model Armor API error (upstream 500)" + assert "Internal Server Error" not in str(exc_info.value.detail) @pytest.mark.asyncio @@ -622,7 +627,7 @@ async def test_model_armor_streaming_block_yields_sse_error(): @pytest.mark.asyncio -async def test_model_armor_api_failure_returns_400(): +async def test_model_armor_api_failure_raises_sanitized_error(): """Test that Model Armor API failures raise HTTP 400, not the upstream status code.""" guardrail = ModelArmorGuardrail( template_id="test-template", @@ -643,15 +648,544 @@ async def test_model_armor_api_failure_returns_400(): with patch.object( guardrail.async_handler, "post", AsyncMock(return_value=mock_response) ): - with pytest.raises(HTTPException) as exc_info: + with pytest.raises(ModelArmorAPIError) as exc_info: await guardrail.make_model_armor_request( content="test content", source="user_prompt", ) - # Should be 400, NOT the upstream 500 - assert exc_info.value.status_code == 400 - assert "upstream 500" in str(exc_info.value.detail) + assert exc_info.value.detail == "Model Armor API error (upstream 500)" + assert "Internal Server Error" not in str(exc_info.value.detail) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sanitize", [True, False]) +async def test_model_armor_error_output_sanitization(sanitize: bool): + marker = "SYNTHETIC_MODEL_ARMOR_MARKER" + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + sanitize_error_detail=sanitize, + ) + guardrail._ensure_access_token_async = AsyncMock( + return_value=("test-token", "test-project") + ) + + error_response = AsyncMock(status_code=500, text=marker) + with patch.object( + guardrail.async_handler, "post", AsyncMock(return_value=error_response) + ), patch.object(verbose_proxy_logger, "debug") as debug_log, patch.object( + verbose_proxy_logger, "error" + ) as error_log, pytest.raises(ModelArmorAPIError) as exc_info: + await guardrail.make_model_armor_request(content=marker) + + direct_log = f"{debug_log.call_args_list} {error_log.call_args_list}" + if sanitize: + assert marker not in str(exc_info.value.detail) + assert marker not in direct_log + else: + assert marker in str(exc_info.value.detail) + assert marker in direct_log + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail_on_error", [True, False]) +async def test_model_armor_api_error_honors_fail_open(fail_on_error: bool): + """An upstream API failure (raised by the real handler as MaskedHTTPStatusError) + must block with a sanitized 400 when fail_on_error is true and let the request + proceed when the operator configured fail-open.""" + marker = "SYNTHETIC_FAIL_OPEN_MARKER" + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + fail_on_error=fail_on_error, + ) + guardrail._ensure_access_token_async = AsyncMock( + return_value=("test-token", "test-project") + ) + guardrail.should_run_guardrail = Mock(return_value=True) + + request = httpx.Request("POST", "https://modelarmor.example.test/v1") + upstream = httpx.Response(503, content=marker.encode(), request=request) + original = httpx.HTTPStatusError("Service Unavailable", request=request, response=upstream) + masked = MaskedHTTPStatusError(original, message=marker, text=marker) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "synthetic input"}], + "metadata": {}, + } + + with patch.object(guardrail.async_handler, "post", AsyncMock(side_effect=masked)): + if fail_on_error: + with pytest.raises(ModelArmorAPIError) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + assert exc_info.value.detail == "Model Armor API error (upstream 503)" + assert marker not in str(exc_info.value.detail) + else: + result = await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type="completion", + ) + assert result is request_data + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail_on_error", [True, False]) +async def test_model_armor_api_error_fail_open_moderation_and_post_call(fail_on_error: bool): + """The during-call and post-call hooks route API failures through fail_on_error + exactly like pre-call: sanitized 400 when failing closed, pass-through when open.""" + api_error = ModelArmorAPIError("Model Armor API error (upstream 503)") + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + fail_on_error=fail_on_error, + ) + guardrail.make_model_armor_request = AsyncMock(side_effect=api_error) + guardrail.should_run_guardrail = Mock(return_value=True) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "synthetic input"}], + "metadata": {}, + } + mock_llm_response = litellm.ModelResponse() + mock_llm_response.choices = [ + litellm.Choices(message=litellm.Message(content="model output")) + ] + + if fail_on_error: + with pytest.raises(ModelArmorAPIError) as mod_exc: + await guardrail.async_moderation_hook( + data=dict(request_data), + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + assert mod_exc.value.detail == "Model Armor API error (upstream 503)" + + with pytest.raises(ModelArmorAPIError) as post_exc: + await guardrail.async_post_call_success_hook( + data=dict(request_data), + user_api_key_dict=UserAPIKeyAuth(), + response=mock_llm_response, + ) + assert post_exc.value.detail == "Model Armor API error (upstream 503)" + else: + moderated = await guardrail.async_moderation_hook( + data=dict(request_data), + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + assert moderated is not None + + result = await guardrail.async_post_call_success_hook( + data=dict(request_data), + user_api_key_dict=UserAPIKeyAuth(), + response=mock_llm_response, + ) + assert result is mock_llm_response + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail_on_error", [True, False]) +async def test_model_armor_api_error_fail_open_streaming(fail_on_error: bool): + """A streaming-path API failure yields a sanitized SSE error frame when failing + closed and passes the original chunks through when the operator opted into fail-open.""" + api_error = ModelArmorAPIError("Model Armor API error (upstream 503)") + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + fail_on_error=fail_on_error, + ) + guardrail.make_model_armor_request = AsyncMock(side_effect=api_error) + guardrail.should_run_guardrail = Mock(return_value=True) + + async def mock_stream(): + yield litellm.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + delta=litellm.types.utils.Delta(content="streamed output") + ) + ] + ) + + chunks = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=mock_stream(), + request_data={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "synthetic input"}], + "metadata": {}, + }, + ): + chunks.append(chunk) + + if fail_on_error: + assert len(chunks) == 1 + assert isinstance(chunks[0], str) + assert "Model Armor API error (upstream 503)" in chunks[0] + assert '"code": "500"' in chunks[0] + else: + assert len(chunks) == 1 + assert isinstance(chunks[0], litellm.ModelResponseStream) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fail_on_error", [True, False]) +async def test_model_armor_api_error_fail_open_file_scan(fail_on_error: bool): + """A file-scan API failure blocks with the sanitized detail when failing closed + and skips the attachment when the operator opted into fail-open.""" + api_error = ModelArmorAPIError("Model Armor API error (upstream 503)") + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + fail_on_error=fail_on_error, + ) + guardrail.make_model_armor_request = AsyncMock(side_effect=api_error) + + pdf_b64 = base64.b64encode(b"%PDF-1.4 synthetic").decode() + messages = [ + { + "role": "user", + "content": [ + { + "type": "file", + "file": { + "file_data": f"data:application/pdf;base64,{pdf_b64}", + "filename": "synthetic.pdf", + "format": "application/pdf", + }, + } + ], + } + ] + data = {"metadata": {}} + + if fail_on_error: + with pytest.raises(ModelArmorAPIError) as exc_info: + await guardrail._scan_request_files(messages=messages, data=data) + assert exc_info.value.detail == "Model Armor API error (upstream 503)" + else: + assert await guardrail._scan_request_files(messages=messages, data=data) is None + + +def test_model_armor_hot_reload_null_stays_sanitized(): + """update_in_memory_litellm_params assigns raw fields; an explicit null in a + hot-reloaded config must not disable sanitization.""" + from litellm.types.guardrails import LitellmParams + + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + ) + guardrail.update_in_memory_litellm_params( + LitellmParams(guardrail="model_armor", mode="pre_call", sanitize_error_detail=None) + ) + assert guardrail.sanitize_error_detail is True + + guardrail.update_in_memory_litellm_params( + LitellmParams(guardrail="model_armor", mode="pre_call", sanitize_error_detail=False) + ) + assert guardrail.sanitize_error_detail is False + + +def test_model_armor_redactor_depth_cap_fails_closed(): + """Past the recursion cap the redactor must return the redaction sentinel, + never raw content, and must not raise RecursionError.""" + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.guardrails.guardrail_hooks.model_armor.model_armor import ( + _redact_scanned_content, + ) + + marker = "SYNTHETIC_DEEP_MARKER" + payload: dict = {"safe_key": marker, "items": [{"safe_key": marker}]} + for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 5): + payload = {"nested": payload} + + redacted = _redact_scanned_content(payload) + assert marker not in str(redacted) + + shallow = _redact_scanned_content({"filterResults": [{"text": marker, "matchState": "MATCH_FOUND"}]}) + assert shallow == {"filterResults": [{"text": "[REDACTED]", "matchState": "MATCH_FOUND"}]} + + uri_payload = _redact_scanned_content( + { + "maliciousUriFilterResult": { + "matchState": "MATCH_FOUND", + "maliciousUriMatchedItems": [{"uri": f"https://evil.example/{marker}"}], + } + } + ) + assert uri_payload == { + "maliciousUriFilterResult": { + "matchState": "MATCH_FOUND", + "maliciousUriMatchedItems": "[REDACTED]", + } + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sanitize", [True, False]) +async def test_model_armor_handler_raised_http_error_sanitized(sanitize: bool): + """The real AsyncHTTPHandler raises on non-2xx via raise_for_status, so a non-200 + never returns a response object. The raised MaskedHTTPStatusError carries the raw + upstream body in its message; the guardrail must convert it to a sanitized + HTTPException instead of letting it bubble raw to callers and logs.""" + marker = "SYNTHETIC_MODEL_ARMOR_MARKER" + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + sanitize_error_detail=sanitize, + ) + guardrail._ensure_access_token_async = AsyncMock( + return_value=("test-token", "test-project") + ) + + request = httpx.Request("POST", "https://modelarmor.example.test/v1") + upstream = httpx.Response(403, content=marker.encode(), request=request) + original = httpx.HTTPStatusError("Forbidden", request=request, response=upstream) + masked = MaskedHTTPStatusError(original, message=marker, text=marker) + + with patch.object( + guardrail.async_handler, "post", AsyncMock(side_effect=masked) + ), patch.object(verbose_proxy_logger, "debug") as debug_log, patch.object( + verbose_proxy_logger, "error" + ) as error_log, pytest.raises(ModelArmorAPIError) as exc_info: + await guardrail.make_model_armor_request(content=marker) + + direct_log = f"{debug_log.call_args_list} {error_log.call_args_list}" + assert "403" in str(exc_info.value.detail) + if sanitize: + assert marker not in str(exc_info.value.detail) + assert marker not in direct_log + else: + assert marker in str(exc_info.value.detail) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sanitize", [True, False]) +async def test_model_armor_post_call_logging_redacts_scanned_content(sanitize: bool): + marker = "SYNTHETIC_POST_CALL_MARKER" + armor_response = { + "sanitizationResult": { + "filterMatchState": "NO_MATCH_FOUND", + "filterResults": { + "sdp": { + "sdpFilterResult": { + "deidentifyResult": { + "matchState": "MATCH_FOUND", + "data": {"text": marker}, + } + } + } + }, + } + } + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + mask_response_content=True, + sanitize_error_detail=sanitize, + ) + guardrail.make_model_armor_request = AsyncMock(return_value=armor_response) + guardrail.should_run_guardrail = Mock(return_value=True) + + mock_llm_response = litellm.ModelResponse() + mock_llm_response.choices = [ + litellm.Choices(message=litellm.Message(content="model output")) + ] + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "synthetic input"}], + "metadata": {}, + "litellm_logging_obj": MagicMock(), + } + + with patch( + "litellm.proxy.common_utils.callback_utils.add_guardrail_response_to_standard_logging_object" + ) as add_logging: + await guardrail.async_post_call_success_hook( + data=request_data, + user_api_key_dict=UserAPIKeyAuth(), + response=mock_llm_response, + ) + + logged = add_logging.call_args.kwargs["guardrail_response"] + assert logged["guardrail_status"] == "success" + logged_armor_response = logged["guardrail_response"]["model_armor_response"] + if sanitize: + assert marker not in str(logged_armor_response) + assert ( + logged_armor_response["sanitizationResult"]["filterResults"]["sdp"][ + "sdpFilterResult" + ]["deidentifyResult"]["matchState"] + == "MATCH_FOUND" + ) + else: + assert logged_armor_response == armor_response + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sanitize", [True, False]) +async def test_model_armor_streaming_logging_redacts_scanned_content(sanitize: bool): + marker = "SYNTHETIC_STREAMING_MARKER" + armor_response = { + "sanitizationResult": { + "filterMatchState": "NO_MATCH_FOUND", + "sanitizedText": marker, + } + } + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + sanitize_error_detail=sanitize, + ) + guardrail.make_model_armor_request = AsyncMock(return_value=armor_response) + guardrail.should_run_guardrail = Mock(return_value=True) + + async def mock_stream(): + yield litellm.ModelResponseStream( + choices=[ + litellm.types.utils.StreamingChoices( + delta=litellm.types.utils.Delta(content="streamed output") + ) + ] + ) + + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "synthetic input"}], + "metadata": {}, + } + + async for _ in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=mock_stream(), + request_data=request_data, + ): + pass + + logged_response = request_data["metadata"]["_model_armor_response"] + if sanitize: + assert logged_response == { + "sanitizationResult": { + "filterMatchState": "NO_MATCH_FOUND", + "sanitizedText": "[REDACTED]", + } + } + assert marker not in str(logged_response) + else: + assert logged_response == armor_response + + +@pytest.mark.asyncio +@pytest.mark.parametrize("sanitize", [True, False]) +async def test_model_armor_match_found_sanitizes_caller_and_logging(sanitize: bool): + marker = "SYNTHETIC_MATCH_FOUND_MARKER" + armor_response = { + "sanitizationResult": { + "filterResults": { + "sdp": { + "sdpFilterResult": { + "inspectResult": { + "matchState": "MATCH_FOUND", + "findings": [{"marker": marker}], + } + } + } + } + } + } + guardrail = ModelArmorGuardrail( + template_id="test-template", + project_id="test-project", + guardrail_name="model-armor-test", + event_hook=[GuardrailEventHooks.pre_mcp_call], + sanitize_error_detail=sanitize, + ) + guardrail.make_model_armor_request = AsyncMock(return_value=armor_response) + guardrail.should_run_guardrail = Mock(return_value=True) + request_data = { + "messages": [{"role": "user", "content": "synthetic input"}], + "metadata": {}, + } + + with pytest.raises(HTTPException) as exc_info: + await guardrail.async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), + cache=MagicMock(spec=DualCache), + data=request_data, + call_type=litellm.types.utils.CallTypes.call_mcp_tool.value, + ) + + detail = exc_info.value.detail + logged_response = request_data["metadata"]["_model_armor_response"] + if sanitize: + assert detail == {"error": "Content blocked by Model Armor"} + assert logged_response == { + "sanitizationResult": { + "filterResults": { + "sdp": { + "sdpFilterResult": { + "inspectResult": { + "matchState": "MATCH_FOUND", + "findings": "[REDACTED]", + } + } + } + } + } + } + assert marker not in str(detail) + assert marker not in str(logged_response) + else: + assert detail["model_armor_response"] == armor_response + assert logged_response == armor_response + assert marker in str(detail) + assert marker in str(logged_response) + + +def test_model_armor_sanitize_error_detail_config_wiring(): + from litellm.proxy.guardrails.guardrail_hooks.model_armor import ( + initialize_guardrail, + ) + from litellm.types.guardrails import LitellmParams + + config = {"guardrail_name": "model-armor-test"} + params = { + "guardrail": "model_armor", + "mode": "pre_mcp_call", + "template_id": "test-template", + "project_id": "test-project", + } + opted_out = initialize_guardrail( + LitellmParams(**params, sanitize_error_detail=False), config + ) + explicit_null = initialize_guardrail( + LitellmParams(**params, sanitize_error_detail=None), config + ) + default = initialize_guardrail(LitellmParams(**params), config) + + assert opted_out.sanitize_error_detail is False + assert explicit_null.sanitize_error_detail is True + assert default.sanitize_error_detail is True def test_model_armor_ui_friendly_name(): @@ -1394,7 +1928,10 @@ async def test_model_armor_guardrail_status_intervened_vs_failed(): ) info = request_data["metadata"]["standard_logging_guardrail_information"] + assert info[0]["guardrail_name"] == guardrail.guardrail_name assert info[0]["guardrail_status"] == "guardrail_intervened" + assert "model_armor_response" not in info[0]["guardrail_response"] + assert "sanitizationResult" not in info[0]["guardrail_response"] # 2: if an API error - guardrail status should be guardrail_failed_to_respond" guardrail2 = ModelArmorGuardrail( diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py index f8995a6f4da..c3af4208d37 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_patch_user.py @@ -465,3 +465,58 @@ def test_apply_patch_ops_filtered_path_raises_400_instead_of_junk_metadata(): ) assert exc_info.value.status_code == 400 + + +def test_apply_patch_ops_remove_group_filtered_path_without_value(): + """Okta removes a user from a team with groups[value eq "..."] and no body + value; the team id must be parsed from the filter so the remove takes effect""" + user = LiteLLM_UserTable( + user_id="user-fp", + user_email="fp@example.com", + teams=["team-1", "team-2"], + metadata={}, + ) + patch_ops = SCIMPatchOp( + Operations=[SCIMPatchOperation(op="remove", path='groups[value eq "team-1"]')] + ) + + _, final_team_set = _apply_patch_ops(existing_user=user, patch_ops=patch_ops) + + assert final_team_set == {"team-2"} + + +def test_apply_patch_ops_add_group_filtered_path_without_value(): + """A filtered add path with no body value adds the team id from the filter.""" + user = LiteLLM_UserTable( + user_id="user-fp", + user_email="fp@example.com", + teams=["team-1"], + metadata={}, + ) + patch_ops = SCIMPatchOp( + Operations=[SCIMPatchOperation(op="add", path="groups[value eq 'team-3']")] + ) + + _, final_team_set = _apply_patch_ops(existing_user=user, patch_ops=patch_ops) + + assert final_team_set == {"team-1", "team-3"} + + +def test_apply_patch_ops_replace_groups_empty_value_does_not_use_path_filter(): + """A filtered replace with an explicit empty value must not resurrect the + filter id; the team set is replaced with the empty value as given.""" + user = LiteLLM_UserTable( + user_id="user-fp", + user_email="fp@example.com", + teams=["team-1", "team-2"], + metadata={}, + ) + patch_ops = SCIMPatchOp( + Operations=[ + SCIMPatchOperation(op="replace", path='groups[value eq "team-1"]', value=[]) + ] + ) + + _, final_team_set = _apply_patch_ops(existing_user=user, patch_ops=patch_ops) + + assert final_team_set == set() diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index f27f1197090..7bb74285ac6 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -1,24 +1,32 @@ +import time from unittest.mock import AsyncMock import pytest from fastapi import HTTPException from litellm.proxy._types import ( + LiteLLM_TeamTable, LiteLLM_UserTable, LitellmUserRoles, + Member, NewUserRequest, NewUserResponse, + ProxyErrorTypes, ProxyException, ) from litellm.proxy.management_endpoints.scim.scim_v2 import ( UserProvisionerHelpers, + _apply_group_patch_updates, _extract_group_member_ids, + _extract_ids_from_path_filter, _handle_team_membership_changes, _process_group_patch_operations, _recompute_scim_member_roles, create_group, create_user, delete_group, + delete_user, + get_groups, get_users, get_service_provider_config, patch_group, @@ -55,9 +63,7 @@ async def test_create_user_existing_user_conflict(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value={"user_id": "existing-user"} - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value={"user_id": "existing-user"}) # Mock the _get_prisma_client_or_raise_exception to return our mock mocker.patch( @@ -231,9 +237,7 @@ async def test_create_user_ingests_entitlements_and_roles(mocker, monkeypatch): }, {"value": "bare-entitlement"}, ] - assert created_metadata["scim_roles"] == [ - {"value": "engineering-admin", "type": "role"} - ] + assert created_metadata["scim_roles"] == [{"value": "engineering-admin", "type": "role"}] @pytest.mark.asyncio @@ -257,9 +261,7 @@ async def test_create_user_uses_default_internal_user_params_role(mocker, monkey default_params = { "user_role": LitellmUserRoles.PROXY_ADMIN, } - monkeypatch.setattr( - "litellm.default_internal_user_params", default_params, raising=False - ) + monkeypatch.setattr("litellm.default_internal_user_params", default_params, raising=False) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -339,16 +341,13 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp "BUG: _update_litellm_setting did not update litellm.default_internal_user_params in memory. " "The local variable reassignment (in_memory_var = ...) doesn't propagate back." ) - assert ( - litellm.default_internal_user_params.get("user_role") - == LitellmUserRoles.INTERNAL_USER - ) + assert litellm.default_internal_user_params.get("user_role") == LitellmUserRoles.INTERNAL_USER # Step 3: Create a user via SCIM scim_user = SCIMUser( schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], - userName="idontexist@krakentest.tech", - emails=[SCIMUserEmail(value="idontexist@krakentest.tech")], + userName="idontexist@example.com", + emails=[SCIMUserEmail(value="idontexist@example.com")], ) mock_prisma_client = mocker.MagicMock() @@ -364,7 +363,7 @@ async def test_scim_create_user_respects_default_role_set_via_ui(mocker, monkeyp new_user_mock = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.new_user", - AsyncMock(return_value=NewUserRequest(user_id="idontexist@krakentest.tech")), + AsyncMock(return_value=NewUserRequest(user_id="idontexist@example.com")), ) mocker.patch( @@ -438,9 +437,7 @@ async def test_get_users_filters_username_by_exposed_scim_username_for_okta(mock take=10, order={"created_at": "desc"}, ) - mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with( - where=expected_where - ) + mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with(where=expected_where) assert response.totalResults == 1 assert response.Resources[0].id == "internal-user-id" @@ -494,9 +491,7 @@ async def test_get_users_filters_email_value_by_user_email(mocker): take=10, order={"created_at": "desc"}, ) - mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with( - where=expected_where - ) + mock_prisma_client.db.litellm_usertable.count.assert_awaited_once_with(where=expected_where) assert response.totalResults == 1 assert response.Resources[0].id == "internal-user-id" @@ -544,15 +539,12 @@ async def test_handle_existing_user_by_email_no_existing_user(mocker): ) assert result is None - mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with( - where={"user_email": "test@example.com"} - ) + mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with(where={"user_email": "test@example.com"}) @pytest.mark.asyncio async def test_handle_existing_user_by_email_existing_user_updated(mocker): - """Should update existing user and return SCIMUser when user with email exists""" - # Mock existing user - create a proper mock object with attributes + """Should rename the existing user, sync team roster, and return SCIMUser""" existing_user = mocker.MagicMock() existing_user.user_id = "old-user-id" existing_user.user_email = "test@example.com" @@ -560,7 +552,6 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): existing_user.teams = ["old-team"] existing_user.metadata = {"old": "data"} - # Mock updated user updated_user = { "user_id": "new-user-id", "user_email": "test@example.com", @@ -569,7 +560,6 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): "metadata": '{"new": "data"}', } - # Mock SCIM user to be returned mock_scim_user = SCIMUser( schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], id="new-user-id", @@ -581,18 +571,17 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=existing_user - ) - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) - # Mock the transformation function mock_transform = mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", AsyncMock(return_value=mock_scim_user), ) + mock_membership = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock(), + ) new_user_request = NewUserRequest( user_id="new-user-id", @@ -607,29 +596,276 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): prisma_client=mock_prisma_client, new_user_request=new_user_request ) - # Verify the result assert result == mock_scim_user - # Verify database operations - mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with( - where={"user_email": "test@example.com"} - ) + mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with(where={"user_email": "test@example.com"}) - mock_prisma_client.db.litellm_usertable.update.assert_called_once_with( - where={"user_id": "old-user-id"}, - data={ - "user_id": "new-user-id", + update_calls = mock_prisma_client.db.litellm_usertable.update.call_args_list + assert len(update_calls) == 2 + assert update_calls[0].kwargs == { + "where": {"user_id": "old-user-id"}, + "data": {"user_id": "new-user-id"}, + } + assert update_calls[1].kwargs == { + "where": {"user_id": "new-user-id"}, + "data": { "user_email": "test@example.com", "user_alias": "New Name", "teams": ["new-team"], "metadata": '{"new": "data"}', }, + } + + mock_membership.assert_awaited_once_with( + user_id="new-user-id", + existing_teams=["old-team"], + new_teams=["new-team"], + raise_on_error=True, ) - # Verify transformation was called mock_transform.assert_called_once_with(updated_user) +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_syncs_roster_and_dedups_teams(mocker): + """Existing-email upsert must add the user to the team roster via the shared + team_member_add path and dedup the teams built from repeated SCIM groups. + + Regression: previously the user's ``teams`` array was raw-written (with + duplicates) and the team roster (members_with_roles / LiteLLM_TeamMembership) + was never touched, so the user appeared in the group on their profile but was + absent from the team directly. + """ + existing_user = mocker.MagicMock() + existing_user.user_id = "same-id" + existing_user.user_email = "member@example.com" + existing_user.user_alias = "Member" + existing_user.teams = [] + existing_user.metadata = {} + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={}) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=None), + ) + mock_membership = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock(), + ) + + new_user_request = NewUserRequest( + user_id="same-id", + user_email="member@example.com", + user_alias="Member", + teams=["team-a", "team-a", "team-b"], + metadata={}, + auto_create_key=False, + ) + + await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, new_user_request=new_user_request + ) + + mock_membership.assert_awaited_once_with( + user_id="same-id", + existing_teams=[], + new_teams=["team-a", "team-b"], + raise_on_error=True, + ) + + update_calls = mock_prisma_client.db.litellm_usertable.update.call_args_list + assert len(update_calls) == 1 + assert update_calls[0].kwargs["where"] == {"user_id": "same-id"} + assert update_calls[0].kwargs["data"]["teams"] == ["team-a", "team-b"] + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_roster_add_failure_blocks_teams_write(mocker): + """A genuine roster add failure must propagate and must not persist the teams array. + + Regression: the roster sync went through patch_team_membership which swallowed + real team_member_add failures, so the endpoint reported success and wrote a + teams array listing a team the roster never received. The strict path now + surfaces the failure so user.teams and members_with_roles cannot diverge. + """ + existing_user = mocker.MagicMock() + existing_user.user_id = "uid" + existing_user.user_email = "member@example.com" + existing_user.user_alias = "Member" + existing_user.teams = [] + existing_user.metadata = {} + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={}) + + mock_team_member_add = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "Team not found"})), + ) + + new_user_request = NewUserRequest( + user_id="uid", + user_email="member@example.com", + user_alias="Member", + teams=["missing-team"], + metadata={}, + auto_create_key=False, + ) + + with pytest.raises(HTTPException): + await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, new_user_request=new_user_request + ) + + mock_team_member_add.assert_awaited_once() + assert mock_prisma_client.db.litellm_usertable.update.await_count == 0 + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_roster_add_already_member_is_noop(mocker): + """Being already in the team is benign even under the strict path: the upsert + succeeds and the deduped teams array is still persisted.""" + existing_user = mocker.MagicMock() + existing_user.user_id = "uid" + existing_user.user_email = "member@example.com" + existing_user.user_alias = "Member" + existing_user.teams = [] + existing_user.metadata = {} + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={}) + + mock_team_member_add = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", + AsyncMock( + side_effect=ProxyException( + message="already in team", + type=ProxyErrorTypes.team_member_already_in_team.value, + param=None, + code=400, + ) + ), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=None), + ) + + new_user_request = NewUserRequest( + user_id="uid", + user_email="member@example.com", + user_alias="Member", + teams=["team-x"], + metadata={}, + auto_create_key=False, + ) + + await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, new_user_request=new_user_request + ) + + mock_team_member_add.assert_awaited_once() + update_calls = mock_prisma_client.db.litellm_usertable.update.call_args_list + assert len(update_calls) == 1 + assert update_calls[0].kwargs["data"]["teams"] == ["team-x"] + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_roster_remove_failure_blocks_teams_write(mocker): + """A genuine roster removal failure must propagate and must not persist the teams array, + symmetrically with add failures, so user.teams cannot drop a team the roster still holds.""" + existing_user = mocker.MagicMock() + existing_user.user_id = "uid" + existing_user.user_email = "member@example.com" + existing_user.user_alias = "Member" + existing_user.teams = ["old-team"] + existing_user.metadata = {} + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={}) + + mock_team_member_delete = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(side_effect=HTTPException(status_code=500, detail={"error": "No db connected"})), + ) + + new_user_request = NewUserRequest( + user_id="uid", + user_email="member@example.com", + user_alias="Member", + teams=[], + metadata={}, + auto_create_key=False, + ) + + with pytest.raises(HTTPException): + await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, new_user_request=new_user_request + ) + + mock_team_member_delete.assert_awaited_once() + assert mock_prisma_client.db.litellm_usertable.update.await_count == 0 + + +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_roster_remove_already_absent_is_noop(mocker): + """A user already absent from the team is the idempotent removal no-op even under the + strict path: the upsert succeeds and the deduped teams array is still persisted.""" + existing_user = mocker.MagicMock() + existing_user.user_id = "uid" + existing_user.user_email = "member@example.com" + existing_user.user_alias = "Member" + existing_user.teams = ["old-team"] + existing_user.metadata = {} + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={}) + + mock_team_member_delete = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(side_effect=HTTPException(status_code=400, detail={"error": "User not found in team"})), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=None), + ) + + new_user_request = NewUserRequest( + user_id="uid", + user_email="member@example.com", + user_alias="Member", + teams=[], + metadata={}, + auto_create_key=False, + ) + + await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, new_user_request=new_user_request + ) + + mock_team_member_delete.assert_awaited_once() + update_calls = mock_prisma_client.db.litellm_usertable.update.call_args_list + assert len(update_calls) == 1 + assert update_calls[0].kwargs["data"]["teams"] == [] + + @pytest.mark.asyncio async def test_handle_team_membership_changes_no_changes(mocker): """Should not call patch_team_membership when existing teams equal new teams""" @@ -762,9 +998,7 @@ async def test_update_user_success(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) # Mock dependencies mocker.patch( @@ -815,11 +1049,7 @@ async def test_update_user_not_found(mocker): ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", - AsyncMock( - side_effect=HTTPException( - status_code=404, detail={"error": "User not found"} - ) - ), + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "User not found"})), ) # Should raise ProxyException (which wraps the HTTPException) @@ -864,9 +1094,7 @@ async def test_patch_user_success(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) # Mock dependencies mocker.patch( @@ -903,9 +1131,7 @@ async def test_patch_user_not_found(mocker): """Should raise 404 when user doesn't exist for patch""" patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation(op="replace", path="displayName", value="New Name") - ], + Operations=[SCIMPatchOperation(op="replace", path="displayName", value="New Name")], ) # Mock dependencies to raise HTTPException for user not found @@ -915,11 +1141,7 @@ async def test_patch_user_not_found(mocker): ) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", - AsyncMock( - side_effect=HTTPException( - status_code=404, detail={"error": "User not found"} - ) - ), + AsyncMock(side_effect=HTTPException(status_code=404, detail={"error": "User not found"})), ) # Should raise ProxyException (which wraps the HTTPException) @@ -939,9 +1161,7 @@ async def test_get_service_provider_config(mocker): # Verify it returns the correct response assert isinstance(result, SCIMServiceProviderConfig) - assert result.schemas == [ - "urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig" - ] + assert result.schemas == ["urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig"] assert result.patch.supported is True assert result.bulk.supported is False assert result.meta is not None @@ -993,21 +1213,15 @@ async def test_update_group_metadata_serialization_issue(mocker): mock_prisma_client.db.litellm_usertable = mocker.MagicMock() # Mock team operations - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) # Mock user operations mock_user = mocker.MagicMock() mock_user.user_id = "user1" mock_user.user_email = "user1@example.com" # Add proper string value for user_email mock_user.teams = [group_id] - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mock_user - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=mock_user) # Mock the _get_prisma_client_or_raise_exception to return our mock @@ -1044,9 +1258,7 @@ async def test_update_group_metadata_serialization_issue(mocker): metadata = update_data["metadata"] # The fix should ensure metadata is serialized as a JSON string - assert isinstance( - metadata, str - ), f"metadata should be a JSON string, but got {type(metadata)}" + assert isinstance(metadata, str), f"metadata should be a JSON string, but got {type(metadata)}" # Verify we can parse it back to verify it contains the expected data import json @@ -1103,9 +1315,7 @@ async def test_team_membership_management(mocker): # Check calls for adding members add_calls = [ - call - for call in mock_patch_team_membership.call_args_list - if call[1]["teams_ids_to_add_user_to"] == [group_id] + call for call in mock_patch_team_membership.call_args_list if call[1]["teams_ids_to_add_user_to"] == [group_id] ] assert len(add_calls) == 2 # user3 and user4 @@ -1131,9 +1341,7 @@ async def test_team_membership_management(mocker): # Each call should either add OR remove, not both add_teams = call[1]["teams_ids_to_add_user_to"] remove_teams = call[1]["teams_ids_to_remove_user_from"] - assert (len(add_teams) > 0) != ( - len(remove_teams) > 0 - ) # XOR - one should be empty + assert (len(add_teams) > 0) != (len(remove_teams) > 0) # XOR - one should be empty @pytest.mark.asyncio @@ -1184,9 +1392,7 @@ async def test_update_group_e2e(mocker): mock_prisma_client.db.litellm_usertable = mocker.MagicMock() # Mock database operations - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) # Mock the updated team that gets returned from database updated_team = LiteLLM_TeamTable( @@ -1203,16 +1409,12 @@ async def test_update_group_e2e(mocker): "scim_data": scim_group_update.model_dump(), }, ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=updated_team - ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated_team) # Mock user validation (all users exist) mock_user = mocker.MagicMock() mock_user.user_id = "test-user" - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mock_user - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) # Mock dependencies mocker.patch( @@ -1265,29 +1467,19 @@ async def test_update_group_e2e(mocker): assert metadata["scim_data"]["displayName"] == "Updated Team Name" # Verify team membership changes were handled correctly - assert ( - mock_patch_team_membership.call_count == 3 - ) # Remove user1, add user3, add user4 + assert mock_patch_team_membership.call_count == 3 # Remove user1, add user3, add user4 # Check membership changes call_args_list = mock_patch_team_membership.call_args_list # Find remove operation (user1) - remove_calls = [ - call - for call in call_args_list - if call[1]["teams_ids_to_remove_user_from"] == [group_id] - ] + remove_calls = [call for call in call_args_list if call[1]["teams_ids_to_remove_user_from"] == [group_id]] assert len(remove_calls) == 1 assert remove_calls[0][1]["user_id"] == "user1" assert remove_calls[0][1]["teams_ids_to_add_user_to"] == [] # Find add operations (user3, user4) - add_calls = [ - call - for call in call_args_list - if call[1]["teams_ids_to_add_user_to"] == [group_id] - ] + add_calls = [call for call in call_args_list if call[1]["teams_ids_to_add_user_to"] == [group_id]] assert len(add_calls) == 2 add_user_ids = {call[1]["user_id"] for call in add_calls} assert add_user_ids == {"user3", "user4"} @@ -1302,9 +1494,7 @@ async def test_update_group_e2e(mocker): assert len(result.members) == 3 # Verify SCIM transformation was called with updated team - ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with( - updated_team - ) + ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team) @pytest.mark.asyncio @@ -1330,15 +1520,9 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch): id=group_id, displayName="Test Group", members=[ - SCIMMember( - value="existing-user", display="Existing User" - ), # This user exists - SCIMMember( - value="new-user-1", display="New User 1" - ), # This user doesn't exist - SCIMMember( - value="new-user-2", display="New User 2" - ), # This user doesn't exist + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist + SCIMMember(value="new-user-2", display="New User 2"), # This user doesn't exist ], ) @@ -1364,9 +1548,7 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch): return mock_user return None # new-user-1 and new-user-2 don't exist - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - side_effect=mock_user_lookup - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) # Mock dependencies mocker.patch( @@ -1381,9 +1563,7 @@ async def test_create_group_with_nonexistent_users_rejects(mocker, monkeypatch): # Verify it's a 400 Bad Request assert int(exc_info.value.code) == 400 assert "does not exist" in str(exc_info.value.message) - assert "new-user-1" in str(exc_info.value.message) or "new-user-2" in str( - exc_info.value.message - ) + assert "new-user-1" in str(exc_info.value.message) or "new-user-2" in str(exc_info.value.message) @pytest.mark.asyncio @@ -1418,15 +1598,9 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): id=group_id, displayName="Updated Group Name", members=[ - SCIMMember( - value="existing-user", display="Existing User" - ), # This user exists - SCIMMember( - value="new-user-3", display="New User 3" - ), # This user doesn't exist - SCIMMember( - value="new-user-4", display="New User 4" - ), # This user doesn't exist + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-3", display="New User 3"), # This user doesn't exist + SCIMMember(value="new-user-4", display="New User 4"), # This user doesn't exist ], ) @@ -1437,18 +1611,14 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): mock_prisma_client.db.litellm_usertable = mocker.MagicMock() # Mock team operations - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=mock_existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team) # Mock updated team response mock_updated_team = mocker.MagicMock() mock_updated_team.team_id = group_id mock_updated_team.team_alias = "Updated Group Name" mock_updated_team.members = ["existing-user", "new-user-3", "new-user-4"] - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=mock_updated_team - ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team) # Mock user lookup - only existing-user exists def mock_user_lookup(where): @@ -1459,9 +1629,7 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): return mock_user return None # new-user-3 and new-user-4 don't exist - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - side_effect=mock_user_lookup - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) # Mock dependencies mocker.patch( @@ -1481,15 +1649,11 @@ async def test_update_group_with_nonexistent_users_rejects(mocker, monkeypatch): # Verify it's a 400 Bad Request assert int(exc_info.value.code) == 400 assert "does not exist" in str(exc_info.value.message) - assert "new-user-3" in str(exc_info.value.message) or "new-user-4" in str( - exc_info.value.message - ) + assert "new-user-3" in str(exc_info.value.message) or "new-user-4" in str(exc_info.value.message) @pytest.mark.asyncio -async def test_create_group_with_nonexistent_users_creates_when_flag_true( - mocker, monkeypatch -): +async def test_create_group_with_nonexistent_users_creates_when_flag_true(mocker, monkeypatch): """ Test that creating a group with non-existent users creates them when scim_upsert_user is True. This preserves backward compatible behavior. @@ -1510,15 +1674,9 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true( id=group_id, displayName="Test Group", members=[ - SCIMMember( - value="existing-user", display="Existing User" - ), # This user exists - SCIMMember( - value="new-user-1", display="New User 1" - ), # This user doesn't exist - should be created - SCIMMember( - value="new-user-2", display="New User 2" - ), # This user doesn't exist - should be created + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist - should be created + SCIMMember(value="new-user-2", display="New User 2"), # This user doesn't exist - should be created ], ) @@ -1540,9 +1698,7 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true( return mock_user return None # new-user-1 and new-user-2 don't exist - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - side_effect=mock_user_lookup - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) # Mock user creation created_user_1 = NewUserResponse(user_id="new-user-1", key="test-key-1") @@ -1591,9 +1747,7 @@ async def test_create_group_with_nonexistent_users_creates_when_flag_true( @pytest.mark.asyncio -async def test_extract_group_member_ids_with_flag_true_creates_users( - mocker, monkeypatch -): +async def test_extract_group_member_ids_with_flag_true_creates_users(mocker, monkeypatch): """ Test that _extract_group_member_ids creates users when scim_upsert_user is True. """ @@ -1612,12 +1766,8 @@ async def test_extract_group_member_ids_with_flag_true_creates_users( id="test-group", displayName="Test Group", members=[ - SCIMMember( - value="existing-user", display="Existing User" - ), # This user exists - SCIMMember( - value="new-user-1", display="New User 1" - ), # This user doesn't exist - should be created + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist - should be created ], ) @@ -1635,9 +1785,7 @@ async def test_extract_group_member_ids_with_flag_true_creates_users( return mock_user return None # new-user-1 doesn't exist - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - side_effect=mock_user_lookup - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) # Mock user creation created_user = NewUserResponse(user_id="new-user-1", key="test-key-1") @@ -1662,9 +1810,7 @@ async def test_extract_group_member_ids_with_flag_true_creates_users( assert len(result.created_users) == 1 # Verify user was created - mock_create_user.assert_called_once_with( - user_id="new-user-1", created_via="scim_group_membership" - ) + mock_create_user.assert_called_once_with(user_id="new-user-1", created_via="scim_group_membership") @pytest.mark.asyncio @@ -1687,12 +1833,8 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa id="test-group", displayName="Test Group", members=[ - SCIMMember( - value="existing-user", display="Existing User" - ), # This user exists - SCIMMember( - value="new-user-1", display="New User 1" - ), # This user doesn't exist - should be rejected + SCIMMember(value="existing-user", display="Existing User"), # This user exists + SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist - should be rejected ], ) @@ -1710,9 +1852,7 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa return mock_user return None # new-user-1 doesn't exist - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - side_effect=mock_user_lookup - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup) # Mock dependencies mocker.patch( @@ -1731,9 +1871,7 @@ async def test_extract_group_member_ids_with_flag_false_rejects(mocker, monkeypa @pytest.mark.asyncio -async def test_process_group_patch_operations_with_flag_true_creates_users( - mocker, monkeypatch -): +async def test_process_group_patch_operations_with_flag_true_creates_users(mocker, monkeypatch): """ Test that _process_group_patch_operations creates users when scim_upsert_user is True. """ @@ -1749,11 +1887,7 @@ async def test_process_group_patch_operations_with_flag_true_creates_users( # Test data patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation( - op="add", path="members", value=[{"value": "new-user-1"}] - ) - ], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "new-user-1"}])], ) # Mock existing team @@ -1777,7 +1911,7 @@ async def test_process_group_patch_operations_with_flag_true_creates_users( ) # Execute the function - update_data, final_members = await _process_group_patch_operations( + update_data, final_members, _ = await _process_group_patch_operations( patch_ops=patch_ops, existing_team=mock_existing_team, prisma_client=mock_prisma_client, @@ -1787,15 +1921,11 @@ async def test_process_group_patch_operations_with_flag_true_creates_users( assert "new-user-1" in final_members # Verify user was created - mock_create_user.assert_called_once_with( - user_id="new-user-1", created_via="scim_group_patch" - ) + mock_create_user.assert_called_once_with(user_id="new-user-1", created_via="scim_group_patch") @pytest.mark.asyncio -async def test_process_group_patch_operations_with_flag_false_rejects( - mocker, monkeypatch -): +async def test_process_group_patch_operations_with_flag_false_rejects(mocker, monkeypatch): """ Test that _process_group_patch_operations rejects non-existent users when scim_upsert_user is False. """ @@ -1811,11 +1941,7 @@ async def test_process_group_patch_operations_with_flag_false_rejects( # Test data patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation( - op="add", path="members", value=[{"value": "new-user-1"}] - ) - ], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "new-user-1"}])], ) # Mock existing team @@ -1890,9 +2016,7 @@ async def test_create_user_grants_admin_when_in_scim_admin_group(mocker, monkeyp @pytest.mark.asyncio -async def test_create_user_keeps_default_when_not_in_scim_admin_group( - mocker, monkeypatch -): +async def test_create_user_keeps_default_when_not_in_scim_admin_group(mocker, monkeypatch): """When scim_admin_group is configured but the user's groups don't include it, the user keeps the non-admin default role.""" from litellm.proxy.proxy_server import proxy_config @@ -1936,9 +2060,7 @@ async def test_create_user_keeps_default_when_not_in_scim_admin_group( @pytest.mark.asyncio -async def test_update_user_demotes_admin_when_removed_from_scim_admin_group( - mocker, monkeypatch -): +async def test_update_user_demotes_admin_when_removed_from_scim_admin_group(mocker, monkeypatch): """Core demotion test: a PUT whose new groups no longer include the configured admin group must re-evaluate the role and write the non-admin default, so an admin removed from the IdP group is demoted without re-login.""" @@ -1972,9 +2094,7 @@ async def test_update_user_demotes_admin_when_removed_from_scim_admin_group( mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2000,9 +2120,7 @@ async def test_update_user_demotes_admin_when_removed_from_scim_admin_group( @pytest.mark.asyncio -async def test_update_user_does_not_force_role_when_scim_admin_group_unset( - mocker, monkeypatch -): +async def test_update_user_does_not_force_role_when_scim_admin_group_unset(mocker, monkeypatch): """When scim_admin_group is unset, PUT must not touch user_role (current behavior preserved).""" from litellm.proxy.proxy_server import proxy_config @@ -2035,9 +2153,7 @@ async def test_update_user_does_not_force_role_when_scim_admin_group_unset( mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2063,9 +2179,7 @@ async def test_update_user_does_not_force_role_when_scim_admin_group_unset( @pytest.mark.asyncio -async def test_update_user_demotes_when_default_params_lack_user_role( - mocker, monkeypatch -): +async def test_update_user_demotes_when_default_params_lack_user_role(mocker, monkeypatch): """Regression: default_internal_user_params set without a user_role key must still resolve to the non-admin default on demotion, not silently skip and leave the user PROXY_ADMIN.""" @@ -2075,9 +2189,7 @@ async def test_update_user_demotes_when_default_params_lack_user_role( return {"litellm_settings": {"scim_admin_group": "litellm-admins"}} monkeypatch.setattr(proxy_config, "get_config", mock_get_config) - monkeypatch.setattr( - "litellm.default_internal_user_params", {"max_budget": 10}, raising=False - ) + monkeypatch.setattr("litellm.default_internal_user_params", {"max_budget": 10}, raising=False) existing_user = mocker.MagicMock() existing_user.teams = ["litellm-admins"] @@ -2101,9 +2213,7 @@ async def test_update_user_demotes_when_default_params_lack_user_role( mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2129,9 +2239,7 @@ async def test_update_user_demotes_when_default_params_lack_user_role( @pytest.mark.asyncio -async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group( - mocker, monkeypatch -): +async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group(mocker, monkeypatch): """PATCH that drops the admin team from the resulting team set must write the non-admin default, mirroring the PUT demotion path.""" from litellm.proxy.proxy_server import proxy_config @@ -2148,11 +2256,7 @@ async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group( patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation( - op="replace", path="groups", value=[{"value": "engineering"}] - ) - ], + Operations=[SCIMPatchOperation(op="replace", path="groups", value=[{"value": "engineering"}])], ) updated_user = { @@ -2168,13 +2272,9 @@ async def test_patch_user_demotes_admin_when_removed_from_scim_admin_group( mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=engineering_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=engineering_team) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2223,11 +2323,7 @@ async def test_patch_user_grants_admin_by_team_display_name(mocker, monkeypatch) patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation( - op="replace", path="groups", value=[{"value": "team-abc-123"}] - ) - ], + Operations=[SCIMPatchOperation(op="replace", path="groups", value=[{"value": "team-abc-123"}])], ) updated_user = { @@ -2243,13 +2339,9 @@ async def test_patch_user_grants_admin_by_team_display_name(mocker, monkeypatch) mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value=updated_user - ) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value=updated_user) mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=admin_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=admin_team) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2302,9 +2394,7 @@ def _scim_admin_prisma(mocker, *, user_teams): @pytest.mark.asyncio -async def test_recompute_scim_member_roles_demotes_when_not_in_admin_group( - mocker, monkeypatch -): +async def test_recompute_scim_member_roles_demotes_when_not_in_admin_group(mocker, monkeypatch): """The shared recompute helper writes the non-admin default for a member whose resulting teams no longer include the configured admin group.""" from litellm.proxy.proxy_server import proxy_config @@ -2324,9 +2414,7 @@ async def test_recompute_scim_member_roles_demotes_when_not_in_admin_group( @pytest.mark.asyncio -async def test_recompute_scim_member_roles_grants_when_in_admin_group( - mocker, monkeypatch -): +async def test_recompute_scim_member_roles_grants_when_in_admin_group(mocker, monkeypatch): """The shared recompute helper grants PROXY_ADMIN when a member's resulting teams include the configured admin group.""" from litellm.proxy.proxy_server import proxy_config @@ -2346,9 +2434,7 @@ async def test_recompute_scim_member_roles_grants_when_in_admin_group( @pytest.mark.asyncio -async def test_recompute_scim_member_roles_noop_when_admin_group_unset( - mocker, monkeypatch -): +async def test_recompute_scim_member_roles_noop_when_admin_group_unset(mocker, monkeypatch): """With scim_admin_group unset the recompute helper must not touch any role, preserving current behavior for SCIM group writes.""" from litellm.proxy.proxy_server import proxy_config @@ -2395,16 +2481,10 @@ async def test_update_group_recomputes_roles_for_changed_members(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mocker.MagicMock() - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2452,24 +2532,16 @@ async def test_patch_group_recomputes_roles_for_changed_members(mocker): ) patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation(op="remove", path="members", value=[{"value": "user1"}]) - ], + Operations=[SCIMPatchOperation(op="remove", path="members", value=[{"value": "user1"}])], ) mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mocker.MagicMock() - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2519,9 +2591,7 @@ async def test_delete_group_recomputes_roles_for_members(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) mock_prisma_client.db.litellm_teamtable.delete = AsyncMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=member) @@ -2552,16 +2622,16 @@ async def test_handle_existing_user_by_email_applies_role_when_admin_group_set(m mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=existing_user - ) - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value={"user_id": "new-user-id"} - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={"user_id": "new-user-id"}) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", AsyncMock(return_value=mocker.MagicMock()), ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock(), + ) new_user_request = NewUserRequest( user_id="new-user-id", @@ -2592,16 +2662,16 @@ async def test_handle_existing_user_by_email_leaves_role_when_admin_group_unset( mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=existing_user - ) - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value={"user_id": "new-user-id"} - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={"user_id": "new-user-id"}) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", AsyncMock(return_value=mocker.MagicMock()), ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock(), + ) new_user_request = NewUserRequest( user_id="new-user-id", @@ -2622,9 +2692,7 @@ async def test_handle_existing_user_by_email_leaves_role_when_admin_group_unset( @pytest.mark.asyncio -async def test_create_user_existing_email_upsert_demotes_when_admin_group_set( - mocker, monkeypatch -): +async def test_create_user_existing_email_upsert_demotes_when_admin_group_set(mocker, monkeypatch): """End-to-end create wiring: a SCIM POST that upserts an existing email while the user is not in the admin group must write the non-admin default, not leave a stale PROXY_ADMIN.""" @@ -2650,12 +2718,8 @@ async def test_create_user_existing_email_upsert_demotes_when_admin_group_set( mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_usertable = mocker.MagicMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) - mock_prisma_client.db.litellm_usertable.find_first = AsyncMock( - return_value=existing_user - ) - mock_prisma_client.db.litellm_usertable.update = AsyncMock( - return_value={"user_id": "returning-user"} - ) + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={"user_id": "returning-user"}) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2669,6 +2733,10 @@ async def test_create_user_existing_email_upsert_demotes_when_admin_group_set( "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", AsyncMock(return_value=scim_user), ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._handle_team_membership_changes", + AsyncMock(), + ) await create_user(user=scim_user) @@ -2698,9 +2766,7 @@ async def test_create_group_recomputes_roles_for_members(mocker): mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None) mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mocker.MagicMock() - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2754,16 +2820,10 @@ async def test_update_group_rename_recomputes_retained_members(mocker): mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mocker.MagicMock() - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2808,24 +2868,16 @@ async def test_patch_group_rename_recomputes_retained_members(mocker): ) patch_ops = SCIMPatchOp( schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], - Operations=[ - SCIMPatchOperation(op="replace", path="displayName", value="Engineering") - ], + Operations=[SCIMPatchOperation(op="replace", path="displayName", value="Engineering")], ) mock_prisma_client = mocker.MagicMock() mock_prisma_client.db = mocker.MagicMock() mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() - mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( - return_value=existing_team - ) - mock_prisma_client.db.litellm_teamtable.update = AsyncMock( - return_value=existing_team - ) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=existing_team) mock_prisma_client.db.litellm_usertable = mocker.MagicMock() - mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( - return_value=mocker.MagicMock() - ) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) mocker.patch( "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", @@ -2855,3 +2907,615 @@ async def test_patch_group_rename_recomputes_retained_members(mocker): recompute_mock.assert_awaited_once() assert set(recompute_mock.call_args[0][1]) == {"user1"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_add_retains_existing_members( + mocker, monkeypatch +): + """A SCIM group ``add`` operation must not drop members already in the team. + + Team membership lives in members_with_roles; team creation leaves the legacy + ``members`` column empty. Seeding the patch result from that empty column + made an ``add`` recompute the member set from scratch and remove everyone + already in the team. The result set must be seeded from members_with_roles so + existing members survive an add of a new one. + """ + + async def mock_get_config(): + return {"litellm_settings": {"scim_upsert_user": True}} + + from litellm.proxy.proxy_server import proxy_config + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + existing_team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], # legacy column intentionally empty, as real teams leave it + members_with_roles=[Member(user_id="existing-user", role="user")], + ) + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation(op="add", path="members", value=[{"value": "new-user"}]) + ], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + # new-user already exists in the DB + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mocker.MagicMock(user_id="new-user") + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=mock_prisma_client, + ) + + assert final_members == {"existing-user", "new-user"} + + +@pytest.mark.asyncio +async def test_process_group_patch_operations_remove_uses_members_with_roles( + mocker, monkeypatch +): + """A ``remove`` op must diff against members_with_roles, so removing one + member leaves the rest of the team intact rather than emptying it.""" + + async def mock_get_config(): + return {"litellm_settings": {"scim_upsert_user": True}} + + from litellm.proxy.proxy_server import proxy_config + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + + existing_team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], + members_with_roles=[ + Member(user_id="keep-user", role="user"), + Member(user_id="drop-user", role="user"), + ], + ) + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation( + op="remove", path="members", value=[{"value": "drop-user"}] + ) + ], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mocker.MagicMock(user_id="drop-user") + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=mock_prisma_client, + ) + + assert final_members == {"keep-user"} + + +@pytest.mark.asyncio +async def test_get_groups_reports_members_from_members_with_roles(mocker): + """GET /Groups must report members from members_with_roles (the source of + truth), not the legacy ``members`` column that team creation leaves empty. + Reporting an empty member list makes the IdP repeatedly re-provision.""" + team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], # legacy column empty + members_with_roles=[Member(user_id="member-1", role="user")], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team]) + mock_prisma_client.db.litellm_teamtable.count = AsyncMock(return_value=1) + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=mocker.MagicMock(user_id="member-1", user_email="member-1@example.com") + ) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + + response = await get_groups(startIndex=1, count=10, filter=None) + + assert [m.value for m in response.Resources[0].members] == ["member-1"] + + +@pytest.mark.asyncio +async def test_apply_group_patch_updates_does_not_write_legacy_members(mocker): + """The group PATCH apply must not write the legacy ``members`` column. + + Membership is reconciled onto the source of truth (members_with_roles and + each member's user.teams) separately; writing the legacy column here too + would create a second, unread copy of membership that can drift from the + source of truth, which is the inconsistency this PR removes. + """ + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + updated = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated) + + result = await _apply_group_patch_updates( + group_id="team-1", + update_data={"team_alias": "Renamed"}, + prisma_client=mock_prisma_client, + ) + + assert result is updated + mock_prisma_client.db.litellm_teamtable.update.assert_awaited_once() + written = mock_prisma_client.db.litellm_teamtable.update.call_args.kwargs["data"] + assert "members" not in written + assert written["team_alias"] == "Renamed" + + +def _mock_prisma_for_delete_user(mocker, team): + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.delete = AsyncMock() + return mock_prisma_client + + +def _patch_delete_user_dependencies(mocker, mock_prisma_client, existing_user): + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._check_user_exists", + AsyncMock(return_value=existing_user), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._set_user_keys_blocked", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._delete_rows_referencing_user", + AsyncMock(), + ) + + +@pytest.mark.asyncio +async def test_delete_user_prunes_members_with_roles(mocker): + """Deleting a SCIM user must remove them from every team they belong to via + team_member_delete, which prunes members_with_roles (the source of truth for + SCIM group membership) so GET /Groups no longer returns a dangling reference + to the now-deleted user.""" + user_id = "scim-del-user" + + existing_user = mocker.MagicMock() + existing_user.teams = ["team-1"] + + team = LiteLLM_TeamTable( + team_id="team-1", + members=[user_id, "other-user"], + members_with_roles=[Member(user_id=user_id, role="user"), Member(user_id="other-user", role="admin")], + ) + + mock_prisma_client = _mock_prisma_for_delete_user(mocker, team) + _patch_delete_user_dependencies(mocker, mock_prisma_client, existing_user) + team_member_delete_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(), + ) + + await delete_user(user_id=user_id) + + team_member_delete_mock.assert_awaited_once() + call = team_member_delete_mock.call_args + assert call.kwargs["data"].team_id == "team-1" + assert call.kwargs["data"].user_id == user_id + assert call.kwargs["user_api_key_dict"].user_role == LitellmUserRoles.PROXY_ADMIN + mock_prisma_client.db.litellm_usertable.delete.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_delete_user_surfaces_prune_failure_and_keeps_user(mocker): + """A genuine failure while pruning members_with_roles must surface: the + endpoint fails loudly and the user row is NOT deleted, so we never report a + successful delete while leaving a dangling member (SCIM DELETE is idempotent, + so the IdP retries).""" + user_id = "scim-del-user" + + existing_user = mocker.MagicMock() + existing_user.teams = ["team-1"] + + team = LiteLLM_TeamTable( + team_id="team-1", + members=[user_id], + members_with_roles=[Member(user_id=user_id, role="user")], + ) + + mock_prisma_client = _mock_prisma_for_delete_user(mocker, team) + _patch_delete_user_dependencies(mocker, mock_prisma_client, existing_user) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(side_effect=Exception("database connection lost")), + ) + + with pytest.raises(Exception): + await delete_user(user_id=user_id) + + mock_prisma_client.db.litellm_usertable.delete.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_delete_user_skips_teams_where_not_a_member(mocker): + """If the user is not in a team's members_with_roles, deletion must treat that + team as a no-op (no team_member_delete call, no error) and still delete the + user, so a stale legacy membership can't block the delete.""" + user_id = "scim-del-user" + + existing_user = mocker.MagicMock() + existing_user.teams = ["team-1"] + + team = LiteLLM_TeamTable( + team_id="team-1", + members=[user_id], + members_with_roles=[Member(user_id="someone-else", role="admin")], + ) + + mock_prisma_client = _mock_prisma_for_delete_user(mocker, team) + _patch_delete_user_dependencies(mocker, mock_prisma_client, existing_user) + team_member_delete_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(), + ) + + await delete_user(user_id=user_id) + + team_member_delete_mock.assert_not_awaited() + mock_prisma_client.db.litellm_usertable.delete.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_patch_group_add_applies_delta_and_keeps_concurrent_add(mocker): + """A group PATCH op:add must be applied as a delta against the live roster, + not as a snapshot-based absolute target. + + When a concurrent PATCH has already added a member between this request's + initial read and its post-write refresh, that member shows up in the + refreshed roster but not in this request's snapshot-derived target. Diffing + the refreshed roster against the snapshot target would issue a spurious + team_member_delete for the concurrently-added member. Applying only this + request's intended delta on top of the refreshed roster must retain them. + """ + from litellm.proxy.management_endpoints.scim.scim_transformations import ( + ScimTransformations, + ) + + group_id = "team-concurrent" + + snapshot_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[Member(user_id="zed", role="user")], + metadata={"externalId": "grp-ext"}, + ) + refreshed_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[ + Member(user_id="zed", role="user"), + Member(user_id="alice", role="user"), + ], + metadata={"externalId": "grp-ext"}, + ) + final_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[ + Member(user_id="zed", role="user"), + Member(user_id="alice", role="user"), + Member(user_id="bob", role="user"), + ], + metadata={"externalId": "grp-ext"}, + ) + + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path="members", value=[{"value": "bob"}])], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + side_effect=[snapshot_team, refreshed_team, final_team] + ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=final_team) + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + patch_membership_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + mocker.patch.object( + ScimTransformations, + "transform_litellm_team_to_scim_group", + AsyncMock( + return_value=SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Group", + ) + ), + ) + + await patch_group(group_id=group_id, patch_ops=patch_ops) + + calls = patch_membership_mock.call_args_list + + removed_user_ids = { + call.kwargs["user_id"] for call in calls if call.kwargs.get("teams_ids_to_remove_user_from") == [group_id] + } + assert removed_user_ids == set() + + added_user_ids = { + call.kwargs["user_id"] for call in calls if call.kwargs.get("teams_ids_to_add_user_to") == [group_id] + } + assert added_user_ids == {"bob"} + + +@pytest.mark.asyncio +async def test_patch_group_replace_stays_absolute_against_concurrent_roster(mocker): + """A group PATCH ``replace`` op declares the roster is exactly the given set, + so it must reconcile as a set-to-target, not as a delta. + + Unlike ``add``/``remove``, ``replace`` is absolute. A member that another + request added concurrently is present in the refreshed roster but not in the + replace target, and ``replace`` must drop it. Rebasing the replace onto the + refreshed roster (the delta behavior correct only for add/remove) would + wrongly retain that concurrently-added member. + """ + from litellm.proxy.management_endpoints.scim.scim_transformations import ( + ScimTransformations, + ) + + group_id = "team-replace-concurrent" + + snapshot_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[Member(user_id="zed", role="user")], + metadata={"externalId": "grp-ext"}, + ) + refreshed_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[ + Member(user_id="alice", role="user"), + Member(user_id="bob", role="user"), + ], + metadata={"externalId": "grp-ext"}, + ) + final_team = LiteLLM_TeamTable( + team_id=group_id, + team_alias="Group", + members_with_roles=[Member(user_id="alice", role="user")], + metadata={"externalId": "grp-ext"}, + ) + + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="replace", path="members", value=[{"value": "alice"}])], + ) + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable = mocker.MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + side_effect=[snapshot_team, refreshed_team, final_team] + ) + mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=final_team) + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mocker.MagicMock()) + + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception", + AsyncMock(return_value=mock_prisma_client), + ) + patch_membership_mock = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._recompute_scim_member_roles", + AsyncMock(), + ) + mocker.patch.object( + ScimTransformations, + "transform_litellm_team_to_scim_group", + AsyncMock( + return_value=SCIMGroup( + schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"], + id=group_id, + displayName="Group", + ) + ), + ) + + await patch_group(group_id=group_id, patch_ops=patch_ops) + + calls = patch_membership_mock.call_args_list + + removed_user_ids = { + call.kwargs["user_id"] for call in calls if call.kwargs.get("teams_ids_to_remove_user_from") == [group_id] + } + assert removed_user_ids == {"bob"} + + added_user_ids = { + call.kwargs["user_id"] for call in calls if call.kwargs.get("teams_ids_to_add_user_to") == [group_id] + } + assert added_user_ids == set() + + +@pytest.mark.parametrize( + "path, attribute, expected", + [ + ('members[value eq "user-1"]', "members", ["user-1"]), + ("members[value eq 'user-1']", "members", ["user-1"]), + ('members[value EQ "user-1"]', "members", ["user-1"]), + ('members[ value eq "user-1" ]', "members", ["user-1"]), + ('groups[value eq "team-1"]', "groups", ["team-1"]), + ('members[value eq "Mixed-CASE-Id"]', "members", ["Mixed-CASE-Id"]), + ('members[value eq "a\\"b"]', "members", ['a"b']), + ('members[value eq "a\\\\b"]', "members", ["a\\b"]), + ("members[value eq 'a\\'b']", "members", ["a'b"]), + ("members", "members", []), + ('groups[value eq "team-1"]', "members", []), + (None, "members", []), + ('members[value eq ""]', "members", []), + ("members[value eq user-1]", "members", []), + ("members[value eq unintendeduser]", "members", []), + ], +) +def test_extract_ids_from_path_filter(path, attribute, expected): + assert _extract_ids_from_path_filter(path, attribute) == expected + + +def test_extract_ids_from_path_filter_unterminated_is_linear(): + """A pathological unterminated quoted filter must not trigger super-linear + backtracking; it returns no id and completes near-instantly.""" + pathological = 'members[value eq "' + ("\\" * 200) + + start = time.perf_counter() + result = _extract_ids_from_path_filter(pathological, "members") + elapsed = time.perf_counter() - start + + assert result == [] + assert elapsed < 1.0 + + +@pytest.mark.asyncio +async def test_process_group_patch_remove_filtered_path_without_value(mocker): + """Okta sends group membership removals as a filtered path with no request + body value; the member id must be parsed out of members[value eq "..."]""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="remove", path='members[value eq "user-1"]')], + ) + + existing_team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], + members_with_roles=[ + Member(user_id="user-1", role="user"), + Member(user_id="user-2", role="user"), + ], + ) + + prisma_client = mocker.MagicMock() + prisma_client.db = mocker.MagicMock() + prisma_client.db.litellm_usertable = mocker.MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="user-1") + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=prisma_client, + ) + + assert final_members == {"user-2"} + + +@pytest.mark.asyncio +async def test_process_group_patch_add_filtered_path_without_value(mocker): + """A filtered add path with no body value adds the id parsed from the filter.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[SCIMPatchOperation(op="add", path='members[value eq "user-3"]')], + ) + + existing_team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], + members_with_roles=[Member(user_id="user-1", role="user")], + ) + + prisma_client = mocker.MagicMock() + prisma_client.db = mocker.MagicMock() + prisma_client.db.litellm_usertable = mocker.MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="user-3") + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=prisma_client, + ) + + assert final_members == {"user-1", "user-3"} + + +@pytest.mark.asyncio +async def test_process_group_patch_replace_empty_value_does_not_use_path_filter(mocker): + """An explicit empty replace value must clear membership rather than pull an + id from the filtered path, which would retain one member and drop the rest.""" + patch_ops = SCIMPatchOp( + schemas=["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + Operations=[ + SCIMPatchOperation(op="replace", path='members[value eq "user-1"]', value=[]) + ], + ) + + existing_team = LiteLLM_TeamTable( + team_id="team-1", + team_alias="Team One", + members=[], + members_with_roles=[ + Member(user_id="user-1", role="user"), + Member(user_id="user-2", role="user"), + ], + ) + + prisma_client = mocker.MagicMock() + prisma_client.db = mocker.MagicMock() + prisma_client.db.litellm_usertable = mocker.MagicMock() + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + return_value=LiteLLM_UserTable(user_id="user-1") + ) + + _, final_members, _ = await _process_group_patch_operations( + patch_ops=patch_ops, + existing_team=existing_team, + prisma_client=prisma_client, + ) + + assert final_members == set() diff --git a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py index f4c6d4f8d15..2504b5744fc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_cache_settings_endpoints.py @@ -17,9 +17,14 @@ from litellm.proxy._types import LitellmTableNames, LitellmUserRoles from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth from litellm.proxy.management_endpoints.cache_settings_endpoints import ( _CACHE_SENSITIVE_FIELDS, + _REDACTED_VALUE, CacheSettingsManager, CacheSettingsUpdateRequest, CacheTestRequest, + _merge_over_saved, + _overlay_environment, + _parse_stored_settings, + _redact_credentials, _resolve_cache_url_precedence, get_cache_settings, test_cache_connection, @@ -610,3 +615,510 @@ async def test_update_cache_settings_no_audit_when_disabled(monkeypatch): ) assert audit_calls == [] + + +class TestParseStoredSettings: + """The stored blob arrives as a JSON string or a parsed dict; both must + normalize to a dict so the secret-preservation read never silently drops it.""" + + def test_parses_a_json_string(self): + assert _parse_stored_settings('{"host": "h", "password": "pw"}') == {"host": "h", "password": "pw"} + + def test_passes_a_dict_through(self): + assert _parse_stored_settings({"host": "h", "password": "pw"}) == {"host": "h", "password": "pw"} + + def test_non_mapping_becomes_empty(self): + assert _parse_stored_settings(None) == {} + assert _parse_stored_settings("[1, 2]") == {} + + +class TestMergeOverSaved: + """The secret-preservation contract behind the redacted-resubmit fix.""" + + def test_redacted_secret_restores_stored_value(self): + # same connection target, an unrelated field edited: the stored secret + # is restored behind the redacted resubmit + merged = _merge_over_saved( + incoming={"type": "redis", "host": "samehost", "namespace": "new", "password": _REDACTED_VALUE}, + saved={"type": "redis", "host": "samehost", "password": "realpw"}, + ) + assert merged["namespace"] == "new" + assert merged["password"] == "realpw" + + def test_stored_secret_not_replayed_to_a_different_target(self): + # credential replay guard: omitting the password while pointing at a new + # host must NOT resurrect the stored secret (it would be sent elsewhere) + merged = _merge_over_saved( + incoming={"type": "redis", "host": "attacker.example.com", "password": _REDACTED_VALUE}, + saved={"type": "redis", "host": "real-redis", "password": "realpw"}, + ) + assert "password" not in merged + + def test_omitted_secret_restores_stored_value(self): + # same host (target unchanged), password field omitted entirely + merged = _merge_over_saved( + incoming={"type": "redis", "host": "samehost", "namespace": "n"}, + saved={"type": "redis", "host": "samehost", "password": "realpw"}, + ) + assert merged["password"] == "realpw" + + def test_sentinel_password_not_replayed_to_different_sentinel_nodes(self): + # sentinel target change with an omitted sentinel_password must not + # resurrect the stored one and send it to the caller's sentinels + merged = _merge_over_saved( + incoming={"type": "redis", "sentinel_nodes": [["attacker", 26379]], "service_name": "mymaster"}, + saved={ + "type": "redis", + "sentinel_nodes": [["real", 26379]], + "service_name": "mymaster", + "sentinel_password": "realsp", + }, + ) + assert "sentinel_password" not in merged + + def test_sentinel_password_preserved_when_sentinel_target_unchanged(self): + merged = _merge_over_saved( + incoming={"type": "redis", "sentinel_nodes": [["real", 26379]], "service_name": "mymaster"}, + saved={ + "type": "redis", + "sentinel_nodes": [["real", 26379]], + "service_name": "mymaster", + "sentinel_password": "realsp", + }, + ) + assert merged["sentinel_password"] == "realsp" + + def test_password_not_replayed_to_different_cluster_nodes(self): + merged = _merge_over_saved( + incoming={"type": "redis", "redis_startup_nodes": [{"host": "attacker", "port": "7001"}]}, + saved={ + "type": "redis", + "redis_startup_nodes": [{"host": "real", "port": "7001"}], + "password": "realpw", + }, + ) + assert "password" not in merged + + def test_equivalent_target_representations_still_preserve_secret(self): + # the client sends port as a string, storage holds it as an int: the + # target is unchanged, so the untouched password must not be dropped + merged = _merge_over_saved( + incoming={"type": "redis", "host": "h", "port": "6379", "password": _REDACTED_VALUE}, + saved={"type": "redis", "host": "h", "port": 6379, "password": "realpw"}, + ) + assert merged["password"] == "realpw" + + def test_explicit_empty_string_clears_the_secret(self): + merged = _merge_over_saved( + incoming={"type": "redis", "host": "h", "password": ""}, + saved={"type": "redis", "host": "h", "password": "realpw"}, + ) + assert merged.get("password") == "" + + def test_explicit_null_clears_the_secret(self): + # an explicit null is a clear, not an omission, so it must not restore + merged = _merge_over_saved( + incoming={"type": "redis", "host": "h", "password": None}, + saved={"type": "redis", "host": "h", "password": "realpw"}, + ) + assert merged.get("password") is None + + def test_secret_not_reused_when_a_pinned_target_field_is_omitted(self): + # omitting the host (a pinned target) means the request does not describe + # the stored target, so the stored secret must not be restored (and thus + # cannot be sent to whatever host the incomplete request resolves to) + merged = _merge_over_saved( + incoming={"type": "redis", "port": "6379"}, + saved={"type": "redis", "host": "real", "port": 6379, "password": "realpw"}, + ) + assert "password" not in merged + + def test_redacted_secret_with_no_stored_value_is_dropped(self): + # env-sourced secret: nothing stored to restore, so the marker must not + # be persisted; the environment stays the source at runtime + merged = _merge_over_saved( + incoming={"type": "redis", "host": "h", "password": _REDACTED_VALUE}, + saved={}, + ) + assert "password" not in merged + + def test_new_secret_value_wins(self): + merged = _merge_over_saved( + incoming={"password": "brandnewpw"}, + saved={"password": "realpw"}, + ) + assert merged["password"] == "brandnewpw" + + def test_switching_from_url_to_host_port_drops_stored_url(self): + # admin migrates a url-mode cache to discrete host/port: the stored url + # must not be resurrected (url precedence would then discard host/port) + merged = _merge_over_saved( + incoming={"type": "redis", "host": "newhost", "port": "6379"}, + saved={"type": "redis", "url": "redis://:pw@oldhost:6379/0"}, + ) + assert "url" not in merged + assert merged["host"] == "newhost" + assert merged["port"] == "6379" + + def test_untouched_url_is_preserved_without_a_discrete_target(self): + # a url-mode save that touches nothing keeps the stored url + merged = _merge_over_saved( + incoming={"type": "redis", "namespace": "ns"}, + saved={"type": "redis", "url": "redis://:pw@host:6379/0"}, + ) + assert merged["url"] == "redis://:pw@host:6379/0" + + +def test_overlay_environment_fills_unset_connection_fields(monkeypatch): + """A cache with no stored connection resolves REDIS_* env for the UI.""" + for var in ("REDIS_URL", "REDIS_HOST", "REDIS_PORT", "REDIS_PASSWORD", "REDIS_USERNAME"): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("REDIS_HOST", "redis.internal") + monkeypatch.setenv("REDIS_PORT", "6380") + monkeypatch.setenv("REDIS_PASSWORD", "env-password") + + effective = _overlay_environment({}) + + assert effective["host"] == "redis.internal" + assert effective["port"] == "6380" + assert effective["password"] == "env-password" + assert effective["type"] == "redis" + + +def test_overlay_environment_stored_value_wins(monkeypatch): + monkeypatch.setenv("REDIS_HOST", "env-host") + effective = _overlay_environment({"type": "redis", "host": "stored-host"}) + assert effective["host"] == "stored-host" + + +@pytest.mark.asyncio +async def test_get_cache_settings_falls_back_to_redis_env(monkeypatch): + """A cache configured purely through REDIS_* env vars shows its effective + connection instead of a blank page, with the password redacted.""" + for var in ("REDIS_URL", "REDIS_HOST", "REDIS_PORT", "REDIS_PASSWORD", "REDIS_USERNAME"): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("REDIS_HOST", "redis.internal") + monkeypatch.setenv("REDIS_PORT", "6380") + monkeypatch.setenv("REDIS_PASSWORD", "env-password") + + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + ): + response = await get_cache_settings(user_api_key_dict=_admin_auth()) + + values = response.current_values + assert values["host"] == "redis.internal" + assert values["port"] == "6380" + assert values["type"] == "redis" + # the env password is surfaced as configured, not leaked in plaintext + assert values["password"] == _REDACTED_VALUE + + +@pytest.mark.asyncio +async def test_get_cache_settings_redacts_password_with_marker(monkeypatch): + for var in ("REDIS_URL", "REDIS_HOST", "REDIS_PORT", "REDIS_PASSWORD", "REDIS_USERNAME"): + monkeypatch.delenv(var, raising=False) + cache_row = MagicMock() + cache_row.cache_settings = json.dumps( + {"type": "redis", "host": "h", "password": "supersecret", "namespace": "ns"} + ) + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + ): + response = await get_cache_settings(user_api_key_dict=_admin_auth()) + + assert response.current_values["password"] == _REDACTED_VALUE + assert response.current_values["namespace"] == "ns" + + +@pytest.mark.asyncio +async def test_get_cache_settings_url_mode_hides_env_discrete_fields(monkeypatch): + """A url-mode stored config must not surface env-overlaid host/port. + + Otherwise a no-op save would submit the env host and, via url precedence, + silently switch the cache off its configured url. + """ + for var in ("REDIS_URL", "REDIS_HOST", "REDIS_PORT", "REDIS_PASSWORD", "REDIS_USERNAME"): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("REDIS_HOST", "env-host") + monkeypatch.setenv("REDIS_PORT", "6380") + + cache_row = MagicMock() + cache_row.cache_settings = {"type": "redis", "url": "redis://:pw@stored-host:6379/0"} + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=cache_row) + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + ): + response = await get_cache_settings(user_api_key_dict=_admin_auth()) + + values = response.current_values + assert values["url"] == _REDACTED_VALUE + # the env host/port must not leak in and shadow the url + assert "host" not in values + assert "port" not in values + + +def _mock_proxy_config_identity_crypto(): + proxy_config = MagicMock() + proxy_config._encrypt_env_variables = MagicMock( + side_effect=lambda environment_variables: dict(environment_variables) + ) + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + proxy_config._init_cache = MagicMock() + proxy_config.switch_on_llm_response_caching = MagicMock() + return proxy_config + + +@pytest.mark.asyncio +async def test_update_preserves_stored_password_on_redacted_resubmit(monkeypatch): + """Editing an unrelated field and re-submitting the redacted password must + keep the stored secret, not persist the marker over a working password.""" + monkeypatch.setattr(litellm, "store_audit_logs", False) + + existing = MagicMock() + # prisma returns the Json column as an already-parsed dict, not a JSON + # string; a reader that json.loads unconditionally would drop the whole row + existing.cache_settings = {"type": "redis", "host": "oldhost", "password": "realpw"} + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() + proxy_config = _mock_proxy_config_identity_crypto() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + result = await update_cache_settings( + request=CacheSettingsUpdateRequest( + # same host (the target is unchanged), an unrelated field edited + cache_settings={"type": "redis", "host": "oldhost", "namespace": "edited", "password": _REDACTED_VALUE} + ), + user_api_key_dict=_admin_auth(), + litellm_changed_by=None, + ) + + persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + assert persisted["host"] == "oldhost" + assert persisted["namespace"] == "edited" + assert persisted["password"] == "realpw" + # the response never echoes the plaintext secret back either + assert result["settings"]["password"] == _REDACTED_VALUE + + +@pytest.mark.asyncio +async def test_update_drops_env_sourced_redacted_secret(monkeypatch): + """With no stored row, a re-submitted redacted secret is env-sourced; the + marker must not be persisted so the environment stays the source.""" + monkeypatch.setattr(litellm, "store_audit_logs", False) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() + proxy_config = _mock_proxy_config_identity_crypto() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + await update_cache_settings( + request=CacheSettingsUpdateRequest( + cache_settings={"type": "redis", "host": "h", "password": _REDACTED_VALUE} + ), + user_api_key_dict=_admin_auth(), + litellm_changed_by=None, + ) + + persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + assert "password" not in persisted + + +@pytest.mark.asyncio +async def test_update_applies_new_password(monkeypatch): + """A real new secret value replaces the stored one.""" + monkeypatch.setattr(litellm, "store_audit_logs", False) + + existing = MagicMock() + existing.cache_settings = json.dumps({"type": "redis", "host": "h", "password": "oldpw"}) + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + mock_prisma.db.litellm_cacheconfig.upsert = AsyncMock() + proxy_config = _mock_proxy_config_identity_crypto() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + ): + await update_cache_settings( + request=CacheSettingsUpdateRequest( + cache_settings={"type": "redis", "host": "h", "password": "brandnewpw"} + ), + user_api_key_dict=_admin_auth(), + litellm_changed_by=None, + ) + + persisted = proxy_config._encrypt_env_variables.call_args.kwargs["environment_variables"] + assert persisted["password"] == "brandnewpw" + + +@pytest.mark.asyncio +async def test_test_cache_connection_survives_saved_lookup_failure(monkeypatch): + """A failed saved-settings lookup must not block the connection test. + + The test endpoint reads the stored row to resolve a redacted credential, but + that read can raise (a misconfigured or unavailable client), and it must fall + back to the submitted settings rather than abort — otherwise a shared client + left in an odd state by another test would break every connection test. + """ + monkeypatch.setattr(litellm, "store_audit_logs", False) + + # a client whose find_unique is not awaitable, so the saved read raises + bad_prisma = MagicMock() + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + cache_instance = MagicMock() + cache_instance.cache = MagicMock() + cache_instance.cache.test_connection = AsyncMock(return_value={"status": "success", "message": "ok"}) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", bad_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.Cache") as mock_cache_class, + ): + mock_cache_class.return_value = cache_instance + result = await test_cache_connection( + request=CacheTestRequest(cache_settings={"type": "redis", "host": "h", "port": "6379", "password": "pw"}), + user_api_key_dict=_admin_auth(), + ) + + mock_cache_class.assert_called_once() + assert result.status == "success" + + +@pytest.mark.asyncio +async def test_get_cache_settings_does_not_surface_non_display_env_credentials(monkeypatch): + """The env overlay must not leak credential kwargs the UI does not manage. + + _redis_kwargs_from_environment resolves every redis.Redis kwarg, including + secrets like azure_client_secret; only cache display fields may be surfaced, + so a non-admin reading /cache/settings never retrieves such a credential. + """ + for var in ("REDIS_URL", "REDIS_HOST", "REDIS_PORT", "REDIS_PASSWORD", "REDIS_AZURE_CLIENT_SECRET"): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("REDIS_HOST", "redis.internal") + monkeypatch.setenv("REDIS_AZURE_CLIENT_SECRET", "super-azure-secret") + + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=None) + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + ): + response = await get_cache_settings(user_api_key_dict=_admin_auth()) + + values = response.current_values + assert values.get("host") == "redis.internal" + # the non-display credential must not appear in the response at all + assert "azure_client_secret" not in values + assert "super-azure-secret" not in values.values() + + +@pytest.mark.asyncio +async def test_test_cache_connection_does_not_log_plaintext_credentials(monkeypatch, caplog): + """The connection test must not write the resolved plaintext secret to logs. + + _merge_over_saved substitutes the stored password for a redacted resubmit, so + the settings dict carries the real secret; the debug log must redact it. + """ + import logging + + existing = MagicMock() + existing.cache_settings = {"type": "redis", "host": "h", "port": "6379", "password": "realredispw"} + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + cache_instance = MagicMock() + cache_instance.cache = MagicMock() + cache_instance.cache.test_connection = AsyncMock(return_value={"status": "success", "message": "ok"}) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.Cache") as mock_cache_class, + caplog.at_level(logging.DEBUG, logger="LiteLLM Proxy"), + ): + mock_cache_class.return_value = cache_instance + # resubmit the redacted marker; the merge resolves it to the stored secret + await test_cache_connection( + request=CacheTestRequest( + cache_settings={"type": "redis", "host": "h", "port": "6379", "password": _REDACTED_VALUE} + ), + user_api_key_dict=_admin_auth(), + ) + + # the real password was used to build the client but never written to the log + assert mock_cache_class.call_args.kwargs["password"] == "realredispw" + assert "realredispw" not in caplog.text + + +@pytest.mark.asyncio +async def test_test_cache_connection_does_not_replay_saved_password_to_new_host(monkeypatch): + """Credential-replay guard on the connection test. + + A caller that submits a different host while omitting the password must not + have the stored password restored and sent to the caller-chosen host. + """ + existing = MagicMock() + existing.cache_settings = {"type": "redis", "host": "real-redis", "port": "6379", "password": "realredispw"} + mock_prisma = MagicMock() + mock_prisma.db.litellm_cacheconfig.find_unique = AsyncMock(return_value=existing) + proxy_config = MagicMock() + proxy_config._decrypt_db_variables = MagicMock(side_effect=lambda variables_dict: dict(variables_dict)) + + cache_instance = MagicMock() + cache_instance.cache = MagicMock() + cache_instance.cache.test_connection = AsyncMock(return_value={"status": "success", "message": "ok"}) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch("litellm.proxy.proxy_server.proxy_config", proxy_config), + patch("litellm.Cache") as mock_cache_class, + ): + mock_cache_class.return_value = cache_instance + await test_cache_connection( + request=CacheTestRequest( + cache_settings={"type": "redis", "host": "attacker.example.com", "port": "6379"} + ), + user_api_key_dict=_admin_auth(), + ) + + called_kwargs = mock_cache_class.call_args.kwargs + # the stored password is NOT sent to the attacker-chosen host + assert called_kwargs.get("password") != "realredispw" + assert "password" not in called_kwargs diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index dffca3093fa..51f72f91dc3 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -14994,3 +14994,76 @@ async def test_list_keys_without_expires_param_forwards_none(): mock_helper.assert_called_once() assert mock_helper.call_args.kwargs["expires_filter"] is None + + +@pytest.mark.asyncio +@patch( + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_sso_identity_assertions_master_key" +) +@patch( + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_env_vars_master_key" +) +@patch( + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_user_credentials_master_key" +) +@patch( + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key" +) +async def test_rotate_master_key_rotates_sso_identity_assertions( + mock_rotate_mcp_server, + mock_rotate_mcp_user, + mock_rotate_env_vars, + mock_rotate_sso, +): + """Master-key rotation must re-encrypt the SSO identity assertion store alongside + the sibling per-user encrypted tables, or a salt rotation orphans every stored + assertion (step 4d).""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _rotate_master_key, + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_tx = AsyncMock() + mock_tx.litellm_proxymodeltable = MagicMock() + mock_tx.litellm_proxymodeltable.delete_many = AsyncMock() + mock_tx.litellm_proxymodeltable.create_many = AsyncMock() + mock_prisma_client.db.tx = MagicMock( + return_value=AsyncMock( + __aenter__=AsyncMock(return_value=mock_tx), + __aexit__=AsyncMock(return_value=False), + ) + ) + mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[]) + mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock( + return_value=[] + ) + + mock_proxy_config = MagicMock() + mock_proxy_config.decrypt_model_list_from_db.return_value = [] + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="test-user", + ) + + with patch( + "litellm.proxy.proxy_server.proxy_config", + mock_proxy_config, + ): + await _rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=user_api_key_dict, + current_master_key="sk-old-master-key", + new_master_key="sk-new-master-key", + ) + + mock_rotate_sso.assert_awaited_once_with( + prisma_client=mock_prisma_client, + new_master_key="sk-new-master-key", + ) diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 3e5bd3e9b7f..e1aaf398f97 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -2,6 +2,7 @@ import os import sys import types import json +from contextlib import ExitStack from datetime import datetime, timedelta from types import SimpleNamespace from typing import List, Optional @@ -2025,8 +2026,469 @@ class TestTemporaryMCPSessionEndpoints: code_challenge_method="S256", response_type="code", scope="scope1", + ephemeral_dcr_client=None, ) + async def _authorize_without_client_id( + self, server, mint_mock=None, code_challenge="chal", code_challenge_method="S256" + ): + """Drive mcp_authorize with no caller client_id against ``server``, returning the + (authorize_with_server mock, raised HTTPException or None) pair. Sends a valid S256 PKCE + pair by default because the ephemeral mint requires it.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_authorize, + ) + + request = MagicMock() + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + patches = [ + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.authorize_with_server", + AsyncMock(return_value=MagicMock()), + ), + ] + if mint_mock is not None: + patches.append( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client", + mint_mock, + ) + ) + with ExitStack() as stack: + entered = [stack.enter_context(p) for p in patches] + authorize_mock = entered[1] + try: + await mcp_authorize( + request=request, + server_id=server.server_id, + user_api_key_dict=admin_auth, + client_id=None, + redirect_uri="http://127.0.0.1:60108/callback", + state="state123", + code_challenge=code_challenge, + code_challenge_method=code_challenge_method, + ) + except HTTPException as exc: + return authorize_mock, exc + return authorize_mock, None + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "code_challenge, code_challenge_method", + [(None, None), ("chal", "plain"), ("chal", None)], + ) + async def test_mcp_authorize_mint_requires_s256_pkce(self, code_challenge, code_challenge_method): + """Without PKCE the sealed code would be bearer-redeemable by any authenticated caller who + intercepts the redirect, so the ephemeral mint refuses to run for a downgraded flow (no + challenge, or a non-S256 method) before any upstream registration happens.""" + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = "https://idp.example.com/register" + mint_mock = AsyncMock() + + authorize_mock, exc = await self._authorize_without_client_id( + server, mint_mock=mint_mock, code_challenge=code_challenge, code_challenge_method=code_challenge_method + ) + + assert exc is not None + assert exc.status_code == 400 + assert "PKCE" in str(exc.detail) + mint_mock.assert_not_awaited() + authorize_mock.assert_not_awaited() + + @pytest.mark.asyncio + @pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate]) + async def test_mcp_authorize_client_forwarded_modes_mint_ephemeral_dcr_client_when_none_supplied(self, auth_type): + """LIT-4581 regression: a client-forwarded-token server created without an auth step has no + stored client_id and the tools-tab browser flow supplies none, so authorize must fall + through to a gateway-side DCR mint and proceed with the minted client instead of + dead-ending on a 400 missing_client_id. Both modes share the caller-held-client contract, + so both get the fall-through.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + EphemeralDcrClient, + ) + + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = auth_type + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = "https://idp.example.com/register" + minted = EphemeralDcrClient(client_id="minted-77", client_secret="mint-secret") + mint_mock = AsyncMock(return_value=minted) + + authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock) + + assert exc is None + mint_mock.assert_awaited_once() + assert authorize_mock.await_args.kwargs["client_id"] == "minted-77" + assert authorize_mock.await_args.kwargs["ephemeral_dcr_client"] is minted + + @pytest.mark.asyncio + async def test_mcp_authorize_rejects_untrusted_redirect_before_minting(self): + """An untrusted redirect_uri must be rejected before the gateway performs any upstream + registration, so bad-redirect requests cannot be used to generate orphan clients at the + IdP.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_authorize, + ) + + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = "https://idp.example.com/register" + mint_mock = AsyncMock() + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + request = MagicMock() + request.base_url = "https://litellm.example.com/" + request.headers = {} + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.mint_ephemeral_dcr_client", + mint_mock, + ), + ): + with pytest.raises(HTTPException) as exc: + await mcp_authorize( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + client_id=None, + redirect_uri="https://evil.example.net/steal", + state="state123", + code_challenge="chal", + code_challenge_method="S256", + ) + + assert exc.value.status_code == 400 + mint_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_authorize_true_passthrough_without_authorization_url_reports_the_real_fault(self): + """A passthrough server whose discovery never yielded an authorize endpoint cannot start any + flow, minted client or not, so the error names the missing authorization url instead of the + misleading missing_client_id remedy.""" + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + server.authorization_url = None + server.registration_url = "https://idp.example.com/register" + mint_mock = AsyncMock() + + authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock) + + assert exc is not None + assert exc.status_code == 400 + assert "authorization url" in str(exc.detail) + mint_mock.assert_not_awaited() + authorize_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_authorize_true_passthrough_without_registration_endpoint_keeps_missing_client_id(self): + """When the upstream exposes no registration endpoint the mint is impossible, so the + authorize fails closed with the existing missing_client_id 400 instead of proceeding with an + empty client.""" + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = None + + authorize_mock, exc = await self._authorize_without_client_id(server) + + assert exc is not None + assert exc.status_code == 400 + assert exc.detail["error"] == "missing_client_id" + authorize_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_authorize_oauth2_server_does_not_mint(self): + """The ephemeral mint is scoped to the client-forwarded-token modes: a plain oauth2 server + keeps the gateway-held-client contract (its client is persisted by the admin register flow), + so an empty client_id stays a 400 and no upstream registration is attempted.""" + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.oauth2 + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = "https://idp.example.com/register" + mint_mock = AsyncMock() + + authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock) + + assert exc is not None + assert exc.status_code == 400 + assert exc.detail["error"] == "missing_client_id" + mint_mock.assert_not_awaited() + authorize_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_authorize_true_passthrough_dcr_bridge_mints_too(self): + """The UI creates passthrough servers with dcr_bridge enabled by default, so the default + clientless tools-page authorize is a bridge server; it must mint exactly like a non-bridge + one (the minted flow runs the bridge short-circuit arm) instead of dead-ending on + missing_client_id. The relay front door stays reserved for clients that present their own + client_id.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + EphemeralDcrClient, + ) + + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + server.dcr_bridge = True + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = "https://idp.example.com/register" + minted = EphemeralDcrClient(client_id="minted-77", client_secret=None) + mint_mock = AsyncMock(return_value=minted) + + authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock) + + assert exc is None + mint_mock.assert_awaited_once() + assert authorize_mock.await_args.kwargs["client_id"] == "minted-77" + assert authorize_mock.await_args.kwargs["ephemeral_dcr_client"] is minted + + @pytest.mark.asyncio + async def test_mcp_authorize_oauth_delegate_dcr_bridge_does_not_mint(self): + """The interactive oauth_delegate dcr_bridge sign-in has its own sealed-identity flow that + captures the SSO user at authorize; the ephemeral mint must not preempt it.""" + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.oauth_delegate + server.dcr_bridge = True + server.authorization_url = "https://idp.example.com/authorize" + server.registration_url = "https://idp.example.com/register" + mint_mock = AsyncMock() + + authorize_mock, exc = await self._authorize_without_client_id(server, mint_mock=mint_mock) + + assert exc is not None + assert exc.status_code == 400 + assert exc.detail["error"] == "missing_client_id" + mint_mock.assert_not_awaited() + authorize_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_token_opens_sealed_passthrough_code_and_exchanges_with_minted_client(self): + """LIT-4581 regression, token leg: the client echoes back the sealed passthrough code the + callback forwarded, so the token endpoint recovers the ephemeral client and the real + upstream code from it and authenticates the exchange with them, with no client_id supplied + by the caller and none stored on the server.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + seal_passthrough_authorization_code, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_token, + ) + + request = MagicMock() + request.base_url = "https://litellm.example.com/" + request.headers = {} + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server", + AsyncMock(return_value={"access_token": "token"}), + ) as exchange_mock, + ): + sealed = seal_passthrough_authorization_code( + upstream_code="up-code", + client_id="minted-77", + client_secret="mint-secret", + mcp_server_id="server-1", + token_endpoint_auth_method="client_secret_basic", + ) + result = await mcp_token( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + grant_type="authorization_code", + code=sealed, + redirect_uri="https://example.com/callback", + client_id=None, + client_secret=None, + code_verifier="verifier", + refresh_token=None, + scope=None, + ) + + assert result == {"access_token": "token"} + assert exchange_mock.await_args.kwargs["code"] == "up-code" + assert exchange_mock.await_args.kwargs["client_id"] == "minted-77" + assert exchange_mock.await_args.kwargs["client_secret"] == "mint-secret" + assert exchange_mock.await_args.kwargs["redirect_uri"] == "https://litellm.example.com/callback" + assert exchange_mock.await_args.kwargs["client_token_endpoint_auth_method"] == "client_secret_basic" + + @pytest.mark.asyncio + async def test_mcp_token_refresh_grant_never_opens_sealed_code(self): + """The minted client is unrecoverable outside the single authorization_code flow by + contract: a refresh_token grant that echoes a leftover sealed passthrough code (plus any + verifier) must not recover the minted credentials, so a clientless server answers + missing_client_id and the client re-runs authorize instead.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + seal_passthrough_authorization_code, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_token, + ) + + request = MagicMock() + request.base_url = "https://litellm.example.com/" + request.headers = {} + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server", + AsyncMock(return_value={"access_token": "token"}), + ) as exchange_mock, + ): + sealed = seal_passthrough_authorization_code( + upstream_code="up-code", + client_id="minted-77", + client_secret="mint-secret", + mcp_server_id="server-1", + token_endpoint_auth_method="client_secret_basic", + ) + with pytest.raises(HTTPException) as exc: + await mcp_token( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + grant_type="refresh_token", + code=sealed, + redirect_uri="https://example.com/callback", + client_id=None, + client_secret=None, + code_verifier="verifier", + refresh_token="leftover-refresh", + scope=None, + ) + + assert exc.value.status_code == 400 + assert exc.value.detail["error"] == "missing_client_id" + exchange_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_token_sealed_code_requires_code_verifier(self): + """A sealed code is minted only for S256 PKCE flows, so redeeming one without the + corresponding verifier is refused at the gateway rather than trusting the upstream to + enforce the binding.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + seal_passthrough_authorization_code, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_token, + ) + + request = MagicMock() + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server", + AsyncMock(), + ) as exchange_mock, + ): + sealed = seal_passthrough_authorization_code( + upstream_code="up-code", client_id="minted-77", client_secret=None, mcp_server_id="server-1" + ) + with pytest.raises(HTTPException) as exc: + await mcp_token( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + grant_type="authorization_code", + code=sealed, + redirect_uri="https://example.com/callback", + client_id=None, + client_secret=None, + code_verifier=None, + refresh_token=None, + scope=None, + ) + + assert exc.value.status_code == 400 + assert "code_verifier" in str(exc.value.detail) + exchange_mock.assert_not_awaited() + + @pytest.mark.asyncio + async def test_mcp_token_rejects_sealed_code_for_another_server(self): + """A sealed passthrough code is bound to the server it was minted for: presenting it at + another server's token endpoint is a 400 before any upstream exchange, so a code cannot be + replayed across a server boundary.""" + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + seal_passthrough_authorization_code, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_token, + ) + + request = MagicMock() + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.true_passthrough + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch("litellm.proxy.proxy_server.master_key", "sk-lit4581-test-master-key"), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server", + AsyncMock(), + ) as exchange_mock, + ): + sealed = seal_passthrough_authorization_code( + upstream_code="up-code", + client_id="minted-77", + client_secret=None, + mcp_server_id="a-different-server", + ) + with pytest.raises(HTTPException) as exc: + await mcp_token( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + grant_type="authorization_code", + code=sealed, + redirect_uri="https://example.com/callback", + client_id=None, + client_secret=None, + code_verifier="verifier", + refresh_token=None, + scope=None, + ) + + assert exc.value.status_code == 400 + exchange_mock.assert_not_awaited() + @pytest.mark.asyncio async def test_mcp_authorize_rejects_non_oauth2_server(self): """mcp_authorize must reject a none-auth server with an accurate 'does not use OAuth' @@ -2163,6 +2625,7 @@ class TestTemporaryMCPSessionEndpoints: code_verifier="verifier", refresh_token=None, scope=None, + client_token_endpoint_auth_method=None, ) @pytest.mark.asyncio @@ -2216,6 +2679,7 @@ class TestTemporaryMCPSessionEndpoints: code_verifier=None, refresh_token="rt-123", scope=None, + client_token_endpoint_auth_method=None, ) @pytest.mark.asyncio @@ -2270,8 +2734,59 @@ class TestTemporaryMCPSessionEndpoints: token_endpoint_auth_method="client_secret_basic", fallback_client_id="server-1", persist_credentials=True, + client_redirect_uris=None, ) + @pytest.mark.asyncio + @pytest.mark.parametrize( + "raw_redirect_uris, forwarded", + [ + (["https://app.example.com/ui/callback"], ["https://app.example.com/ui/callback"]), + (["https://app.example.com/ui/callback", 42, "", None], None), + ("not-a-list", None), + ([], None), + ([123], None), + ], + ) + async def test_mcp_register_forwards_validated_redirect_uris(self, raw_redirect_uris, forwarded): + """dcr_bridge servers relay the registration upstream and require the browser client's own + redirect_uris, so mcp_register must forward them; the value is caller-controlled and is + validated by the same client_supplied_redirect_uris boundary helper as the root /register + door, so a malformed list is rejected whole at both doors (RFC 7591 redirect_uris is + all-or-nothing) rather than silently forwarding the surviving entries here and rejecting + them there.""" + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + mcp_register, + ) + + request = MagicMock() + server = generate_mock_mcp_server_config_record(server_id="server-1") + server.auth_type = MCPAuth.oauth2 + request_body = {"client_name": "LiteLLM", "redirect_uris": raw_redirect_uris} + admin_auth = generate_mock_user_api_key_auth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404", + return_value=server, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body", + AsyncMock(return_value=request_body), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.register_client_with_server", + AsyncMock(return_value={"client_id": "generated"}), + ) as register_mock, + ): + await mcp_register( + request=request, + server_id="server-1", + user_api_key_dict=admin_auth, + ) + + assert register_mock.await_args.kwargs["client_redirect_uris"] == forwarded + @pytest.mark.asyncio async def test_mcp_register_does_not_persist_for_non_admin(self): """A non-admin caller (who may have access to a real server) must not persist the DCR diff --git a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py index f4470e7e83d..7ed123f6cdf 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_organization_endpoints.py @@ -10,9 +10,7 @@ import pytest from fastapi import HTTPException from fastapi.testclient import TestClient -sys.path.insert( - 0, os.path.abspath("../../../") -) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("../../../")) # Adds the parent directory to the system path @pytest.mark.asyncio @@ -58,16 +56,12 @@ async def test_organization_update_object_permissions_existing_permission(monkey "vector_stores": ["old_store_1", "old_store_2"], } - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=existing_object_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=existing_object_permission) # Mock upsert operation updated_permission = MagicMock() updated_permission.object_permission_id = "existing_perm_id_123" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=updated_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=updated_permission) # Test data with new object permission data_json = { @@ -107,9 +101,7 @@ async def test_get_organization_daily_activity_admin_param_passing(monkeypatch): # Mock prisma client mock_prisma_client = AsyncMock() - mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[] - ) + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Admin view -> skip membership restriction @@ -121,9 +113,7 @@ async def test_get_organization_daily_activity_admin_param_passing(monkeypatch): # Patch downstream common function and verify call args mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse") get_daily_activity_mock = AsyncMock(return_value=mocked_response) - monkeypatch.setattr( - organization_endpoints, "get_daily_activity", get_daily_activity_mock - ) + monkeypatch.setattr(organization_endpoints, "get_daily_activity", get_daily_activity_mock) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin1") result = await get_organization_daily_activity( @@ -172,17 +162,11 @@ async def test_get_organization_daily_activity_non_admin_defaults_to_admin_orgs( # Mock prisma client and memberships mock_prisma_client = AsyncMock() - mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[] - ) + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[]) mock_prisma_client.db.litellm_organizationmembership.find_many = AsyncMock( return_value=[ - SimpleNamespace( - organization_id="orgA", user_role=LitellmUserRoles.ORG_ADMIN.value - ), - SimpleNamespace( - organization_id="orgB", user_role=LitellmUserRoles.ORG_ADMIN.value - ), + SimpleNamespace(organization_id="orgA", user_role=LitellmUserRoles.ORG_ADMIN.value), + SimpleNamespace(organization_id="orgB", user_role=LitellmUserRoles.ORG_ADMIN.value), ] ) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) @@ -196,13 +180,9 @@ async def test_get_organization_daily_activity_non_admin_defaults_to_admin_orgs( # Patch downstream aggregator mocked_response = MagicMock(name="SpendAnalyticsPaginatedResponse") get_daily_activity_mock = AsyncMock(return_value=mocked_response) - monkeypatch.setattr( - organization_endpoints, "get_daily_activity", get_daily_activity_mock - ) + monkeypatch.setattr(organization_endpoints, "get_daily_activity", get_daily_activity_mock) - auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id="regular-user" - ) + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="regular-user") await get_organization_daily_activity( organization_ids=None, start_date="2024-02-01", @@ -238,15 +218,9 @@ async def test_get_organization_daily_activity_non_admin_unauthorized_org_raises # Mock prisma client and memberships (only orgA is admin) mock_prisma_client = AsyncMock() mock_prisma_client.db.litellm_organizationmembership.find_many = AsyncMock( - return_value=[ - SimpleNamespace( - organization_id="orgA", user_role=LitellmUserRoles.ORG_ADMIN.value - ) - ] - ) - mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[] + return_value=[SimpleNamespace(organization_id="orgA", user_role=LitellmUserRoles.ORG_ADMIN.value)] ) + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Non-admin view @@ -255,9 +229,7 @@ async def test_get_organization_daily_activity_non_admin_unauthorized_org_raises lambda _: False, ) - auth = UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id="regular-user" - ) + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="regular-user") with pytest.raises(HTTPException) as exc: await get_organization_daily_activity( @@ -312,21 +284,17 @@ async def test_organization_update_object_permissions_no_existing_permission( ) # Mock find_unique to return None (no existing permission) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) # Mock upsert to create new record new_permission = MagicMock() new_permission.object_permission_id = "new_perm_id_456" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=new_permission) data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["brand_new_store"] - ).model_dump(exclude_unset=True, exclude_none=True), + "object_permission": LiteLLM_ObjectPermissionBase(vector_stores=["brand_new_store"]).model_dump( + exclude_unset=True, exclude_none=True + ), "organization_alias": "updated_org_2", } @@ -381,21 +349,17 @@ async def test_organization_update_object_permissions_missing_permission_record( ) # Mock find_unique to return None (permission record not found) - mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None) # Mock upsert to create new record new_permission = MagicMock() new_permission.object_permission_id = "recreated_perm_id_789" - mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( - return_value=new_permission - ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(return_value=new_permission) data_json = { - "object_permission": LiteLLM_ObjectPermissionBase( - vector_stores=["recreated_store"] - ).model_dump(exclude_unset=True, exclude_none=True), + "object_permission": LiteLLM_ObjectPermissionBase(vector_stores=["recreated_store"]).model_dump( + exclude_unset=True, exclude_none=True + ), "organization_alias": "updated_org_3", } @@ -446,18 +410,14 @@ async def test_list_organization_filter_by_org_id(monkeypatch): ) # Mock find_many to return filtered results - mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[mock_org1] - ) + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[mock_org1]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Test as proxy admin auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") - result = await list_organization( - org_id="org-123", org_alias=None, user_api_key_dict=auth - ) + result = await list_organization(org_id="org-123", org_alias=None, user_api_key_dict=auth) # Verify the correct organization was returned assert len(result) == 1 @@ -512,18 +472,14 @@ async def test_list_organization_filter_by_org_alias(monkeypatch): ) # Mock find_many to return filtered results - mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock( - return_value=[mock_org1, mock_org2] - ) + mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[mock_org1, mock_org2]) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) # Test as proxy admin with org_alias filter auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-user") - result = await list_organization( - org_id=None, org_alias="test", user_api_key_dict=auth - ) + result = await list_organization(org_id=None, org_alias="test", user_api_key_dict=auth) # Verify organizations with "test" in alias were returned assert len(result) == 2 @@ -532,9 +488,7 @@ async def test_list_organization_filter_by_org_alias(monkeypatch): # Verify find_many was called with correct where conditions (case-insensitive contains) mock_prisma_client.db.litellm_organizationtable.find_many.assert_called_once() call_args = mock_prisma_client.db.litellm_organizationtable.find_many.call_args - assert call_args.kwargs["where"] == { - "organization_alias": {"contains": "test", "mode": "insensitive"} - } + assert call_args.kwargs["where"] == {"organization_alias": {"contains": "test", "mode": "insensitive"}} assert call_args.kwargs["include"] == { "litellm_budget_table": True, "members": True, @@ -612,16 +566,12 @@ def patched_org_prisma(): ), patch("litellm.proxy.proxy_server.proxy_logging_obj"), ): - mock_prisma.db.litellm_organizationtable.find_unique = AsyncMock( - return_value=victim_row - ) + mock_prisma.db.litellm_organizationtable.find_unique = AsyncMock(return_value=victim_row) yield mock_prisma @pytest.mark.asyncio -async def test_organization_member_add_rejects_unauthorized_caller( - patched_org_prisma, unauthorized_caller -): +async def test_organization_member_add_rejects_unauthorized_caller(patched_org_prisma, unauthorized_caller): # ``organization_member_add`` catches HTTPException in its # catch-all and re-wraps as ProxyException with the original status # code preserved. @@ -653,9 +603,7 @@ async def test_organization_member_add_rejects_unauthorized_caller( @pytest.mark.asyncio -async def test_organization_member_update_rejects_unauthorized_caller( - patched_org_prisma, unauthorized_caller -): +async def test_organization_member_update_rejects_unauthorized_caller(patched_org_prisma, unauthorized_caller): from litellm.proxy._types import OrganizationMemberUpdateRequest from litellm.proxy.management_endpoints.organization_endpoints import ( organization_member_update, @@ -676,9 +624,7 @@ async def test_organization_member_update_rejects_unauthorized_caller( @pytest.mark.asyncio -async def test_organization_member_delete_rejects_unauthorized_caller( - patched_org_prisma, unauthorized_caller -): +async def test_organization_member_delete_rejects_unauthorized_caller(patched_org_prisma, unauthorized_caller): from litellm.proxy._types import OrganizationMemberDeleteRequest from litellm.proxy.management_endpoints.organization_endpoints import ( organization_member_delete, @@ -695,3 +641,354 @@ async def test_organization_member_delete_rejects_unauthorized_caller( user_api_key_dict=unauthorized_caller, ) assert exc.value.status_code == 403 + + +@pytest.mark.parametrize( + "body", + [{"tpm_limit": ""}, {"tmp_limit": None}], + ids=["non-numeric-limit", "unknown-key"], +) +def test_v2_model_rejects_invalid_body(body): + """A non-numeric limit and an unknown/misspelled key are both rejected at model validation (422 at the route).""" + from pydantic import ValidationError + + from litellm.proxy._types import OrganizationUpdateRequestV2 + + with pytest.raises(ValidationError): + OrganizationUpdateRequestV2.model_validate(body) + + +class _FakeTxContext: + def __init__(self, tx): + self._tx = tx + + async def __aenter__(self): + return self._tx + + async def __aexit__(self, exc_type, exc, tb): + return False + + +async def _run_update_organization_v2( + monkeypatch, + *, + body: dict, + existing_budget_id, + existing_metadata, + existing_object_permission_id=None, + existing_object_permission_row=None, +): + from litellm.proxy._types import ( + LitellmUserRoles, + OrganizationUpdateRequestV2, + UserAPIKeyAuth, + ) + from litellm.proxy.management_endpoints import organization_endpoints + from litellm.proxy.management_endpoints.organization_endpoints import ( + update_organization_v2, + ) + from litellm.proxy.utils import jsonify_object + + mock_prisma_client = AsyncMock() + mock_prisma_client.jsonify_object = jsonify_object + + existing_org = MagicMock() + existing_org.budget_id = existing_budget_id + existing_org.object_permission_id = existing_object_permission_id + existing_org.metadata = existing_metadata + + mock_prisma_client.db.litellm_organizationtable.find_unique = AsyncMock(return_value=existing_org) + mock_prisma_client.db.litellm_organizationtable.update = AsyncMock(return_value=MagicMock()) + mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=existing_object_permission_row + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock() + + tx = MagicMock() + tx.litellm_organizationtable = mock_prisma_client.db.litellm_organizationtable + tx.litellm_budgettable = mock_prisma_client.db.litellm_budgettable + tx.litellm_objectpermissiontable.upsert = AsyncMock() + mock_prisma_client.db.tx = MagicMock(return_value=_FakeTxContext(tx)) + mock_prisma_client.tx = tx + + call_order = MagicMock() + call_order.attach_mock(tx.litellm_objectpermissiontable.upsert, "permission_upsert") + call_order.attach_mock(mock_prisma_client.db.litellm_organizationtable.update, "org_update") + mock_prisma_client.call_order = call_order + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") + await update_organization_v2( + organization_id="org-1", + data=OrganizationUpdateRequestV2.model_validate(body), + user_api_key_dict=auth, + ) + return mock_prisma_client + + +@pytest.mark.asyncio +async def test_v2_update_clears_tpm_limit_and_metadata(monkeypatch): + """A cleared tpm_limit is written to the budget row as None; a cleared metadata is written as {}.""" + prisma = await _run_update_organization_v2( + monkeypatch, + body={"tpm_limit": None, "metadata": None}, + existing_budget_id="budget-1", + existing_metadata={"stale": "value"}, + ) + + budget_write = prisma.db.litellm_budgettable.update.await_args + assert budget_write.kwargs["where"] == {"budget_id": "budget-1"} + assert budget_write.kwargs["data"]["tpm_limit"] is None + assert "soft_budget" not in budget_write.kwargs["data"] + + write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] + assert json.loads(write_data["metadata"]) == {} + assert "budget_id" not in write_data + + +@pytest.mark.asyncio +async def test_v2_update_untouched_fields_not_written(monkeypatch): + """Omitted fields are left untouched: only organization_alias is written, no budget-row write.""" + prisma = await _run_update_organization_v2( + monkeypatch, + body={"organization_alias": "renamed"}, + existing_budget_id="budget-1", + existing_metadata={"keep": "me"}, + ) + + prisma.db.litellm_budgettable.update.assert_not_awaited() + write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] + assert write_data["organization_alias"] == "renamed" + assert "metadata" not in write_data + assert "tpm_limit" not in write_data + + +@pytest.mark.asyncio +async def test_v2_update_metadata_replaces_not_merges(monkeypatch): + """Sending metadata replaces the stored blob wholesale; a previously-present key is gone.""" + prisma = await _run_update_organization_v2( + monkeypatch, + body={"metadata": {"a": 1}}, + existing_budget_id="budget-1", + existing_metadata={"stale": "value"}, + ) + write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] + assert json.loads(write_data["metadata"]) == {"a": 1} + + +@pytest.mark.asyncio +async def test_v2_rejects_null_clear_of_non_nullable_fields(monkeypatch): + """organization_alias and models are non-nullable columns, so a null clear is a 422, not a 500.""" + from litellm.proxy._types import LitellmUserRoles, OrganizationUpdateRequestV2, UserAPIKeyAuth + from litellm.proxy.management_endpoints.organization_endpoints import update_organization_v2 + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", AsyncMock()) + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") + + for body in ({"organization_alias": None}, {"models": None}): + with pytest.raises(HTTPException) as exc: + await update_organization_v2( + organization_id="org-1", + data=OrganizationUpdateRequestV2.model_validate(body), + user_api_key_dict=auth, + ) + assert exc.value.status_code == 422 + + +@pytest.mark.asyncio +async def test_v2_rejects_negative_max_budget(monkeypatch): + """v2 rejects a negative max_budget with a 422 before touching the DB.""" + from litellm.proxy._types import LitellmUserRoles, OrganizationUpdateRequestV2, UserAPIKeyAuth + from litellm.proxy.management_endpoints.organization_endpoints import update_organization_v2 + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", AsyncMock()) + + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") + with pytest.raises(HTTPException) as exc: + await update_organization_v2( + organization_id="org-1", + data=OrganizationUpdateRequestV2.model_validate({"max_budget": -5}), + user_api_key_dict=auth, + ) + assert exc.value.status_code == 422 + assert "max_budget" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_v2_rejects_caller_without_org_access(monkeypatch): + """v2 runs the real _verify_org_access guard: a non-admin without ORG_ADMIN on the org gets 403 and no write.""" + from litellm.proxy._types import LitellmUserRoles, OrganizationUpdateRequestV2, UserAPIKeyAuth + from litellm.proxy.management_endpoints import organization_endpoints + from litellm.proxy.management_endpoints.organization_endpoints import update_organization_v2 + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr(organization_endpoints, "_user_has_admin_view", lambda _: False) + + caller = MagicMock() + caller.organization_memberships = [] + monkeypatch.setattr(organization_endpoints, "get_user_object", AsyncMock(return_value=caller)) + + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1") + with pytest.raises(HTTPException) as exc: + await update_organization_v2( + organization_id="org-1", + data=OrganizationUpdateRequestV2.model_validate({"tpm_limit": 5}), + user_api_key_dict=auth, + ) + assert exc.value.status_code == 403 + mock_prisma_client.db.litellm_organizationtable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_v2_wires_object_permission_onto_org_write(monkeypatch): + """A sent object_permission merges over the existing permission row and its id is linked onto the org write.""" + existing_row = MagicMock() + existing_row.model_dump.return_value = { + "object_permission_id": "op-123", + "mcp_servers": ["server-1"], + } + + prisma = await _run_update_organization_v2( + monkeypatch, + body={"object_permission": {"vector_stores": ["vs-1"]}}, + existing_budget_id="budget-1", + existing_metadata={}, + existing_object_permission_id="op-123", + existing_object_permission_row=existing_row, + ) + + upsert = prisma.tx.litellm_objectpermissiontable.upsert.await_args.kwargs + assert upsert["where"] == {"object_permission_id": "op-123"} + assert upsert["data"]["update"]["mcp_servers"] == ["server-1"] + assert upsert["data"]["update"]["vector_stores"] == ["vs-1"] + write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] + assert write_data["object_permission_id"] == "op-123" + + +@pytest.mark.asyncio +async def test_v2_object_permission_upsert_runs_inside_transaction(monkeypatch): + """The permission upsert runs on the tx client, before the org write that links it, so a rollback cannot + leave merged grants live on a row the org still points at.""" + prisma = await _run_update_organization_v2( + monkeypatch, + body={"object_permission": {"vector_stores": ["vs-1"]}}, + existing_budget_id="budget-1", + existing_metadata={}, + ) + + prisma.tx.litellm_objectpermissiontable.upsert.assert_awaited_once() + prisma.db.litellm_objectpermissiontable.upsert.assert_not_awaited() + + upsert = prisma.tx.litellm_objectpermissiontable.upsert.await_args.kwargs + linked_id = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"]["object_permission_id"] + assert upsert["where"] == {"object_permission_id": linked_id} + assert upsert["data"]["create"]["object_permission_id"] == linked_id + + ordered = [name for name, _, _ in prisma.call_order.mock_calls if name in ("permission_upsert", "org_update")] + assert ordered == ["permission_upsert", "org_update"] + + +@pytest.mark.asyncio +async def test_v2_clears_object_permission_when_sent_null(monkeypatch): + """object_permission: null detaches the org's permission row (object_permission_id -> None), no merge.""" + prisma = await _run_update_organization_v2( + monkeypatch, + body={"object_permission": None}, + existing_budget_id="budget-1", + existing_metadata={}, + ) + + prisma.tx.litellm_objectpermissiontable.upsert.assert_not_awaited() + prisma.db.litellm_objectpermissiontable.find_unique.assert_not_awaited() + write_data = prisma.db.litellm_organizationtable.update.await_args.kwargs["data"] + assert write_data["object_permission_id"] is None + + +@pytest.mark.asyncio +async def test_v2_rejects_empty_object_permission(monkeypatch): + """object_permission: {} merges nothing, so it is rejected (send null to clear) rather than silently leaving grants.""" + from litellm.proxy._types import LitellmUserRoles, OrganizationUpdateRequestV2, UserAPIKeyAuth + from litellm.proxy.management_endpoints import organization_endpoints + from litellm.proxy.management_endpoints.organization_endpoints import update_organization_v2 + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr(organization_endpoints, "_verify_org_access", AsyncMock()) + + auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1") + with pytest.raises(HTTPException) as exc: + await update_organization_v2( + organization_id="org-1", + data=OrganizationUpdateRequestV2.model_validate({"object_permission": {}}), + user_api_key_dict=auth, + ) + assert exc.value.status_code == 422 + assert "object_permission" in str(exc.value.detail) + mock_prisma_client.db.litellm_organizationtable.update.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_v2_writes_budget_and_org_in_one_transaction(monkeypatch): + """A change touching both the budget row and the org row runs both writes inside one prisma transaction.""" + prisma = await _run_update_organization_v2( + monkeypatch, + body={"tpm_limit": 500, "metadata": {"a": 1}}, + existing_budget_id="budget-1", + existing_metadata={}, + ) + + prisma.db.tx.assert_called_once() + prisma.db.litellm_budgettable.update.assert_awaited_once() + prisma.db.litellm_organizationtable.update.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_v2_serializes_model_max_budget_on_budget_write(monkeypatch): + """model_max_budget is a Json column, so it is JSON-serialized on the budget-row write like new_budget/metadata.""" + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_model_max_budget", + lambda _: None, + ) + + prisma = await _run_update_organization_v2( + monkeypatch, + body={"model_max_budget": {"gpt-4o": {"max_budget": 10}}}, + existing_budget_id="budget-1", + existing_metadata={}, + ) + + written = prisma.db.litellm_budgettable.update.await_args.kwargs["data"]["model_max_budget"] + assert isinstance(written, str) + assert json.loads(written) == {"gpt-4o": {"max_budget": 10}} + + +def test_build_budget_write_data_recomputes_reset_at_on_duration(): + """A sent budget_duration recomputes budget_reset_at so the reset window follows the new duration.""" + from litellm.proxy.management_endpoints.organization_endpoints import build_budget_write_data + + data = build_budget_write_data({"budget_duration": "30d"}, "admin-1") + assert data["budget_duration"] == "30d" + assert "budget_reset_at" in data + assert data["updated_by"] == "admin-1" + + +def test_build_budget_write_data_no_reset_at_without_duration(): + """Clearing a limit writes it through untouched and does not recompute budget_reset_at.""" + from litellm.proxy.management_endpoints.organization_endpoints import build_budget_write_data + + data = build_budget_write_data({"tpm_limit": None}, "admin-1") + assert data["tpm_limit"] is None + assert "budget_reset_at" not in data + + +def test_build_budget_write_data_clears_reset_at_with_null_duration(): + """Clearing budget_duration also nulls budget_reset_at so no stale reset timestamp survives.""" + from litellm.proxy.management_endpoints.organization_endpoints import build_budget_write_data + + data = build_budget_write_data({"budget_duration": None}, "admin-1") + assert data["budget_duration"] is None + assert data["budget_reset_at"] is None 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 4936191c344..5202c8cbfc0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -1693,6 +1693,78 @@ async def test_update_team_members_list_duplicate_prevention(): assert len(mock_team.members_with_roles) == 1 +@pytest.mark.asyncio +async def test_add_team_members_reconciles_against_freshly_locked_row(): + """ + Regression: _add_team_members_to_team must build the new members_with_roles + from the row it re-reads under a lock inside the write transaction, not from + the stale complete_team_data snapshot captured at the start of the request. + + Two concurrent /team/member_add calls for the same team read the same + snapshot; without the locked re-read the losing write rewrites the whole + JSON array from its stale copy and silently drops the member the other call + already committed. Here the snapshot holds only "zed", a concurrent writer + has already committed "alice" (returned by the locked SELECT), and this call + adds "bob". The write must contain all three. + """ + from litellm.proxy.management_endpoints.team_endpoints import ( + _add_team_members_to_team, + ) + + stale_snapshot = LiteLLM_TeamTable( + team_id="test-team-lock", + members_with_roles=[Member(user_id="zed", role="user")], + ) + + freshly_committed = [ + {"user_id": "zed", "user_email": None, "role": "user"}, + {"user_id": "alice", "user_email": None, "role": "user"}, + ] + + captured: dict = {} + + async def _capture_update(where, data): + captured["data"] = data + return LiteLLM_TeamTable( + team_id="test-team-lock", + members_with_roles=json.loads(data["members_with_roles"]), + ) + + tx = MagicMock() + tx.query_raw = AsyncMock(return_value=[{"members_with_roles": freshly_committed}]) + tx.litellm_teamtable.update = AsyncMock(side_effect=_capture_update) + + tx_cm = MagicMock() + tx_cm.__aenter__ = AsyncMock(return_value=tx) + tx_cm.__aexit__ = AsyncMock(return_value=None) + + prisma_client = MagicMock() + prisma_client.tx = MagicMock(return_value=tx_cm) + + with patch( + "litellm.proxy.management_endpoints.team_endpoints._process_team_members", + new=AsyncMock(return_value=([], [])), + ): + updated_team, _, _ = await _add_team_members_to_team( + data=TeamMemberAddRequest( + team_id="test-team-lock", + member=Member(user_id="bob", role="user"), + ), + complete_team_data=stale_snapshot, + prisma_client=cast(object, prisma_client), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + litellm_proxy_admin_name="admin", + ) + + written_ids = sorted(m["user_id"] for m in json.loads(captured["data"]["members_with_roles"])) + assert written_ids == ["alice", "bob", "zed"] + + lock_reads = [call for call in tx.query_raw.call_args_list if "FOR UPDATE" in str(call.args[0])] + assert lock_reads, "expected a SELECT ... FOR UPDATE row-lock read before the write" + + assert [m.user_id for m in updated_team.members_with_roles] == ["zed", "alice", "bob"] + + def test_add_new_models_to_team_with_existing_models(): """ Test add_new_models_to_team function with existing models @@ -4106,6 +4178,8 @@ async def test_new_team_max_budget_within_user_limit(): } mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user) + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_prisma.db.litellm_usertable.update = AsyncMock(return_value=mock_user) # Mock team membership table @@ -4247,6 +4321,8 @@ async def test_new_team_org_scoped_budget_bypasses_user_limit(): } mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user) + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_prisma.db.litellm_usertable.update = AsyncMock(return_value=mock_user) # Mock team membership table @@ -4393,6 +4469,8 @@ async def test_new_team_org_scoped_models_bypasses_user_limit(): } mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user) + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_prisma.db.litellm_usertable.update = AsyncMock(return_value=mock_user) # Mock team membership table @@ -7245,6 +7323,8 @@ async def test_new_team_soft_budget_validation( } mock_prisma.db.litellm_usertable = MagicMock() mock_prisma.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user) + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user) mock_prisma.db.litellm_usertable.update = AsyncMock(return_value=mock_user) # Mock team membership table @@ -9733,7 +9813,6 @@ async def _drive_team_write( raw_body=None, user=None, find_returns_none=False, - json_side_effect=None, ): """Drive POST ``update_team`` or PATCH ``patch_team`` against a mocked team. @@ -9748,6 +9827,7 @@ async def _drive_team_write( from litellm.proxy._types import ( LiteLLM_TeamTable, LitellmUserRoles, + PatchTeamRequest, UpdateTeamRequest, UserAPIKeyAuth, ) @@ -9793,14 +9873,10 @@ async def _drive_team_write( litellm_changed_by=None, ) else: - if json_side_effect is not None: - req.json = AsyncMock(side_effect=json_side_effect) - else: - req.json = AsyncMock( - return_value=raw_body if raw_body is not None else dict(payload or {}) - ) + body = raw_body if raw_body is not None else dict(payload or {}) result = await patch_team( team_id=_PATCH_TEAM_ID, + data=PatchTeamRequest.model_validate(body), http_request=req, user_api_key_dict=auth, litellm_changed_by=None, @@ -9948,25 +10024,36 @@ async def test_patch_strips_system_managed_metadata_key_like_post(): assert patch_meta == {"cost_center": "9999"} -@pytest.mark.asyncio -@pytest.mark.parametrize("raw_body", [["not", "an", "object"], "a-string", 42, True]) -async def test_patch_rejects_non_object_body(raw_body): - from litellm.proxy._types import ProxyException +@pytest.mark.parametrize( + "kwargs", + [ + {"json": ["not", "an", "object"]}, + {"json": "a-string"}, + {"json": 42}, + {"content": b"{not json"}, + {"json": {"tpm_limit": "not-an-int"}}, + ], + ids=["list", "string", "number", "malformed-json", "wrong-field-type"], +) +def test_patch_rejects_a_malformed_body_with_422(kwargs): + """The body is a declared parameter, so FastAPI rejects a malformed one before the + handler runs. This is the same 422 POST /team/update already returns; the route + previously answered 400 here and 500 for a wrongly typed field, reporting a caller + mistake as a server fault.""" + from fastapi import FastAPI + from fastapi.testclient import TestClient - with pytest.raises(ProxyException) as exc: - await _drive_team_write("patch", existing_metadata={"a": 1}, raw_body=raw_body) - assert exc.value.code == "400" or exc.value.code == 400 + from litellm.proxy._types import PatchTeamRequest + app = FastAPI() -@pytest.mark.asyncio -async def test_patch_rejects_invalid_json_body(): - from litellm.proxy._types import ProxyException + @app.patch("/team/{team_id}") + async def _route(team_id: str, data: PatchTeamRequest): # pragma: no cover - schema only + return {} - with pytest.raises(ProxyException) as exc: - await _drive_team_write( - "patch", existing_metadata={"a": 1}, json_side_effect=ValueError("no body") - ) - assert exc.value.code == "400" or exc.value.code == 400 + response = TestClient(app).patch("/team/abc", **kwargs) + + assert response.status_code == 422 @pytest.mark.asyncio @@ -10036,3 +10123,103 @@ async def test_patch_returns_full_team_object_not_wrapper(): ) assert isinstance(result, LiteLLM_TeamTable) assert result.team_id == _PATCH_TEAM_ID + + +# --------------------------------------------------------------------------- +# PATCH body is validated through PatchTeamRequest before it is handed to +# update_team. The write below must stay byte-identical to what the untyped +# **body construction produced, or a partial update starts writing columns the +# caller never mentioned. +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_patch_writes_only_the_keys_the_caller_sent(): + """An omitted field must not reach the DB write at all. If validation ever + materialises defaults, every unmentioned column gets overwritten with null.""" + _, update_mock = await _drive_team_write("patch", raw_body={"tpm_limit": 5}) + written = update_mock.call_args.kwargs["data"] + + assert written["tpm_limit"] == 5 + for untouched in ("rpm_limit", "max_budget", "models", "blocked", "budget_duration"): + assert untouched not in written, f"{untouched} was written despite not being sent" + + +@pytest.mark.asyncio +async def test_patch_preserves_explicit_null_as_a_clear(): + """null is a clear, not an omission: it has to survive validation and reach the write.""" + _, update_mock = await _drive_team_write("patch", raw_body={"max_budget": None}) + written = update_mock.call_args.kwargs["data"] + + assert "max_budget" in written + assert written["max_budget"] is None + + +def _patch_body_to_update_request(body: dict): + """The exact reshaping patch_team performs between the raw body and update_team.""" + from litellm.proxy._types import PatchTeamRequest, UpdateTeamRequest + + parsed = PatchTeamRequest.model_validate(body) + return UpdateTeamRequest( + team_id=_PATCH_TEAM_ID, + **parsed.model_dump(exclude_unset=True, exclude={"team_id"}), + ) + + +@pytest.mark.parametrize( + "body", + [ + {"tpm_limit": 5}, + {"max_budget": None}, + {"object_permission": {"vector_stores": []}}, + {"metadata": {"a": 1, "b": None}}, + {"models": ["gpt-4"], "blocked": False}, + ], + ids=["scalar", "explicit-null", "partial-nested", "metadata-with-null", "list-and-false"], +) +def test_patch_body_reshaping_adds_no_keys_the_caller_did_not_send(body): + """Validating through PatchTeamRequest must be shape-preserving. If it ever + materialises defaults, a partial update silently overwrites untouched columns, + and for the merge-only object_permission it would wipe sibling sub-keys.""" + reshaped = _patch_body_to_update_request(body) + dumped = reshaped.model_dump(exclude_unset=True, exclude={"team_id"}) + + assert dumped == body + assert reshaped.model_fields_set == set(body) | {"team_id"} + + +@pytest.mark.asyncio +async def test_patch_ignores_unknown_body_keys(): + """Unknown keys were silently dropped by the previous construction; keep that.""" + _, update_mock = await _drive_team_write( + "patch", raw_body={"tpm_limit": 5, "not_a_team_field": "x"} + ) + written = update_mock.call_args.kwargs["data"] + + assert written["tpm_limit"] == 5 + assert "not_a_team_field" not in written + + +def test_patch_team_request_makes_team_id_optional(): + """PATCH takes team_id from the path, so the body model must not require it, + while still inheriting every UpdateTeamRequest field.""" + from litellm.proxy._types import PatchTeamRequest, UpdateTeamRequest + + parsed = PatchTeamRequest.model_validate({"tpm_limit": 5}) + + assert parsed.team_id is None + assert parsed.model_fields_set == {"tpm_limit"} + assert set(UpdateTeamRequest.model_fields).issubset(set(PatchTeamRequest.model_fields)) + + +def test_patch_team_route_publishes_its_request_body_schema(): + """The dashboard's generated client types this call off the OpenAPI spec, which + FastAPI can only emit because the body is a declared parameter.""" + from litellm.proxy.proxy_server import app + + operation = app.openapi()["paths"]["/team/{team_id}"]["patch"] + schema = operation["requestBody"]["content"]["application/json"]["schema"] + + assert schema == {"$ref": "#/components/schemas/PatchTeamRequest"} + properties = app.openapi()["components"]["schemas"]["PatchTeamRequest"]["properties"] + assert "tpm_limit" in properties and "metadata" in properties diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 5631aa69102..c693017e134 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1458,7 +1458,7 @@ async def test_get_generic_sso_response_with_additional_headers(): "fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class ): # Act - result, received_response, _ = await get_generic_sso_response( + result, received_response, _, _ = await get_generic_sso_response( request=mock_request, jwt_handler=mock_jwt_handler, generic_client_id=generic_client_id, @@ -1522,7 +1522,7 @@ async def test_get_generic_sso_response_with_empty_headers(): "fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class ): # Act - result, received_response, _ = await get_generic_sso_response( + result, received_response, _, _ = await get_generic_sso_response( request=mock_request, jwt_handler=mock_jwt_handler, generic_client_id=generic_client_id, @@ -2214,7 +2214,95 @@ class TestCLIKeyRegenerationFlow: _get_cli_sso_flow_or_raise(login_id="cli-test_1234567890", cache=mock_cache) assert expired_exc.value.status_code == 400 assert "session not found or expired" in expired_exc.value.detail - assert "enable_redis_auth_cache" in expired_exc.value.detail + assert "configure a Redis cache" in expired_exc.value.detail + assert "enable_redis_auth_cache" not in expired_exc.value.detail + + def test_cli_sso_flow_is_redis_authoritative_when_redis_attached(self): + """ + When Redis is attached, the CLI SSO flow must be read from and written to + Redis directly, never the in-memory layer. Otherwise the worker that served + /sso/cli/start keeps serving its stale in-memory flow and never sees the + sso_complete/session_data update another worker wrote, which is exactly the + multi-worker failure this fix targets. + """ + from litellm.proxy.management_endpoints.ui_sso import ( + CLI_SSO_SESSION_TTL_SECONDS, + _get_cli_sso_flow_cache_key, + _get_cli_sso_flow_or_raise, + _set_cli_sso_flow, + ) + + login_id = "cli-redis_authoritative_1234567890" + cache_key = _get_cli_sso_flow_cache_key(login_id) + fresh_flow = {"poll_secret_hash": "fresh", "sso_complete": True} + stale_flow = {"poll_secret_hash": "stale", "sso_complete": False} + + redis_cache = MagicMock() + redis_cache.get_cache.return_value = fresh_flow + cache = MagicMock() + cache.redis_cache = redis_cache + cache.get_cache.return_value = stale_flow + + result = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cache) + + assert result == fresh_flow + redis_cache.get_cache.assert_called_once_with(key=cache_key) + cache.get_cache.assert_not_called() + + _set_cli_sso_flow(login_id=login_id, cache=cache, flow=fresh_flow) + + redis_cache.set_cache.assert_called_once_with( + key=cache_key, value=json.dumps(fresh_flow), ttl=CLI_SSO_SESSION_TTL_SECONDS + ) + cache.set_cache.assert_not_called() + + def test_cli_sso_flow_with_enum_survives_redis_round_trip(self): + """ + RedisCache stores values via str(value) and reads them back through + json.loads/ast.literal_eval. A raw flow dict containing a Python enum + (session_data.user_role after the SSO callback) produces an unparseable + repr, so every worker reading the completed flow from Redis got a + SyntaxError and returned 400 "session not found". The flow must survive + a real Redis serialization round trip. + """ + from litellm.caching.redis_cache import RedisCache + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.management_endpoints.ui_sso import ( + _get_cli_sso_flow_or_raise, + _set_cli_sso_flow, + ) + + login_id = "cli-enum_round_trip_1234567890" + completed_flow = { + "poll_secret_hash": "hash", + "sso_complete": True, + "user_code_verified": False, + "session_data": { + "user_id": "user-1", + "user_role": LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, + "models": [], + "teams": ["team-1"], + "team_details": [{"team_id": "team-1", "team_alias": "alias"}], + }, + } + + redis_store: dict = {} + redis_cache = MagicMock() + redis_cache.set_cache.side_effect = lambda key, value, ttl: redis_store.__setitem__( + key, str(value).encode("utf-8") + ) + redis_cache.get_cache.side_effect = lambda key: RedisCache._get_cache_logic( + MagicMock(), redis_store.get(key) + ) + cache = MagicMock() + cache.redis_cache = redis_cache + + _set_cli_sso_flow(login_id=login_id, cache=cache, flow=completed_flow) + flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cache) + + assert flow["sso_complete"] is True + assert flow["session_data"]["user_role"] == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value + assert flow["session_data"]["team_details"] == [{"team_id": "team-1", "team_alias": "alias"}] @pytest.mark.asyncio async def test_cli_sso_start_creates_bound_flow(self): @@ -2228,10 +2316,13 @@ class TestCLIKeyRegenerationFlow: mock_request = MagicMock(spec=Request) mock_request.client = SimpleNamespace(host="127.0.0.1") mock_request.headers = {} - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.increment_cache.return_value = 1 - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): result = await cli_sso_start(request=mock_request) assert result["login_id"].startswith("cli-") @@ -2259,10 +2350,13 @@ class TestCLIKeyRegenerationFlow: mock_request = MagicMock(spec=Request) mock_request.client = SimpleNamespace(host="127.0.0.1") mock_request.headers = {} - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.increment_cache.return_value = 31 - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): with pytest.raises(HTTPException) as exc_info: await cli_sso_start(request=mock_request) @@ -2281,7 +2375,7 @@ class TestCLIKeyRegenerationFlow: mock_request.client = SimpleNamespace(host="127.0.0.1") mock_request.headers = {} mock_request.base_url = "https://proxy.example.com/" - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.increment_cache.return_value = 1 with ( @@ -2315,7 +2409,7 @@ class TestCLIKeyRegenerationFlow: mock_request.client = SimpleNamespace(host="127.0.0.1") mock_request.headers = {} mock_request.base_url = "https://proxy.example.com/" - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.increment_cache.return_value = 1 with ( @@ -2349,7 +2443,7 @@ class TestCLIKeyRegenerationFlow: mock_request = MagicMock(spec=Request) mock_request.base_url = "https://proxy.example.com/" - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = {"poll_secret_hash": "h"} async def drive(enabled: bool): @@ -2358,6 +2452,7 @@ class TestCLIKeyRegenerationFlow: patch("litellm.proxy.proxy_server.premium_user", True), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch( "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", None, @@ -2525,7 +2620,7 @@ class TestCLIKeyRegenerationFlow: ) mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"} - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": "poll-secret-hash", "user_code_hash": "user-code-hash", @@ -2544,6 +2639,7 @@ class TestCLIKeyRegenerationFlow: ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), ): result = await cli_sso_callback( request=mock_request, @@ -2568,7 +2664,7 @@ class TestCLIKeyRegenerationFlow: mock_request.body = AsyncMock( return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token" ) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "user_code_hash": _hash_cli_sso_secret( @@ -2582,6 +2678,7 @@ class TestCLIKeyRegenerationFlow: with ( patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch( "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success", @@ -2606,7 +2703,7 @@ class TestCLIKeyRegenerationFlow: mock_request = MagicMock(spec=Request) mock_request.body = AsyncMock(return_value=b"user_code=ABCD-EFGH") - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "user_code_hash": _hash_cli_sso_secret( @@ -2618,7 +2715,10 @@ class TestCLIKeyRegenerationFlow: "session_data": {"user_id": "test-user-123"}, } - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): with pytest.raises(HTTPException) as exc_info: await cli_sso_complete( request=mock_request, login_id="cli-session-4567890" @@ -2640,7 +2740,7 @@ class TestCLIKeyRegenerationFlow: mock_request.body = AsyncMock( return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token" ) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "user_code_hash": _hash_cli_sso_secret( @@ -2651,7 +2751,10 @@ class TestCLIKeyRegenerationFlow: "session_data": None, } - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): with pytest.raises(HTTPException) as exc_info: await cli_sso_complete( request=mock_request, login_id="cli-session-4567890" @@ -2687,7 +2790,7 @@ class TestCLIKeyRegenerationFlow: mock_sso_result = {"user_email": "test@example.com", "user_id": "test-user-123"} # Mock cache - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": "poll-secret-hash", "user_code_hash": "user-code-hash", @@ -2709,6 +2812,7 @@ class TestCLIKeyRegenerationFlow: ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch( "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", return_value="Success", @@ -2769,7 +2873,7 @@ class TestCLIKeyRegenerationFlow: } # Mock cache - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -2777,7 +2881,10 @@ class TestCLIKeyRegenerationFlow: "session_data": session_data, } - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): # Act - First poll without team_id result = await cli_poll_key( key_id=session_key, @@ -2803,7 +2910,7 @@ class TestCLIKeyRegenerationFlow: cli_poll_key, ) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -2816,7 +2923,10 @@ class TestCLIKeyRegenerationFlow: }, } - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): with pytest.raises(HTTPException) as exc_info: await cli_poll_key(key_id="cli-session-789123", team_id=None) @@ -2830,7 +2940,7 @@ class TestCLIKeyRegenerationFlow: cli_poll_key, ) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -2843,7 +2953,10 @@ class TestCLIKeyRegenerationFlow: }, } - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): result = await cli_poll_key( key_id="cli-session-789123", team_id=None, @@ -2893,6 +3006,7 @@ class TestCLIKeyRegenerationFlow: prefill_user_code=None, result=mock_result, received_response=None, + sso_assertion=None, ) @pytest.mark.asyncio @@ -2933,6 +3047,7 @@ class TestCLIKeyRegenerationFlow: prefill_user_code="WXYZ-2345", result=mock_result, received_response=None, + sso_assertion=None, ) def test_get_redirect_url_does_not_include_existing_key_in_url(self): @@ -3009,7 +3124,7 @@ class TestCLIKeyRegenerationFlow: ) # Mock cache - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -3021,6 +3136,7 @@ class TestCLIKeyRegenerationFlow: with ( patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch("litellm.proxy.proxy_server.prisma_client"), patch( "litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token", @@ -3084,7 +3200,7 @@ class TestCLIKeyRegenerationFlow: models=["gpt-4"], max_budget=100.0, ) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -3095,6 +3211,7 @@ class TestCLIKeyRegenerationFlow: with ( patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch("litellm.proxy.proxy_server.prisma_client"), patch( "litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token", @@ -3140,7 +3257,7 @@ class TestCLIKeyRegenerationFlow: "models": ["gpt-4"], "user_email": "unbudgeted@example.com", } - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -3151,6 +3268,7 @@ class TestCLIKeyRegenerationFlow: with ( patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch( "litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token", return_value=mock_jwt_token, @@ -4080,7 +4198,7 @@ class TestPKCEFunctionality: mock_request.query_params = {"state": test_state} # Mock cache with async methods — use dict format (primary path) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) test_code_verifier = "test_code_verifier_abc123xyz" mock_cache.async_get_cache = AsyncMock( return_value={"code_verifier": test_code_verifier} @@ -4131,7 +4249,7 @@ class TestPKCEFunctionality: mock_sso.__exit__ = MagicMock(return_value=False) test_state = "test456" - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.async_set_cache = AsyncMock() @@ -4655,7 +4773,7 @@ class TestPKCEFunctionality: from litellm.proxy._types import ProxyException from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found mock_request = MagicMock(spec=Request) @@ -4781,7 +4899,7 @@ class TestPKCEFunctionality: from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler # Cache returns an integer — unexpected format - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.async_get_cache = AsyncMock(return_value=12345) mock_cache.async_delete_cache = AsyncMock() @@ -4823,7 +4941,7 @@ class TestPKCEFunctionality: from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.async_get_cache = AsyncMock(return_value=None) # verifier not found mock_request = MagicMock(spec=Request) @@ -4911,7 +5029,7 @@ class TestPKCEFunctionality: from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler # Cache returns an integer — unexpected format - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.async_get_cache = AsyncMock(return_value=12345) mock_cache.async_delete_cache = AsyncMock() @@ -4963,7 +5081,7 @@ class TestPKCEFunctionality: from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler legacy_verifier = "legacy_plain_string_verifier_abc123" - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.async_get_cache = AsyncMock(return_value=legacy_verifier) mock_request = MagicMock(spec=Request) @@ -6247,7 +6365,7 @@ class TestCliSsoAttributionMetadata: provider="generic", team_ids=[], ) - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": "poll-secret-hash", "user_code_hash": "user-code-hash", @@ -6264,6 +6382,7 @@ class TestCliSsoAttributionMetadata: ), patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch("litellm.proxy.proxy_server.user_custom_sso", None), ): await ui_sso.cli_sso_callback( @@ -6288,7 +6407,7 @@ class TestCliSsoAttributionMetadata: mock_request = MagicMock(spec=Request) mock_request.base_url = "http://internal-proxy.local/" - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": "poll-secret-hash", "user_code_hash": "user-code-hash", @@ -6311,6 +6430,7 @@ class TestCliSsoAttributionMetadata: ) as get_user_info_mock, patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch("litellm.proxy.proxy_server.user_custom_sso", None), patch( "litellm.proxy.proxy_server.general_settings", @@ -6357,7 +6477,7 @@ class TestCliSsoAttributionMetadata: "user_id": "test-user-123", "employment_type": "contractor", } - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": "poll-secret-hash", "user_code_hash": "user-code-hash", @@ -6385,6 +6505,7 @@ class TestCliSsoAttributionMetadata: ), patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch("litellm.proxy.proxy_server.user_custom_sso", None), patch( "litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page", @@ -6426,7 +6547,7 @@ class TestCliSsoAttributionMetadata: "org": {"cost_center": "CC-42"}, }, } - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -6434,7 +6555,10 @@ class TestCliSsoAttributionMetadata: "session_data": session_data, } - with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache): + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), + ): result = await cli_poll_key( key_id=session_key, team_id=None, @@ -7019,7 +7143,7 @@ class TestPKCEStateCookieBinding: ): jwt_handler = MagicMock(spec=JWTHandler) jwt_handler.get_team_ids_from_jwt.return_value = [] - result, _, _ = await get_generic_sso_response( + result, _, _, _ = await get_generic_sso_response( request=mock_request, jwt_handler=jwt_handler, generic_client_id="cid", @@ -7078,7 +7202,7 @@ async def test_debug_sso_callback_renders_full_jwt_claims(): } async def fake_get_generic_sso_response(**kwargs): - return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload + return parsed_openid, raw_userinfo_with_leaked_token, access_token_payload, None with ( patch.dict( @@ -7285,7 +7409,7 @@ async def test_cli_poll_key_tolerates_missing_user_row(): "models": ["gpt-4"], } - mock_cache = MagicMock() + mock_cache = MagicMock(redis_cache=None) mock_cache.get_cache.return_value = { "poll_secret_hash": _hash_cli_sso_secret("poll-secret"), "sso_complete": True, @@ -7297,6 +7421,7 @@ async def test_cli_poll_key_tolerates_missing_user_row(): with ( patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.cli_sso_session_cache", mock_cache), patch("litellm.proxy.proxy_server.prisma_client"), patch( "litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token", @@ -7374,3 +7499,267 @@ async def test_auth_callback_without_oauth_error_proceeds_to_normal_flow(): assert exc_info.value.status_code == 500 assert "DB not connected" in str(exc_info.value.detail) + + +# ── SSO identity assertion capture + persist wiring (EMA) ───────────────────── + + +def _ema_id_token(sub: str = "u1") -> str: + import time as _time + + import jwt as _pyjwt + + return _pyjwt.encode( + {"iss": "https://idp.example.com", "sub": sub, "exp": int(_time.time()) + 3600}, + "test-idp-signing-key-32-bytes-long-xxxx", + algorithm="HS256", + ) + + +@pytest.mark.asyncio +async def test_pkce_arm_captures_sso_assertion(): + """The PKCE token exchange strips bearer fields from received_response for safety; + the typed assertion carrier must still capture id_token + refresh_token.""" + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + get_generic_sso_response, + ) + + id_token = _ema_id_token() + mock_request = MagicMock(spec=Request) + mock_request.query_params = {"state": "matched-state", "code": "auth-code"} + mock_request.cookies = {"litellm_oauth_state": "matched-state"} + + with ( + patch.object( + SSOAuthenticationHandler, + "prepare_token_exchange_parameters", + AsyncMock( + return_value={ + "code_verifier": "verifier", + "_pkce_cache_key": "pkce_verifier:matched-state", + } + ), + ), + patch.object( + SSOAuthenticationHandler, + "_pkce_token_exchange", + AsyncMock( + return_value={ + "access_token": "tok", + "id_token": id_token, + "refresh_token": "rt_from_idp", + "sub": "user@example.com", + "email": "user@example.com", + } + ), + ), + patch.object(SSOAuthenticationHandler, "_delete_pkce_verifier", AsyncMock()), + patch("fastapi_sso.sso.base.DiscoveryDocument"), + patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()), + patch.dict( + os.environ, + { + "GENERIC_CLIENT_SECRET": "x", + "GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth", + "GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token", + "GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo", + "GENERIC_CLIENT_USE_PKCE": "true", + }, + ), + ): + jwt_handler = MagicMock(spec=JWTHandler) + jwt_handler.get_team_ids_from_jwt.return_value = [] + result, received_response, _, sso_assertion = await get_generic_sso_response( + request=mock_request, + jwt_handler=jwt_handler, + generic_client_id="cid", + redirect_url="https://proxy.example.com/sso/callback", + sso_jwt_handler=None, + ) + + assert sso_assertion is not None + assert sso_assertion.id_token.get_secret_value() == id_token + assert sso_assertion.refresh_token is not None + assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp" + # The sanitized received_response must still not carry bearer material. + assert "id_token" not in (received_response or {}) + assert "refresh_token" not in (received_response or {}) + + +@pytest.mark.asyncio +async def test_verify_and_process_arm_captures_sso_assertion(): + """The non-PKCE generic arm reads the raw bearer fields off the fastapi-sso client.""" + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + id_token = _ema_id_token() + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + mock_jwt_handler.get_team_ids_from_jwt.return_value = [] + + mock_sso_instance = MagicMock() + mock_sso_instance.verify_and_process = AsyncMock( + return_value={"sub": "u1", "email": "u@example.com"} + ) + mock_sso_instance.access_token = None + mock_sso_instance.id_token = id_token + mock_sso_instance.refresh_token = "rt_from_idp" + mock_sso_class = MagicMock(return_value=mock_sso_instance) + + with patch.dict( + os.environ, + { + "GENERIC_CLIENT_SECRET": "test_secret", + "GENERIC_AUTHORIZATION_ENDPOINT": "https://auth.example.com/auth", + "GENERIC_TOKEN_ENDPOINT": "https://auth.example.com/token", + "GENERIC_USERINFO_ENDPOINT": "https://auth.example.com/userinfo", + }, + ): + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch( + "fastapi_sso.sso.generic.create_provider", return_value=mock_sso_class + ): + _, _, _, sso_assertion = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test_client_id", + redirect_url="http://test.com/callback", + sso_jwt_handler=None, + ) + + assert sso_assertion is not None + assert sso_assertion.id_token.get_secret_value() == id_token + assert sso_assertion.refresh_token is not None + assert sso_assertion.refresh_token.get_secret_value() == "rt_from_idp" + + +@pytest.mark.asyncio +async def test_redirect_from_openid_persists_assertion_under_canonical_user_id(): + """The browser funnel persists the captured assertion AFTER canonical user + resolution, keyed by the user_id admission will later resolve (the key-generation + response user_id), not the raw IdP subject.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import ( + assertion_from_sso_login, + ) + + assertion = assertion_from_sso_login(_ema_id_token(), "rt_1") + assert assertion is not None + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:4000/" + mock_request.cookies = {} + + retain_mock = AsyncMock() + with ( + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.premium_user", False), + patch("litellm.proxy.proxy_server.user_custom_sso", None), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.redis_usage_cache", None), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch( + "litellm.proxy.proxy_server.generate_key_helper_fn", + AsyncMock( + return_value={"token": "sk-ui-key", "user_id": "canonical-user-id"} + ), + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db", + AsyncMock(return_value=None), + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.check_and_update_if_proxy_admin_id", + AsyncMock(return_value="internal_user"), + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema", + retain_mock, + ), + ): + response = await SSOAuthenticationHandler.get_redirect_response_from_openid( + result=CustomOpenID( + id="raw-idp-subject", + email="u@example.com", + first_name="U", + last_name="Ser", + display_name="U Ser", + provider="generic", + team_ids=[], + user_role=None, + ), + request=mock_request, + received_response=None, + generic_client_id="cid", + ui_access_mode=None, + access_token_payload=None, + jwt_handler=None, + sso_assertion=assertion, + ) + + retain_mock.assert_awaited_once_with( + user_id="canonical-user-id", assertion=assertion + ) + assert response is not None + + +@pytest.mark.asyncio +async def test_cli_completion_persists_assertion_under_db_user_id(): + """The CLI funnel persists the captured assertion under the DB-resolved user_id.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.sso_assertion_store import ( + assertion_from_sso_login, + ) + from litellm.proxy.management_endpoints.ui_sso import ( + _complete_cli_sso_callback_session, + ) + + assertion = assertion_from_sso_login(_ema_id_token(), None) + assert assertion is not None + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:4000/" + + user_info = MagicMock() + user_info.user_id = "cli-user-id" + user_info.user_role = "internal_user" + user_info.models = [] + user_info.teams = [] + + retain_mock = AsyncMock() + with ( + patch( + "litellm.proxy.management_endpoints.ui_sso.get_user_info_from_db", + AsyncMock(return_value=user_info), + ), + patch( + "litellm.proxy.management_endpoints.ui_sso._fetch_cli_sso_team_details", + AsyncMock(return_value=[]), + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.build_cli_sso_attribution_metadata", + return_value={}, + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.retain_sso_identity_assertion_for_ema", + retain_mock, + ), + ): + response = await _complete_cli_sso_callback_session( + request=mock_request, + key="cli-login-id", + flow={}, + result={"sub": "raw-idp-subject"}, + parsed_openid_result={ + "user_id": "raw-idp-subject", + "user_email": "u@example.com", + "user_role": None, + }, + user_defined_values=None, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + cli_sso_session_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + sso_assertion=assertion, + ) + + retain_mock.assert_awaited_once_with(user_id="cli-user-id", assertion=assertion) + assert response.status_code == 200 diff --git a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py index dbbdca65cc6..01e5414a469 100644 --- a/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py +++ b/tests/test_litellm/proxy/management_helpers/test_management_helpers_utils.py @@ -202,7 +202,8 @@ async def test_add_new_member_clones_default_team_budget_id(): "teams": [test_team_id], "user_role": "internal_user", } - mock_prisma_client.db.litellm_usertable.upsert = AsyncMock( + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -305,7 +306,8 @@ async def test_add_new_member_budget_duration_only_clones_default_max_budget(): "teams": ["team-dc"], "user_role": "internal_user", } - mock_prisma_client.db.litellm_usertable.upsert = AsyncMock( + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) mock_default_budget_row = MagicMock() @@ -388,7 +390,8 @@ async def test_add_new_member_no_budget_when_no_default_and_no_max_budget(): "teams": [test_team_id], "user_role": "internal_user", } - mock_prisma_client.db.litellm_usertable.upsert = AsyncMock( + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -455,7 +458,8 @@ async def test_add_new_member_creates_new_budget_when_max_budget_provided(): "teams": [test_team_id], "user_role": "internal_user", } - mock_prisma_client.db.litellm_usertable.upsert = AsyncMock( + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) @@ -531,7 +535,8 @@ async def test_add_new_member_persists_budget_duration(): "teams": ["team-dur"], "user_role": "internal_user", } - mock_prisma_client.db.litellm_usertable.upsert = AsyncMock( + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) mock_budget_response = MagicMock() @@ -594,7 +599,8 @@ async def test_add_new_member_persists_budget_duration_without_max_budget(): "teams": ["team-dur2"], "user_role": "internal_user", } - mock_prisma_client.db.litellm_usertable.upsert = AsyncMock( + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response) + mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( return_value=mock_user_response ) mock_budget_response = MagicMock() @@ -997,3 +1003,116 @@ async def test_attach_object_permission_to_dict_with_none_object_permission_id() # Verify no database query was made mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_not_called() + + +@pytest.mark.asyncio +async def test_add_new_member_appends_team_only_if_absent_for_existing_user(): + """Adding an existing user to a team must append the team id only if it is + not already present. + + add_new_member is the single writer of user.teams for every team add + (/team/member_add, /user/new, SSO, SCIM). An unconditional append let + repeated or concurrent adds accumulate duplicate team ids in user.teams, + which also breaks auth logic that keys off the number of teams a user + belongs to. The append must go through a filtered update that no-ops when + the team is already present, and it must not fall through to creating a new + user row for a user that already exists. + """ + from litellm.proxy._types import LitellmUserRoles + + new_member = Member(user_id="existing-user", role="user") + user_api_key_dict = UserAPIKeyAuth( + user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + mock_prisma_client = AsyncMock() + + mock_user_after = MagicMock() + mock_user_after.model_dump.return_value = { + "user_id": "existing-user", + "user_email": None, + "teams": ["team-1"], + "user_role": "internal_user", + } + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_after) + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() + # no team default budget and no explicit budget -> no team membership row + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + + result_user, _ = await add_new_member( + new_member=new_member, + max_budget_in_team=None, + prisma_client=mock_prisma_client, + team_id="team-1", + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name="admin", + ) + + assert result_user is not None + assert result_user.user_id == "existing-user" + + # the append must be a filtered, idempotent update keyed off the team id, so + # a repeated or concurrent add of a team the user already has is a no-op + mock_prisma_client.db.litellm_usertable.update_many.assert_called_once() + where = mock_prisma_client.db.litellm_usertable.update_many.call_args.kwargs["where"] + assert where["user_id"] == "existing-user" + assert where["NOT"] == {"teams": {"has": "team-1"}} + data = mock_prisma_client.db.litellm_usertable.update_many.call_args.kwargs["data"] + assert data == {"teams": {"push": ["team-1"]}} + + # upsert (not an unconditional teams push) is what ensures the row exists, so + # its update branch must not carry a teams push that would duplicate + mock_prisma_client.db.litellm_usertable.upsert.assert_called_once() + upsert_update = mock_prisma_client.db.litellm_usertable.upsert.call_args.kwargs["data"]["update"] + assert "teams" not in upsert_update + + +@pytest.mark.asyncio +async def test_add_new_member_creates_missing_user_atomically_via_upsert(): + """A brand-new user added to a team must be created via an atomic upsert, not + a separate existence check followed by create. + + Concurrent provisioning of the same new user (which SCIM group reconciles do) + would race a check-then-create into a duplicate-key failure. The upsert seeds + teams on create, and the filtered append is a no-op because the team is + already present on the freshly created row. + """ + from litellm.proxy._types import LitellmUserRoles + + new_member = Member(user_id="brand-new-user", role="user") + user_api_key_dict = UserAPIKeyAuth( + user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN + ) + + mock_prisma_client = AsyncMock() + + mock_created = MagicMock() + mock_created.model_dump.return_value = { + "user_id": "brand-new-user", + "user_email": None, + "teams": ["team-1"], + "user_role": "internal_user", + } + mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_created) + mock_prisma_client.db.litellm_usertable.update_many = AsyncMock() + mock_prisma_client.db.litellm_usertable.create = AsyncMock() + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None) + + result_user, _ = await add_new_member( + new_member=new_member, + max_budget_in_team=None, + prisma_client=mock_prisma_client, + team_id="team-1", + user_api_key_dict=user_api_key_dict, + litellm_proxy_admin_name="admin", + ) + + assert result_user is not None + assert result_user.user_id == "brand-new-user" + + # existence is established by an atomic upsert (create-or-update), never a + # non-atomic standalone create that could race under concurrent provisioning + mock_prisma_client.db.litellm_usertable.upsert.assert_called_once() + mock_prisma_client.db.litellm_usertable.create.assert_not_called() + create_data = mock_prisma_client.db.litellm_usertable.upsert.call_args.kwargs["data"]["create"] + assert create_data["teams"] == ["team-1"] 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 93b21c9d3c1..8d1d8185e4d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -8,6 +8,7 @@ Pins covered: from __future__ import annotations +import json import os from types import SimpleNamespace from typing import Any, Dict @@ -407,6 +408,124 @@ async def test_ProxyConfig_save_config_invalid_path_raises(monkeypatch): await pc.save_config({"x": 1}) +@pytest.mark.asyncio +async def test_ProxyConfig_save_config_db_omits_environment_variables_by_default(monkeypatch): + """A save_config after get_config() (which resolves os.environ/ placeholders + to plaintext and merges the environment_variables section) must not snapshot + those env vars into the DB config row. Persisting them would make a stale DB + row shadow YAML/container env on every subsequent restart.""" + mock_prisma = MagicMock() + mock_prisma.insert_data = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + # a valid salt so the env-var encryption path (reached only if the pop + # regresses) runs cleanly, making this fail on the assertion below rather + # than on an incidental encryption crash + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-key") + + pc = ProxyConfig() + cfg = { + "model_list": [{"model_name": "gpt-4o"}], + "litellm_settings": {"success_callback": ["langfuse"]}, + "environment_variables": {"OPENAI_API_KEY": "sk-from-yaml"}, + } + await pc.save_config(cfg) + + mock_prisma.insert_data.assert_awaited_once() + written = mock_prisma.insert_data.await_args.kwargs["data"] + assert "environment_variables" not in written + # unrelated sections are still persisted; model_list is stripped as before + assert written["litellm_settings"] == {"success_callback": ["langfuse"]} + assert "model_list" not in written + # the caller's dict is not mutated (save_config works on a copy) + assert cfg["environment_variables"] == {"OPENAI_API_KEY": "sk-from-yaml"} + + +@pytest.mark.asyncio +async def test_ProxyConfig_save_config_db_persists_environment_variables_when_opted_in(monkeypatch): + """The explicit opt-in path (include_env_vars=True) still persists env vars, + encrypted, so the dedicated config-update flow can write them.""" + mock_prisma = MagicMock() + mock_prisma.insert_data = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-key") + + pc = ProxyConfig() + cfg = {"litellm_settings": {}, "environment_variables": {"OPENAI_API_KEY": "sk-explicit"}} + await pc.save_config(cfg, include_env_vars=True) + + mock_prisma.insert_data.assert_awaited_once() + written = mock_prisma.insert_data.await_args.kwargs["data"] + assert set(written["environment_variables"].keys()) == {"OPENAI_API_KEY"} + # value is encrypted at rest, not the plaintext it came in as + assert written["environment_variables"]["OPENAI_API_KEY"] != "sk-explicit" + + +def _install_fake_config_repo(monkeypatch, existing_row): + """Route ProxyConfig's ConfigRepository through an in-memory fake that + records the value written to the environment_variables row.""" + captured: dict = {} + + class _FakeTable: + async def find_first(self, where): + return SimpleNamespace(param_value=existing_row) if existing_row is not None else None + + async def upsert(self, where, data): + captured["value"] = json.loads(data["update"]["param_value"]) + + class _FakeRepo: + def __init__(self, client): + self.table = _FakeTable() + + monkeypatch.setattr("litellm.proxy.proxy_server.ConfigRepository", _FakeRepo) + monkeypatch.setattr("litellm.proxy.proxy_server.invalidate_config_param", AsyncMock()) + return captured + + +@pytest.mark.asyncio +async def test_ProxyConfig_save_environment_variables_merges_sets_and_deletes(monkeypatch): + """The per-key env-var write updates/deletes only the named keys and leaves + every other stored key untouched, so an unrelated env var is never lost or + snapshotted.""" + captured = _install_fake_config_repo( + monkeypatch, + existing_row={"EXISTING_KEY": "ciphertext-existing", "UI_LOGO_PATH": "old-logo", "LITELLM_FAVICON_URL": "old"}, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-test-salt-key") + + pc = ProxyConfig() + await pc.save_environment_variables({"UI_LOGO_PATH": "new-logo", "LITELLM_FAVICON_URL": None}) + + written = captured["value"] + # unrelated key preserved byte-for-byte + assert written["EXISTING_KEY"] == "ciphertext-existing" + # set key updated and encrypted (not the plaintext) + assert "UI_LOGO_PATH" in written and written["UI_LOGO_PATH"] != "new-logo" + # None-valued key deleted + assert "LITELLM_FAVICON_URL" not in written + + +@pytest.mark.asyncio +async def test_ProxyConfig_save_environment_variables_noop_without_db(monkeypatch): + """With no DB configured the per-key write must do nothing (never touch the + config repository).""" + captured = _install_fake_config_repo(monkeypatch, existing_row={}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {}) + + pc = ProxyConfig() + await pc.save_environment_variables({"UI_LOGO_PATH": "x"}) + + assert "value" not in captured + + # --------------------------------------------------------------------------- # ProxyConfig._check_for_os_environ_vars # --------------------------------------------------------------------------- @@ -950,6 +1069,32 @@ async def test_ProxyConfig_load_config_wires_general_settings_url_validation(tmp litellm.provider_url_destination_allowed_hosts = original_provider_hosts +@pytest.mark.asyncio +async def test_ProxyConfig_load_config_wires_config_reload_interval(tmp_path, monkeypatch): + """general_settings.proxy_config_reload_interval_seconds must reach the proxy_server + module global that schedules the DB config-reload jobs, so operators can tune multi-pod + convergence from config.yaml.""" + import litellm.proxy.proxy_server as proxy_server + + f = tmp_path / "c.yaml" + f.write_text( + "model_list: []\n" + "general_settings:\n" + " proxy_config_reload_interval_seconds: 47\n" + "litellm_settings: {}\n" + ) + 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) + + original = proxy_server.proxy_config_reload_interval_seconds + try: + await ProxyConfig().load_config(router=None, config_file_path=str(f)) + assert proxy_server.proxy_config_reload_interval_seconds == 47 + finally: + proxy_server.proxy_config_reload_interval_seconds = original + + @pytest.mark.asyncio async def test_ProxyConfig_load_config_missing_file_raises(monkeypatch): monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_config.py b/tests/test_litellm/proxy/proxy_server/test_routes_config.py index 4ac6fc46a61..ad3c470acf3 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_config.py @@ -13,6 +13,7 @@ Routes covered: from __future__ import annotations +import json from unittest.mock import AsyncMock, MagicMock from .conftest import VOLATILE_KEYS, normalize @@ -473,6 +474,83 @@ def test_config_list_happy_admin(client, auth_as, mock_prisma, monkeypatch): } +def test_config_list_exposes_config_reload_interval(client, auth_as, mock_prisma, monkeypatch): + """proxy_config_reload_interval_seconds must surface in the admin UI general-settings + list as an Integer field defaulting to 30, so operators can tune multi-pod convergence + from the dashboard.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + row = MagicMock() + row.param_value = {} + table.find_first = AsyncMock(return_value=row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.get("/config/list", params={"config_type": "general_settings"}) + assert response.status_code == 200 + by_name = {entry["field_name"]: entry for entry in response.json()} + assert "proxy_config_reload_interval_seconds" in by_name + entry = by_name["proxy_config_reload_interval_seconds"] + assert entry["field_type"] == "Integer" + assert entry["field_default_value"] == 30 + + +def test_config_field_update_accepts_config_reload_interval(client, auth_as, mock_prisma, monkeypatch): + """POST /config/field/update accepts proxy_config_reload_interval_seconds and persists + it to the DB general_settings row for all pods to pick up.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + table.find_first = AsyncMock(return_value=None) + upsert_row = { + "param_name": "general_settings", + "param_value": {"proxy_config_reload_interval_seconds": 45}, + "id": "row-1", + } + table.upsert = AsyncMock(return_value=upsert_row) + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/field/update", + json={ + "field_name": "proxy_config_reload_interval_seconds", + "field_value": 45, + "config_type": "general_settings", + }, + ) + assert response.status_code == 200 + upserted = table.upsert.call_args.kwargs["data"]["create"]["param_value"] + assert json.loads(upserted)["proxy_config_reload_interval_seconds"] == 45 + + +def test_config_field_update_rejects_non_positive_config_reload_interval(client, auth_as, mock_prisma, monkeypatch): + """A non-positive proxy_config_reload_interval_seconds from the UI is rejected with a 400 + and never persisted, since APScheduler requires a positive interval.""" + from litellm.proxy import proxy_server as ps + from litellm.proxy._types import LitellmUserRoles + + table = _install_litellm_config(mock_prisma) + table.find_first = AsyncMock(return_value=None) + table.upsert = AsyncMock() + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + + with auth_as(LitellmUserRoles.PROXY_ADMIN): + response = client.post( + "/config/field/update", + json={ + "field_name": "proxy_config_reload_interval_seconds", + "field_value": 0, + "config_type": "general_settings", + }, + ) + assert response.status_code == 400 + table.upsert.assert_not_called() + + def test_config_list_non_admin_rejected(client, auth_as, mock_prisma, monkeypatch): """Non-admin gets a 400 with the role embedded in the error message.""" from litellm.proxy import proxy_server as ps diff --git a/tests/test_litellm/proxy/test_provider_url_destination_guard.py b/tests/test_litellm/proxy/test_provider_url_destination_guard.py index 51cd76105d0..c8771abbc8e 100644 --- a/tests/test_litellm/proxy/test_provider_url_destination_guard.py +++ b/tests/test_litellm/proxy/test_provider_url_destination_guard.py @@ -39,6 +39,46 @@ class TestRejectUrlValuedDestinations: assert exc_info.value.status_code == 400 assert exc_info.value.detail["param"] == "model" + def test_provider_prefixed_url_rejected(self): + with pytest.raises(HTTPException) as exc_info: + _reject_url_valued_destinations( + {"model": "huggingface/https://attacker.example/v1"} + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["param"] == "model" + + def test_comma_batch_smuggled_url_rejected(self): + with pytest.raises(HTTPException) as exc_info: + _reject_url_valued_destinations( + {"model": "gpt-4,huggingface/https://attacker.example/v1"} + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["param"] == "model" + + def test_provider_prefixed_uppercase_scheme_url_rejected(self): + with pytest.raises(HTTPException) as exc_info: + _reject_url_valued_destinations( + {"model": "huggingface/HTTPS://evil.example/v1"} + ) + assert exc_info.value.status_code == 400 + assert exc_info.value.detail["param"] == "model" + + def test_provider_prefixed_plain_model_passes(self): + _reject_url_valued_destinations({"model": "huggingface/BAAI/bge-small-en"}) + + def test_comma_batch_plain_models_pass(self): + _reject_url_valued_destinations({"model": "gpt-4,huggingface/BAAI/bge-small-en"}) + + def test_provider_prefixed_url_respects_allowlist(self, monkeypatch): + monkeypatch.setattr( + litellm, + "provider_url_destination_allowed_hosts", + ["trusted.example"], + ) + _reject_url_valued_destinations( + {"model": "huggingface/https://trusted.example/v1"} + ) + def test_url_valued_file_id_rejected(self): with pytest.raises(HTTPException) as exc_info: _reject_url_valued_destinations( diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index a100e7837f4..5b67780dc58 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -751,6 +751,97 @@ async def test_initialize_scheduled_jobs_credentials(monkeypatch): assert len(mock_scheduler_calls) > 0 +@pytest.mark.asyncio +async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(monkeypatch): + """ + The DB config-reload jobs (add_deployment, get_credentials) that keep multi-pod + deployments in sync must be scheduled at the configured + proxy_config_reload_interval_seconds, not a hardcoded value. + """ + monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False) + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + from litellm.proxy.proxy_server import ProxyStartupEvent + from litellm.proxy.utils import ProxyLogging + + mock_prisma_client = MagicMock() + mock_proxy_logging = MagicMock(spec=ProxyLogging) + mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_config = AsyncMock() + mock_scheduler = MagicMock() + + configured_interval = 47 + + with ( + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True), + patch( + "litellm.proxy.proxy_server.proxy_config_reload_interval_seconds", + configured_interval, + ), + patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler), + ): + await ProxyStartupEvent.initialize_scheduled_background_jobs( + general_settings={}, + prisma_client=mock_prisma_client, + proxy_budget_rescheduler_min_time=1, + proxy_budget_rescheduler_max_time=2, + proxy_batch_write_at=5, + proxy_logging_obj=mock_proxy_logging, + ) + + scheduled_seconds = { + job_call.kwargs["id"]: job_call.kwargs.get("seconds") + for job_call in mock_scheduler.add_job.call_args_list + if "id" in job_call.kwargs + } + assert scheduled_seconds["add_deployment_job"] == configured_interval + assert scheduled_seconds["get_credentials_job"] == configured_interval + + +@pytest.mark.asyncio +async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_interval(monkeypatch): + """ + A non-positive proxy_config_reload_interval_seconds (misconfig via env/config/DB) would + make APScheduler reject the job and crash startup, so the scheduler must fall back to the + 30s default instead of forwarding the bad value. + """ + monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False) + monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False) + from litellm.proxy.proxy_server import ProxyStartupEvent + from litellm.proxy.utils import ProxyLogging + + mock_prisma_client = MagicMock() + mock_proxy_logging = MagicMock(spec=ProxyLogging) + mock_proxy_logging.slack_alerting_instance = MagicMock() + mock_proxy_config = AsyncMock() + mock_scheduler = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config), + patch("litellm.proxy.proxy_server.store_model_in_db", True), + patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True), + patch("litellm.proxy.proxy_server.proxy_config_reload_interval_seconds", 0), + patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler), + ): + await ProxyStartupEvent.initialize_scheduled_background_jobs( + general_settings={}, + prisma_client=mock_prisma_client, + proxy_budget_rescheduler_min_time=1, + proxy_budget_rescheduler_max_time=2, + proxy_batch_write_at=5, + proxy_logging_obj=mock_proxy_logging, + ) + + scheduled_seconds = { + job_call.kwargs["id"]: job_call.kwargs.get("seconds") + for job_call in mock_scheduler.add_job.call_args_list + if "id" in job_call.kwargs + } + assert scheduled_seconds["add_deployment_job"] == 30 + assert scheduled_seconds["get_credentials_job"] == 30 + + @pytest.mark.asyncio async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_false(monkeypatch): """ @@ -924,6 +1015,102 @@ def test_get_config_custom_callback_api_env_vars(monkeypatch): } +@patch( + "litellm.proxy.common_utils.callback_utils.CustomLogger.get_callback_env_vars", + return_value=["LANGFUSE_PUBLIC_KEY", "LANGFUSE_SECRET_KEY", "LANGFUSE_HOST"], +) +def test_get_config_callbacks_fall_back_to_process_env(mock_env_vars, monkeypatch): + """A callback configured purely via process env vars is surfaced. + + An IaC deployment sets LANGFUSE_* on the gateway and never touches the UI, + so nothing is stored in the config environment_variables overlay. The read + endpoint must still report the live values instead of blanks. + """ + from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth + + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-env-only") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-env-only") + monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com") + + config_data = { + "litellm_settings": {"success_callback": ["langfuse"]}, + "general_settings": {}, + "environment_variables": {}, + } + mock_router = MagicMock() + mock_router.get_settings.return_value = {} + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data)) + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234" + ) + + client = TestClient(app) + try: + response = client.get("/get/config/callbacks") + finally: + app.dependency_overrides = original_overrides + + assert response.status_code == 200 + langfuse_cb = next( + (cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None + ) + assert langfuse_cb is not None + assert langfuse_cb["variables"] == { + "LANGFUSE_PUBLIC_KEY": "pk-env-only", + "LANGFUSE_SECRET_KEY": "sk-env-only", + "LANGFUSE_HOST": "https://cloud.langfuse.com", + } + + +@patch( + "litellm.proxy.common_utils.callback_utils.CustomLogger.get_callback_env_vars", + return_value=["LANGFUSE_SECRET_KEY", "LANGFUSE_HOST"], +) +def test_get_config_callback_env_secrets_redacted_for_non_admin(mock_env_vars, monkeypatch): + """Surfacing env vars must not widen who can read secret values. + + The callback role gate redacts sensitive keys for anyone below full admin, + and that must hold whether the value came from the stored config or the + process env. A non-secret var (LANGFUSE_HOST) still resolves for context. + """ + from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth + + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-env-only-secret") + monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com") + + config_data = { + "litellm_settings": {"success_callback": ["langfuse"]}, + "general_settings": {}, + "environment_variables": {}, + } + mock_router = MagicMock() + mock_router.get_settings.return_value = {} + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data)) + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user" + ) + + client = TestClient(app) + try: + response = client.get("/get/config/callbacks") + finally: + app.dependency_overrides = original_overrides + + assert response.status_code == 200 + langfuse_cb = next( + (cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None + ) + assert langfuse_cb is not None + assert langfuse_cb["variables"]["LANGFUSE_SECRET_KEY"] == "REDACTED" + assert langfuse_cb["variables"]["LANGFUSE_HOST"] == "https://cloud.langfuse.com" + + def test_get_config_returns_email_settings(monkeypatch): """ Regression for https://github.com/BerriAI/litellm/issues/19221 @@ -987,6 +1174,113 @@ def test_get_config_returns_email_settings(monkeypatch): assert "*" in variables["SMTP_PASSWORD"] +def _get_email_alert_variables(monkeypatch, config_data): + from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth + + mock_router = MagicMock() + mock_router.get_settings.return_value = {} + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router) + monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data)) + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234" + ) + + client = TestClient(app) + try: + response = client.get("/get/config/callbacks") + finally: + app.dependency_overrides = original_overrides + + assert response.status_code == 200 + email_alert = next((a for a in response.json()["alerts"] if a["name"] == "email"), None) + assert email_alert is not None + return email_alert["variables"] + + +def test_get_config_returns_email_settings_set_only_in_process_env(monkeypatch): + """ + Regression for LIT-4165. + + SMTP supplied purely as process env vars (helm/terraform, no UI writes) is + live at runtime because litellm/proxy/utils.py::send_email resolves every + field from os.getenv. The /get/config/callbacks email block only read the + config/DB environment_variables overlay though, so those deployments saw an + empty Email Server Settings page and could not tell SMTP was configured. + The slack block one branch above already fell back to os.getenv. + """ + smtp_password = "env-only-app-password" + monkeypatch.setenv("SMTP_HOST", "smtp.env-host.com") + monkeypatch.setenv("SMTP_PORT", "2525") + monkeypatch.setenv("SMTP_TLS", "False") + monkeypatch.setenv("SMTP_USERNAME", "env-user") + monkeypatch.setenv("SMTP_PASSWORD", smtp_password) + monkeypatch.setenv("SMTP_SENDER_EMAIL", "alerts@env-host.com") + monkeypatch.setenv("TEST_EMAIL_ADDRESS", "admin@env-host.com") + + variables = _get_email_alert_variables( + monkeypatch, + { + "litellm_settings": {}, + "general_settings": {"alerting": ["email"]}, + "environment_variables": {}, + }, + ) + + # Every one of these was None before the fix, despite SMTP working. + assert variables["SMTP_HOST"] == "smtp.env-host.com" + assert variables["SMTP_PORT"] == "2525" + assert variables["SMTP_TLS"] == "False" + assert variables["SMTP_USERNAME"] == "env-user" + assert variables["SMTP_SENDER_EMAIL"] == "alerts@env-host.com" + assert variables["TEST_EMAIL_ADDRESS"] == "admin@env-host.com" + + # An env-sourced secret is masked exactly like a stored one. + assert variables["SMTP_PASSWORD"] not in (None, smtp_password) + assert "*" in variables["SMTP_PASSWORD"] + + +def test_get_config_email_settings_prefer_stored_over_process_env(monkeypatch): + """ + Stored environment_variables win over the process environment, matching the + load order in ProxyConfig.get_config, which pushes stored values into + os.environ. Only a field with no stored entry falls back to os.getenv. + """ + monkeypatch.setenv("SMTP_HOST", "smtp.env-host.com") + monkeypatch.setenv("SMTP_SENDER_EMAIL", "alerts@env-host.com") + + variables = _get_email_alert_variables( + monkeypatch, + { + "litellm_settings": {}, + "general_settings": {"alerting": ["email"]}, + "environment_variables": {"SMTP_HOST": "smtp.stored-host.com"}, + }, + ) + + assert variables["SMTP_HOST"] == "smtp.stored-host.com" + assert variables["SMTP_SENDER_EMAIL"] == "alerts@env-host.com" + + +def test_get_config_email_settings_absent_everywhere_stay_none(monkeypatch): + """A field set in neither source is reported unset rather than invented.""" + for var in ("SMTP_HOST", "SMTP_PORT", "SMTP_TLS", "SMTP_USERNAME", "SMTP_PASSWORD", "SMTP_SENDER_EMAIL"): + monkeypatch.delenv(var, raising=False) + + variables = _get_email_alert_variables( + monkeypatch, + { + "litellm_settings": {}, + "general_settings": {"alerting": ["email"]}, + "environment_variables": {}, + }, + ) + + assert variables["SMTP_HOST"] is None + assert variables["SMTP_PASSWORD"] is None + + def test_get_config_returns_slack_webhook(monkeypatch): """ Same double-decryption regression as the email block (issue #19221): the @@ -2588,6 +2882,69 @@ async def test_load_config_max_budget_env_var_coerced_to_float(tmp_path, monkeyp litellm.max_budget = original_max_budget +def test_max_ui_session_budget_default_is_one_dollar(): + """LIT-4662: the dashboard session budget default is a product decision; the + old 0.25 default locked admins out of auto router Test Connection and the + playground mid-session with an error that looked like a hardcoded cap.""" + assert litellm.max_ui_session_budget == 1.0 + + +@pytest.mark.asyncio +async def test_load_config_max_ui_session_budget_applied_and_coerced(tmp_path, monkeypatch): + """ + max_ui_session_budget configured via os.environ resolves to a string; + load_config must coerce it to float so every dashboard session key is + minted with a numeric max_budget. + """ + from litellm.proxy.proxy_server import ProxyConfig + + monkeypatch.setenv("UI_SESSION_BUDGET", "2.5") + test_config = { + "model_list": [], + "litellm_settings": {"max_ui_session_budget": "os.environ/UI_SESSION_BUDGET"}, + } + config_file = tmp_path / "config.yaml" + config_file.write_text(yaml.dump(test_config)) + + original_budget = litellm.max_ui_session_budget + try: + proxy_config = ProxyConfig() + await proxy_config.load_config( + router=MagicMock(), config_file_path=str(config_file) + ) + assert isinstance(litellm.max_ui_session_budget, float) + assert litellm.max_ui_session_budget == 2.5 + finally: + litellm.max_ui_session_budget = original_budget + + +@pytest.mark.asyncio +async def test_load_config_max_ui_session_budget_none_disables_cap(tmp_path): + """ + max_ui_session_budget: null in config disables the dashboard session cap + entirely (session keys minted with no max_budget); load_config must pass + None through instead of raising on float(None). + """ + from litellm.proxy.proxy_server import ProxyConfig + + test_config = { + "model_list": [], + "litellm_settings": {"max_ui_session_budget": None}, + } + config_file = tmp_path / "config.yaml" + config_file.write_text(yaml.dump(test_config)) + + original_budget = litellm.max_ui_session_budget + try: + proxy_config = ProxyConfig() + await proxy_config.load_config( + router=MagicMock(), config_file_path=str(config_file) + ) + assert litellm.max_ui_session_budget is None + finally: + litellm.max_ui_session_budget = original_budget + + @pytest.mark.asyncio async def test_load_config_default_internal_user_params_max_budget_scientific_notation(tmp_path): """ @@ -9048,6 +9405,85 @@ def test_general_settings_ui_fields_are_db_overridable(): ) +@pytest.mark.asyncio +async def test_update_config_field_max_ui_session_budget_sets_live_value(monkeypatch): + """LIT-4662: the dashboard session budget is editable from the Admin UI General tab. + A Dollar field must accept values above 1 (the old Float type capped at 1, which cannot + express a dollar budget), apply live via setattr, and persist under litellm_settings.""" + from unittest.mock import MagicMock + + import litellm.proxy.proxy_server as ps + from litellm.proxy._types import ( + ConfigFieldUpdate, + LitellmUserRoles, + UserAPIKeyAuth, + ) + from litellm.proxy.proxy_server import update_config_general_settings + + saved: dict = {} + + async def fake_get_config(): + return {"litellm_settings": {}} + + async def fake_save_config(new_config=None): + saved.update(new_config or {}) + + monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config) + monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config) + monkeypatch.setattr(ps, "prisma_client", MagicMock()) + monkeypatch.setattr(litellm, "store_audit_logs", False) + monkeypatch.setattr(litellm, "max_ui_session_budget", 1.0) + + admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN) + await update_config_general_settings( + data=ConfigFieldUpdate( + field_name="max_ui_session_budget", + field_value=25.0, + config_type="general_settings", + ), + user_api_key_dict=admin, + ) + + assert litellm.max_ui_session_budget == 25.0 + assert saved["litellm_settings"]["max_ui_session_budget"] == 25.0 + + +@pytest.mark.parametrize("bad_value", [True, "abc", -1, 0, [2.5]]) +def test_validate_max_ui_session_budget_rejects_malformed(bad_value): + """A Dollar field accepts only positive numbers; zero would block every dashboard + LLM call at mint and non-numerics would break session key generation.""" + from fastapi import HTTPException + + from litellm.proxy.proxy_server import _validate_general_settings_ui_litellm_value + + with pytest.raises(HTTPException) as exc_info: + _validate_general_settings_ui_litellm_value("max_ui_session_budget", bad_value) + assert exc_info.value.status_code == 400 + + +@pytest.mark.parametrize("empty_value", [None, ""]) +def test_validate_max_ui_session_budget_empty_restores_default(empty_value): + """Clearing the field in the UI restores the shipped $1 default rather than None; + None would silently remove the session spend guardrail (unlimited budget), which + must stay a deliberate config.yaml act (max_ui_session_budget: null).""" + from litellm.proxy.proxy_server import _validate_general_settings_ui_litellm_value + + assert _validate_general_settings_ui_litellm_value("max_ui_session_budget", empty_value) == 1.0 + + +def test_general_settings_ui_defaults_unchanged_for_existing_fields(): + """The spec-default mechanism added for max_ui_session_budget must not change what + clearing the pre-existing fields restores (None for Float/Select, False for Boolean).""" + from litellm.proxy.proxy_server import ( + _GENERAL_SETTINGS_UI_LITELLM_FIELDS, + _general_settings_ui_litellm_default, + ) + + assert _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["budget_exceeded_throttle_percentage"]) is None + assert _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["enable_anthropic_prompt_caching"]) is False + assert _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["anthropic_prompt_caching_ttl"]) is None + + @pytest.mark.parametrize( "field_name, db_value", [ diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 5ace46fc775..4673807a135 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -709,6 +709,39 @@ def test_create_model_info_response_reads_real_cost_map(): assert response["max_output_tokens"] > 0 +def test_create_model_info_response_includes_mode_from_lookup(): + response = create_model_info_response( + model_id="text-embedding-3-small", + provider="openai", + llm_router=None, + get_model_info=lambda _model: _fake_model_info(mode="embedding"), + ) + + assert response["mode"] == "embedding" + + +def test_create_model_info_response_omits_mode_when_lookup_raises(): + response = create_model_info_response( + model_id="my-custom-deployment", + provider="openai", + llm_router=None, + get_model_info=_raise_unmapped, + ) + + assert "mode" not in response + + +def test_create_model_info_response_omits_non_string_mode(): + response = create_model_info_response( + model_id="some-model", + provider="openai", + llm_router=None, + get_model_info=lambda _model: _fake_model_info(mode=None), + ) + + assert "mode" not in response + + class TestPostCallFailureHookLLMExceptionAlerting: """The llm_exceptions alert is for infra / LLM-API failures, not user errors (https://github.com/BerriAI/litellm/issues/3395). Already-normalized diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py index d0cb5ec5465..849d5494c6e 100644 --- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -54,8 +54,8 @@ def _patched_init_cache(litellm_settings: dict, cache_params: dict): _FakeRedisCache (passes the isinstance guard in _init_cache). 3. Extracts enable_redis_auth_cache from litellm_settings and passes it as the second argument to _init_cache (matching production behaviour). - 4. Yields (user_api_key_cache, spend_counter_cache) after calling - _init_cache, then restores everything. + 4. Yields (user_api_key_cache, spend_counter_cache, cli_sso_session_cache) + after calling _init_cache, then restores everything. """ fake_redis = _FakeRedisCache() @@ -64,19 +64,21 @@ def _patched_init_cache(litellm_settings: dict, cache_params: dict): fresh_user_cache = DualCache() fresh_spend_cache = DualCache() + fresh_cli_sso_cache = DualCache() enable_redis_auth_cache = litellm_settings.get("enable_redis_auth_cache", False) with ( patch.object(ps, "user_api_key_cache", fresh_user_cache), patch.object(ps, "spend_counter_cache", fresh_spend_cache), + patch.object(ps, "cli_sso_session_cache", fresh_cli_sso_cache), patch.object(ps, "llm_router", None), # Cache is locally imported inside _init_cache: patch it at source. patch("litellm.Cache", return_value=mock_litellm_cache), ): litellm.cache = None ps.ProxyConfig()._init_cache(cache_params, enable_redis_auth_cache) - yield fresh_user_cache, fresh_spend_cache + yield fresh_user_cache, fresh_spend_cache, fresh_cli_sso_cache # --------------------------------------------------------------------------- @@ -90,7 +92,7 @@ class TestRedisAuthCacheFlag: with _patched_init_cache( litellm_settings={"enable_redis_auth_cache": True}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, - ) as (user_cache, _): + ) as (user_cache, _, _cli_sso_cache): assert user_cache.redis_cache is not None, ( "Redis should be attached to user_api_key_cache when " "enable_redis_auth_cache=True" @@ -101,7 +103,7 @@ class TestRedisAuthCacheFlag: with _patched_init_cache( litellm_settings={"enable_redis_auth_cache": False}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, - ) as (user_cache, _): + ) as (user_cache, _, _cli_sso_cache): assert user_cache.redis_cache is None, ( "user_api_key_cache must remain in-memory-only when " "enable_redis_auth_cache=False" @@ -112,7 +114,7 @@ class TestRedisAuthCacheFlag: with _patched_init_cache( litellm_settings={}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, - ) as (user_cache, _): + ) as (user_cache, _, _cli_sso_cache): assert user_cache.redis_cache is None, ( "user_api_key_cache must remain in-memory-only when " "enable_redis_auth_cache is absent from litellm_settings" @@ -129,7 +131,7 @@ class TestRedisAuthCacheFlag: with _patched_init_cache( litellm_settings=ls, cache_params={"type": "redis", "host": "localhost", "port": 6379}, - ) as (_, spend_cache): + ) as (_, spend_cache, _cli_sso_cache): assert spend_cache.redis_cache is not None, ( f"spend_counter_cache must always get Redis " f"(enable_redis_auth_cache={flag_value!r})" @@ -140,6 +142,28 @@ class TestRedisAuthCacheFlag: with _patched_init_cache( litellm_settings={"enable_redis_auth_cache": False}, cache_params={"type": "redis", "host": "localhost", "port": 6379}, - ) as (user_cache, spend_cache): + ) as (user_cache, spend_cache, _cli_sso_cache): assert spend_cache.redis_cache is not None assert user_cache.redis_cache is None + + def test_cli_sso_session_cache_always_gets_redis_regardless_of_flag(self): + """ + cli_sso_session_cache must receive Redis regardless of the auth-cache + flag so that `lite login` works on multi-worker deployments without + enable_redis_auth_cache (regression for the CLI SSO "Invalid CLI login + session" bug) + """ + for flag_value in (True, False, None): + ls = ( + {"enable_redis_auth_cache": flag_value} + if flag_value is not None + else {} + ) + with _patched_init_cache( + litellm_settings=ls, + cache_params={"type": "redis", "host": "localhost", "port": 6379}, + ) as (_, _, cli_sso_cache): + assert cli_sso_cache.redis_cache is not None, ( + f"cli_sso_session_cache must always get Redis " + f"(enable_redis_auth_cache={flag_value!r})" + ) diff --git a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py index 69845ec59c2..85dbf70b452 100644 --- a/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py +++ b/tests/test_litellm/proxy/ui_crud_endpoints/test_proxy_setting_endpoints.py @@ -53,6 +53,7 @@ def mock_proxy_config(monkeypatch): # Add a counter to track save_config calls save_config_call_count = 0 + saved_env_updates: list = [] async def mock_save_config(new_config=None): nonlocal mock_config, save_config_call_count @@ -61,13 +62,22 @@ def mock_proxy_config(monkeypatch): mock_config = new_config return mock_config + async def mock_save_environment_variables(updates): + saved_env_updates.append(updates) + from litellm.proxy.proxy_server import proxy_config monkeypatch.setattr(proxy_config, "get_config", mock_get_config) monkeypatch.setattr(proxy_config, "save_config", mock_save_config) + monkeypatch.setattr(proxy_config, "save_environment_variables", mock_save_environment_variables) - # Return both the config and the call counter - return {"config": mock_config, "save_call_count": lambda: save_config_call_count} + # Return the config, the save_config call counter, and any env-var updates + # the endpoint routed through the dedicated save_environment_variables path + return { + "config": mock_config, + "save_call_count": lambda: save_config_call_count, + "env_updates": lambda: saved_env_updates, + } @pytest.fixture @@ -386,6 +396,146 @@ class TestProxySettingEndpoints: call_args = mock_prisma.db.litellm_ssoconfig.find_unique.call_args assert call_args.kwargs["where"]["id"] == "sso_config" + def _mock_sso_db_record(self, monkeypatch, sso_settings): + """Point /get/sso_settings at a stored SSO row (or None for no row).""" + from unittest.mock import AsyncMock, MagicMock + + mock_prisma = MagicMock() + if sso_settings is None: + mock_db_record = None + else: + mock_db_record = MagicMock() + mock_db_record.sso_settings = sso_settings + mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_db_record) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + + # The resolver decrypts stored values via decrypt_value_helper; make it an + # identity so the plaintext fixtures round-trip. + monkeypatch.setattr( + "litellm.proxy.config_resolvers.sso.decrypt_value_helper", + lambda value, key, exception_type="error", return_original_value=False: value, + ) + + def test_get_sso_settings_falls_back_to_process_env( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """ + Regression for LIT-4165. + + SSO configured purely as process env vars (helm/terraform, no UI writes) + logs users in successfully, because ui_sso.py resolves every setting from + os.environ. /get/sso_settings read only the sso_config table though, so + the Admin UI showed "not configured" for a working SSO deployment and hid + the Edit/Delete controls behind an empty-state placeholder. + """ + self._mock_sso_db_record(monkeypatch, None) + monkeypatch.setenv("GENERIC_CLIENT_ID", "env-client-id") + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "env-client-secret-value") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://idp.example.com/authorize") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://idp.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://idp.example.com/userinfo") + monkeypatch.setenv("GENERIC_SCOPE", "openid email profile groups") + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com") + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + values = response.json()["values"] + + # Every one of these was None before the fix, despite SSO working. + assert values["generic_client_id"] == "env-client-id" + assert values["generic_authorization_endpoint"] == "https://idp.example.com/authorize" + assert values["generic_token_endpoint"] == "https://idp.example.com/token" + assert values["generic_userinfo_endpoint"] == "https://idp.example.com/userinfo" + assert values["generic_scope"] == "openid email profile groups" + assert values["proxy_base_url"] == "https://gateway.example.com" + + # An env-sourced secret is masked exactly like a stored one. + assert values["generic_client_secret"] not in (None, "env-client-secret-value") + assert "*" in values["generic_client_secret"] + + def test_get_sso_settings_does_not_mutate_os_environ( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """A GET must not write os.environ. The legacy read path decrypted DB + values straight into the environment, so opening the settings page + repopulated env and masked any consumer that stopped reading it.""" + self._mock_sso_db_record(monkeypatch, {"generic_client_id": "db-only-id"}) + monkeypatch.delenv("GENERIC_CLIENT_ID", raising=False) + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + assert response.json()["values"]["generic_client_id"] == "db-only-id" + # The DB value must NOT have leaked into the process environment. + assert "GENERIC_CLIENT_ID" not in os.environ + + def test_get_sso_settings_prefers_stored_over_process_env( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """A stored value wins; only fields absent from the row fall back to env.""" + self._mock_sso_db_record(monkeypatch, {"generic_client_id": "stored-client-id"}) + monkeypatch.setenv("GENERIC_CLIENT_ID", "env-client-id") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://idp.example.com/token") + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + values = response.json()["values"] + assert values["generic_client_id"] == "stored-client-id" + assert values["generic_token_endpoint"] == "https://idp.example.com/token" + + def test_get_sso_settings_blank_stored_value_falls_back_to_process_env( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """ + Blank means absent. update_sso_settings clears the env var for a blank + field, so a blank row entry cannot describe a live setting; os.environ is + the effective config and is what the UI must report. + """ + self._mock_sso_db_record(monkeypatch, {"generic_client_id": " ", "generic_token_endpoint": ""}) + monkeypatch.setenv("GENERIC_CLIENT_ID", "env-client-id") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://idp.example.com/token") + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + values = response.json()["values"] + assert values["generic_client_id"] == "env-client-id" + assert values["generic_token_endpoint"] == "https://idp.example.com/token" + + def test_get_sso_settings_unset_everywhere_reports_source( + self, mock_proxy_config, mock_auth, monkeypatch + ): + """A field set in neither source is unset (or its effective default), + and provenance reports which.""" + self._mock_sso_db_record(monkeypatch, None) + for env_var in ( + "GENERIC_CLIENT_ID", + "GENERIC_CLIENT_SECRET", + "GENERIC_TOKEN_ENDPOINT", + "GENERIC_SCOPE", + "GOOGLE_CLIENT_ID", + "MICROSOFT_CLIENT_ID", + "PROXY_BASE_URL", + ): + monkeypatch.delenv(env_var, raising=False) + + response = client.get("/get/sso_settings") + + assert response.status_code == 200 + body = response.json() + values = body["values"] + provenance = body["provenance"] + assert values["generic_client_id"] is None + assert provenance["generic_client_id"] == "unset" + assert values["generic_client_secret"] is None + assert values["google_client_id"] is None + # generic_scope carries the same effective default the login path applies, + # so the settings page shows the scope logins would actually request. + assert values["generic_scope"] == "openid email profile" + assert provenance["generic_scope"] == "default" + def test_update_sso_settings(self, mock_proxy_config, mock_auth, monkeypatch): """Test updating the SSO settings to the dedicated database table""" import json @@ -840,11 +990,18 @@ class TestProxySettingEndpoints: assert data["status"] == "success" assert data["theme_config"]["logo_url"] == "https://example.com/new-logo.png" - # Verify config was updated - updated_config = mock_proxy_config["config"] - assert "UI_LOGO_PATH" in updated_config["environment_variables"] + # The logo path is applied to the live process immediately + assert os.environ["UI_LOGO_PATH"] == "https://example.com/new-logo.png" assert mock_proxy_config["save_call_count"]() == 1 + # env vars are persisted through the dedicated per-key path, and ONLY + # the two keys this endpoint owns are touched. The unrelated SSO env + # vars in the merged config are never snapshotted. + env_updates = mock_proxy_config["env_updates"]() + assert env_updates == [ + {"UI_LOGO_PATH": "https://example.com/new-logo.png", "LITELLM_FAVICON_URL": None} + ] + def test_update_ui_theme_settings_with_favicon( self, mock_proxy_config, mock_auth, monkeypatch ): @@ -869,13 +1026,15 @@ class TestProxySettingEndpoints: == "https://example.com/custom-favicon.ico" ) - updated_config = mock_proxy_config["config"] - assert "UI_LOGO_PATH" in updated_config["environment_variables"] - assert "LITELLM_FAVICON_URL" in updated_config["environment_variables"] - assert ( - updated_config["environment_variables"]["LITELLM_FAVICON_URL"] - == "https://example.com/custom-favicon.ico" - ) + assert os.environ["UI_LOGO_PATH"] == "https://example.com/new-logo.png" + assert os.environ["LITELLM_FAVICON_URL"] == "https://example.com/custom-favicon.ico" + # Only the two owned keys are persisted, both with their new values + assert mock_proxy_config["env_updates"]() == [ + { + "UI_LOGO_PATH": "https://example.com/new-logo.png", + "LITELLM_FAVICON_URL": "https://example.com/custom-favicon.ico", + } + ] def test_update_ui_theme_settings_clear_favicon( self, mock_proxy_config, mock_auth, monkeypatch @@ -925,6 +1084,88 @@ class TestProxySettingEndpoints: assert data["values"]["logo_url"] == "https://example.com/logo.png" assert data["values"]["favicon_url"] == "https://example.com/favicon.ico" + def test_get_ui_theme_settings_falls_back_to_process_env( + self, mock_proxy_config, monkeypatch + ): + """Branding supplied only as process env vars must surface in the read. + + A deployment that sets UI_LOGO_PATH / LITELLM_FAVICON_URL via IaC and + never touches the UI has no stored ui_theme_config, yet the branding is + live, so the settings page must reflect it rather than reading blank. + """ + monkeypatch.delenv("UI_LOGO_PATH", raising=False) + monkeypatch.delenv("LITELLM_FAVICON_URL", raising=False) + monkeypatch.setenv("UI_LOGO_PATH", "https://cdn.example.com/logo.png") + monkeypatch.setenv("LITELLM_FAVICON_URL", "https://cdn.example.com/favicon.ico") + + response = client.get("/get/ui_theme_settings") + + assert response.status_code == 200 + values = response.json()["values"] + assert values["logo_url"] == "https://cdn.example.com/logo.png" + assert values["favicon_url"] == "https://cdn.example.com/favicon.ico" + + def test_get_ui_theme_settings_stored_value_wins_over_env( + self, mock_auth, monkeypatch + ): + """A stored ui_theme_config field outranks the env var for that field. + + The env fallback only fills fields the stored config leaves blank, so the + UI-driven flow is unchanged while an unstored field still resolves. + """ + from litellm.proxy.proxy_server import proxy_config + + stored_config = { + "litellm_settings": { + "ui_theme_config": {"logo_url": "https://db.example.com/logo.png"} + } + } + + async def mock_get_config(): + return stored_config + + monkeypatch.setattr(proxy_config, "get_config", mock_get_config) + monkeypatch.setenv("UI_LOGO_PATH", "https://env.example.com/logo.png") + monkeypatch.setenv("LITELLM_FAVICON_URL", "https://env.example.com/favicon.ico") + + response = client.get("/get/ui_theme_settings") + + assert response.status_code == 200 + values = response.json()["values"] + assert values["logo_url"] == "https://db.example.com/logo.png" + assert values["favicon_url"] == "https://env.example.com/favicon.ico" + + def test_get_ui_theme_settings_reports_unset_when_absent_everywhere( + self, mock_proxy_config, monkeypatch + ): + """A field set in neither the stored config nor the env stays null.""" + monkeypatch.delenv("UI_LOGO_PATH", raising=False) + monkeypatch.delenv("LITELLM_FAVICON_URL", raising=False) + + response = client.get("/get/ui_theme_settings") + + assert response.status_code == 200 + values = response.json()["values"] + assert values["logo_url"] is None + assert values["favicon_url"] is None + + def test_get_ui_theme_settings_does_not_disclose_local_path_env_value( + self, mock_proxy_config, monkeypatch + ): + """This endpoint is public, so an env-configured local filesystem branding + path must never be surfaced to anonymous callers; only public http(s) URLs. + """ + monkeypatch.setenv("UI_LOGO_PATH", "/mnt/secret/internal/logo.png") + monkeypatch.setenv("LITELLM_FAVICON_URL", "file:///etc/favicon.ico") + + response = client.get("/get/ui_theme_settings") + + assert response.status_code == 200 + values = response.json()["values"] + # the local path / file scheme is withheld rather than disclosed + assert values["logo_url"] is None + assert values["favicon_url"] is None + def test_get_ui_settings(self, mock_auth, monkeypatch): """Test retrieving UI settings with allowlist sanitization""" from unittest.mock import AsyncMock, MagicMock @@ -1362,19 +1603,20 @@ class TestProxySettingEndpoints: monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) - # Mock the decryption method to return decrypted values - def mock_decrypt_and_set(environment_variables): - return { - "google_client_id": "decrypted_google_id", - "google_client_secret": "decrypted_google_secret", - "microsoft_client_id": "decrypted_microsoft_id", - "proxy_base_url": "https://decrypted.example.com", - } + # The resolver decrypts each stored value via decrypt_value_helper; map + # the ciphertext fixtures to their plaintext. + decrypted_by_ciphertext = { + "encrypted_google_id": "decrypted_google_id", + "encrypted_google_secret": "decrypted_google_secret", + "encrypted_microsoft_id": "decrypted_microsoft_id", + "encrypted_proxy_url": "https://decrypted.example.com", + } - from litellm.proxy.proxy_server import proxy_config + def mock_decrypt(value, key, exception_type="error", return_original_value=False): + return decrypted_by_ciphertext.get(value, value) monkeypatch.setattr( - proxy_config, "_decrypt_and_set_db_env_variables", mock_decrypt_and_set + "litellm.proxy.config_resolvers.sso.decrypt_value_helper", mock_decrypt ) response = client.get("/get/sso_settings") diff --git a/tests/test_litellm/repositories/test_repositories.py b/tests/test_litellm/repositories/test_repositories.py index af2eea823f4..6308faf8fc7 100644 --- a/tests/test_litellm/repositories/test_repositories.py +++ b/tests/test_litellm/repositories/test_repositories.py @@ -5,7 +5,7 @@ Tests for gateway repository layer. import json from datetime import datetime from typing import Any, Dict, List, Optional -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -499,6 +499,42 @@ class TestTeamRepository: assert team.team_id == "team-123" assert team.team_alias == "Engineering" + @pytest.mark.asyncio + @pytest.mark.parametrize( + "raw_value, expected_ids", + [ + ( + [ + {"user_id": "a", "role": "user"}, + {"user_id": "b", "role": "admin"}, + ], + ["a", "b"], + ), + (json.dumps([{"user_id": "a", "role": "user"}]), ["a"]), + ({}, []), + (None, []), + ], + ) + async def test_get_members_with_roles_locked(self, repo, raw_value, expected_ids): + tx = MagicMock() + tx.query_raw = AsyncMock(return_value=[{"members_with_roles": raw_value}]) + + members = await repo.get_members_with_roles_locked(tx, "team-1") + + assert [m.user_id for m in members] == expected_ids + sql = tx.query_raw.call_args.args[0] + assert "FOR UPDATE" in sql + assert tx.query_raw.call_args.args[1] == "team-1" + + @pytest.mark.asyncio + async def test_get_members_with_roles_locked_missing_row(self, repo): + tx = MagicMock() + tx.query_raw = AsyncMock(return_value=[]) + + members = await repo.get_members_with_roles_locked(tx, "missing") + + assert members == [] + @pytest.mark.asyncio async def test_create_team_all_fields(self, repo): team = await repo.create_team( diff --git a/tests/test_litellm/responses/test_responses_prompt_management.py b/tests/test_litellm/responses/test_responses_prompt_management.py index 84e98390268..7044d8384f8 100644 --- a/tests/test_litellm/responses/test_responses_prompt_management.py +++ b/tests/test_litellm/responses/test_responses_prompt_management.py @@ -14,13 +14,19 @@ Covers: """ import asyncio -from typing import List +from typing import List, cast from unittest.mock import AsyncMock, MagicMock, patch import pytest +from litellm.integrations.anthropic_cache_control_hook import ( + AnthropicCacheControlHook, +) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import ( + AllMessageValues, + ResponseInputParam, +) # --------------------------------------------------------------------------- # Helpers @@ -71,6 +77,56 @@ def _patch_responses_dispatch(): ] +def _make_cache_control_case() -> tuple[ + ResponseInputParam, + list[AllMessageValues], + dict[str, object], +]: + system_message = cast( + AllMessageValues, + {"role": "system", "content": "Analyze the request"}, + ) + assistant_message = cast( + AllMessageValues, + { + "type": "message", + "id": "msg_1", + "role": "assistant", + "status": "completed", + "content": [ + { + "type": "output_text", + "text": "The code has a bug", + "annotations": [], + } + ], + }, + ) + user_message = cast( + AllMessageValues, + {"role": "user", "content": "Check for security issues"}, + ) + reasoning_item = { + "type": "reasoning", + "id": "rs_1", + "summary": [], + "encrypted_content": "encrypted", + } + original_input = cast( + ResponseInputParam, + [system_message, reasoning_item, assistant_message, user_message], + ) + _, merged_messages, _ = AnthropicCacheControlHook().get_chat_completion_prompt( + model="azure/gpt-5-codex", + messages=[system_message, assistant_message, user_message], + non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]}, + prompt_id=None, + prompt_variables=None, + dynamic_callback_params={}, + ) + return original_input, merged_messages, reasoning_item + + # --------------------------------------------------------------------------- # Tests # --------------------------------------------------------------------------- @@ -256,6 +312,66 @@ class TestResponsesAPIPromptManagement: assert all(isinstance(m, dict) and "role" in m for m in passed_messages) assert len(passed_messages) == 1 + def test_cache_control_hook_preserves_reasoning_items(self): + original_input, merged_messages, reasoning_item = _make_cache_control_case() + logging_obj = _make_logging_obj( + merged_model="azure/gpt-5-codex", + merged_messages=merged_messages, + ) + + patches = _patch_responses_dispatch() + with patches[0], patches[1], patches[2], patches[3] as mock_handler: + import litellm + + litellm.responses( + input=original_input, + model="azure/gpt-5-codex", + litellm_logging_obj=logging_obj, + cache_control_injection_points=[{"location": "message", "role": "system"}], + ) + + sent_input = mock_handler.call_args.kwargs["input"] + assert [item.get("type") for item in sent_input] == [ + None, + "reasoning", + "message", + None, + ] + assert sent_input[0]["cache_control"] == {"type": "ephemeral"} + assert sent_input[1] == reasoning_item + assert sent_input[2]["id"] == "msg_1" + + def test_all_non_message_input_items_remain_unchanged(self): + reasoning_item = { + "type": "reasoning", + "id": "rs_1", + "summary": [], + "encrypted_content": "encrypted", + } + original_input = cast(ResponseInputParam, [reasoning_item]) + logging_obj = _make_logging_obj( + merged_model="openai/gpt-4o", + merged_messages=[ + cast( + AllMessageValues, + {"role": "system", "content": "Analyze the request"}, + ) + ], + ) + + patches = _patch_responses_dispatch() + with patches[0], patches[1], patches[2], patches[3] as mock_handler: + import litellm + + litellm.responses( + input=original_input, + model="gpt-4o", + prompt_id="all-non-message", + litellm_logging_obj=logging_obj, + ) + + assert mock_handler.call_args.kwargs["input"] == original_input + def test_model_override_re_resolves_provider(self): """[G] When the prompt template overrides the model to a different provider, custom_llm_provider is re-resolved so downstream routing uses the correct provider. @@ -393,3 +509,33 @@ class TestAsyncResponsesAPIPromptManagement: passed_messages = call_kwargs["messages"] assert all(isinstance(m, dict) and "role" in m for m in passed_messages) assert len(passed_messages) == 1 + + @pytest.mark.asyncio + async def test_async_cache_control_hook_preserves_reasoning_items(self): + original_input, merged_messages, reasoning_item = _make_cache_control_case() + logging_obj = _make_logging_obj( + merged_model="azure/gpt-5-codex", + merged_messages=merged_messages, + ) + + patches = _patch_responses_dispatch() + with patches[0], patches[1], patches[2], patches[3] as mock_handler: + import litellm + + await litellm.aresponses( + input=original_input, + model="azure/gpt-5-codex", + litellm_logging_obj=logging_obj, + cache_control_injection_points=[{"location": "message", "role": "system"}], + ) + + sent_input = mock_handler.call_args.kwargs["input"] + assert [item.get("type") for item in sent_input] == [ + None, + "reasoning", + "message", + None, + ] + assert sent_input[0]["cache_control"] == {"type": "ephemeral"} + assert sent_input[1] == reasoning_item + assert sent_input[2]["id"] == "msg_1" diff --git a/tests/test_litellm/test_router_per_deployment_num_retries.py b/tests/test_litellm/test_router_per_deployment_num_retries.py index af2372616a6..25574fcb268 100644 --- a/tests/test_litellm/test_router_per_deployment_num_retries.py +++ b/tests/test_litellm/test_router_per_deployment_num_retries.py @@ -3,11 +3,15 @@ Unit tests for per-deployment num_retries in litellm_params GitHub Issue: #18968 - Per-deployment max_retries/num_retries in litellm_params is not used in retry logic """ +import httpx import pytest +import pytest_asyncio from unittest.mock import patch import litellm from litellm import Router +from litellm.types.router import RetryPolicy +from litellm.integrations.custom_logger import CustomLogger class TestPerDeploymentNumRetries: @@ -319,3 +323,255 @@ class TestNumRetriesNoneGuard: # 1 initial attempt + at least 1 retry -> proves None fell back to a positive int assert calls["n"] >= 2 + + +class TestNoProviderRetryAmplification: + """ + A routed request must reach the upstream provider exactly ``1 + `` + times. The Router is the sole retry owner for routed calls, so the provider SDK + must never retry on top of it. Otherwise a per-deployment ``num_retries`` set in + ``litellm_params`` is applied twice - once by the Router loop and once as the + provider client's ``max_retries`` - turning one request into ``(1 + num_retries) ** 2`` + upstream requests. + + These tests count actual upstream HTTP requests through the full Router completion + path by injecting a counting transport via ``litellm.aclient_session`` (the + documented seam the OpenAI client builder reads), so both Router-level and any + provider-SDK-level retries are observed. + """ + + @staticmethod + def _install_counting_upstream() -> dict: + """Route every upstream POST to a 500 and count it. ``retry-after: 0`` keeps + provider-SDK backoff at zero so a mutated (double-retrying) build stays fast.""" + counter = {"n": 0} + + def handler(request: httpx.Request) -> httpx.Response: + counter["n"] += 1 + return httpx.Response( + 500, + headers={"retry-after": "0"}, + json={"error": {"message": "boom", "type": "server_error"}}, + ) + + litellm.aclient_session = httpx.AsyncClient(transport=httpx.MockTransport(handler)) + return counter + + @pytest_asyncio.fixture(autouse=True) + async def _isolate_clients(self): + litellm.in_memory_llm_clients_cache.flush_cache() + yield + session = litellm.aclient_session + litellm.aclient_session = None + litellm.in_memory_llm_clients_cache.flush_cache() + if session is not None: + await session.aclose() + + @staticmethod + def _router(api_base: str, litellm_params: dict, **router_kwargs) -> Router: + params = {"model": "openai/gpt-4o-mini", "api_base": api_base, "api_key": "sk-fake"} + params.update(litellm_params) + return Router(model_list=[{"model_name": "mock", "litellm_params": params}], **router_kwargs) + + async def _call_and_count(self, router: Router, **call_kwargs) -> int: + counter = self._install_counting_upstream() + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.InternalServerError): + await router.acompletion( + model="mock", messages=[{"role": "user", "content": "hi"}], **call_kwargs + ) + return counter["n"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("num_retries", [2, 5]) + async def test_deployment_num_retries_sends_no_extra_provider_requests(self, num_retries): + """ + Deployment ``num_retries=N`` (every attempt failing) must send exactly ``N + 1`` + upstream requests, not ``(N + 1) ** 2``. This is the amplification regression: + an unfixed build sends 9 (N=2) or 36 (N=5). + """ + counter = self._install_counting_upstream() + router = self._router( + f"https://amp-{num_retries}.local/v1", {"num_retries": num_retries}, num_retries=1 + ) + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.InternalServerError): + await router.acompletion(model="mock", messages=[{"role": "user", "content": "hi"}]) + assert counter["n"] == num_retries + 1 + + @pytest.mark.asyncio + async def test_request_max_retries_does_not_nest_with_router_retries(self): + """ + A request-body ``max_retries`` must not make the provider SDK retry on top of the + Router. With deployment ``num_retries=5`` and request ``max_retries=3`` the count + stays ``6``; a build that lets either value reach the provider SDK sends 24 or 36. + """ + router = self._router("https://nest-req.local/v1", {"num_retries": 5}, num_retries=1) + assert await self._call_and_count(router, max_retries=3) == 6 + + @pytest.mark.asyncio + async def test_deployment_max_retries_does_not_nest_with_router_retries(self): + """ + A deployment-level ``max_retries`` is likewise never applied on top of the Router's + retries for a routed call: deployment ``num_retries=5`` plus ``max_retries=3`` still + sends exactly ``6`` upstream requests. + """ + router = self._router( + "https://nest-dep.local/v1", {"num_retries": 5, "max_retries": 3}, num_retries=1 + ) + assert await self._call_and_count(router) == 6 + + @pytest.mark.asyncio + async def test_retry_policy_configured_does_not_reintroduce_amplification(self): + """ + With a retry policy configured alongside a per-deployment ``num_retries=5``, the + provider SDK still must not retry: exactly ``6`` upstream requests, not 36. + """ + router = self._router( + "https://policy.local/v1", + {"num_retries": 5}, + num_retries=1, + retry_policy=RetryPolicy(InternalServerErrorRetries=2), + ) + assert await self._call_and_count(router) == 6 + + @pytest.mark.asyncio + async def test_global_num_retries_not_amplified(self): + """ + Global ``num_retries`` (no per-deployment setting) already behaves correctly and + must stay that way: ``num_retries=3`` sends ``4`` upstream requests. + """ + router = self._router("https://global.local/v1", {}, num_retries=3) + assert await self._call_and_count(router) == 4 + + @pytest.mark.asyncio + async def test_direct_completion_still_forwards_num_retries_to_provider(self): + """ + For a NON-routed direct ``litellm.acompletion`` call, ``num_retries`` remains an + alias for the provider client's ``max_retries`` (the instructor use case). The + provider SDK therefore retries in addition to litellm's own retry wrapper, so the + upstream count exceeds ``num_retries + 1`` - proving the routed-call fix did not + change direct-call behaviour. + """ + counter = self._install_counting_upstream() + num_retries = 2 + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.InternalServerError): + await litellm.acompletion( + model="openai/gpt-4o-mini", + api_base="https://direct.local/v1", + api_key="sk-fake", + messages=[{"role": "user", "content": "hi"}], + num_retries=num_retries, + ) + assert counter["n"] > num_retries + 1 + + +class _AttemptCounter(CustomLogger): + """Counts upstream call attempts via the pre-call hook (one per attempt).""" + + def __init__(self): + self.attempts = 0 + + def log_pre_api_call(self, model, messages, kwargs): + self.attempts += 1 + + +class TestRequestNumRetriesBeatsGlobal: + """ + A per-request num_retries (request body or the x-litellm-num-retries header, both of + which arrive as the num_retries kwarg) must take precedence over the global + litellm.num_retries (litellm_settings.num_retries on the proxy) during retry handling. + + The regression: the @client wrapper stamped the global litellm.num_retries onto the + raised exception, and async_function_with_retries then adopted that stamped value, + overwriting the request-level num_retries it had already resolved. This exercises the + real retry loop end to end (the failing call flows through the wrapped litellm.acompletion), + which the kwargs-merge-only test above does not. + """ + + @pytest.fixture(autouse=True) + def _restore_litellm_globals(self): + prev_num_retries = litellm.num_retries + prev_callbacks = litellm.callbacks + yield + litellm.num_retries = prev_num_retries + litellm.callbacks = prev_callbacks + + @staticmethod + def _router(global_num_retries): + return Router( + model_list=[ + { + "model_name": "mock", + "litellm_params": { + "model": "openai/mock", + "api_key": "sk-fake", + "mock_response": "litellm.InternalServerError", + }, + } + ], + num_retries=global_num_retries, + ) + + async def _count_attempts(self, *, global_num_retries, request_num_retries): + counter = _AttemptCounter() + litellm.callbacks = [counter] + litellm.num_retries = global_num_retries + router = self._router(global_num_retries) + kwargs = {"model": "mock", "messages": [{"role": "user", "content": "hi"}]} + if request_num_retries is not None: + kwargs["num_retries"] = request_num_retries + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.InternalServerError): + await router.acompletion(**kwargs) + return counter.attempts + + @pytest.mark.asyncio + async def test_request_num_retries_overrides_global(self): + """global=3 + request=1 -> 2 attempts (1 initial + 1 retry), not 4 (1 + global 3).""" + attempts = await self._count_attempts(global_num_retries=3, request_num_retries=1) + assert attempts == 2 + + @pytest.mark.asyncio + async def test_request_num_retries_zero_disables_retries_despite_global(self): + """global=3 + request=0 -> a single attempt (retries disabled by the request).""" + attempts = await self._count_attempts(global_num_retries=3, request_num_retries=0) + assert attempts == 1 + + @pytest.mark.asyncio + async def test_global_num_retries_applies_when_request_omits_it(self): + """No request num_retries -> the global still applies: 1 initial + 3 retries = 4.""" + attempts = await self._count_attempts(global_num_retries=3, request_num_retries=None) + assert attempts == 4 + + @pytest.mark.asyncio + async def test_deployment_num_retries_reaches_wrapper_when_no_request_value(self): + """ + With no request value and the router default at 0, a deployment's + litellm_params.num_retries reaches the wrapped call, is carried on the raised + exception, and is applied: deployment 2 -> 1 initial + 2 retries = 3 (not 1). + """ + counter = _AttemptCounter() + litellm.callbacks = [counter] + litellm.num_retries = None + router = Router( + model_list=[ + { + "model_name": "mock", + "litellm_params": { + "model": "openai/mock", + "api_key": "sk-fake", + "mock_response": "litellm.InternalServerError", + "num_retries": 2, + }, + } + ], + num_retries=0, + ) + with patch("asyncio.sleep", return_value=None): + with pytest.raises(litellm.InternalServerError): + await router.acompletion( + model="mock", messages=[{"role": "user", "content": "hi"}] + ) + assert counter.attempts == 3 diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index f935af8907d..12f19eaa1fe 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -4,9 +4,34 @@ "count": 1 } }, - "src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx": { + "src/app/(dashboard)/access-groups/_components/AccessGroupsDetailsPage.tsx": { "no-restricted-imports": { "count": 1 + } + }, + "src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupBaseForm.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupCreateModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/access-groups/_components/AccessGroupsModal/AccessGroupEditModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/access-groups/_components/AccessGroupsPage.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx": { + "no-restricted-imports": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -15,17 +40,25 @@ "src/app/(dashboard)/agents/_components/AgentsPanel.tsx": { "no-restricted-imports": { "count": 1 - }, - "react-hooks/set-state-in-effect": { + } + }, + "src/app/(dashboard)/agents/_components/AgentsTable.tsx": { + "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/agents/_components/add_agent_form.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 3 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 2 @@ -35,6 +68,12 @@ } }, "src/app/(dashboard)/agents/_components/agent_card_discovery.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/refs": { "count": 3 }, @@ -43,44 +82,77 @@ } }, "src/app/(dashboard)/agents/_components/agent_cost_view.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/agents/_components/agent_form_fields.tsx": { - "no-nested-ternary": { + "local/filename-pascal-case": { "count": 1 - } - }, - "src/app/(dashboard)/agents/_components/agent_info.tsx": { + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/agents/_components/agent_info.tsx": { + "local/filename-pascal-case": { "count": 1 }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/immutability": { "count": 1 } }, "src/app/(dashboard)/agents/_components/agent_virtual_keys.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/agents/_components/cost_config_fields.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/app/(dashboard)/agents/_components/dynamic_agent_form_fields.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 2 - } - }, - "src/app/(dashboard)/api-reference/_components/APIReferenceView.tsx": { + }, "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/budgets/_components/budget_modal.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/budgets/_components/budget_panel.test.tsx": { @@ -89,19 +161,31 @@ } }, "src/app/(dashboard)/budgets/_components/budget_panel.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/budgets/_components/edit_budget_modal.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/caching/_components/cache_dashboard.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, + "prefer-const": { + "count": 3 + }, "react-hooks/purity": { "count": 1 }, @@ -110,6 +194,14 @@ } }, "src/app/(dashboard)/caching/_components/cache_health.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/caching/_components/cache_settings/CacheFormField.tsx": { "no-restricted-imports": { "count": 1 } @@ -119,33 +211,109 @@ "count": 1 } }, - "src/app/(dashboard)/caching/_components/cache_settings/index.tsx": { + "src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsFields.ts": { "no-restricted-imports": { "count": 1 + } + }, + "src/app/(dashboard)/caching/_components/cache_settings/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx": { + "src/app/(dashboard)/caching/_components/coordination_redis_settings/CoordinationRedisFormField.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx": { + "src/app/(dashboard)/caching/_components/coordination_redis_settings/CoordinationRedisTypeSelector.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.tsx": { - "no-nested-ternary": { - "count": 2 + "src/app/(dashboard)/caching/_components/coordination_redis_settings/coordinationRedisFields.ts": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/caching/_components/coordination_redis_settings/index.tsx": { + "local/filename-pascal-case": { + "count": 1 }, "no-restricted-imports": { "count": 1 } }, + "src/app/(dashboard)/caching/_components/response_time_indicator.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/app/(dashboard)/cost-optimization/_components/AutorouterTab.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/cost-optimization/_components/PromptCompressionTab.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-nested-ternary": { + "count": 2 + }, + "no-restricted-imports": { + "count": 2 + } + }, "src/app/(dashboard)/cost-tracking/_components/how_it_works.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/cost-tracking/_components/pricing_calculator/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -156,8 +324,11 @@ } }, "src/app/(dashboard)/cost-tracking/_components/pricing_calculator/multi_cost_results.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/cost-tracking/_components/pricing_calculator/multi_export_dropdown.test.tsx": { @@ -166,6 +337,9 @@ } }, "src/app/(dashboard)/cost-tracking/_components/pricing_calculator/multi_export_dropdown.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -181,16 +355,17 @@ } }, "src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, - "src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.test.ts": { - "unused-imports/no-unused-imports": { - "count": 1 - } - }, "src/app/(dashboard)/cost-tracking/_components/provider_margin_table.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -206,13 +381,24 @@ } }, "src/app/(dashboard)/guardrails-monitor/_components/EvaluationSettingsModal.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "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 + }, + "no-restricted-imports": { + "count": 1 } }, "src/app/(dashboard)/guardrails-monitor/_components/GuardrailsMonitorView.tsx": { @@ -223,37 +409,64 @@ "src/app/(dashboard)/guardrails-monitor/_components/GuardrailsOverview.tsx": { "no-nested-ternary": { "count": 8 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/guardrails/_components/GuardrailTestPanel.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/guardrails/_components/GuardrailTestPlayground.tsx": { "no-nested-ternary": { "count": 1 - } - }, - "src/app/(dashboard)/guardrails/_components/GuardrailTestResults.tsx": { + }, "no-restricted-imports": { "count": 1 } }, + "src/app/(dashboard)/guardrails/_components/GuardrailTestResults.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, "src/app/(dashboard)/guardrails/_components/GuardrailsPanel.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/guardrails/_components/TeamGuardrailsTab.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 2 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 }, @@ -261,22 +474,44 @@ "count": 2 } }, + "src/app/(dashboard)/guardrails/_components/content_filter/CategoryTable.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/guardrails/_components/content_filter/CompetitorIntentConfiguration.tsx": { "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/guardrails/_components/content_filter/ContentCategoryConfiguration.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 3 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/guardrails/_components/content_filter/ContentFilterConfiguration.tsx": { + "local/no-complex-jsx-arrow": { + "count": 3 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/guardrails/_components/content_filter/ContentFilterDisplay.tsx": { "no-restricted-imports": { "count": 1 @@ -286,51 +521,157 @@ "max-params": { "count": 2 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/guardrails/_components/content_filter/CustomPatternModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/content_filter/KeywordModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/content_filter/KeywordTable.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/content_filter/PatternModal.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/content_filter/PatternTable.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/guardrails/_components/custom_code/CustomCodeModal.tsx": { "no-nested-ternary": { "count": 6 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/app/(dashboard)/guardrails/_components/guardrail_info.tsx": { - "max-params": { + "src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/guardrail_garden.tsx": { + "local/filename-pascal-case": { "count": 1 }, "no-restricted-imports": { "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/guardrail_info.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "max-params": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 3 } }, + "src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/app/(dashboard)/guardrails/_components/guardrail_optional_params.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 5 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/guardrails/_components/guardrail_provider_fields.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 5 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/app/(dashboard)/guardrails/_components/tool_permission/ToolPermissionRulesEditor.tsx": { + "src/app/(dashboard)/guardrails/_components/guardrail_table.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/llm_judge/LLMJudgeFields.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/guardrails/_components/pii_components.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/pii_configuration.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/guardrails/_components/tool_permission/ToolPermissionRulesEditor.tsx": { + "no-restricted-imports": { + "count": 2 }, "react-hooks/purity": { "count": 1 @@ -496,16 +837,54 @@ "count": 2 } }, + "src/app/(dashboard)/hooks/useTeams.tsx": { + "local/filename-pascal-case": { + "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/MCPLogoSelector.test.tsx": { "unused-imports/no-unused-imports": { "count": 1 } }, + "src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/mcp-servers/_components/MCPNetworkSettings.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/immutability": { "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/MCPServerCard.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, "src/app/(dashboard)/mcp-servers/_components/MCPSubmissionsTab.tsx": { "react-hooks/set-state-in-effect": { "count": 1 @@ -516,89 +895,176 @@ "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "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": 2 + } + }, + "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": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "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": { "no-nested-ternary": { "count": 3 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/mcp-servers/_components/TruePassthroughWarning.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/mcp-servers/_components/UserEnvVarsModal.tsx": { "no-nested-ternary": { "count": 2 + }, + "no-restricted-imports": { + "count": 1 } }, "src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 4 } }, - "src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx": { - "no-restricted-imports": { + "src/app/(dashboard)/mcp-servers/_components/index.tsx": { + "local/filename-pascal-case": { "count": 1 + } + }, + "src/app/(dashboard)/mcp-servers/_components/mcp_connect.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 }, "react-hooks/static-components": { "count": 4 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_connection_status.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 3 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_discovery.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 2 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_config.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_server_cost_display.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_server_edit.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/immutability": { "count": 1 @@ -608,48 +1074,73 @@ } }, "src/app/(dashboard)/mcp-servers/_components/mcp_server_view.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 3 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 2 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_tool_configuration.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 2 } }, - "src/app/(dashboard)/memory/_components/MemoryView.tsx": { - "react-hooks/set-state-in-effect": { + "src/app/(dashboard)/mcp-servers/_components/utils.tsx": { + "local/filename-pascal-case": { "count": 1 } }, - "src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx": { + "src/app/(dashboard)/memory/_components/MemoryDetailDrawer.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/memory/_components/MemoryEditModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/memory/_components/MemoryView.tsx": { "no-restricted-imports": { "count": 1 - }, - "react-hooks/preserve-manual-memoization": { - "count": 4 } }, "src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.test.tsx": { @@ -661,8 +1152,11 @@ } }, "src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 3 @@ -678,7 +1172,7 @@ "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx": { @@ -686,9 +1180,20 @@ "count": 1 } }, + "src/app/(dashboard)/models-and-endpoints/utils/modelDataTransformer.ts": { + "prefer-const": { + "count": 6 + } + }, "src/app/(dashboard)/old-usage/_components/usage.tsx": { - "no-restricted-imports": { - "count": 2 + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, + "prefer-const": { + "count": 6 }, "react-hooks/immutability": { "count": 1 @@ -697,14 +1202,19 @@ "count": 1 } }, - "src/app/(dashboard)/organizations/_components/organizations.tsx": { + "src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/playground/components/chat_ui/A2AMetrics.tsx": { "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/playground/components/chat_ui/AdditionalModelSettings.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 2 @@ -714,10 +1224,18 @@ "no-nested-ternary": { "count": 2 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 5 } }, + "src/app/(dashboard)/playground/components/chat_ui/ChatImageUpload.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/playground/components/chat_ui/ChatImageUtils.test.tsx": { "max-nested-callbacks": { "count": 1 @@ -729,10 +1247,19 @@ } }, "src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 7 }, "no-restricted-imports": { + "count": 2 + }, + "prefer-const": { "count": 1 }, "react-hooks/set-state-in-effect": { @@ -746,11 +1273,19 @@ "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 1 + }, "no-restricted-syntax": { "count": 2 } }, "src/app/(dashboard)/playground/components/chat_ui/CodeInterpreterTool.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/playground/components/chat_ui/EndpointSelector.tsx": { "no-restricted-imports": { "count": 1 } @@ -759,6 +1294,9 @@ "no-nested-ternary": { "count": 2 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/immutability": { "count": 2 }, @@ -766,25 +1304,70 @@ "count": 1 } }, + "src/app/(dashboard)/playground/components/chat_ui/ResponsesImageUpload.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/playground/components/chat_ui/SearchResultsDisplay.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/playground/components/chat_ui/SessionManagement.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/playground/components/compareUI/CompareUI.tsx": { + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 4 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/playground/components/compareUI/components/ComparisonPanel.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/playground/components/compareUI/components/MessageDisplay.tsx": { "no-nested-ternary": { "count": 1 } }, + "src/app/(dashboard)/playground/components/compareUI/components/MessageInput.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/playground/components/compareUI/components/ModelSelector.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/playground/components/compareUI/components/UnifiedSelector.tsx": { "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/playground/components/complianceUI/ComplianceUI.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 8 }, @@ -793,6 +1376,9 @@ } }, "src/app/(dashboard)/playground/llm_calls/a2a_send_message.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 2 }, @@ -801,21 +1387,33 @@ } }, "src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, "src/app/(dashboard)/playground/llm_calls/audio_speech.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, "src/app/(dashboard)/playground/llm_calls/audio_transcriptions.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, "src/app/(dashboard)/playground/llm_calls/embeddings_api.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 }, @@ -824,21 +1422,33 @@ } }, "src/app/(dashboard)/playground/llm_calls/fetch_agents.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-syntax": { "count": 1 } }, "src/app/(dashboard)/playground/llm_calls/image_edits.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, "src/app/(dashboard)/playground/llm_calls/image_generation.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, "src/app/(dashboard)/playground/llm_calls/interactions_api.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 }, @@ -852,17 +1462,26 @@ } }, "src/app/(dashboard)/policies/_components/add_attachment_form.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/immutability": { "count": 1 } }, "src/app/(dashboard)/policies/_components/add_policy_form.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, + "prefer-const": { + "count": 2 + }, "react-hooks/immutability": { "count": 2 }, @@ -871,20 +1490,35 @@ } }, "src/app/(dashboard)/policies/_components/ai_suggestion_modal.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 3 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 10 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/immutability": { "count": 1 } }, "src/app/(dashboard)/policies/_components/guardrail_selection_modal.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -895,9 +1529,20 @@ } }, "src/app/(dashboard)/policies/_components/impact_popover.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/policies/_components/impact_preview_alert.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -908,44 +1553,73 @@ } }, "src/app/(dashboard)/policies/_components/index.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/policies/_components/pipeline_flow_builder.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 2 } }, "src/app/(dashboard)/policies/_components/policy_info.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/app/(dashboard)/policies/_components/policy_test_panel.tsx": { + "src/app/(dashboard)/policies/_components/policy_templates.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 + } + }, + "src/app/(dashboard)/policies/_components/policy_test_panel.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 }, "react-hooks/immutability": { "count": 1 } }, "src/app/(dashboard)/policies/_components/template_parameter_modal.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/immutability": { "count": 1 }, @@ -956,34 +1630,74 @@ "src/app/(dashboard)/projects/_components/ProjectDetailsPage.tsx": { "no-nested-ternary": { "count": 3 + }, + "no-restricted-imports": { + "count": 1 } }, "src/app/(dashboard)/projects/_components/ProjectKeysSection.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/projects/_components/ProjectModals/CreateProjectModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/projects/_components/ProjectModals/EditProjectModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/projects/_components/ProjectModals/ProjectBaseForm.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/projects/_components/ProjectModals/ProjectBaseForm.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 2 } }, - "src/app/(dashboard)/prompts/_components/add_prompt_form.tsx": { + "src/app/(dashboard)/projects/_components/ProjectsPage.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/app/(dashboard)/prompts/_components/index.tsx": { - "no-nested-ternary": { + "src/app/(dashboard)/prompts/_components/add_prompt_form.tsx": { + "local/filename-pascal-case": { "count": 1 }, "no-restricted-imports": { + "count": 3 + } + }, + "src/app/(dashboard)/prompts/_components/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-nested-ternary": { "count": 1 }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/app/(dashboard)/prompts/_components/prompt_editor_view.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/DeveloperMessageCard.tsx": { "no-restricted-imports": { "count": 1 @@ -991,7 +1705,7 @@ }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/ModelConfigCard.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/PromptCodeSnippets.tsx": { @@ -999,7 +1713,7 @@ "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -1007,17 +1721,17 @@ }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/PromptEditorHeader.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/PromptMessagesCard.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/PublishModal.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/ToolsCard.tsx": { @@ -1031,19 +1745,38 @@ } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/VersionHistorySidePanel.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/immutability": { "count": 1 } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/conversation_panel/MessageInput.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/prompts/_components/prompt_editor_view/conversation_panel/MessageList.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/prompts/_components/prompt_editor_view/conversation_panel/VariableInput.tsx": { "no-restricted-imports": { "count": 1 } }, "src/app/(dashboard)/prompts/_components/prompt_editor_view/conversation_panel/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } @@ -1053,70 +1786,117 @@ "count": 1 } }, + "src/app/(dashboard)/prompts/_components/prompt_editor_view/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/app/(dashboard)/prompts/_components/prompt_info.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 3 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 2 } }, + "src/app/(dashboard)/prompts/_components/prompt_utils.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/app/(dashboard)/prompts/_components/tool_modal.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/prompts/_components/variable_textarea.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/router-settings/_components/general_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { + "count": 3 + }, + "prefer-const": { "count": 2 } }, "src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/app/(dashboard)/search-tools/_components/SearchToolTester.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/search-tools/_components/SearchToolView.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/search-tools/_components/SearchTools.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/static-components": { "count": 1 } }, + "src/app/(dashboard)/search-tools/_components/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/app/(dashboard)/search-tools/_components/types.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/app/(dashboard)/skills/_components/ClaudeCodePluginsPanel.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/skills/_components/add_plugin_form.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/tag-management/_components/components/CreateTagModal.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/tag-management/_components/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1124,16 +1904,19 @@ "count": 1 } }, + "src/app/(dashboard)/tag-management/_components/tagTableColumns.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/app/(dashboard)/tag-management/_components/tag_info.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, - "src/app/(dashboard)/transform-request/TransformRequestPanel.tsx": { "no-restricted-imports": { + "count": 3 + }, + "react-hooks/set-state-in-effect": { "count": 1 } }, @@ -1151,14 +1934,25 @@ "src/app/(dashboard)/usage/_components/components/EndpointUsage/components/EndpointUsageTable.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/usage/_components/components/EntityUsage/SpendByProvider.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/usage/_components/components/EntityUsage/TopModelView.tsx": { "no-restricted-imports": { "count": 1 } @@ -1167,16 +1961,25 @@ "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/immutability": { "count": 1 } }, "src/app/(dashboard)/usage/_components/components/UsagePageView.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/purity": { "count": 1 @@ -1185,6 +1988,14 @@ "count": 3 } }, + "src/app/(dashboard)/usage/_components/components/UsageViewSelect/UsageViewSelect.tsx": { + "local/no-complex-jsx-arrow": { + "count": 2 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity.ts": { "react-hooks/refs": { "count": 1 @@ -1193,13 +2004,29 @@ "count": 1 } }, - "src/app/(dashboard)/users/_components/DefaultUserSettings.tsx": { + "src/app/(dashboard)/users/_components/BulkEditUsers.tsx": { "no-restricted-imports": { "count": 1 + }, + "prefer-const": { + "count": 1 } }, - "src/app/(dashboard)/users/_components/edit_user.tsx": { + "src/app/(dashboard)/users/_components/DefaultUserSettings.tsx": { "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/users/_components/edit_user.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/app/(dashboard)/users/_components/index.tsx": { + "local/filename-pascal-case": { "count": 1 } }, @@ -1212,49 +2039,55 @@ } }, "src/app/(dashboard)/users/_components/user_edit_view.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/users/_components/view_users.tsx": { - "no-nested-ternary": { + "local/filename-pascal-case": { "count": 1 }, "no-restricted-imports": { + "count": 3 + }, + "prefer-const": { "count": 1 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/app/(dashboard)/users/_components/view_users/columns.tsx": { - "max-params": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, - "src/app/(dashboard)/users/_components/view_users/table.tsx": { - "no-nested-ternary": { - "count": 1 - }, - "no-restricted-imports": { - "count": 1 - } - }, "src/app/(dashboard)/users/_components/view_users/user_info_view.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/app/(dashboard)/vector-stores/_components/CreateVectorStore.tsx": { + "no-restricted-imports": { + "count": 3 + } + }, + "src/app/(dashboard)/vector-stores/_components/S3VectorsConfig.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/(dashboard)/vector-stores/_components/TestVectorStoreTab.tsx": { "no-restricted-imports": { "count": 1 } @@ -1264,13 +2097,21 @@ "count": 2 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react/no-unescaped-entities": { "count": 1 } }, + "src/app/(dashboard)/vector-stores/_components/VectorStoreTester.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/app/(dashboard)/vector-stores/_components/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -1279,9 +2120,12 @@ } }, "src/app/(dashboard)/vector-stores/_components/vector_store_info.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -1290,6 +2134,9 @@ "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 1 + }, "no-restricted-syntax": { "count": 3 }, @@ -1303,6 +2150,12 @@ } }, "src/app/login/LoginPage.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 2 } @@ -1317,15 +2170,41 @@ "count": 1 } }, + "src/app/onboarding/OnboardingErrorView.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/onboarding/OnboardingFormBody.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/app/onboarding/OnboardingLoadingView.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/AIHub/ModelHubTable.test.tsx": { "max-params": { "count": 1 } }, "src/components/AIHub/ModelHubTable.tsx": { + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, + "prefer-const": { + "count": 4 + } + }, + "src/components/AIHub/SkillHubDashboard.tsx": { "no-restricted-imports": { "count": 1 } @@ -1340,7 +2219,7 @@ }, "src/components/AIHub/forms/MakeAgentPublicForm.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -1356,7 +2235,7 @@ "count": 2 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -1364,28 +2243,93 @@ }, "src/components/AIHub/forms/MakeModelPublicForm.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/components/CreateUserButton.tsx": { + "src/components/BetaBadge.tsx": { "no-restricted-imports": { "count": 1 + } + }, + "src/components/CloudZeroCostTracking/CloudZeroCostTracking.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/CloudZeroCostTracking/CloudZeroCreateModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/CloudZeroCostTracking/CloudZeroEmptyPlaceholder.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/CloudZeroCostTracking/CloudZeroIntegrationSettings.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/CloudZeroCostTracking/CloudZeroUpdateModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/CreateUserButton.tsx": { + "no-restricted-imports": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/DebugWarningBanner.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/DeletedKeysPage/DeletedKeysPage.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/DeletedTeamsPage/DeletedTeamsPage.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/DeprecationBanner.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/EntityUsageExport/EntityUsageExportModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/EntityUsageExport/ExportFormatSelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/EntityUsageExport/ExportSummary.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/EntityUsageExport/ExportTypeSelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/EntityUsageExport/UsageExportHeader.tsx": { "no-restricted-imports": { - "count": 2 + "count": 3 } }, "src/components/EntityUsageExport/types.ts": { @@ -1409,11 +2353,17 @@ "src/components/GuardrailSettingsView.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/GuardrailsMonitor/LogViewer.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/HelpLink.test.tsx": { @@ -1421,9 +2371,27 @@ "count": 1 } }, + "src/components/KeyAliasSelect/PaginatedKeyAliasSelect/PaginatedKeyAliasSelect.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/LicenseExpiryBanner.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/ModelSelect/ModelSelect.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/ModelSelect/PaginatedModelSelect/PaginatedModelSelect.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/Navbar/BlogDropdown/BlogDropdown.test.tsx": { @@ -1431,20 +2399,58 @@ "count": 12 } }, - "src/components/Navbar/UserDropdown/UserDropdown.tsx": { - "react-hooks/set-state-in-effect": { + "src/components/Navbar/BlogDropdown/BlogDropdown.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/Navbar/CommunityEngagementButtons/CommunityEngagementButtons.tsx": { + "no-restricted-imports": { "count": 1 } }, - "src/components/SCIM.tsx": { + "src/components/Navbar/NotificationsBell/NotificationsBell.tsx": { "no-restricted-imports": { "count": 1 + } + }, + "src/components/Navbar/UserDropdown/UserDropdown.tsx": { + "no-restricted-imports": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/Navbar/ViewSwitcher.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/Navbar/WorkerDropdown/WorkerDropdown.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/SCIM.tsx": { + "no-restricted-imports": { + "count": 2 + }, + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, + "src/components/SSOModals.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/SSOModals.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx": { "no-restricted-imports": { "count": 1 } @@ -1452,19 +2458,55 @@ "src/components/Settings/AdminSettings/HashicorpVault/HashicorpVault.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/HashicorpVault/HashicorpVaultEmptyPlaceholder.tsx": { + "no-restricted-imports": { + "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 } }, + "src/components/Settings/AdminSettings/MCPSemanticFilterSettings/MCPSemanticFilterTestPanel.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/PluginSettings/PluginSettings.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/SSOSettings/Modals/AddSSOSettingsModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.test.tsx": { @@ -1472,9 +2514,32 @@ "count": 1 } }, + "src/components/Settings/AdminSettings/SSOSettings/Modals/EditSSOSettingsModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/Settings/AdminSettings/SSOSettings/RedactableField.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/SSOSettings/RoleMappings.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/SSOSettings/SSOSettings.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/Settings/AdminSettings/SSOSettings/SSOSettingsEmptyPlaceholder.tsx": { + "no-restricted-imports": { + "count": 1 } }, "src/components/Settings/AdminSettings/SSOSettings/SSOSettingsLoadingSkeleton.test.tsx": { @@ -1482,7 +2547,15 @@ "count": 4 } }, + "src/components/Settings/AdminSettings/SSOSettings/SSOSettingsLoadingSkeleton.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/Settings/AdminSettings/UISettings/PageVisibilitySettings.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-render": { "count": 2 } @@ -1490,19 +2563,40 @@ "src/components/Settings/AdminSettings/UISettings/UISettings.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/Settings/RouterSettings/Fallbacks/AddFallbacks.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/components/Settings/RouterSettings/Fallbacks/FallbackSelectionForm.tsx": { + "src/components/Settings/RouterSettings/Fallbacks/AddFallbacksModal.tsx": { "no-restricted-imports": { "count": 1 + } + }, + "src/components/Settings/RouterSettings/Fallbacks/EditFallbacks.tsx": { + "no-restricted-imports": { + "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": 2 }, "react-hooks/set-state-in-effect": { "count": 1 @@ -1510,7 +2604,10 @@ }, "src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 + }, + "prefer-const": { + "count": 2 } }, "src/components/TeamSSOSettings.test.tsx": { @@ -1518,66 +2615,146 @@ "count": 1 } }, + "src/components/TeamSSOSettings.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/Teams.test.tsx": { "max-nested-callbacks": { "count": 4 + }, + "prefer-const": { + "count": 6 } }, "src/components/Teams.tsx": { + "local/no-complex-jsx-arrow": { + "count": 4 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 2 }, "no-restricted-imports": { - "count": 1 + "count": 2 + }, + "prefer-const": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 3 } }, + "src/components/TeamsPage/teamTableColumns.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/ToolDetail.tsx": { + "no-restricted-imports": { + "count": 1 + }, "unused-imports/no-unused-imports": { "count": 2 } }, - "src/components/ToolPolicies.tsx": { - "no-nested-ternary": { - "count": 1 - }, + "src/components/ToolPolicies/PolicySelect.tsx": { "no-restricted-imports": { "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 - }, - "react-hooks/static-components": { - "count": 7 - }, - "unused-imports/no-unused-imports": { + } + }, + "src/components/ToolPolicies/ToolPoliciesTableColumns.tsx": { + "no-restricted-imports": { "count": 1 } }, "src/components/UIAccessControlForm.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/UsagePage/components/EntityUsage/TopKeyView.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/components/activity_metrics.tsx": { - "no-nested-ternary": { + "src/components/UsagePage/components/KeyModelUsageView.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/UsagePage/utils/value_formatters.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/VirtualKeysPage/keyTableColumns.tsx": { + "local/filename-pascal-case": { "count": 1 }, "no-restricted-imports": { "count": 1 } }, - "src/components/add_model/AddModelForm.tsx": { + "src/components/activity_metrics.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/add_model/AdaptiveRoutingConfig.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/add_model/AddModelForm.test.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/add_model/AddModelForm.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-nested-ternary": { + "count": 1 + }, + "no-restricted-imports": { + "count": 4 + } + }, + "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 } }, "src/components/add_model/RouterConfigBuilder.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/purity": { "count": 1 }, @@ -1585,56 +2762,153 @@ "count": 1 } }, - "src/components/add_model/add_auto_router_tab.tsx": { + "src/components/add_model/SemanticKeywordMatching.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/add_model/add_auto_router_tab.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/add_model/add_auto_router_tab.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 3 + } + }, + "src/components/add_model/add_model_modes.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/add_model/add_model_tab.test.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, "src/components/add_model/add_model_tab.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 4 + } + }, + "src/components/add_model/advanced_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 4 + }, + "prefer-const": { + "count": 2 + } + }, + "src/components/add_model/auto_router_connection_test.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, - "src/components/add_model/advanced_settings.tsx": { + "src/components/add_model/cache_control_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "prefer-const": { + "count": 1 + } + }, + "src/components/add_model/conditional_public_model_name.test.tsx": { "no-restricted-imports": { "count": 1 } }, "src/components/add_model/conditional_public_model_name.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 2 } }, + "src/components/add_model/handle_add_auto_router_submit.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/add_model/handle_add_model_submit.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/add_model/litellm_model_name.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/add_model/litellm_model_name.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, + "no-restricted-imports": { + "count": 3 + } + }, + "src/components/add_model/model_connection_test.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-nested-ternary": { + "count": 2 + }, "no-restricted-imports": { "count": 1 } }, - "src/components/add_model/model_connection_test.tsx": { - "no-nested-ternary": { - "count": 2 + "src/components/add_model/provider_specific_fields.test.tsx": { + "no-restricted-imports": { + "count": 1 } }, "src/components/add_model/provider_specific_fields.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 5 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/immutability": { "count": 3 } }, "src/components/add_pass_through.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { - "count": 2 + "count": 3 } }, "src/components/agent_management/AgentSelector.test.tsx": { @@ -1645,28 +2919,58 @@ "count": 1 } }, + "src/components/agent_management/AgentSelector.tsx": { + "no-restricted-imports": { + "count": 1 + }, + "prefer-const": { + "count": 1 + } + }, + "src/components/alerting/alerting_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "prefer-const": { + "count": 1 + } + }, "src/components/alerting/dynamic_form.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 4 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/bulk_create_users_button.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/callback_info_helpers.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/chat/KeysPanel.tsx": { "no-nested-ternary": { "count": 1 } }, "src/components/chat/MCPAppsPanel.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 7 } @@ -1686,17 +2990,45 @@ "count": 2 } }, - "src/components/claude_code_plugins/MakeSkillPublicForm.tsx": { + "src/components/chat_ui/MCPEventsDisplay.tsx": { "no-restricted-imports": { "count": 1 + } + }, + "src/components/chat_ui/ReasoningContent.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/chat_ui/ResponseMetrics.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/chat_ui/mode_endpoint_mapping.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/claude_code_plugins/MakeSkillPublicForm.tsx": { + "no-restricted-imports": { + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/components/cloudzero_export_modal.tsx": { - "no-restricted-imports": { + "src/components/claude_code_plugins/skill_detail.tsx": { + "local/filename-pascal-case": { "count": 1 + } + }, + "src/components/cloudzero_export_modal.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 }, "no-restricted-syntax": { "count": 3 @@ -1707,7 +3039,7 @@ }, "src/components/common_components/AccessGroupSelector.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/common_components/AutoRotationView.tsx": { @@ -1715,26 +3047,70 @@ "count": 1 } }, + "src/components/common_components/DefaultProxyAdminTag.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/common_components/DeleteResourceModal.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/common_components/DurationSelect.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/common_components/Filters/FilterInput.tsx": { + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/common_components/Filters/FiltersButton.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/Filters/ResetFiltersButton.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/common_components/IconActionButton/BaseActionButton.tsx": { "no-restricted-imports": { "count": 1 } }, - "src/components/common_components/KeyLifecycleSettings.tsx": { + "src/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/common_components/KeyLifecycleSettings.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/common_components/LabeledField.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/MemberTable.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, "src/components/common_components/ModelAliasManager.tsx": { "no-restricted-imports": { "count": 1 @@ -1745,41 +3121,93 @@ }, "src/components/common_components/ModelSelector.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/common_components/NewBadge.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/OrganizationDropdown.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/components/common_components/PassThroughGuardrailsSection.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/components/common_components/PassThroughSecuritySection.tsx": { + "src/components/common_components/PassThroughRoutesSelector.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/common_components/PassThroughSecuritySection.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, "src/components/common_components/PremiumLoggingSettings.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/common_components/ProjectDropdown.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/RateLimitTypeFormItem.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/RateLimitTypeFormItem.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/common_components/RouterSettingsAccordion.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/common_components/TableHeaderSortDropdown/TableHeaderSortDropdown.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/budget_duration_dropdown.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/components/common_components/chartUtils.test.tsx": { "no-restricted-imports": { "count": 1 } }, "src/components/common_components/chartUtils.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, @@ -1788,16 +3216,25 @@ } }, "src/components/common_components/check_openapi_schema.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 3 } }, "src/components/common_components/fetch_teams.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, "src/components/common_components/simple_table.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, @@ -1805,56 +3242,182 @@ "count": 1 } }, + "src/components/common_components/team_dropdown.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/team_multi_select.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/common_components/user_search_modal.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/components/constants.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/components/edit_auto_router/edit_auto_router_modal.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/immutability": { "count": 1 } }, "src/components/email_events/email_event_settings.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/immutability": { "count": 1 } }, "src/components/email_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { + "count": 2 + }, + "prefer-const": { + "count": 1 + } + }, + "src/components/guardrails/GuardrailSelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/key_info_utils.tsx": { + "local/filename-pascal-case": { + "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/TagRateLimitEditor.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/key_team_helpers/fetch_available_models_team_key.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "prefer-const": { "count": 1 } }, "src/components/key_team_helpers/key_list.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/key_team_helpers/transform_key_info.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/key_value_input.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { + "count": 2 + } + }, + "src/components/leftnav.tsx": { + "local/filename-pascal-case": { "count": 1 } }, "src/components/llm_calls/chat_completion.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 }, "no-nested-ternary": { "count": 1 + }, + "prefer-const": { + "count": 1 + } + }, + "src/components/llm_calls/fetch_models.tsx": { + "local/filename-pascal-case": { + "count": 1 } }, "src/components/llm_calls/responses_api.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 } }, + "src/components/logging_settings_view.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "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 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/mcp_tools/ByokCredentialModal.tsx": { "no-restricted-imports": { "count": 1 } @@ -1862,6 +3425,9 @@ "src/components/mcp_tools/MCPToolArgumentsForm.tsx": { "no-nested-ternary": { "count": 5 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/mcp_tools/McpCrudPermissionPanel.tsx": { @@ -1869,31 +3435,61 @@ "count": 3 }, "no-restricted-imports": { + "count": 2 + } + }, + "src/components/mcp_tools/types.tsx": { + "local/filename-pascal-case": { "count": 1 } }, "src/components/model_add/CredentialModal.tsx": { + "no-restricted-imports": { + "count": 3 + } + }, + "src/components/model_add/CredentialsPanel.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/model_add/CredentialsPanel.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/model_add/credential_form_helpers.test.ts": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/model_add/credential_form_helpers.ts": { "no-restricted-imports": { "count": 1 } }, "src/components/model_add/reuse_credentials.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/components/model_dashboard/HealthCheckComponent.tsx": { - "no-nested-ternary": { - "count": 3 - }, "no-restricted-imports": { - "count": 1 - }, - "react-hooks/immutability": { + "count": 2 + } + }, + "src/components/model_dashboard/ModelSettingsModal/ModelSettingsModal.tsx": { + "no-restricted-imports": { "count": 1 } }, "src/components/model_dashboard/all_models_table.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, @@ -1901,8 +3497,55 @@ "count": 1 } }, - "src/components/model_dashboard/health_check_columns.tsx": { - "max-params": { + "src/components/model_filters.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/model_group_alias_settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, + "src/components/model_info_view.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, + "no-nested-ternary": { + "count": 14 + }, + "no-restricted-imports": { + "count": 2 + }, + "prefer-const": { + "count": 5 + }, + "react-hooks/set-state-in-effect": { + "count": 1 + } + }, + "src/components/molecules/cost_optimization_feedback_banner.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/molecules/filter.tsx": { + "local/filename-pascal-case": { "count": 1 }, "no-nested-ternary": { @@ -1912,40 +3555,11 @@ "count": 1 } }, - "src/components/model_dashboard/table.tsx": { - "no-nested-ternary": { + "src/components/molecules/message_manager.tsx": { + "local/filename-pascal-case": { "count": 1 }, "no-restricted-imports": { - "count": 1 - } - }, - "src/components/model_filters.tsx": { - "no-restricted-imports": { - "count": 1 - } - }, - "src/components/model_group_alias_settings.tsx": { - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, - "src/components/model_info_view.tsx": { - "no-nested-ternary": { - "count": 14 - }, - "no-restricted-imports": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 - } - }, - "src/components/molecules/filter.tsx": { - "no-nested-ternary": { "count": 2 } }, @@ -1958,22 +3572,55 @@ } }, "src/components/molecules/models/columns.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "max-params": { "count": 1 }, "no-nested-ternary": { "count": 2 }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/molecules/notifications_manager.test.tsx": { "no-restricted-imports": { "count": 1 } }, + "src/components/molecules/notifications_manager.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 3 + } + }, "src/components/navbar.test.tsx": { + "prefer-const": { + "count": 1 + }, "unused-imports/no-unused-imports": { "count": 1 } }, + "src/components/navbar.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, "src/components/networking.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "max-params": { "count": 23 }, @@ -1982,14 +3629,28 @@ }, "no-restricted-syntax": { "count": 154 + }, + "prefer-const": { + "count": 33 } }, "src/components/object_permissions_view.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, "src/components/onboarding_link.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/organisms/RegenerateKeyModal.tsx": { "no-restricted-imports": { "count": 1 } @@ -2003,17 +3664,32 @@ } }, "src/components/organisms/create_key_button.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "local/no-complex-jsx-arrow": { + "count": 2 + }, + "max-lines": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + }, + "prefer-const": { + "count": 4 + }, "react-hooks/set-state-in-effect": { "count": 4 } }, "src/components/organization/organization_view.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 3 + }, "unused-imports/no-unused-imports": { "count": 1 } @@ -2024,11 +3700,17 @@ } }, "src/components/pass_through_info.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 + }, + "no-restricted-imports": { + "count": 2 } }, "src/components/per_user_usage.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -2038,7 +3720,7 @@ }, "src/components/permissions/AgentPermissions.tsx": { "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/permissions/MCPServerPermissions.tsx": { @@ -2046,7 +3728,7 @@ "count": 3 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/permissions/VectorStorePermissions.tsx": { @@ -2057,19 +3739,58 @@ "src/components/policies/PolicySelector.tsx": { "no-nested-ternary": { "count": 1 - } - }, - "src/components/price_data_reload.tsx": { - "react-hooks/immutability": { - "count": 2 - } - }, - "src/components/public_model_hub.tsx": { + }, "no-restricted-imports": { "count": 1 } }, + "src/components/price_data_reload.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "react-hooks/immutability": { + "count": 2 + } + }, + "src/components/provider_info_helpers.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "prefer-const": { + "count": 3 + } + }, + "src/components/public_model_hub.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, "src/components/query_param_input.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/route_preview.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/router_settings/LatencyBasedConfiguration.tsx": { "no-restricted-imports": { "count": 1 } @@ -2077,9 +3798,49 @@ "src/components/router_settings/ReliabilityRetriesSection.tsx": { "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/router_settings/RoutingStrategySelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/router_settings/TagFilteringToggle.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/router_settings/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, + "prefer-const": { + "count": 2 + } + }, + "src/components/routing_groups/RoutingGroupModal.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/routing_groups/RoutingGroupsTable.tsx": { + "no-restricted-imports": { + "count": 2 } }, "src/components/routing_groups/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/preserve-manual-memoization": { "count": 1 } @@ -2087,17 +3848,42 @@ "src/components/search_tools/SearchToolSelector.tsx": { "no-nested-ternary": { "count": 1 - } - }, - "src/components/settings.tsx": { - "no-nested-ternary": { - "count": 2 }, "no-restricted-imports": { "count": 1 } }, + "src/components/settings.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/settings.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 4 + }, + "no-nested-ternary": { + "count": 2 + }, + "no-restricted-imports": { + "count": 3 + }, + "prefer-const": { + "count": 7 + } + }, + "src/components/shared/CreatedKeyDisplay.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/shared/advanced_date_picker.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -2105,67 +3891,208 @@ "count": 3 } }, + "src/components/shared/chart_loader.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/charts/area_chart.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/charts/bar_chart.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/charts/chart_legend.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/charts/chart_tooltip.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/charts/donut_chart.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/charts/line_chart.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/errorUtils.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/form/FormField.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + } + }, + "src/components/shared/form/field.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/shared/numerical_input.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, + "src/components/shared/table_cells/cell_tooltip.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/date_cell.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/id_cell.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/identity_cell.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/models_cell.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/money_cell.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/spend_budget_cell.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/shared/table_cells/status_badge.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/shared/usage_date_picker.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 } }, + "src/components/tag_management/TagSelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/tag_management/types.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/team/EditMembership.tsx": { "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/team/LoggingSettings.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, + "src/components/team/MyUserTab.tsx": { "no-restricted-imports": { "count": 1 } }, "src/components/team/TeamInfo.tsx": { + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 3 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/team/TeamMemberTab.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, + "no-restricted-imports": { + "count": 2 + } + }, "src/components/team/TeamVirtualKeysTable.tsx": { "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/team/member_permissions.tsx": { - "no-restricted-imports": { + "local/filename-pascal-case": { "count": 1 }, + "no-restricted-imports": { + "count": 2 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, + "src/components/team/permission_definitions.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/team/useMyTeamMember.ts": { "no-restricted-syntax": { "count": 1 } }, + "src/components/templates/KeyInfoHeader.tsx": { + "no-restricted-imports": { + "count": 2 + } + }, "src/components/templates/key_edit_view.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "local/no-complex-jsx-arrow": { + "count": 2 + }, "no-nested-ternary": { "count": 2 }, "no-restricted-imports": { - "count": 1 + "count": 2 } }, "src/components/templates/key_info_view.test.tsx": { @@ -2174,41 +4101,253 @@ } }, "src/components/templates/key_info_view.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "max-lines": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, "no-restricted-imports": { - "count": 1 + "count": 2 }, "react-hooks/set-state-in-effect": { "count": 1 } }, - "src/components/user_agent_activity.tsx": { + "src/components/ui/AntDLoadingSpinner.tsx": { "no-restricted-imports": { - "count": 2 + "count": 1 + } + }, + "src/components/ui/alert-dialog.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/avatar.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/badge.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/breadcrumb.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/button.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/card.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/chart.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/checkbox.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/collapsible.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/combobox.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/dialog.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/dropdown-menu.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/input-group.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/input.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/label.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/meter.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/popover.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/scroll-area.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/select.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/separator.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/sheet.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/sidebar.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/skeleton.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/switch.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/table.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/tabs.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/textarea.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/tooltip.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/ui/ui-loading-spinner.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/update_model_credentials_modal.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/user_agent_activity.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 3 }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/components/user_dashboard.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, + "prefer-const": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 2 } }, + "src/components/vector_store_management/VectorStoreSelector.test.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/vector_store_management/VectorStoreSelector.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/vector_store_management/types.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/vector_store_providers.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/view_logs/AuditLogDrawer/AuditLogDrawer.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/CostBreakdownViewer.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/view_logs/EvalViewer/EvalViewer.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/view_logs/GuardrailViewer/CompliancePanel.tsx": { "no-nested-ternary": { "count": 2 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } @@ -2221,26 +4360,90 @@ "src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx": { "no-nested-ternary": { "count": 4 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/CollapsibleMessage.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/DrawerHeader.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/HistoryTree.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/JsonViewer.tsx": { + "no-restricted-imports": { + "count": 1 } }, "src/components/view_logs/LogDetailsDrawer/LogDetailContent.tsx": { "no-nested-ternary": { - "count": 4 + "count": 3 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx": { "no-nested-ternary": { "count": 3 }, + "no-restricted-imports": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 2 } }, + "src/components/view_logs/LogDetailsDrawer/OutputCard.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.test.tsx": { "unused-imports/no-unused-imports": { "count": 2 } }, + "src/components/view_logs/LogDetailsDrawer/RealtimePrettyView.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/SectionHeader.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/SimpleMessageBlock.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/SimpleToolCallBlock.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/TokenFlow.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/LogDetailsDrawer/TruncatedValue.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/components/view_logs/LogDetailsDrawer/prettyMessagesUtils.ts": { "no-nested-ternary": { "count": 1 @@ -2252,11 +4455,53 @@ } }, "src/components/view_logs/LogsTableToolbar.tsx": { + "local/no-complex-jsx-arrow": { + "count": 1 + }, "no-nested-ternary": { "count": 4 + }, + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/ToolsSection/FormattedToolView.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/ToolsSection/ToolExpandedContent.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/ToolsSection/ToolItem.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/ToolsSection/ToolsSection.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/VectorStoreViewer.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, + "src/components/view_logs/columns.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "no-restricted-imports": { + "count": 1 } }, "src/components/view_logs/index.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-restricted-imports": { "count": 1 }, @@ -2264,16 +4509,45 @@ "count": 1 } }, + "src/components/view_logs/log_filter_logic.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, + "src/components/view_logs/logs_utils.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/view_logs/table.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 2 } }, + "src/components/view_model/model_name_display.tsx": { + "local/filename-pascal-case": { + "count": 1 + } + }, "src/components/view_user_spend.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, + "prefer-const": { + "count": 3 + }, "react-hooks/set-state-in-effect": { "count": 2 } }, + "src/contexts/AntdGlobalProvider.tsx": { + "no-restricted-imports": { + "count": 1 + } + }, "src/contexts/AuthContext.tsx": { "react-hooks/set-state-in-effect": { "count": 1 @@ -2300,11 +4574,17 @@ } }, "src/hooks/useMcpOAuthFlow.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "react-hooks/set-state-in-effect": { "count": 1 } }, "src/hooks/useTestMCPConnection.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "no-nested-ternary": { "count": 1 }, @@ -2313,6 +4593,9 @@ } }, "src/hooks/useToolsOAuthFlow.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "react-hooks/refs": { "count": 1 }, @@ -2321,6 +4604,9 @@ } }, "src/hooks/useUserMcpOAuthFlow.tsx": { + "local/filename-pascal-case": { + "count": 1 + }, "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 0cf5b4ff655..b5e876bfbba 100644 --- a/ui/litellm-dashboard/eslint.config.mjs +++ b/ui/litellm-dashboard/eslint.config.mjs @@ -19,12 +19,13 @@ const eslintConfig = [ "unused-imports/no-unused-imports": "error", "local/no-large-inline-object-arg": "warn", "local/no-long-condition-chain": "warn", + "local/no-complex-jsx-arrow": ["error", { maxStatements: 2 }], "@typescript-eslint/no-explicit-any": "warn", "no-console": ["warn", { allow: ["warn", "error"] }], "@typescript-eslint/no-unused-vars": "off", "@typescript-eslint/no-unused-expressions": "off", "@typescript-eslint/ban-ts-comment": "off", - "prefer-const": "off", + "prefer-const": "error", "no-empty": "off", "no-prototype-builtins": "off", "no-useless-catch": "off", @@ -51,13 +52,32 @@ const eslintConfig = [ patterns: [ { group: ["@tremor/react", "@tremor/react/*"], - message: "@tremor/react is being phased out; build new UI with antd instead of adding tremor imports.", + message: + "@tremor/react is being phased out; build new UI with shadcn/ui primitives instead of adding tremor imports.", + }, + { + group: ["antd", "antd/*"], + message: + "antd is being phased out; build new UI with shadcn/ui primitives instead of adding antd imports.", }, ], }, ], }, }, + { + files: ["src/**/*.tsx"], + rules: { + "local/filename-pascal-case": "error", + }, + }, + { + files: ["src/**/*.{ts,tsx}"], + ignores: ["src/**/*.test.{ts,tsx}", "src/**/*.spec.{ts,tsx}", "src/data/**"], + rules: { + "max-lines": ["error", { max: 800, skipBlankLines: true, skipComments: true }], + }, + }, { files: ["src/lib/http/**"], rules: { diff --git a/ui/litellm-dashboard/knip.json b/ui/litellm-dashboard/knip.json index afed6b0f90e..48b39e8122d 100644 --- a/ui/litellm-dashboard/knip.json +++ b/ui/litellm-dashboard/knip.json @@ -1,7 +1,7 @@ { "$schema": "https://unpkg.com/knip@5/schema.json", "entry": ["scripts/**/*.{ts,mjs}", "src/components/ui/**/*.{ts,tsx}"], - "project": ["src/**/*.{ts,tsx}", "tests/**/*.{ts,tsx}", "scripts/**/*.{ts,mjs}", "e2e_tests/**/*.ts"], + "project": ["src/**/*.{ts,tsx}", "tests/**/*.{ts,tsx}", "scripts/**/*.{ts,mjs}"], "ignore": ["src/lib/http/schema.d.ts"], "ignoreDependencies": [ "openapi-typescript", @@ -10,14 +10,6 @@ "tailwindcss", "tw-animate-css" ], - "playwright": { - "config": [ - "e2e_tests/playwright.config.ts", - "e2e_tests/serverRootPath.config.ts", - "e2e_tests/migration.serverRootPath.config.ts" - ], - "entry": ["e2e_tests/**/*.spec.ts", "e2e_tests/**/*.setup.ts", "e2e_tests/globalSetup.ts"] - }, "vitest": { "config": ["vitest.config.ts"] }, diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 7a65b63b33c..5f6b4b889b1 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -14,6 +14,7 @@ "@base-ui/react": "^1.6.0", "@headlessui/tailwindcss": "0.2.2", "@heroicons/react": "1.0.6", + "@hookform/resolvers": "5.4.0", "@tanstack/react-pacer": "0.22.1", "@tanstack/react-query": "5.100.7", "@tanstack/react-table": "8.21.3", @@ -26,7 +27,7 @@ "jwt-decode": "4.0.0", "lucide-react": "0.513.0", "moment": "2.30.1", - "next": "16.2.6", + "next": "16.2.11", "openai": "4.104.0", "openapi-fetch": "^0.17.0", "openapi-react-query": "^0.5.4", @@ -34,17 +35,18 @@ "react": "18.3.1", "react-copy-to-clipboard": "5.1.1", "react-dom": "18.3.1", + "react-hook-form": "7.82.0", "react-json-view-lite": "2.5.0", "react-markdown": "9.1.0", "react-syntax-highlighter": "15.6.6", "recharts": "3.9.2", "remark-gfm": "4.0.1", "tailwind-merge": "3.4.0", - "uuid": "14.0.0" + "uuid": "14.0.0", + "zod": "3.25.76" }, "devDependencies": { "@eslint/js": "9.39.2", - "@playwright/test": "1.58.1", "@tailwindcss/forms": "0.5.11", "@tailwindcss/postcss": "4.3.2", "@testing-library/dom": "10.4.1", @@ -59,7 +61,7 @@ "@vitest/coverage-v8": "3.2.6", "@vitest/ui": "3.2.6", "eslint": "9.39.2", - "eslint-config-next": "16.2.6", + "eslint-config-next": "16.2.11", "eslint-config-prettier": "10.1.8", "eslint-plugin-unused-imports": "4.3.0", "jsdom": "27.4.0", @@ -799,9 +801,9 @@ } }, "node_modules/@emnapi/runtime": { - "version": "1.10.0", - "resolved": "https://registry.npmjs.org/@emnapi/runtime/-/runtime-1.10.0.tgz", - "integrity": "sha512-ewvYlk86xUoGI0zQRNq/mC+16R1QeDlKQy21Ki3oSYXNgLb45GV1P6A0M+/s6nyCuNDqe5VpaY84BzXGwVbwFA==", + "version": "1.11.2", + "resolved": "https://registry.npmjs.org/@emnapi/runtime/-/runtime-1.11.2.tgz", + "integrity": "sha512-kyOl3X0DuTiT1h2ft8r2fYO8JYtU9a9Xis/zBSiGArNaagCOWx90N1k2wxp18czFDH+OgcWGb5ZP/XMt3dcyPA==", "license": "MIT", "optional": true, "dependencies": { @@ -1556,6 +1558,18 @@ "react": ">= 16" } }, + "node_modules/@hookform/resolvers": { + "version": "5.4.0", + "resolved": "https://registry.npmjs.org/@hookform/resolvers/-/resolvers-5.4.0.tgz", + "integrity": "sha512-EIsqr/t/qbinPIhGjMdtvutIN1Kk4uwbROE9/UQ93CAVGR7GkA7Y92+fX80OzXi/OB67jVFYwKGO1WzkxmkFZw==", + "license": "MIT", + "dependencies": { + "@standard-schema/utils": "^0.3.0" + }, + "peerDependencies": { + "react-hook-form": "^7.55.0" + } + }, "node_modules/@humanfs/core": { "version": "0.19.2", "resolved": "https://registry.npmjs.org/@humanfs/core/-/core-0.19.2.tgz", @@ -1633,9 +1647,9 @@ } }, "node_modules/@img/sharp-darwin-arm64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-darwin-arm64/-/sharp-darwin-arm64-0.34.5.tgz", - "integrity": "sha512-imtQ3WMJXbMY4fxb/Ndp6HBTNVtWCUI0WdobyheGf5+ad6xX8VIDO8u2xE4qc/fr08CKG/7dDseFtn6M6g/r3w==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-darwin-arm64/-/sharp-darwin-arm64-0.35.3.tgz", + "integrity": "sha512-RMnFX7YQsMoh7lWfcM4NEHHymBX/rLuKNPVM84XE9ONPcaSCDgE7CHIHpSgPcO2xcRthgBy1HfNO319mwhIAkg==", "cpu": [ "arm64" ], @@ -1645,19 +1659,19 @@ "darwin" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-darwin-arm64": "1.2.4" + "@img/sharp-libvips-darwin-arm64": "1.3.2" } }, "node_modules/@img/sharp-darwin-x64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-darwin-x64/-/sharp-darwin-x64-0.34.5.tgz", - "integrity": "sha512-YNEFAF/4KQ/PeW0N+r+aVVsoIY0/qxxikF2SWdp+NRkmMB7y9LBZAVqQ4yhGCm/H3H270OSykqmQMKLBhBJDEw==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-darwin-x64/-/sharp-darwin-x64-0.35.3.tgz", + "integrity": "sha512-Xo+5uFBtLN0BKqieTxiFzFPQAUlBbbH5iBKyRX/z1JrbnYsHTfKJnUfL8+p2TPXr1pXqao4eeL4Rl144uDpK9w==", "cpu": [ "x64" ], @@ -1667,19 +1681,38 @@ "darwin" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-darwin-x64": "1.2.4" + "@img/sharp-libvips-darwin-x64": "1.3.2" + } + }, + "node_modules/@img/sharp-freebsd-wasm32": { + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-freebsd-wasm32/-/sharp-freebsd-wasm32-0.35.3.tgz", + "integrity": "sha512-lUxcqWIj2wMQ9BrwNjngcr1gWUr5xgaGThBRqPPalIC2n67Cqj1uPh8NnA/ZhAg8hUbKl+kVHKwgUIwe6ZYPrg==", + "license": "Apache-2.0", + "optional": true, + "os": [ + "freebsd" + ], + "dependencies": { + "@img/sharp-wasm32": "0.35.3" + }, + "engines": { + "node": ">=20.9.0" + }, + "funding": { + "url": "https://opencollective.com/libvips" } }, "node_modules/@img/sharp-libvips-darwin-arm64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-darwin-arm64/-/sharp-libvips-darwin-arm64-1.2.4.tgz", - "integrity": "sha512-zqjjo7RatFfFoP0MkQ51jfuFZBnVE2pRiaydKJ1G/rHZvnsrHAOcQALIi9sA5co5xenQdTugCvtb1cuf78Vf4g==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-darwin-arm64/-/sharp-libvips-darwin-arm64-1.3.2.tgz", + "integrity": "sha512-9J6ypZFpQBj4YnePGoq/S38w6nz+vqg5WZLrLGY4YuSemdMq47GMLBPO42MzwdGwpg/agZ7xzZcFHa48xlywfg==", "cpu": [ "arm64" ], @@ -1693,9 +1726,9 @@ } }, "node_modules/@img/sharp-libvips-darwin-x64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-darwin-x64/-/sharp-libvips-darwin-x64-1.2.4.tgz", - "integrity": "sha512-1IOd5xfVhlGwX+zXv2N93k0yMONvUlANylbJw1eTah8K/Jtpi15KC+WSiaX/nBmbm2HxRM1gZ0nSdjSsrZbGKg==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-darwin-x64/-/sharp-libvips-darwin-x64-1.3.2.tgz", + "integrity": "sha512-m2pW1n6cns9VaubNwsZ+c3CRYjxNQWgJ5gPlnL1nbBcpkBvFm6SCFN5o0psFHI8w9n11NKhFkeEDns98tiqbEw==", "cpu": [ "x64" ], @@ -1709,9 +1742,9 @@ } }, "node_modules/@img/sharp-libvips-linux-arm": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-arm/-/sharp-libvips-linux-arm-1.2.4.tgz", - "integrity": "sha512-bFI7xcKFELdiNCVov8e44Ia4u2byA+l3XtsAj+Q8tfCwO6BQ8iDojYdvoPMqsKDkuoOo+X6HZA0s0q11ANMQ8A==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-arm/-/sharp-libvips-linux-arm-1.3.2.tgz", + "integrity": "sha512-1eMLzy92I4J6rmi4mAT8yC3HxOtniyGELlzGbNMLLeqe052ahFQ0h6LFq+lh5DsDIdYViIDst08abvSbcEdLXQ==", "cpu": [ "arm" ], @@ -1725,9 +1758,9 @@ } }, "node_modules/@img/sharp-libvips-linux-arm64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-arm64/-/sharp-libvips-linux-arm64-1.2.4.tgz", - "integrity": "sha512-excjX8DfsIcJ10x1Kzr4RcWe1edC9PquDRRPx3YVCvQv+U5p7Yin2s32ftzikXojb1PIFc/9Mt28/y+iRklkrw==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-arm64/-/sharp-libvips-linux-arm64-1.3.2.tgz", + "integrity": "sha512-dqVSFynCox4C/J8kT16V7SIFAns0IjgLwkvYT7p8LQVmJ5OS5b6tI9IGflxTeuBS//zXeFIUbwt5dwxyZ17cnA==", "cpu": [ "arm64" ], @@ -1741,9 +1774,9 @@ } }, "node_modules/@img/sharp-libvips-linux-ppc64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-ppc64/-/sharp-libvips-linux-ppc64-1.2.4.tgz", - "integrity": "sha512-FMuvGijLDYG6lW+b/UvyilUWu5Ayu+3r2d1S8notiGCIyYU/76eig1UfMmkZ7vwgOrzKzlQbFSuQfgm7GYUPpA==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-ppc64/-/sharp-libvips-linux-ppc64-1.3.2.tgz", + "integrity": "sha512-3z0NHDxD6n5I9gc05U1eW1AyRm+Gznzq3naMrthPNqE6oYykcogW0l/jfpJdjYnuNl8R7yI9pNbE1XiUeyq0Aw==", "cpu": [ "ppc64" ], @@ -1757,9 +1790,9 @@ } }, "node_modules/@img/sharp-libvips-linux-riscv64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-riscv64/-/sharp-libvips-linux-riscv64-1.2.4.tgz", - "integrity": "sha512-oVDbcR4zUC0ce82teubSm+x6ETixtKZBh/qbREIOcI3cULzDyb18Sr/Wcyx7NRQeQzOiHTNbZFF1UwPS2scyGA==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-riscv64/-/sharp-libvips-linux-riscv64-1.3.2.tgz", + "integrity": "sha512-bsb4rI+NldGOsXuej2r8OdSS8+zXDVaCWxyWrcv6kneTOlgAHtZABRzBBCwdsPiD90J4myNJuHpg6kA20ImW/w==", "cpu": [ "riscv64" ], @@ -1773,9 +1806,9 @@ } }, "node_modules/@img/sharp-libvips-linux-s390x": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-s390x/-/sharp-libvips-linux-s390x-1.2.4.tgz", - "integrity": "sha512-qmp9VrzgPgMoGZyPvrQHqk02uyjA0/QrTO26Tqk6l4ZV0MPWIW6LTkqOIov+J1yEu7MbFQaDpwdwJKhbJvuRxQ==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-s390x/-/sharp-libvips-linux-s390x-1.3.2.tgz", + "integrity": "sha512-/ABshyj8gCpyIrNXnHn4LorDJ0HHm1VhXPBlxZ8zAtfVPAaSafXPGn+sUSIRiwaSBy0mmFjSjiXI5mkcwdChKQ==", "cpu": [ "s390x" ], @@ -1789,9 +1822,9 @@ } }, "node_modules/@img/sharp-libvips-linux-x64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-x64/-/sharp-libvips-linux-x64-1.2.4.tgz", - "integrity": "sha512-tJxiiLsmHc9Ax1bz3oaOYBURTXGIRDODBqhveVHonrHJ9/+k89qbLl0bcJns+e4t4rvaNBxaEZsFtSfAdquPrw==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linux-x64/-/sharp-libvips-linux-x64-1.3.2.tgz", + "integrity": "sha512-ITPEtgffGJ0S6G9dRyw/366tJQqFRcHWPHhC+Stpg3Z8AEMrDrTr2lhdz4f/Y/HMbRh//7Z5mBzEpVdi62Oc3w==", "cpu": [ "x64" ], @@ -1805,9 +1838,9 @@ } }, "node_modules/@img/sharp-libvips-linuxmusl-arm64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linuxmusl-arm64/-/sharp-libvips-linuxmusl-arm64-1.2.4.tgz", - "integrity": "sha512-FVQHuwx1IIuNow9QAbYUzJ+En8KcVm9Lk5+uGUQJHaZmMECZmOlix9HnH7n1TRkXMS0pGxIJokIVB9SuqZGGXw==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linuxmusl-arm64/-/sharp-libvips-linuxmusl-arm64-1.3.2.tgz", + "integrity": "sha512-zE9EdiUzUmg5mDT5a1rk5fYJ6GWPloTwWBYDS14naqHsL+EaMpDj1AWnpLgh3u0YCORv2Tt50wrcrpYqkP97Kw==", "cpu": [ "arm64" ], @@ -1821,9 +1854,9 @@ } }, "node_modules/@img/sharp-libvips-linuxmusl-x64": { - "version": "1.2.4", - "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linuxmusl-x64/-/sharp-libvips-linuxmusl-x64-1.2.4.tgz", - "integrity": "sha512-+LpyBk7L44ZIXwz/VYfglaX/okxezESc6UxDSoyo2Ks6Jxc4Y7sGjpgU9s4PMgqgjj1gZCylTieNamqA1MF7Dg==", + "version": "1.3.2", + "resolved": "https://registry.npmjs.org/@img/sharp-libvips-linuxmusl-x64/-/sharp-libvips-linuxmusl-x64-1.3.2.tgz", + "integrity": "sha512-m0lrLiUt+lBYnCFr8qV/65yMR4E/c7/wf78I5eKTdkEakFAlZ9QlzEM3QIhhAwVeUhLAHLcCq7a7Vszq/oFNZQ==", "cpu": [ "x64" ], @@ -1837,9 +1870,9 @@ } }, "node_modules/@img/sharp-linux-arm": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linux-arm/-/sharp-linux-arm-0.34.5.tgz", - "integrity": "sha512-9dLqsvwtg1uuXBGZKsxem9595+ujv0sJ6Vi8wcTANSFpwV/GONat5eCkzQo/1O6zRIkh0m/8+5BjrRr7jDUSZw==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linux-arm/-/sharp-linux-arm-0.35.3.tgz", + "integrity": "sha512-affVWCTLooy8TSxbDx2qkzuDeaWLNVBA+P//FNBirHsXpP2fuBhk5AuboYUnrDnzoXes8GFjpTx0SBFOCRg+FA==", "cpu": [ "arm" ], @@ -1849,19 +1882,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linux-arm": "1.2.4" + "@img/sharp-libvips-linux-arm": "1.3.2" } }, "node_modules/@img/sharp-linux-arm64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linux-arm64/-/sharp-linux-arm64-0.34.5.tgz", - "integrity": "sha512-bKQzaJRY/bkPOXyKx5EVup7qkaojECG6NLYswgktOZjaXecSAeCWiZwwiFf3/Y+O1HrauiE3FVsGxFg8c24rZg==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linux-arm64/-/sharp-linux-arm64-0.35.3.tgz", + "integrity": "sha512-QgKDspHPnrU+GQ55XPhGwyhC8acLVOOSyAvo1oVfFmrIXLkDNmGWzAfDZ4xK8oSA1qBQrALcHX0G5UZni/SuFQ==", "cpu": [ "arm64" ], @@ -1871,19 +1904,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linux-arm64": "1.2.4" + "@img/sharp-libvips-linux-arm64": "1.3.2" } }, "node_modules/@img/sharp-linux-ppc64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linux-ppc64/-/sharp-linux-ppc64-0.34.5.tgz", - "integrity": "sha512-7zznwNaqW6YtsfrGGDA6BRkISKAAE1Jo0QdpNYXNMHu2+0dTrPflTLNkpc8l7MUP5M16ZJcUvysVWWrMefZquA==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linux-ppc64/-/sharp-linux-ppc64-0.35.3.tgz", + "integrity": "sha512-sMd8rDxmpLOwv/7N44klFjOD5DUO7FLdjiXDI0hoxYaf7Ar262dQIEkosE98bps+5HPLtp/EvNqeqQtOycP/IA==", "cpu": [ "ppc64" ], @@ -1893,19 +1926,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linux-ppc64": "1.2.4" + "@img/sharp-libvips-linux-ppc64": "1.3.2" } }, "node_modules/@img/sharp-linux-riscv64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linux-riscv64/-/sharp-linux-riscv64-0.34.5.tgz", - "integrity": "sha512-51gJuLPTKa7piYPaVs8GmByo7/U7/7TZOq+cnXJIHZKavIRHAP77e3N2HEl3dgiqdD/w0yUfiJnII77PuDDFdw==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linux-riscv64/-/sharp-linux-riscv64-0.35.3.tgz", + "integrity": "sha512-0Eob78yjlYPfL5vMNWAW55l3R9Y6BQS/gOfe0ZcP9mEz9ohhKSt4im1hayiknXgf8AWrFqMvJcKIdmLmEe7yeQ==", "cpu": [ "riscv64" ], @@ -1915,19 +1948,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linux-riscv64": "1.2.4" + "@img/sharp-libvips-linux-riscv64": "1.3.2" } }, "node_modules/@img/sharp-linux-s390x": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linux-s390x/-/sharp-linux-s390x-0.34.5.tgz", - "integrity": "sha512-nQtCk0PdKfho3eC5MrbQoigJ2gd1CgddUMkabUj+rBevs8tZ2cULOx46E7oyX+04WGfABgIwmMC0VqieTiR4jg==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linux-s390x/-/sharp-linux-s390x-0.35.3.tgz", + "integrity": "sha512-KgAxQ0DxpNOq1rG2t5cgTgShJFGSuU7XO45cqC+1NVOuZnP6tlgZRuSYOfNupGkHID0o3cJOsw4DVeJpMovcGw==", "cpu": [ "s390x" ], @@ -1937,19 +1970,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linux-s390x": "1.2.4" + "@img/sharp-libvips-linux-s390x": "1.3.2" } }, "node_modules/@img/sharp-linux-x64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linux-x64/-/sharp-linux-x64-0.34.5.tgz", - "integrity": "sha512-MEzd8HPKxVxVenwAa+JRPwEC7QFjoPWuS5NZnBt6B3pu7EG2Ge0id1oLHZpPJdn3OQK+BQDiw9zStiHBTJQQQQ==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linux-x64/-/sharp-linux-x64-0.35.3.tgz", + "integrity": "sha512-8pqvxubL2PGdhlPy6GLqzDYMUjyRmKAwKHYKixpdJYBUK7PJ0C029XdsnpFIdgRZG68fZiGdHVWcKPvtiPB4cA==", "cpu": [ "x64" ], @@ -1959,19 +1992,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linux-x64": "1.2.4" + "@img/sharp-libvips-linux-x64": "1.3.2" } }, "node_modules/@img/sharp-linuxmusl-arm64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linuxmusl-arm64/-/sharp-linuxmusl-arm64-0.34.5.tgz", - "integrity": "sha512-fprJR6GtRsMt6Kyfq44IsChVZeGN97gTD331weR1ex1c1rypDEABN6Tm2xa1wE6lYb5DdEnk03NZPqA7Id21yg==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linuxmusl-arm64/-/sharp-linuxmusl-arm64-0.35.3.tgz", + "integrity": "sha512-Vz0iQjzzcSX3HCbfwFfCSG/9SCIqyO0mH2sXyiHaAYfBk0cRsCWXRyQYX0ovCK/PAQBbTzQ0dsPQHh5MAFL59w==", "cpu": [ "arm64" ], @@ -1981,19 +2014,19 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linuxmusl-arm64": "1.2.4" + "@img/sharp-libvips-linuxmusl-arm64": "1.3.2" } }, "node_modules/@img/sharp-linuxmusl-x64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-linuxmusl-x64/-/sharp-linuxmusl-x64-0.34.5.tgz", - "integrity": "sha512-Jg8wNT1MUzIvhBFxViqrEhWDGzqymo3sV7z7ZsaWbZNDLXRJZoRGrjulp60YYtV4wfY8VIKcWidjojlLcWrd8Q==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-linuxmusl-x64/-/sharp-linuxmusl-x64-0.35.3.tgz", + "integrity": "sha512-6O1NPKcDVj9QEdg7Hx549EX8U0rp6yXQERqru6yRN7fGBn32UvIRJUlWnk+8xDCiG76hXVBbX82NZ/ZKr0euIg==", "cpu": [ "x64" ], @@ -2003,38 +2036,54 @@ "linux" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-libvips-linuxmusl-x64": "1.2.4" + "@img/sharp-libvips-linuxmusl-x64": "1.3.2" } }, "node_modules/@img/sharp-wasm32": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-wasm32/-/sharp-wasm32-0.34.5.tgz", - "integrity": "sha512-OdWTEiVkY2PHwqkbBI8frFxQQFekHaSSkUIJkwzclWZe64O1X4UlUjqqqLaPbUpMOQk6FBu/HtlGXNblIs0huw==", - "cpu": [ - "wasm32" - ], + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-wasm32/-/sharp-wasm32-0.35.3.tgz", + "integrity": "sha512-cZ0XkcYGpHZkqW6iCkqTcmUC0CD9DhD5d/qeZlZkfRBn6GnHniZXLUo5+9xw8Iv76YE6LQFN9YNBlKREcCG76w==", "license": "Apache-2.0 AND LGPL-3.0-or-later AND MIT", "optional": true, "dependencies": { - "@emnapi/runtime": "^1.7.0" + "@emnapi/runtime": "^1.11.1" }, "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" + }, + "funding": { + "url": "https://opencollective.com/libvips" + } + }, + "node_modules/@img/sharp-webcontainers-wasm32": { + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-webcontainers-wasm32/-/sharp-webcontainers-wasm32-0.35.3.tgz", + "integrity": "sha512-2rnq7bX3NzeR2T4YWgz8qiG4h3TSdMe+vN1iQXpJleSJ3SM5zQ8Fy2SyyXAWlbxpEZ2Y+Z4u1BePgJEYbSy80Q==", + "cpu": [ + "wasm32" + ], + "license": "Apache-2.0", + "optional": true, + "dependencies": { + "@img/sharp-wasm32": "0.35.3" + }, + "engines": { + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" } }, "node_modules/@img/sharp-win32-arm64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-win32-arm64/-/sharp-win32-arm64-0.34.5.tgz", - "integrity": "sha512-WQ3AgWCWYSb2yt+IG8mnC6Jdk9Whs7O0gxphblsLvdhSpSTtmu69ZG1Gkb6NuvxsNACwiPV6cNSZNzt0KPsw7g==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-win32-arm64/-/sharp-win32-arm64-0.35.3.tgz", + "integrity": "sha512-4bPwFdMbeC4JQ8L8LOyWp6nsHcboP5fxkp6iPOXz2Vg49R42TuMs2whkJ5OAP4/Ul035qOzy0AecOF9VOscn4w==", "cpu": [ "arm64" ], @@ -2044,16 +2093,16 @@ "win32" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" } }, "node_modules/@img/sharp-win32-ia32": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-win32-ia32/-/sharp-win32-ia32-0.34.5.tgz", - "integrity": "sha512-FV9m/7NmeCmSHDD5j4+4pNI8Cp3aW+JvLoXcTUo0IqyjSfAZJ8dIUmijx1qaJsIiU+Hosw6xM5KijAWRJCSgNg==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-win32-ia32/-/sharp-win32-ia32-0.35.3.tgz", + "integrity": "sha512-r53mXsBN6lFUDiST764SvgwUdHAqM4rPAiDzAmf4fLoB6X/rkfyTrLCg6+g17wJJiCmB3JYgHuUldCWUIRFSXw==", "cpu": [ "ia32" ], @@ -2063,16 +2112,16 @@ "win32" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": "^20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" } }, "node_modules/@img/sharp-win32-x64": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/@img/sharp-win32-x64/-/sharp-win32-x64-0.34.5.tgz", - "integrity": "sha512-+29YMsqY2/9eFEiW93eqWnuLcWcufowXewwSNIT6UwZdUUCrM3oFjMWH/Z6/TMmb4hlFenmfAVbpWeup2jryCw==", + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/@img/sharp-win32-x64/-/sharp-win32-x64-0.35.3.tgz", + "integrity": "sha512-D4y1vNeZrIIJCN+uHaWVtH86B+aCrdMYYjicy9pXHvbGZeGYLLSd3wdVuC37FxVXlU1ARsk84eKWfWMXGYEqvA==", "cpu": [ "x64" ], @@ -2082,7 +2131,7 @@ "win32" ], "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" @@ -2195,15 +2244,15 @@ } }, "node_modules/@next/env": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/env/-/env-16.2.6.tgz", - "integrity": "sha512-gd8HoHN4ufj73WmR3JmVolrpJR47ILK6LouP5xElPglaVxir6e1a7VzvTvDWkOoPXT9rkkTzyCxBu4yeZfZwcw==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/env/-/env-16.2.11.tgz", + "integrity": "sha512-0do5A3BJ2gxWr0ZCMcD6BhW+e595jyxdTl3rXTS6lOtD8ektMiW6CO+EPwt1Eca1DBnm90r/7GdiKWBKxH++DA==", "license": "MIT" }, "node_modules/@next/eslint-plugin-next": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.2.6.tgz", - "integrity": "sha512-Z8l6o4JWKUl755x4R+wogD86KPeU+Ckw4K+SYG4kHeOJtRenDeK+OSbGcqZpDtbwn9DsJVdir2UxmwXuinUbUw==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/eslint-plugin-next/-/eslint-plugin-next-16.2.11.tgz", + "integrity": "sha512-vMEf/aXOpzFFdtIvFYOnIDPKb0xBbrXONsz83CcKdRrekfxNdL8PNkq5qHqAHSXVlIifnX68LOMaxr3z5PkeLQ==", "dev": true, "license": "MIT", "dependencies": { @@ -2211,9 +2260,9 @@ } }, "node_modules/@next/swc-darwin-arm64": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-darwin-arm64/-/swc-darwin-arm64-16.2.6.tgz", - "integrity": "sha512-ZJGkkcNfYgrrMkqOdZ7zoLa1TOy0qpcMfk/z4Mh/FKUz40gVO+HNQWqmLxf67Z5WB64DRp0dhEbyHfel+6sJUg==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-darwin-arm64/-/swc-darwin-arm64-16.2.11.tgz", + "integrity": "sha512-wryL4pjKmDwGv2ox6+GZDFxvmtSRLqApBR8kL1j4+vhB7Z5vJC/zAnXpiR9Xkfzl0AS8WLMnsuGV/UKI67/rrw==", "cpu": [ "arm64" ], @@ -2227,9 +2276,9 @@ } }, "node_modules/@next/swc-darwin-x64": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-16.2.6.tgz", - "integrity": "sha512-v/YLBHIY132Ced3puBJ7YJKw1lqsCrgcNo2aRJlCEyQrrCeRJlvGlnmxhPxNQI3KE3N1DN5r9TPNPvka3nq5RQ==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-darwin-x64/-/swc-darwin-x64-16.2.11.tgz", + "integrity": "sha512-aZl2j4f/fLyjQvOhv0Oe9UaMAQHolYpKhctsoYzplSumKJKPUmgjcf6545aBtysLTcu994TREd0+pSgNE4ohmg==", "cpu": [ "x64" ], @@ -2243,9 +2292,9 @@ } }, "node_modules/@next/swc-linux-arm64-gnu": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-16.2.6.tgz", - "integrity": "sha512-RPOvqlYBbcQjkz9VQQDZ2T2bARIjXZV1KFlt+V2Mr6SW/e4I9fcKsaA0hdyf2FHoTlsV2xnBd5Y912rP/1Ce6w==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-gnu/-/swc-linux-arm64-gnu-16.2.11.tgz", + "integrity": "sha512-5jEriyEnH/LWFy27L2ZG0XaLlyEJIjhsImEsiS9P563PKEVp2BVups/xfOucIrsvVntp11oNcZwjHvaDPYVB5g==", "cpu": [ "arm64" ], @@ -2259,9 +2308,9 @@ } }, "node_modules/@next/swc-linux-arm64-musl": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-16.2.6.tgz", - "integrity": "sha512-URUTu1+dMkxJsPFgm+OeEvq9wf5sujw0EvgYy80TDGHTSLTnIHeqb0Eu8A3sC95IRgjejQL+kC4mw+4yPxiAXA==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-linux-arm64-musl/-/swc-linux-arm64-musl-16.2.11.tgz", + "integrity": "sha512-eIjcpx2fnnFSSkZDbTxy74KnokUXDjfoLClpWelfgHLf621aTqswhwXQ7GkD5K5rplrS6LZ/Bj+mVuvzluBOEg==", "cpu": [ "arm64" ], @@ -2275,9 +2324,9 @@ } }, "node_modules/@next/swc-linux-x64-gnu": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-16.2.6.tgz", - "integrity": "sha512-DOj182mPV8G3UkrayLoREM5YEYI+Dk5wv7Ox9xl1fFibAELEsFD0lDPfHIeILlutMMfdyhlzYPELG3peuKaurw==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-linux-x64-gnu/-/swc-linux-x64-gnu-16.2.11.tgz", + "integrity": "sha512-8WgzpaWMs46qJT9kiV47cje86L0x/Mu9t8/Gwj+pnbgW3rETVfCnaScPjlYUwNScpOozdcIMHWmAvuZJUonR2w==", "cpu": [ "x64" ], @@ -2291,9 +2340,9 @@ } }, "node_modules/@next/swc-linux-x64-musl": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-16.2.6.tgz", - "integrity": "sha512-HKQ5SP/V/ub73UvF7n/zeJlxk2kLmtL7Wzrg4WfmkjmNos5onJ2tKu7yZOPdL18A6Svfn3max29ym+ry7NkK4g==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-linux-x64-musl/-/swc-linux-x64-musl-16.2.11.tgz", + "integrity": "sha512-I3UgPds7G4ZYnTb/H+5GBGuUT2DhAk6j0mL6A4s63RjFs74wB2hOWP0vaxsK+3NJraExt3eYEPQ/UtT0x/64Nw==", "cpu": [ "x64" ], @@ -2307,9 +2356,9 @@ } }, "node_modules/@next/swc-win32-arm64-msvc": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-16.2.6.tgz", - "integrity": "sha512-LZXpTlPyS5v7HhSmnvsLGP3iIYgYOBnc8r8ArlT55sGHV89bR2HlDdBjWQ+PY6SJMmk8TuVGFuxalnP3k/0Dwg==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-win32-arm64-msvc/-/swc-win32-arm64-msvc-16.2.11.tgz", + "integrity": "sha512-n89CjtcThnjrwgJMAiI5xbqwLY51zvwC9tSlArmVndAJLYVl9T9UAdlkXTmZvE++idoXe8KdglQlhNRdUp1c6g==", "cpu": [ "arm64" ], @@ -2323,9 +2372,9 @@ } }, "node_modules/@next/swc-win32-x64-msvc": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-16.2.6.tgz", - "integrity": "sha512-F0+4i0h9J6C4eE3EAPWsoCk7UW/dbzOjyzxY0qnDUOYFu6FFmdZ6l97/XdV3/Nz3VYyO7UWjyEJUXkGqcoXfMA==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/@next/swc-win32-x64-msvc/-/swc-win32-x64-msvc-16.2.11.tgz", + "integrity": "sha512-md8CLNggS1Dx9pUgApzps5uAf+N8GN9xywzmNx9vHAWo94HtBwCCqkSnhIrdfQe83Dhz8Lfo/20Nb1Zxal092w==", "cpu": [ "x64" ], @@ -2673,8 +2722,9 @@ "version": "1.58.1", "resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.58.1.tgz", "integrity": "sha512-6LdVIUERWxQMmUSSQi0I53GgCBYgM2RpGngCPY7hSeju+VrKjq3lvs7HpJoPbDiY5QM5EYRtRX5fvrinnMAz3w==", - "devOptional": true, "license": "Apache-2.0", + "optional": true, + "peer": true, "dependencies": { "playwright": "1.58.1" }, @@ -3570,6 +3620,72 @@ "node": ">=14.0.0" } }, + "node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@emnapi/core": { + "version": "1.11.1", + "dev": true, + "inBundle": true, + "license": "MIT", + "optional": true, + "dependencies": { + "@emnapi/wasi-threads": "1.2.2", + "tslib": "^2.4.0" + } + }, + "node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@emnapi/runtime": { + "version": "1.11.1", + "dev": true, + "inBundle": true, + "license": "MIT", + "optional": true, + "dependencies": { + "tslib": "^2.4.0" + } + }, + "node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@emnapi/wasi-threads": { + "version": "1.2.2", + "dev": true, + "inBundle": true, + "license": "MIT", + "optional": true, + "dependencies": { + "tslib": "^2.4.0" + } + }, + "node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@napi-rs/wasm-runtime": { + "version": "1.1.4", + "dev": true, + "inBundle": true, + "license": "MIT", + "optional": true, + "dependencies": { + "@tybys/wasm-util": "^0.10.1" + }, + "funding": { + "type": "github", + "url": "https://github.com/sponsors/Brooooooklyn" + }, + "peerDependencies": { + "@emnapi/core": "^1.7.1", + "@emnapi/runtime": "^1.7.1" + } + }, + "node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/@tybys/wasm-util": { + "version": "0.10.2", + "dev": true, + "inBundle": true, + "license": "MIT", + "optional": true, + "dependencies": { + "tslib": "^2.4.0" + } + }, + "node_modules/@tailwindcss/oxide-wasm32-wasi/node_modules/tslib": { + "version": "2.8.1", + "dev": true, + "inBundle": true, + "license": "0BSD", + "optional": true + }, "node_modules/@tailwindcss/oxide-win32-arm64-msvc": { "version": "4.3.2", "resolved": "https://registry.npmjs.org/@tailwindcss/oxide-win32-arm64-msvc/-/oxide-win32-arm64-msvc-4.3.2.tgz", @@ -6592,13 +6708,13 @@ } }, "node_modules/eslint-config-next": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.2.6.tgz", - "integrity": "sha512-z2ELYSkyrrJ6cuunTU8vhsT/RpouPkjaSah06nVW6Rg2Hpg0Vs8s497/e5s8G8qtdp4ccsiovz5P1rv+5VSW2Q==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/eslint-config-next/-/eslint-config-next-16.2.11.tgz", + "integrity": "sha512-FIpbK/dUyxUExchDB7eBg3k+VU8R2iR/Cx9/kqTBUTFv2bOIR9aRrpno4rvAQ9VhiPQAyFKNA2NlZwouGWtclA==", "dev": true, "license": "MIT", "dependencies": { - "@next/eslint-plugin-next": "16.2.6", + "@next/eslint-plugin-next": "16.2.11", "eslint-import-resolver-node": "^0.3.6", "eslint-import-resolver-typescript": "^3.5.2", "eslint-plugin-import": "^2.32.0", @@ -7311,7 +7427,6 @@ "version": "2.3.2", "resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz", "integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==", - "dev": true, "hasInstallScript": true, "license": "MIT", "optional": true, @@ -10260,12 +10375,12 @@ "license": "MIT" }, "node_modules/next": { - "version": "16.2.6", - "resolved": "https://registry.npmjs.org/next/-/next-16.2.6.tgz", - "integrity": "sha512-qOVgKJg1+At15NpeUP+eJgCHvTCgXsogweq87Ri/Ix7PkqQHg4sdaXmSFqKlgaIXE4kW0g25LE68W87UANlHtw==", + "version": "16.2.11", + "resolved": "https://registry.npmjs.org/next/-/next-16.2.11.tgz", + "integrity": "sha512-B339zaqbyK8cmxhoAvLrcwoabwCP1wz21zSzfqxqXAemTu2BXnH7tQnfcglKv1vnMUIDBc+Hth7XODQriTZiRQ==", "license": "MIT", "dependencies": { - "@next/env": "16.2.6", + "@next/env": "16.2.11", "@swc/helpers": "0.5.15", "baseline-browser-mapping": "^2.9.19", "caniuse-lite": "^1.0.30001579", @@ -10279,14 +10394,14 @@ "node": ">=20.9.0" }, "optionalDependencies": { - "@next/swc-darwin-arm64": "16.2.6", - "@next/swc-darwin-x64": "16.2.6", - "@next/swc-linux-arm64-gnu": "16.2.6", - "@next/swc-linux-arm64-musl": "16.2.6", - "@next/swc-linux-x64-gnu": "16.2.6", - "@next/swc-linux-x64-musl": "16.2.6", - "@next/swc-win32-arm64-msvc": "16.2.6", - "@next/swc-win32-x64-msvc": "16.2.6", + "@next/swc-darwin-arm64": "16.2.11", + "@next/swc-darwin-x64": "16.2.11", + "@next/swc-linux-arm64-gnu": "16.2.11", + "@next/swc-linux-arm64-musl": "16.2.11", + "@next/swc-linux-x64-gnu": "16.2.11", + "@next/swc-linux-x64-musl": "16.2.11", + "@next/swc-win32-arm64-msvc": "16.2.11", + "@next/swc-win32-x64-msvc": "16.2.11", "sharp": "^0.34.5" }, "peerDependencies": { @@ -10898,8 +11013,9 @@ "version": "1.58.1", "resolved": "https://registry.npmjs.org/playwright/-/playwright-1.58.1.tgz", "integrity": "sha512-+2uTZHxSCcxjvGc5C891LrS1/NlxglGxzrC4seZiVjcYVQfUa87wBL6rTDqzGjuoWNjnBzRqKmF6zRYGMvQUaQ==", - "devOptional": true, "license": "Apache-2.0", + "optional": true, + "peer": true, "dependencies": { "playwright-core": "1.58.1" }, @@ -10917,8 +11033,9 @@ "version": "1.58.1", "resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.58.1.tgz", "integrity": "sha512-bcWzOaTxcW+VOOGBCQgnaKToLJ65d6AqfLVKEWvexyS3AS6rbXl+xdpYRMGSRBClPvyj44njOWoxjNdL/H9UNg==", - "devOptional": true, "license": "Apache-2.0", + "optional": true, + "peer": true, "bin": { "playwright-core": "cli.js" }, @@ -11780,6 +11897,22 @@ "react": "^18.3.1" } }, + "node_modules/react-hook-form": { + "version": "7.82.0", + "resolved": "https://registry.npmjs.org/react-hook-form/-/react-hook-form-7.82.0.tgz", + "integrity": "sha512-Zw/uFZ2dO+02GHlBn7JFGn8kZJ7LdM33B/0BXOovzFay+CMhf94JMw5BVu+F1tVkUKjNvBuaE3fz5BJhga10Tg==", + "license": "MIT", + "engines": { + "node": ">=18.0.0" + }, + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/react-hook-form" + }, + "peerDependencies": { + "react": "^16.8.0 || ^17 || ^18 || ^19" + } + }, "node_modules/react-is": { "version": "17.0.2", "resolved": "https://registry.npmjs.org/react-is/-/react-is-17.0.2.tgz", @@ -12503,48 +12636,53 @@ } }, "node_modules/sharp": { - "version": "0.34.5", - "resolved": "https://registry.npmjs.org/sharp/-/sharp-0.34.5.tgz", - "integrity": "sha512-Ou9I5Ft9WNcCbXrU9cMgPBcCK8LiwLqcbywW3t4oDV37n1pzpuNLsYiAV8eODnjbtQlSDwZ2cUEeQz4E54Hltg==", - "hasInstallScript": true, + "version": "0.35.3", + "resolved": "https://registry.npmjs.org/sharp/-/sharp-0.35.3.tgz", + "integrity": "sha512-ej0zVHuZGHCiABXcNxeYhpRnPNPAcvbG8RMdBAhDAxLKkCRVSpK3Iyu7qbqw3JMzoj0REeM6f3tJLtVwl0023Q==", "license": "Apache-2.0", "optional": true, "dependencies": { - "@img/colour": "^1.0.0", + "@img/colour": "^1.1.0", "detect-libc": "^2.1.2", - "semver": "^7.7.3" + "semver": "^7.8.5" }, "engines": { - "node": "^18.17.0 || ^20.3.0 || >=21.0.0" + "node": ">=20.9.0" }, "funding": { "url": "https://opencollective.com/libvips" }, "optionalDependencies": { - "@img/sharp-darwin-arm64": "0.34.5", - "@img/sharp-darwin-x64": "0.34.5", - "@img/sharp-libvips-darwin-arm64": "1.2.4", - "@img/sharp-libvips-darwin-x64": "1.2.4", - "@img/sharp-libvips-linux-arm": "1.2.4", - "@img/sharp-libvips-linux-arm64": "1.2.4", - "@img/sharp-libvips-linux-ppc64": "1.2.4", - "@img/sharp-libvips-linux-riscv64": "1.2.4", - "@img/sharp-libvips-linux-s390x": "1.2.4", - "@img/sharp-libvips-linux-x64": "1.2.4", - "@img/sharp-libvips-linuxmusl-arm64": "1.2.4", - "@img/sharp-libvips-linuxmusl-x64": "1.2.4", - "@img/sharp-linux-arm": "0.34.5", - "@img/sharp-linux-arm64": "0.34.5", - "@img/sharp-linux-ppc64": "0.34.5", - "@img/sharp-linux-riscv64": "0.34.5", - "@img/sharp-linux-s390x": "0.34.5", - "@img/sharp-linux-x64": "0.34.5", - "@img/sharp-linuxmusl-arm64": "0.34.5", - "@img/sharp-linuxmusl-x64": "0.34.5", - "@img/sharp-wasm32": "0.34.5", - "@img/sharp-win32-arm64": "0.34.5", - "@img/sharp-win32-ia32": "0.34.5", - "@img/sharp-win32-x64": "0.34.5" + "@img/sharp-darwin-arm64": "0.35.3", + "@img/sharp-darwin-x64": "0.35.3", + "@img/sharp-freebsd-wasm32": "0.35.3", + "@img/sharp-libvips-darwin-arm64": "1.3.2", + "@img/sharp-libvips-darwin-x64": "1.3.2", + "@img/sharp-libvips-linux-arm": "1.3.2", + "@img/sharp-libvips-linux-arm64": "1.3.2", + "@img/sharp-libvips-linux-ppc64": "1.3.2", + "@img/sharp-libvips-linux-riscv64": "1.3.2", + "@img/sharp-libvips-linux-s390x": "1.3.2", + "@img/sharp-libvips-linux-x64": "1.3.2", + "@img/sharp-libvips-linuxmusl-arm64": "1.3.2", + "@img/sharp-libvips-linuxmusl-x64": "1.3.2", + "@img/sharp-linux-arm": "0.35.3", + "@img/sharp-linux-arm64": "0.35.3", + "@img/sharp-linux-ppc64": "0.35.3", + "@img/sharp-linux-riscv64": "0.35.3", + "@img/sharp-linux-s390x": "0.35.3", + "@img/sharp-linux-x64": "0.35.3", + "@img/sharp-linuxmusl-arm64": "0.35.3", + "@img/sharp-linuxmusl-x64": "0.35.3", + "@img/sharp-webcontainers-wasm32": "0.35.3", + "@img/sharp-win32-arm64": "0.35.3", + "@img/sharp-win32-ia32": "0.35.3", + "@img/sharp-win32-x64": "0.35.3" + }, + "peerDependenciesMeta": { + "@types/node": { + "optional": true + } } }, "node_modules/shebang-command": { @@ -14156,7 +14294,6 @@ "version": "3.25.76", "resolved": "https://registry.npmjs.org/zod/-/zod-3.25.76.tgz", "integrity": "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ==", - "devOptional": true, "license": "MIT", "funding": { "url": "https://github.com/sponsors/colinhacks" diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index b5e93d175bf..63bcdfb6076 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -14,10 +14,6 @@ "test:coverage": "vitest run --coverage", "format": "prettier --write .", "format:check": "prettier --check .", - "e2e": "playwright test --config e2e_tests/playwright.config.ts", - "e2e:ui": "playwright test --ui --config e2e_tests/playwright.config.ts", - "e2e:migration": "playwright test e2e_tests/tests/migration/migratedPages.spec.ts --config e2e_tests/playwright.config.ts", - "e2e:migration:root": "playwright test --config e2e_tests/migration.serverRootPath.config.ts", "knip": "knip", "knip:ci": "knip --exclude exports,nsExports,types,nsTypes,enumMembers,classMembers,duplicates", "knip:fix": "knip --fix", @@ -30,6 +26,7 @@ "@base-ui/react": "^1.6.0", "@headlessui/tailwindcss": "0.2.2", "@heroicons/react": "1.0.6", + "@hookform/resolvers": "5.4.0", "@tanstack/react-pacer": "0.22.1", "@tanstack/react-query": "5.100.7", "@tanstack/react-table": "8.21.3", @@ -42,7 +39,7 @@ "jwt-decode": "4.0.0", "lucide-react": "0.513.0", "moment": "2.30.1", - "next": "16.2.6", + "next": "16.2.11", "openai": "4.104.0", "openapi-fetch": "^0.17.0", "openapi-react-query": "^0.5.4", @@ -50,17 +47,18 @@ "react": "18.3.1", "react-copy-to-clipboard": "5.1.1", "react-dom": "18.3.1", + "react-hook-form": "7.82.0", "react-json-view-lite": "2.5.0", "react-markdown": "9.1.0", "react-syntax-highlighter": "15.6.6", "recharts": "3.9.2", "remark-gfm": "4.0.1", "tailwind-merge": "3.4.0", - "uuid": "14.0.0" + "uuid": "14.0.0", + "zod": "3.25.76" }, "devDependencies": { "@eslint/js": "9.39.2", - "@playwright/test": "1.58.1", "@tailwindcss/forms": "0.5.11", "@tailwindcss/postcss": "4.3.2", "@testing-library/dom": "10.4.1", @@ -75,7 +73,7 @@ "@vitest/coverage-v8": "3.2.6", "@vitest/ui": "3.2.6", "eslint": "9.39.2", - "eslint-config-next": "16.2.6", + "eslint-config-next": "16.2.11", "eslint-config-prettier": "10.1.8", "eslint-plugin-unused-imports": "4.3.0", "jsdom": "27.4.0", @@ -100,7 +98,8 @@ "axios": "1.13.6", "postcss": "8.5.13", "esbuild": "0.28.1", - "date-fns": "^4.4.0" + "date-fns": "^4.4.0", + "sharp": "^0.35.0" }, "engines": { "node": ">=20.9.0", diff --git a/ui/litellm-dashboard/public/assets/logos/ai21.svg b/ui/litellm-dashboard/public/assets/logos/ai21.svg index 7e62a9517af..3c8c75e6d6f 100644 --- a/ui/litellm-dashboard/public/assets/logos/ai21.svg +++ b/ui/litellm-dashboard/public/assets/logos/ai21.svg @@ -1 +1 @@ -AI21 \ No newline at end of file +AI21 \ No newline at end of file diff --git a/ui/litellm-dashboard/public/assets/logos/promptguard.svg b/ui/litellm-dashboard/public/assets/logos/promptguard.svg index 44cdd52eae3..4b2fd3c386e 100644 --- a/ui/litellm-dashboard/public/assets/logos/promptguard.svg +++ b/ui/litellm-dashboard/public/assets/logos/promptguard.svg @@ -1,5 +1,5 @@ + viewBox="0 0 1024 1024" enable-background="new 0 0 1024 1024" xml:space="preserve"> Soniox +Soniox diff --git a/ui/litellm-dashboard/scripts/eslint-rules/filename-pascal-case.mjs b/ui/litellm-dashboard/scripts/eslint-rules/filename-pascal-case.mjs new file mode 100644 index 00000000000..7477925238e --- /dev/null +++ b/ui/litellm-dashboard/scripts/eslint-rules/filename-pascal-case.mjs @@ -0,0 +1,59 @@ +import { basename } from "path"; + +const NEXT_RESERVED = new Set([ + "page", + "layout", + "route", + "template", + "default", + "loading", + "error", + "global-error", + "not-found", + "middleware", + "instrumentation", + "sitemap", + "robots", + "manifest", + "icon", + "apple-icon", + "favicon", + "opengraph-image", + "twitter-image", +]); + +const PASCAL_CASE = /^[A-Z][A-Za-z0-9]*$/; + +const rule = { + meta: { + type: "suggestion", + docs: { + description: "Require PascalCase filenames for .tsx modules; exempt Next.js reserved files and test/spec files.", + }, + schema: [], + messages: { + notPascalCase: "Filename '{{name}}' should be PascalCase (e.g. '{{suggestion}}.tsx').", + }, + }, + create(context) { + const filename = context.filename; + const stem = basename(filename).replace(/\.tsx$/, ""); + const [head, ...rest] = stem.split("."); + if (rest.includes("test") || rest.includes("spec")) return {}; + if (NEXT_RESERVED.has(head)) return {}; + if (PASCAL_CASE.test(head)) return {}; + const pascalHead = head + .split(/[-_]/) + .filter(Boolean) + .map((part) => part.charAt(0).toUpperCase() + part.slice(1)) + .join(""); + const suggestion = [pascalHead, ...rest].join("."); + return { + Program(node) { + context.report({ node, messageId: "notPascalCase", data: { name: `${stem}.tsx`, suggestion } }); + }, + }; + }, +}; + +export default rule; diff --git a/ui/litellm-dashboard/scripts/eslint-rules/index.mjs b/ui/litellm-dashboard/scripts/eslint-rules/index.mjs index 150ba1d02e9..9e9f901a6df 100644 --- a/ui/litellm-dashboard/scripts/eslint-rules/index.mjs +++ b/ui/litellm-dashboard/scripts/eslint-rules/index.mjs @@ -1,10 +1,14 @@ import noLargeInlineObjectArg from "./no-large-inline-object-arg.mjs"; import noLongConditionChain from "./no-long-condition-chain.mjs"; +import noComplexJsxArrow from "./no-complex-jsx-arrow.mjs"; +import filenamePascalCase from "./filename-pascal-case.mjs"; const plugin = { rules: { "no-large-inline-object-arg": noLargeInlineObjectArg, "no-long-condition-chain": noLongConditionChain, + "no-complex-jsx-arrow": noComplexJsxArrow, + "filename-pascal-case": filenamePascalCase, }, }; diff --git a/ui/litellm-dashboard/scripts/eslint-rules/no-complex-jsx-arrow.mjs b/ui/litellm-dashboard/scripts/eslint-rules/no-complex-jsx-arrow.mjs new file mode 100644 index 00000000000..b3dabe03a21 --- /dev/null +++ b/ui/litellm-dashboard/scripts/eslint-rules/no-complex-jsx-arrow.mjs @@ -0,0 +1,41 @@ +const DEFAULT_MAX_STATEMENTS = 2; + +const isJsxAttributeValue = (node) => { + const parent = node.parent; + if (parent == null) return false; + return parent.type === "JSXExpressionContainer" && parent.parent?.type === "JSXAttribute"; +}; + +const rule = { + meta: { + type: "suggestion", + docs: { + description: + "Disallow arrow functions with block bodies over a few statements passed inline as JSX attributes; extract them into a named handler.", + }, + schema: [ + { + type: "object", + properties: { maxStatements: { type: "integer", minimum: 1 } }, + additionalProperties: false, + }, + ], + messages: { + tooComplex: "Inline JSX arrow handler has {{count}} statements; extract it into a named function (max {{max}}).", + }, + }, + create(context) { + const maxStatements = context.options[0]?.maxStatements ?? DEFAULT_MAX_STATEMENTS; + return { + ArrowFunctionExpression(node) { + if (node.body.type !== "BlockStatement") return; + if (!isJsxAttributeValue(node)) return; + const count = node.body.body.length; + if (count <= maxStatements) return; + context.report({ node, messageId: "tooComplex", data: { count, max: maxStatements } }); + }, + }; + }, +}; + +export default rule; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.test.tsx index 7c8aaa2b785..a1484ffb5c5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsPage.test.tsx @@ -38,6 +38,7 @@ const mockAccessGroups: AccessGroupResponse[] = [ const mockUseAccessGroups = vi.fn(); const mockUseDeleteAccessGroup = vi.fn(); const mockMutate = vi.fn(); +const mockUseAuthorized = vi.fn(); vi.mock("@/app/(dashboard)/hooks/accessGroups/useAccessGroups", () => ({ useAccessGroups: () => mockUseAccessGroups(), @@ -47,6 +48,10 @@ vi.mock("@/app/(dashboard)/hooks/accessGroups/useDeleteAccessGroup", () => ({ useDeleteAccessGroup: () => mockUseDeleteAccessGroup(), })); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + vi.mock("./AccessGroupsDetailsPage", () => ({ AccessGroupDetail: ({ accessGroupId, onBack }: { accessGroupId: string; onBack: () => void }) => (
@@ -65,49 +70,42 @@ vi.mock("./AccessGroupsModal/AccessGroupCreateModal", () => ({ ) : null, })); -vi.mock("@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton", () => ({ - default: ({ variant, tooltipText, onClick }: { variant: string; tooltipText: string; onClick: () => void }) => ( - - ), -})); +const makeGroups = (count: number): AccessGroupResponse[] => + Array.from({ length: count }, (_, index) => { + const suffix = String(index + 1).padStart(2, "0"); + return { + ...mockAccessGroups[0], + access_group_id: `ag-${suffix}`, + access_group_name: `Group ${suffix}`, + description: `Group ${suffix} description`, + }; + }); + +const openRowMenu = async (user: ReturnType, groupId: string) => { + await user.click(screen.getByTestId(`access-group-actions-${groupId}`)); + return screen.findByTestId("access-group-action-delete"); +}; describe("AccessGroupsPage", () => { beforeEach(() => { vi.clearAllMocks(); - mockUseAccessGroups.mockReturnValue({ - data: mockAccessGroups, - isLoading: false, - }); - mockUseDeleteAccessGroup.mockReturnValue({ - mutate: mockMutate, - isPending: false, - }); + mockUseAccessGroups.mockReturnValue({ data: mockAccessGroups, isLoading: false }); + mockUseDeleteAccessGroup.mockReturnValue({ mutate: mockMutate, isPending: false }); + mockUseAuthorized.mockReturnValue({ userRole: "Admin", accessToken: "sk-test" }); }); - it("should render", () => { - renderWithProviders(); - expect(screen.getByRole("heading", { name: "Access Groups" })).toBeInTheDocument(); - }); - - it("should display page title and subtitle", () => { + it("renders the page title and subtitle", () => { renderWithProviders(); expect(screen.getByRole("heading", { name: "Access Groups" })).toBeInTheDocument(); expect(screen.getByText("Manage resource permissions for your organization")).toBeInTheDocument(); }); - it("should display Create Access Group button", () => { + it("shows the Create Access Group button for an admin", () => { renderWithProviders(); expect(screen.getByRole("button", { name: /create access group/i })).toBeInTheDocument(); }); - it("should display search input with placeholder", () => { - renderWithProviders(); - expect(screen.getByPlaceholderText("Search groups by name, ID, or description...")).toBeInTheDocument(); - }); - - it("should display access groups in table", () => { + it("renders every access group row", () => { renderWithProviders(); expect(screen.getByText("ag-1")).toBeInTheDocument(); expect(screen.getByText("Admin Group")).toBeInTheDocument(); @@ -115,57 +113,70 @@ describe("AccessGroupsPage", () => { expect(screen.getByText("Read Only")).toBeInTheDocument(); }); - it("should display resource counts for each group", () => { + it("renders resource counts for each group", () => { renderWithProviders(); - const table = screen.getByRole("table"); - expect(table).toHaveTextContent("2"); - expect(table).toHaveTextContent("1"); + // ag-1 has 2 models, 1 mcp server, 1 agent. + const adminRow = screen.getByText("ag-1").closest("tr") as HTMLElement; + expect(within(adminRow).getByTitle("2 Models")).toHaveTextContent("2"); + expect(within(adminRow).getByTitle("1 MCP Servers")).toHaveTextContent("1"); + expect(within(adminRow).getByTitle("1 Agents")).toHaveTextContent("1"); }); - it("should filter groups by search text matching name", async () => { + it("shows the expected column headers", () => { + renderWithProviders(); + expect(screen.getByRole("columnheader", { name: /^ID$/i })).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: /Name/i })).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: /Resources/i })).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: /Created/i })).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: /Updated/i })).toBeInTheDocument(); + }); + + it("filters by name", async () => { const user = userEvent.setup(); renderWithProviders(); - const searchInput = screen.getByPlaceholderText("Search groups by name, ID, or description..."); - await user.type(searchInput, "Admin"); + await user.type(screen.getByPlaceholderText("Search groups by name, ID, or description..."), "Admin"); expect(screen.getByText("Admin Group")).toBeInTheDocument(); expect(screen.queryByText("Read Only")).not.toBeInTheDocument(); }); - it("should filter groups by search text matching ID", async () => { + it("filters by ID", async () => { const user = userEvent.setup(); renderWithProviders(); - const searchInput = screen.getByPlaceholderText("Search groups by name, ID, or description..."); - await user.type(searchInput, "ag-2"); + await user.type(screen.getByPlaceholderText("Search groups by name, ID, or description..."), "ag-2"); expect(screen.getByText("Read Only")).toBeInTheDocument(); expect(screen.queryByText("Admin Group")).not.toBeInTheDocument(); }); - it("should filter groups by search text matching description", async () => { + it("filters by description", async () => { const user = userEvent.setup(); renderWithProviders(); - const searchInput = screen.getByPlaceholderText("Search groups by name, ID, or description..."); - await user.type(searchInput, "read-only"); + await user.type(screen.getByPlaceholderText("Search groups by name, ID, or description..."), "read-only"); expect(screen.getByText("Read Only")).toBeInTheDocument(); expect(screen.queryByText("Admin Group")).not.toBeInTheDocument(); }); - it("should reset to first page when search text changes", async () => { + it("shows the filtered empty state when nothing matches", async () => { const user = userEvent.setup(); renderWithProviders(); - const searchInput = screen.getByPlaceholderText("Search groups by name, ID, or description..."); - await user.type(searchInput, "Admin"); - const pagination = screen.getByText(/groups/); - expect(pagination).toHaveTextContent("1 groups"); + await user.type(screen.getByPlaceholderText("Search groups by name, ID, or description..."), "no-such-group"); + expect(screen.getByText("No matching access groups")).toBeInTheDocument(); + expect(screen.queryByText("Admin Group")).not.toBeInTheDocument(); }); - it("should open create modal when Create Access Group button is clicked", async () => { - const user = userEvent.setup(); + it("shows the empty state when there are no groups", () => { + mockUseAccessGroups.mockReturnValue({ data: [], isLoading: false }); renderWithProviders(); - await user.click(screen.getByRole("button", { name: /create access group/i })); - expect(screen.getByTestId("create-access-group-modal")).toBeInTheDocument(); + expect(screen.getByText("No access groups yet")).toBeInTheDocument(); }); - it("should close create modal when cancel is clicked", async () => { + it("renders loading skeletons on the initial load", () => { + mockUseAccessGroups.mockReturnValue({ data: undefined, isLoading: true }); + renderWithProviders(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("Admin Group")).not.toBeInTheDocument(); + }); + + it("opens and closes the create modal", async () => { const user = userEvent.setup(); renderWithProviders(); await user.click(screen.getByRole("button", { name: /create access group/i })); @@ -174,33 +185,22 @@ describe("AccessGroupsPage", () => { expect(screen.queryByTestId("create-access-group-modal")).not.toBeInTheDocument(); }); - it("should navigate to detail view when group ID is clicked", async () => { + it("opens the detail view when the ID cell is clicked and returns via Back", async () => { const user = userEvent.setup(); renderWithProviders(); await user.click(screen.getByText("ag-1")); expect(screen.getByTestId("access-group-detail")).toBeInTheDocument(); expect(screen.getByText("Detail for ag-1")).toBeInTheDocument(); - }); - - it("should return to list view when Back is clicked from detail", async () => { - const user = userEvent.setup(); - renderWithProviders(); - await user.click(screen.getByText("ag-1")); - expect(screen.getByTestId("access-group-detail")).toBeInTheDocument(); await user.click(screen.getByRole("button", { name: "Back" })); expect(screen.queryByTestId("access-group-detail")).not.toBeInTheDocument(); expect(screen.getByText("Admin Group")).toBeInTheDocument(); }); - it("should open delete modal when delete action is clicked", async () => { + it("opens the delete modal from the row actions menu", async () => { const user = userEvent.setup(); renderWithProviders(); - const deleteButtons = screen.getAllByRole("button", { - name: "Delete access group", - }); - await user.click(deleteButtons[0]); + await user.click(await openRowMenu(user, "ag-1")); const dialog = screen.getByRole("dialog", { name: "Delete Access Group" }); - expect(dialog).toBeInTheDocument(); expect( within(dialog).getByText("Are you sure you want to delete this access group? This action cannot be undone."), ).toBeInTheDocument(); @@ -209,71 +209,49 @@ describe("AccessGroupsPage", () => { expect(within(dialog).getByText("Admin Group")).toBeInTheDocument(); }); - it("should close delete modal when cancel is clicked", async () => { + it("closes the delete modal on cancel without deleting", async () => { const user = userEvent.setup(); renderWithProviders(); - const deleteButtons = screen.getAllByRole("button", { - name: "Delete access group", - }); - await user.click(deleteButtons[0]); + await user.click(await openRowMenu(user, "ag-1")); const dialog = screen.getByRole("dialog", { name: "Delete Access Group" }); await user.click(within(dialog).getByRole("button", { name: "Cancel" })); expect(screen.queryByRole("dialog", { name: "Delete Access Group" })).not.toBeInTheDocument(); + expect(mockMutate).not.toHaveBeenCalled(); }); - it("should call delete mutation when delete is confirmed", async () => { + it("calls the delete mutation with the group ID when confirmed", async () => { const user = userEvent.setup(); mockMutate.mockImplementation((_id: string, opts?: { onSuccess?: () => void }) => { opts?.onSuccess?.(); }); renderWithProviders(); - const deleteButtons = screen.getAllByRole("button", { - name: "Delete access group", - }); - await user.click(deleteButtons[0]); + await user.click(await openRowMenu(user, "ag-1")); const dialog = screen.getByRole("dialog", { name: "Delete Access Group" }); - const deleteConfirmButton = within(dialog).getByRole("button", { name: /delete/i }); - await user.click(deleteConfirmButton); + await user.click(within(dialog).getByRole("button", { name: /delete/i })); expect(mockMutate).toHaveBeenCalledWith("ag-1", expect.any(Object)); }); - it("should display pagination with total count", () => { - renderWithProviders(); - expect(screen.getByText("2 groups")).toBeInTheDocument(); - }); - - it("should show table headers for ID, Name, Resources, and Actions", () => { - renderWithProviders(); - expect(screen.getByRole("columnheader", { name: /ID/i })).toBeInTheDocument(); - expect(screen.getByRole("columnheader", { name: /Name/i })).toBeInTheDocument(); - expect(screen.getByRole("columnheader", { name: /Resources/i })).toBeInTheDocument(); - expect(screen.getByRole("columnheader", { name: /Actions/i })).toBeInTheDocument(); - }); - - it("should display loading state when data is loading", () => { - mockUseAccessGroups.mockReturnValue({ - data: undefined, - isLoading: true, - }); - renderWithProviders(); - const table = screen.getByRole("table"); - expect(table).toBeInTheDocument(); - }); - - it("should display empty state when no groups match search", async () => { + it("still shows matches when searching from a later page", async () => { const user = userEvent.setup(); + mockUseAccessGroups.mockReturnValue({ data: makeGroups(25), isLoading: false }); renderWithProviders(); - const searchInput = screen.getByPlaceholderText("Search groups by name, ID, or description..."); - await user.type(searchInput, "nonexistent-group-xyz"); - expect(screen.getByRole("table")).toBeInTheDocument(); + + await user.click(screen.getByTestId("pagination-next")); + expect(screen.getByText("ag-11")).toBeInTheDocument(); + expect(screen.queryByText("ag-01")).not.toBeInTheDocument(); + + // The only match lives on page 1, so the page index must reset or the table reads as empty. + await user.type(screen.getByPlaceholderText("Search groups by name, ID, or description..."), "ag-01"); + expect(await screen.findByText("ag-01")).toBeInTheDocument(); + expect(screen.queryByText("No matching access groups")).not.toBeInTheDocument(); }); - it("should display empty data when useAccessGroups returns empty array", () => { - mockUseAccessGroups.mockReturnValue({ - data: [], - isLoading: false, - }); + it("hides the Create button and row actions for a non-admin", () => { + mockUseAuthorized.mockReturnValue({ userRole: "Admin Viewer", accessToken: "sk-test" }); renderWithProviders(); - expect(screen.getByRole("table")).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /create access group/i })).not.toBeInTheDocument(); + expect(screen.queryByTestId("access-group-actions-ag-1")).not.toBeInTheDocument(); + // The read-only view still lists the groups. + expect(screen.getByText("Admin Group")).toBeInTheDocument(); }); }); 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 dbbf4e35900..0de6596f57c 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 @@ -1,38 +1,17 @@ import { AccessGroupResponse, useAccessGroups } from "@/app/(dashboard)/hooks/accessGroups/useAccessGroups"; import { useDeleteAccessGroup } from "@/app/(dashboard)/hooks/accessGroups/useDeleteAccessGroup"; import { PlusOutlined } from "@ant-design/icons"; -import { - ColumnDef, - flexRender, - getCoreRowModel, - getSortedRowModel, - Row, - SortingState, - useReactTable, -} from "@tanstack/react-table"; -import { Button, Card, Flex, Input, Layout, Pagination, Space, Table, Tag, theme, Tooltip, Typography } from "antd"; -import { BotIcon, LayersIcon, SearchIcon, ServerIcon } from "lucide-react"; -import { useEffect, useMemo, useState } from "react"; +import { Button, Flex, Input, Layout, Space, theme, Typography } from "antd"; +import { SearchIcon } from "lucide-react"; +import { useMemo, useState } from "react"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; -import TableIconActionButton from "@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; -import { - SortState, - TableHeaderSortDropdown, -} from "@/components/common_components/TableHeaderSortDropdown/TableHeaderSortDropdown"; -import { DateCell, IdCell } from "@/components/shared/table_cells"; import { AccessGroupDetail } from "./AccessGroupsDetailsPage"; import { AccessGroupCreateModal } from "./AccessGroupsModal/AccessGroupCreateModal"; +import { AccessGroupsTable } from "./AccessGroupsTable"; import { AccessGroup } from "./types"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { isProxyAdminRole } from "@/utils/roles"; -declare module "@tanstack/react-table" { - // eslint-disable-next-line @typescript-eslint/no-unused-vars - interface ColumnMeta { - responsive?: string[]; - } -} - const { Title, Text } = Typography; const { Content } = Layout; @@ -52,55 +31,6 @@ function mapResponseToAccessGroup(r: AccessGroupResponse): AccessGroup { updatedBy: r.updated_by ?? "", }; } -function buildAntdColumns( - table: ReturnType>, - rowLookup: Map>, - onSortingChange: (s: SortingState) => void, -) { - const headers = table.getHeaderGroups()[0]?.headers ?? []; - - return headers.map((header) => { - const canSort = header.column.getCanSort(); - const isSorted = header.column.getIsSorted(); - const meta = header.column.columnDef.meta as { responsive?: string[] } | undefined; - - const col: Record = { - title: ( -
- {header.isPlaceholder ? null : flexRender(header.column.columnDef.header, header.getContext())} - {canSort && ( - { - if (newState === false) { - onSortingChange([]); - } else { - onSortingChange([{ id: header.column.id, desc: newState === "desc" }]); - } - }} - columnId={header.column.id} - /> - )} -
- ), - key: header.id, - width: header.column.columnDef.size, - render: (_: unknown, record: AccessGroup) => { - const row = rowLookup.get(record.id); - if (!row) return null; - const cell = row.getVisibleCells().find((c) => c.column.id === header.id); - if (!cell) return null; - return flexRender(cell.column.columnDef.cell, cell.getContext()); - }, - }; - - if (meta?.responsive) { - col.responsive = meta.responsive; - } - - return col; - }); -} export function AccessGroupsPage() { const { token } = theme.useToken(); @@ -113,151 +43,19 @@ export function AccessGroupsPage() { const [selectedGroupId, setSelectedGroupId] = useState(null); const [isCreateModalVisible, setIsCreateModalVisible] = useState(false); const [searchText, setSearchText] = useState(""); - const [currentPage, setCurrentPage] = useState(1); - const [sorting, setSorting] = useState([]); const [groupToDelete, setGroupToDelete] = useState(null); const deleteMutation = useDeleteAccessGroup(); - const pageSize = 10; - useEffect(() => { - setCurrentPage(1); - }, [searchText]); - - // ---------- filtered data ---------- - const filteredGroups = useMemo( - () => - groups.filter( - (group) => - group.name.toLowerCase().includes(searchText.toLowerCase()) || - group.id.toLowerCase().includes(searchText.toLowerCase()) || - group.description.toLowerCase().includes(searchText.toLowerCase()), - ), - [groups, searchText], - ); - - // ---------- TanStack column definitions ---------- - const columnDefs = useMemo[]>( - () => [ - { - id: "id", - accessorKey: "id", - header: () => ID, - enableSorting: false, - size: 170, - cell: ({ row }) => , - }, - { - id: "name", - accessorKey: "name", - header: () => Name, - enableSorting: true, - cell: ({ getValue }) => getValue() as string, - }, - { - id: "resources", - header: () => Resources, - enableSorting: false, - cell: ({ row }) => { - const record = row.original; - const modelIds = record.modelIds ?? []; - const mcpServerIds = record.mcpServerIds ?? []; - const agentIds = record.agentIds ?? []; - return ( - - - - - - {modelIds?.length} - - - - - - - - {mcpServerIds?.length} - - - - - - - - {agentIds?.length} - - - - - ); - }, - }, - { - id: "createdAt", - accessorKey: "createdAt", - header: () => Created, - enableSorting: true, - sortingFn: "datetime", - cell: ({ getValue }) => , - meta: { responsive: ["lg"] }, - }, - { - id: "updatedAt", - accessorKey: "updatedAt", - header: () => Updated, - enableSorting: false, - cell: ({ getValue }) => , - meta: { responsive: ["xl"] }, - }, - ...(canModify - ? [ - { - id: "actions", - header: () => Actions, - enableSorting: false, - cell: ({ row }: { row: Row }) => ( - - setGroupToDelete(row.original)} - /> - - ), - }, - ] - : []), - ], - // setSelectedGroup is stable (useState setter) - // eslint-disable-next-line react-hooks/exhaustive-deps - [canModify], - ); - - // ---------- TanStack table instance ---------- - const table = useReactTable({ - data: filteredGroups, - columns: columnDefs, - state: { sorting }, - onSortingChange: setSorting, - getCoreRowModel: getCoreRowModel(), - getSortedRowModel: getSortedRowModel(), - getRowId: (row) => row.id, - }); - - // All sorted rows from TanStack - const sortedRows = table.getRowModel().rows; - - // Paginated slice - const paginatedRows = sortedRows.slice((currentPage - 1) * pageSize, currentPage * pageSize); - - // Map for O(1) lookup by record id in antd render() - const rowLookup = useMemo(() => new Map(paginatedRows.map((row) => [row.original.id, row])), [paginatedRows]); - - // Convert TanStack headers → antd columns - const antdColumns = buildAntdColumns(table, rowLookup, setSorting); - - // antd dataSource (just the originals for the current page) - const dataSource = paginatedRows.map((row) => row.original); + const filteredGroups = useMemo(() => { + const query = searchText.trim().toLowerCase(); + if (!query) return groups; + return groups.filter( + (group) => + group.name.toLowerCase().includes(query) || + group.id.toLowerCase().includes(query) || + group.description.toLowerCase().includes(query), + ); + }, [groups, searchText]); if (selectedGroupId) { return setSelectedGroupId(null)} />; @@ -279,34 +77,25 @@ export function AccessGroupsPage() { )} - - - } - placeholder="Search groups by name, ID, or description..." - style={{ maxWidth: 400 }} - value={searchText} - onChange={(e) => setSearchText(e.target.value)} - allowClear - /> - setCurrentPage(page)} - size="small" - showTotal={(total) => `${total} groups`} - showSizeChanger={false} - /> - - - + + } + placeholder="Search groups by name, ID, or description..." + style={{ maxWidth: 400 }} + value={searchText} + onChange={(e) => setSearchText(e.target.value)} + allowClear + /> + + + 0} + canModify={canModify} + onGroupClick={setSelectedGroupId} + onDeleteClick={setGroupToDelete} + /> setIsCreateModalVisible(false)} /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsTable.tsx new file mode 100644 index 00000000000..10d1735d3e7 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsTable.tsx @@ -0,0 +1,72 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Layers } from "lucide-react"; +import { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; + +import { getAccessGroupsTableColumns } from "./AccessGroupsTableColumns"; +import { AccessGroup } from "./types"; + +interface AccessGroupsTableProps { + groups: AccessGroup[]; + isLoading: boolean; + isFiltered: boolean; + canModify: boolean; + onGroupClick: (id: string) => void; + onDeleteClick: (group: AccessGroup) => void; +} + +const PAGE_SIZE_OPTIONS = [10, 25, 50]; + +function EmptyState({ isFiltered }: { isFiltered: boolean }) { + return ( +
+
+ +
+
+ {isFiltered ? "No matching access groups" : "No access groups yet"} +
+
+ {isFiltered + ? "Try a different search term." + : "Create an access group to manage resource permissions for your organization."} +
+
+ ); +} + +export function AccessGroupsTable({ + groups, + isLoading, + isFiltered, + canModify, + onGroupClick, + onDeleteClick, +}: AccessGroupsTableProps) { + const [sorting, setSorting] = useState([]); + + const columns = useMemo(() => { + const deps = { canModify, onGroupClick, onDeleteClick }; + return getAccessGroupsTableColumns(deps); + }, [canModify, onGroupClick, onDeleteClick]); + + return ( + group.id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + paginationMode="client" + pageSizeOptions={PAGE_SIZE_OPTIONS} + isLoading={isLoading} + loadingMessage="Loading access groups…" + noDataMessage={} + size="compact" + /> + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsTableColumns.tsx new file mode 100644 index 00000000000..ae65f161b1e --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/access-groups/_components/AccessGroupsTableColumns.tsx @@ -0,0 +1,182 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Bot, Layers, MoreHorizontal, Server, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell } from "@/components/shared/table_cells"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +import { AccessGroup } from "./types"; + +interface ResourceTone { + icon: typeof Layers; + className: string; +} + +const RESOURCE_TONES: Record<"models" | "mcpServers" | "agents", ResourceTone> = { + models: { icon: Layers, className: "bg-blue-50 text-blue-700 ring-blue-600/20" }, + mcpServers: { icon: Server, className: "bg-cyan-50 text-cyan-700 ring-cyan-600/20" }, + agents: { icon: Bot, className: "bg-purple-50 text-purple-700 ring-purple-600/20" }, +}; + +function ResourcesCell({ group }: { group: AccessGroup }) { + const items = [ + { key: "models" as const, label: "Models", count: group.modelIds.length }, + { key: "mcpServers" as const, label: "MCP Servers", count: group.mcpServerIds.length }, + { key: "agents" as const, label: "Agents", count: group.agentIds.length }, + ]; + + return ( +
+ {items.map((item) => { + const tone = RESOURCE_TONES[item.key]; + const Icon = tone.icon; + return ( + + + {item.count} + + ); + })} +
+ ); +} + +function AccessGroupRowActions({ + group, + onDeleteClick, +}: { + group: AccessGroup; + onDeleteClick: (group: AccessGroup) => void; +}) { + return ( + + + + + + onDeleteClick(group)} + > + + Delete access group + + + + ); +} + +interface AccessGroupsTableColumnsDeps { + canModify: boolean; + onGroupClick: (id: string) => void; + onDeleteClick: (group: AccessGroup) => void; +} + +export const getAccessGroupsTableColumns = ({ + canModify, + onGroupClick, + onDeleteClick, +}: AccessGroupsTableColumnsDeps): ColumnDef[] => { + const columns: ColumnDef[] = [ + { + id: "id", + accessorKey: "id", + meta: { title: "ID" }, + header: "ID", + size: 200, + enableSorting: false, + cell: ({ row }) => ( + onGroupClick(row.original.id)} + /> + ), + }, + { + id: "name", + accessorKey: "name", + meta: { title: "Name" }, + header: ({ column }) => , + size: 220, + enableSorting: true, + cell: ({ row }) => { + const name = row.original.name; + return ( + + {name || "-"} + + ); + }, + }, + { + id: "resources", + meta: { title: "Resources" }, + header: "Resources", + size: 220, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "createdAt", + accessorKey: "createdAt", + meta: { title: "Created" }, + header: ({ column }) => , + size: 150, + enableSorting: true, + sortingFn: "datetime", + cell: ({ row }) => , + }, + { + id: "updatedAt", + accessorKey: "updatedAt", + meta: { title: "Updated" }, + header: "Updated", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + ]; + + if (!canModify) { + return columns; + } + + return [ + ...columns, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, + ]; +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx index 48674f21883..441d300436a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.test.tsx @@ -1,12 +1,13 @@ import React from "react"; -import { render, screen, waitFor, act, fireEvent, within } from "@testing-library/react"; +import { act, render, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { describe, it, expect, vi, beforeEach } from "vitest"; import AgentsPanel from "./AgentsPanel"; import * as networking from "@/components/networking"; vi.mock("@/components/networking", () => ({ getAgentsList: vi.fn().mockResolvedValue({ agents: [] }), - deleteAgentCall: vi.fn(), + deleteAgentCall: vi.fn().mockResolvedValue({}), })); vi.mock("./add_agent_form", () => ({ @@ -19,56 +20,54 @@ vi.mock("./agent_info", () => ({ describe("AgentsPanel", () => { beforeEach(() => { - vi.clearAllMocks(); + // mockReset (not mockClear) so an unconsumed *Once queue cannot leak into the next test + vi.mocked(networking.getAgentsList).mockReset().mockResolvedValue({ agents: [] }); + vi.mocked(networking.deleteAgentCall).mockReset().mockResolvedValue({}); }); - it("should render the Agents panel title", async () => { + it("should render the Agents panel title", () => { render(); expect(screen.getByText("Agents")).toBeInTheDocument(); }); - it("should show Add New Agent button for admin users", async () => { + it("should show Add New Agent button for admin users", () => { render(); - expect(screen.getByText("+ Add New Agent")).toBeInTheDocument(); + expect(screen.getByText("Add New Agent")).toBeInTheDocument(); }); - it("should show Add New Agent button for proxy_admin users", async () => { + it("should show Add New Agent button for proxy_admin users", () => { render(); - expect(screen.getByText("+ Add New Agent")).toBeInTheDocument(); + expect(screen.getByText("Add New Agent")).toBeInTheDocument(); }); - it("should not show Add New Agent button for internal_user role", async () => { + it("should not show Add New Agent button for internal_user role", () => { render(); - expect(screen.queryByText("+ Add New Agent")).not.toBeInTheDocument(); + expect(screen.queryByText("Add New Agent")).not.toBeInTheDocument(); }); - it("should not show Add New Agent button for internal_user_viewer role", async () => { + it("should not show Add New Agent button for internal_user_viewer role", () => { render(); - expect(screen.queryByText("+ Add New Agent")).not.toBeInTheDocument(); + expect(screen.queryByText("Add New Agent")).not.toBeInTheDocument(); }); - it("should show Actions column header for admin role", async () => { + it("should show the Actions column for admin role", async () => { render(); - await waitFor(() => { - expect(screen.getByRole("columnheader", { name: /actions/i })).toBeInTheDocument(); - }); + expect(await screen.findByRole("columnheader", { name: /actions/i })).toBeInTheDocument(); }); - it("should not show Actions column header for internal user role", async () => { + it("should not show the Actions column for internal user role", async () => { render(); await waitFor(() => { expect(screen.queryByRole("columnheader", { name: /actions/i })).not.toBeInTheDocument(); - // confirm table is rendered (not still loading) expect(screen.getByRole("table")).toBeInTheDocument(); }); }); - it("should render the Health Check toggle", async () => { - render(); + it("should render the Health Check toggle for admins and non-admins", () => { + const { unmount } = render(); expect(screen.getByText("Health Check")).toBeInTheDocument(); - }); + unmount(); - it("should render the Health Check toggle for non-admin users too", async () => { render(); expect(screen.getByText("Health Check")).toBeInTheDocument(); }); @@ -108,19 +107,187 @@ describe("AgentsPanel", () => { expect(within(keylessRow).getByText("Needs Setup")).toBeInTheDocument(); }); - it("should call getAgentsList with health_check=true when toggle is enabled", async () => { + it("should refetch with health_check=true when the toggle is enabled", async () => { + const user = userEvent.setup(); render(); await waitFor(() => { expect(networking.getAgentsList).toHaveBeenCalledWith("test-token", false); }); - const toggle = screen.getByRole("switch"); - await act(async () => { - fireEvent.click(toggle); - }); + await user.click(screen.getByRole("switch")); await waitFor(() => { expect(networking.getAgentsList).toHaveBeenCalledWith("test-token", true); }); }); + + it("should delete an agent through the ⋯ menu and confirm modal, then refetch", async () => { + const user = userEvent.setup(); + vi.mocked(networking.getAgentsList).mockResolvedValue({ + agents: [ + { + agent_id: "agent-9", + agent_name: "Doomed Agent", + litellm_params: { model: "gpt-4" }, + spend: 0, + keys: [], + }, + ], + }); + + render(); + + await user.click(await screen.findByTestId("agent-actions-agent-9")); + await user.click(await screen.findByTestId("agent-action-delete")); + + const modal = await screen.findByRole("dialog"); + await user.click(within(modal).getByRole("button", { name: /^delete$/i })); + + await waitFor(() => { + expect(networking.deleteAgentCall).toHaveBeenCalledWith("test-token", "agent-9"); + }); + // one initial load + one post-delete refetch + await waitFor(() => { + expect(vi.mocked(networking.getAgentsList).mock.calls.length).toBeGreaterThanOrEqual(2); + }); + }); + + it("should show a loading skeleton on initial load and clear it once agents arrive", async () => { + render(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + await waitFor(() => { + expect(screen.queryByTestId("skeleton-row")).not.toBeInTheDocument(); + }); + }); + + it("should clear the loading state when there is no access token rather than skeleton forever", async () => { + render(); + await waitFor(() => { + expect(screen.queryByTestId("skeleton-row")).not.toBeInTheDocument(); + }); + expect(screen.getByText("No agents yet")).toBeInTheDocument(); + expect(networking.getAgentsList).not.toHaveBeenCalled(); + }); + + it("should not show rows fetched with a previous access token after the token changes", async () => { + const agentFor = (name: string) => ({ + agent_id: `id-${name}`, + agent_name: name, + litellm_params: { model: "gpt-4" }, + spend: 0, + keys: [], + }); + let resolveSecond: (value: { agents: ReturnType[] }) => void = () => {}; + vi.mocked(networking.getAgentsList) + .mockResolvedValueOnce({ agents: [agentFor("first-token-agent")] }) + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveSecond = resolve; + }), + ); + + const { rerender } = render(); + expect(await screen.findByText("first-token-agent")).toBeInTheDocument(); + + rerender(); + + // the previous token's rows must not linger while the new token loads + expect(screen.queryByText("first-token-agent")).not.toBeInTheDocument(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + + await act(async () => { + resolveSecond({ agents: [agentFor("second-token-agent")] }); + }); + expect(await screen.findByText("second-token-agent")).toBeInTheDocument(); + }); + + it("should drop previous rows when the fetch for a new token fails", async () => { + vi.mocked(networking.getAgentsList) + .mockResolvedValueOnce({ + agents: [ + { agent_id: "stale", agent_name: "Stale Agent", litellm_params: { model: "gpt-4" }, spend: 0, keys: [] }, + ], + }) + .mockRejectedValueOnce(new Error("unauthorized")); + + const { rerender } = render(); + expect(await screen.findByText("Stale Agent")).toBeInTheDocument(); + + rerender(); + + await waitFor(() => { + expect(screen.getByText("No agents yet")).toBeInTheDocument(); + }); + expect(screen.queryByText("Stale Agent")).not.toBeInTheDocument(); + }); + + it("should ignore a superseded response so it cannot overwrite the current token's rows", async () => { + let resolveFirst: (value: { + agents: { agent_id: string; agent_name: string; litellm_params: { model: string }; spend: number; keys: [] }[]; + }) => void = () => {}; + vi.mocked(networking.getAgentsList) + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveFirst = resolve; + }), + ) + .mockResolvedValueOnce({ + agents: [ + { agent_id: "current", agent_name: "Current Agent", litellm_params: { model: "gpt-4" }, spend: 0, keys: [] }, + ], + }); + + const { rerender } = render(); + rerender(); + + expect(await screen.findByText("Current Agent")).toBeInTheDocument(); + + // the slow token-a response lands last and must be discarded + await act(async () => { + resolveFirst({ + agents: [ + { agent_id: "stale", agent_name: "Superseded Agent", litellm_params: { model: "gpt-4" }, spend: 0, keys: [] }, + ], + }); + }); + + expect(screen.queryByText("Superseded Agent")).not.toBeInTheDocument(); + expect(screen.getByText("Current Agent")).toBeInTheDocument(); + }); + + it("should keep rows visible during a health-check refetch instead of re-showing the skeleton", async () => { + const user = userEvent.setup(); + const agents = [ + { + agent_id: "agent-1", + agent_name: "Stable Agent", + litellm_params: { model: "gpt-4" }, + spend: 0, + keys: [], + }, + ]; + let resolveRefetch: (value: { agents: typeof agents }) => void = () => {}; + vi.mocked(networking.getAgentsList) + .mockResolvedValueOnce({ agents }) + .mockImplementationOnce( + () => + new Promise((resolve) => { + resolveRefetch = resolve; + }), + ); + + render(); + expect(await screen.findByText("Stable Agent")).toBeInTheDocument(); + + await user.click(screen.getByRole("switch")); + + expect(screen.getByText("Stable Agent")).toBeInTheDocument(); + expect(screen.queryByTestId("skeleton-row")).not.toBeInTheDocument(); + + await act(async () => { + resolveRefetch({ agents }); + }); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx index 84634620426..a4a71530c84 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsPanel.tsx @@ -1,27 +1,15 @@ import React, { useState, useEffect } from "react"; -import { - Button, - Card, - Table, - TableBody, - TableCell, - TableHead, - TableHeaderCell, - TableRow, - Badge, - Text, -} from "@tremor/react"; -import { Modal, Alert, Tooltip, Skeleton, Switch } from "antd"; -import { CheckCircleOutlined } from "@ant-design/icons"; +import { Modal, Alert } from "antd"; +import { Plus } from "lucide-react"; import { getAgentsList, deleteAgentCall } from "@/components/networking"; import AddAgentForm from "./add_agent_form"; import { isAdminRole } from "@/utils/roles"; import AgentInfoView from "./agent_info"; +import AgentsTable from "./AgentsTable"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { Agent } from "@/components/agents/types"; import { Team } from "@/components/key_team_helpers/key_list"; -import { DateCell, IdCell, MoneyCell, StatusBadge } from "@/components/shared/table_cells"; -import TableIconActionButton from "@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; +import { Button } from "@/components/ui/button"; interface AgentsPanelProps { accessToken: string | null; @@ -36,37 +24,66 @@ interface AgentsResponse { const AgentsPanel: React.FC = ({ accessToken, userRole, teams }) => { const [agentsList, setAgentsList] = useState([]); const [isAddModalVisible, setIsAddModalVisible] = useState(false); - const [isLoading, setIsLoading] = useState(false); + const [isLoading, setIsLoading] = useState(true); const [isDeleting, setIsDeleting] = useState(false); + const [isHealthCheckLoading, setIsHealthCheckLoading] = useState(false); const [agentToDelete, setAgentToDelete] = useState<{ id: string; name: string } | null>(null); const [selectedAgentId, setSelectedAgentId] = useState(null); const [healthCheckEnabled, setHealthCheckEnabled] = useState(false); const isAdmin = userRole ? isAdminRole(userRole) : false; - const fetchAgents = async (healthCheck?: boolean) => { + useEffect(() => { + let cancelled = false; + const loadForToken = async () => { + if (!accessToken) { + setAgentsList([]); + setIsLoading(false); + return; + } + setIsLoading(true); + try { + const response: AgentsResponse = await getAgentsList(accessToken, false); + if (!cancelled) { + setAgentsList(response.agents || []); + } + } catch (error) { + console.error("Error fetching agents:", error); + if (!cancelled) { + setAgentsList([]); + } + } finally { + if (!cancelled) { + setIsLoading(false); + } + } + }; + loadForToken(); + return () => { + cancelled = true; + }; + }, [accessToken]); + + const refetchAgents = async (healthCheck: boolean) => { if (!accessToken) { return; } - - setIsLoading(true); try { - const response: AgentsResponse = await getAgentsList(accessToken, healthCheck ?? healthCheckEnabled); + const response: AgentsResponse = await getAgentsList(accessToken, healthCheck); setAgentsList(response.agents || []); } catch (error) { console.error("Error fetching agents:", error); - } finally { - setIsLoading(false); } }; - useEffect(() => { - fetchAgents(); - }, [accessToken]); - - const handleHealthCheckToggle = (checked: boolean) => { + const handleHealthCheckToggle = async (checked: boolean) => { setHealthCheckEnabled(checked); - fetchAgents(checked); + setIsHealthCheckLoading(true); + try { + await refetchAgents(checked); + } finally { + setIsHealthCheckLoading(false); + } }; const handleAddAgent = () => { @@ -81,7 +98,7 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams }; const handleSuccess = () => { - fetchAgents(); + refetchAgents(healthCheckEnabled); }; const handleDeleteClick = (agentId: string, agentName: string) => { @@ -95,7 +112,7 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams try { await deleteAgentCall(accessToken, agentToDelete.id); NotificationsManager.success(`Agent "${agentToDelete.name}" deleted successfully`); - fetchAgents(); + await refetchAgents(healthCheckEnabled); } catch (error) { console.error("Error deleting agent:", error); NotificationsManager.fromBackend("Failed to delete agent"); @@ -109,14 +126,6 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams setAgentToDelete(null); }; - const sortedAgents = [...agentsList].sort((a, b) => { - const dateA = a.created_at ? new Date(a.created_at).getTime() : 0; - const dateB = b.created_at ? new Date(b.created_at).getTime() : 0; - return dateB - dateA; - }); - - const columnCount = isAdmin ? 7 : 6; - return (
@@ -132,25 +141,14 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams showIcon className="mb-3" /> -
- {isAdmin && ( + {isAdmin && ( +
- )} - -
- - Health Check - -
-
-
+
+ )}
{selectedAgentId ? ( @@ -161,73 +159,16 @@ const AgentsPanel: React.FC = ({ accessToken, userRole, teams isAdmin={isAdmin} /> ) : ( - - {isLoading ? ( - - ) : ( -
- - - Agent Name - Agent ID - Spend (USD) - Model - Created - Status - {isAdmin && Actions} - - - - {sortedAgents.length === 0 ? ( - - - - No agents found. Click "+ Add New Agent" to create one. - - - - ) : ( - sortedAgents.map((agent) => ( - - - {agent.agent_name} - - - setSelectedAgentId(id)} /> - - - - - - - {agent.litellm_params?.model || "N/A"} - - - - - - - {(agent.keys?.length ?? 0) > 0 ? ( - - ) : ( - - )} - - {isAdmin && ( - - handleDeleteClick(agent.agent_id, agent.agent_name)} - /> - - )} - - )) - )} - -
- )} -
+ setSelectedAgentId(id)} + onDeleteClick={handleDeleteClick} + /> )} = {}): Agent => ({ + agent_id: "agent-1", + agent_name: "Test Agent", + litellm_params: { model: "gpt-4" }, + spend: 0, + keys: [{ token: "hash-1", key_alias: "primary", key_name: "sk-...1" }], + created_at: "2023-01-01T00:00:00Z", + ...overrides, +}); + +describe("AgentsTable", () => { + it("renders every column header", () => { + render(); + for (const header of ["Agent Name", "Agent ID", "Spend (USD)", "Model", "Created", "Status"]) { + expect(screen.getByText(header)).toBeInTheDocument(); + } + }); + + it("renders the agent's model and opens the detail view when the ID cell is clicked", async () => { + const user = userEvent.setup(); + const onAgentClick = vi.fn(); + const agent = makeAgent({ agent_id: "agent-xyz", agent_name: "Router", litellm_params: { model: "claude-3-5" } }); + render(); + + expect(screen.getByText("claude-3-5")).toBeInTheDocument(); + + await user.click(screen.getByText("agent-xyz")); + expect(onAgentClick).toHaveBeenCalledWith("agent-xyz"); + }); + + it("marks agents Active when they have keys and Needs Setup when they have none", () => { + render( + , + ); + + const keyedRow = screen.getByText("Keyed Agent").closest("tr")!; + const keylessRow = screen.getByText("Keyless Agent").closest("tr")!; + expect(within(keyedRow).getByText("Active")).toBeInTheDocument(); + expect(within(keylessRow).getByText("Needs Setup")).toBeInTheDocument(); + }); + + it("deletes an agent through the ⋯ actions menu", async () => { + const user = userEvent.setup(); + const onDeleteClick = vi.fn(); + const agent = makeAgent({ agent_id: "agent-9", agent_name: "Doomed Agent" }); + render(); + + await user.click(screen.getByTestId("agent-actions-agent-9")); + await user.click(await screen.findByTestId("agent-action-delete")); + + expect(onDeleteClick).toHaveBeenCalledWith("agent-9", "Doomed Agent"); + }); + + it("hides the actions column entirely for non-admins", () => { + const agent = makeAgent({ agent_id: "agent-2" }); + render(); + + expect(screen.queryByTestId("agent-actions-agent-2")).not.toBeInTheDocument(); + expect(screen.queryByRole("columnheader", { name: /actions/i })).not.toBeInTheDocument(); + expect(screen.getByRole("table")).toBeInTheDocument(); + }); + + it("shows the actions column for admins", () => { + render(); + expect(screen.getByRole("columnheader", { name: /actions/i })).toBeInTheDocument(); + expect(screen.getByTestId("agent-actions-agent-3")).toBeInTheDocument(); + }); + + it("defaults to sorting by created_at descending (newest first)", () => { + render( + , + ); + + const bodyRows = screen.getAllByRole("row").slice(1); + expect(bodyRows[0].textContent).toContain("Beta Agent"); + expect(bodyRows[1].textContent).toContain("Alpha Agent"); + }); + + it("sorts agents with no created_at last, never ahead of dated ones", () => { + render( + , + ); + + const bodyRows = screen.getAllByRole("row").slice(1); + expect(bodyRows[0].textContent).toContain("Beta Agent"); + expect(bodyRows[1].textContent).toContain("Alpha Agent"); + expect(bodyRows[2].textContent).toContain("Undated Agent"); + }); + + it("shows a rich empty state when there are no agents", () => { + render(); + expect(screen.getByText("No agents yet")).toBeInTheDocument(); + expect(screen.queryByTestId("skeleton-row")).not.toBeInTheDocument(); + }); + + it("renders loading skeleton rows on initial load instead of the empty state", () => { + render(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("No agents yet")).not.toBeInTheDocument(); + }); + + it("invokes the health-check toggle from the toolbar", async () => { + const user = userEvent.setup(); + const onHealthCheckToggle = vi.fn(); + render(); + + expect(screen.getByText("Health Check")).toBeInTheDocument(); + await user.click(screen.getByRole("switch")); + expect(onHealthCheckToggle).toHaveBeenCalledWith(true, expect.anything()); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx new file mode 100644 index 00000000000..824ae47f3e6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTable.tsx @@ -0,0 +1,88 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Tooltip, Switch } from "antd"; +import { CheckCircleOutlined } from "@ant-design/icons"; +import { Bot } from "lucide-react"; +import React, { useMemo, useState } from "react"; + +import { Agent } from "@/components/agents/types"; +import { DataTable } from "@/components/shared/DataTable"; + +import { getAgentsTableColumns } from "./AgentsTableColumns"; + +interface AgentsTableProps { + agents: Agent[]; + isLoading: boolean; + isAdmin: boolean; + healthCheckEnabled: boolean; + isHealthCheckLoading: boolean; + onHealthCheckToggle: (checked: boolean) => void; + onAgentClick: (agentId: string) => void; + onDeleteClick: (agentId: string, agentName: string) => void; +} + +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; + +function EmptyState() { + return ( +
+
+ +
+
No agents yet
+
Add an agent to make it available in your organization.
+
+ ); +} + +const AgentsTable: React.FC = ({ + agents, + isLoading, + isAdmin, + healthCheckEnabled, + isHealthCheckLoading, + onHealthCheckToggle, + onAgentClick, + onDeleteClick, +}) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const columns = useMemo( + () => getAgentsTableColumns({ isAdmin, onAgentClick, onDeleteClick }), + [isAdmin, onAgentClick, onDeleteClick], + ); + + return ( + agent.agent_id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading agents…" + noDataMessage={} + size="compact" + toolbar={() => ( +
+ +
+ + Health Check + +
+
+
+ )} + /> + ); +}; + +export default AgentsTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx new file mode 100644 index 00000000000..a8fe3973a42 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/AgentsTableColumns.tsx @@ -0,0 +1,163 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { MoreHorizontal, Trash2 } from "lucide-react"; + +import { Agent } from "@/components/agents/types"; +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell, MoneyCell, StatusBadge } from "@/components/shared/table_cells"; +import { Badge } from "@/components/ui/badge"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +interface AgentRowActionsProps { + agent: Agent; + onDeleteClick: (agentId: string, agentName: string) => void; +} + +function AgentRowActions({ agent, onDeleteClick }: AgentRowActionsProps) { + return ( + + + + + + onDeleteClick(agent.agent_id, agent.agent_name)} + > + + Delete + + + + ); +} + +interface AgentsTableColumnsDeps { + isAdmin: boolean; + onAgentClick: (agentId: string) => void; + onDeleteClick: (agentId: string, agentName: string) => void; +} + +export const getAgentsTableColumns = ({ + isAdmin, + onAgentClick, + onDeleteClick, +}: AgentsTableColumnsDeps): ColumnDef[] => [ + { + id: "agent_name", + accessorKey: "agent_name", + meta: { title: "Agent Name" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => { + const name = row.original.agent_name; + return ( + + {name || "-"} + + ); + }, + }, + { + id: "agent_id", + accessorKey: "agent_id", + meta: { title: "Agent ID" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => ( + onAgentClick(row.original.agent_id)} + /> + ), + }, + { + id: "spend", + accessorKey: "spend", + meta: { title: "Spend (USD)" }, + header: ({ column }) => , + size: 130, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "model", + meta: { title: "Model" }, + header: "Model", + size: 170, + enableSorting: false, + cell: ({ row }) => { + const model = row.original.litellm_params?.model; + if (!model) { + return N/A; + } + return ( + + + {model} + + + ); + }, + }, + { + id: "created_at", + accessorFn: (agent) => { + const timestamp = agent.created_at ? new Date(agent.created_at).getTime() : 0; + return Number.isNaN(timestamp) ? 0 : timestamp; + }, + meta: { title: "Created" }, + header: ({ column }) => , + size: 150, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "status", + meta: { title: "Status" }, + header: "Status", + size: 130, + enableSorting: false, + cell: ({ row }) => { + const hasKeys = (row.original.keys?.length ?? 0) > 0; + return hasKeys ? ( + + ) : ( + + ); + }, + }, + ...(isAdmin + ? [ + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + } satisfies ColumnDef, + ] + : []), +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx new file mode 100644 index 00000000000..767e7c2ae5f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.test.tsx @@ -0,0 +1,89 @@ +import React from "react"; +import { render, screen, fireEvent, within } from "@testing-library/react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import AddAgentForm from "./add_agent_form"; +import * as networking from "@/components/networking"; +import type { AgentCreateInfo } from "@/components/networking"; + +vi.mock("@/components/networking", () => ({ + createAgentCall: vi.fn(), + getAgentCreateMetadata: vi.fn(), + getAgentsList: vi.fn(), + keyCreateForAgentCall: vi.fn(), + keyListCall: vi.fn(), + keyUpdateCall: vi.fn(), + modelAvailableCall: vi.fn(), +})); + +vi.mock("./agent_card_discovery", () => ({ + default: () =>
, +})); + +vi.mock("./agent_form_fields", () => ({ + default: () =>
, +})); + +const a2aInfo: AgentCreateInfo = { + agent_type: "a2a", + agent_type_display_name: "A2A Agent", + description: "Agent-to-agent protocol", + logo_url: "/ui/assets/logos/a2a_agent.png", + credential_fields: [], + use_a2a_form_fields: true, +}; + +const renderForm = () => + render(); + +describe("AddAgentForm logos", () => { + beforeEach(() => { + vi.mocked(networking.getAgentCreateMetadata).mockReset().mockResolvedValue([a2aInfo]); + vi.mocked(networking.getAgentsList).mockReset().mockResolvedValue({ agents: [] }); + vi.mocked(networking.keyListCall).mockReset().mockResolvedValue({ keys: [] }); + vi.mocked(networking.modelAvailableCall).mockReset().mockResolvedValue({ data: [] }); + }); + + it("renders the modal title and agent type selection logos as images from logo_url", async () => { + renderForm(); + + const titleLogo = await screen.findByAltText("Agent logo"); + expect(titleLogo).toBeInstanceOf(HTMLImageElement); + expect(titleLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png")); + + const selectionLogo = await screen.findByAltText("A2A Agent logo"); + expect(selectionLogo).toBeInstanceOf(HTMLImageElement); + expect(selectionLogo).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png")); + }); + + it("renders the option logo when the agent type dropdown is opened", async () => { + renderForm(); + + await screen.findByAltText("A2A Agent logo"); + fireEvent.mouseDown(screen.getByRole("combobox")); + + const optionLogos = await screen.findAllByAltText("A2A Agent logo"); + expect(optionLogos.length).toBeGreaterThanOrEqual(2); + optionLogos.forEach((img) => { + expect(img).toHaveAttribute("src", expect.stringContaining("assets/logos/a2a_agent.png")); + }); + }); + + it("swaps a failing logo for a letter avatar and warns with the url", async () => { + const warnSpy = vi.spyOn(console, "warn").mockImplementation(() => {}); + renderForm(); + + const titleLogo = await screen.findByAltText("Agent logo"); + const header = screen.getByText("Add New Agent").parentElement!; + fireEvent.error(titleLogo); + + expect(warnSpy).toHaveBeenCalledWith(expect.stringContaining("assets/logos/a2a_agent.png")); + expect(screen.queryByAltText("Agent logo")).not.toBeInTheDocument(); + expect(within(header).getByText("A")).toBeInTheDocument(); + + const selectionLogo = screen.getByAltText("A2A Agent logo"); + fireEvent.error(selectionLogo); + expect(screen.queryByAltText("A2A Agent logo")).not.toBeInTheDocument(); + expect(warnSpy).toHaveBeenCalledTimes(2); + warnSpy.mockRestore(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx index 8ca2b5afe16..e35388b78da 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/agents/_components/add_agent_form.tsx @@ -1,7 +1,7 @@ import React, { useState, useEffect } from "react"; import { Modal, Form, Select, Input, Steps, Radio, Tag, Divider, Switch, InputNumber, Collapse } from "antd"; import MessageManager from "@/components/molecules/message_manager"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import { Button } from "@tremor/react"; import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined, InfoCircleOutlined } from "@ant-design/icons"; import CreatedKeyDisplay from "@/components/shared/CreatedKeyDisplay"; @@ -712,17 +712,13 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok value={info.agent_type} label={
- + {info.agent_type_display_name}
} >
- {info.agent_type_display_name} +
{info.agent_type_display_name}
{info.description &&
{info.description}
} @@ -948,7 +944,7 @@ const AddAgentForm: React.FC = ({ visible, onClose, accessTok title={
{selectedLogo && currentStep < 1 && ( - Agent + )}

Add New Agent

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.test.tsx index 66fa0dfa63f..ad1cf28dc54 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.test.tsx @@ -1,4 +1,5 @@ -import { render } from "@testing-library/react"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import APIReferenceView from "./APIReferenceView"; @@ -44,4 +45,48 @@ describe("APIReferenceView", () => { expect(renderedCode).toContain(apiDocUrl); expect(renderedCode).not.toContain(proxyUrl); }); + + it("renders the page title, blurb and docs link", () => { + render(); + + expect(screen.getByText("OpenAI Compatible Proxy: API Reference")).toBeTruthy(); + expect(screen.getByText(/LiteLLM is OpenAI Compatible/)).toBeTruthy(); + + const docsLink = screen.getByRole("link", { name: /API Reference Docs/ }); + expect(docsLink.getAttribute("href")).toBe("https://docs.litellm.ai/docs/proxy/user_keys"); + expect(docsLink.getAttribute("target")).toBe("_blank"); + }); + + it("exposes the three SDK tabs with the first selected by default", () => { + render(); + + expect(screen.getAllByRole("tab").map((tab) => tab.textContent)).toEqual([ + "OpenAI Python SDK", + "LlamaIndex", + "Langchain Py", + ]); + expect(screen.getAllByRole("tab").map((tab) => tab.getAttribute("aria-selected"))).toEqual([ + "true", + "false", + "false", + ]); + }); + + it.each([ + ["OpenAI Python SDK", "import openai"], + ["LlamaIndex", "from llama_index.llms import AzureOpenAI"], + ["Langchain Py", "from langchain.chat_models import ChatOpenAI"], + ])("selecting %s shows its snippet wired to the base url", async (tabName, marker) => { + const proxyUrl = "https://proxy.litellm.test"; + const user = userEvent.setup(); + render(); + + await user.click(screen.getByRole("tab", { name: tabName })); + + expect(screen.getByRole("tab", { name: tabName }).getAttribute("aria-selected")).toBe("true"); + + const selectedPanel = screen.getByRole("tabpanel"); + expect(selectedPanel.textContent).toContain(marker); + expect(selectedPanel.textContent).toContain(proxyUrl); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.tsx index 333bd1cad13..9342017ed3f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-reference/_components/APIReferenceView.tsx @@ -1,7 +1,7 @@ "use client"; import React from "react"; -import { Text, Tab, TabGroup, TabList, TabPanel, TabPanels, Grid } from "@tremor/react"; import CodeBlock from "@/components/CodeBlock"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import DocLink from "./DocLink"; interface ApiRefProps { @@ -21,33 +21,35 @@ const APIReferenceView: React.FC = ({ proxySettings }) => { } return ( - <> - -
- {/* Header row with Docs link on the right */} -
-

- OpenAI Compatible Proxy: API Reference -

- -
+
+
+ {/* Header row with Docs link on the right */} +
+

OpenAI Compatible Proxy: API Reference

+ +
- - LiteLLM is OpenAI Compatible. This means your API Key works with the OpenAI SDK. Just replace the base_url - to point to your litellm proxy. Example Below{" "} - +

+ LiteLLM is OpenAI Compatible. This means your API Key works with the OpenAI SDK. Just replace the base_url to + point to your litellm proxy. Example Below{" "} +

- - - OpenAI Python SDK - LlamaIndex - Langchain Py - - - - + + + OpenAI Python SDK + + + LlamaIndex + + + Langchain Py + + + + - + /> + - - + - + /> + - - + - - - -
- - + /> + + +
+
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx index 17d14cd7fac..13472a3d1df 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.test.tsx @@ -76,6 +76,22 @@ describe("CacheDashboard cache analytics charts", () => { expect(screen.getByText("Cached Completion Tokens vs Generated Completion Tokens")).toBeInTheDocument(); }); + it("scopes the analytics tab to the response cache, not provider prompt caching", async () => { + renderDashboard(); + + expect(await screen.findByText(/is not shown here/)).toBeInTheDocument(); + expect(screen.getByRole("link", { name: "response cache" })).toHaveAttribute( + "href", + "https://docs.litellm.ai/docs/proxy/caching", + ); + expect(screen.getByRole("link", { name: "prompt caching" })).toHaveAttribute( + "href", + "https://docs.litellm.ai/docs/completion/prompt_caching", + ); + expect(screen.queryByText("Cached Tokens")).not.toBeInTheDocument(); + expect(screen.getAllByText("Cached Completion Tokens").length).toBeGreaterThan(0); + }); + it("renders the requests chart with each category legend-bound to its fill and stacked in order", async () => { renderDashboard(); const { requestsCard } = await findChartCards(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx index b8e8dc8adb1..51c0b85cedb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_dashboard.tsx @@ -282,6 +282,28 @@ const CacheDashboard: React.FC = ({ accessToken, token, userRole + + Analytics for LiteLLM's{" "} + + response cache + {" "} + (e.g. Redis / in-memory): requests answered from cache without calling the LLM provider. Provider-side{" "} + + prompt caching + {" "} + (cached input tokens from Anthropic, OpenAI, etc.) is not shown here; see "Prompt Caching + Metrics" on the Usage page or individual requests in the Logs page. + = ({ accessToken, token, userRole

- Cached Tokens + Cached Completion Tokens

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFieldSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFieldSection.tsx index ced822cd796..96106869009 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFieldSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFieldSection.tsx @@ -10,6 +10,7 @@ interface CacheFieldSectionProps { embeddingModels: EmbeddingModelOption[]; gridCols?: string; headingLevel?: "h4" | "h5"; + configuredSecrets?: ReadonlySet; } const CacheFieldSection: React.FC = ({ @@ -19,6 +20,7 @@ const CacheFieldSection: React.FC = ({ embeddingModels, gridCols = "grid-cols-1 gap-6 sm:grid-cols-2", headingLevel = "h4", + configuredSecrets, }) => { const fields = fieldsForSection(section, redisType); if (fields.length === 0) { @@ -32,7 +34,12 @@ const CacheFieldSection: React.FC = ({ {title}

{fields.map((field) => ( - + ))}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFormField.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFormField.tsx index d92ca302901..dbd8c32d18d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFormField.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/CacheFormField.tsx @@ -7,22 +7,29 @@ export interface EmbeddingModelOption { label: string; } +export const SECRET_ALREADY_SET_PLACEHOLDER = "Already set. Enter a new value to replace it."; + interface CacheFormFieldProps { field: CacheField; embeddingModels: EmbeddingModelOption[]; + isSecretConfigured?: boolean; } -const renderControl = (field: CacheField, embeddingModels: EmbeddingModelOption[]): React.ReactNode => { +const renderControl = ( + field: CacheField, + embeddingModels: EmbeddingModelOption[], + placeholder: string, +): React.ReactNode => { switch (field.type) { case "boolean": return ; case "password": - return ; + return ; case "integer": case "float": - return ; + return ; case "list": - return ; + return ; case "model-select": return ( ; + return ; } }; -const CacheFormField: React.FC = ({ field, embeddingModels }) => ( +const CacheFormField: React.FC = ({ field, embeddingModels, isSecretConfigured = false }) => ( = ({ field, embeddingModels rules={field.rules} valuePropName={field.type === "boolean" ? "checked" : "value"} > - {renderControl(field, embeddingModels)} + {renderControl(field, embeddingModels, isSecretConfigured ? SECRET_ALREADY_SET_PLACEHOLDER : field.helpText)} ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsFields.ts b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsFields.ts index 1f5b566fc5f..e33b525c3ef 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsFields.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsFields.ts @@ -8,6 +8,10 @@ export type CacheSection = "connection" | "cluster" | "sentinel" | "semantic" | export type CacheFieldRule = NonNullable[number]; +// Marker the backend returns for a configured credential and maps back to the +// stored secret on save, so the plaintext never round-trips through the form. +export const REDACTED_VALUE = "***REDACTED***"; + export interface CacheField { readonly name: string; readonly label: string; @@ -17,6 +21,9 @@ export interface CacheField { readonly redisType: RedisType | null; readonly defaultValue?: string | number | boolean; readonly rules?: CacheFieldRule[]; + // Credential field: never prefilled into the form, and dropped from the save + // payload when left untouched so the redacted marker is never persisted. + readonly secret?: boolean; } export const REDIS_TYPES: readonly RedisType[] = ["node", "cluster", "sentinel", "semantic"]; @@ -93,6 +100,7 @@ export const CACHE_FIELDS: readonly CacheField[] = [ helpText: "Full Redis/Valkey connection URL (e.g. redis://:password@host:6379/1). When set, it takes precedence over Host, Port, Password, and Database Index.", redisType: null, + secret: true, }, { name: "host", @@ -128,6 +136,7 @@ export const CACHE_FIELDS: readonly CacheField[] = [ section: "connection", helpText: "Redis server password", redisType: null, + secret: true, }, { name: "username", @@ -170,6 +179,7 @@ export const CACHE_FIELDS: readonly CacheField[] = [ section: "sentinel", helpText: "Password for Redis Sentinel authentication", redisType: "sentinel", + secret: true, }, { name: "similarity_threshold", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.test.ts index 79f28a97842..c530519ee06 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.test.ts @@ -1,5 +1,6 @@ import { describe, it, expect } from "vitest"; -import { buildCachePayload, buildInitialValues, fieldsForSection } from "./cacheSettingsUtils"; +import { buildCachePayload, buildInitialValues, configuredSecretFields, fieldsForSection } from "./cacheSettingsUtils"; +import { REDACTED_VALUE } from "./cacheSettingsFields"; describe("fieldsForSection", () => { it("should only include a redis-type-specific field when that type is selected", () => { @@ -83,4 +84,49 @@ describe("buildCachePayload", () => { const payload = buildCachePayload("node", { sentinel_nodes: '[["localhost",26379]]' }, { forTesting: false }); expect(payload).not.toHaveProperty("sentinel_nodes"); }); + + it("should drop a secret whose value is the redacted marker so it is never persisted", () => { + const payload = buildCachePayload( + "node", + { host: "localhost", password: REDACTED_VALUE, url: REDACTED_VALUE }, + { forTesting: false }, + ); + expect(payload).not.toHaveProperty("password"); + expect(payload).not.toHaveProperty("url"); + expect(payload.host).toBe("localhost"); + }); + + it("should send a real new secret value the admin typed", () => { + const payload = buildCachePayload("node", { password: "brandnewpw" }, { forTesting: false }); + expect(payload.password).toBe("brandnewpw"); + }); +}); + +describe("secret handling", () => { + it("buildInitialValues never prefills a credential, even when the server reports it configured", () => { + const serverValues = { + host: "localhost", + password: REDACTED_VALUE, + url: REDACTED_VALUE, + sentinel_password: REDACTED_VALUE, + }; + const values = buildInitialValues(serverValues); + expect(values.password).toBe(""); + expect(values.url).toBe(""); + expect(values.sentinel_password).toBe(""); + // non-secret fields are still prefilled + expect(values.host).toBe("localhost"); + }); + + it("configuredSecretFields reports which credentials the server marked as set", () => { + const configured = configuredSecretFields({ + password: REDACTED_VALUE, + url: "", + host: "localhost", + }); + expect(configured.has("password")).toBe(true); + expect(configured.has("url")).toBe(false); + // a non-secret field is never reported as a configured secret + expect(configured.has("host")).toBe(false); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.ts b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.ts index 088da21961c..7b9454a37c3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/cacheSettingsUtils.ts @@ -1,4 +1,4 @@ -import { CACHE_FIELDS, CacheField, CacheSection, RedisType } from "./cacheSettingsFields"; +import { CACHE_FIELDS, CacheField, CacheSection, REDACTED_VALUE, RedisType } from "./cacheSettingsFields"; export type CacheFormValue = string | number | boolean | undefined; export type CacheFormValues = Record; @@ -11,7 +11,20 @@ export const isFieldVisible = (field: CacheField, redisType: RedisType): boolean export const fieldsForSection = (section: CacheSection, redisType: RedisType): CacheField[] => CACHE_FIELDS.filter((field) => field.section === section && isFieldVisible(field, redisType)); +const hasValue = (raw: unknown): boolean => raw !== undefined && raw !== null && raw !== ""; + +// Credential fields the server reports as configured (returned as the redacted +// marker). Used to show an "already set" hint without ever holding the secret. +export const configuredSecretFields = (currentValues: Record): ReadonlySet => + new Set(CACHE_FIELDS.filter((field) => field.secret && hasValue(currentValues[field.name])).map((f) => f.name)); + const initialValueForField = (field: CacheField, raw: unknown): CacheFormValue => { + // Never prefill a credential: the server sends the redacted marker for a + // configured secret, and echoing it back would persist the marker. + if (field.secret) { + return ""; + } + const source = raw ?? field.defaultValue; if (field.type === "boolean") { @@ -35,6 +48,11 @@ export const buildInitialValues = (currentValues: Record): Cach Object.fromEntries(CACHE_FIELDS.map((field) => [field.name, initialValueForField(field, currentValues[field.name])])); const saveValueForField = (field: CacheField, raw: CacheFormValue): CacheSavePayloadValue | undefined => { + // A redacted secret echoed back untouched must never be persisted as a value. + if (field.secret && raw === REDACTED_VALUE) { + return undefined; + } + if (field.type === "boolean") { return Boolean(raw); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/index.tsx b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/index.tsx index 4382769ae9c..fea2c04015b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/index.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/caching/_components/cache_settings/index.tsx @@ -8,7 +8,7 @@ import RedisTypeSelector from "./RedisTypeSelector"; import CacheFieldSection from "./CacheFieldSection"; import { EmbeddingModelOption } from "./CacheFormField"; import { REDIS_TYPES, REDIS_TYPE_DESCRIPTIONS, RedisType } from "./cacheSettingsFields"; -import { buildCachePayload, buildInitialValues, CacheFormValues } from "./cacheSettingsUtils"; +import { buildCachePayload, buildInitialValues, CacheFormValues, configuredSecretFields } from "./cacheSettingsUtils"; interface CacheSettingsProps { accessToken: string | null; @@ -25,6 +25,7 @@ const CacheSettings: React.FC = ({ accessToken }) => { const [embeddingModels, setEmbeddingModels] = useState([]); const [isTesting, setIsTesting] = useState(false); const [isSaving, setIsSaving] = useState(false); + const [configuredSecrets, setConfiguredSecrets] = useState>(new Set()); const loadCacheSettings = useCallback(async () => { if (!accessToken) { @@ -34,6 +35,7 @@ const CacheSettings: React.FC = ({ accessToken }) => { const data = (await getCacheSettingsCall(accessToken)) as { current_values?: Record }; const currentValues = data.current_values ?? {}; form.setFieldsValue(buildInitialValues(currentValues)); + setConfiguredSecrets(configuredSecretFields(currentValues)); setRedisType(toRedisType(currentValues.redis_type)); } catch (error) { console.error("Failed to load cache settings:", error); @@ -144,6 +146,7 @@ const CacheSettings: React.FC = ({ accessToken }) => { section="connection" redisType={redisType} embeddingModels={embeddingModels} + configuredSecrets={configuredSecrets} />
@@ -166,6 +169,7 @@ const CacheSettings: React.FC = ({ accessToken }) => { section="sentinel" redisType={redisType} embeddingModels={embeddingModels} + configuredSecrets={configuredSecrets} />
)} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx index 21ee41936c1..1ededd9e4b1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.test.tsx @@ -6,25 +6,6 @@ import { renderWithProviders } from "../../../../../tests/test-utils"; import AddMarginForm from "./add_margin_form"; import { MarginConfig } from "./types"; -vi.mock("@/components/provider_info_helpers", () => ({ - Providers: { - OpenAI: "OpenAI", - Anthropic: "Anthropic", - }, - provider_map: { - OpenAI: "openai", - Anthropic: "anthropic", - }, - providerLogoMap: { - OpenAI: "https://example.com/openai.png", - Anthropic: "https://example.com/anthropic.png", - }, -})); - -vi.mock("./provider_display_helpers", () => ({ - handleImageError: vi.fn(), -})); - const DEFAULT_PROPS = { marginConfig: {} as MarginConfig, selectedProvider: undefined, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx index f2c06387301..a17b7fc4ac3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_margin_form.tsx @@ -2,10 +2,9 @@ import React from "react"; import { TextInput, Button } from "@tremor/react"; import { Select as AntdSelect, Form, Tooltip, Radio } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; -import { Providers, provider_map, providerLogoMap } from "@/components/provider_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Providers, provider_map } from "@/components/provider_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; import { MarginConfig } from "./types"; -import { handleImageError } from "./provider_display_helpers"; interface AddMarginFormProps { marginConfig: MarginConfig; @@ -73,12 +72,7 @@ const AddMarginForm: React.FC = ({ return (
- {`${providerEnum} handleImageError(e, providerDisplayName)} - /> + {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx index 48d23d4645d..08fb63c32b9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.test.tsx @@ -5,25 +5,7 @@ import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "../../../../../tests/test-utils"; import AddProviderForm from "./add_provider_form"; import { DiscountConfig } from "./types"; - -vi.mock("@/components/provider_info_helpers", () => ({ - Providers: { - OpenAI: "OpenAI", - Anthropic: "Anthropic", - }, - provider_map: { - OpenAI: "openai", - Anthropic: "anthropic", - }, - providerLogoMap: { - OpenAI: "https://example.com/openai.png", - Anthropic: "https://example.com/anthropic.png", - }, -})); - -vi.mock("./provider_display_helpers", () => ({ - handleImageError: vi.fn(), -})); +import { Providers, providerLogoMap } from "@/components/provider_info_helpers"; const DEFAULT_PROPS = { discountConfig: {} as DiscountConfig, @@ -84,4 +66,18 @@ describe("AddProviderForm", () => { renderWithProviders(); expect(screen.getByText("%")).toBeInTheDocument(); }); + + it("renders the selected provider's bundled logo via the shared Logo component", async () => { + renderWithProviders(); + + const logo = await screen.findByRole("img", { name: `${Providers.OpenAI} logo` }); + expect(logo.getAttribute("src")).toBe(providerLogoMap[Providers.OpenAI]); + }); + + it("falls back to a letter avatar for a selected provider that has no bundled logo", () => { + renderWithProviders(); + + expect(screen.queryByRole("img", { name: `${Providers.PG_VECTOR} logo` })).not.toBeInTheDocument(); + expect(screen.getByText(Providers.PG_VECTOR.charAt(0))).toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx index c4961263533..0fdaed8814b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/add_provider_form.tsx @@ -2,10 +2,9 @@ import React from "react"; import { TextInput, Button } from "@tremor/react"; import { Select as AntdSelect, Form, Tooltip } from "antd"; import { InfoCircleOutlined } from "@ant-design/icons"; -import { Providers, provider_map, providerLogoMap } from "@/components/provider_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Providers, provider_map } from "@/components/provider_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; import { DiscountConfig } from "./types"; -import { handleImageError } from "./provider_display_helpers"; interface AddProviderFormProps { discountConfig: DiscountConfig; @@ -60,12 +59,7 @@ const AddProviderForm: React.FC = ({ return (
- {`${providerEnum} handleImageError(e, providerDisplayName)} - /> + {providerDisplayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx index 0e1c7da92ba..0dae83ba808 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/cost_tracking_settings.test.tsx @@ -49,11 +49,7 @@ vi.mock("@/components/provider_info_helpers", () => ({ Providers: { OpenAI: "OpenAI" }, provider_map: { OpenAI: "openai" }, providerLogoMap: {}, -})); - -vi.mock("./provider_display_helpers", () => ({ - getProviderDisplayInfo: vi.fn(() => ({ displayName: "OpenAI", logo: "", enumKey: "OpenAI" })), - handleImageError: vi.fn(), + getProviderLogoAndName: (providerValue: string) => ({ logo: "", displayName: providerValue }), })); const ADMIN_PROPS = { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts index 8de7fdd7271..90701dd8f1f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/index.ts @@ -11,7 +11,6 @@ export type { MarginConfig, CostMarginResponse, } from "./types"; -export type { ProviderDisplayInfo } from "./provider_display_helpers"; export * from "./provider_display_helpers"; export { useDiscountConfig } from "./use_discount_config"; export { useMarginConfig } from "./use_margin_config"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx index c1c43ebdb4f..2e8dbb429f0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.test.tsx @@ -43,15 +43,6 @@ vi.mock("@tremor/react", () => ({ }, })); -vi.mock("./provider_display_helpers", () => ({ - getProviderDisplayInfo: vi.fn((providerValue: string) => ({ - displayName: providerValue === "openai" ? "OpenAI" : providerValue, - logo: providerValue === "openai" ? "https://example.com/openai.png" : "", - enumKey: providerValue === "openai" ? "OpenAI" : null, - })), - handleImageError: vi.fn(), -})); - const DEFAULT_DISCOUNT_CONFIG = { openai: 0.05, anthropic: 0.1, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx index d802f6d83dd..8727d6cb33c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_discount_table.tsx @@ -3,7 +3,8 @@ import { TextInput, Icon, Text } from "@tremor/react"; import { TrashIcon, PencilAltIcon, CheckIcon, XIcon } from "@heroicons/react/outline"; import { SimpleTable } from "@/components/common_components/simple_table"; import { DiscountConfig } from "./types"; -import { getProviderDisplayInfo, handleImageError } from "./provider_display_helpers"; +import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; interface ProviderDiscountTableProps { discountConfig: DiscountConfig; @@ -55,8 +56,8 @@ const ProviderDiscountTable: React.FC = ({ const data: ProviderDiscountRow[] = Object.entries(discountConfig) .map(([provider, discount]) => ({ provider, discount })) .sort((a, b) => { - const displayA = getProviderDisplayInfo(a.provider).displayName; - const displayB = getProviderDisplayInfo(b.provider).displayName; + const displayA = getProviderLogoAndName(a.provider).displayName; + const displayB = getProviderLogoAndName(b.provider).displayName; return displayA.localeCompare(displayB); }); @@ -67,17 +68,10 @@ const ProviderDiscountTable: React.FC = ({ { header: "Provider", cell: (row) => { - const { displayName, logo } = getProviderDisplayInfo(row.provider); + const { displayName } = getProviderLogoAndName(row.provider); return (
- {logo && ( - {`${displayName} handleImageError(e, displayName)} - /> - )} + {displayName}
); @@ -129,7 +123,7 @@ const ProviderDiscountTable: React.FC = ({ { header: "Actions", cell: (row) => { - const { displayName } = getProviderDisplayInfo(row.provider); + const { displayName } = getProviderLogoAndName(row.provider); return ( ({ - Providers: { - OpenAI: "OpenAI", - Anthropic: "Anthropic", - Azure: "Azure", - }, provider_map: { OpenAI: "openai", Anthropic: "anthropic", Azure: "azure", }, - providerLogoMap: { - OpenAI: "https://example.com/openai.png", - Anthropic: "https://example.com/anthropic.png", - Azure: "https://example.com/azure.png", - }, })); -describe("getProviderDisplayInfo", () => { - it("should return display name and logo for a known backend provider value", () => { - const info = getProviderDisplayInfo("openai"); - expect(info.displayName).toBe("OpenAI"); - expect(info.logo).toBe("https://example.com/openai.png"); - expect(info.enumKey).toBe("OpenAI"); - }); - - it("should return the raw value as display name for an unknown provider", () => { - const info = getProviderDisplayInfo("my-custom-provider"); - expect(info.displayName).toBe("my-custom-provider"); - expect(info.logo).toBe(""); - expect(info.enumKey).toBeNull(); - }); - - it("should match a provider by its backend value regardless of casing", () => { - const info = getProviderDisplayInfo("anthropic"); - expect(info.displayName).toBe("Anthropic"); - expect(info.enumKey).toBe("Anthropic"); - }); -}); - describe("getProviderBackendValue", () => { it("should return the backend value for a known provider enum key", () => { expect(getProviderBackendValue("OpenAI")).toBe("openai"); @@ -54,38 +22,3 @@ describe("getProviderBackendValue", () => { expect(getProviderBackendValue("UnknownProvider")).toBeNull(); }); }); - -describe("handleImageError", () => { - it("should replace the img element with a fallback div showing the first letter", () => { - const img = document.createElement("img"); - const parent = document.createElement("div"); - parent.appendChild(img); - - const event = { target: img } as any; - handleImageError(event, "OpenAI"); - - expect(parent.querySelector("img")).toBeNull(); - const fallback = parent.firstChild as HTMLElement; - expect(fallback.tagName).toBe("DIV"); - expect(fallback.textContent).toBe("O"); - }); - - it("should use the first character of the fallback text as the label", () => { - const img = document.createElement("img"); - const parent = document.createElement("div"); - parent.appendChild(img); - - const event = { target: img } as any; - handleImageError(event, "Anthropic"); - - const fallback = parent.firstChild as HTMLElement; - expect(fallback.textContent).toBe("A"); - }); - - it("should do nothing if the image has no parent element", () => { - const img = document.createElement("img"); - const event = { target: img } as any; - // Should not throw - expect(() => handleImageError(event, "OpenAI")).not.toThrow(); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts index 5489eb12487..ed98ba3586b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_display_helpers.ts @@ -1,28 +1,4 @@ -import { Providers, provider_map, providerLogoMap } from "@/components/provider_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; - -export interface ProviderDisplayInfo { - displayName: string; - logo: string; - enumKey: string | null; -} - -/** - * Convert backend provider value (e.g., "openai") to display info - */ -export const getProviderDisplayInfo = (providerValue: string): ProviderDisplayInfo => { - const enumKey = Object.keys(provider_map).find( - (key) => provider_map[key as keyof typeof provider_map] === providerValue, - ); - - if (enumKey) { - const displayName = Providers[enumKey as keyof typeof Providers]; - const logo = resolveLogoSrc(providerLogoMap[displayName]) ?? ""; - return { displayName, logo, enumKey }; - } - - return { displayName: providerValue, logo: "", enumKey: null }; -}; +import { provider_map } from "@/components/provider_info_helpers"; /** * Convert provider enum key (e.g., "OpenAI") to backend value (e.g., "openai") @@ -30,17 +6,3 @@ export const getProviderDisplayInfo = (providerValue: string): ProviderDisplayIn export const getProviderBackendValue = (providerEnum: string): string | null => { return provider_map[providerEnum as keyof typeof provider_map] || null; }; - -/** - * Handle image error by replacing with fallback div - */ -export const handleImageError = (e: React.SyntheticEvent, fallbackText: string) => { - const target = e.target as HTMLImageElement; - const parent = target.parentElement; - if (parent) { - const fallbackDiv = document.createElement("div"); - fallbackDiv.className = "w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs"; - fallbackDiv.textContent = fallbackText.charAt(0); - parent.replaceChild(fallbackDiv, target); - } -}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx index e1b17dea23d..170e61141b6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-tracking/_components/provider_margin_table.test.tsx @@ -4,6 +4,7 @@ import { screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "../../../../../tests/test-utils"; import ProviderMarginTable from "./provider_margin_table"; +import { Providers, providerLogoMap } from "@/components/provider_info_helpers"; vi.mock("@heroicons/react/outline", () => ({ TrashIcon: function TrashIcon() { @@ -43,15 +44,6 @@ vi.mock("@tremor/react", () => ({ }, })); -vi.mock("./provider_display_helpers", () => ({ - getProviderDisplayInfo: vi.fn((providerValue: string) => { - if (providerValue === "openai") return { displayName: "OpenAI", logo: "", enumKey: "OpenAI" }; - if (providerValue === "anthropic") return { displayName: "Anthropic", logo: "", enumKey: "Anthropic" }; - return { displayName: providerValue, logo: "", enumKey: null }; - }), - handleImageError: vi.fn(), -})); - describe("ProviderMarginTable", () => { const onMarginChange = vi.fn(); const onRemoveProvider = vi.fn(); @@ -95,6 +87,30 @@ describe("ProviderMarginTable", () => { expect(screen.getByText("OpenAI")).toBeInTheDocument(); }); + it("should render the provider's bundled logo via the shared Logo component", () => { + renderWithProviders( + , + ); + const logo = screen.getByRole("img", { name: `${Providers.OpenAI} logo` }); + expect(logo.getAttribute("src")).toBe(providerLogoMap[Providers.OpenAI]); + }); + + it("should fall back to a letter avatar for a provider with no bundled logo", () => { + renderWithProviders( + , + ); + expect(screen.queryByRole("img")).not.toBeInTheDocument(); + expect(screen.getByText("m")).toBeInTheDocument(); + }); + it("should display the global provider as 'Global (All Providers)'", () => { renderWithProviders( = ({ .sort((a, b) => { if (a.provider === "global") return -1; if (b.provider === "global") return 1; - const displayA = getProviderDisplayInfo(a.provider).displayName; - const displayB = getProviderDisplayInfo(b.provider).displayName; + const displayA = getProviderLogoAndName(a.provider).displayName; + const displayB = getProviderLogoAndName(b.provider).displayName; return displayA.localeCompare(displayB); }); @@ -115,17 +116,10 @@ const ProviderMarginTable: React.FC = ({
); } - const { displayName, logo } = getProviderDisplayInfo(row.provider); + const { displayName } = getProviderLogoAndName(row.provider); return (
- {logo && ( - {`${displayName} handleImageError(e, displayName)} - /> - )} + {displayName}
); @@ -186,7 +180,7 @@ const ProviderMarginTable: React.FC = ({ { header: "Actions", cell: (row) => { - const displayName = row.provider === "global" ? "Global" : getProviderDisplayInfo(row.provider).displayName; + const displayName = row.provider === "global" ? "Global" : getProviderLogoAndName(row.provider).displayName; return ( ({ getGuardrailsList: vi.fn(), @@ -48,7 +48,8 @@ vi.mock("@/utils/roles", () => ({ isAdminRole: vi.fn((role: string) => role === "admin"), })); -vi.mock("./guardrail_info_helpers", () => ({ +vi.mock("./guardrail_info_helpers", async (importOriginal) => ({ + ...(await importOriginal()), getGuardrailLogoAndName: vi.fn(() => ({ logo: null, displayName: "Test Provider", @@ -78,6 +79,7 @@ describe("GuardrailsPanel", () => { }; const mockGetGuardrailsList = vi.mocked(getGuardrailsList); + const mockDeleteGuardrailCall = vi.mocked(deleteGuardrailCall); beforeEach(() => { vi.clearAllMocks(); @@ -107,4 +109,35 @@ describe("GuardrailsPanel", () => { fireEvent.click(screen.getByText("Guardrails")); expect(screen.getByText("Add New Guardrail")).toBeInTheDocument(); }); + + it("should delete the clicked guardrail after confirming in the modal", async () => { + render(); + fireEvent.click(screen.getByText("Guardrails")); + + fireEvent.click(await screen.findByTestId("delete-button")); + + const modal = within(await screen.findByRole("dialog")); + expect(modal.getByText("Delete Guardrail")).toBeInTheDocument(); + expect(modal.getByText("test-guardrail-1")).toBeInTheDocument(); + expect(modal.getByText("Test Provider")).toBeInTheDocument(); + + fireEvent.click(modal.getByRole("button", { name: "Delete" })); + + await waitFor(() => { + expect(mockDeleteGuardrailCall).toHaveBeenCalledWith("test-token", "test-guardrail-1"); + }); + expect(mockGetGuardrailsList).toHaveBeenCalledTimes(2); + }); + + it("should not delete anything when the modal is cancelled", async () => { + render(); + fireEvent.click(screen.getByText("Guardrails")); + + fireEvent.click(await screen.findByTestId("delete-button")); + const modal = within(await screen.findByRole("dialog")); + + fireEvent.click(modal.getByRole("button", { name: "Cancel" })); + + expect(mockDeleteGuardrailCall).not.toHaveBeenCalled(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.test.tsx index 8fc0d36c2b4..91d95155e82 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.test.tsx @@ -41,3 +41,17 @@ describe("AddGuardrailForm close behavior", () => { expect(onClose).toHaveBeenCalledTimes(1); }); }); + +describe("AddGuardrailForm provider options", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("renders provider options with logos from the bundled guardrail logo map", async () => { + renderForm(); + fireEvent.mouseDown(screen.getByLabelText("Guardrail Provider")); + + const logo = await screen.findByAltText("Presidio PII logo"); + expect(logo.getAttribute("src")).toContain("microsoft_azure.svg"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx index 202568b478d..17331014c57 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/add_guardrail_form.tsx @@ -12,10 +12,10 @@ import { type CompetitorIntentConfig } from "./content_filter/CompetitorIntentCo import { choiceToSkipSystemForCreate, choiceToSkipToolForCreate, + getGuardrailLogo, getGuardrailProviders, getSupportedModesForProvider, guardrail_provider_map, - guardrailLogoMap, populateGuardrailProviderMap, populateGuardrailProviders, shouldRenderContentFilterConfigSettings, @@ -23,7 +23,7 @@ import { shouldRenderPIIConfigSettings, toModeArray, } from "./guardrail_info_helpers"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import GuardrailOptionalParams from "./guardrail_optional_params"; import GuardrailProviderFields from "./guardrail_provider_fields"; import LLMJudgeFields from "./llm_judge/LLMJudgeFields"; @@ -725,53 +725,19 @@ const AddGuardrailForm: React.FC = ({ visible, onClose, a dropdownRender={(menu) => menu} showSearch={true} > - {Object.entries(getGuardrailProviders()).map(([key, value]) => ( -
- } - > + {Object.entries(getGuardrailProviders()).map(([key, value]) => { + const optionContent = (
- {guardrailLogoMap[value] && ( - { - // Hide broken image icon if image fails to load - e.currentTarget.style.display = "none"; - }} - /> - )} + {value}
- - ))} + ); + return ( + + ); + })} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx index 9ceb6ba244b..ec3d05a6907 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrailTableColumns.tsx @@ -16,6 +16,7 @@ import { import { cn } from "@/lib/cva.config"; import { getGuardrailLogoAndName } from "./guardrail_info_helpers"; +import { Logo } from "@/components/molecules/logo/Logo"; const CONFIG_DELETE_HINT = "Config guardrails are defined in the config file and cannot be deleted from the dashboard."; @@ -23,16 +24,7 @@ function GuardrailProviderCell({ provider }: { provider: string }) { const { logo, displayName } = getGuardrailLogoAndName(provider); return (
- {logo ? ( - { - (event.currentTarget as HTMLImageElement).style.display = "none"; - }} - /> - ) : null} + {displayName}
); 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 0fa5d2ffcd2..2d1f35e456c 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 @@ -53,15 +53,28 @@ describe("GuardrailCard", () => { expect(screen.queryByText(/F1:/)).not.toBeInTheDocument(); }); + it("should render the logo through the shared Logo component with the card src", () => { + render(); + const img = screen.getByAltText("Test Guardrail logo"); + expect(img.getAttribute("src")).toContain("/logos/test.svg"); + }); + + it("should pass a bundled static-import src through unchanged", () => { + const bundledCard: GuardrailCardInfo = { ...baseCard, logo: "/_next/static/media/akto.svg" }; + render(); + expect(screen.getByAltText("Test Guardrail logo")).toHaveAttribute("src", "/_next/static/media/akto.svg"); + }); + it("should show fallback initial when logo fails to load", () => { render(); - const img = screen.getByRole("presentation"); + const img = screen.getByAltText("Test Guardrail logo"); act(() => { fireEvent.error(img); }); expect(screen.getByText("T")).toBeInTheDocument(); + expect(screen.queryByAltText("Test Guardrail logo")).not.toBeInTheDocument(); }); it("should show fallback initial when logo src is empty", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx index 8e9fcc21dfe..53abf3eb81c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_card.tsx @@ -1,42 +1,7 @@ import React, { useState } from "react"; import { CheckCircleFilled } from "@ant-design/icons"; import { GuardrailCardInfo } from "./guardrail_garden_data"; -import { resolveLogoSrc } from "@/lib/assetPaths"; - -const LogoWithFallback: React.FC<{ src: string; name: string }> = ({ src, name }) => { - const [hasError, setHasError] = useState(false); - - if (hasError || !src) { - return ( -
- {name?.charAt(0) || "?"} -
- ); - } - - return ( - setHasError(true)} - /> - ); -}; +import { Logo } from "@/components/molecules/logo/Logo"; const GuardrailCard: React.FC<{ card: GuardrailCardInfo; onClick: () => void }> = ({ card, onClick }) => { const [hovered, setHovered] = useState(false); @@ -61,7 +26,7 @@ const GuardrailCard: React.FC<{ card: GuardrailCardInfo; onClick: () => void }> > {/* Icon + Name row */}
- + {card.name}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts new file mode 100644 index 00000000000..13909e48185 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.test.ts @@ -0,0 +1,54 @@ +import { describe, expect, it } from "vitest"; +import { ALL_CARDS, LITELLM_CONTENT_FILTER_CARDS, PARTNER_GUARDRAIL_CARDS } from "./guardrail_garden_data"; + +const EXPECTED_PARTNER_LOGO_FILES: Record = { + presidio: "microsoft_azure.svg", + bedrock: "bedrock.svg", + lakera: "lakeraai.jpeg", + openai_moderation: "openai_small.svg", + google_model_armor: "google.svg", + guardrails_ai: "guardrails_ai.jpeg", + zscaler: "zscaler.svg", + panw: "palo_alto_networks.jpeg", + cisco_ai_defense: "cisco.png", + noma: "noma_security.png", + aporia: "aporia.png", + aim: "aim_security.jpeg", + cato_networks: "cato_networks.svg", + prompt_security: "prompt_security.png", + lasso: "lasso.png", + pangea: "pangea.png", + enkryptai: "enkrypt_ai.avif", + javelin: "javelin.png", + pillar: "pillar.jpeg", + akto: "akto.svg", + promptguard: "promptguard.svg", + xecguard: "xecguard.svg", + deepkeep: "deepkeep.svg", + repelloai: "repelloai.png", + straiker: "straiker.svg", +}; + +describe("guardrail_garden_data logos", () => { + it("points every partner card at its own provider's bundled logo file", () => { + expect(new Set(PARTNER_GUARDRAIL_CARDS.map((card) => card.id))).toEqual( + new Set(Object.keys(EXPECTED_PARTNER_LOGO_FILES)), + ); + for (const card of PARTNER_GUARDRAIL_CARDS) { + expect(card.logo, `card ${card.id}`).toContain(EXPECTED_PARTNER_LOGO_FILES[card.id]); + } + }); + + it("uses the LiteLLM logo for every content filter card", () => { + for (const card of LITELLM_CONTENT_FILTER_CARDS) { + expect(card.logo, `card ${card.id}`).toContain("litellm_logo.jpg"); + } + }); + + it("bundles every card logo instead of referencing runtime /ui asset paths", () => { + for (const card of ALL_CARDS) { + expect(card.logo, `card ${card.id}`).not.toBe(""); + expect(card.logo, `card ${card.id}`).not.toContain("/ui/assets/logos/"); + } + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts index a29d12f53f9..744af89a357 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts @@ -1,3 +1,5 @@ +import { guardrailLogoMap } from "./guardrail_info_helpers"; + export interface GuardrailCardInfo { id: string; name: string; @@ -16,7 +18,7 @@ export interface GuardrailCardInfo { providerKey?: string; } -const ASSET_PREFIX = "/ui/assets/logos/"; +const litellmContentFilterLogo = guardrailLogoMap["LiteLLM Content Filter"]; export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ { @@ -26,7 +28,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Detects requests for personalized financial advice, investment recommendations, or financial planning.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], eval: { f1: 100.0, @@ -42,7 +44,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects insults, name-calling, and personal attacks directed at the chatbot, staff, or other people.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], eval: { f1: 100.0, @@ -58,7 +60,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects requests for unauthorized legal advice, case analysis, or legal recommendations.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], }, { @@ -67,7 +69,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects requests for medical diagnosis, treatment recommendations, or health advice.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Topic Blocker"], }, { @@ -76,7 +78,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content related to violence, criminal planning, attacks, and violent threats.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -85,7 +87,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content related to self-harm, suicide, and dangerous self-destructive behavior.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -94,7 +96,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content that could endanger child safety or exploit minors.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -103,7 +105,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects content related to illegal weapons manufacturing, distribution, or acquisition.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Safety"], }, { @@ -112,7 +114,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects gender-based discrimination, stereotypes, and biased language.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -121,7 +123,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects racial discrimination, stereotypes, and racially biased content.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -130,7 +132,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects religious discrimination, intolerance, and religiously biased content.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -139,7 +141,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects discrimination based on sexual orientation and related biased content.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Bias"], }, { @@ -148,7 +150,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects jailbreak attempts designed to bypass AI safety guidelines and restrictions.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -157,7 +159,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects attempts to extract sensitive data through prompt manipulation.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -166,7 +168,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects SQL injection attempts embedded in prompts.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -175,7 +177,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects attempts to inject malicious code through prompts.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -184,7 +186,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects attempts to extract or override system prompts.", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Prompt Injection"], }, { @@ -193,7 +195,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ description: "Detects toxic, abusive, and hateful language across multiple languages (EN, AU, DE, ES, FR).", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Toxicity"], }, { @@ -203,7 +205,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Detect and block sensitive data patterns like SSNs, credit card numbers, API keys, and custom regex patterns.", category: "litellm", subcategory: "Patterns", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["PII", "Regex", "Data Protection"], }, { @@ -213,7 +215,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Block or mask content containing specific keywords or phrases. Upload custom word lists or add individual terms.", category: "litellm", subcategory: "Keywords", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Keywords", "Blocklist"], }, { @@ -223,7 +225,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Detects markdown fenced code blocks in requests and responses. Block or mask executable code (e.g. Python, JavaScript, Bash) by language with configurable confidence.", category: "litellm", subcategory: "Code Safety", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Code", "Safety", "Prompt Injection"], }, { @@ -233,7 +235,7 @@ export const LITELLM_CONTENT_FILTER_CARDS: GuardrailCardInfo[] = [ "Block or reframe competitor comparison and ranking intent. Detect when users ask to compare or recommend competitors (airline or generic competitor lists).", category: "litellm", subcategory: "Content Category", - logo: `${ASSET_PREFIX}litellm_logo.jpg`, + logo: litellmContentFilterLogo, tags: ["Content Category", "Competitor", "Topic Blocker"], }, ]; @@ -245,7 +247,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "Microsoft Presidio for PII detection and anonymization. Supports 30+ entity types with configurable actions.", category: "partner", - logo: `${ASSET_PREFIX}microsoft_azure.svg`, + logo: guardrailLogoMap["Presidio PII"], tags: ["PII", "Microsoft"], providerKey: "PresidioPII", }, @@ -254,7 +256,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Bedrock Guardrail", description: "AWS Bedrock Guardrails for content filtering, topic avoidance, and sensitive information detection.", category: "partner", - logo: `${ASSET_PREFIX}bedrock.svg`, + logo: guardrailLogoMap["Bedrock Guardrail"], tags: ["AWS", "Content Safety"], providerKey: "Bedrock", }, @@ -263,7 +265,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Lakera", description: "AI security platform protecting against prompt injections, data leakage, and harmful content.", category: "partner", - logo: `${ASSET_PREFIX}lakeraai.jpeg`, + logo: guardrailLogoMap["Lakera"], tags: ["Security", "Prompt Injection"], providerKey: "Lakera", }, @@ -272,7 +274,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "OpenAI Moderation", description: "OpenAI's content moderation API for detecting harmful content across multiple categories.", category: "partner", - logo: `${ASSET_PREFIX}openai_small.svg`, + logo: guardrailLogoMap["OpenAI Moderation"], tags: ["Content Moderation", "OpenAI"], }, { @@ -280,7 +282,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Google Cloud Model Armor", description: "Google Cloud's model protection service for safe and responsible AI deployments.", category: "partner", - logo: `${ASSET_PREFIX}google.svg`, + logo: guardrailLogoMap["Google Cloud Model Armor"], tags: ["Google Cloud", "Safety"], }, { @@ -288,7 +290,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Guardrails AI", description: "Open-source framework for adding structural, type, and quality guarantees to LLM outputs.", category: "partner", - logo: `${ASSET_PREFIX}guardrails_ai.jpeg`, + logo: guardrailLogoMap["Guardrails AI"], tags: ["Open Source", "Validation"], }, { @@ -296,7 +298,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Zscaler AI Guard", description: "Enterprise AI security from Zscaler for monitoring and protecting AI/ML workloads.", category: "partner", - logo: `${ASSET_PREFIX}zscaler.svg`, + logo: guardrailLogoMap["Zscaler AI Guard"], tags: ["Enterprise", "Security"], }, { @@ -304,7 +306,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "PANW Prisma AIRS", description: "Palo Alto Networks Prisma AI Runtime Security for securing AI applications in production.", category: "partner", - logo: `${ASSET_PREFIX}palo_alto_networks.jpeg`, + logo: guardrailLogoMap["PANW Prisma AIRS"], tags: ["Enterprise", "Security"], }, { @@ -313,7 +315,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "Cisco AI Defense Inspection API for runtime protection: prompt injection, PII/PCI/PHI, harassment, hate speech, profanity, violence, and code detection.", category: "partner", - logo: `${ASSET_PREFIX}cisco.png`, + logo: guardrailLogoMap["Cisco AI Defense"], tags: ["Enterprise", "Security", "Prompt Injection", "PII"], providerKey: "CiscoAiDefense", }, @@ -322,7 +324,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Noma Security", description: "AI security platform for detecting and preventing AI-specific threats and vulnerabilities.", category: "partner", - logo: `${ASSET_PREFIX}noma_security.png`, + logo: guardrailLogoMap["Noma Security"], tags: ["Security", "Threat Detection"], }, { @@ -330,7 +332,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Aporia AI", description: "Real-time AI guardrails for hallucination detection, topic control, and policy enforcement.", category: "partner", - logo: `${ASSET_PREFIX}aporia.png`, + logo: guardrailLogoMap["Aporia AI"], tags: ["Hallucination", "Policy"], }, { @@ -338,7 +340,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "AIM Guardrail", description: "AIM Security guardrails for comprehensive AI threat detection and mitigation.", category: "partner", - logo: `${ASSET_PREFIX}aim_security.jpeg`, + logo: guardrailLogoMap["AIM Guardrail"], tags: ["Security", "Threat Detection"], }, { @@ -346,7 +348,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Cato Networks Guardrail", description: "Cato Networks guardrails for comprehensive AI threat detection and mitigation.", category: "partner", - logo: `${ASSET_PREFIX}cato_networks.svg`, + logo: guardrailLogoMap["Cato Networks Guardrail"], tags: ["Security", "Threat Detection"], }, { @@ -354,7 +356,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Prompt Security", description: "Protect against prompt injection attacks, data leakage, and other LLM security threats.", category: "partner", - logo: `${ASSET_PREFIX}prompt_security.png`, + logo: guardrailLogoMap["Prompt Security"], tags: ["Prompt Injection", "Security"], }, { @@ -362,7 +364,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Lasso Guardrail", description: "Content moderation and safety guardrails for responsible AI deployments.", category: "partner", - logo: `${ASSET_PREFIX}lasso.png`, + logo: guardrailLogoMap["Lasso Guardrail"], tags: ["Content Moderation"], }, { @@ -370,7 +372,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Pangea Guardrail", description: "Pangea's AI guardrails for secure, compliant, and trustworthy AI applications.", category: "partner", - logo: `${ASSET_PREFIX}pangea.png`, + logo: guardrailLogoMap["Pangea Guardrail"], tags: ["Compliance", "Security"], }, { @@ -378,7 +380,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "EnkryptAI", description: "AI security and governance platform for enterprise AI safety and compliance.", category: "partner", - logo: `${ASSET_PREFIX}enkrypt_ai.avif`, + logo: guardrailLogoMap["EnkryptAI"], tags: ["Enterprise", "Governance"], }, { @@ -386,7 +388,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Javelin Guardrails", description: "AI gateway with built-in guardrails for secure and compliant AI operations.", category: "partner", - logo: `${ASSET_PREFIX}javelin.png`, + logo: guardrailLogoMap["Javelin Guardrails"], tags: ["Gateway", "Security"], }, { @@ -394,7 +396,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Pillar Guardrail", description: "AI safety platform for monitoring, testing, and securing AI systems.", category: "partner", - logo: `${ASSET_PREFIX}pillar.jpeg`, + logo: guardrailLogoMap["Pillar Guardrail"], tags: ["Monitoring", "Safety"], }, { @@ -402,7 +404,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ name: "Akto Guardrail", description: "AI security platform from Akto.io with automatic monitoring and guardrails for AI/ML applications.", category: "partner", - logo: `${ASSET_PREFIX}akto.svg`, + logo: guardrailLogoMap["Akto"], tags: ["Security", "Safety", "Monitoring"], }, { @@ -411,7 +413,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "AI security gateway with prompt injection detection, PII redaction, topic filtering, entity blocklists, and hallucination detection. Self-hostable with drop-in proxy integration.", category: "partner", - logo: `${ASSET_PREFIX}promptguard.svg`, + logo: guardrailLogoMap["PromptGuard"], tags: ["Security", "Prompt Injection", "PII"], providerKey: "Promptguard", eval: { @@ -428,7 +430,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "CyCraft XecGuard AI security gateway. Multi-policy scanning (prompt injection, harmful content, PII, system-prompt enforcement) plus RAG context grounding.", category: "partner", - logo: `${ASSET_PREFIX}xecguard.svg`, + logo: guardrailLogoMap["XecGuard"], tags: ["Security", "Policy", "Grounding", "RAG"], providerKey: "Xecguard", }, @@ -438,7 +440,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "DeepKeep AI Firewall for comprehensive LLM security — prompt injection detection, PII protection, content moderation, and policy enforcement with configurable guardrail pipelines.", category: "partner", - logo: `${ASSET_PREFIX}deepkeep.svg`, + logo: guardrailLogoMap["DeepKeep AI Firewall"], tags: ["Security", "Prompt Injection", "PII", "Firewall"], providerKey: "Deepkeep", }, @@ -448,7 +450,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "RepelloAI Argus scans prompts and responses against policies configured per asset in the Repello dashboard.", category: "partner", - logo: `${ASSET_PREFIX}repelloai.png`, + logo: guardrailLogoMap["RepelloAI Argus"], tags: ["Security", "Policy", "Prompt Injection"], providerKey: "Repelloai", }, @@ -458,7 +460,7 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ description: "Defend AI Agentic Guardrails: Indirect/Direct Prompt Injection, Tool Misuse, Malicious MCP and Skills", category: "partner", - logo: `${ASSET_PREFIX}straiker.svg`, + logo: guardrailLogoMap["Straiker"], tags: ["Agentic", "Prompt Injection", "Tool Misuse", "MCP", "Skills"], providerKey: "Straiker", }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.test.tsx new file mode 100644 index 00000000000..e17e739267d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.test.tsx @@ -0,0 +1,32 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; +import GuardrailDetailView from "./guardrail_garden_detail"; +import type { GuardrailCardInfo } from "./guardrail_garden_data"; + +vi.mock("./add_guardrail_form", () => ({ default: () => null })); + +const makeCard = (overrides: Partial = {}): GuardrailCardInfo => ({ + id: "bedrock", + name: "Bedrock Guardrail", + description: "AWS Bedrock Guardrails for content filtering.", + category: "partner", + logo: "/_next/static/media/bedrock.svg", + tags: ["AWS"], + ...overrides, +}); + +const renderDetail = (card: GuardrailCardInfo) => + render(); + +describe("GuardrailDetailView logo", () => { + it("renders the card logo through the shared Logo component with the bundled src", () => { + renderDetail(makeCard()); + expect(screen.getByAltText("Bedrock Guardrail logo")).toHaveAttribute("src", "/_next/static/media/bedrock.svg"); + }); + + it("falls back to a letter avatar when the card has no logo", () => { + renderDetail(makeCard({ logo: "" })); + expect(screen.queryByAltText("Bedrock Guardrail logo")).not.toBeInTheDocument(); + expect(screen.getByText("B")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx index c92486bbad9..71c7a527614 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_detail.tsx @@ -2,7 +2,7 @@ import React, { useState } from "react"; import { Button } from "antd"; import { ArrowLeftOutlined } from "@ant-design/icons"; import AddGuardrailForm from "./add_guardrail_form"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import { GUARDRAIL_PRESETS } from "./guardrail_garden_configs"; import { GuardrailCardInfo } from "./guardrail_garden_data"; @@ -60,14 +60,7 @@ const GuardrailDetailView: React.FC = ({ card, onBack, {/* ── Header block (Vertex-style) ── */}
- { - (e.target as HTMLImageElement).style.display = "none"; - }} - /> +

{card.name}

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 c89fe7277c9..7bb7737e152 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 @@ -81,6 +81,37 @@ describe("Guardrail Info", () => { expect(getByText("Settings")).toBeInTheDocument(); }); + it("should render the provider logo from the bundled guardrail logo map", async () => { + vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ + guardrail_id: "123", + guardrail_name: "Test Guardrail", + litellm_params: { + guardrail: "presidio", + mode: "pre_call", + default_on: true, + }, + created_at: "2024-01-01T00:00:00Z", + updated_at: "2024-01-01T00:00:00Z", + guardrail_definition_location: "database", + }); + + vi.mocked(networking.getGuardrailUISettings).mockResolvedValue({ + supported_entities: [], + supported_actions: [], + pii_entity_categories: [], + supported_modes: ["pre_call", "post_call"], + }); + + vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({}); + + const { findByAltText } = render( + {}} accessToken="123" isAdmin={true} />, + ); + + const logo = await findByAltText("Presidio PII logo"); + expect(logo.getAttribute("src")).toContain("microsoft_azure.svg"); + }); + it("should not render the edit button for config guardrails", async () => { // Mock the network responses vi.mocked(networking.getGuardrailInfo).mockResolvedValue({ 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 1941ec94a60..07df6ff15d9 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 @@ -12,6 +12,7 @@ import { Button, Divider, Form, Input, Select, Tooltip } from "antd"; import { CheckIcon, CopyIcon } from "lucide-react"; import React, { useCallback, useEffect, useState } from "react"; import NotificationsManager from "@/components/molecules/notifications_manager"; +import { Logo } from "@/components/molecules/logo/Logo"; import ContentFilterManager, { formatContentFilterDataForAPI } from "./content_filter/ContentFilterManager"; import CustomCodeModal, { EditGuardrailData } from "./custom_code/CustomCodeModal"; import { @@ -524,17 +525,7 @@ const GuardrailInfoView: React.FC = ({ guardrailId, onClose, Provider
- {logo && ( - {`${displayName} { - // Hide broken image - (e.target as HTMLImageElement).style.display = "none"; - }} - /> - )} + {displayName}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx index dd70bb2cf51..12aaba0d696 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx @@ -1,4 +1,30 @@ -import { resolveLogoSrc } from "@/lib/assetPaths"; +import aimSecurityLogo from "../../../../../public/assets/logos/aim_security.jpeg"; +import aktoLogo from "../../../../../public/assets/logos/akto.svg"; +import aporiaLogo from "../../../../../public/assets/logos/aporia.png"; +import bedrockLogo from "../../../../../public/assets/logos/bedrock.svg"; +import catoNetworksLogo from "../../../../../public/assets/logos/cato_networks.svg"; +import ciscoLogo from "../../../../../public/assets/logos/cisco.png"; +import deepkeepLogo from "../../../../../public/assets/logos/deepkeep.svg"; +import enkryptAiLogo from "../../../../../public/assets/logos/enkrypt_ai.avif"; +import googleLogo from "../../../../../public/assets/logos/google.svg"; +import guardrailsAiLogo from "../../../../../public/assets/logos/guardrails_ai.jpeg"; +import javelinLogo from "../../../../../public/assets/logos/javelin.png"; +import lakeraAiLogo from "../../../../../public/assets/logos/lakeraai.jpeg"; +import lassoLogo from "../../../../../public/assets/logos/lasso.png"; +import litellmLogo from "../../../../../public/assets/logos/litellm_logo.jpg"; +import microsoftAzureLogo from "../../../../../public/assets/logos/microsoft_azure.svg"; +import nomaSecurityLogo from "../../../../../public/assets/logos/noma_security.png"; +import openaiSmallLogo from "../../../../../public/assets/logos/openai_small.svg"; +import paloAltoNetworksLogo from "../../../../../public/assets/logos/palo_alto_networks.jpeg"; +import pangeaLogo from "../../../../../public/assets/logos/pangea.png"; +import pillarLogo from "../../../../../public/assets/logos/pillar.jpeg"; +import promptSecurityLogo from "../../../../../public/assets/logos/prompt_security.png"; +import promptguardLogo from "../../../../../public/assets/logos/promptguard.svg"; +import qohashLogo from "../../../../../public/assets/logos/qohash.jpg"; +import repelloAiLogo from "../../../../../public/assets/logos/repelloai.png"; +import straikerLogo from "../../../../../public/assets/logos/straiker.svg"; +import xecguardLogo from "../../../../../public/assets/logos/xecguard.svg"; +import zscalerLogo from "../../../../../public/assets/logos/zscaler.svg"; // Legacy enum - keeping for backward compatibility export enum GuardrailProviders { @@ -136,40 +162,43 @@ export const shouldRenderLLMJudgeFields = (provider: string | null) => { return guardrail_provider_map[provider] === "llm_as_a_judge"; }; -const asset_logos_folder = "/ui/assets/logos/"; +export const guardrailLogoMap = { + "Zscaler AI Guard": zscalerLogo.src, + "Presidio PII": microsoftAzureLogo.src, + "Bedrock Guardrail": bedrockLogo.src, + Lakera: lakeraAiLogo.src, + "Azure Content Safety Prompt Shield": microsoftAzureLogo.src, + "Azure Content Safety Text Moderation": microsoftAzureLogo.src, + "Aporia AI": aporiaLogo.src, + "PANW Prisma AIRS": paloAltoNetworksLogo.src, + "Cisco AI Defense": ciscoLogo.src, + "Noma Security": nomaSecurityLogo.src, + "Javelin Guardrails": javelinLogo.src, + "Pillar Guardrail": pillarLogo.src, + "Google Cloud Model Armor": googleLogo.src, + "Guardrails AI": guardrailsAiLogo.src, + "Lasso Guardrail": lassoLogo.src, + "Pangea Guardrail": pangeaLogo.src, + "AIM Guardrail": aimSecurityLogo.src, + "Cato Networks Guardrail": catoNetworksLogo.src, + "OpenAI Moderation": openaiSmallLogo.src, + EnkryptAI: enkryptAiLogo.src, + "Prompt Security": promptSecurityLogo.src, + PromptGuard: promptguardLogo.src, + XecGuard: xecguardLogo.src, + "LiteLLM Content Filter": litellmLogo.src, + "LiteLLM LLM as a Judge": litellmLogo.src, + Akto: aktoLogo.src, + "DeepKeep AI Firewall": deepkeepLogo.src, + "Qostodian Nexus": qohashLogo.src, + "RepelloAI Argus": repelloAiLogo.src, + Straiker: straikerLogo.src, +} satisfies Record; -export const guardrailLogoMap: Record = { - "Zscaler AI Guard": `${asset_logos_folder}zscaler.svg`, - "Presidio PII": `${asset_logos_folder}microsoft_azure.svg`, - "Bedrock Guardrail": `${asset_logos_folder}bedrock.svg`, - Lakera: `${asset_logos_folder}lakeraai.jpeg`, - "Azure Content Safety Prompt Shield": `${asset_logos_folder}microsoft_azure.svg`, - "Azure Content Safety Text Moderation": `${asset_logos_folder}microsoft_azure.svg`, - "Aporia AI": `${asset_logos_folder}aporia.png`, - "PANW Prisma AIRS": `${asset_logos_folder}palo_alto_networks.jpeg`, - "Cisco AI Defense": `${asset_logos_folder}cisco.png`, - "Noma Security": `${asset_logos_folder}noma_security.png`, - "Javelin Guardrails": `${asset_logos_folder}javelin.png`, - "Pillar Guardrail": `${asset_logos_folder}pillar.jpeg`, - "Google Cloud Model Armor": `${asset_logos_folder}google.svg`, - "Guardrails AI": `${asset_logos_folder}guardrails_ai.jpeg`, - "Lasso Guardrail": `${asset_logos_folder}lasso.png`, - "Pangea Guardrail": `${asset_logos_folder}pangea.png`, - "AIM Guardrail": `${asset_logos_folder}aim_security.jpeg`, - "Cato Networks Guardrail": `${asset_logos_folder}cato_networks.svg`, - "OpenAI Moderation": `${asset_logos_folder}openai_small.svg`, - EnkryptAI: `${asset_logos_folder}enkrypt_ai.avif`, - "Prompt Security": `${asset_logos_folder}prompt_security.png`, - PromptGuard: `${asset_logos_folder}promptguard.svg`, - XecGuard: `${asset_logos_folder}xecguard.svg`, - "LiteLLM Content Filter": `${asset_logos_folder}litellm_logo.jpg`, - "LiteLLM LLM as a Judge": `${asset_logos_folder}litellm_logo.jpg`, - Akto: `${asset_logos_folder}akto.svg`, - "DeepKeep AI Firewall": `${asset_logos_folder}deepkeep.svg`, - "Qostodian Nexus": `${asset_logos_folder}qohash.jpg`, - "RepelloAI Argus": `${asset_logos_folder}repelloai.png`, - Straiker: `${asset_logos_folder}straiker.svg`, -}; +export const getGuardrailLogo = (displayName: string): string | undefined => + Object.prototype.hasOwnProperty.call(guardrailLogoMap, displayName) + ? guardrailLogoMap[displayName as keyof typeof guardrailLogoMap] + : undefined; export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; displayName: string } => { if (!guardrailValue) { @@ -188,7 +217,7 @@ export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; // Get the display name from current GuardrailProviders and logo from map const currentProviders = getGuardrailProviders(); const displayName = currentProviders[enumKey as keyof typeof currentProviders]; - const logo = resolveLogoSrc(guardrailLogoMap[displayName as keyof typeof guardrailLogoMap]) ?? ""; + const logo = getGuardrailLogo(displayName ?? "") ?? ""; return { logo, displayName: displayName || guardrailValue }; }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx index 4f556e74c16..7612b702391 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_table.test.tsx @@ -30,6 +30,22 @@ describe("GuardrailTable", () => { } }); + it("renders the provider logo from the bundled guardrail logo map", () => { + render(); + const logo = screen.getByAltText("Presidio PII logo"); + expect(logo.getAttribute("src")).toContain("microsoft_azure.svg"); + }); + + it("falls back to a letter avatar for an unknown provider slug", () => { + const guardrail = makeGuardrail({ + litellm_params: { guardrail: "mystery_guard", mode: "pre_call", default_on: false }, + }); + render(); + expect(screen.getByText("mystery_guard")).toBeInTheDocument(); + expect(screen.queryByAltText("mystery_guard logo")).not.toBeInTheDocument(); + expect(screen.getByText("m")).toBeInTheDocument(); + }); + it("deletes a DB guardrail through the actions menu", async () => { const user = userEvent.setup(); const onDeleteClick = vi.fn(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.test.ts new file mode 100644 index 00000000000..5eb3bdc105d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.test.ts @@ -0,0 +1,103 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { renderHook, waitFor } from "@testing-library/react"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import React, { ReactNode } from "react"; +import { useSetKeyBlockedState, setKeyBlockedState } from "./useSetKeyBlockedState"; +import { apiClient } from "@/components/networking"; + +vi.mock("@/components/networking", () => ({ + apiClient: { post: vi.fn() }, +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => mockUseAuthorized(), +})); + +const mockPost = vi.mocked(apiClient.post); + +const createWrapper = () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false }, mutations: { retry: false } } }); + const wrapper = ({ children }: { children: ReactNode }) => + React.createElement(QueryClientProvider, { client: queryClient }, children); + return { queryClient, wrapper }; +}; + +describe("setKeyBlockedState", () => { + beforeEach(() => { + mockPost.mockReset(); + }); + + it("POSTs the key hash to /key/block when blocking", async () => { + mockPost.mockResolvedValueOnce({ blocked: true }); + + const result = await setKeyBlockedState("sk-access", { keyToken: "hashed-token", blocked: true }); + + expect(mockPost).toHaveBeenCalledWith("/key/block", { + accessToken: "sk-access", + body: { key: "hashed-token" }, + }); + expect(result).toEqual({ blocked: true }); + }); + + it("POSTs the key hash to /key/unblock when unblocking", async () => { + mockPost.mockResolvedValueOnce({ blocked: false }); + + const result = await setKeyBlockedState("sk-access", { keyToken: "hashed-token", blocked: false }); + + expect(mockPost).toHaveBeenCalledWith("/key/unblock", { + accessToken: "sk-access", + body: { key: "hashed-token" }, + }); + expect(result).toEqual({ blocked: false }); + }); + + it("falls back to the requested state when the response has no blocked field", async () => { + mockPost.mockResolvedValueOnce(null); + + const result = await setKeyBlockedState("sk-access", { keyToken: "hashed-token", blocked: true }); + + expect(result).toEqual({ blocked: true }); + }); +}); + +describe("useSetKeyBlockedState", () => { + beforeEach(() => { + mockPost.mockReset(); + mockUseAuthorized.mockReturnValue({ accessToken: "sk-access" }); + }); + + it("invalidates key queries after a successful mutation", async () => { + mockPost.mockResolvedValueOnce({ blocked: true }); + const { queryClient, wrapper } = createWrapper(); + const invalidateSpy = vi.spyOn(queryClient, "invalidateQueries"); + + const { result } = renderHook(() => useSetKeyBlockedState(), { wrapper }); + result.current.mutate({ keyToken: "hashed-token", blocked: true }); + + await waitFor(() => expect(result.current.isSuccess).toBe(true)); + expect(invalidateSpy).toHaveBeenCalledWith({ queryKey: ["keys"] }); + }); + + it("surfaces request failures as mutation errors", async () => { + mockPost.mockRejectedValueOnce(new Error("Key not found.")); + const { wrapper } = createWrapper(); + + const { result } = renderHook(() => useSetKeyBlockedState(), { wrapper }); + result.current.mutate({ keyToken: "missing", blocked: true }); + + await waitFor(() => expect(result.current.isError).toBe(true)); + expect(result.current.error?.message).toBe("Key not found."); + }); + + it("errors without an access token", async () => { + mockUseAuthorized.mockReturnValue({ accessToken: null }); + const { wrapper } = createWrapper(); + + const { result } = renderHook(() => useSetKeyBlockedState(), { wrapper }); + result.current.mutate({ keyToken: "hashed-token", blocked: true }); + + await waitFor(() => expect(result.current.isError).toBe(true)); + expect(mockPost).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.ts new file mode 100644 index 00000000000..792ef567f99 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useSetKeyBlockedState.ts @@ -0,0 +1,45 @@ +import { useMutation, useQueryClient } from "@tanstack/react-query"; +import { apiClient } from "@/components/networking"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { keyKeys } from "./useKeys"; + +export interface SetKeyBlockedStateInput { + keyToken: string; + blocked: boolean; +} + +export interface SetKeyBlockedStateResult { + blocked: boolean; +} + +interface BlockKeyResponse { + blocked?: boolean | null; +} + +export const setKeyBlockedState = async ( + accessToken: string, + { keyToken, blocked }: SetKeyBlockedStateInput, +): Promise => { + const response = await apiClient.post(blocked ? "/key/block" : "/key/unblock", { + accessToken, + body: { key: keyToken }, + }); + return { blocked: response?.blocked ?? blocked }; +}; + +export const useSetKeyBlockedState = () => { + const { accessToken } = useAuthorized(); + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: async (input) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return setKeyBlockedState(accessToken, input); + }, + onSuccess: () => { + queryClient.invalidateQueries({ queryKey: keyKeys.all }); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts index 0431a8d39f7..1a02e363de9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/sso/useSSOSettings.ts @@ -24,6 +24,7 @@ export interface SSOSettingsValues { generic_authorization_endpoint: string | null; generic_token_endpoint: string | null; generic_userinfo_endpoint: string | null; + generic_scope: string | null; proxy_base_url: string | null; user_email: string | null; ui_access_mode: string | null; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx index 94b9058b372..67b5d6bfe92 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.test.tsx @@ -52,4 +52,23 @@ describe("MCPLogoSelector", () => { await user.click(githubButton); expect(onChange).toHaveBeenCalledWith(undefined); }); + + it("should render grid logos from bundled static assets instead of public paths", () => { + render(); + const src = screen.getByAltText("GitHub").getAttribute("src"); + expect(src).toMatch(/^\/_next\//); + expect(src).toContain("github.svg"); + }); + + it("should preview a stored well-known path via its bundled asset", () => { + render(); + const src = screen.getByAltText("Selected logo").getAttribute("src"); + expect(src).toMatch(/^\/_next\//); + expect(src).toContain("github.svg"); + }); + + it("should preview a custom external URL untouched", () => { + render(); + expect(screen.getByAltText("Selected logo").getAttribute("src")).toBe("https://cdn.example.com/logo.png"); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx index 6f626a1a70b..a67a0dc882d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPLogoSelector.tsx @@ -1,31 +1,51 @@ -import React, { useState } from "react"; +import React from "react"; import { Input, Tooltip } from "antd"; import { InfoCircleOutlined, LinkOutlined } from "@ant-design/icons"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; +import githubLogo from "../../../../../public/assets/logos/github.svg"; +import slackLogo from "../../../../../public/assets/logos/slack.svg"; +import notionLogo from "../../../../../public/assets/logos/notion.svg"; +import linearLogo from "../../../../../public/assets/logos/linear.svg"; +import jiraLogo from "../../../../../public/assets/logos/jira.svg"; +import figmaLogo from "../../../../../public/assets/logos/figma.svg"; +import gmailLogo from "../../../../../public/assets/logos/gmail.svg"; +import googleDriveLogo from "../../../../../public/assets/logos/google_drive.svg"; +import stripeLogo from "../../../../../public/assets/logos/stripe.svg"; +import shopifyLogo from "../../../../../public/assets/logos/shopify.svg"; +import salesforceLogo from "../../../../../public/assets/logos/salesforce.svg"; +import hubspotLogo from "../../../../../public/assets/logos/hubspot.svg"; +import twilioLogo from "../../../../../public/assets/logos/twilio.svg"; +import cloudflareLogo from "../../../../../public/assets/logos/cloudflare.svg"; +import sentryLogo from "../../../../../public/assets/logos/sentry.svg"; +import postgresqlLogo from "../../../../../public/assets/logos/postgresql.svg"; +import snowflakeLogo from "../../../../../public/assets/logos/snowflake.svg"; +import zapierLogo from "../../../../../public/assets/logos/zapier.svg"; +import googleLogo from "../../../../../public/assets/logos/google.svg"; +import gitlabLogo from "../../../../../public/assets/logos/gitlab.svg"; const logos = "/ui/assets/logos/"; -const WELL_KNOWN_LOGOS: { name: string; url: string }[] = [ - { name: "GitHub", url: `${logos}github.svg` }, - { name: "Slack", url: `${logos}slack.svg` }, - { name: "Notion", url: `${logos}notion.svg` }, - { name: "Linear", url: `${logos}linear.svg` }, - { name: "Jira", url: `${logos}jira.svg` }, - { name: "Figma", url: `${logos}figma.svg` }, - { name: "Gmail", url: `${logos}gmail.svg` }, - { name: "Google Drive", url: `${logos}google_drive.svg` }, - { name: "Stripe", url: `${logos}stripe.svg` }, - { name: "Shopify", url: `${logos}shopify.svg` }, - { name: "Salesforce", url: `${logos}salesforce.svg` }, - { name: "HubSpot", url: `${logos}hubspot.svg` }, - { name: "Twilio", url: `${logos}twilio.svg` }, - { name: "Cloudflare", url: `${logos}cloudflare.svg` }, - { name: "Sentry", url: `${logos}sentry.svg` }, - { name: "PostgreSQL", url: `${logos}postgresql.svg` }, - { name: "Snowflake", url: `${logos}snowflake.svg` }, - { name: "Zapier", url: `${logos}zapier.svg` }, - { name: "Google", url: `${logos}google.svg` }, - { name: "GitLab", url: `${logos}gitlab.svg` }, +const WELL_KNOWN_LOGOS: { name: string; url: string; src: string }[] = [ + { name: "GitHub", url: `${logos}github.svg`, src: githubLogo.src }, + { name: "Slack", url: `${logos}slack.svg`, src: slackLogo.src }, + { name: "Notion", url: `${logos}notion.svg`, src: notionLogo.src }, + { name: "Linear", url: `${logos}linear.svg`, src: linearLogo.src }, + { name: "Jira", url: `${logos}jira.svg`, src: jiraLogo.src }, + { name: "Figma", url: `${logos}figma.svg`, src: figmaLogo.src }, + { name: "Gmail", url: `${logos}gmail.svg`, src: gmailLogo.src }, + { name: "Google Drive", url: `${logos}google_drive.svg`, src: googleDriveLogo.src }, + { name: "Stripe", url: `${logos}stripe.svg`, src: stripeLogo.src }, + { name: "Shopify", url: `${logos}shopify.svg`, src: shopifyLogo.src }, + { name: "Salesforce", url: `${logos}salesforce.svg`, src: salesforceLogo.src }, + { name: "HubSpot", url: `${logos}hubspot.svg`, src: hubspotLogo.src }, + { name: "Twilio", url: `${logos}twilio.svg`, src: twilioLogo.src }, + { name: "Cloudflare", url: `${logos}cloudflare.svg`, src: cloudflareLogo.src }, + { name: "Sentry", url: `${logos}sentry.svg`, src: sentryLogo.src }, + { name: "PostgreSQL", url: `${logos}postgresql.svg`, src: postgresqlLogo.src }, + { name: "Snowflake", url: `${logos}snowflake.svg`, src: snowflakeLogo.src }, + { name: "Zapier", url: `${logos}zapier.svg`, src: zapierLogo.src }, + { name: "Google", url: `${logos}google.svg`, src: googleLogo.src }, + { name: "GitLab", url: `${logos}gitlab.svg`, src: gitlabLogo.src }, ]; interface MCPLogoSelectorProps { @@ -34,16 +54,12 @@ interface MCPLogoSelectorProps { } const MCPLogoSelector: React.FC = ({ value, onChange }) => { - const [imgErrors, setImgErrors] = useState>(new Set()); + const selectedWellKnown = WELL_KNOWN_LOGOS.find((l) => l.url === value); const handleSelect = (url: string) => { onChange?.(value === url ? undefined : url); }; - const handleImgError = (url: string) => { - setImgErrors((prev) => new Set(prev).add(url)); - }; - return (
@@ -56,13 +72,10 @@ const MCPLogoSelector: React.FC = ({ value, onChange }) => {/* Preview */} {value && (
- Selected logo { - (e.target as HTMLImageElement).style.display = "none"; - }} />
{value}
@@ -81,8 +94,6 @@ const MCPLogoSelector: React.FC = ({ value, onChange }) =>
{WELL_KNOWN_LOGOS.map((logo) => { const isSelected = value === logo.url; - const hasFailed = imgErrors.has(logo.url); - if (hasFailed) return null; return ( ); @@ -112,7 +118,7 @@ const MCPLogoSelector: React.FC = ({ value, onChange }) => } placeholder="Or paste a custom logo URL..." - value={value && !WELL_KNOWN_LOGOS.some((l) => l.url === value) ? value : ""} + value={value && !selectedWellKnown ? value : ""} onChange={(e) => { const v = e.target.value.trim(); onChange?.(v || undefined); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx index a0998b587fb..d6343afe219 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.test.tsx @@ -1,8 +1,9 @@ import React from "react"; import { render, screen } from "@testing-library/react"; -import { describe, it, expect, vi } from "vitest"; +import { describe, it, expect, vi, afterEach } from "vitest"; import MCPServerCard from "./MCPServerCard"; import type { MCPServer } from "@/components/mcp_tools/types"; +import { setServerRootPath } from "@/lib/serverRootPath"; const baseServer: MCPServer = { server_id: "srv-1", @@ -43,3 +44,26 @@ describe("MCPServerCard OAuth flow indicator", () => { expect(screen.queryByText("OAuth flow not set")).not.toBeInTheDocument(); }); }); + +describe("MCPServerCard logo", () => { + afterEach(() => { + setServerRootPath("/"); + }); + + it("passes an external logo_url through untouched", () => { + renderCard({ mcp_info: { server_name: "demo_server", logo_url: "https://cdn.example.com/logo.png" } }); + expect(screen.getByAltText("demo_server logo").getAttribute("src")).toBe("https://cdn.example.com/logo.png"); + }); + + it("prefixes a stored asset path with the server root path under a non-root mount", () => { + setServerRootPath("/litellm"); + renderCard({ mcp_info: { server_name: "demo_server", logo_url: "/ui/assets/logos/github.svg" } }); + expect(screen.getByAltText("demo_server logo").getAttribute("src")).toBe("/litellm/ui/assets/logos/github.svg"); + }); + + it("renders a letter avatar when no logo_url is set", () => { + renderCard({ mcp_info: { server_name: "demo_server" } }); + expect(screen.queryByAltText("demo_server logo")).not.toBeInTheDocument(); + expect(screen.getByText("DE")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx index 4282cdba278..c7dd6e47f76 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/MCPServerCard.tsx @@ -1,4 +1,4 @@ -import { useState, type FC, type KeyboardEvent, type MouseEvent } from "react"; +import { type FC, type KeyboardEvent, type MouseEvent } from "react"; import { Dropdown, Tooltip, Typography, Tag } from "antd"; import type { MenuProps } from "antd"; import { @@ -9,6 +9,7 @@ import { ThunderboltOutlined, } from "@ant-design/icons"; import { AUTH_TYPE, type MCPServer } from "@/components/mcp_tools/types"; +import { Logo } from "@/components/molecules/logo/Logo"; import { getMaskedAndFullUrl } from "./utils"; const { Text } = Typography; @@ -52,8 +53,6 @@ const MCPServerCard: FC = ({ const name = server.server_name || alias || server.server_id; // Logo is sourced exclusively from the admin-set `mcp_info.logo_url`. const candidateLogo = server.mcp_info?.logo_url ?? undefined; - const [failedLogoUrl, setFailedLogoUrl] = useState(null); - const logoUrl = candidateLogo && failedLogoUrl !== candidateLogo ? candidateLogo : undefined; const transport = server.transport || "http"; const displayTransport = server.spec_path && transport !== "stdio" ? "openapi" : transport; const authType = server.auth_type || "none"; @@ -148,13 +147,8 @@ const MCPServerCard: FC = ({ className={`group relative flex h-full cursor-pointer flex-col gap-3 rounded-lg p-4 transition-all duration-150 focus:outline-hidden focus-visible:ring-2 focus-visible:ring-blue-400 ${cardClass}`} >
- {logoUrl ? ( - {`${name} setFailedLogoUrl(logoUrl)} - /> + {candidateLogo ? ( + ) : (
{(name || "?").slice(0, 2).toUpperCase()} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx index 9f7639d00c7..b21a5218c20 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/create_mcp_server.tsx @@ -39,10 +39,9 @@ import NotificationsManager from "@/components/molecules/notifications_manager"; import { useMcpOAuthFlow } from "@/hooks/useMcpOAuthFlow"; import { useTestMCPConnection } from "@/hooks/useTestMCPConnection"; import { getSecureItem, setSecureItem } from "@/utils/secureStorage"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import mcpLogo from "../../../../../public/assets/logos/mcp_logo.png"; -const asset_logos_folder = "/ui/assets/logos/"; -export const mcpLogoImg = `${asset_logos_folder}mcp_logo.png`; +export const mcpLogoImg = mcpLogo.src; interface CreateMCPServerProps { userRole: string; @@ -791,7 +790,7 @@ const CreateMCPServer: React.FC = ({ )} MCP Logo = ({ auth_type={mcpServer.auth_type} oauth2_flow={mcpServer.oauth2_flow} delegate_auth_to_upstream={mcpServer.delegate_auth_to_upstream} + dcr_bridge={mcpServer.dcr_bridge} tokenUrl={mcpServer.token_url} userRole={userRole} userID={userID} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx index 8b0e6d62f66..3a189f50264 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.test.tsx @@ -17,8 +17,12 @@ vi.mock("@/utils/mcpTokenStore", () => ({ removeToken: vi.fn(), })); +const { toolsOAuthFlowSpy } = vi.hoisted(() => ({ + toolsOAuthFlowSpy: vi.fn(() => ({ startOAuthFlow: vi.fn(), status: "idle", error: null })), +})); + vi.mock("@/hooks/useToolsOAuthFlow", () => ({ - useToolsOAuthFlow: () => ({ startOAuthFlow: vi.fn(), status: "idle", error: null }), + useToolsOAuthFlow: toolsOAuthFlowSpy, })); vi.mock("@/hooks/useUserMcpOAuthFlow", () => ({ @@ -54,6 +58,27 @@ const credStatus = (overrides: Record = {}) => ({ ...overrides, }); +describe("MCPToolsViewer gatewayMintsClient wiring", () => { + // Pins the call site (not just the helper): the viewer must pass the bridge-AWARE + // gatewayMintsClientFor value to useToolsOAuthFlow, so the browser skips its own register exactly + // when the gateway mints. The oauth_delegate + dcr_bridge cell is the regression guard: with the + // old bridge-blind predicate it would have passed true here and dead-ended. + beforeEach(() => toolsOAuthFlowSpy.mockClear()); + + it.each([ + { auth_type: "true_passthrough", dcr_bridge: true, gatewayMintsClient: true }, + { auth_type: "true_passthrough", dcr_bridge: false, gatewayMintsClient: true }, + { auth_type: "oauth_delegate", dcr_bridge: false, gatewayMintsClient: true }, + { auth_type: "oauth_delegate", dcr_bridge: true, gatewayMintsClient: false }, + ])( + "passes gatewayMintsClient=$gatewayMintsClient for $auth_type dcr_bridge=$dcr_bridge", + ({ auth_type, dcr_bridge, gatewayMintsClient }) => { + renderViewer({ auth_type, dcr_bridge, tokenUrl: null }); + expect(toolsOAuthFlowSpy).toHaveBeenCalledWith(expect.objectContaining({ gatewayMintsClient })); + }, + ); +}); + describe("MCPToolsViewer auth gate routing", () => { beforeEach(() => { vi.mocked(listMCPTools).mockReset().mockResolvedValue({ tools: [], error: null }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx index 428c10da284..28c17d1e41c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_tools.tsx @@ -4,6 +4,7 @@ import { ToolTestPanel } from "./ToolTestPanel"; import { resolveLogoSrc } from "@/lib/assetPaths"; import { isClientForwardedTokenMode, + gatewayMintsClientFor, MCPTool, MCPToolsViewerProps, MCPContent, @@ -28,6 +29,7 @@ const MCPToolsViewer = ({ auth_type, oauth2_flow, delegate_auth_to_upstream, + dcr_bridge, userRole, userID, serverAlias, @@ -76,6 +78,7 @@ const MCPToolsViewer = ({ serverId, serverAlias, userId: userID, + gatewayMintsClient: gatewayMintsClientFor({ auth_type, dcr_bridge }), onSuccess: setOauthToken, }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryDetailDrawer.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryDetailDrawer.tsx new file mode 100644 index 00000000000..970e088ec00 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryDetailDrawer.tsx @@ -0,0 +1,114 @@ +"use client"; + +import { Drawer, Space, Typography } from "antd"; +import React from "react"; + +import { MemoryRow } from "@/components/networking"; + +const { Text, Paragraph } = Typography; + +interface MemoryDetailDrawerProps { + row: MemoryRow | null; + onClose: () => void; +} + +function formatTimestamp(ts?: string): string { + if (!ts) return "—"; + try { + const d = new Date(ts); + return d.toLocaleString(); + } catch { + return ts; + } +} + +export function MemoryDetailDrawer({ row, onClose }: MemoryDetailDrawerProps) { + return ( + + {row.key} + + ) : ( + "Memory" + ) + } + width={720} + destroyOnClose + > + {row && ( + + +
+ + Memory ID + + + {row.memory_id} + +
+
+ + User ID + + {row.user_id ?? "-"} +
+
+ + Team ID + + {row.team_id ?? "-"} +
+
+
+ Value + + {row.value} + +
+ {row.metadata !== undefined && row.metadata !== null && ( +
+ Metadata + + {JSON.stringify(row.metadata, null, 2)} + +
+ )} + ·} wrap size="small" style={{ color: "rgba(0,0,0,0.45)" }}> + + Created {formatTimestamp(row.created_at)} + {row.created_by ? ` by ${row.created_by}` : ""} + + + Updated {formatTimestamp(row.updated_at)} + {row.updated_by ? ` by ${row.updated_by}` : ""} + + +
+ )} +
+ ); +} + +export default MemoryDetailDrawer; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTable.test.tsx new file mode 100644 index 00000000000..f664c650cd4 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTable.test.tsx @@ -0,0 +1,170 @@ +import { PaginationState } from "@tanstack/react-table"; +import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import React, { useState } from "react"; +import { describe, expect, it, vi } from "vitest"; + +import { MemoryRow } from "@/components/networking"; + +import { MemoryTable } from "./MemoryTable"; + +const makeMemory = (overrides: Partial = {}): MemoryRow => ({ + memory_id: "mem-1", + key: "user:profile", + value: "The user prefers concise answers.", + metadata: null, + user_id: "user-42", + team_id: "team-7", + updated_at: "2024-05-01T12:00:00Z", + ...overrides, +}); + +const baseProps = { + data: [makeMemory()], + isLoading: false, + rowCount: 1, + pagination: { pageIndex: 0, pageSize: 50 } as PaginationState, + onPaginationChange: vi.fn(), + searchValue: "", + onSearchChange: vi.fn(), + isRefreshing: false, + onRefresh: vi.fn(), + hasActiveSearch: false, + onViewClick: vi.fn(), + onEditClick: vi.fn(), + onDeleteClick: vi.fn(), +}; + +describe("MemoryTable", () => { + it("renders every column header", () => { + render(); + for (const header of ["ID", "Name", "Preview", "User ID", "Team ID", "Updated"]) { + expect(screen.getByText(header)).toBeInTheDocument(); + } + }); + + it("opens the detail view when the ID identity cell is clicked", async () => { + const user = userEvent.setup(); + const onViewClick = vi.fn(); + const row = makeMemory({ memory_id: "mem-click" }); + render(); + + await user.click(screen.getByText("mem-click")); + + expect(onViewClick).toHaveBeenCalledTimes(1); + expect(onViewClick).toHaveBeenCalledWith(row); + }); + + it("routes each overflow-menu action to its callback with the row", async () => { + const user = userEvent.setup(); + const onViewClick = vi.fn(); + const onEditClick = vi.fn(); + const onDeleteClick = vi.fn(); + const row = makeMemory({ memory_id: "mem-9" }); + render( + , + ); + + await user.click(screen.getByTestId("memory-actions-mem-9")); + await user.click(await screen.findByTestId("memory-action-edit")); + expect(onEditClick).toHaveBeenCalledWith(row); + expect(onViewClick).not.toHaveBeenCalled(); + expect(onDeleteClick).not.toHaveBeenCalled(); + + await user.click(screen.getByTestId("memory-actions-mem-9")); + await user.click(await screen.findByTestId("memory-action-delete")); + expect(onDeleteClick).toHaveBeenCalledWith(row); + + await user.click(screen.getByTestId("memory-actions-mem-9")); + await user.click(await screen.findByTestId("memory-action-view")); + expect(onViewClick).toHaveBeenCalledWith(row); + }); + + it("shows the empty-only copy when there is no data and no active search", () => { + render(); + expect(screen.getByText("No memories stored yet")).toBeInTheDocument(); + expect(screen.queryByText("No matching memories")).not.toBeInTheDocument(); + }); + + it("shows the filtered-empty copy when a search is active", () => { + render(); + expect(screen.getByText("No matching memories")).toBeInTheDocument(); + expect(screen.queryByText("No memories stored yet")).not.toBeInTheDocument(); + }); + + it("renders loading skeleton rows instead of the empty state while loading", () => { + render(); + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("No memories stored yet")).not.toBeInTheDocument(); + }); + + it("drives the pagination footer from the server rowCount, not the page's row length", () => { + render(); + const range = screen.getByTestId("pagination-range"); + expect(range).toHaveTextContent("Showing 1-50 of 120"); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 3"); + expect(screen.getByTestId("pagination-next")).toBeEnabled(); + }); + + it("advances the page through the server pagination handler", async () => { + const user = userEvent.setup(); + const onPaginationChange = vi.fn(); + render(); + + await user.click(screen.getByTestId("pagination-next")); + + expect(onPaginationChange).toHaveBeenCalled(); + }); + + it("forwards toolbar search input and refresh to their callbacks", async () => { + const user = userEvent.setup(); + const onSearchChange = vi.fn(); + const onRefresh = vi.fn(); + render(); + + await user.type(screen.getByTestId("datatable-search"), "u"); + expect(onSearchChange).toHaveBeenCalledWith("u"); + + await user.click(screen.getByTestId("datatable-refresh")); + expect(onRefresh).toHaveBeenCalledTimes(1); + }); + + it("keeps the page in range when the rows-per-page selector shrinks the page count", async () => { + const user = userEvent.setup(); + const rowCount = 120; + const seen: PaginationState[] = []; + + function Harness() { + const [pagination, setPagination] = useState({ pageIndex: 4, pageSize: 25 }); + seen.push(pagination); + return ( + + ); + } + + render(); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 5 of 5"); + + await user.click(screen.getByTestId("pagination-page-size")); + await user.click(await screen.findByRole("option", { name: "100" })); + + const final = seen[seen.length - 1]; + expect(final.pageSize).toBe(100); + expect(final.pageIndex).toBeLessThanOrEqual(Math.ceil(rowCount / final.pageSize) - 1); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 2 of 2"); + }); + + it("renders secondary id and date cells for the row", () => { + render(); + const table = screen.getByRole("table"); + expect(within(table).getByText("user-42")).toBeInTheDocument(); + expect(within(table).getByText("team-7")).toBeInTheDocument(); + expect(within(table).getByText("user:profile")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTable.tsx new file mode 100644 index 00000000000..50dd04ee14c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTable.tsx @@ -0,0 +1,94 @@ +"use client"; + +import { OnChangeFn, PaginationState } from "@tanstack/react-table"; +import { Database } from "lucide-react"; +import React, { useMemo } from "react"; + +import { MemoryRow } from "@/components/networking"; +import { DataTable, DataTableToolbar } from "@/components/shared/DataTable"; + +import { getMemoryTableColumns } from "./MemoryTableColumns"; + +interface MemoryTableProps { + data: MemoryRow[]; + isLoading: boolean; + rowCount: number; + pagination: PaginationState; + onPaginationChange: OnChangeFn; + searchValue: string; + onSearchChange: (value: string) => void; + isRefreshing: boolean; + onRefresh: () => void; + hasActiveSearch: boolean; + onViewClick: (row: MemoryRow) => void; + onEditClick: (row: MemoryRow) => void; + onDeleteClick: (row: MemoryRow) => void; +} + +function MemoryEmptyState({ hasActiveSearch }: { hasActiveSearch: boolean }) { + return ( +
+
+ +
+
+ {hasActiveSearch ? "No matching memories" : "No memories stored yet"} +
+
+ {hasActiveSearch + ? "No memories have keys starting with your search." + : "Memories your agents store under /v1/memory will appear here."} +
+
+ ); +} + +export function MemoryTable({ + data, + isLoading, + rowCount, + pagination, + onPaginationChange, + searchValue, + onSearchChange, + isRefreshing, + onRefresh, + hasActiveSearch, + onViewClick, + onEditClick, + onDeleteClick, +}: MemoryTableProps) { + const columns = useMemo(() => { + const columnDeps = { onViewClick, onEditClick, onDeleteClick }; + return getMemoryTableColumns(columnDeps); + }, [onViewClick, onEditClick, onDeleteClick]); + + return ( + row.memory_id} + paginationMode="server" + pagination={pagination} + onPaginationChange={onPaginationChange} + rowCount={rowCount} + isLoading={isLoading} + loadingMessage="Loading memories…" + noDataMessage={} + size="compact" + toolbar={(table) => ( + + )} + /> + ); +} + +export default MemoryTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTableColumns.tsx new file mode 100644 index 00000000000..6b2a6b08704 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryTableColumns.tsx @@ -0,0 +1,150 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { Eye, MoreHorizontal, Pencil, Trash2 } from "lucide-react"; + +import { MemoryRow } from "@/components/networking"; +import { DateCell, IdCell, IdentityCell } from "@/components/shared/table_cells"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +interface MemoryRowActionsProps { + row: MemoryRow; + onViewClick: (row: MemoryRow) => void; + onEditClick: (row: MemoryRow) => void; + onDeleteClick: (row: MemoryRow) => void; +} + +function MemoryRowActions({ row, onViewClick, onEditClick, onDeleteClick }: MemoryRowActionsProps) { + return ( + + + + + + onViewClick(row)}> + + View + + onEditClick(row)}> + + Edit + + + onDeleteClick(row)}> + + Delete + + + + ); +} + +export interface MemoryTableColumnsDeps { + onViewClick: (row: MemoryRow) => void; + onEditClick: (row: MemoryRow) => void; + onDeleteClick: (row: MemoryRow) => void; +} + +export const getMemoryTableColumns = ({ + onViewClick, + onEditClick, + onDeleteClick, +}: MemoryTableColumnsDeps): ColumnDef[] => [ + { + id: "memory_id", + accessorKey: "memory_id", + meta: { title: "ID" }, + header: "ID", + size: 180, + enableSorting: false, + cell: ({ row }) => ( + onViewClick(row.original)} + /> + ), + }, + { + id: "key", + accessorKey: "key", + meta: { title: "Name" }, + header: "Name", + size: 200, + enableSorting: false, + cell: ({ row }) => ( + + {row.original.key} + + ), + }, + { + id: "value", + accessorKey: "value", + meta: { title: "Preview" }, + header: "Preview", + enableSorting: false, + cell: ({ row }) => ( + + {row.original.value || "-"} + + ), + }, + { + id: "user_id", + accessorKey: "user_id", + meta: { title: "User ID" }, + header: "User ID", + size: 160, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "team_id", + accessorKey: "team_id", + meta: { title: "Team ID" }, + header: "Team ID", + size: 160, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "updated_at", + accessorKey: "updated_at", + meta: { title: "Updated" }, + header: "Updated", + size: 170, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => ( +
+ +
+ ), + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.test.tsx new file mode 100644 index 00000000000..f415c99225a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.test.tsx @@ -0,0 +1,45 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render } from "@testing-library/react"; +import React from "react"; +import { describe, expect, it, vi } from "vitest"; + +import { MemoryRow } from "@/components/networking"; + +import { MemoryView } from "./MemoryView"; + +interface CapturedTableProps { + isLoading: boolean; + rowCount: number; + data: MemoryRow[]; + hasActiveSearch: boolean; +} + +const captured = vi.hoisted(() => ({ current: null as CapturedTableProps | null })); + +vi.mock("./MemoryTable", () => ({ + MemoryTable: function MemoryTableMock(props: CapturedTableProps) { + captured.current = props; + return
; + }, +})); + +const renderView = (accessToken: string | null) => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + return render( + + + , + ); +}; + +describe("MemoryView", () => { + it("keeps the table out of the skeleton state when the token is null (disabled query)", () => { + renderView(null); + + expect(captured.current).not.toBeNull(); + expect(captured.current?.isLoading).toBe(false); + expect(captured.current?.data).toEqual([]); + expect(captured.current?.rowCount).toBe(0); + expect(captured.current?.hasActiveSearch).toBe(false); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.tsx index 4ee784f4664..fcb15978f47 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/memory/_components/MemoryView.tsx @@ -1,21 +1,19 @@ "use client"; -import React, { useMemo, useState } from "react"; +import { useDebouncedValue } from "@tanstack/react-pacer/debouncer"; import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query"; -import { Button, Card, Drawer, Empty, Input, Space, Table, Typography, message } from "antd"; -import type { ColumnsType } from "antd/es/table"; -import { - DeleteOutlined, - EditOutlined, - EyeOutlined, - PlusOutlined, - ReloadOutlined, - SearchOutlined, -} from "@ant-design/icons"; +import type { PaginationState } from "@tanstack/react-table"; +import { PlusOutlined } from "@ant-design/icons"; +import { Button, Space, Typography, message } from "antd"; +import React, { useCallback, useMemo, useState } from "react"; + import { MemoryRow, createMemory, deleteMemory, fetchMemoryList, updateMemory } from "@/components/networking"; -import { DateCell, IdCell } from "@/components/shared/table_cells"; -import { MemoryEditModal } from "./MemoryEditModal"; import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; +import { DEBOUNCE_WAIT_MS } from "@/utils/debounceConstants"; + +import { MemoryDetailDrawer } from "./MemoryDetailDrawer"; +import { MemoryEditModal } from "./MemoryEditModal"; +import { MemoryTable } from "./MemoryTable"; const { Text, Paragraph, Title } = Typography; @@ -25,38 +23,16 @@ interface MemoryViewProps { userRole: string | null; } -function previewValue(value: string, max = 120): string { - if (!value) return ""; - const trimmed = value.trim(); - if (trimmed.length <= max) return trimmed; - return `${trimmed.slice(0, max)}…`; -} - -function formatTimestamp(ts?: string): string { - if (!ts) return "—"; - try { - const d = new Date(ts); - return d.toLocaleString(); - } catch { - return ts; - } -} - -const PAGE_SIZE = 50; +const DEFAULT_PAGE_SIZE = 50; export const MemoryView: React.FC = ({ accessToken }) => { const [searchInput, setSearchInput] = useState(""); - const [appliedSearch, setAppliedSearch] = useState(""); + const [debouncedSearch] = useDebouncedValue(searchInput, { wait: DEBOUNCE_WAIT_MS }); + const [pagination, setPagination] = useState({ pageIndex: 0, pageSize: DEFAULT_PAGE_SIZE }); const [detailRow, setDetailRow] = useState(null); const [editRow, setEditRow] = useState(null); const [deleteRow, setDeleteRow] = useState(null); const [isCreateOpen, setIsCreateOpen] = useState(false); - const [currentPage, setCurrentPage] = useState(1); - - // Reset to page 1 whenever the filter changes. - React.useEffect(() => { - setCurrentPage(1); - }, [appliedSearch]); const queryClient = useQueryClient(); // React Query key prefix for all memory-list variants (paged + filtered). @@ -65,15 +41,15 @@ export const MemoryView: React.FC = ({ accessToken }) => { const MEMORY_LIST_KEY = "memoryList" as const; const { data, isLoading, isFetching } = useQuery({ - queryKey: [MEMORY_LIST_KEY, appliedSearch, currentPage], + queryKey: [MEMORY_LIST_KEY, debouncedSearch, pagination.pageIndex, pagination.pageSize], queryFn: () => { if (!accessToken) throw new Error("Access token required"); // Prefix search matches the Redis-style mental model (namespace scan): // typing "user:" finds "user:profile", "user:prefs", etc. return fetchMemoryList(accessToken, { - keyPrefix: appliedSearch || undefined, - page: currentPage, - pageSize: PAGE_SIZE, + keyPrefix: debouncedSearch || undefined, + page: pagination.pageIndex + 1, + pageSize: pagination.pageSize, }); }, enabled: !!accessToken, @@ -88,7 +64,10 @@ export const MemoryView: React.FC = ({ accessToken }) => { // refetches from scratch (pagination + filter-aware). // - on error: surface the message via antd `message.error`. - const invalidateList = () => queryClient.invalidateQueries({ queryKey: [MEMORY_LIST_KEY] }); + const invalidateList = useCallback( + () => queryClient.invalidateQueries({ queryKey: [MEMORY_LIST_KEY] }), + [queryClient], + ); const createMutation = useMutation({ mutationFn: (args: { key: string; value: string; metadata: unknown }) => { @@ -133,9 +112,14 @@ export const MemoryView: React.FC = ({ accessToken }) => { }, }); - const handleDelete = (row: MemoryRow) => { - setDeleteRow(row); - }; + const handleSearchChange = useCallback((value: string) => { + setSearchInput(value); + setPagination((prev) => ({ ...prev, pageIndex: 0 })); + }, []); + + const handleView = useCallback((row: MemoryRow) => setDetailRow(row), []); + const handleEdit = useCallback((row: MemoryRow) => setEditRow(row), []); + const handleDelete = useCallback((row: MemoryRow) => setDeleteRow(row), []); const confirmDelete = async () => { if (!deleteRow) return; @@ -192,242 +176,43 @@ export const MemoryView: React.FC = ({ accessToken }) => { } }; - const columns: ColumnsType = [ - { - title: "ID", - dataIndex: "memory_id", - key: "memory_id", - width: 140, - render: (_: unknown, r: MemoryRow) => setDetailRow(r)} />, - }, - { - title: "Name", - dataIndex: "key", - key: "key", - width: 200, - render: (k: string) => {k}, - // No client-side sorter: pagination is server-side, so a client sort - // would only reorder the current page and mislead users into thinking - // the whole list is sorted. Backend returns rows ordered by - // `updated_at DESC`; use the prefix filter for discovery by name. - }, - { - title: "Preview", - dataIndex: "value", - key: "value", - render: (v: string) => ( - - {previewValue(v)} - - ), - }, - { - title: "User ID", - dataIndex: "user_id", - key: "user_id", - width: 160, - render: (uid?: string | null) => , - }, - { - title: "Team ID", - dataIndex: "team_id", - key: "team_id", - width: 160, - render: (tid?: string | null) => , - }, - { - title: "Updated", - dataIndex: "updated_at", - key: "updated_at", - width: 180, - render: (ts?: string) => , - // No sorter — backend already returns rows in `updated_at DESC` order, - // and a client-side sorter on a paginated view would only affect the - // current page. - }, - { - title: "", - key: "actions", - width: 140, - render: (_: unknown, r: MemoryRow) => ( - -
- - - - } - value={searchInput} - onChange={(e) => setSearchInput(e.target.value)} - onPressEnter={() => setAppliedSearch(searchInput.trim())} - onClear={() => { - setSearchInput(""); - setAppliedSearch(""); - }} - style={{ width: 280 }} - /> - - - - - - - `${range[0]}–${range[1]} of ${n}`, - onChange: (page) => setCurrentPage(page), - }} - locale={{ - emptyText: ( - - ), - }} - /> - + {/* Detail drawer */} - setDetailRow(null)} - title={ - detailRow ? ( - - {detailRow.key} - - ) : ( - "Memory" - ) - } - width={720} - destroyOnClose - > - {detailRow && ( - - -
- - Memory ID - - - {detailRow.memory_id} - -
-
- - User ID - - {detailRow.user_id ?? "-"} -
-
- - Team ID - - {detailRow.team_id ?? "-"} -
-
-
- Value - - {detailRow.value} - -
- {detailRow.metadata !== undefined && detailRow.metadata !== null && ( -
- Metadata - - {JSON.stringify(detailRow.metadata, null, 2)} - -
- )} - ·} wrap size="small" style={{ color: "rgba(0,0,0,0.45)" }}> - - Created {formatTimestamp(detailRow.created_at)} - {detailRow.created_by ? ` by ${detailRow.created_by}` : ""} - - - Updated {formatTimestamp(detailRow.updated_at)} - {detailRow.updated_by ? ` by ${detailRow.updated_by}` : ""} - - -
- )} -
+ setDetailRow(null)} /> {/* Create / edit modal */} { - let store: Record = {}; - return { - getItem: (key: string) => store[key] || null, - setItem: (key: string, value: string) => { - store[key] = value; - }, - removeItem: (key: string) => { - delete store[key]; - }, - clear: () => { - store = {}; - }, - }; -})(); -Object.defineProperty(window, "localStorage", { value: localStorageMock }); - -// Minimal stubs to avoid Next.js router and network usage during render -vi.mock("@/components/networking", () => ({ - credentialListCall: vi.fn().mockResolvedValue({ credentials: [] }), - modelInfoCall: vi.fn().mockResolvedValue({ data: [] }), - modelCostMap: vi.fn().mockResolvedValue({}), - getPassThroughEndpointsCall: vi.fn().mockResolvedValue({ endpoints: {} }), - getCallbacksCall: vi.fn().mockResolvedValue({ router_settings: {} }), - setCallbacksCall: vi.fn().mockResolvedValue(undefined), - getUiSettings: vi.fn().mockResolvedValue({ values: {} }), - latestHealthChecksCall: vi.fn().mockResolvedValue({ latest_health_checks: {} }), - getModelCostMapReloadStatus: vi.fn().mockResolvedValue({}), -})); - -vi.mock("@/app/(dashboard)/models-and-endpoints/components/ModelAnalyticsTab/ModelAnalyticsTab", () => ({ - default: () => null, -})); - -vi.mock("@/components/add_model/add_auto_router_tab", () => ({ - default: () => null, -})); - -vi.mock("@/components/add_model/AddModelForm", () => ({ - default: () => null, -})); - -const mockHealthCheckComponent = vi.fn((_props: { all_models_on_proxy?: string[] }) => null); -vi.mock("@/components/model_dashboard/HealthCheckComponent", () => ({ - default: (props: { all_models_on_proxy?: string[] }) => { - mockHealthCheckComponent(props); - return null; - }, -})); - -vi.mock("@/app/(dashboard)/hooks/useTeams", () => ({ - default: () => ({ - teams: [], - setTeams: vi.fn(), - }), -})); - -const mockUseModelsInfo = vi.fn(); -vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ - useModelsInfo: () => mockUseModelsInfo(), -})); - -const mockUseUISettings = vi.fn(); -vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ - useUISettings: () => mockUseUISettings(), -})); - -const mockUseModelCostMap = vi.fn(); -vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ - useModelCostMap: () => mockUseModelCostMap(), -})); - -const mockUseAuthorized = vi.fn(); -vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ - default: () => mockUseAuthorized(), -})); - -const createQueryClient = () => - new QueryClient({ - defaultOptions: { queries: { retry: false, gcTime: 0 } }, - }); - -describe("ModelsAndEndpointsView", () => { - beforeEach(() => { - mockUseModelsInfo.mockReturnValue({ - data: { data: [] }, - isLoading: false, - refetch: vi.fn(), - }); - mockUseUISettings.mockReturnValue({ - data: { values: {} }, - }); - mockUseModelCostMap.mockReturnValue({ - data: {}, - isLoading: false, - error: null, - }); - mockUseAuthorized.mockReturnValue({ - accessToken: "123", - token: "123", - userRole: "Admin", - userId: "123", - }); - // eslint-disable-next-line @typescript-eslint/no-explicit-any - (global as any).ResizeObserver = class { - observe() {} - unobserve() {} - disconnect() {} - }; - }); - - it("should render the models and endpoints view", async () => { - const queryClient = createQueryClient(); - const { findByText } = render( - - - , - ); - expect(await findByText("Model Management", {}, { timeout: 10000 })).toBeInTheDocument(); - }); - - it("should show Cost Optimization feedback banner by default", async () => { - localStorageMock.clear(); - const queryClient = createQueryClient(); - const { findByText } = render( - - - , - ); - expect(await findByText("Help shape cost optimization", {}, { timeout: 10000 })).toBeInTheDocument(); - }); - - it("should hide Cost Optimization feedback banner when dismiss button is clicked and persist to localStorage", async () => { - localStorageMock.clear(); - const queryClient = createQueryClient(); - const { findByText, queryByText, container } = render( - - - , - ); - - // Wait for banner to appear - expect(await findByText("Help shape cost optimization", {}, { timeout: 10000 })).toBeInTheDocument(); - - // Find and click dismiss button (X button) - const dismissButton = container.querySelector('button[aria-label="Dismiss banner"]'); - expect(dismissButton).not.toBeNull(); - fireEvent.click(dismissButton!); - - // Banner should be hidden - expect(queryByText("Help shape cost optimization")).not.toBeInTheDocument(); - - // LocalStorage should be updated - expect(localStorageMock.getItem("hideCostOptimizationFeedbackBanner")).toBe("true"); - }); - - it("should keep Cost Optimization feedback banner hidden across remounts once dismissed", async () => { - // Set localStorage to hide banner - localStorageMock.setItem("hideCostOptimizationFeedbackBanner", "true"); - const queryClient = createQueryClient(); - const { findByText, queryByText } = render( - - - , - ); - - // Wait for component to render - await findByText("Model Management", {}, { timeout: 10000 }); - - // Banner should not be visible - expect(queryByText("Help shape cost optimization")).not.toBeInTheDocument(); - }); - - it("should pass model IDs (not model names) to HealthCheckComponent as all_models_on_proxy", async () => { - mockHealthCheckComponent.mockClear(); - const modelDataWithIds = { - data: [ - { model_name: "gpt-4", model_info: { id: "deployment-id-1" } }, - { model_name: "gpt-4", model_info: { id: "deployment-id-2" } }, - ], - }; - mockUseModelsInfo.mockReturnValue({ - data: { data: modelDataWithIds.data }, - isLoading: false, - refetch: vi.fn(), - }); - - const queryClient = createQueryClient(); - const { getByRole } = render( - - - , - ); - - const healthStatusTab = getByRole("tab", { name: "Health Status" }); - await act(async () => { - healthStatusTab.click(); - }); - - expect(mockHealthCheckComponent).toHaveBeenCalled(); - const healthCheckProps = mockHealthCheckComponent.mock.calls[0][0]; - expect(healthCheckProps.all_models_on_proxy).toEqual(["deployment-id-1", "deployment-id-2"]); - expect(healthCheckProps.all_models_on_proxy).not.toContain("gpt-4"); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx deleted file mode 100644 index 672bcc2aa95..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ /dev/null @@ -1,491 +0,0 @@ -import { useCredentials } from "@/app/(dashboard)/hooks/credentials/useCredentials"; -import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; -import { useModelsInfo } from "@/app/(dashboard)/hooks/models/useModels"; -import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; -import { useUpdateRetryPolicy } from "@/app/(dashboard)/hooks/routerSettings/useUpdateRetryPolicy"; -import AllModelsTab from "@/app/(dashboard)/models-and-endpoints/components/AllModelsTab"; -import CostOptimizationFeedbackBanner from "@/components/molecules/cost_optimization_feedback_banner"; -import ModelRetrySettingsTab from "@/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab"; -import PriceDataManagementTab from "@/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab"; -import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit"; -import { Team } from "@/components/key_team_helpers/key_list"; -import CredentialsPanel from "@/components/model_add/CredentialsPanel"; -import { getCallbacksCall } from "@/components/networking"; -import { Providers, getPlaceholder, getProviderModels } from "@/components/provider_info_helpers"; -import { getDisplayModelName } from "@/components/view_model/model_name_display"; -import { transformModelData } from "./utils/modelDataTransformer"; -import { all_admin_roles, internalUserRoles, isProxyAdminRole, isUserTeamAdminForAnyTeam } from "@/utils/roles"; -import { RefreshIcon } from "@heroicons/react/outline"; -import { useQueryClient } from "@tanstack/react-query"; -import { Col, Grid, Icon, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; -import type { UploadProps } from "antd"; -import { Form } from "antd"; -import React, { useCallback, useEffect, useMemo, useState } from "react"; -import AddModelTab from "../../../components/add_model/add_model_tab"; -import HealthCheckComponent from "../../../components/model_dashboard/HealthCheckComponent"; -import ModelGroupAliasSettings from "../../../components/model_group_alias_settings"; -import ModelInfoView from "../../../components/model_info_view"; -import NotificationsManager from "../../../components/molecules/notifications_manager"; -import PassThroughSettings from "../../../components/PassThroughSettings/PassThroughSettings"; -import TeamInfoView from "../../../components/team/TeamInfo"; -import useAuthorized from "../hooks/useAuthorized"; - -interface ModelDashboardProps { - premiumUser: boolean; - teams: Team[] | null; -} - -interface RetryPolicyObject { - [key: string]: { [retryPolicyKey: string]: number } | undefined; -} - -interface GlobalRetryPolicyObject { - [retryPolicyKey: string]: number; -} - -interface RouterSettings { - model_group_retry_policy?: RetryPolicyObject | null; - retry_policy?: GlobalRetryPolicyObject | null; - num_retries?: number | null; - model_group_alias?: { [key: string]: string } | null; -} - -const HEALTH_PAGE_SIZE = 50; - -const ModelsAndEndpointsView: React.FC = ({ premiumUser, teams }) => { - const { accessToken, token, userRole, userId: userID } = useAuthorized(); - const [addModelForm] = Form.useForm(); - const [lastRefreshed, setLastRefreshed] = useState(""); - const [providerModels, setProviderModels] = useState>([]); - const [selectedProvider, setSelectedProvider] = useState(Providers.Anthropic); - const [selectedModelGroup, setSelectedModelGroup] = useState(null); - - const [retryScope, setRetryScope] = useState("global"); - const [modelGroupRetryPolicy, setModelGroupRetryPolicy] = useState(null); - const [globalRetryPolicy, setGlobalRetryPolicy] = useState(null); - const [defaultRetry, setDefaultRetry] = useState(0); - const [modelGroupAlias, setModelGroupAlias] = useState<{ [key: string]: string }>({}); - const [showAdvancedSettings, setShowAdvancedSettings] = useState(false); - const [selectedModelId, setSelectedModelId] = useState(null); - const [selectedTeamId, setSelectedTeamId] = useState(null); - const [selectedTabIndex, setSelectedTabIndex] = useState(0); - const [healthCurrentPage, setHealthCurrentPage] = useState(1); - - const queryClient = useQueryClient(); - const { data: modelDataResponse, isLoading: isLoadingModels, refetch: refetchModels } = useModelsInfo(); - const { data: healthModelDataResponse, isLoading: isLoadingHealthModels } = useModelsInfo( - healthCurrentPage, - HEALTH_PAGE_SIZE, - ); - const { data: modelCostMapData, isLoading: isLoadingModelCostMap } = useModelCostMap(); - const { data: credentialsResponse, isLoading: isLoadingCredentials } = useCredentials(); - const credentialsList = credentialsResponse?.credentials || []; - const { data: uiSettings, isLoading: isLoadingUISettings } = useUISettings(); - const updateRetryPolicy = useUpdateRetryPolicy(accessToken); - - const availableModelGroups = useMemo(() => { - if (!modelDataResponse?.data) return []; - const allModelGroups = new Set(); - for (const model of modelDataResponse.data) { - allModelGroups.add(model.model_name); - } - return Array.from(allModelGroups).sort(); - }, [modelDataResponse?.data]); - - const availableModelAccessGroups = useMemo(() => { - if (!modelDataResponse?.data) return []; - const allModelAccessGroups = new Set(); - for (const model of modelDataResponse.data) { - const modelInfo = model.model_info; - if (modelInfo?.access_groups) { - for (const group of modelInfo.access_groups) { - allModelAccessGroups.add(group); - } - } - } - return Array.from(allModelAccessGroups); - }, [modelDataResponse?.data]); - - const allModelsOnProxy = useMemo(() => { - if (!modelDataResponse?.data) return []; - return modelDataResponse.data.map((model: any) => model.model_name); - }, [modelDataResponse?.data]); - - const healthModelIdsOnProxy = useMemo(() => { - if (!healthModelDataResponse?.data) return []; - return healthModelDataResponse.data - .map((model: any) => model.model_info?.id) - .filter((id: string | undefined): id is string => Boolean(id)); - }, [healthModelDataResponse?.data]); - - const getProviderFromModel = (model: string) => { - if (modelCostMapData !== null && modelCostMapData !== undefined) { - if (typeof modelCostMapData == "object" && model in modelCostMapData) { - return modelCostMapData[model]["litellm_provider"]; - } - } - return "openai"; - }; - - const processedModelData = useMemo(() => { - if (!modelDataResponse?.data) return { data: [] }; - return transformModelData(modelDataResponse, getProviderFromModel); - }, [modelDataResponse?.data, getProviderFromModel]); - - const processedHealthModelData = useMemo(() => { - if (!healthModelDataResponse?.data) return { data: [] }; - return transformModelData(healthModelDataResponse, getProviderFromModel); - }, [healthModelDataResponse?.data, getProviderFromModel]); - - const healthPaginationMeta = useMemo(() => { - return { - total_count: healthModelDataResponse?.total_count ?? 0, - current_page: healthModelDataResponse?.current_page ?? healthCurrentPage, - total_pages: healthModelDataResponse?.total_pages ?? 1, - size: healthModelDataResponse?.size ?? HEALTH_PAGE_SIZE, - }; - }, [healthModelDataResponse, healthCurrentPage]); - - const isProxyAdmin = userRole && isProxyAdminRole(userRole); - const isInternalUser = userRole && internalUserRoles.includes(userRole); - const isUserTeamAdmin = userID && isUserTeamAdminForAnyTeam(teams, userID); - const addModelDisabledForInternalUsers = - isInternalUser && uiSettings?.values?.disable_model_add_for_internal_users === true; - // Hide tab if user is NOT a proxy admin AND (internal user with setting enabled OR not a team admin) - const shouldHideAddModelTab = !isProxyAdmin && (addModelDisabledForInternalUsers || !isUserTeamAdmin); - - const setProviderModelsFn = (provider: Providers) => { - const _providerModels = getProviderModels(provider, modelCostMapData); - setProviderModels(_providerModels); - }; - - const uploadProps: UploadProps = { - name: "file", - accept: ".json", - pastable: false, - beforeUpload: (file) => { - if (file.type === "application/json") { - const reader = new FileReader(); - reader.onload = (e) => { - if (e.target) { - const jsonStr = e.target.result as string; - addModelForm.setFieldsValue({ vertex_credentials: jsonStr }); - } - }; - reader.readAsText(file); - } - return false; - }, - onChange(info) { - if (info.file.status === "done") { - NotificationsManager.success(`${info.file.name} file uploaded successfully`); - } else if (info.file.status === "error") { - NotificationsManager.fromBackend(`${info.file.name} file upload failed.`); - } - }, - }; - - const handleRefreshClick = () => { - const currentDate = new Date(); - setLastRefreshed(currentDate.toLocaleTimeString([], { hour: "2-digit", minute: "2-digit" })); - setHealthCurrentPage(1); - queryClient.invalidateQueries({ queryKey: ["models", "list"] }); - refetchModels(); - }; - - const fetchRouterSettings = useCallback(async (): Promise => { - if (!accessToken || !userID || !userRole) { - return null; - } - try { - const routerSettingsInfo = await getCallbacksCall(accessToken, userID, userRole); - return routerSettingsInfo.router_settings; - } catch (error) { - console.error("Error fetching model data:", error); - return null; - } - }, [accessToken, userID, userRole]); - - const applyRouterSettings = useCallback((routerSettings: RouterSettings) => { - setModelGroupRetryPolicy(routerSettings.model_group_retry_policy ?? null); - setGlobalRetryPolicy(routerSettings.retry_policy ?? null); - setDefaultRetry(routerSettings.num_retries ?? 2); - setModelGroupAlias(routerSettings.model_group_alias || {}); - }, []); - - const loadRetrySettings = useCallback(async () => { - const routerSettings = await fetchRouterSettings(); - if (routerSettings) { - applyRouterSettings(routerSettings); - } - }, [fetchRouterSettings, applyRouterSettings]); - - const handleSaveRetrySettings = () => { - updateRetryPolicy.mutate( - { - retry_policy: globalRetryPolicy, - model_group_retry_policy: modelGroupRetryPolicy, - }, - { - onSuccess: () => { - NotificationsManager.success("Retry settings saved successfully"); - loadRetrySettings(); - }, - onError: () => { - NotificationsManager.fromBackend("Failed to save retry settings"); - }, - }, - ); - }; - - useEffect(() => { - if (!accessToken || !token || !userRole || !userID || !modelDataResponse) { - return; - } - let active = true; - void (async () => { - const routerSettings = await fetchRouterSettings(); - if (active && routerSettings) { - applyRouterSettings(routerSettings); - } - })(); - return () => { - active = false; - }; - }, [accessToken, token, userRole, userID, modelDataResponse, fetchRouterSettings, applyRouterSettings]); - - const isLoading = isLoadingModels || isLoadingModelCostMap || isLoadingCredentials || isLoadingUISettings; - - // Admin Viewer can view all models read-only — page render proceeds; the - // individual write-action tabs (Add Model, LLM Credentials, etc.) are - // gated separately below. - - const handleOk = async () => { - try { - const values = await addModelForm.validateFields(); - await handleAddModelSubmit(values, accessToken, addModelForm, handleRefreshClick); - } catch (error: any) { - const errorMessages = - error.errorFields - ?.map((field: any) => { - return `${field.name.join(".")}: ${field.errors.join(", ")}`; - }) - .join(" | ") || "Unknown validation error"; - NotificationsManager.fromBackend(`Please fill in the following required fields: ${errorMessages}`); - } - }; - - Object.keys(Providers).find((key) => (Providers as { [index: string]: any })[key] === selectedProvider); - // If a team is selected, render TeamInfoView in full page layout - if (selectedTeamId) { - return ( -
- setSelectedTeamId(null)} - accessToken={accessToken} - is_team_admin={userRole === "Admin"} - is_proxy_admin={userRole === "Proxy Admin"} - userModels={allModelsOnProxy} - editTeam={false} - onUpdate={handleRefreshClick} - premiumUser={premiumUser} - /> -
- ); - } - - return ( -
- -
- {/* Model Management Header */} -
-
-

Model Management

- {!all_admin_roles.includes(userRole) ? ( -

Add models for teams you are an admin for.

- ) : ( -

Add and manage models for the proxy

- )} -
-
- - {/* Cost Optimization Feedback Banner */} - - {selectedModelId && !isLoading ? ( - { - setSelectedModelId(null); - }} - accessToken={accessToken} - userID={userID} - userRole={userRole} - onModelUpdate={(updatedModel) => { - queryClient.invalidateQueries({ queryKey: ["models", "list"] }); - handleRefreshClick(); - }} - modelAccessGroups={availableModelAccessGroups} - /> - ) : ( - (() => { - // Build a single source-of-truth list of {tab, panel} pairs. - // Conditionally-hidden tabs (e.g. "Add Model" for non-admin) get - // filtered out as a unit so tab indices and panel indices can - // never drift apart — Tremor's TabList and TabPanels filter - // falsy children inconsistently, which previously caused - // "click LLM Credentials, see nothing" for Admin Viewer. - const isAdmin = all_admin_roles.includes(userRole); - const visibleTabs: Array<{ tab: React.ReactElement; panel: React.ReactElement }> = [ - { - tab: {isAdmin ? "All Models" : "Your Models"}, - panel: ( - - ), - }, - ]; - if (!shouldHideAddModelTab) { - visibleTabs.push({ - tab: Add Model, - panel: ( - - - - ), - }); - } - if (isAdmin) { - visibleTabs.push( - { - tab: LLM Credentials, - panel: ( - - - - ), - }, - { - tab: Pass-Through Endpoints, - panel: ( - - - - ), - }, - { - tab: Health Status, - panel: ( - - - - ), - }, - { - tab: Model Retry Settings, - panel: ( - - ), - }, - { - tab: Model Group Alias, - panel: ( - - - - ), - }, - { - tab: Price Data Reload, - panel: , - }, - ); - } - return ( - - -
{visibleTabs.map((t) => t.tab)}
- -
- {lastRefreshed && Last Refreshed: {lastRefreshed}} - -
-
- {visibleTabs.map((t) => t.panel)} -
- ); - })() - )} - - - - ); -}; - -export default ModelsAndEndpointsView; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/add/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/add/page.tsx new file mode 100644 index 00000000000..7e60ca58fc3 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/add/page.tsx @@ -0,0 +1,59 @@ +"use client"; + +import { Form } from "antd"; +import { useState } from "react"; +import { useQueryClient } from "@tanstack/react-query"; +import AddModelTab from "@/components/add_model/add_model_tab"; +import { handleAddModelSubmit } from "@/components/add_model/handle_add_model_submit"; +import { Providers, getPlaceholder, getProviderModels } from "@/components/provider_info_helpers"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; +import { useCredentials } from "@/app/(dashboard)/hooks/credentials/useCredentials"; +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { vertexCredentialsUploadProps } from "@/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload"; + +export default function AddModelPage() { + const { accessToken, userRole } = useAuthorized(); + const [form] = Form.useForm(); + const queryClient = useQueryClient(); + const { data: modelCostMapData } = useModelCostMap(); + const { data: credentialsResponse } = useCredentials(); + const { data: teams } = useTeams(); + const [selectedProvider, setSelectedProvider] = useState(Providers.Anthropic); + const [providerModels, setProviderModels] = useState([]); + const [showAdvancedSettings, setShowAdvancedSettings] = useState(false); + + const refresh = () => queryClient.invalidateQueries({ queryKey: ["models", "list"] }); + + const handleOk = async () => { + try { + const values = await form.validateFields(); + await handleAddModelSubmit(values, accessToken, form, refresh); + } catch (error: any) { + const errorMessages = + error.errorFields?.map((field: any) => `${field.name.join(".")}: ${field.errors.join(", ")}`).join(" | ") || + "Unknown validation error"; + NotificationsManager.fromBackend(`Please fill in the following required fields: ${errorMessages}`); + } + }; + + return ( + setProviderModels(getProviderModels(provider, modelCostMapData))} + getPlaceholder={getPlaceholder} + uploadProps={vertexCredentialsUploadProps(form)} + showAdvancedSettings={showAdvancedSettings} + setShowAdvancedSettings={setShowAdvancedSettings} + teams={teams ?? null} + credentials={credentialsResponse?.credentials || []} + accessToken={accessToken} + userRole={userRole} + /> + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx index 20c9e805a0d..6efc85c019b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/AllModelsTab.tsx @@ -11,7 +11,7 @@ import { modelDeleteCall, modelPatchUpdateCall } from "@/components/networking"; import { InfoCircleOutlined, SettingOutlined } from "@ant-design/icons"; import { PaginationState, SortingState } from "@tanstack/react-table"; import { useQueryClient } from "@tanstack/react-query"; -import { Grid, TabPanel } from "@tremor/react"; +import { Grid } from "@tremor/react"; import { Badge, Button, Select, Skeleton, Space, Typography } from "antd"; import ModelSettingsModal from "@/components/model_dashboard/ModelSettingsModal/ModelSettingsModal"; import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer"; @@ -232,7 +232,7 @@ const AllModelsTab = ({ }; return ( - +
@@ -600,7 +600,7 @@ const AllModelsTab = ({ onCancel={() => setIsModelSettingsModalVisible(false)} onSuccess={() => setIsModelSettingsModalVisible(false)} /> - +
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx index 5ff761b0663..a4e3c4b958c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab.tsx @@ -1,4 +1,4 @@ -import { Button, Select, SelectItem, TabPanel, Text, Title } from "@tremor/react"; +import { Button, Select, SelectItem, Text, Title } from "@tremor/react"; import { InputNumber } from "antd"; import React from "react"; @@ -64,7 +64,7 @@ const ModelRetrySettingsTab = ({ }; return ( - +
Retry Policy Scope: @@ -132,7 +132,7 @@ const ModelRetrySettingsTab = ({ - +
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx new file mode 100644 index 00000000000..282cc2722db --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.test.tsx @@ -0,0 +1,22 @@ +/* @vitest-environment jsdom */ +import { render } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; +import PriceDataManagementTab from "./PriceDataManagementTab"; + +// Deliberately do NOT mock @tremor/react. These tab components render standalone +// (inside antd Tabs / directly as a route page), no longer inside a Tremor +// . A Tremor root renders nothing without that context, so +// this asserts the component's content is visible on its own — reverting the root +// back to makes the title disappear and fails this test. +vi.mock("@/components/price_data_reload", () => ({ default: () =>
reload
})); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => ({ accessToken: "sk-test" }) })); +vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ + useModelCostMap: () => ({ refetch: vi.fn() }), +})); + +describe("PriceDataManagementTab", () => { + it("renders its content standalone, without a Tremor TabGroup ancestor", () => { + const { getByText } = render(); + expect(getByText("Price Data Management")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx index d44d19879d5..9420643578c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab.tsx @@ -1,4 +1,4 @@ -import { TabPanel, Text, Title } from "@tremor/react"; +import { Text, Title } from "@tremor/react"; import PriceDataReload from "@/components/price_data_reload"; import React from "react"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; @@ -9,7 +9,7 @@ const PriceDataManagementTab = () => { const { refetch: refetchModelCostMap } = useModelCostMap(); return ( - +
Price Data Management @@ -28,7 +28,7 @@ const PriceDataManagementTab = () => { className="w-full" />
- +
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/detailNavigation.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/detailNavigation.test.ts new file mode 100644 index 00000000000..717fdc85e28 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/detailNavigation.test.ts @@ -0,0 +1,52 @@ +/* @vitest-environment jsdom */ +import { act, renderHook } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { useModelDetailRouting } from "./detailNavigation"; + +// The detail overlay is driven by ?model=/?team= on the current path. Under the +// /ui static mount a router.push to the same path (query-only change) is a no-op, +// so navigation goes through history.pushState (client-side, no full reload). +vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams(window.location.search) })); + +describe("useModelDetailRouting", () => { + beforeEach(() => { + window.history.pushState(null, "", "/models-and-endpoints/"); + }); + + it("openModel sets ?model= via history.pushState (no full navigation)", () => { + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useModelDetailRouting()); + act(() => result.current.openModel("abc-1")); + expect(spy).toHaveBeenCalledWith(null, "", expect.stringContaining("model=abc-1")); + spy.mockRestore(); + }); + + it("openTeam sets ?team= and drops any model param", () => { + window.history.pushState(null, "", "/models-and-endpoints/?model=abc-1"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useModelDetailRouting()); + act(() => result.current.openTeam("team-9")); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).toContain("team=team-9"); + expect(url).not.toContain("model="); + spy.mockRestore(); + }); + + it("close removes both model and team params", () => { + window.history.pushState(null, "", "/models-and-endpoints/?model=abc-1"); + const spy = vi.spyOn(window.history, "pushState"); + const { result } = renderHook(() => useModelDetailRouting()); + act(() => result.current.close()); + const url = spy.mock.calls.at(-1)?.[2] as string; + expect(url).not.toContain("model="); + expect(url).not.toContain("team="); + spy.mockRestore(); + }); + + it("reads modelId and teamId from the query string", () => { + window.history.pushState(null, "", "/models-and-endpoints/?model=xyz"); + const { result } = renderHook(() => useModelDetailRouting()); + expect(result.current.modelId).toBe("xyz"); + expect(result.current.teamId).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/detailNavigation.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/detailNavigation.ts new file mode 100644 index 00000000000..5e120e42740 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/detailNavigation.ts @@ -0,0 +1,51 @@ +import { useSearchParams } from "next/navigation"; +import { useCallback } from "react"; + +export interface ModelDetailRouting { + modelId: string | null; + teamId: string | null; + openModel: (id: string) => void; + openTeam: (id: string) => void; + close: () => void; +} + +function navigateWithParams(mutate: (params: URLSearchParams) => void): void { + const params = new URLSearchParams(window.location.search); + mutate(params); + const qs = params.toString(); + const url = qs ? `${window.location.pathname}?${qs}` : window.location.pathname; + window.history.pushState(null, "", url); +} + +export function useModelDetailRouting(): ModelDetailRouting { + const searchParams = useSearchParams(); + + const openModel = useCallback((id: string) => { + navigateWithParams((params) => { + params.delete("team"); + params.set("model", id); + }); + }, []); + + const openTeam = useCallback((id: string) => { + navigateWithParams((params) => { + params.delete("model"); + params.set("team", id); + }); + }, []); + + const close = useCallback(() => { + navigateWithParams((params) => { + params.delete("model"); + params.delete("team"); + }); + }, []); + + return { + modelId: searchParams?.get("model") ?? null, + teamId: searchParams?.get("team") ?? null, + openModel, + openTeam, + close, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/health/page.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/health/page.test.tsx new file mode 100644 index 00000000000..677796957cc --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/health/page.test.tsx @@ -0,0 +1,54 @@ +/* @vitest-environment jsdom */ +import { render } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import HealthStatusPage from "./page"; + +vi.mock("next/navigation", () => ({ + usePathname: () => "/models-and-endpoints/health", + useRouter: () => ({ push: vi.fn(), replace: vi.fn() }), + useSearchParams: () => new URLSearchParams(""), +})); + +const mockHealthCheckComponent = vi.fn((_props: { all_models_on_proxy?: string[] }) => null); +vi.mock("@/components/model_dashboard/HealthCheckComponent", () => ({ + default: (props: { all_models_on_proxy?: string[] }) => { + mockHealthCheckComponent(props); + return null; + }, +})); + +vi.mock("@/app/(dashboard)/models-and-endpoints/utils/modelDataTransformer", () => ({ + transformModelData: () => ({ data: [] }), +})); + +const mockUseModelsInfo = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/models/useModels", () => ({ useModelsInfo: () => mockUseModelsInfo() })); +vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({ useModelCostMap: () => ({ data: {} }) })); +vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ useTeams: () => ({ data: [] }) })); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => ({ accessToken: "123" }) })); + +describe("HealthStatusPage", () => { + beforeEach(() => { + mockHealthCheckComponent.mockClear(); + }); + + it("passes deployment ids (not model names) to HealthCheckComponent as all_models_on_proxy", () => { + mockUseModelsInfo.mockReturnValue({ + data: { + data: [ + { model_name: "gpt-4", model_info: { id: "deployment-id-1" } }, + { model_name: "gpt-4", model_info: { id: "deployment-id-2" } }, + ], + total_count: 2, + }, + isLoading: false, + }); + + render(); + + expect(mockHealthCheckComponent).toHaveBeenCalled(); + const props = mockHealthCheckComponent.mock.calls[0][0]; + expect(props.all_models_on_proxy).toEqual(["deployment-id-1", "deployment-id-2"]); + expect(props.all_models_on_proxy).not.toContain("gpt-4"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/health/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/health/page.tsx new file mode 100644 index 00000000000..8942db5f002 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/health/page.tsx @@ -0,0 +1,63 @@ +"use client"; + +import { useCallback, useMemo, useState } from "react"; +import type { PaginationState } from "@tanstack/react-table"; +import HealthCheckComponent from "@/components/model_dashboard/HealthCheckComponent"; +import { getDisplayModelName } from "@/components/view_model/model_name_display"; +import { useModelsInfo } from "@/app/(dashboard)/hooks/models/useModels"; +import { useModelCostMap } from "@/app/(dashboard)/hooks/models/useModelCostMap"; +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { transformModelData } from "@/app/(dashboard)/models-and-endpoints/utils/modelDataTransformer"; +import { useModelDetailRouting } from "@/app/(dashboard)/models-and-endpoints/detailNavigation"; + +const HEALTH_PAGE_SIZE = 50; + +export default function HealthStatusPage() { + const { accessToken } = useAuthorized(); + const { data: teams } = useTeams(); + const { data: modelCostMapData } = useModelCostMap(); + const { openModel } = useModelDetailRouting(); + const [pagination, setPagination] = useState({ pageIndex: 0, pageSize: HEALTH_PAGE_SIZE }); + const { data: healthModelDataResponse, isLoading } = useModelsInfo(pagination.pageIndex + 1, pagination.pageSize); + + const getProviderFromModel = useCallback( + (model: string) => { + if (modelCostMapData && typeof modelCostMapData === "object" && model in modelCostMapData) { + return modelCostMapData[model]["litellm_provider"]; + } + return "openai"; + }, + [modelCostMapData], + ); + + const processedHealthModelData = useMemo(() => { + if (!healthModelDataResponse?.data) { + return { data: [] }; + } + return transformModelData(healthModelDataResponse, getProviderFromModel); + }, [healthModelDataResponse, getProviderFromModel]); + + const healthModelIdsOnProxy = useMemo( + () => + healthModelDataResponse?.data + ?.map((model: any) => model.model_info?.id) + .filter((id: string | undefined): id is string => Boolean(id)) ?? [], + [healthModelDataResponse?.data], + ); + + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/layout.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/layout.test.tsx new file mode 100644 index 00000000000..d47d8b9a04e --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/layout.test.tsx @@ -0,0 +1,126 @@ +/* @vitest-environment jsdom */ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { act, render } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import ModelsAndEndpointsLayout from "./layout"; + +const { mockPush, mockReplace, navState } = vi.hoisted(() => ({ + mockPush: vi.fn(), + mockReplace: vi.fn(), + navState: { pathname: "/models-and-endpoints", search: "" }, +})); +vi.mock("next/navigation", () => ({ + usePathname: () => navState.pathname, + useRouter: () => ({ push: mockPush, replace: mockReplace }), + useSearchParams: () => new URLSearchParams(navState.search), +})); + +vi.mock("@/components/networking", () => ({ serverRootPath: "" })); + +vi.mock("@/components/molecules/cost_optimization_feedback_banner", () => ({ default: () => null })); +vi.mock("@/components/model_info_view", () => ({ + default: ({ modelId }: { modelId: string }) =>
model:{modelId}
, +})); +vi.mock("@/components/team/TeamInfo", () => ({ + default: ({ teamId }: { teamId: string }) =>
team:{teamId}
, +})); + +const mockUseAuthorized = vi.fn(); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: () => mockUseAuthorized() })); +vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ useTeams: () => ({ data: [] }) })); +vi.mock("@/app/(dashboard)/hooks/uiSettings/useUISettings", () => ({ + useUISettings: () => ({ data: { values: {} } }), +})); +vi.mock("@/app/(dashboard)/models-and-endpoints/useModelDashboardData", () => ({ + useModelDashboardData: () => ({ + availableModelGroups: [], + availableModelAccessGroups: [], + allModelsOnProxy: [], + isLoading: false, + }), +})); + +const renderLayout = () => { + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false, gcTime: 0 } } }); + return render( + + +
CHILD
+
+
, + ); +}; + +describe("ModelsAndEndpointsLayout", () => { + beforeEach(() => { + navState.pathname = "/models-and-endpoints"; + navState.search = ""; + mockPush.mockClear(); + mockReplace.mockClear(); + mockUseAuthorized.mockReturnValue({ + accessToken: "123", + token: "123", + userRole: "Admin", + userId: "123", + premiumUser: false, + }); + // eslint-disable-next-line @typescript-eslint/no-explicit-any + (global as any).ResizeObserver = class { + observe() {} + unobserve() {} + disconnect() {} + }; + }); + + it("renders the admin tab bar and the active tab's page content", () => { + const { getByRole, getByTestId } = renderLayout(); + expect(getByRole("tab", { name: "LLM Credentials" })).toBeInTheDocument(); + expect(getByRole("tab", { name: "Health Status" })).toBeInTheDocument(); + expect(getByTestId("tab-content")).toHaveTextContent("CHILD"); + }); + + it("navigates to a tab's path when its tab is clicked", async () => { + const { getByRole } = renderLayout(); + await act(async () => { + getByRole("tab", { name: "Health Status" }).click(); + }); + expect(mockPush).toHaveBeenCalledWith(expect.stringMatching(/\/models-and-endpoints\/health\/$/)); + }); + + it("redirects to the base models path when the tab path is not permitted for the role", async () => { + const replaceMock = vi.fn(); + const originalLocation = window.location; + Object.defineProperty(window, "location", { + configurable: true, + value: { replace: replaceMock, assign: vi.fn(), href: "http://localhost/", pathname: "/", search: "" }, + }); + mockUseAuthorized.mockReturnValue({ + accessToken: "123", + token: "123", + userRole: "Internal User", + userId: "123", + premiumUser: false, + }); + navState.pathname = "/models-and-endpoints/llm-credentials"; + await act(async () => { + renderLayout(); + }); + expect(replaceMock).toHaveBeenCalledWith(expect.stringMatching(/\/models-and-endpoints\/$/)); + Object.defineProperty(window, "location", { configurable: true, value: originalLocation }); + }); + + it("renders the model detail overlay from ?model and hides the tabs and page content", () => { + navState.search = "model=abc-123"; + const { getByTestId, queryByTestId, queryByRole } = renderLayout(); + expect(getByTestId("model-info")).toHaveTextContent("model:abc-123"); + expect(queryByTestId("tab-content")).toBeNull(); + expect(queryByRole("tab", { name: "Health Status" })).toBeNull(); + }); + + it("renders the team detail overlay from ?team", () => { + navState.search = "team=team-9"; + const { getByTestId, queryByTestId } = renderLayout(); + expect(getByTestId("team-info")).toHaveTextContent("team:team-9"); + expect(queryByTestId("tab-content")).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/layout.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/layout.tsx new file mode 100644 index 00000000000..1aea5330c5a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/layout.tsx @@ -0,0 +1,162 @@ +"use client"; + +import type { ReactNode } from "react"; +import { useEffect, useMemo, useState } from "react"; +import { usePathname, useRouter } from "next/navigation"; +import { Tabs } from "antd"; +import { RefreshIcon } from "@heroicons/react/outline"; +import { useQueryClient } from "@tanstack/react-query"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { useUISettings } from "@/app/(dashboard)/hooks/uiSettings/useUISettings"; +import { all_admin_roles, internalUserRoles, isProxyAdminRole, isUserTeamAdminForAnyTeam } from "@/utils/roles"; +import CostOptimizationFeedbackBanner from "@/components/molecules/cost_optimization_feedback_banner"; +import ModelInfoView from "@/components/model_info_view"; +import TeamInfoView from "@/components/team/TeamInfo"; +import { modelTabHref, slugFromPathname, type ModelTabSlug } from "@/app/(dashboard)/models-and-endpoints/tabRoutes"; +import { useModelDetailRouting } from "@/app/(dashboard)/models-and-endpoints/detailNavigation"; +import { useModelDashboardData } from "@/app/(dashboard)/models-and-endpoints/useModelDashboardData"; + +const BASE_TAB_KEY = "all-models"; + +const TAB_LABELS: Record = { + add: "Add Model", + "llm-credentials": "LLM Credentials", + "pass-through": "Pass-Through Endpoints", + health: "Health Status", + "retry-settings": "Model Retry Settings", + "model-group-alias": "Model Group Alias", + "price-data": "Price Data Reload", +}; + +export default function ModelsAndEndpointsLayout({ children }: { children: ReactNode }) { + const { accessToken, userRole, userId: userID, premiumUser } = useAuthorized(); + const { data: teams, isLoading: teamsLoading } = useTeams(); + const { data: uiSettings, isLoading: uiSettingsLoading } = useUISettings(); + const pathname = usePathname(); + const router = useRouter(); + const queryClient = useQueryClient(); + const { modelId, teamId, close } = useModelDetailRouting(); + const { availableModelAccessGroups, allModelsOnProxy } = useModelDashboardData(); + + const [lastRefreshed, setLastRefreshed] = useState(""); + + const isProxyAdmin = userRole && isProxyAdminRole(userRole); + const isInternalUser = userRole && internalUserRoles.includes(userRole); + const isUserTeamAdmin = userID && isUserTeamAdminForAnyTeam(teams ?? null, userID); + const addModelDisabledForInternalUsers = + isInternalUser && uiSettings?.values?.disable_model_add_for_internal_users === true; + const shouldHideAddModelTab = !isProxyAdmin && (addModelDisabledForInternalUsers || !isUserTeamAdmin); + const isAdmin = all_admin_roles.includes(userRole); + + const visibleSlugs = useMemo>( + () => [ + "", + ...(shouldHideAddModelTab ? [] : (["add"] as const)), + ...(isAdmin + ? (["llm-credentials", "pass-through", "health", "retry-settings", "model-group-alias", "price-data"] as const) + : []), + ], + [shouldHideAddModelTab, isAdmin], + ); + + const activeSlug = slugFromPathname(pathname); + const isKnownSlug = visibleSlugs.some((slug) => slug === activeSlug); + const activeKey = isKnownSlug ? activeSlug || BASE_TAB_KEY : BASE_TAB_KEY; + + useEffect(() => { + if (teamsLoading || uiSettingsLoading) { + return; + } + if (activeSlug !== "" && !isKnownSlug) { + window.location.replace(modelTabHref("")); + } + }, [activeSlug, isKnownSlug, teamsLoading, uiSettingsLoading]); + + const allModelsLabel = isAdmin ? "All Models" : "Your Models"; + const tabItems = visibleSlugs.map((slug) => { + const key = slug || BASE_TAB_KEY; + return { + key, + label: slug ? TAB_LABELS[slug] : allModelsLabel, + children: key === activeKey ? children : null, + }; + }); + + const handleRefreshClick = () => { + setLastRefreshed(new Date().toLocaleTimeString([], { hour: "2-digit", minute: "2-digit" })); + queryClient.invalidateQueries({ queryKey: ["models", "list"] }); + }; + + const invalidateModels = () => queryClient.invalidateQueries({ queryKey: ["models", "list"] }); + + if (teamId) { + return ( +
+ +
+ ); + } + + return ( +
+
+
+
+

Model Management

+ {isAdmin ? ( +

Add and manage models for the proxy

+ ) : ( +

Add models for teams you are an admin for.

+ )} +
+
+ + + + {modelId ? ( + + ) : ( + router.push(modelTabHref(key === BASE_TAB_KEY ? "" : key))} + items={tabItems} + tabBarExtraContent={{ + right: ( +
+ {lastRefreshed && Last Refreshed: {lastRefreshed}} + +
+ ), + }} + /> + )} +
+
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/llm-credentials/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/llm-credentials/page.tsx new file mode 100644 index 00000000000..207ced5be0d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/llm-credentials/page.tsx @@ -0,0 +1,10 @@ +"use client"; + +import { Form } from "antd"; +import CredentialsPanel from "@/components/model_add/CredentialsPanel"; +import { vertexCredentialsUploadProps } from "@/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload"; + +export default function LlmCredentialsPage() { + const [form] = Form.useForm(); + return ; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/model-group-alias/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/model-group-alias/page.tsx new file mode 100644 index 00000000000..c06de353ddf --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/model-group-alias/page.tsx @@ -0,0 +1,39 @@ +"use client"; + +import { useEffect, useState } from "react"; +import ModelGroupAliasSettings from "@/components/model_group_alias_settings"; +import { getCallbacksCall } from "@/components/networking"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; + +export default function ModelGroupAliasPage() { + const { accessToken, userId: userID, userRole } = useAuthorized(); + const [modelGroupAlias, setModelGroupAlias] = useState<{ [key: string]: string }>({}); + + useEffect(() => { + if (!accessToken || !userID || !userRole) { + return; + } + let active = true; + void (async () => { + try { + const info = await getCallbacksCall(accessToken, userID, userRole); + if (active) { + setModelGroupAlias(info.router_settings?.model_group_alias || {}); + } + } catch (error) { + console.error("Error fetching model group alias:", error); + } + })(); + return () => { + active = false; + }; + }, [accessToken, userID, userRole]); + + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx index 7594ee2f492..546309cfcc1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/page.tsx @@ -1,11 +1,23 @@ "use client"; -import ModelsAndEndpointsView from "@/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { useState } from "react"; +import AllModelsTab from "@/app/(dashboard)/models-and-endpoints/components/AllModelsTab"; +import { useModelDashboardData } from "@/app/(dashboard)/models-and-endpoints/useModelDashboardData"; +import { useModelDetailRouting } from "@/app/(dashboard)/models-and-endpoints/detailNavigation"; -export default function ModelsAndEndpointsPage() { - const { premiumUser } = useAuthorized(); - const { data: teams } = useTeams(); - return ; +export default function AllModelsPage() { + const [selectedModelGroup, setSelectedModelGroup] = useState(null); + const { availableModelGroups, availableModelAccessGroups } = useModelDashboardData(); + const { openModel, openTeam } = useModelDetailRouting(); + + return ( + + ); } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/pass-through/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/pass-through/page.tsx new file mode 100644 index 00000000000..4ba7b8b260b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/pass-through/page.tsx @@ -0,0 +1,11 @@ +"use client"; + +import PassThroughSettings from "@/components/PassThroughSettings/PassThroughSettings"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; + +export default function PassThroughPage() { + const { accessToken, userRole, userId: userID, premiumUser } = useAuthorized(); + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/price-data/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/price-data/page.tsx new file mode 100644 index 00000000000..b8f385be13f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/price-data/page.tsx @@ -0,0 +1,7 @@ +"use client"; + +import PriceDataManagementTab from "@/app/(dashboard)/models-and-endpoints/components/PriceDataManagementTab"; + +export default function PriceDataPage() { + return ; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/retry-settings/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/retry-settings/page.tsx new file mode 100644 index 00000000000..6442be3e54d --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/retry-settings/page.tsx @@ -0,0 +1,100 @@ +"use client"; + +import { useCallback, useEffect, useState } from "react"; +import ModelRetrySettingsTab from "@/app/(dashboard)/models-and-endpoints/components/ModelRetrySettingsTab"; +import { getCallbacksCall } from "@/components/networking"; +import { useUpdateRetryPolicy } from "@/app/(dashboard)/hooks/routerSettings/useUpdateRetryPolicy"; +import { useModelDashboardData } from "@/app/(dashboard)/models-and-endpoints/useModelDashboardData"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import NotificationsManager from "@/components/molecules/notifications_manager"; + +interface RetryPolicyObject { + [key: string]: { [retryPolicyKey: string]: number } | undefined; +} + +interface GlobalRetryPolicyObject { + [retryPolicyKey: string]: number; +} + +interface RouterSettings { + model_group_retry_policy?: RetryPolicyObject | null; + retry_policy?: GlobalRetryPolicyObject | null; + num_retries?: number | null; +} + +export default function ModelRetrySettingsPage() { + const { accessToken, userId: userID, userRole } = useAuthorized(); + const { availableModelGroups } = useModelDashboardData(); + const updateRetryPolicy = useUpdateRetryPolicy(accessToken); + + const [retryScope, setRetryScope] = useState("global"); + const [modelGroupRetryPolicy, setModelGroupRetryPolicy] = useState(null); + const [globalRetryPolicy, setGlobalRetryPolicy] = useState(null); + const [defaultRetry, setDefaultRetry] = useState(0); + + const fetchRetrySettings = useCallback(async () => { + if (!accessToken || !userID || !userRole) { + return null; + } + try { + const info = await getCallbacksCall(accessToken, userID, userRole); + return info.router_settings; + } catch (error) { + console.error("Error fetching router settings:", error); + return null; + } + }, [accessToken, userID, userRole]); + + const applyRetrySettings = useCallback((routerSettings: RouterSettings) => { + setModelGroupRetryPolicy(routerSettings.model_group_retry_policy ?? null); + setGlobalRetryPolicy(routerSettings.retry_policy ?? null); + setDefaultRetry(routerSettings.num_retries ?? 2); + }, []); + + useEffect(() => { + let active = true; + void (async () => { + const routerSettings = await fetchRetrySettings(); + if (active && routerSettings) { + applyRetrySettings(routerSettings); + } + })(); + return () => { + active = false; + }; + }, [fetchRetrySettings, applyRetrySettings]); + + const handleSaveRetrySettings = () => { + updateRetryPolicy.mutate( + { retry_policy: globalRetryPolicy, model_group_retry_policy: modelGroupRetryPolicy }, + { + onSuccess: () => { + NotificationsManager.success("Retry settings saved successfully"); + void fetchRetrySettings().then((routerSettings) => { + if (routerSettings) { + applyRetrySettings(routerSettings); + } + }); + }, + onError: () => { + NotificationsManager.fromBackend("Failed to save retry settings"); + }, + }, + ); + }; + + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/tabRoutes.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/tabRoutes.test.ts new file mode 100644 index 00000000000..920bd3dc156 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/tabRoutes.test.ts @@ -0,0 +1,38 @@ +/* @vitest-environment jsdom */ +import { describe, expect, it, vi } from "vitest"; + +vi.mock("@/components/networking", () => ({ serverRootPath: "" })); + +import { MODEL_TAB_SLUGS, modelTabHref, slugFromPathname } from "./tabRoutes"; + +describe("slugFromPathname", () => { + it("returns empty string for the base path with or without a trailing slash", () => { + expect(slugFromPathname("/models-and-endpoints")).toBe(""); + expect(slugFromPathname("/models-and-endpoints/")).toBe(""); + }); + + it("extracts the tab slug from dev and proxy-mounted (/ui) paths", () => { + expect(slugFromPathname("/models-and-endpoints/add")).toBe("add"); + expect(slugFromPathname("/ui/models-and-endpoints/llm-credentials/")).toBe("llm-credentials"); + }); + + it("returns the raw segment for an unknown tab so the view can redirect to base", () => { + expect(slugFromPathname("/ui/models-and-endpoints/bogus")).toBe("bogus"); + }); + + it("returns empty string when the models base segment is not in the path", () => { + expect(slugFromPathname("/teams")).toBe(""); + }); +}); + +describe("modelTabHref", () => { + it("builds the trailing-slash base href for the empty slug", () => { + expect(modelTabHref("")).toBe("/ui/models-and-endpoints/"); + }); + + it("builds a trailing-slash href for every tab slug (required by static export)", () => { + for (const slug of MODEL_TAB_SLUGS) { + expect(modelTabHref(slug)).toBe(`/ui/models-and-endpoints/${slug}/`); + } + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/tabRoutes.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/tabRoutes.ts new file mode 100644 index 00000000000..ddf9546c5c8 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/tabRoutes.ts @@ -0,0 +1,29 @@ +import { migratedHref } from "@/utils/migratedPages"; + +export const MODELS_BASE_SEGMENT = "models-and-endpoints"; + +export const MODEL_TAB_SLUGS = [ + "add", + "llm-credentials", + "pass-through", + "health", + "retry-settings", + "model-group-alias", + "price-data", +] as const; + +export type ModelTabSlug = (typeof MODEL_TAB_SLUGS)[number]; + +export function modelTabHref(slug: string): string { + const base = migratedHref(MODELS_BASE_SEGMENT); + return slug ? `${base}/${slug}/` : `${base}/`; +} + +export function slugFromPathname(pathname: string): string { + const parts = pathname.split("/").filter(Boolean); + const idx = parts.indexOf(MODELS_BASE_SEGMENT); + if (idx === -1) { + return ""; + } + return parts[idx + 1] ?? ""; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/useModelDashboardData.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/useModelDashboardData.ts new file mode 100644 index 00000000000..c793e41bfd1 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/useModelDashboardData.ts @@ -0,0 +1,32 @@ +import { useMemo } from "react"; +import { useModelsInfo } from "@/app/(dashboard)/hooks/models/useModels"; + +export interface ModelDashboardData { + availableModelGroups: string[]; + availableModelAccessGroups: string[]; + allModelsOnProxy: string[]; + isLoading: boolean; +} + +export function useModelDashboardData(): ModelDashboardData { + const { data: modelDataResponse, isLoading } = useModelsInfo(); + + const availableModelGroups = useMemo(() => { + const groups = new Set(modelDataResponse?.data?.map((model) => model.model_name) ?? []); + return Array.from(groups).sort(); + }, [modelDataResponse?.data]); + + const availableModelAccessGroups = useMemo(() => { + const groups = new Set( + modelDataResponse?.data?.flatMap((model) => model.model_info?.access_groups ?? []) ?? [], + ); + return Array.from(groups); + }, [modelDataResponse?.data]); + + const allModelsOnProxy = useMemo( + () => modelDataResponse?.data?.map((model) => model.model_name) ?? [], + [modelDataResponse?.data], + ); + + return { availableModelGroups, availableModelAccessGroups, allModelsOnProxy, isLoading }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts new file mode 100644 index 00000000000..61bfbd7a99f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/vertexCredentialsUpload.ts @@ -0,0 +1,29 @@ +import type { FormInstance, UploadProps } from "antd"; +import NotificationsManager from "@/components/molecules/notifications_manager"; + +export function vertexCredentialsUploadProps(form: FormInstance): UploadProps { + return { + name: "file", + accept: ".json", + pastable: false, + beforeUpload: (file) => { + if (file.type === "application/json") { + const reader = new FileReader(); + reader.onload = (event) => { + if (event.target) { + form.setFieldsValue({ vertex_credentials: event.target.result as string }); + } + }; + reader.readAsText(file); + } + return false; + }, + onChange(info) { + if (info.file.status === "done") { + NotificationsManager.success(`${info.file.name} file uploaded successfully`); + } else if (info.file.status === "error") { + NotificationsManager.fromBackend(`${info.file.name} file upload failed.`); + } + }, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx new file mode 100644 index 00000000000..e3db50b7300 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.test.tsx @@ -0,0 +1,202 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../../../tests/test-utils"; +import UsagePage from "./usage"; + +const networking = vi.hoisted(() => ({ + adminSpendLogsCall: vi.fn(), + adminTopKeysCall: vi.fn(), + adminTopModelsCall: vi.fn(), + adminTopEndUsersCall: vi.fn(), + teamSpendLogsCall: vi.fn(), + tagsSpendLogsCall: vi.fn(), + allTagNamesCall: vi.fn(), + adminspendByProvider: vi.fn(), + adminGlobalActivity: vi.fn(), + adminGlobalActivityPerModel: vi.fn(), + getProxyUISettings: vi.fn(), + modelAvailableCall: vi.fn(), + keyInfoV1Call: vi.fn(), +})); + +vi.mock("@/components/networking", () => networking); +vi.mock("../../../../components/networking", () => networking); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => ({ + accessToken: "sk-test", + token: "tok", + userRole: "Admin", + userId: "u1", + premiumUser: true, + }), +})); + +const UNLIMITED_SETTINGS = { DISABLE_EXPENSIVE_DB_QUERIES: false, NUM_SPEND_LOGS_ROWS: 10 }; + +const renderUsage = (overrides: Partial> = {}) => + renderWithProviders( + , + ); + +beforeEach(() => { + vi.clearAllMocks(); + networking.getProxyUISettings.mockResolvedValue(UNLIMITED_SETTINGS); + networking.adminSpendLogsCall.mockResolvedValue([{ date: "2026-07-01", spend: 12.5 }]); + networking.adminTopKeysCall.mockResolvedValue([ + { api_key: "sk-abcdefghijk", key_alias: "prod-key", total_spend: 9.5 }, + ]); + networking.adminTopModelsCall.mockResolvedValue([{ model: "gpt-5.1", total_spend: 7.25 }]); + networking.adminTopEndUsersCall.mockResolvedValue([ + { end_user: "customer-alpha", total_spend: 3.5, total_count: 42 }, + ]); + networking.teamSpendLogsCall.mockResolvedValue({ + daily_spend: [{ date: "2026-07-01", "team-a": 5 }], + teams: ["team-a"], + total_spend_per_team: [{ team_id: "team-a", total_spend: 5 }], + }); + networking.tagsSpendLogsCall.mockResolvedValue({ spend_per_tag: [{ name: "prod", spend: 4 }] }); + networking.allTagNamesCall.mockResolvedValue({ tag_names: ["prod", "staging"] }); + networking.adminspendByProvider.mockResolvedValue([{ provider: "openai", spend: 6.75 }]); + networking.adminGlobalActivity.mockResolvedValue({ + sum_api_requests: 120, + sum_total_tokens: 4500, + daily_data: [{ date: "2026-07-01", api_requests: 120, total_tokens: 4500 }], + }); + networking.adminGlobalActivityPerModel.mockResolvedValue([]); + networking.modelAvailableCall.mockResolvedValue({ data: [] }); + networking.keyInfoV1Call.mockResolvedValue({ info: {} }); +}); + +describe("old usage page", () => { + describe("when the proxy has disabled expensive DB queries", () => { + beforeEach(() => { + networking.getProxyUISettings.mockResolvedValue({ + DISABLE_EXPENSIVE_DB_QUERIES: true, + NUM_SPEND_LOGS_ROWS: 2500000, + }); + }); + + it("shows the database query limit warning instead of the usage dashboard", async () => { + renderUsage(); + + expect(await screen.findByText("Database Query Limit Reached")).toBeInTheDocument(); + expect(screen.getByText(/SpendLogs in DB has/)).toHaveTextContent("2500000"); + expect(screen.getByText(/Please follow our guide to view usage when SpendLogs has more than 1M rows/i)); + expect(screen.queryByRole("tab", { name: "All Up" })).not.toBeInTheDocument(); + }); + + it("links to the cost tracking guide in a new tab", async () => { + renderUsage(); + + const link = await screen.findByRole("link", { name: "View Usage Guide" }); + expect(link).toHaveAttribute("href", "https://docs.litellm.ai/docs/proxy/cost_tracking"); + expect(link).toHaveAttribute("target", "_blank"); + }); + + it("skips every expensive usage query", async () => { + renderUsage(); + + await screen.findByText("Database Query Limit Reached"); + await waitFor(() => expect(networking.getProxyUISettings).toHaveBeenCalled()); + + expect(networking.adminSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.adminspendByProvider).not.toHaveBeenCalled(); + expect(networking.adminTopKeysCall).not.toHaveBeenCalled(); + expect(networking.adminTopModelsCall).not.toHaveBeenCalled(); + expect(networking.adminGlobalActivity).not.toHaveBeenCalled(); + expect(networking.adminGlobalActivityPerModel).not.toHaveBeenCalled(); + expect(networking.teamSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled(); + expect(networking.tagsSpendLogsCall).not.toHaveBeenCalled(); + }); + }); + + describe("as an admin", () => { + it("renders the admin tabs", async () => { + renderUsage(); + + expect(await screen.findByRole("tab", { name: "All Up" })).toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Team Based Usage" })).toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Customer Usage" })).toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Tag Based Usage" })).toBeInTheDocument(); + }); + + it("renders the cost panel cards", async () => { + renderUsage(); + + expect(await screen.findByText("Monthly Spend")).toBeInTheDocument(); + expect(screen.getByText("Top Virtual Keys")).toBeInTheDocument(); + expect(screen.getByText("Top Models")).toBeInTheDocument(); + expect(screen.getByText("Spend by Provider")).toBeInTheDocument(); + }); + + it("lists spend by provider in a table", async () => { + renderUsage(); + + const providerCell = await screen.findByText("openai"); + const row = providerCell.closest("tr"); + expect(row).not.toBeNull(); + expect(within(row as HTMLElement).getByText("$6.75")).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: "Provider" })).toBeInTheDocument(); + }); + + it("shows the customer usage table when its tab is selected", async () => { + const user = userEvent.setup(); + renderUsage(); + + await user.click(await screen.findByRole("tab", { name: "Customer Usage" })); + + const customerCell = await screen.findByText("customer-alpha"); + const row = customerCell.closest("tr"); + expect(row).not.toBeNull(); + expect(within(row as HTMLElement).getByText("$3.50")).toBeInTheDocument(); + expect(within(row as HTMLElement).getByText("42")).toBeInTheDocument(); + expect(screen.getByRole("columnheader", { name: "Total Events" })).toBeInTheDocument(); + }); + + it("shows the tag spend panel when its tab is selected", async () => { + const user = userEvent.setup(); + renderUsage(); + + await user.click(await screen.findByRole("tab", { name: "Tag Based Usage" })); + + expect(await screen.findByText("Spend Per Tag")).toBeInTheDocument(); + }); + + it("shows the team spend panel when its tab is selected", async () => { + const user = userEvent.setup(); + renderUsage(); + + await user.click(await screen.findByRole("tab", { name: "Team Based Usage" })); + + expect(await screen.findByText("Total Spend Per Team")).toBeInTheDocument(); + expect(screen.getByText("Daily Spend Per Team")).toBeInTheDocument(); + }); + }); + + describe("as a non-admin", () => { + it("renders only the All Up tab and skips admin-only queries", async () => { + renderUsage({ userRole: "Internal User" }); + + expect(await screen.findByRole("tab", { name: "All Up" })).toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Team Based Usage" })).not.toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Customer Usage" })).not.toBeInTheDocument(); + expect(screen.queryByRole("tab", { name: "Tag Based Usage" })).not.toBeInTheDocument(); + + await waitFor(() => expect(networking.adminSpendLogsCall).toHaveBeenCalled()); + expect(networking.teamSpendLogsCall).not.toHaveBeenCalled(); + expect(networking.adminTopEndUsersCall).not.toHaveBeenCalled(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx index 01f8cb1cd45..3d55f9bb698 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/old-usage/_components/usage.tsx @@ -1,40 +1,26 @@ -import { - BarChart, - BarList, - Card, - Title, - Table, - TableHead, - TableHeaderCell, - TableRow, - TableCell, - TableBody, - Subtitle, -} from "@tremor/react"; - import React, { useState, useEffect } from "react"; import ViewUserSpend from "@/components/view_user_spend"; import { ProxySettings } from "@/components/user_dashboard"; import UsageDatePicker from "@/components/shared/usage_date_picker"; +import { Button } from "@/components/ui/button"; +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { - Grid, - Col, - Text, - TabPanel, - TabPanels, - TabGroup, - TabList, - Tab, - Select, - SelectItem, - DateRangePickerValue, - DonutChart, - AreaChart, - Button, - MultiSelect, - MultiSelectItem, -} from "@tremor/react"; + Combobox, + ComboboxChip, + ComboboxChips, + ComboboxChipsInput, + ComboboxContent, + ComboboxEmpty, + ComboboxItem, + ComboboxList, + ComboboxValue, +} from "@/components/ui/combobox"; +import { Meter, MeterIndicator, MeterTrack } from "@/components/ui/meter"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; +import { AreaChart, BarChart, DonutChart } from "@/components/shared/charts"; import { adminSpendLogsCall, @@ -68,69 +54,41 @@ interface GlobalActivityData { daily_data: { date: string; api_requests: number; total_tokens: number }[]; } -type CustomTooltipTypeBar = { - payload: any; - active: boolean | undefined; - label: any; -}; +type UsageDateRange = { from?: Date; to?: Date }; -const customTooltip = (props: CustomTooltipTypeBar) => { - const { payload, active } = props; - if (!active || !payload) return null; +type TeamSpendTotal = { name: string; value: number }; - const value = payload[0].payload; - const date = value["startTime"]; - const model_values = value["models"]; - const entries: [string, number][] = Object.entries(model_values).map(([key, value]) => [key, value as number]); +type TagOption = { value: string; label: string; disabled: boolean }; - entries.sort((a, b) => b[1] - a[1]); - const topEntries = entries.slice(0, 5); - - return ( -
- {date} - {topEntries.map(([key, value]) => ( -
-
-

- {key} - {":"} - - {" "} - {value ? `$${formatNumberWithCommas(value, 2)}` : ""} - -

-
-
- ))} -
- ); -}; - -function getTopKeys(data: Array<{ [key: string]: unknown }>): any[] { - const spendKeys: { key: string; spend: unknown }[] = []; - - data.forEach((dict) => { - Object.entries(dict).forEach(([key, value]) => { - if (key !== "spend" && key !== "startTime" && key !== "models" && key !== "users") { - spendKeys.push({ key, spend: value }); - } - }); - }); - - spendKeys.sort((a, b) => Number(b.spend) - Number(a.spend)); - - const topKeys = spendKeys.slice(0, 5).map((k) => k.key); - return topKeys; -} -type DataDict = { [key: string]: unknown }; -type UserData = { user_id: string; spend: number }; +const ALL_TAGS = "all-tags"; const isAdminOrAdminViewer = (role: string | null): boolean => { if (role === null) return false; return role === "Admin" || role === "Admin Viewer"; }; +const TeamSpendBarList: React.FC<{ data: TeamSpendTotal[] }> = ({ data }) => { + const max = Math.max(0, ...data.map((team) => team.value)); + + return ( +
+ {data.map((team) => ( +
+

{team.name}

+ + + + + +

+ {formatNumberWithCommas(team.value, 2)} +

+
+ ))} +
+ ); +}; + const UsagePage: React.FC = ({ accessToken, token, userRole, userID, keys, premiumUser }) => { const currentDate = new Date(); const [keySpendData, setKeySpendData] = useState([]); @@ -141,13 +99,13 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use const [topTagsData, setTopTagsData] = useState([]); const [allTagNames, setAllTagNames] = useState([]); const [uniqueTeamIds, setUniqueTeamIds] = useState([]); - const [totalSpendPerTeam, setTotalSpendPerTeam] = useState([]); + const [totalSpendPerTeam, setTotalSpendPerTeam] = useState([]); const [spendByProvider, setSpendByProvider] = useState([]); const [globalActivity, setGlobalActivity] = useState({} as GlobalActivityData); const [globalActivityPerModel, setGlobalActivityPerModel] = useState([]); - const [selectedKeyID, setSelectedKeyID] = useState(""); - const [selectedTags, setSelectedTags] = useState(["all-tags"]); - const [dateValue, setDateValue] = useState({ + const [selectedKeyToken, setSelectedKeyToken] = useState(null); + const [selectedTags, setSelectedTags] = useState([ALL_TAGS]); + const [dateValue, setDateValue] = useState({ from: new Date(Date.now() - 7 * 24 * 60 * 60 * 1000), to: new Date(), }); @@ -160,6 +118,21 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use let startTime = formatDate(firstDay); let endTime = formatDate(lastDay); + const selectableKeys: { token: string; alias: string }[] = (keys ?? []) + .filter((key: any) => key && typeof key["key_alias"] === "string" && key["key_alias"].length > 0) + .map((key: any) => ({ token: String(key["token"]), alias: String(key["key_alias"]) })); + + const tagOptions: TagOption[] = [ + { value: ALL_TAGS, label: "All Tags", disabled: false }, + ...allTagNames + .filter((tag) => tag !== ALL_TAGS) + .map((tag) => ({ + value: tag, + label: premiumUser ? tag : `✨ ${tag} (Enterprise only Feature)`, + disabled: !premiumUser, + })), + ]; + function valueFormatterNumbers(number: number) { const formatter = new Intl.NumberFormat("en-US", { maximumFractionDigits: 0, @@ -405,7 +378,7 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use setUniqueTeamIds(teamSpend.teams); return teamSpend.total_spend_per_team.map((tspt: any) => ({ name: tspt["team_id"] || "", - value: formatNumberWithCommas(tspt["total_spend"] || 0, 2), + value: Number(tspt["total_spend"] || 0), })); }, setTotalSpendPerTeam, @@ -524,223 +497,252 @@ const UsagePage: React.FC = ({ accessToken, token, userRole, use if (proxySettings?.DISABLE_EXPENSIVE_DB_QUERIES) { return ( -
+
- Database Query Limit Reached - - SpendLogs in DB has {proxySettings.NUM_SPEND_LOGS_ROWS} rows. -

- Please follow our guide to view usage when SpendLogs has more than 1M rows. -
- + + Database Query Limit Reached + + +

+ SpendLogs in DB has {proxySettings.NUM_SPEND_LOGS_ROWS} rows. +

+ Please follow our guide to view usage when SpendLogs has more than 1M rows. +

+
); } return ( -
- - - All Up +
+ + + All Up - {isAdminOrAdminViewer(userRole) ? ( + {isAdminOrAdminViewer(userRole) && ( <> - Team Based Usage - Customer Usage - Tag Based Usage - - ) : ( - <> -
+ Team Based Usage + Customer Usage + Tag Based Usage )} - - - - - - Cost - Activity - - - - -
- - Project Spend {new Date().toLocaleString("default", { month: "long" })} 1 -{" "} - {new Date(new Date().getFullYear(), new Date().getMonth() + 1, 0).getDate()} - - - - - - Monthly Spend - + + + + + Cost + Activity + + + +
+
+

+ Project Spend {new Date().toLocaleString("default", { month: "long" })} 1 -{" "} + {new Date(new Date().getFullYear(), new Date().getMonth() + 1, 0).getDate()} +

+ +
+
+ + + Monthly Spend + + + + + +
+
+ + + Top Virtual Keys + + + {}} /> + + +
+
+ + + Top Models + + + `$${formatNumberWithCommas(value, 2)}`} + /> + + +
+
+
+ + + Spend by Provider + + +
+
+ `$${formatNumberWithCommas(value, 2)}`} + /> +
+
+
+ + + Provider + Spend + + + + {spendByProvider.map((provider) => ( + + {provider.provider} + + + + + ))} + +
+
+
+ + +
+
+ + + +
+ + + All Up + + +
+
+

+ API Requests {valueFormatterNumbers(globalActivity.sum_api_requests)} +

+ - - - - - Top Virtual Keys - {}} /> - - - - - Top Models +
+
+

+ Tokens {valueFormatterNumbers(globalActivity.sum_total_tokens)} +

`$${formatNumberWithCommas(value, 2)}`} + categories={["total_tokens"]} /> - - - - - - Spend by Provider - <> - - - `$${formatNumberWithCommas(value, 2)}`} - /> - - - - - - Provider - Spend - - - - {spendByProvider.map((provider) => ( - - {provider.provider} - - - - - ))} - -
- -
- -
- - - - - - - All Up - - - +
+
+
+
+ + {globalActivityPerModel.map((globalActivity, index) => ( + + + {globalActivity.model} + + +
+
+

API Requests {valueFormatterNumbers(globalActivity.sum_api_requests)} - +

- - - +
+
+

Tokens {valueFormatterNumbers(globalActivity.sum_total_tokens)} - +

- - - +
+
+
+
+ ))} +
+
+ + - <> - {globalActivityPerModel.map((globalActivity, index) => ( - - {globalActivity.model} - - - - API Requests {valueFormatterNumbers(globalActivity.sum_api_requests)} - - - - - - Tokens {valueFormatterNumbers(globalActivity.sum_total_tokens)} - - - - - - ))} - - - - - - - - - - - Total Spend Per Team - - - - Daily Spend Per Team + +
+
+ + + Total Spend Per Team + + + + + + + + Daily Spend Per Team + + = ({ accessToken, token, userRole, use yAxisWidth={80} stack={true} /> - - - - - - -

- Customers of your LLM API calls. Tracked when a `user` param is passed in your LLM calls{" "} - - docs here - -

- - - { - setDateValue(value); - updateEndUserData(value.from, value.to, null); - }} - /> - - - Select Key - - - + + +
+
+
- - - - - Customer - Spend - Total Events - - - - - {topUsers?.map((user: any, index: number) => ( - - {user.end_user} - - - - {user.total_count} - + +

+ Customers of your LLM API calls. Tracked when a `user` param is passed in your LLM calls{" "} + + docs here + +

+
+
+ { + setDateValue(value); + updateEndUserData(value.from, value.to, null); + }} + /> +
+
+

Select Key

+
-
-
- - - - { - setDateValue(value); - updateTagSpendData(value.from, value.to); - }} - /> - + + +
+
- - {premiumUser ? ( -
- setSelectedTags(value as string[])}> - setSelectedTags(["all-tags"])} - > - All Tags - - {allTagNames && - allTagNames - .filter((tag) => tag !== "all-tags") - .map((tag: any, index: number) => { - return ( - - {tag} - - ); - })} - -
- ) : ( -
- setSelectedTags(value as string[])}> - setSelectedTags(["all-tags"])} - > - All Tags - - {allTagNames && - allTagNames - .filter((tag) => tag !== "all-tags") - .map((tag: any, index: number) => { - return ( - - ✨ {tag} (Enterprise only Feature) - - ); - })} - -
- )} - - - - - - Spend Per Tag - + + +
+ + + + Customer + Spend + Total Events + + + + + {topUsers?.map((user: any, index: number) => ( + + {user.end_user} + + + + {user.total_count} + + ))} + +
+
+
+
+ + + +
+
+ { + setDateValue(value); + updateTagSpendData(value.from, value.to); + }} + /> +
+ +
+ selectedTags.includes(option.value))} + onValueChange={(options: TagOption[]) => setSelectedTags(options.map((option) => option.value))} + isItemEqualToValue={(a: TagOption, b: TagOption) => a.value === b.value} + itemToStringLabel={(option: TagOption) => option.label} + > + + + {(options: TagOption[]) => + options.map((option) => ( + + {option.label} + + )) + } + + + + + No tags found + + {(option: TagOption) => ( + + {option.label} + + )} + + + +
+
+
+
+ + + Spend Per Tag + + +

Get Started by Tracking cost per tag{" "} here - - - - - - - - - +

+ +
+
+
+
+
+
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.test.tsx index 814625ff6be..37eeaf4c2af 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.test.tsx @@ -7,8 +7,6 @@ describe("OrganizationFilters", () => { const defaultFilters: FilterState = { org_id: "", org_alias: "", - sort_by: "", - sort_order: "asc", }; it("should render", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.tsx index 5643a4bc51a..6ad2f00fdb0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/OrganizationFilters.tsx @@ -14,8 +14,6 @@ interface OrganizationFiltersProps { type FilterState = { org_id: string; org_alias: string; - sort_by: string; - sort_order: "asc" | "desc"; }; const OrganizationFilters = ({ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx new file mode 100644 index 00000000000..d381e5e65ca --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.test.tsx @@ -0,0 +1,57 @@ +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen } from "@testing-library/react"; +import React from "react"; +import { describe, expect, it, vi } from "vitest"; + +vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ + __esModule: true, + default: () => null, +})); +vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({ + __esModule: true, + default: () => null, +})); +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => ({ + accessToken: null, + userId: null, + userRole: null, + }), +})); +vi.mock("./OrganizationsTable", () => ({ + __esModule: true, + default: (props: { isLoading: boolean }) => ( +
isLoading:{String(props.isLoading)}
+ ), +})); + +import OrganizationsPanel from "./OrganizationsPanel"; + +const renderWithQueryClient = (ui: React.ReactElement) => { + const queryClient = new QueryClient({ + defaultOptions: { queries: { retry: false } }, + }); + return render({ui}); +}; + +describe("OrganizationsPanel", () => { + it("gates non-premium users behind the enterprise notice", () => { + renderWithQueryClient(); + + expect(screen.getByText(/LiteLLM Enterprise feature/i)).toBeInTheDocument(); + expect(screen.queryByText("+ Create New Organization")).not.toBeInTheDocument(); + }); + + it("shows the create button for a premium admin", () => { + renderWithQueryClient(); + + expect(screen.getByText("+ Create New Organization")).toBeInTheDocument(); + }); + + it("resolves the loading skeleton to false when the query is disabled (no token)", () => { + renderWithQueryClient(); + + // A disabled React Query keeps isPending true forever; feeding isLoading avoids a stuck skeleton. + expect(screen.getByTestId("organizations-table")).toHaveTextContent("isLoading:false"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx new file mode 100644 index 00000000000..9f7e029a1d4 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsPanel.tsx @@ -0,0 +1,299 @@ +import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; +import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels"; +import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters"; +import { InfoCircleOutlined } from "@ant-design/icons"; +import { Form, Input, Modal, Select as Select2, Tooltip } from "antd"; +import { useQueryClient } from "@tanstack/react-query"; +import React, { useState } from "react"; +import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; +import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector"; +import { ModelSelect } from "@/components/ModelSelect/ModelSelect"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { organizationCreateCall, organizationDeleteCall } from "@/components/networking"; +import OrganizationInfoView from "@/components/organization/organization_view"; +import NumericalInput from "@/components/shared/numerical_input"; +import { Button } from "@/components/ui/button"; +import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector"; + +import OrganizationsTable from "./OrganizationsTable"; + +interface OrganizationsPanelProps { + userRole: string; + accessToken: string | null; + premiumUser: boolean; +} + +const OrganizationsPanel: React.FC = ({ userRole, accessToken, premiumUser }) => { + const [selectedOrgId, setSelectedOrgId] = useState(null); + const [editOrg, setEditOrg] = useState(false); + const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); + const [orgToDelete, setOrgToDelete] = useState(null); + const [isDeleting, setIsDeleting] = useState(false); + const [isOrgModalVisible, setIsOrgModalVisible] = useState(false); + const [form] = Form.useForm(); + const [showFilters, setShowFilters] = useState(false); + const [filters, setFilters] = useState({ org_id: "", org_alias: "" }); + + const queryClient = useQueryClient(); + const { data: organizations = [], isLoading } = useOrganizations({ + org_id: filters.org_id, + org_alias: filters.org_alias, + }); + const { data: userModels = [] } = useUserModels(); + + const searchActive = Boolean(filters.org_id || filters.org_alias); + + const refetchOrganizations = () => queryClient.invalidateQueries({ queryKey: organizationKeys.lists() }); + + const handleFilterChange = (key: keyof FilterState, value: string) => { + setFilters((previousFilters) => ({ ...previousFilters, [key]: value })); + }; + + const handleFilterReset = () => { + setFilters({ org_id: "", org_alias: "" }); + }; + + const handleDelete = (orgId: string | null) => { + if (!orgId) return; + + setOrgToDelete(orgId); + setIsDeleteModalOpen(true); + }; + + const confirmDelete = async () => { + if (!orgToDelete || !accessToken) return; + + try { + setIsDeleting(true); + await organizationDeleteCall(accessToken, orgToDelete); + NotificationsManager.success("Organization deleted successfully"); + + setIsDeleteModalOpen(false); + setOrgToDelete(null); + await refetchOrganizations(); + } catch (error) { + console.error("Error deleting organization:", error); + } finally { + setIsDeleting(false); + } + }; + + const cancelDelete = () => { + setIsDeleteModalOpen(false); + setOrgToDelete(null); + }; + + const handleCreate = async (values: any) => { + try { + if (!accessToken) return; + + // Transform allowed_vector_store_ids and allowed_mcp_servers_and_groups into object_permission + if ( + (values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) || + (values.allowed_mcp_servers_and_groups && + (values.allowed_mcp_servers_and_groups.servers?.length > 0 || + values.allowed_mcp_servers_and_groups.accessGroups?.length > 0)) + ) { + values.object_permission = {}; + if (values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) { + values.object_permission.vector_stores = values.allowed_vector_store_ids; + delete values.allowed_vector_store_ids; + } + if (values.allowed_mcp_servers_and_groups) { + if (values.allowed_mcp_servers_and_groups.servers?.length > 0) { + values.object_permission.mcp_servers = values.allowed_mcp_servers_and_groups.servers; + } + if (values.allowed_mcp_servers_and_groups.accessGroups?.length > 0) { + values.object_permission.mcp_access_groups = values.allowed_mcp_servers_and_groups.accessGroups; + } + delete values.allowed_mcp_servers_and_groups; + } + } + + await organizationCreateCall(accessToken, values); + NotificationsManager.success("Organization created successfully"); + setIsOrgModalVisible(false); + form.resetFields(); + await refetchOrganizations(); + } catch (error) { + console.error("Error creating organization:", error); + } + }; + + const handleCancel = () => { + setIsOrgModalVisible(false); + form.resetFields(); + }; + + if (!premiumUser) { + return ( +
+

+ This is a LiteLLM Enterprise feature, and requires a valid key to use. Get a trial key{" "} + + here + + . +

+
+ ); + } + + return ( +
+ {(userRole === "Admin" || userRole === "Org Admin") && ( + + )} + + {selectedOrgId ? ( + { + setSelectedOrgId(null); + setEditOrg(false); + }} + accessToken={accessToken} + is_org_admin={true} + is_proxy_admin={userRole === "Admin"} + userModels={userModels} + editOrg={editOrg} + /> + ) : ( + <> +

Click on an organization ID to view its details.

+ + { + setSelectedOrgId(organizationId); + setEditOrg(true); + }} + onDeleteClick={handleDelete} + /> + + )} + + +
+ + + + + form.setFieldValue("models", values)} + context="organization" + /> + + + + + + + + daily + weekly + monthly + + + + + + + + + + + Allowed Vector Stores{" "} + + + + + } + name="allowed_vector_store_ids" + className="mt-4" + help="Select vector stores this organization can access. Leave empty for access to all vector stores" + > + form.setFieldValue("allowed_vector_store_ids", values)} + value={form.getFieldValue("allowed_vector_store_ids")} + accessToken={accessToken || ""} + placeholder="Select vector stores (optional)" + /> + + + + Allowed MCP Servers{" "} + + + + + } + name="allowed_mcp_servers_and_groups" + className="mt-4" + help="Select MCP servers and access groups this organization can access." + > + form.setFieldValue("allowed_mcp_servers_and_groups", values)} + value={form.getFieldValue("allowed_mcp_servers_and_groups")} + accessToken={accessToken || ""} + placeholder="Select MCP servers and access groups (optional)" + /> + + + + + + +
+ +
+
+
+ + +
+ ); +}; + +export default OrganizationsPanel; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx new file mode 100644 index 00000000000..a06c5c885e3 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.test.tsx @@ -0,0 +1,188 @@ +import { render, screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import React from "react"; +import { describe, expect, it, vi } from "vitest"; + +import { Organization } from "@/components/networking"; + +import OrganizationsTable from "./OrganizationsTable"; + +const makeOrganization = (overrides: Partial = {}): Organization => ({ + organization_id: "org-alpha", + organization_alias: "Alpha", + budget_id: "budget-1", + metadata: {}, + models: [], + spend: 0, + model_spend: {}, + created_at: "2023-01-01T00:00:00Z", + created_by: "someone", + updated_at: "2023-01-01T00:00:00Z", + updated_by: "someone", + litellm_budget_table: null, + teams: null, + users: null, + members: null, + ...overrides, +}); + +const baseProps = { + isLoading: false, + userRole: "Admin", + searchActive: false, + onOrganizationClick: vi.fn(), + onEditClick: vi.fn(), + onDeleteClick: vi.fn(), +}; + +describe("OrganizationsTable", () => { + it("renders every column header", () => { + render(); + for (const header of [ + "Organization ID", + "Organization Name", + "Created", + "Spend (USD)", + "Budget (USD)", + "Models", + "TPM / RPM Limits", + "Members", + ]) { + expect(screen.getByText(header)).toBeInTheDocument(); + } + }); + + it("opens the detail view when the organization ID cell is clicked", async () => { + const user = userEvent.setup(); + const onOrganizationClick = vi.fn(); + render( + , + ); + + await user.click(screen.getByText("org-123")); + + expect(onOrganizationClick).toHaveBeenCalledWith("org-123"); + }); + + it("edits and deletes an organization through the ⋯ actions menu (admin)", async () => { + const user = userEvent.setup(); + const onEditClick = vi.fn(); + const onDeleteClick = vi.fn(); + render( + , + ); + + await user.click(screen.getByTestId("organization-actions-org-9")); + await user.click(await screen.findByTestId("organization-action-edit")); + expect(onEditClick).toHaveBeenCalledWith("org-9"); + + await user.click(screen.getByTestId("organization-actions-org-9")); + await user.click(await screen.findByTestId("organization-action-delete")); + expect(onDeleteClick).toHaveBeenCalledWith("org-9"); + }); + + it("hides the row actions menu from non-admins", () => { + render( + , + ); + + expect(screen.queryByTestId("organization-actions-org-9")).not.toBeInTheDocument(); + }); + + it("sorts by created_at descending by default", () => { + render( + , + ); + + const rows = screen.getAllByRole("row"); + // rows[0] is the header row; the newest organization must lead the body. + expect(within(rows[1]).getByText("Newer")).toBeInTheDocument(); + expect(within(rows[2]).getByText("Older")).toBeInTheDocument(); + }); + + it("renders budget, limits, members, and models for a fully-populated organization", () => { + render( + , + ); + + expect(screen.getByText("$100.00")).toBeInTheDocument(); + expect(screen.getByText("TPM: 1000")).toBeInTheDocument(); + expect(screen.getByText("RPM: 60")).toBeInTheDocument(); + expect(screen.getByText("3 Members")).toBeInTheDocument(); + // Five models, three visible -> the shared ModelsCell collapses the rest. + expect(screen.getByText("+2 more")).toBeInTheDocument(); + }); + + it("shows Unlimited budget and All Proxy Models when unset", () => { + render( + , + ); + + expect(screen.getByText("All Proxy Models")).toBeInTheDocument(); + // Budget shows a standalone "Unlimited"; the limits fall back inline. + expect(screen.getByText("Unlimited")).toBeInTheDocument(); + expect(screen.getByText("TPM: Unlimited")).toBeInTheDocument(); + expect(screen.getByText("RPM: Unlimited")).toBeInTheDocument(); + }); + + it("renders loading skeletons instead of rows while loading", () => { + render( + , + ); + + expect(screen.getAllByTestId("skeleton-row").length).toBeGreaterThan(0); + expect(screen.queryByText("ShouldNotShow")).not.toBeInTheDocument(); + }); + + it("uses a search-aware empty state", () => { + const { rerender } = render(); + expect(screen.getByText("No organizations yet")).toBeInTheDocument(); + + rerender(); + expect(screen.getByText("No matching organizations")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.tsx new file mode 100644 index 00000000000..8e68a57d2f7 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTable.tsx @@ -0,0 +1,75 @@ +"use client"; + +import { SortingState } from "@tanstack/react-table"; +import { Building2, SearchX } from "lucide-react"; +import React, { useMemo, useState } from "react"; + +import { DataTable } from "@/components/shared/DataTable"; +import { Organization } from "@/components/networking"; + +import { getOrganizationsTableColumns } from "./OrganizationsTableColumns"; + +interface OrganizationsTableProps { + organizations: Organization[]; + isLoading: boolean; + userRole: string; + searchActive: boolean; + onOrganizationClick: (organizationId: string) => void; + onEditClick: (organizationId: string) => void; + onDeleteClick: (organizationId: string) => void; +} + +const DEFAULT_SORTING: SortingState = [{ id: "created_at", desc: true }]; + +function EmptyState({ searchActive }: { searchActive: boolean }) { + const Icon = searchActive ? SearchX : Building2; + return ( +
+
+ +
+
+ {searchActive ? "No matching organizations" : "No organizations yet"} +
+
+ {searchActive + ? "No organizations match your search. Try a different name or ID." + : "Create an organization to group teams, models, and budgets."} +
+
+ ); +} + +const OrganizationsTable: React.FC = ({ + organizations, + isLoading, + userRole, + searchActive, + onOrganizationClick, + onEditClick, + onDeleteClick, +}) => { + const [sorting, setSorting] = useState(DEFAULT_SORTING); + + const columns = useMemo(() => { + const deps = { userRole, onOrganizationClick, onEditClick, onDeleteClick }; + return getOrganizationsTableColumns(deps); + }, [userRole, onOrganizationClick, onEditClick, onDeleteClick]); + + return ( + organization.organization_id || String(index)} + sortingMode="client" + sorting={sorting} + onSortingChange={setSorting} + isLoading={isLoading} + loadingMessage="Loading organizations…" + noDataMessage={} + size="compact" + /> + ); +}; + +export default OrganizationsTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx new file mode 100644 index 00000000000..31f6a00916c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/OrganizationsTableColumns.tsx @@ -0,0 +1,186 @@ +"use client"; + +import { ColumnDef } from "@tanstack/react-table"; +import { MoreHorizontal, Pencil, Trash2 } from "lucide-react"; + +import { DataTableSortHeader } from "@/components/shared/DataTable"; +import { DateCell, IdentityCell, ModelsCell, MoneyCell } from "@/components/shared/table_cells"; +import { Organization } from "@/components/networking"; +import { buttonVariants } from "@/components/ui/button"; +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu"; +import { cn } from "@/lib/cva.config"; + +interface OrganizationBudget { + max_budget?: number | null; + tpm_limit?: number | null; + rpm_limit?: number | null; +} + +const getOrganizationBudget = (organization: Organization): OrganizationBudget => + (organization.litellm_budget_table ?? {}) as OrganizationBudget; + +function OrganizationLimitsCell({ organization }: { organization: Organization }) { + const { tpm_limit, rpm_limit } = getOrganizationBudget(organization); + return ( +
+ TPM: {tpm_limit ? tpm_limit : "Unlimited"} + RPM: {rpm_limit ? rpm_limit : "Unlimited"} +
+ ); +} + +interface OrganizationRowActionsProps { + organization: Organization; + onEditClick: (organizationId: string) => void; + onDeleteClick: (organizationId: string) => void; +} + +function OrganizationRowActions({ organization, onEditClick, onDeleteClick }: OrganizationRowActionsProps) { + return ( + + + + + + onEditClick(organization.organization_id)} + > + + Edit + + onDeleteClick(organization.organization_id)} + > + + Delete + + + + ); +} + +export interface OrganizationsTableColumnsDeps { + userRole: string; + onOrganizationClick: (organizationId: string) => void; + onEditClick: (organizationId: string) => void; + onDeleteClick: (organizationId: string) => void; +} + +export const getOrganizationsTableColumns = ({ + userRole, + onOrganizationClick, + onEditClick, + onDeleteClick, +}: OrganizationsTableColumnsDeps): ColumnDef[] => [ + { + id: "organization_id", + accessorKey: "organization_id", + meta: { title: "Organization ID" }, + header: ({ column }) => , + size: 220, + enableSorting: true, + cell: ({ row }) => ( + onOrganizationClick(row.original.organization_id)} + /> + ), + }, + { + id: "organization_alias", + accessorKey: "organization_alias", + meta: { title: "Organization Name" }, + header: ({ column }) => , + size: 200, + enableSorting: true, + cell: ({ row }) => { + const alias = row.original.organization_alias; + return ( + + {alias || "-"} + + ); + }, + }, + { + id: "created_at", + accessorKey: "created_at", + sortingFn: "datetime", + meta: { title: "Created" }, + header: ({ column }) => , + size: 130, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "spend", + accessorKey: "spend", + meta: { title: "Spend (USD)" }, + header: ({ column }) => , + size: 120, + enableSorting: true, + cell: ({ row }) => , + }, + { + id: "max_budget", + meta: { title: "Budget (USD)" }, + header: "Budget (USD)", + size: 120, + enableSorting: false, + cell: ({ row }) => ( + + ), + }, + { + id: "models", + meta: { title: "Models", skeleton: "chips" }, + header: "Models", + size: 260, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "limits", + meta: { title: "TPM / RPM Limits" }, + header: "TPM / RPM Limits", + size: 150, + enableSorting: false, + cell: ({ row }) => , + }, + { + id: "members", + meta: { title: "Members" }, + header: "Members", + size: 100, + enableSorting: false, + cell: ({ row }) => {row.original.members?.length ?? 0} Members, + }, + { + id: "actions", + meta: { className: "text-right", headerClassName: "text-right" }, + header: () => Actions, + size: 64, + enableSorting: false, + enableHiding: false, + cell: ({ row }) => + userRole === "Admin" ? ( +
+ +
+ ) : null, + }, +]; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/organizations.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/organizations.test.tsx deleted file mode 100644 index 75a6d30ac2e..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/organizations.test.tsx +++ /dev/null @@ -1,39 +0,0 @@ -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { render } from "@testing-library/react"; -import React from "react"; -import { describe, expect, it, vi } from "vitest"; - -vi.mock("@/components/vector_store_management/VectorStoreSelector", () => ({ - __esModule: true, - default: () => null, -})); -vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({ - __esModule: true, - default: () => null, -})); -vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ - default: () => ({ - accessToken: null, - userId: null, - userRole: null, - }), -})); - -import OrganizationsTable from "./organizations"; - -const renderWithQueryClient = (ui: React.ReactElement) => { - const queryClient = new QueryClient({ - defaultOptions: { queries: { retry: false } }, - }); - return render({ui}); -}; - -describe("OrganizationsTable", () => { - it("should render the OrganizationsTable component", () => { - const { getByText } = renderWithQueryClient( - , - ); - - expect(getByText("+ Create New Organization")).toBeInTheDocument(); - }); -}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/organizations.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/organizations.tsx deleted file mode 100644 index 87d8010759d..00000000000 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/_components/organizations.tsx +++ /dev/null @@ -1,535 +0,0 @@ -import { organizationKeys, useOrganizations } from "@/app/(dashboard)/hooks/organizations/useOrganizations"; -import { useUserModels } from "@/app/(dashboard)/hooks/models/useModels"; -import OrganizationFilters, { FilterState } from "@/app/(dashboard)/organizations/OrganizationFilters"; -import { InfoCircleOutlined } from "@ant-design/icons"; -import { ChevronDownIcon, ChevronRightIcon, RefreshIcon } from "@heroicons/react/outline"; -import { - Badge, - Button, - Card, - Col, - Grid, - Icon, - Tab, - TabGroup, - Table, - TableBody, - TableCell, - TableHead, - TableHeaderCell, - TableRow, - TabList, - TabPanel, - TabPanels, - Text, - TextInput, -} from "@tremor/react"; -import { Form, Input, Modal, Select as Select2, Tooltip } from "antd"; -import { useQueryClient } from "@tanstack/react-query"; -import React, { useState } from "react"; -import { DateCell, IdCell, MoneyCell } from "@/components/shared/table_cells"; -import DeleteResourceModal from "@/components/common_components/DeleteResourceModal"; -import TableIconActionButton from "@/components/common_components/IconActionButton/TableIconActionButtons/TableIconActionButton"; -import { getModelDisplayName } from "@/components/key_team_helpers/fetch_available_models_team_key"; -import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector"; -import { ModelSelect } from "@/components/ModelSelect/ModelSelect"; -import NotificationsManager from "@/components/molecules/notifications_manager"; -import { - Organization, - organizationCreateCall, - organizationDeleteCall, - organizationListCall, -} from "@/components/networking"; -import OrganizationInfoView from "@/components/organization/organization_view"; -import NumericalInput from "@/components/shared/numerical_input"; -import VectorStoreSelector from "@/components/vector_store_management/VectorStoreSelector"; - -interface OrganizationsTableProps { - userRole: string; - accessToken: string | null; - lastRefreshed?: string; - handleRefreshClick?: () => void; - premiumUser: boolean; -} - -export const fetchOrganizations = async ( - accessToken: string, - setOrganizations: (organizations: Organization[]) => void, - org_id: string | null = null, - org_alias: string | null = null, -) => { - const organizations = await organizationListCall(accessToken, org_id, org_alias); - setOrganizations(organizations); -}; - -const OrganizationsTable: React.FC = ({ - userRole, - accessToken, - lastRefreshed, - handleRefreshClick, - premiumUser, -}) => { - const [selectedOrgId, setSelectedOrgId] = useState(null); - const [editOrg, setEditOrg] = useState(false); - const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); - const [orgToDelete, setOrgToDelete] = useState(null); - const [isDeleting, setIsDeleting] = useState(false); - const [isOrgModalVisible, setIsOrgModalVisible] = useState(false); - const [form] = Form.useForm(); - const [expandedAccordions, setExpandedAccordions] = useState>({}); - const [showFilters, setShowFilters] = useState(false); - const [filters, setFilters] = useState({ - org_id: "", - org_alias: "", - sort_by: "created_at", - sort_order: "desc", - }); - - const queryClient = useQueryClient(); - const { data: organizations = [] } = useOrganizations({ org_id: filters.org_id, org_alias: filters.org_alias }); - const { data: userModels = [] } = useUserModels(); - - const refetchOrganizations = () => queryClient.invalidateQueries({ queryKey: organizationKeys.lists() }); - - const handleFilterChange = (key: keyof FilterState, value: string) => { - setFilters((previousFilters) => ({ ...previousFilters, [key]: value })); - }; - - const handleFilterReset = () => { - setFilters({ - org_id: "", - org_alias: "", - sort_by: "created_at", - sort_order: "desc", - }); - }; - - const handleDelete = (orgId: string | null) => { - if (!orgId) return; - - setOrgToDelete(orgId); - setIsDeleteModalOpen(true); - }; - - const confirmDelete = async () => { - if (!orgToDelete || !accessToken) return; - - try { - setIsDeleting(true); - await organizationDeleteCall(accessToken, orgToDelete); - NotificationsManager.success("Organization deleted successfully"); - - setIsDeleteModalOpen(false); - setOrgToDelete(null); - await refetchOrganizations(); - } catch (error) { - console.error("Error deleting organization:", error); - } finally { - setIsDeleting(false); - } - }; - - const cancelDelete = () => { - setIsDeleteModalOpen(false); - setOrgToDelete(null); - }; - - const handleCreate = async (values: any) => { - try { - if (!accessToken) return; - - // Transform allowed_vector_store_ids and allowed_mcp_servers_and_groups into object_permission - if ( - (values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) || - (values.allowed_mcp_servers_and_groups && - (values.allowed_mcp_servers_and_groups.servers?.length > 0 || - values.allowed_mcp_servers_and_groups.accessGroups?.length > 0)) - ) { - values.object_permission = {}; - if (values.allowed_vector_store_ids && values.allowed_vector_store_ids.length > 0) { - values.object_permission.vector_stores = values.allowed_vector_store_ids; - delete values.allowed_vector_store_ids; - } - if (values.allowed_mcp_servers_and_groups) { - if (values.allowed_mcp_servers_and_groups.servers?.length > 0) { - values.object_permission.mcp_servers = values.allowed_mcp_servers_and_groups.servers; - } - if (values.allowed_mcp_servers_and_groups.accessGroups?.length > 0) { - values.object_permission.mcp_access_groups = values.allowed_mcp_servers_and_groups.accessGroups; - } - delete values.allowed_mcp_servers_and_groups; - } - } - - await organizationCreateCall(accessToken, values); - NotificationsManager.success("Organization created successfully"); - setIsOrgModalVisible(false); - form.resetFields(); - await refetchOrganizations(); - } catch (error) { - console.error("Error creating organization:", error); - } - }; - - const handleCancel = () => { - setIsOrgModalVisible(false); - form.resetFields(); - }; - - if (!premiumUser) { - return ( -
- - This is a LiteLLM Enterprise feature, and requires a valid key to use. Get a trial key{" "} - - here - - . - -
- ); - } - - return ( -
- - - {(userRole === "Admin" || userRole === "Org Admin") && ( - - )} - {selectedOrgId ? ( - { - setSelectedOrgId(null); - setEditOrg(false); - }} - accessToken={accessToken} - is_org_admin={true} // You'll need to implement proper org admin check - is_proxy_admin={userRole === "Admin"} - userModels={userModels} - editOrg={editOrg} - /> - ) : ( - - -
- Your Organizations -
-
- {lastRefreshed && Last Refreshed: {lastRefreshed}} - -
-
- - - Click on “Organization ID” to view organization details. - - - -
-
- -
-
- - - - Organization ID - Organization Name - Created - Spend (USD) - Budget (USD) - Models - TPM / RPM Limits - Info - Actions - - - - - {organizations && organizations.length > 0 - ? organizations - .sort((a, b) => new Date(b.created_at).getTime() - new Date(a.created_at).getTime()) - .map((org: Organization) => ( - - - - - {org.organization_alias} - - - - - - - - - - 3 ? "px-0" : ""} - > -
- {Array.isArray(org.models) ? ( -
- {org.models.length === 0 ? ( - - All Proxy Models - - ) : ( - <> -
- {org.models.length > 3 && ( -
- { - setExpandedAccordions((prev) => ({ - ...prev, - [org.organization_id || ""]: - !prev[org.organization_id || ""], - })); - }} - /> -
- )} -
- {org.models.slice(0, 3).map((model, index) => - model === "all-proxy-models" ? ( - - All Proxy Models - - ) : ( - - - {model.length > 30 - ? `${getModelDisplayName(model).slice(0, 30)}...` - : getModelDisplayName(model)} - - - ), - )} - {org.models.length > 3 && - !expandedAccordions[org.organization_id || ""] && ( - - - +{org.models.length - 3}{" "} - {org.models.length - 3 === 1 - ? "more model" - : "more models"} - - - )} - {expandedAccordions[org.organization_id || ""] && ( -
- {org.models.slice(3).map((model, index) => - model === "all-proxy-models" ? ( - - All Proxy Models - - ) : ( - - - {model.length > 30 - ? `${getModelDisplayName(model).slice(0, 30)}...` - : getModelDisplayName(model)} - - - ), - )} -
- )} -
-
- - )} -
- ) : null} -
-
- - - TPM:{" "} - {org.litellm_budget_table?.tpm_limit - ? org.litellm_budget_table?.tpm_limit - : "Unlimited"} -
- RPM:{" "} - {org.litellm_budget_table?.rpm_limit - ? org.litellm_budget_table?.rpm_limit - : "Unlimited"} -
-
- - {org.members?.length || 0} Members - - - {userRole === "Admin" && ( - <> - { - setSelectedOrgId(org.organization_id); - setEditOrg(true); - }} - /> - handleDelete(org.organization_id)} - /> - - )} - -
- )) - : null} -
-
-
- -
-
-
-
- )} - -
- -
- - - - - form.setFieldValue("models", values)} - context="organization" - /> - - - - - - - - daily - weekly - monthly - - - - - - - - - - - Allowed Vector Stores{" "} - - - - - } - name="allowed_vector_store_ids" - className="mt-4" - help="Select vector stores this organization can access. Leave empty for access to all vector stores" - > - form.setFieldValue("allowed_vector_store_ids", values)} - value={form.getFieldValue("allowed_vector_store_ids")} - accessToken={accessToken || ""} - placeholder="Select vector stores (optional)" - /> - - - - Allowed MCP Servers{" "} - - - - - } - name="allowed_mcp_servers_and_groups" - className="mt-4" - help="Select MCP servers and access groups this organization can access." - > - form.setFieldValue("allowed_mcp_servers_and_groups", values)} - value={form.getFieldValue("allowed_mcp_servers_and_groups")} - accessToken={accessToken || ""} - placeholder="Select MCP servers and access groups (optional)" - /> - - - - - - -
- -
-
-
- - -
- ); -}; - -export default OrganizationsTable; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/organizations/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/organizations/page.tsx index 649e54f63eb..a492a572580 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/organizations/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/organizations/page.tsx @@ -1,9 +1,9 @@ "use client"; -import OrganizationsTable from "./_components/organizations"; +import OrganizationsPanel from "./_components/OrganizationsPanel"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; export default function OrganizationsPage() { const { accessToken, userRole, premiumUser } = useAuthorized(); - return ; + return ; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx index 99c58e2b98f..d6a4aaea2e1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.test.tsx @@ -1,7 +1,8 @@ -import { render, screen } from "@testing-library/react"; +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; -import { getPromptsList } from "@/components/networking"; +import { deletePromptCall, getPromptsList } from "@/components/networking"; import PromptsPanel from "./index"; @@ -12,20 +13,39 @@ vi.mock("@/components/networking", () => ({ vi.mock("./PromptTable", () => ({ __esModule: true, - default: ({ isLoading }: { isLoading: boolean }) => ( -
{isLoading ? "table-loading" : "table-loaded"}
+ default: ({ + isLoading, + onDeleteClick, + }: { + isLoading: boolean; + onDeleteClick: (id: string, name: string) => void; + }) => ( +
+ {isLoading ? "table-loading" : "table-loaded"} + +
), })); -vi.mock("./prompt_info", () => ({ __esModule: true, default: () => null })); -vi.mock("./add_prompt_form", () => ({ __esModule: true, default: () => null })); -vi.mock("./prompt_editor_view", () => ({ __esModule: true, default: () => null })); +vi.mock("./prompt_info", () => ({ __esModule: true, default: () =>
prompt-info-view
})); +vi.mock("./add_prompt_form", () => ({ + __esModule: true, + default: ({ visible }: { visible: boolean }) => (visible ?
add-prompt-form
: null), +})); +vi.mock("./prompt_editor_view", () => ({ __esModule: true, default: () =>
prompt-editor-view
})); const mockGetPromptsList = vi.mocked(getPromptsList); +const mockDeletePromptCall = vi.mocked(deletePromptCall); + +const renderPanel = (userRole?: string) => + render(); describe("PromptsPanel loading state", () => { beforeEach(() => { vi.clearAllMocks(); + mockGetPromptsList.mockResolvedValue({ prompts: [] } as never); }); it("should resolve the loading state when accessToken is null instead of showing the skeleton forever", async () => { @@ -39,7 +59,7 @@ describe("PromptsPanel loading state", () => { mockGetPromptsList.mockReturnValue( new Promise((resolve) => { resolveFetch = resolve; - }), + }) as never, ); render(); expect(screen.getByText("table-loading")).toBeInTheDocument(); @@ -49,3 +69,134 @@ describe("PromptsPanel loading state", () => { expect(mockGetPromptsList).toHaveBeenCalledWith("sk-test", undefined); }); }); + +describe("PromptsPanel toolbar", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockGetPromptsList.mockResolvedValue({ prompts: [] } as never); + }); + + it("should offer both create actions to a proxy admin", async () => { + renderPanel("Admin"); + + expect(await screen.findByRole("button", { name: /add new prompt/i })).toBeEnabled(); + expect(screen.getByRole("button", { name: /upload \.prompt file/i })).toBeEnabled(); + }); + + it("should hide both create actions from a read-only viewer", async () => { + renderPanel("Admin Viewer"); + + expect(await screen.findByText("table-loaded")).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /add new prompt/i })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: /upload \.prompt file/i })).not.toBeInTheDocument(); + }); + + it("should open the editor view when the add action is used", async () => { + const user = userEvent.setup(); + renderPanel("Admin"); + + await user.click(await screen.findByRole("button", { name: /add new prompt/i })); + + expect(screen.getByText("prompt-editor-view")).toBeInTheDocument(); + expect(screen.queryByTestId("prompt-table")).not.toBeInTheDocument(); + }); + + it("should open the upload form when the upload action is used", async () => { + const user = userEvent.setup(); + renderPanel("Admin"); + + expect(screen.queryByText("add-prompt-form")).not.toBeInTheDocument(); + await user.click(await screen.findByRole("button", { name: /upload \.prompt file/i })); + + expect(screen.getByText("add-prompt-form")).toBeInTheDocument(); + }); + + it("should refetch scoped to the environment picked in the filter", async () => { + const user = userEvent.setup(); + renderPanel("Admin"); + await screen.findByText("table-loaded"); + + expect(screen.getByText("All Environments")).toBeInTheDocument(); + + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByText("Production")); + + await waitFor(() => expect(mockGetPromptsList).toHaveBeenLastCalledWith("sk-test", "production")); + }); + + it("should show the picked environment by label and clear back to the unfiltered list", async () => { + // Base UI's exit animation never completes in jsdom, so the closing popup keeps + // pointer-events: none and blocks the second open. The clicks still dispatch. + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + renderPanel("Admin"); + await screen.findByText("table-loaded"); + + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByText("Production")); + await waitFor(() => expect(screen.getByRole("combobox")).toHaveTextContent("Production")); + + await user.click(screen.getByRole("combobox")); + await user.click(await screen.findByText("All Environments")); + + await waitFor(() => expect(screen.getByRole("combobox")).toHaveTextContent("All Environments")); + await waitFor(() => expect(mockGetPromptsList).toHaveBeenLastCalledWith("sk-test", undefined)); + }); +}); + +describe("PromptsPanel delete confirmation", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockGetPromptsList.mockResolvedValue({ prompts: [] } as never); + mockDeletePromptCall.mockResolvedValue(undefined as never); + }); + + it("should not delete until the confirmation is accepted", async () => { + const user = userEvent.setup(); + renderPanel("Admin"); + + await user.click(await screen.findByRole("button", { name: "row-delete" })); + + expect(await screen.findByText(/delete prompt: my-prompt/i)).toBeInTheDocument(); + expect(screen.getByText(/cannot be undone/i)).toBeInTheDocument(); + expect(mockDeletePromptCall).not.toHaveBeenCalled(); + + await user.click(screen.getByRole("button", { name: /^delete$/i })); + + await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1")); + }); + + it("should abandon the delete when the confirmation is dismissed", async () => { + const user = userEvent.setup(); + renderPanel("Admin"); + + await user.click(await screen.findByRole("button", { name: "row-delete" })); + await screen.findByText(/delete prompt: my-prompt/i); + + await user.click(screen.getByRole("button", { name: /cancel/i })); + + await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument()); + expect(mockDeletePromptCall).not.toHaveBeenCalled(); + }); + + it("should keep the confirmation up while the delete request is still in flight", async () => { + const user = userEvent.setup(); + let finishDelete: () => void = () => {}; + mockDeletePromptCall.mockReturnValue( + new Promise((resolve) => { + finishDelete = () => resolve(); + }) as never, + ); + renderPanel("Admin"); + + await user.click(await screen.findByRole("button", { name: "row-delete" })); + await screen.findByText(/delete prompt: my-prompt/i); + await user.click(screen.getByRole("button", { name: /^delete$/i })); + await waitFor(() => expect(mockDeletePromptCall).toHaveBeenCalledWith("sk-test", "prompt-1")); + + await user.keyboard("{Escape}"); + expect(screen.getByText(/delete prompt: my-prompt/i)).toBeInTheDocument(); + + finishDelete(); + await waitFor(() => expect(screen.queryByText(/delete prompt: my-prompt/i)).not.toBeInTheDocument()); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx index de461ebd86d..9bebabb8cf2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/prompts/_components/index.tsx @@ -1,7 +1,6 @@ import React, { useState, useEffect } from "react"; -import { Button } from "@tremor/react"; -import { Modal, Select } from "antd"; +import { Plus, Upload } from "lucide-react"; import { getPromptsList, PromptSpec, ListPromptsResponse, deletePromptCall } from "@/components/networking"; import PromptTable from "./PromptTable"; import PromptInfoView from "./prompt_info"; @@ -9,6 +8,28 @@ import AddPromptForm from "./add_prompt_form"; import PromptEditorView from "./prompt_editor_view"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { isAdminRole, isProxyAdminRole } from "@/utils/roles"; +import { Button } from "@/components/ui/button"; +import { + AlertDialog, + AlertDialogCancel, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "@/components/ui/alert-dialog"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; + +const ALL_ENVIRONMENTS_LABEL = "All Environments"; + +const ENVIRONMENT_OPTIONS = [ + { label: "Development", value: "development" }, + { label: "Staging", value: "staging" }, + { label: "Production", value: "production" }, +]; + +// SelectValue falls back to the raw value unless the root can map it to a label. +const ENVIRONMENT_ITEMS = [{ label: ALL_ENVIRONMENTS_LABEL, value: null }, ...ENVIRONMENT_OPTIONS]; interface PromptsProps { accessToken: string | null; @@ -141,26 +162,33 @@ const PromptsPanel: React.FC = ({ accessToken, userRole }) => { {canModify && ( <> )}
= ({ accessToken, userRole }) => { /> {promptToDelete && ( - { + if (!open && !isDeleting) handleDeleteCancel(); + }} > -

Are you sure you want to delete prompt: {promptToDelete.name} ?

-

This action cannot be undone.

-
+ + + Delete Prompt + + Are you sure you want to delete prompt: {promptToDelete.name} ? This action cannot be undone. + + + + Cancel + + + + )} ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx new file mode 100644 index 00000000000..c4cfa98b2a1 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.test.tsx @@ -0,0 +1,101 @@ +import { renderWithProviders, screen, within } from "../../../../../tests/test-utils"; +import userEvent from "@testing-library/user-event"; +import { vi } from "vitest"; +import GeneralSettings from "./general_settings"; +import { deleteConfigFieldSetting, getGeneralSettingsCall, updateConfigFieldSetting } from "@/components/networking"; + +vi.mock("@/components/networking", () => ({ + getGeneralSettingsCall: vi.fn(), + updateConfigFieldSetting: vi.fn().mockResolvedValue({}), + deleteConfigFieldSetting: vi.fn().mockResolvedValue({}), +})); + +vi.mock("@/components/router_settings", () => ({ default: () => null })); +vi.mock("@/components/Settings/RouterSettings/Fallbacks/Fallbacks", () => ({ default: () => null })); +vi.mock("@/components/routing_groups", () => ({ default: () => null })); + +// Mirrors the /config/list ordering: the two prompt-caching rows sit between the +// General-tab rows in the unfiltered response but are filtered out of the General +// tab's table, so any index-based lookup into the unfiltered array reads the wrong +// row for every field rendered after them. +const SETTINGS_FIXTURE = [ + { + field_name: "budget_exceeded_throttle_percentage", + field_type: "Float", + field_value: null, + field_description: "throttle fraction", + stored_in_db: null, + field_default_value: null, + }, + { + field_name: "enable_anthropic_prompt_caching", + field_type: "Boolean", + field_value: true, + field_description: "prompt caching toggle", + stored_in_db: true, + field_tab: "prompt_caching", + field_default_value: false, + }, + { + field_name: "anthropic_prompt_caching_ttl", + field_type: "Select", + field_value: "5m", + field_description: "prompt caching ttl", + stored_in_db: true, + field_options: ["5m", "1h"], + field_tab: "prompt_caching", + field_default_value: null, + }, + { + field_name: "max_ui_session_budget", + field_type: "Dollar", + field_value: 7.5, + field_description: "dashboard session budget", + stored_in_db: true, + field_default_value: 1.0, + }, +]; + +const settingsRow = async (fieldName: string) => { + const cell = await screen.findByText(fieldName); + const row = cell.closest("tr"); + expect(row).not.toBeNull(); + return row as HTMLElement; +}; + +describe("GeneralSettings General tab", () => { + beforeEach(() => { + vi.mocked(getGeneralSettingsCall).mockResolvedValue([...SETTINGS_FIXTURE.map((s) => ({ ...s }))]); + vi.mocked(updateConfigFieldSetting).mockClear(); + vi.mocked(deleteConfigFieldSetting).mockClear(); + }); + + it("updates max_ui_session_budget with its own value, not the value at its filtered index", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByText("General")); + const row = await settingsRow("max_ui_session_budget"); + + await user.click(within(row).getByRole("button", { name: /update/i })); + + expect(updateConfigFieldSetting).toHaveBeenCalledWith("token", "max_ui_session_budget", 7.5); + }); + + it("reset shows the field's default value instead of an empty input", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.click(screen.getByText("General")); + const row = await settingsRow("max_ui_session_budget"); + expect(within(row).getByRole("spinbutton")).toHaveValue("7.50"); + + const actionCell = row.querySelectorAll("td")[3]; + const resetIcon = actionCell.querySelector("svg"); + expect(resetIcon).not.toBeNull(); + await user.click(resetIcon as unknown as Element); + + expect(deleteConfigFieldSetting).toHaveBeenCalledWith("token", "max_ui_session_budget"); + expect(within(row).getByRole("spinbutton")).toHaveValue("1.00"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx index dda7a23a8d4..fa3447e0cbf 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx @@ -41,6 +41,7 @@ export interface generalSettingsItem { stored_in_db: boolean | null; field_options?: string[] | null; field_tab?: string | null; + field_default_value?: any; } const SettingValueEditor: React.FC<{ @@ -75,6 +76,17 @@ const SettingValueEditor: React.FC<{ /> ); } + if (setting.field_type === "Dollar") { + return ( + onChange(setting.field_name, newValue)} + /> + ); + } if (setting.field_type === "Select") { return ( = ({ accessToken, user setGeneralSettings(updatedSettings); }; - const handleUpdateField = (fieldName: string, idx: number) => { + const handleUpdateField = (fieldName: string) => { if (!accessToken) { return; } - let fieldValue = generalSettings[idx].field_value; + let fieldValue = generalSettings.find((setting) => setting.field_name === fieldName)?.field_value; if (fieldValue == null || fieldValue == undefined) { return; @@ -194,7 +206,7 @@ const GeneralSettings: React.FC = ({ accessToken, user } }; - const handleResetField = (fieldName: string, idx: number) => { + const handleResetField = (fieldName: string) => { if (!accessToken) { return; } @@ -204,7 +216,9 @@ const GeneralSettings: React.FC = ({ accessToken, user // update value in state const updatedSettings = generalSettings.map((setting) => - setting.field_name === fieldName ? { ...setting, stored_in_db: null, field_value: null } : setting, + setting.field_name === fieldName + ? { ...setting, stored_in_db: null, field_value: setting.field_default_value ?? null } + : setting, ); setGeneralSettings(updatedSettings); } catch (error) { @@ -281,8 +295,8 @@ const GeneralSettings: React.FC = ({ accessToken, user )} - - handleResetField(value.field_name, index)}> + + handleResetField(value.field_name)}> Reset diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.test.tsx new file mode 100644 index 00000000000..7f1d7edc97f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.test.tsx @@ -0,0 +1,34 @@ +import { render, screen } from "@testing-library/react"; +import { describe, expect, it } from "vitest"; +import { SearchProviderLabel } from "./CreateSearchTools"; + +describe("SearchProviderLabel", () => { + it("renders the tavily logo from the static bundle, untouched by server-root prefixing", () => { + render(); + const img = screen.getByRole("img", { name: "Tavily logo" }); + expect(img).toHaveAttribute("src", "/_next/static/media/tavily.png"); + }); + + it("renders the exa_ai logo file for the exa_ai slug", () => { + render(); + const img = screen.getByRole("img", { name: "Exa AI logo" }); + expect(img.getAttribute("src")).toContain("exa_ai.png"); + }); + + it("renders the google_pse logo file for the google_pse slug", () => { + render(); + expect(screen.getByRole("img", { name: "Google PSE logo" }).getAttribute("src")).toContain("google_pse.png"); + }); + + it("falls back to a letter avatar for a provider with no bundled logo", () => { + render(); + expect(screen.queryByRole("img")).toBeNull(); + expect(screen.getByText("B")).toBeInTheDocument(); + expect(screen.getByText("Brave Search")).toBeInTheDocument(); + }); + + it("does not guess a legacy /ui/assets/logos/.png url for unknown providers", () => { + const { container } = render(); + expect(container.querySelector("img")).toBeNull(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx index b1cb5eb5581..1eeff00cb1b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/CreateSearchTools.tsx @@ -4,44 +4,37 @@ import { useQuery } from "@tanstack/react-query"; import { Button, TextInput } from "@tremor/react"; import { Form, Input, Modal, Select, Tooltip, Typography } from "antd"; import React, { useState } from "react"; -import { resolveLogoSrc } from "@/lib/assetPaths"; +import { Logo } from "@/components/molecules/logo/Logo"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { createSearchTool, fetchAvailableSearchProviders } from "@/components/networking"; import SearchConnectionTest from "./SearchConnectionTest"; import { AvailableSearchProvider, SearchTool } from "./types"; +import dataforseoLogo from "../../../../../public/assets/logos/dataforseo.png"; +import exaAiLogo from "../../../../../public/assets/logos/exa_ai.png"; +import googlePseLogo from "../../../../../public/assets/logos/google_pse.png"; +import parallelAiLogo from "../../../../../public/assets/logos/parallel_ai.png"; +import perplexityLogo from "../../../../../public/assets/logos/perplexity.png"; +import tavilyLogo from "../../../../../public/assets/logos/tavily.png"; const { TextArea } = Input; -// Search provider logos folder path (matches existing provider logo pattern) -const searchProviderLogosFolder = "/ui/assets/logos/"; - -// Helper function to get logo path for a search provider -const getSearchProviderLogo = (providerName: string): string => { - return `${searchProviderLogosFolder}${providerName}.png`; +const searchProviderLogoMap: Record = { + perplexity: perplexityLogo.src, + tavily: tavilyLogo.src, + parallel_ai: parallelAiLogo.src, + exa_ai: exaAiLogo.src, + google_pse: googlePseLogo.src, + dataforseo: dataforseoLogo.src, }; -// Component to display search provider logo and name interface SearchProviderLabelProps { providerName: string; displayName: string; } -const SearchProviderLabel: React.FC = ({ providerName, displayName }) => ( -
- {/* eslint-disable-next-line @next/next/no-img-element */} - { - e.currentTarget.style.display = "none"; - }} - /> +export const SearchProviderLabel: React.FC = ({ providerName, displayName }) => ( +
+ {displayName}
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.test.tsx new file mode 100644 index 00000000000..cfe5d4d843e --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.test.tsx @@ -0,0 +1,130 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import SearchConnectionTest from "./SearchConnectionTest"; +import * as networking from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; + +vi.mock("@/components/networking", () => ({ + testSearchToolConnection: vi.fn(), +})); + +const defaultProps = { + litellmParams: { search_provider: "tavily" }, + accessToken: "test-token", +}; + +describe("SearchConnectionTest", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("passes the access token and params to the connection test", async () => { + vi.mocked(networking.testSearchToolConnection).mockResolvedValue({ + status: "success", + message: "ok", + }); + + render(); + + await waitFor(() => { + expect(networking.testSearchToolConnection).toHaveBeenCalledWith( + defaultProps.accessToken, + defaultProps.litellmParams, + ); + }); + }); + + it("shows a loading state naming the provider while the test is pending", () => { + vi.mocked(networking.testSearchToolConnection).mockReturnValue(new Promise(() => {})); + + render(); + + expect(screen.getByText(/Testing connection to tavily/i)).toBeInTheDocument(); + }); + + it("renders a success state with the test query and result count", async () => { + vi.mocked(networking.testSearchToolConnection).mockResolvedValue({ + status: "success", + message: "ok", + test_query: "hello world", + results_count: 3, + }); + + render(); + + expect(await screen.findByText(/Connection to tavily successful/i)).toBeInTheDocument(); + expect(screen.getByText("hello world")).toBeInTheDocument(); + expect(screen.getByText(/Results retrieved: 3/i)).toBeInTheDocument(); + }); + + it("fires a success notification and completion callback on a successful test", async () => { + const onTestComplete = vi.fn(); + vi.mocked(networking.testSearchToolConnection).mockResolvedValue({ + status: "success", + message: "ok", + }); + + render(); + + await waitFor(() => { + expect(NotificationsManager.success).toHaveBeenCalledWith("Connection test successful!"); + }); + expect(onTestComplete).toHaveBeenCalledTimes(1); + }); + + it("renders a failure state with a cleaned error message and error type", async () => { + vi.mocked(networking.testSearchToolConnection).mockResolvedValue({ + status: "error", + message: "litellm.AuthenticationError: Invalid API key\nstack trace: deep internals", + error_type: "AuthenticationError", + }); + + render(); + + expect(await screen.findByText(/Connection to tavily failed/i)).toBeInTheDocument(); + expect(screen.getByText("Invalid API key")).toBeInTheDocument(); + expect(screen.getByText("AuthenticationError")).toBeInTheDocument(); + expect(screen.getByText("Verify your API key is correct and active")).toBeInTheDocument(); + }); + + it("reveals the raw error details when Show Details is toggled", async () => { + const user = userEvent.setup(); + vi.mocked(networking.testSearchToolConnection).mockResolvedValue({ + status: "error", + message: "litellm.AuthenticationError: Invalid API key\nstack trace: deep internals", + error_type: "AuthenticationError", + }); + + render(); + + const toggle = await screen.findByRole("button", { name: /show details/i }); + expect(screen.queryByText("Full Error Details")).not.toBeInTheDocument(); + + await user.click(toggle); + + expect(screen.getByText("Full Error Details")).toBeInTheDocument(); + expect(screen.getByText(/stack trace: deep internals/i)).toBeInTheDocument(); + }); + + it("treats a rejected request as a connection failure", async () => { + vi.mocked(networking.testSearchToolConnection).mockRejectedValue(new Error("network down")); + + render(); + + expect(await screen.findByText(/Connection to tavily failed/i)).toBeInTheDocument(); + expect(screen.getByText("network down")).toBeInTheDocument(); + }); + + it("links out to the search documentation", async () => { + vi.mocked(networking.testSearchToolConnection).mockResolvedValue({ + status: "success", + message: "ok", + }); + + render(); + + const docLink = await screen.findByRole("link", { name: /View Search Documentation/i }); + expect(docLink).toHaveAttribute("href", "https://docs.litellm.ai/docs/search"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.tsx index 4e8678ded71..446d9a71517 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchConnectionTest.tsx @@ -1,10 +1,10 @@ -import { InfoCircleOutlined, WarningOutlined } from "@ant-design/icons"; -import { Button, Divider, Typography } from "antd"; +import { AlertTriangle, CheckCircle2, Info } from "lucide-react"; import React, { useEffect, useState } from "react"; import NotificationsManager from "@/components/molecules/notifications_manager"; import { testSearchToolConnection } from "@/components/networking"; - -const { Text } = Typography; +import { Button } from "@/components/ui/button"; +import { Separator } from "@/components/ui/separator"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; interface SearchConnectionTestProps { litellmParams: Record; @@ -52,30 +52,23 @@ const SearchConnectionTest: React.FC = ({ litellmPara const getCleanErrorMessage = (errorMsg: string) => { if (!errorMsg) return "Unknown error"; - // Remove stack traces const mainError = errorMsg.split("stack trace:")[0].trim(); - // Remove litellm error prefixes const cleanedError = mainError.replace(/^litellm\.(.*?)Error:\s*/, ""); - // Remove AuthenticationError prefix if it exists const finalError = cleanedError.replace(/^AuthenticationError:\s*/, ""); - // If the error contains HTML (like a 401 page), extract just the key info if (finalError.includes("") || finalError.includes("(.*?)<\/title>/); if (titleMatch) { return titleMatch[1]; } - // If it's a 401 error if (finalError.includes("401") || finalError.includes("Authorization Required")) { return "Authentication failed: Invalid API key or credentials"; } return "Authentication error - please check your API key"; } - // Limit very long error messages if (finalError.length > 200) { return finalError.substring(0, 200) + "..."; } @@ -87,34 +80,12 @@ const SearchConnectionTest: React.FC = ({ litellmPara if (isLoading) { return ( -
-
-
-
-
- +
+
+ +

Testing connection to {litellmParams.search_provider || "search provider"}... - - +

); @@ -125,147 +96,88 @@ const SearchConnectionTest: React.FC = ({ litellmPara } return ( -
+
{testResult.status === "success" ? ( -
-
- -
-
- +
+ +
+

Connection to {litellmParams.search_provider} successful! - +

{testResult.test_query && ( - - Test query:{" "} - - {testResult.test_query} - - +

+ Test query: {testResult.test_query} +

)} {testResult.results_count !== undefined && ( - - Results retrieved: {testResult.results_count} - +

Results retrieved: {testResult.results_count}

)}
) : ( - <> -
-
- - - Connection to {litellmParams.search_provider || "search provider"} failed - -
+
+
+ +

+ Connection to {litellmParams.search_provider || "search provider"} failed +

+
-
- - Error:{" "} - - - {errorMessage} - +
+

Error:

+

{errorMessage}

- {testResult.error_type && ( -
- - Error type:{" "} - - {testResult.error_type} - - -
- )} - - {testResult.message && ( -
- -
- )} -
- - {showDetails && ( -
- - Full Error Details - -
-                  {testResult.message}
-                
+ {testResult.error_type && ( +
+

+ Error type:{" "} + + {testResult.error_type} + +

)} -
- - Troubleshooting tips: - -
    -
  • Verify your API key is correct and active
  • -
  • Check if the search provider service is operational
  • -
  • Ensure you have sufficient credits/quota with the provider
  • -
  • - Review the provider's documentation for any additional requirements -
  • -
-
+ {testResult.message && ( +
+ +
+ )}
- + + {showDetails && ( +
+

Full Error Details

+
+                {testResult.message}
+              
+
+ )} + +
+

Troubleshooting tips:

+
    +
  • Verify your API key is correct and active
  • +
  • Check if the search provider service is operational
  • +
  • Ensure you have sufficient credits/quota with the provider
  • +
  • Review the provider's documentation for any additional requirements
  • +
+
+
)} - -
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTester.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTester.tsx index 2fe9f3b5b8c..95772608235 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTester.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolTester.tsx @@ -1,12 +1,12 @@ import React, { useState } from "react"; -import { Button, Input, Typography, Spin } from "antd"; +import { ExternalLink, Search } from "lucide-react"; import MessageManager from "@/components/molecules/message_manager"; -import { SearchOutlined, LoadingOutlined } from "@ant-design/icons"; import { searchToolQueryCall } from "@/components/networking"; import NotificationsManager from "@/components/molecules/notifications_manager"; -import { Card, Title as TremorTitle } from "@tremor/react"; - -const { Text } = Typography; +import { Button } from "@/components/ui/button"; +import { Card } from "@/components/ui/card"; +import { Input } from "@/components/ui/input"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; interface SearchResult { title: string; @@ -36,7 +36,6 @@ export const SearchToolTester: React.FC = ({ searchToolNa }[] >([]); const [expandedResults, setExpandedResults] = useState>({}); - const [isInputFocused, setIsInputFocused] = useState(false); const handleSearch = async () => { if (!query.trim()) { @@ -60,7 +59,6 @@ export const SearchToolTester: React.FC = ({ searchToolNa }; setSearchHistory((prev) => [historyEntry, ...prev]); - // Don't clear query after search so user can modify it } catch (error) { console.error("Error querying search tool:", error); NotificationsManager.fromBackend("Failed to query search tool"); @@ -87,113 +85,79 @@ export const SearchToolTester: React.FC = ({ searchToolNa })); }; - const antIcon = ; - const latestResults = searchHistory.length > 0 ? searchHistory[0] : null; return ( - -
- Test Search Tool + +
+

Test Search Tool

-
- {/* Search Bar at Top */} +
-
- +
+ setQuery(e.target.value)} - onFocus={() => setIsInputFocused(true)} - onBlur={() => setIsInputFocused(false)} - onPressEnter={(e) => { - if (!e.shiftKey) { + onKeyDown={(e) => { + if (e.key === "Enter" && !e.shiftKey) { e.preventDefault(); handleSearch(); } }} placeholder="Enter your search query..." disabled={isLoading} - bordered={false} - style={{ fontSize: "15px", padding: 0, height: "100%", boxShadow: "none" }} + className="h-12 pl-11 text-[15px]" />
-
- {/* Results Area */}
{!latestResults && !isLoading ? ( -
-
- +
+
+
- Test your search tool - Enter a query above to see search results +

Test your search tool

+

Enter a query above to see search results

) : (
{isLoading && ( -
- - Searching... +
+ +

Searching...

)} {latestResults && !isLoading && ( <> - {/* Query Info Bar */} -
+
- +

Search Query - -

{latestResults.query}
+

+
{latestResults.query}
-
- {formatTimestamp(latestResults.timestamp)} -
-
+
+

{formatTimestamp(latestResults.timestamp)}

+
+
{latestResults.response?.results?.length || 0}{" "} {latestResults.response?.results?.length === 1 ? "result" : "results"}
{latestResults.latency !== undefined && ( <> - • -
{latestResults.latency}ms
+ • +
{latestResults.latency}ms
)}
@@ -201,7 +165,6 @@ export const SearchToolTester: React.FC = ({ searchToolNa
- {/* Search Results */} {latestResults.response && latestResults.response.results && latestResults.response.results.length > 0 ? ( @@ -212,73 +175,43 @@ export const SearchToolTester: React.FC = ({ searchToolNa return (
{ - e.currentTarget.style.boxShadow = - "0 4px 6px -1px rgba(0, 0, 0, 0.1), 0 2px 4px -1px rgba(0, 0, 0, 0.06)"; - e.currentTarget.style.borderColor = "#e0e7ff"; - }} - onMouseLeave={(e) => { - e.currentTarget.style.boxShadow = "0 1px 2px 0 rgba(0, 0, 0, 0.05)"; - e.currentTarget.style.borderColor = "#e5e7eb"; - }} + className="rounded-lg border border-border bg-card transition-shadow hover:shadow-md" >
- {/* Title and External Link */} -
+
(e.currentTarget.style.textDecoration = "underline")} - onMouseLeave={(e) => (e.currentTarget.style.textDecoration = "none")} + className="flex-1 text-lg leading-snug font-semibold text-primary hover:underline" > {result.title}
- {/* URL */} -
{result.url}
+
{result.url}
- {/* Snippet Preview */} -
+
{isResultExpanded ? result.snippet : `${result.snippet.substring(0, 200)}${result.snippet.length > 200 ? "..." : ""}`}
- {/* Expand/Collapse */} {result.snippet.length > 200 && ( @@ -289,31 +222,22 @@ export const SearchToolTester: React.FC = ({ searchToolNa })}
) : ( -
-
- +
+
+
- No results found - Try a different search query +

No results found

+

Try a different search query

)} )} - {/* Search History Sidebar */} {searchHistory.length > 1 && ( -
-
- Previous Searches -
@@ -321,21 +245,21 @@ export const SearchToolTester: React.FC = ({ searchToolNa {searchHistory.slice(1, 6).map((entry, index) => (
{ setQuery(entry.query); }} > -
{entry.query}
-
- +
{entry.query}
+
+ {entry.response?.results?.length || 0}{" "} {entry.response?.results?.length === 1 ? "result" : "results"} {entry.latency !== undefined && ( <> • - {entry.latency}ms + {entry.latency}ms )} • diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.test.tsx index 049523c0254..e4b2cf62940 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.test.tsx @@ -155,13 +155,8 @@ describe("SearchToolView", () => { const toolNameContainer = screen.getByText("Test Search Tool").closest("div"); expect(toolNameContainer).toBeInTheDocument(); - const copyButtons = within(toolNameContainer!).getAllByRole("button"); - const nameCopyButton = copyButtons.find((button) => { - return button.querySelector("svg") !== null; - }); - - expect(nameCopyButton).toBeInTheDocument(); - await user.click(nameCopyButton!); + const nameCopyButton = within(toolNameContainer!).getByRole("button"); + await user.click(nameCopyButton); await waitFor(() => { expect(copyToClipboard).toHaveBeenCalledWith("Test Search Tool"); @@ -176,13 +171,8 @@ describe("SearchToolView", () => { const toolIdContainer = screen.getByText("test-tool-id-123").closest("div"); expect(toolIdContainer).toBeInTheDocument(); - const copyButtons = within(toolIdContainer!).getAllByRole("button"); - const idCopyButton = copyButtons.find((button) => { - return button.querySelector("svg") !== null; - }); - - expect(idCopyButton).toBeInTheDocument(); - await user.click(idCopyButton!); + const idCopyButton = within(toolIdContainer!).getByRole("button"); + await user.click(idCopyButton); await waitFor(() => { expect(copyToClipboard).toHaveBeenCalledWith("test-tool-id-123"); @@ -197,22 +187,14 @@ describe("SearchToolView", () => { render(); const toolNameContainer = screen.getByText("Test Search Tool").closest("div"); - const copyButtons = within(toolNameContainer!).getAllByRole("button"); - const nameCopyButton = copyButtons.find((button) => { - return button.querySelector("svg") !== null; - }); + const nameCopyButton = within(toolNameContainer!).getByRole("button"); - expect(nameCopyButton).toBeInTheDocument(); + expect(nameCopyButton.querySelector(".lucide-copy")).toBeInTheDocument(); - const initialSvg = nameCopyButton!.querySelector("svg"); - expect(initialSvg).toBeInTheDocument(); - - await user.click(nameCopyButton!); + await user.click(nameCopyButton); await waitFor(() => { - const updatedSvg = nameCopyButton!.querySelector("svg"); - expect(updatedSvg).toBeInTheDocument(); - expect(nameCopyButton).toHaveClass("text-green-600"); + expect(nameCopyButton.querySelector(".lucide-check")).toBeInTheDocument(); }); }); @@ -224,13 +206,9 @@ describe("SearchToolView", () => { render(); const toolNameContainer = screen.getByText("Test Search Tool").closest("div"); - const copyButtons = within(toolNameContainer!).getAllByRole("button"); - const nameCopyButton = copyButtons.find((button) => { - return button.querySelector("svg") !== null; - }); + const nameCopyButton = within(toolNameContainer!).getByRole("button"); - expect(nameCopyButton).toBeInTheDocument(); - await user.click(nameCopyButton!); + await user.click(nameCopyButton); await waitFor( () => { @@ -239,7 +217,7 @@ describe("SearchToolView", () => { { timeout: 3000 }, ); - expect(nameCopyButton).not.toHaveClass("text-green-600"); + expect(nameCopyButton.querySelector(".lucide-check")).not.toBeInTheDocument(); }); it("should render SearchToolTester when accessToken is provided", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.tsx index e77234aa3a0..1f7c992aa98 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/search-tools/_components/SearchToolView.tsx @@ -1,9 +1,8 @@ import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; -import { ArrowLeftIcon } from "@heroicons/react/outline"; -import { Button, Card, Grid, Text, Title } from "@tremor/react"; -import { Button as AntdButton } from "antd"; -import { CheckIcon, CopyIcon } from "lucide-react"; +import { ArrowLeft, Check, Copy } from "lucide-react"; import React, { useState } from "react"; +import { Button } from "@/components/ui/button"; +import { Card, CardContent } from "@/components/ui/card"; import { SearchToolTester } from "./SearchToolTester"; import { AvailableSearchProvider, SearchTool } from "./types"; @@ -43,73 +42,73 @@ export const SearchToolView: React.FC = ({
- -
- {searchTool.search_tool_name} - : } +
+

{searchTool.search_tool_name}

+
-
- {searchTool.search_tool_id} - : } +
+

{searchTool.search_tool_id}

+
- +
- Provider -
- {getProviderDisplayName(searchTool.litellm_params.search_provider)} -
+ +

Provider

+

+ {getProviderDisplayName(searchTool.litellm_params.search_provider)} +

+
- API Key -
- {searchTool.litellm_params.api_key ? "****" : "Not set"} -
+ +

API Key

+

{searchTool.litellm_params.api_key ? "****" : "Not set"}

+
- Created At -
- {searchTool.created_at ? new Date(searchTool.created_at).toLocaleString() : "Unknown"} -
+ +

Created At

+

+ {searchTool.created_at ? new Date(searchTool.created_at).toLocaleString() : "Unknown"} +

+
- +
{searchTool.search_tool_info?.description && ( - Description -
- {searchTool.search_tool_info.description} -
+ +

Description

+

{searchTool.search_tool_info.description}

+
)} - {/* Search Tool Tester */}
{accessToken && }
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/tool-policies/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/tool-policies/page.tsx index 6aaebaab959..08fded8dca6 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/tool-policies/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/tool-policies/page.tsx @@ -4,6 +4,6 @@ import ToolPoliciesView from "@/components/ToolPoliciesView"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; export default function ToolPolicies() { - const { accessToken, userRole } = useAuthorized(); - return ; + const { accessToken } = useAuthorized(); + return ; } diff --git a/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.test.tsx new file mode 100644 index 00000000000..a0add153116 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.test.tsx @@ -0,0 +1,160 @@ +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; +import TransformRequestPanel from "./TransformRequestPanel"; +import { transformRequestCall } from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; + +vi.mock("@/components/networking", () => ({ + transformRequestCall: vi.fn(), +})); + +vi.mock("@/components/molecules/notifications_manager", () => ({ + default: { + success: vi.fn(), + info: vi.fn(), + fromBackend: vi.fn(), + }, +})); + +const transformRequestCallMock = vi.mocked(transformRequestCall); +const notify = vi.mocked(NotificationsManager); + +const ACCESS_TOKEN = "sk-test-token"; + +const getRequestTextarea = () => screen.getByPlaceholderText(/press cmd\/ctrl \+ enter to transform/i); + +const getTransformButton = () => screen.getByRole("button", { name: /transform/i }); + +const getCopyButton = () => screen.getByRole("button", { name: /copy to clipboard/i }); + +describe("TransformRequestPanel", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + afterEach(() => { + vi.restoreAllMocks(); + }); + + it("renders both panels, the prefilled request and the placeholder curl", () => { + render(); + + expect(screen.getByText("Original Request")).toBeInTheDocument(); + expect(screen.getByText("Transformed Request")).toBeInTheDocument(); + expect(screen.getByText(/sensitive headers are not shown/i)).toBeInTheDocument(); + + expect((getRequestTextarea() as HTMLTextAreaElement).value).toContain('"model": "openai/gpt-4o"'); + expect(screen.getByText(/https:\/\/api\.openai\.com\/v1\/chat\/completions/)).toBeInTheDocument(); + + expect(screen.getByRole("link", { name: /here/i })).toHaveAttribute( + "href", + "https://github.com/BerriAI/litellm/issues", + ); + }); + + it("sends the edited request body as a completion call and renders the returned curl", async () => { + const user = userEvent.setup(); + transformRequestCallMock.mockResolvedValue({ + raw_request_api_base: "https://api.anthropic.com/v1/messages", + raw_request_body: { model: "claude-opus-4-8", max_tokens: 42 }, + raw_request_headers: { "x-api-key": "redacted" }, + }); + + render(); + + const textarea = getRequestTextarea(); + await user.clear(textarea); + await user.type(textarea, '{{"model": "claude-opus-4-8"}'); + + await user.click(getTransformButton()); + + await waitFor(() => expect(transformRequestCallMock).toHaveBeenCalledTimes(1)); + expect(transformRequestCallMock).toHaveBeenCalledWith(ACCESS_TOKEN, { + call_type: "completion", + request_body: { model: "claude-opus-4-8" }, + }); + + const output = await screen.findByText(/api\.anthropic\.com\/v1\/messages/); + expect(output.textContent).toContain("curl -X POST"); + expect(output.textContent).toContain("-H 'x-api-key: redacted'"); + expect(output.textContent).toContain('"model": "claude-opus-4-8"'); + expect(output.textContent).toContain('"max_tokens": 42'); + expect(notify.success).toHaveBeenCalledWith("Request transformed successfully"); + }); + + it("transforms on Cmd/Ctrl + Enter without clicking the button", async () => { + const user = userEvent.setup(); + transformRequestCallMock.mockResolvedValue({ + raw_request_api_base: "https://api.openai.com/v1/chat/completions", + raw_request_body: { model: "gpt-4o" }, + raw_request_headers: {}, + }); + + render(); + + getRequestTextarea().focus(); + await user.keyboard("{Meta>}{Enter}{/Meta}"); + + await waitFor(() => expect(transformRequestCallMock).toHaveBeenCalledTimes(1)); + }); + + it("rejects invalid JSON without calling the backend", async () => { + const user = userEvent.setup(); + + render(); + + const textarea = getRequestTextarea(); + await user.clear(textarea); + await user.type(textarea, "not json"); + await user.click(getTransformButton()); + + await waitFor(() => expect(notify.fromBackend).toHaveBeenCalledWith("Invalid JSON in request body")); + expect(transformRequestCallMock).not.toHaveBeenCalled(); + }); + + it("does not call the backend when there is no access token", async () => { + const user = userEvent.setup(); + + render(); + + await user.click(getTransformButton()); + + await waitFor(() => expect(notify.fromBackend).toHaveBeenCalledWith("No access token found")); + expect(transformRequestCallMock).not.toHaveBeenCalled(); + }); + + it("reports a failed transform and leaves the placeholder curl in place", async () => { + const user = userEvent.setup(); + vi.spyOn(console, "error").mockImplementation(() => {}); + transformRequestCallMock.mockRejectedValue(new Error("boom")); + + render(); + + await user.click(getTransformButton()); + + await waitFor(() => expect(notify.fromBackend).toHaveBeenCalledWith("Failed to transform request")); + expect(screen.getByText(/https:\/\/api\.openai\.com\/v1\/chat\/completions/)).toBeInTheDocument(); + }); + + it("copies the transformed request to the clipboard", async () => { + const user = userEvent.setup(); + const writeText = vi.spyOn(navigator.clipboard, "writeText"); + transformRequestCallMock.mockResolvedValue({ + raw_request_api_base: "https://api.anthropic.com/v1/messages", + raw_request_body: { model: "claude-opus-4-8" }, + raw_request_headers: {}, + }); + + render(); + + await user.click(getTransformButton()); + await screen.findByText(/api\.anthropic\.com\/v1\/messages/); + + await user.click(getCopyButton()); + + expect(writeText).toHaveBeenCalledTimes(1); + expect(writeText.mock.calls[0]?.[0]).toContain("https://api.anthropic.com/v1/messages"); + expect(notify.success).toHaveBeenCalledWith("Copied to clipboard"); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.tsx index 04d1701de3f..0c41547b9b7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/transform-request/TransformRequestPanel.tsx @@ -1,9 +1,12 @@ import React, { useState } from "react"; -import { Button } from "antd"; -import { CopyOutlined } from "@ant-design/icons"; -import { Title } from "@tremor/react"; +import { ArrowRight, Copy } from "lucide-react"; +import { Button } from "@/components/ui/button"; +import { Card, CardContent, CardDescription, CardFooter, CardHeader, CardTitle } from "@/components/ui/card"; +import { Textarea } from "@/components/ui/textarea"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { transformRequestCall } from "@/components/networking"; import NotificationsManager from "@/components/molecules/notifications_manager"; + interface TransformRequestPanelProps { accessToken: string | null; } @@ -128,130 +131,50 @@ ${formattedBody} }; return ( -
- Playground -

See how LiteLLM transforms your request for the specified provider.

-
+
+

Playground

+

+ See how LiteLLM transforms your request for the specified provider. +

+
{/* Original Request Panel */} -
-
-

Original Request

-

- The request you would send to LiteLLM /chat/completions endpoint. -

-
+ + + Original Request + The request you would send to LiteLLM /chat/completions endpoint. + -