diff --git a/.circleci/config.yml b/.circleci/config.yml index 4fcc461a3f6..c9407162649 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -3886,7 +3886,7 @@ jobs: command: | cd ~/project # Check pyproject.toml - CURRENT_VERSION=$(python -c "import toml; print(toml.load('pyproject.toml')['tool']['poetry']['dependencies']['litellm-proxy-extras'].split('\"')[1])") + CURRENT_VERSION=$(python -c "import toml; dep = toml.load('pyproject.toml')['tool']['poetry']['dependencies']['litellm-proxy-extras']; print(dep['version'] if isinstance(dep, dict) else dep)") if [ "$CURRENT_VERSION" != "$NEW_VERSION" ]; then echo "Error: Version in pyproject.toml ($CURRENT_VERSION) doesn't match new version ($NEW_VERSION)" exit 1 diff --git a/.github/scripts/close_duplicate_issues.py b/.github/scripts/close_duplicate_issues.py new file mode 100755 index 00000000000..4e17e1d6d8b --- /dev/null +++ b/.github/scripts/close_duplicate_issues.py @@ -0,0 +1,208 @@ +#!/usr/bin/env python3 +""" +Detect and close duplicate GitHub issues using title similarity. + +Modes: + --scan Compare all open issues against each other (batch) + --issue-number N Check a single issue against older open issues + +Requires the `gh` CLI to be authenticated. +""" + +import argparse +import difflib +import json +import re +import subprocess +import sys + + +def normalize_title(title: str) -> str: + """Strip common prefixes, lowercase, and collapse whitespace.""" + title = re.sub( + r"^\[?(bug|feature request|enhancement|question|docs)[:\]]?\s*", + "", + title, + flags=re.IGNORECASE, + ) + return " ".join(title.lower().split()) + + +def gh(*args: str) -> str: + """Run a gh CLI command and return stdout.""" + result = subprocess.run( + ["gh", *args], + capture_output=True, + text=True, + check=True, + ) + return result.stdout + + +def fetch_open_issues(repo: str | None) -> list[dict]: + """Fetch all open issues (excluding PRs) via gh api --paginate.""" + if repo: + endpoint = f"repos/{repo}/issues?state=open&per_page=100&sort=created&direction=asc" + else: + endpoint = "repos/{owner}/{repo}/issues?state=open&per_page=100&sort=created&direction=asc" + cmd = ["api", "--paginate", endpoint] + + raw = gh(*cmd) + # gh --paginate concatenates JSON arrays, so we may get multiple arrays + issues = [] + for line in raw.strip().splitlines(): + line = line.strip() + if not line: + continue + parsed = json.loads(line) + if isinstance(parsed, list): + issues.extend(parsed) + else: + issues.append(parsed) + + # Filter out pull requests (they also appear in the issues endpoint) + return [i for i in issues if "pull_request" not in i] + + +def close_as_duplicate( + issue_number: int, duplicate_of: int, repo: str | None, dry_run: bool +) -> None: + """Close an issue as duplicate of another, adding a comment and label.""" + repo_args = ["--repo", repo] if repo else [] + + if dry_run: + print(f" [DRY RUN] Would close #{issue_number} as duplicate of #{duplicate_of}") + return + + # Add comment + comment_body = ( + f"Closing as duplicate of #{duplicate_of}.\n\n" + "If you believe this is not a duplicate, please reopen and add context " + "explaining how this differs." + ) + gh("issue", "comment", str(issue_number), "--body", comment_body, *repo_args) + + # Add label + gh("issue", "edit", str(issue_number), "--add-label", "duplicate", *repo_args) + + # Close with not_planned reason + gh( + "api", + f"repos/{repo or '{owner}/{repo}'}/issues/{issue_number}", + "-X", + "PATCH", + "-f", + "state=closed", + "-f", + "state_reason=not_planned", + ) + + print(f" Closed #{issue_number} as duplicate of #{duplicate_of}") + + +def find_duplicate( + issue: dict, candidates: list[dict], threshold: float +) -> dict | None: + """Return the first candidate whose normalized title is above threshold.""" + norm = normalize_title(issue["title"]) + for candidate in candidates: + if candidate["number"] == issue["number"]: + continue + cand_norm = normalize_title(candidate["title"]) + ratio = difflib.SequenceMatcher(None, norm, cand_norm).ratio() + if ratio >= threshold: + return candidate + return None + + +def scan_all(issues: list[dict], threshold: float, repo: str | None, dry_run: bool) -> int: + """Compare every issue against all older issues. Returns count of duplicates found.""" + # Sort oldest first + issues.sort(key=lambda i: i["number"]) + closed_count = 0 + + for idx, issue in enumerate(issues): + older = issues[:idx] + if not older: + continue + dup = find_duplicate(issue, older, threshold) + if dup: + ratio = difflib.SequenceMatcher( + None, + normalize_title(issue["title"]), + normalize_title(dup["title"]), + ).ratio() + print( + f"#{issue['number']}: \"{issue['title']}\"\n" + f" -> duplicate of #{dup['number']}: \"{dup['title']}\" " + f"({ratio:.0%} similar)" + ) + close_as_duplicate(issue["number"], dup["number"], repo, dry_run) + closed_count += 1 + + return closed_count + + +def check_single( + issue_number: int, issues: list[dict], threshold: float, repo: str | None, dry_run: bool +) -> bool: + """Check a single issue against all older open issues. Returns True if duplicate found.""" + target = None + for i in issues: + if i["number"] == issue_number: + target = i + break + + if target is None: + print(f"Issue #{issue_number} not found among open issues.") + return False + + older = [i for i in issues if i["number"] < issue_number] + dup = find_duplicate(target, older, threshold) + if dup: + ratio = difflib.SequenceMatcher( + None, + normalize_title(target["title"]), + normalize_title(dup["title"]), + ).ratio() + print( + f"#{target['number']}: \"{target['title']}\"\n" + f" -> duplicate of #{dup['number']}: \"{dup['title']}\" " + f"({ratio:.0%} similar)" + ) + close_as_duplicate(issue_number, dup["number"], repo, dry_run) + return True + + print(f"#{issue_number}: no duplicate found above threshold {threshold}") + return False + + +def main() -> None: + parser = argparse.ArgumentParser(description="Detect and close duplicate GitHub issues") + mode = parser.add_mutually_exclusive_group(required=True) + mode.add_argument("--scan", action="store_true", help="Scan all open issues") + mode.add_argument("--issue-number", type=int, help="Check a single issue number") + parser.add_argument("--threshold", type=float, default=0.85, help="Similarity threshold (0-1)") + parser.add_argument("--close", action="store_true", help="Actually close duplicates (default is dry-run)") + parser.add_argument("--repo", type=str, help="Repository (owner/repo). Auto-detected if omitted.") + args = parser.parse_args() + + dry_run = not args.close + + if dry_run: + print("=== DRY RUN MODE (pass --close to actually close issues) ===\n") + + print("Fetching open issues...") + issues = fetch_open_issues(args.repo) + print(f"Found {len(issues)} open issues.\n") + + if args.scan: + count = scan_all(issues, args.threshold, args.repo, dry_run) + print(f"\nTotal duplicates {'found' if dry_run else 'closed'}: {count}") + else: + found = check_single(args.issue_number, issues, args.threshold, args.repo, dry_run) + sys.exit(0 if found else 0) # Always exit 0; finding no dup is not an error + + +if __name__ == "__main__": + main() diff --git a/.github/workflows/auto_update_price_and_context_window.yml b/.github/workflows/auto_update_price_and_context_window.yml index e7d65242c19..98b9d868e68 100644 --- a/.github/workflows/auto_update_price_and_context_window.yml +++ b/.github/workflows/auto_update_price_and_context_window.yml @@ -7,6 +7,7 @@ on: jobs: auto_update_price_and_context_window: + if: github.repository == 'BerriAI/litellm' runs-on: ubuntu-latest steps: - uses: actions/checkout@v3 diff --git a/.github/workflows/check_duplicate_issues.yml b/.github/workflows/check_duplicate_issues.yml index 9477dd2f8e2..6d11ce573eb 100644 --- a/.github/workflows/check_duplicate_issues.yml +++ b/.github/workflows/check_duplicate_issues.yml @@ -27,3 +27,26 @@ jobs: {{/issues}} Please review the linked issue(s) to see if they address your concern. If this is not a duplicate, please provide additional context to help us understand the difference. + + - name: Checkout close script + if: github.event.action == 'opened' + uses: actions/checkout@v4 + with: + sparse-checkout: .github/scripts + + - name: Set up Python + if: github.event.action == 'opened' + uses: actions/setup-python@v5 + with: + python-version: "3.11" + + - name: Auto-close if high-confidence duplicate + if: github.event.action == 'opened' + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + run: | + python3 .github/scripts/close_duplicate_issues.py \ + --issue-number ${{ github.event.issue.number }} \ + --repo ${{ github.repository }} \ + --threshold 0.85 \ + --close diff --git a/.github/workflows/scan_duplicate_issues.yml b/.github/workflows/scan_duplicate_issues.yml new file mode 100644 index 00000000000..06e8f453a8c --- /dev/null +++ b/.github/workflows/scan_duplicate_issues.yml @@ -0,0 +1,47 @@ +name: Scan Duplicate Issues (One-Time) + +on: + workflow_dispatch: + inputs: + threshold: + description: "Similarity threshold (0-1)" + required: false + default: "0.85" + close: + description: "Actually close duplicates (false = dry run)" + required: false + type: boolean + default: false + +jobs: + scan: + runs-on: ubuntu-latest + permissions: + issues: write + contents: read + steps: + - name: Checkout scripts + uses: actions/checkout@v4 + with: + sparse-checkout: .github/scripts + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.11" + + - name: Scan for duplicate issues + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + INPUT_THRESHOLD: ${{ inputs.threshold }} + INPUT_CLOSE: ${{ inputs.close }} + run: | + CLOSE_FLAG="" + if [ "$INPUT_CLOSE" = "true" ]; then + CLOSE_FLAG="--close" + fi + python3 .github/scripts/close_duplicate_issues.py \ + --scan \ + --repo ${{ github.repository }} \ + --threshold "$INPUT_THRESHOLD" \ + $CLOSE_FLAG diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 7c5c269f899..48bd21e0e3c 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -74,3 +74,32 @@ jobs: - name: Check import safety run: | poetry run python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) + + secret-scan: + runs-on: ubuntu-latest + timeout-minutes: 5 + permissions: + contents: read + + steps: + - uses: actions/checkout@v4 + with: + fetch-depth: 0 + + - name: Set up Python + uses: actions/setup-python@v4 + with: + python-version: '3.12' + + - name: Run secret scan test + run: | + pip install pytest + pytest tests/litellm/test_no_hardcoded_secrets.py -v + + - name: Run ggshield secret scan + if: ${{ secrets.GITGUARDIAN_API_KEY != '' }} + env: + GITGUARDIAN_API_KEY: ${{ secrets.GITGUARDIAN_API_KEY }} + run: | + pip install ggshield + ggshield secret scan repo . diff --git a/docs/my-website/docs/observability/datadog.md b/docs/my-website/docs/observability/datadog.md index 9385b0020cf..e83cfcbafe0 100644 --- a/docs/my-website/docs/observability/datadog.md +++ b/docs/my-website/docs/observability/datadog.md @@ -7,6 +7,7 @@ import TabItem from '@theme/TabItem'; LiteLLM Supports logging to the following Datdog Integrations: - `datadog` [Datadog Logs](https://docs.datadoghq.com/logs/) - `datadog_llm_observability` [Datadog LLM Observability](https://www.datadoghq.com/product/llm-observability/) +- `datadog_metrics` [Datadog Custom Metrics](#datadog-custom-metrics) - `datadog_cost_management` [Datadog Cloud Cost Management](#datadog-cloud-cost-management) - `ddtrace-run` [Datadog Tracing](#datadog-tracing) @@ -168,6 +169,65 @@ On the Datadog LLM Observability page, you should see that both input messages a +## Datadog Custom Metrics + +| Feature | Details | +|---------|---------| +| **What is logged** | Latency metrics, request counts by status code | +| **Events** | Success + Failure | +| **Product Link** | [Datadog Metrics](https://docs.datadoghq.com/metrics/) | + +Publishes the following metrics to Datadog via the `/api/v2/series` endpoint: + +| Metric | Type | Description | +|--------|------|-------------| +| `litellm.request.total_latency` | Gauge | End-to-end request latency (seconds) | +| `litellm.llm_api.latency` | Gauge | Time spent waiting for the LLM provider response (seconds) | +| `litellm.llm_api.request_count` | Count | Request count, tagged with status code | + +Using `total_latency` and `llm_api.latency`, you can derive **internal latency** = `total_latency - llm_api.latency`. + +All metrics include the following tags: `env`, `service`, `version`, `HOSTNAME`, `POD_NAME`, `provider`, `model_name`, `model_group`, `team`, `status_code`. + +**Step 1**: Create a `config.yaml` file + +```yaml +model_list: + - model_name: gpt-3.5-turbo + litellm_params: + model: gpt-3.5-turbo +litellm_settings: + success_callback: ["datadog_metrics"] + failure_callback: ["datadog_metrics"] +``` + +**Step 2**: Set required env variables + +```shell +DD_API_KEY="your-api-key" +DD_SITE="us5.datadoghq.com" # your datadog site +``` + +**Step 3**: Start the proxy and make a test request + +```shell +litellm --config config.yaml +``` + +```shell +curl --location 'http://0.0.0.0:4000/chat/completions' \ + --header 'Content-Type: application/json' \ + --header 'Authorization: Bearer sk-1234' \ + --data '{ + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "hello"}] +}' +``` + +**Step 4**: View metrics in Datadog Metrics Explorer + +Navigate to **Metrics > Explorer** in Datadog and search for `litellm.request.total_latency`, `litellm.llm_api.latency`, or `litellm.llm_api.request_count`. + ## Datadog Cloud Cost Management | Feature | Details | diff --git a/docs/my-website/docs/proxy/health.md b/docs/my-website/docs/proxy/health.md index 6f98265e40a..2764a6f0d4f 100644 --- a/docs/my-website/docs/proxy/health.md +++ b/docs/my-website/docs/proxy/health.md @@ -330,6 +330,22 @@ model_list: health_check_timeout: 10 # 👈 OVERRIDE HEALTH CHECK TIMEOUT ``` +## Health Check Max Tokens + +By default, health checks use `max_tokens=1` to minimize cost and latency. For wildcard models, the default is `max_tokens=10`. + +You can override this per-model by setting `health_check_max_tokens` in the `model_info` section of your config.yaml. + +```yaml +model_list: + - model_name: openai/gpt-4o + litellm_params: + model: openai/gpt-4o + api_key: os.environ/OPENAI_API_KEY + model_info: + health_check_max_tokens: 5 # 👈 OVERRIDE HEALTH CHECK MAX TOKENS +``` + ## `/health/readiness` Unprotected endpoint for checking if proxy is ready to accept requests diff --git a/docs/my-website/release_notes/v1.81.14.md b/docs/my-website/release_notes/v1.81.14.md index b3a0018b162..3a133f092ae 100644 --- a/docs/my-website/release_notes/v1.81.14.md +++ b/docs/my-website/release_notes/v1.81.14.md @@ -489,6 +489,71 @@ graph LR --- +## Security + +We run [Grype](https://github.com/anchore/grype) and [Trivy](https://github.com/aquasecurity/trivy) security scans on every LiteLLM Docker image. Here's the vulnerability report for this release across all published images: + +### Docker Image Scan Summary + +| Image | Critical | High | Medium | Low | +|-------|----------|------|--------|-----| +| `ghcr.io/berriai/litellm:main-latest` | **0** ✅ | 4 unique CVEs | 4 | 1 | +| `ghcr.io/berriai/litellm-ee:main-latest` | **0** ✅ | 4 unique CVEs | 4 | 1 | +| `ghcr.io/berriai/litellm-non_root:main-latest` | **1** | 11 unique CVEs | 6 | 2 | +| `ghcr.io/berriai/litellm-database:main-latest` | **1** | 7 unique CVEs | 5 | 1 | +| `ghcr.io/berriai/litellm-spend_logs:main-latest` | **4** | 35 matches | 40 | 10 | + +:::note +Vulnerability counts are based on full image scans including build-time tooling. High match counts are often inflated by packages like `minimatch` appearing at multiple versions; the unique CVE counts above reflect the actual distinct vulnerabilities. +::: + +### Critical Severity + +**1. Node.js Critical (non-root, database, spend_logs images):** +Node.js 24.12.0 is used **only** for the Admin UI build and Prisma client generation — it is **not** part of the LiteLLM Python application runtime. + +| Package | Vulnerability | Description | Fix Version | +|---------|---------------|-------------|-------------| +| `node` | CVE-2025-55130 | Node.js critical vulnerability | 20.20.0 | + +**2. OpenSSL & Go Critical (spend_logs image only):** +The `spend_logs` image contains additional vulnerabilities in the underlying Go modules and system libraries. + +| Package | Vulnerability | Description | Fix Version | +|---------|---------------|-------------|-------------| +| `libcrypto3`, `libssl3` | CVE-2025-15467 | OpenSSL critical vulnerability | 3.3.6-r0 | +| `stdlib` (Go) | CVE-2025-68121 | Go standard library critical vulnerability | 1.24.13+ | + +### High Severity + +All high-severity vulnerabilities are in **npm/Node.js build-time dependencies** or system-level libraries — they are **not** in the LiteLLM Python application code. + +**Present in all images:** + +| Package | Vulnerability | Description | Fix Version | +|---------|---------------|-------------|-------------| +| `minimatch` | CVE-2026-26996 | DoS via specially crafted glob patterns | 10.2.1+ / 9.0.6+ | +| `minimatch` | CVE-2026-27903 | DoS due to unbounded recursive backtracking | 10.2.3+ / 9.0.7+ | +| `minimatch` | CVE-2026-27904 | DoS via catastrophic backtracking in glob expressions | 10.2.3+ / 9.0.7+ | +| `tar` | CVE-2026-26960 / GHSA-83g3-92jg-28cx | Arbitrary file read/write via malicious archive hardlinks | 7.5.8 | + +### Medium Severity (all images) + +| Package | Vulnerability | Status | +|---------|---------------|--------| +| `pypdf` 6.7.2 | GHSA-x7hp-r3qg-r3cj | Fix available in 6.7.3 | +| Python 3.13 | CVE-2025-15366, CVE-2025-15367, CVE-2025-12781 | No upstream fix available | + +### Recommendations + +- **LiteLLM Main & EE images** (`litellm:main-latest`, `litellm-ee:main-latest`) have the best security posture with **0 critical vulnerabilities**. +- All HIGH/CRITICAL findings in the main images relate to build-time Node.js/npm tooling, not the Python runtime. +- We are actively monitoring upstream Python and system library fixes for remaining medium-severity vulnerabilities. + +To report a security vulnerability, email support@berri.ai with details and steps to reproduce. + +--- + ## Documentation Updates - Add OpenAI Agents SDK with LiteLLM guide - [PR #21311](https://github.com/BerriAI/litellm/pull/21311) diff --git a/litellm/__init__.py b/litellm/__init__.py index 50fa0e76755..3c61aca3b8e 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -105,6 +105,7 @@ _custom_logger_compatible_callbacks_literal = Literal[ "prometheus", "otel", "datadog", + "datadog_metrics", "datadog_llm_observability", "galileo", "braintrust", diff --git a/litellm/constants.py b/litellm/constants.py index 3d2cebf2224..4c38ecd74b5 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -193,9 +193,9 @@ _DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client fo # Aiohttp connection pooling - prevents memory leaks from unbounded connection growth # Set to 0 for unlimited (not recommended for production) -AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 300)) +AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 1000)) AIOHTTP_CONNECTOR_LIMIT_PER_HOST = int( - os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 50) + os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 500) ) AIOHTTP_KEEPALIVE_TIMEOUT = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120)) AIOHTTP_TTL_DNS_CACHE = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300)) diff --git a/litellm/integrations/callback_configs.json b/litellm/integrations/callback_configs.json index 6a003b8c499..c2b0c4ddce9 100644 --- a/litellm/integrations/callback_configs.json +++ b/litellm/integrations/callback_configs.json @@ -83,6 +83,27 @@ }, "description": "Datadog Logging Integration" }, + { + "id": "datadog_metrics", + "displayName": "Datadog Metrics", + "logo": "datadog.png", + "supports_key_team_logging": false, + "dynamic_params": { + "dd_api_key": { + "type": "password", + "ui_name": "API Key", + "description": "Datadog API key for authentication", + "required": true + }, + "dd_site": { + "type": "text", + "ui_name": "Site", + "description": "Datadog site URL (e.g., us5.datadoghq.com)", + "required": true + } + }, + "description": "Datadog Custom Metrics Integration" + }, { "id": "datadog_cost_management", "displayName": "Datadog Cost Management", @@ -434,4 +455,4 @@ }, "description": "SQS Queue (AWS) Logging Integration" } -] \ No newline at end of file +] diff --git a/litellm/integrations/datadog/datadog_metrics.py b/litellm/integrations/datadog/datadog_metrics.py new file mode 100644 index 00000000000..fcf40701e28 --- /dev/null +++ b/litellm/integrations/datadog/datadog_metrics.py @@ -0,0 +1,286 @@ +import asyncio +import gzip +import os +import time +from datetime import datetime +from typing import List, Optional, Union + +from litellm._logging import verbose_logger +from litellm.integrations.custom_batch_logger import CustomBatchLogger +from litellm.integrations.datadog.datadog_handler import ( + get_datadog_env, + get_datadog_hostname, + get_datadog_pod_name, + get_datadog_service, +) +from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.integrations.base_health_check import IntegrationHealthCheckStatus +from litellm.types.integrations.datadog_metrics import ( + DatadogMetricPoint, + DatadogMetricSeries, + DatadogMetricsPayload, +) +from litellm.types.utils import StandardLoggingPayload + + +class DatadogMetricsLogger(CustomBatchLogger): + def __init__(self, start_periodic_flush: bool = True, **kwargs): + self.dd_api_key = os.getenv("DD_API_KEY") + self.dd_app_key = os.getenv("DD_APP_KEY") + self.dd_site = os.getenv("DD_SITE", "datadoghq.com") + + if not self.dd_api_key: + verbose_logger.warning( + "Datadog Metrics: DD_API_KEY is required. Integration will not work." + ) + + self.upload_url = f"https://api.{self.dd_site}/api/v2/series" + + self.async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.LoggingCallback + ) + + # Initialize lock + self.flush_lock = asyncio.Lock() + + # Only set flush_lock if not already provided by caller + if "flush_lock" not in kwargs: + kwargs["flush_lock"] = self.flush_lock + + # Send metrics more quickly to datadog (every 5 seconds) + if "flush_interval" not in kwargs: + kwargs["flush_interval"] = 5 + + super().__init__(**kwargs) + + # Start periodic flush task only if instructed + if start_periodic_flush: + asyncio.create_task(self.periodic_flush()) + + def _extract_tags( + self, + log: StandardLoggingPayload, + status_code: Optional[Union[str, int]] = None, + ) -> List[str]: + """ + Builds the list of tags for a Datadog metric point + """ + # Base tags + tags = [ + f"env:{get_datadog_env()}", + f"service:{get_datadog_service()}", + f"version:{os.getenv('DD_VERSION', 'unknown')}", + f"HOSTNAME:{get_datadog_hostname()}", + f"POD_NAME:{get_datadog_pod_name()}", + ] + + # Add metric-specific tags + if provider := log.get("custom_llm_provider"): + tags.append(f"provider:{provider}") + + if model := log.get("model"): + tags.append(f"model_name:{model}") + + if model_group := log.get("model_group"): + tags.append(f"model_group:{model_group}") + + if status_code is not None: + tags.append(f"status_code:{status_code}") + + # Extract team tag + metadata = log.get("metadata", {}) or {} + team_tag = ( + metadata.get("user_api_key_team_alias") + or metadata.get("team_alias") # type: ignore + or metadata.get("user_api_key_team_id") + or metadata.get("team_id") # type: ignore + ) + + if team_tag: + tags.append(f"team:{team_tag}") + + return tags + + def _add_metrics_from_log( + self, + log: StandardLoggingPayload, + kwargs: dict, + status_code: Union[str, int] = "200", + ): + """ + Extracts latencies and appends Datadog metric series to the queue + """ + tags = self._extract_tags(log, status_code=status_code) + + # We record metrics with the end_time as the timestamp for the point + end_time_dt = kwargs.get("end_time") or datetime.now() + timestamp = int(end_time_dt.timestamp()) + + # 1. Total Request Latency Metric (End to End) + start_time_dt = kwargs.get("start_time") + if start_time_dt and end_time_dt: + total_duration = (end_time_dt - start_time_dt).total_seconds() + series_total_latency: DatadogMetricSeries = { + "metric": "litellm.request.total_latency", + "type": 3, # gauge + "points": [{"timestamp": timestamp, "value": total_duration}], + "tags": tags, + } + self.log_queue.append(series_total_latency) + + # 2. LLM API Latency Metric (Provider alone) + api_call_start_time = kwargs.get("api_call_start_time") + if api_call_start_time and end_time_dt: + llm_api_duration = (end_time_dt - api_call_start_time).total_seconds() + series_llm_latency: DatadogMetricSeries = { + "metric": "litellm.llm_api.latency", + "type": 3, # gauge + "points": [{"timestamp": timestamp, "value": llm_api_duration}], + "tags": tags, + } + self.log_queue.append(series_llm_latency) + + # 3. Request Count / Status Code + series_count: DatadogMetricSeries = { + "metric": "litellm.llm_api.request_count", + "type": 1, # count + "points": [{"timestamp": timestamp, "value": 1.0}], + "tags": tags, + "interval": self.flush_interval, + } + self.log_queue.append(series_count) + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + try: + standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get( + "standard_logging_object", None + ) + + if standard_logging_object is None: + return + + self._add_metrics_from_log( + log=standard_logging_object, kwargs=kwargs, status_code="200" + ) + + if len(self.log_queue) >= self.batch_size: + await self.flush_queue() + + except Exception as e: + verbose_logger.exception( + f"Datadog Metrics: Error in async_log_success_event: {str(e)}" + ) + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): + try: + standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get( + "standard_logging_object", None + ) + + if standard_logging_object is None: + return + + # Extract status code from error information + status_code = "500" # default + error_information = ( + standard_logging_object.get("error_information", {}) or {} + ) + error_code = error_information.get("error_code") # type: ignore + if error_code is not None: + status_code = str(error_code) + + self._add_metrics_from_log( + log=standard_logging_object, kwargs=kwargs, status_code=status_code + ) + + if len(self.log_queue) >= self.batch_size: + await self.flush_queue() + + except Exception as e: + verbose_logger.exception( + f"Datadog Metrics: Error in async_log_failure_event: {str(e)}" + ) + + async def async_send_batch(self): + if not self.log_queue: + return + + batch = self.log_queue.copy() + payload_data: DatadogMetricsPayload = {"series": batch} + + try: + await self._upload_to_datadog(payload_data) + except Exception as e: + verbose_logger.exception( + f"Datadog Metrics: Error in async_send_batch: {str(e)}" + ) + raise + + async def _upload_to_datadog(self, payload: DatadogMetricsPayload): + if not self.dd_api_key: + return + + headers = { + "Content-Type": "application/json", + "DD-API-KEY": self.dd_api_key, + } + + if self.dd_app_key: + headers["DD-APPLICATION-KEY"] = self.dd_app_key + + json_data = safe_dumps(payload) + compressed_data = gzip.compress(json_data.encode("utf-8")) + headers["Content-Encoding"] = "gzip" + + response = await self.async_client.post( + self.upload_url, content=compressed_data, headers=headers # type: ignore + ) + + response.raise_for_status() + + verbose_logger.debug( + f"Datadog Metrics: Uploaded {len(payload['series'])} metric points. Status: {response.status_code}" + ) + + async def async_health_check(self) -> IntegrationHealthCheckStatus: + """ + Check if the service is healthy + """ + try: + # Send a test metric point to Datadog + test_metric_point: DatadogMetricPoint = { + "timestamp": int(time.time()), + "value": 1.0, + } + test_metric_series: DatadogMetricSeries = { + "metric": "litellm.health_check", + "type": 3, # Gauge + "points": [test_metric_point], + "tags": ["env:health_check"], + } + + payload_data: DatadogMetricsPayload = {"series": [test_metric_series]} + + await self._upload_to_datadog(payload_data) + + return IntegrationHealthCheckStatus( + status="healthy", + error_message=None, + ) + except Exception as e: + return IntegrationHealthCheckStatus( + status="unhealthy", + error_message=str(e), + ) + + async def get_request_response_payload( + self, + request_id: str, + start_time_utc: Optional[datetime], + end_time_utc: Optional[datetime], + ) -> Optional[dict]: + pass diff --git a/litellm/litellm_core_utils/custom_logger_registry.py b/litellm/litellm_core_utils/custom_logger_registry.py index fc73701ea9d..2d483f78613 100644 --- a/litellm/litellm_core_utils/custom_logger_registry.py +++ b/litellm/litellm_core_utils/custom_logger_registry.py @@ -20,6 +20,7 @@ from litellm.integrations.braintrust_logging import BraintrustLogger from litellm.integrations.cloudzero.cloudzero import CloudZeroLogger from litellm.integrations.datadog.datadog import DataDogLogger from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger +from litellm.integrations.datadog.datadog_metrics import DatadogMetricsLogger from litellm.integrations.deepeval import DeepEvalLogger from litellm.integrations.dotprompt import DotpromptManager from litellm.integrations.focus.focus_logger import FocusLogger @@ -66,6 +67,7 @@ class CustomLoggerRegistry: "prometheus": PrometheusLogger, "datadog": DataDogLogger, "datadog_llm_observability": DataDogLLMObsLogger, + "datadog_metrics": DatadogMetricsLogger, "gcs_bucket": GCSBucketLogger, "opik": OpikLogger, "argilla": ArgillaLogger, diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 47a27c8ef5b..9e972f1910b 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -14,7 +14,6 @@ TEST_PDF_URL = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9U class HealthCheckHelpers: - @staticmethod async def ahealth_check_wildcard_models( model: str, @@ -44,7 +43,9 @@ class HealthCheckHelpers: model_params["model"] = cheapest_models[0] model_params["litellm_logging_obj"] = litellm_logging_obj model_params["fallbacks"] = fallback_models - model_params["max_tokens"] = 10 # gpt-5-nano throws errors for max_tokens=1 + model_params["max_tokens"] = model_params.get( + "max_tokens", 10 + ) # gpt-5-nano throws errors for max_tokens=1 await acompletion(**model_params) return {} @@ -130,7 +131,7 @@ class HealthCheckHelpers: Callable, ]: """ - Returns a dictionary of mode handlers for health check calls. + Returns a dictionary of mode handlers for health check calls. Mode Handlers are Callables that need to be run for execution of the health check call. @@ -215,4 +216,4 @@ class HealthCheckHelpers: "document_url": TEST_PDF_URL, }, ), - } \ No newline at end of file + } diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e450b233c7e..5e5a6cea1b2 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -133,6 +133,7 @@ from ..integrations.azure_sentinel.azure_sentinel import AzureSentinelLogger from ..integrations.azure_storage.azure_storage import AzureBlobStorageLogger from ..integrations.custom_prompt_management import CustomPromptManagement from ..integrations.datadog.datadog import DataDogLogger +from ..integrations.datadog.datadog_metrics import DatadogMetricsLogger from ..integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger from ..integrations.dotprompt import DotpromptManager from ..integrations.dynamodb import DyanmoDBLogger @@ -3661,6 +3662,14 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 _datadog_logger = DataDogLogger() _in_memory_loggers.append(_datadog_logger) return _datadog_logger # type: ignore + elif logging_integration == "datadog_metrics": + for callback in _in_memory_loggers: + if isinstance(callback, DatadogMetricsLogger): + return callback # type: ignore + + _datadog_metrics_logger = DatadogMetricsLogger() + _in_memory_loggers.append(_datadog_metrics_logger) + return _datadog_metrics_logger # type: ignore elif logging_integration == "datadog_llm_observability": _datadog_llm_obs_logger = DataDogLLMObsLogger() _in_memory_loggers.append(_datadog_llm_obs_logger) @@ -4268,6 +4277,10 @@ def get_custom_logger_compatible_class( # noqa: PLR0915 for callback in _in_memory_loggers: if isinstance(callback, DataDogLogger): return callback + elif logging_integration == "datadog_metrics": + for callback in _in_memory_loggers: + if isinstance(callback, DatadogMetricsLogger): + return callback elif logging_integration == "datadog_llm_observability": for callback in _in_memory_loggers: if isinstance(callback, DataDogLLMObsLogger): diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index baf274f2c62..3b75a56fcc9 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -1968,22 +1968,24 @@ class CustomStreamWrapper: self.rules.post_call_rules( input=self.response_uptil_now, model=self.model ) - # Store a shallow copy so usage stripping below - # does not mutate the stored chunk. - self.chunks.append(processed_chunk.model_copy()) - # Add mcp_list_tools to first chunk if present if not self.sent_first_chunk: processed_chunk = self._add_mcp_list_tools_to_first_chunk(processed_chunk) self.sent_first_chunk = True - if ( + + _has_usage = ( hasattr(processed_chunk, "usage") and getattr(processed_chunk, "usage", None) is not None - ): + ) + + if _has_usage: + # Store a copy ONLY when usage stripping below will mutate + # the chunk. For non-usage chunks (vast majority), store + # directly to avoid expensive model_copy() per chunk. + self.chunks.append(processed_chunk.model_copy()) + # Strip usage from the outgoing chunk so it's not sent twice # (once in the chunk, once in _hidden_params). - # Create a new object without usage, matching sync behavior. - # The copy in self.chunks retains usage for calculate_total_usage(). obj_dict = processed_chunk.model_dump() if "usage" in obj_dict: del obj_dict["usage"] @@ -1995,6 +1997,9 @@ class CustomStreamWrapper: ) if is_empty: continue + else: + # No usage data — safe to store directly without copying + self.chunks.append(processed_chunk) # add usage as hidden param if self.sent_last_chunk is True and self.stream_options is None: diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 60a93b169c8..ec5b942ec1b 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -68,7 +68,7 @@ def make_sync_call( model_response=model_response, json_mode=json_mode ) else: - decoder = AWSEventStreamDecoder(model=model) + decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode) completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size)) # LOGGING diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index a0f2f65fb7f..d210f294c64 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1779,6 +1779,92 @@ class AmazonConverseConfig(BaseConfig): return content_str, tools, reasoningContentBlocks, citationsContentBlocks + @staticmethod + def _unwrap_bedrock_properties(json_str: str) -> str: + """ + Unwrap Bedrock's response_format JSON structure. + + If the JSON has a single "properties" key, extract its value. + Otherwise, return the original string. + + Args: + json_str: JSON string to unwrap + + Returns: + Unwrapped JSON string or original if unwrapping not needed + """ + try: + response_data = json.loads(json_str) + if ( + isinstance(response_data, dict) + and "properties" in response_data + and len(response_data) == 1 + ): + response_data = response_data["properties"] + return json.dumps(response_data) + except json.JSONDecodeError: + pass + return json_str + + @staticmethod + def _filter_json_mode_tools( + json_mode: Optional[bool], + tools: List[ChatCompletionToolCallChunk], + chat_completion_message: ChatCompletionResponseMessage, + ) -> Optional[List[ChatCompletionToolCallChunk]]: + """ + When json_mode is True, Bedrock may return the internal `json_tool_call` + tool alongside real user-defined tools. This method handles 3 scenarios: + + 1. Only json_tool_call present -> convert to text content, return None + 2. Mixed json_tool_call + real -> filter out json_tool_call, return real tools + 3. No json_tool_call / no json_mode -> return tools as-is + """ + if not json_mode or not tools: + return tools if tools else None + + json_tool_indices = [ + i + for i, t in enumerate(tools) + if t["function"].get("name") == RESPONSE_FORMAT_TOOL_NAME + ] + + if not json_tool_indices: + # No json_tool_call found, return tools unchanged + return tools + + if len(json_tool_indices) == len(tools): + # All tools are json_tool_call — convert first one to content + verbose_logger.debug( + "Processing JSON tool call response for response_format" + ) + json_mode_content_str: Optional[str] = tools[0]["function"].get( + "arguments" + ) + if json_mode_content_str is not None: + json_mode_content_str = AmazonConverseConfig._unwrap_bedrock_properties( + json_mode_content_str + ) + chat_completion_message["content"] = json_mode_content_str + return None + + # Mixed: filter out json_tool_call, keep real tools. + # Preserve the json_tool_call content as message text so the structured + # output from response_format is not silently lost. + first_idx = json_tool_indices[0] + json_mode_args = tools[first_idx]["function"].get("arguments") + if json_mode_args is not None: + json_mode_args = AmazonConverseConfig._unwrap_bedrock_properties( + json_mode_args + ) + existing = chat_completion_message.get("content") or "" + chat_completion_message["content"] = ( + existing + json_mode_args if existing else json_mode_args + ) + + real_tools = [t for i, t in enumerate(tools) if i not in json_tool_indices] + return real_tools if real_tools else None + def _transform_response( # noqa: PLR0915 self, model: str, @@ -1801,7 +1887,7 @@ class AmazonConverseConfig(BaseConfig): additional_args={"complete_input_dict": data}, ) - json_mode: Optional[bool] = optional_params.pop("json_mode", None) + json_mode: Optional[bool] = optional_params.get("json_mode", None) ## RESPONSE OBJECT try: completion_response = ConverseResponseBlock(**response.json()) # type: ignore @@ -1885,37 +1971,13 @@ class AmazonConverseConfig(BaseConfig): self._transform_thinking_blocks(reasoningContentBlocks) ) chat_completion_message["content"] = content_str - if ( - json_mode is True - and tools is not None - and len(tools) == 1 - and tools[0]["function"].get("name") == RESPONSE_FORMAT_TOOL_NAME - ): - verbose_logger.debug( - "Processing JSON tool call response for response_format" - ) - json_mode_content_str: Optional[str] = tools[0]["function"].get("arguments") - if json_mode_content_str is not None: - # Bedrock returns the response wrapped in a "properties" object - # We need to extract the actual content from this wrapper - try: - response_data = json.loads(json_mode_content_str) - - # If Bedrock wrapped the response in "properties", extract the content - if ( - isinstance(response_data, dict) - and "properties" in response_data - and len(response_data) == 1 - ): - response_data = response_data["properties"] - json_mode_content_str = json.dumps(response_data) - except json.JSONDecodeError: - # If parsing fails, use the original response - pass - - chat_completion_message["content"] = json_mode_content_str - elif tools: - chat_completion_message["tool_calls"] = tools + filtered_tools = self._filter_json_mode_tools( + json_mode=json_mode, + tools=tools, + chat_completion_message=chat_completion_message, + ) + if filtered_tools: + chat_completion_message["tool_calls"] = filtered_tools ## CALCULATING USAGE - bedrock returns usage in the headers usage = self._transform_usage(completion_response["usage"]) diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 1c58a11eebe..88f7341ed08 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -22,6 +22,7 @@ import litellm from litellm import verbose_logger from litellm._uuid import uuid from litellm.caching.caching import InMemoryCache +from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.logging_utils import track_llm_api_timing @@ -252,7 +253,7 @@ async def make_call( response.aiter_bytes(chunk_size=stream_chunk_size) ) else: - decoder = AWSEventStreamDecoder(model=model) + decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode) completion_stream = decoder.aiter_bytes( response.aiter_bytes(chunk_size=stream_chunk_size) ) @@ -346,7 +347,7 @@ def make_sync_call( response.iter_bytes(chunk_size=stream_chunk_size) ) else: - decoder = AWSEventStreamDecoder(model=model) + decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode) completion_stream = decoder.iter_bytes( response.iter_bytes(chunk_size=stream_chunk_size) ) @@ -1282,7 +1283,7 @@ def get_response_stream_shape(): class AWSEventStreamDecoder: - def __init__(self, model: str) -> None: + def __init__(self, model: str, json_mode: Optional[bool] = False) -> None: from botocore.parsers import EventStreamJSONParser self.model = model @@ -1290,6 +1291,8 @@ class AWSEventStreamDecoder: self.content_blocks: List[ContentBlockDeltaEvent] = [] self.tool_calls_index: Optional[int] = None self.response_id: Optional[str] = None + self.json_mode = json_mode + self._current_tool_name: Optional[str] = None def check_empty_tool_call_args(self) -> bool: """ @@ -1391,6 +1394,16 @@ class AWSEventStreamDecoder: response_tool_name = get_bedrock_tool_name( response_tool_name=_response_tool_name ) + self._current_tool_name = response_tool_name + + # When json_mode is True, suppress the internal json_tool_call + # and convert its content to text in delta events instead + if ( + self.json_mode is True + and response_tool_name == RESPONSE_FORMAT_TOOL_NAME + ): + return tool_use, provider_specific_fields, thinking_blocks + self.tool_calls_index = ( 0 if self.tool_calls_index is None else self.tool_calls_index + 1 ) @@ -1445,19 +1458,27 @@ class AWSEventStreamDecoder: if "text" in delta_obj: text = delta_obj["text"] elif "toolUse" in delta_obj: - tool_use = { - "id": None, - "type": "function", - "function": { - "name": None, - "arguments": delta_obj["toolUse"]["input"], - }, - "index": ( - self.tool_calls_index - if self.tool_calls_index is not None - else index - ), - } + # When json_mode is True and this is the internal json_tool_call, + # convert tool input to text content instead of tool call arguments + if ( + self.json_mode is True + and self._current_tool_name == RESPONSE_FORMAT_TOOL_NAME + ): + text = delta_obj["toolUse"]["input"] + else: + tool_use = { + "id": None, + "type": "function", + "function": { + "name": None, + "arguments": delta_obj["toolUse"]["input"], + }, + "index": ( + self.tool_calls_index + if self.tool_calls_index is not None + else index + ), + } elif "reasoningContent" in delta_obj: provider_specific_fields = { "reasoningContent": delta_obj["reasoningContent"], @@ -1494,6 +1515,17 @@ class AWSEventStreamDecoder: ) -> Optional[ChatCompletionToolCallChunk]: """Handle stop/contentBlockIndex event in converse chunk parsing.""" tool_use: Optional[ChatCompletionToolCallChunk] = None + + # If the ending block was the internal json_tool_call, skip emitting + # the empty-args tool chunk and reset tracking state + if ( + self.json_mode is True + and self._current_tool_name == RESPONSE_FORMAT_TOOL_NAME + ): + self._current_tool_name = None + return tool_use + + self._current_tool_name = None is_empty = self.check_empty_tool_call_args() if is_empty: tool_use = { diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 9921b74b561..553ba4d6c49 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -8,6 +8,7 @@ JWT token must have 'litellm_proxy_admin' in scope. import fnmatch import os +import re from typing import Any, List, Literal, Optional, Set, Tuple, cast from cryptography import x509 @@ -235,7 +236,17 @@ class JWTHandler: return self.litellm_jwtauth.team_id_default else: return default_value - # At this point, team_id is not the sentinel, so it should be a string + # AAD and other IdPs often send roles/groups as a list of strings. + # team_id_jwt_field is singular, so take the first element when a list + # is returned. This avoids "unhashable type: 'list'" errors downstream. + if isinstance(team_id, list): + if not team_id: + return default_value + verbose_proxy_logger.debug( + f"JWT Auth: team_id_jwt_field '{self.litellm_jwtauth.team_id_jwt_field}' " + f"returned a list {team_id}; using first element '{team_id[0]}' automatically." + ) + team_id = team_id[0] return team_id # type: ignore[return-value] elif self.litellm_jwtauth.team_id_default is not None: team_id = self.litellm_jwtauth.team_id_default @@ -453,6 +464,52 @@ class JWTHandler: scopes = [] return scopes + async def _resolve_jwks_url(self, url: str) -> str: + """ + If url points to an OIDC discovery document (*.well-known/openid-configuration), + fetch it and return the jwks_uri contained within. Otherwise return url unchanged. + This lets JWT_PUBLIC_KEY_URL be set to a well-known discovery endpoint instead of + requiring operators to manually find the JWKS URL. + """ + if ".well-known/openid-configuration" not in url: + return url + + cache_key = f"litellm_oidc_discovery_{url}" + cached_jwks_uri = await self.user_api_key_cache.async_get_cache(cache_key) + if cached_jwks_uri is not None: + return cached_jwks_uri + + verbose_proxy_logger.debug( + f"JWT Auth: Fetching OIDC discovery document from {url}" + ) + response = await self.http_handler.get(url) + if response.status_code != 200: + raise Exception( + f"JWT Auth: OIDC discovery endpoint {url} returned status {response.status_code}: {response.text}" + ) + try: + discovery = response.json() + except Exception as e: + raise Exception( + f"JWT Auth: Failed to parse OIDC discovery document at {url}: {e}" + ) + + jwks_uri = discovery.get("jwks_uri") + if not jwks_uri: + raise Exception( + f"JWT Auth: OIDC discovery document at {url} does not contain a 'jwks_uri' field." + ) + + verbose_proxy_logger.debug( + f"JWT Auth: Resolved OIDC discovery {url} -> jwks_uri={jwks_uri}" + ) + await self.user_api_key_cache.async_set_cache( + key=cache_key, + value=jwks_uri, + ttl=self.litellm_jwtauth.public_key_ttl, + ) + return jwks_uri + async def get_public_key(self, kid: Optional[str]) -> dict: keys_url = os.getenv("JWT_PUBLIC_KEY_URL") @@ -462,6 +519,7 @@ class JWTHandler: keys_url_list = [url.strip() for url in keys_url.split(",")] for key_url in keys_url_list: + key_url = await self._resolve_jwks_url(key_url) cache_key = f"litellm_jwt_auth_keys_{key_url}" cached_keys = await self.user_api_key_cache.async_get_cache(cache_key) @@ -913,8 +971,30 @@ class JWTAuthManager: if jwt_handler.is_required_team_id() is True: team_id_field = jwt_handler.litellm_jwtauth.team_id_jwt_field team_alias_field = jwt_handler.litellm_jwtauth.team_alias_jwt_field + hint = "" + if team_id_field: + # "roles.0" — dot-notation numeric indexing is not supported + if "." in team_id_field: + parts = team_id_field.rsplit(".", 1) + if parts[-1].isdigit(): + base_field = parts[0] + hint = ( + f" Hint: dot-notation array indexing (e.g. '{team_id_field}') is not " + f"supported. Use '{base_field}' instead — LiteLLM automatically " + f"uses the first element when the field value is a list." + ) + # "roles[0]" — bracket-notation indexing is also not supported in get_nested_value + elif "[" in team_id_field and team_id_field.endswith("]"): + m = re.match(r"^(\w+)\[(\d+)\]$", team_id_field) + if m: + base_field = m.group(1) + hint = ( + f" Hint: array indexing (e.g. '{team_id_field}') is not supported " + f"in team_id_jwt_field. Use '{base_field}' instead — LiteLLM " + f"automatically uses the first element when the field value is a list." + ) raise Exception( - f"No team found in token. Checked team_id field '{team_id_field}' and team_alias field '{team_alias_field}'" + f"No team found in token. Checked team_id field '{team_id_field}' and team_alias field '{team_alias_field}'.{hint}" ) return individual_team_id, team_object diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 62ca6dc2ae2..9ecae363ed7 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -390,9 +390,7 @@ def get_logging_caching_headers(request_data: Dict) -> Optional[Dict]: ) if "applied_policies" in _metadata: - headers["x-litellm-applied-policies"] = ",".join( - _metadata["applied_policies"] - ) + headers["x-litellm-applied-policies"] = ",".join(_metadata["applied_policies"]) if "policy_sources" in _metadata: sources = _metadata["policy_sources"] @@ -449,9 +447,7 @@ def add_policy_to_applied_policies_header( request_data["metadata"] = _metadata -def add_policy_sources_to_metadata( - request_data: Dict, policy_sources: Dict[str, str] -): +def add_policy_sources_to_metadata(request_data: Dict, policy_sources: Dict[str, str]): """ Store policy match reasons in metadata for x-litellm-policy-sources header. diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 34fbf47253b..adfe69315c6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -427,6 +427,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): analyze_results: Any, output_parse_pii: bool, masked_entity_count: Dict[str, int], + request_data: Optional[Dict] = None, ) -> str: """ Send analysis results to the Presidio anonymizer endpoint to get redacted text @@ -482,10 +483,22 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if item["operator"] == "replace" and output_parse_pii is True: # check if token in dict # if exists, add a uuid to the replacement token for swapping back to the original text in llm response output parsing - if replacement in self.pii_tokens: - replacement = replacement + str(uuid.uuid4()) + if request_data is None: + verbose_proxy_logger.warning( + "Presidio anonymize_text called without request_data — " + "PII tokens cannot be stored per-request. " + "This may indicate a missing caller update." + ) + request_data = {} + if "pii_tokens" not in request_data: + request_data["pii_tokens"] = {} + pii_tokens = request_data["pii_tokens"] - self.pii_tokens[replacement] = new_text[ + # Always append a UUID to ensure the replacement token is unique to this request and session. + # This prevents collisions where the LLM might hallucinate a generic token like [PHONE_NUMBER]. + replacement = f"{replacement}_{str(uuid.uuid4())[:12]}" + + pii_tokens[replacement] = new_text[ start:end ] # get text it'll replace @@ -495,7 +508,13 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): masked_entity_count[entity_type] = ( masked_entity_count.get(entity_type, 0) + 1 ) - return redacted_text["text"] + # When output_parse_pii is True, new_text contains UUID-suffixed + # tokens that match the keys in pii_tokens. Returning + # redacted_text["text"] (Presidio's original output) would send + # un-suffixed tokens to the LLM, making unmasking impossible. + # When output_parse_pii is False, new_text == redacted_text["text"] + # because no UUID suffix is appended. + return new_text else: raise Exception("Invalid anonymizer response: received None") except Exception as e: @@ -525,10 +544,17 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return analyze_results filtered_results: List[PresidioAnalyzeResponseItem] = [] + deny_list_strings = [ + x.value if hasattr(x, "value") else str(x) + for x in self.presidio_entities_deny_list + ] for item in analyze_results: entity_type = item.get("entity_type") - if entity_type and entity_type in self.presidio_entities_deny_list: + str_entity_type = str( + entity_type.value if hasattr(entity_type, "value") else entity_type + ) + if entity_type and str_entity_type in deny_list_strings: continue if self.presidio_score_thresholds: @@ -621,6 +647,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): analyze_results=analyze_results, output_parse_pii=output_parse_pii, masked_entity_count=masked_entity_count, + request_data=request_data, ) return anonymized_text return redacted_text["text"] @@ -866,14 +893,129 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if isinstance(response, ModelResponse) and not isinstance( response.choices[0], StreamingChoices ): # /chat/completions requests - if isinstance(response.choices[0].message.content, str): - verbose_proxy_logger.debug( - f"self.pii_tokens: {self.pii_tokens}; initial response: {response.choices[0].message.content}" - ) - for key, value in self.pii_tokens.items(): - response.choices[0].message.content = response.choices[ - 0 - ].message.content.replace(key, value) + await self._process_response_for_pii( + response=response, + request_data=data, + mode="unmask", + ) + return response + + @staticmethod + def _unmask_pii_text(text: str, pii_tokens: Dict[str, str]) -> str: + """ + Replace PII tokens in *text* with their original values. + + Includes a fallback for tokens that were truncated by ``max_tokens``: + if the *end* of ``text`` matches the *beginning* of a token and the + overlap is long enough, the truncated suffix is replaced with the + original value. The minimum overlap length is + ``min(20, len(token) // 2)`` to reduce the risk of false positives + when multiple tokens share a common prefix. + """ + for token, original_text in pii_tokens.items(): + if token in text: + text = text.replace(token, original_text) + else: + # FALLBACK: Handle truncated tokens (token cut off by max_tokens) + # Only check at the very end of the text. + min_overlap = min(20, len(token) // 2) + for i in range(max(0, len(text) - len(token)), len(text)): + sub = text[i:] + if token.startswith(sub) and len(sub) >= min_overlap: + text = text[:i] + original_text + break + return text + + async def _process_response_for_pii( + self, + response: ModelResponse, + request_data: dict, + mode: Literal["mask", "unmask"], + ) -> ModelResponse: + """ + Helper to recursively process a ModelResponse for PII. + Handles all choices and tool calls. + """ + pii_tokens = request_data.get("pii_tokens", {}) if request_data else {} + if not pii_tokens and mode == "unmask": + verbose_proxy_logger.debug( + "No pii_tokens found in request_data — nothing to unmask" + ) + presidio_config = self.get_presidio_settings_from_request_data( + request_data or {} + ) + + for choice in response.choices: + message = getattr(choice, "message", None) + if message is None: + continue + + # 1. Process content + content = getattr(message, "content", None) + if isinstance(content, str): + if mode == "unmask": + message.content = self._unmask_pii_text(content, pii_tokens) + elif mode == "mask": + message.content = await self.check_pii( + text=content, + output_parse_pii=False, + presidio_config=presidio_config, + request_data=request_data, + ) + elif isinstance(content, list): + for item in content: + if not isinstance(item, dict): + continue + text_value = item.get("text") + if text_value is None: + continue + if mode == "unmask": + item["text"] = self._unmask_pii_text(text_value, pii_tokens) + elif mode == "mask": + item["text"] = await self.check_pii( + text=text_value, + output_parse_pii=False, + presidio_config=presidio_config, + request_data=request_data, + ) + + # 2. Process tool calls + tool_calls = getattr(message, "tool_calls", None) + if tool_calls: + for tool_call in tool_calls: + function = getattr(tool_call, "function", None) + if function and hasattr(function, "arguments"): + args = function.arguments + if isinstance(args, str): + if mode == "unmask": + function.arguments = self._unmask_pii_text( + args, pii_tokens + ) + elif mode == "mask": + function.arguments = await self.check_pii( + text=args, + output_parse_pii=False, + presidio_config=presidio_config, + request_data=request_data, + ) + + # 3. Process legacy function calls + function_call = getattr(message, "function_call", None) + if function_call and hasattr(function_call, "arguments"): + args = function_call.arguments + if isinstance(args, str): + if mode == "unmask": + function_call.arguments = self._unmask_pii_text( + args, pii_tokens + ) + elif mode == "mask": + function_call.arguments = await self.check_pii( + text=args, + output_parse_pii=False, + presidio_config=presidio_config, + request_data=request_data, + ) + return response async def _mask_output_response( @@ -891,38 +1033,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if response.choices and isinstance(response.choices[0], StreamingChoices): return response - presidio_config = self.get_presidio_settings_from_request_data( - request_data or {} + await self._process_response_for_pii( + response=response, + request_data=request_data, + mode="mask", ) - - for choice in response.choices: - # Type narrowing: StreamingChoices doesn't have .message attribute - if not hasattr(choice, "message"): - continue - content = getattr(choice.message, "content", None) # type: ignore - if content is None: - continue - if isinstance(content, str): - choice.message.content = await self.check_pii( # type: ignore - text=content, - output_parse_pii=False, - presidio_config=presidio_config, - request_data=request_data, - ) - elif isinstance(content, list): - for item in content: - if not isinstance(item, dict): - continue - text_value = item.get("text") - if text_value is None: - continue - item["text"] = await self.check_pii( - text=text_value, - output_parse_pii=False, - presidio_config=presidio_config, - request_data=request_data, - ) - return response async def async_post_call_streaming_iterator_hook( @@ -934,7 +1049,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): """ Process streaming response chunks to unmask PII tokens when needed. """ - from litellm.llms.base_llm.base_model_iterator import MockResponseIterator + from litellm.llms.base_llm.base_model_iterator import ( + convert_model_response_to_streaming, + ) from litellm.main import stream_chunk_builder from litellm.types.utils import ModelResponse @@ -959,45 +1076,16 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return # Apply Presidio masking on the assembled response - presidio_config = self.get_presidio_settings_from_request_data( - request_data or {} - ) - - content_to_mask = "" - if ( - hasattr(assembled_model_response, "choices") - and len(assembled_model_response.choices) > 0 - ): - if hasattr( - assembled_model_response.choices[0], "message" - ) and hasattr( - assembled_model_response.choices[0].message, "content" - ): - content_to_mask = ( - assembled_model_response.choices[0].message.content or "" - ) - - masked_content = await self.check_pii( - text=content_to_mask, - output_parse_pii=False, - presidio_config=presidio_config, + await self._process_response_for_pii( + response=assembled_model_response, request_data=request_data, + mode="mask", ) - if ( - hasattr(assembled_model_response, "choices") - and len(assembled_model_response.choices) > 0 - ): - if hasattr(assembled_model_response.choices[0], "message"): - assembled_model_response.choices[ - 0 - ].message.content = masked_content - - mock_response = MockResponseIterator( - model_response=assembled_model_response + mock_response_stream = convert_model_response_to_streaming( + assembled_model_response ) - async for chunk in mock_response: - yield chunk + yield mock_response_stream return except Exception as e: @@ -1011,7 +1099,12 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): return # --- PII unmasking path (output_parse_pii=True) --- - if not (self.output_parse_pii and self.pii_tokens): + pii_tokens = request_data.get("pii_tokens", {}) if request_data else {} + if not pii_tokens and request_data: + verbose_proxy_logger.debug( + "No pii_tokens in request_data for streaming unmask path" + ) + if not (self.output_parse_pii and pii_tokens): async for chunk in response: yield chunk return @@ -1034,20 +1127,27 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): yield chunk return - # Apply PII unmasking to assembled content - for choice in assembled_model_response.choices: - if hasattr(choice, "message") and hasattr(choice.message, "content"): - content = choice.message.content - if isinstance(content, str): - for token, original_text in self.pii_tokens.items(): - content = content.replace(token, original_text) - choice.message.content = content + # --- PRESERVE USAGE METADATA --- + # stream_chunk_builder might miss usage if it's only in the last chunk + if ( + not hasattr(assembled_model_response, "usage") + or not assembled_model_response.usage + ) and remaining_chunks: + last_chunk = remaining_chunks[-1] + if hasattr(last_chunk, "usage") and last_chunk.usage: + assembled_model_response.usage = last_chunk.usage - mock_response = MockResponseIterator( - model_response=assembled_model_response + # Apply PII unmasking to assembled content (unmasking tokens back to original text) + await self._process_response_for_pii( + response=assembled_model_response, + request_data=request_data, + mode="unmask", ) - async for chunk in mock_response: - yield chunk + + mock_response_stream = convert_model_response_to_streaming( + assembled_model_response + ) + yield mock_response_stream except Exception as e: verbose_proxy_logger.error(f"Error in PII streaming processing: {str(e)}") diff --git a/litellm/proxy/guardrails/guardrail_initializers.py b/litellm/proxy/guardrails/guardrail_initializers.py index 639aebf45c9..109f2237165 100644 --- a/litellm/proxy/guardrails/guardrail_initializers.py +++ b/litellm/proxy/guardrails/guardrail_initializers.py @@ -93,6 +93,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail): presidio_analyzer_api_base=litellm_params.presidio_analyzer_api_base, presidio_anonymizer_api_base=litellm_params.presidio_anonymizer_api_base, presidio_language=litellm_params.presidio_language, + presidio_entities_deny_list=litellm_params.presidio_entities_deny_list, apply_to_output=False, ) params.update(overrides) diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index d228bdb2129..a8d0e3e9af2 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -234,6 +234,14 @@ def _update_litellm_params_for_health_check( - for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID """ litellm_params["messages"] = _get_random_llm_message() + _health_check_max_tokens = model_info.get("health_check_max_tokens", None) + if _health_check_max_tokens is not None: + litellm_params["max_tokens"] = _health_check_max_tokens + elif "*" not in ( + model_info.get("health_check_model") or litellm_params.get("model") or "" + ): + litellm_params["max_tokens"] = 1 + _health_check_model = model_info.get("health_check_model", None) if _health_check_model is not None: litellm_params["model"] = _health_check_model @@ -321,7 +329,9 @@ async def perform_health_check( # Filter by model_id first so a single deployment is checked when id is specified if model_id is not None: - _by_id = [x for x in model_list if (x.get("model_info") or {}).get("id") == model_id] + _by_id = [ + x for x in model_list if (x.get("model_info") or {}).get("id") == model_id + ] if _by_id: model_list = _by_id elif model is not None: diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 95b1836d8a9..5a4a8eae385 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -230,6 +230,7 @@ async def health_services_endpoint( # noqa: PLR0915 "custom_callback_api", "langsmith", "datadog", + "datadog_metrics", "datadog_llm_observability", "generic_api", "arize", @@ -284,6 +285,30 @@ async def health_services_endpoint( # noqa: PLR0915 else "Datadog is healthy" ), } + elif service == "datadog_metrics": + from litellm.integrations.datadog.datadog_metrics import ( + DatadogMetricsLogger, + ) + from litellm.litellm_core_utils.litellm_logging import ( + get_custom_logger_compatible_class, + ) + + datadog_metrics_logger = get_custom_logger_compatible_class( + "datadog_metrics" + ) + if datadog_metrics_logger is None: + datadog_metrics_logger = DatadogMetricsLogger( + start_periodic_flush=False + ) + response = await datadog_metrics_logger.async_health_check() + return { + "status": response["status"], + "message": ( + response["error_message"] + if response["status"] == "unhealthy" + else "Datadog Metrics is healthy" + ), + } elif service == "arize": from litellm.integrations.arize.arize import ArizeLogger diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index c07f30f8646..95c9c806120 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -153,7 +153,7 @@ class KeyManagementEventHooks: new_secret_name = ( response.key_alias or data.key_alias - or f"virtual-key-{response.token_id}" + or initial_secret_name ) verbose_proxy_logger.info( "Updating secret in secret manager: secret_name=%s", diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 5b56133f1ce..4d44006369b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -14,6 +14,7 @@ import copy import inspect import json import os +import re import secrets import traceback from datetime import datetime, timedelta, timezone @@ -125,7 +126,7 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime: def _set_key_rotation_fields( - data: dict, auto_rotate: bool, rotation_interval: Optional[str] + data: dict, auto_rotate: bool, rotation_interval: Optional[str], existing_key_alias: Optional[str] = None ) -> None: """ Helper function to set rotation fields in key data if auto_rotate is enabled. @@ -134,8 +135,21 @@ def _set_key_rotation_fields( data: Dictionary to update with rotation fields auto_rotate: Whether auto rotation is enabled rotation_interval: The rotation interval string (required if auto_rotate is True) + existing_key_alias: The existing key alias from the database (if any) """ if auto_rotate and rotation_interval: + if ( + litellm._key_management_settings is not None + and litellm._key_management_settings.store_virtual_keys is True + and data.get("key_alias") is None + and existing_key_alias is None + ): + raise ProxyException( + message="key_alias is required when auto_rotate=True and store_virtual_keys is enabled. This ensures stable secret naming during rotation.", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=400, + ) data.update( { "auto_rotate": auto_rotate, @@ -625,6 +639,8 @@ async def _common_key_generation_helper( # noqa: PLR0915 prisma_client=prisma_client, ) + _validate_key_alias_format(key_alias=data_json.get("key_alias", None)) + await _enforce_unique_key_alias( key_alias=data_json.get("key_alias", None), prisma_client=prisma_client, @@ -1931,6 +1947,8 @@ async def update_key_fn( data=data, existing_key_row=existing_key_row ) + _validate_key_alias_format(key_alias=non_default_values.get("key_alias", None)) + await _enforce_unique_key_alias( key_alias=non_default_values.get("key_alias", None), prisma_client=prisma_client, @@ -1942,6 +1960,7 @@ async def update_key_fn( non_default_values, non_default_values.get("auto_rotate", False), non_default_values.get("rotation_interval"), + existing_key_alias=existing_key_row.key_alias, ) _data = {**non_default_values, "token": key} @@ -3378,6 +3397,7 @@ async def _execute_virtual_key_regeneration( non_default_values = await prepare_key_update_data( data=data, existing_key_row=key_in_db ) + _validate_key_alias_format(key_alias=non_default_values.get("key_alias")) verbose_proxy_logger.debug("non_default_values: %s", non_default_values) update_data.update(non_default_values) update_data = prisma_client.jsonify_object(data=update_data) @@ -3983,6 +4003,10 @@ async def list_keys( status: Optional[str] = Query( None, description="Filter by status (e.g. 'deleted')" ), + project_id: Optional[str] = Query(None, description="Filter keys by project ID"), + access_group_id: Optional[str] = Query( + None, description="Filter keys by access group ID" + ), ) -> KeyListResponseObject: """ List all keys for a given user / team / organization. @@ -4076,6 +4100,8 @@ async def list_keys( sort_order=sort_order, expand=expand, status=status, + project_id=project_id, + access_group_id=access_group_id, ) verbose_proxy_logger.debug("Successfully prepared response") @@ -4252,6 +4278,8 @@ def _build_key_filter_conditions( admin_team_ids: Optional[List[str]], member_team_ids: Optional[List[str]] = None, include_created_by_keys: bool = False, + project_id: Optional[str] = None, + access_group_id: Optional[str] = None, ) -> Dict[str, Union[str, Dict[str, Any], List[Dict[str, Any]]]]: """Build filter conditions for key listing. @@ -4343,6 +4371,13 @@ def _build_key_filter_conditions( elif len(or_conditions) == 1: where.update(or_conditions[0]) + # Apply project_id and access_group_id as global AND filters so they + # narrow results across all visibility conditions (own keys, team keys, etc.) + if project_id: + where = {"AND": [where, {"project_id": project_id}]} + if access_group_id: + where = {"AND": [where, {"access_group_ids": {"hasSome": [access_group_id]}}]} + verbose_proxy_logger.debug(f"Filter conditions: {where}") return where @@ -4369,6 +4404,8 @@ async def _list_key_helper( sort_order: str = "desc", expand: Optional[List[str]] = None, status: Optional[str] = None, + project_id: Optional[str] = None, + access_group_id: Optional[str] = None, ) -> KeyListResponseObject: """ Helper function to list keys @@ -4402,6 +4439,8 @@ async def _list_key_helper( admin_team_ids=admin_team_ids, member_team_ids=member_team_ids, include_created_by_keys=include_created_by_keys, + project_id=project_id, + access_group_id=access_group_id, ) # Calculate skip for pagination @@ -4952,6 +4991,31 @@ async def test_key_logging( ) +_KEY_ALIAS_PATTERN = re.compile(r"^[a-zA-Z0-9][a-zA-Z0-9_\-/\.]{0,253}[a-zA-Z0-9]$") + + +def _validate_key_alias_format(key_alias: Optional[str]) -> None: + """ + Validate the format of the key_alias. + + Rules: + - None is OK (no alias). + - Otherwise must be 2–255 chars + - start/end with alphanumeric + - only allow a-zA-Z0-9_-/. + """ + if key_alias is None: + return + + if not _KEY_ALIAS_PATTERN.match(key_alias): + raise ProxyException( + message="Invalid key_alias format. Must be 2-255 characters, start/end with alphanumeric, and only contain a-zA-Z0-9_-/.", + type=ProxyErrorTypes.bad_request_error, + param="key_alias", + code=400, + ) + + async def _enforce_unique_key_alias( key_alias: Optional[str], prisma_client: Any, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 48025863641..bc2728c2203 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -2166,22 +2166,24 @@ async def _run_background_health_check(): "Error in shared health check, falling back to direct health check: %s", str(e), ) - healthy_endpoints, unhealthy_endpoints = ( - await _run_direct_health_check_with_instrumentation( - _llm_model_list, - health_check_details, - health_check_concurrency, - instrumentation_context, - ) - ) - else: - healthy_endpoints, unhealthy_endpoints = ( - await _run_direct_health_check_with_instrumentation( + ( + healthy_endpoints, + unhealthy_endpoints, + ) = await _run_direct_health_check_with_instrumentation( _llm_model_list, health_check_details, health_check_concurrency, instrumentation_context, ) + else: + ( + healthy_endpoints, + unhealthy_endpoints, + ) = await _run_direct_health_check_with_instrumentation( + _llm_model_list, + health_check_details, + health_check_concurrency, + instrumentation_context, ) # Update the global variable with the health check results @@ -3506,7 +3508,15 @@ class ProxyConfig: combined_id_list.append(model_info.id) ## CONFIG MODELS ## - config = await self.get_config(config_file_path=user_config_file_path) + try: + config = await self.get_config(config_file_path=user_config_file_path) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to load config in _delete_deployment: %s. " + "Skipping deployment cleanup to avoid removing valid models.", + str(e), + ) + return 0 model_list = config.get("model_list", None) if model_list: for model in model_list: @@ -3624,8 +3634,20 @@ class ProxyConfig: proxy_logging_obj: ProxyLogging, ): global llm_router, llm_model_list, master_key, general_settings - config_data = await proxy_config.get_config() - search_tools = self.parse_search_tools(config_data) + + # Load config separately so a timeout here doesn't block model loading + config_data: dict = {} + search_tools = None + try: + config_data = await proxy_config.get_config() + search_tools = self.parse_search_tools(config_data) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to load config in _update_llm_router: %s. " + "Proceeding with model loading using cached/empty config.", + str(e), + ) + try: models_list: list = new_models if isinstance(new_models, list) else [] if llm_router is None and master_key is not None: @@ -5302,13 +5324,15 @@ async def async_data_generator( ): verbose_proxy_logger.debug("inside generator") try: - # Use a list to accumulate response segments to avoid O(n^2) string concatenation - str_so_far_parts: list[str] = [] error_message: Optional[str] = None requested_model_from_client = _get_client_requested_model_for_streaming( request_data=request_data ) model_mismatch_logged = False + # Use a running string instead of list + join to avoid O(n^2) overhead. + # Previously "".join(str_so_far_parts) was called every chunk, re-joining + # the entire accumulated response. String += is O(n) amortized total. + _str_so_far: str = "" async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( user_api_key_dict=user_api_key_dict, response=response, @@ -5319,12 +5343,12 @@ async def async_data_generator( user_api_key_dict=user_api_key_dict, response=chunk, data=request_data, - str_so_far="".join(str_so_far_parts), + str_so_far=_str_so_far if _str_so_far else None, ) if isinstance(chunk, (ModelResponse, ModelResponseStream)): response_str = litellm.get_response_string(response_obj=chunk) - str_so_far_parts.append(response_str) + _str_so_far += response_str chunk, model_mismatch_logged = _restamp_streaming_chunk_model( chunk=chunk, @@ -7136,6 +7160,11 @@ async def audio_speech( ) except Exception as e: + await proxy_logging_obj.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=e, + request_data=data, + ) verbose_proxy_logger.error( "litellm.proxy.proxy_server.audio_speech(): Exception occured - {}".format( str(e) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index f6613b5548f..5e0d5336aa9 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -23,23 +23,31 @@ from typing import ( ) from litellm import _custom_logger_compatible_callbacks_literal -from litellm.constants import (DEFAULT_MODEL_CREATED_AT_TIME, - MAX_TEAM_LIST_LIMIT) -from litellm.proxy._types import (DB_CONNECTION_ERROR_TYPES, CommonProxyErrors, - ProxyErrorTypes, ProxyException, - SpendLogsMetadata, SpendLogsPayload) +from litellm.constants import DEFAULT_MODEL_CREATED_AT_TIME, MAX_TEAM_LIST_LIMIT +from litellm.proxy._types import ( + DB_CONNECTION_ERROR_TYPES, + CommonProxyErrors, + ProxyErrorTypes, + ProxyException, + SpendLogsMetadata, + SpendLogsPayload, +) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypes, CallTypesLiteral try: - from litellm_enterprise.enterprise_callbacks.send_emails.base_email import \ - BaseEmailLogger - from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import \ - ResendEmailLogger - from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import \ - SendGridEmailLogger - from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import \ - SMTPEmailLogger + from litellm_enterprise.enterprise_callbacks.send_emails.base_email import ( + BaseEmailLogger, + ) + from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import ( + ResendEmailLogger, + ) + from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import ( + SendGridEmailLogger, + ) + from litellm_enterprise.enterprise_callbacks.send_emails.smtp_email import ( + SMTPEmailLogger, + ) except ImportError: BaseEmailLogger = None # type: ignore SendGridEmailLogger = None # type: ignore @@ -58,56 +66,70 @@ from fastapi import HTTPException, status import litellm import litellm.litellm_core_utils import litellm.litellm_core_utils.litellm_logging -from litellm import (EmbeddingResponse, ImageResponse, ModelResponse, - ModelResponseStream, Router) +from litellm import ( + EmbeddingResponse, + ImageResponse, + ModelResponse, + ModelResponseStream, + Router, +) from litellm._logging import verbose_proxy_logger from litellm._service_logger import ServiceLogging, ServiceTypes from litellm.caching.caching import DualCache, RedisCache from litellm.caching.dual_cache import LimitedSizeOrderedDict from litellm.exceptions import RejectedRequestError -from litellm.integrations.custom_guardrail import (CustomGuardrail, - ModifyResponseException) +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + ModifyResponseException, +) from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting -from litellm.integrations.SlackAlerting.utils import \ - _add_langfuse_trace_id_to_alert +from litellm.integrations.SlackAlerting.utils import _add_langfuse_trace_id_to_alert from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler -from litellm.proxy._types import (AlertType, CallInfo, - LiteLLM_VerificationTokenView, Member, - UserAPIKeyAuth) +from litellm.proxy._types import ( + AlertType, + CallInfo, + LiteLLM_VerificationTokenView, + Member, + UserAPIKeyAuth, +) from litellm.proxy.auth.route_checks import RouteChecks -from litellm.proxy.db.create_views import (create_missing_views, - should_create_missing_views) +from litellm.proxy.db.create_views import ( + create_missing_views, + should_create_missing_views, +) from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.db.log_db_metrics import log_db_metrics from litellm.proxy.db.prisma_client import PrismaWrapper -from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import \ - UnifiedLLMGuardrails +from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( + UnifiedLLMGuardrails, +) from litellm.proxy.hooks import PROXY_HOOKS, get_proxy_hook from litellm.proxy.hooks.cache_control_check import _PROXY_CacheControlCheck from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter -from litellm.proxy.hooks.parallel_request_limiter import \ - _PROXY_MaxParallelRequestsHandler +from litellm.proxy.hooks.parallel_request_limiter import ( + _PROXY_MaxParallelRequestsHandler, +) from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.policy_engine.pipeline_executor import PipelineExecutor from litellm.secret_managers.main import str_to_bool from litellm.types.integrations.slack_alerting import DEFAULT_ALERT_TYPES -from litellm.types.mcp import (MCPDuringCallResponseObject, - MCPPreCallRequestObject, - MCPPreCallResponseObject) -from litellm.types.proxy.policy_engine.pipeline_types import \ - PipelineExecutionResult +from litellm.types.mcp import ( + MCPDuringCallResponseObject, + MCPPreCallRequestObject, + MCPPreCallResponseObject, +) +from litellm.types.proxy.policy_engine.pipeline_types import PipelineExecutionResult from litellm.types.utils import LLMResponseTypes, LoggedLiteLLMParams if TYPE_CHECKING: from opentelemetry.trace import Span as _Span - from litellm.litellm_core_utils.litellm_logging import \ - Logging as LiteLLMLoggingObj + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj Span = Union[_Span, Any] else: @@ -1050,9 +1072,10 @@ class ProxyLogging: """Process prompt template if applicable.""" from litellm.proxy.prompts.prompt_endpoints import ( - construct_versioned_prompt_id, get_latest_version_prompt_id) - from litellm.proxy.prompts.prompt_registry import \ - IN_MEMORY_PROMPT_REGISTRY + construct_versioned_prompt_id, + get_latest_version_prompt_id, + ) + from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY from litellm.utils import get_non_default_completion_params if prompt_version is None: @@ -1102,8 +1125,9 @@ class ProxyLogging: def _process_guardrail_metadata(self, data: dict) -> None: """Process guardrails from metadata and add to applied_guardrails.""" - from litellm.proxy.common_utils.callback_utils import \ - add_guardrail_to_applied_guardrails_header + from litellm.proxy.common_utils.callback_utils import ( + add_guardrail_to_applied_guardrails_header, + ) metadata_standard = data.get("metadata") or {} metadata_litellm = data.get("litellm_metadata") or {} @@ -2000,27 +2024,32 @@ class ProxyLogging: if isinstance(response, (ModelResponse, ModelResponseStream)): response_str = litellm.get_response_string(response_obj=response) elif isinstance(response, dict) and self.is_a2a_streaming_response(response): - from litellm.llms.a2a.common_utils import \ - extract_text_from_a2a_response + from litellm.llms.a2a.common_utils import extract_text_from_a2a_response response_str = extract_text_from_a2a_response(response) if response_str is not None: + # Cache model-level guardrails check per-request to avoid repeated + # dict lookups + llm_router.get_deployment() per callback per chunk. + _cached_guardrail_data: Optional[dict] = None + _guardrail_data_computed = False + for callback in litellm.callbacks: try: _callback: Optional[CustomLogger] = None if isinstance(callback, CustomGuardrail): # Main - V2 Guardrails implementation - from litellm.types.guardrails import \ - GuardrailEventHooks + from litellm.types.guardrails import GuardrailEventHooks - ## CHECK FOR MODEL-LEVEL GUARDRAILS - modified_data = _check_and_merge_model_level_guardrails( - data=data, llm_router=llm_router - ) + ## CHECK FOR MODEL-LEVEL GUARDRAILS (cached per-request) + if not _guardrail_data_computed: + _cached_guardrail_data = _check_and_merge_model_level_guardrails( + data=data, llm_router=llm_router + ) + _guardrail_data_computed = True if ( callback.should_run_guardrail( - data=modified_data, + data=_cached_guardrail_data, event_type=GuardrailEventHooks.post_call, ) is not True @@ -4626,8 +4655,9 @@ async def update_spend_logs_job( # Guardrail/policy usage tracking (same batch, outside spend-logs update) try: - from litellm.proxy.guardrails.usage_tracking import \ - process_spend_logs_guardrail_usage + from litellm.proxy.guardrails.usage_tracking import ( + process_spend_logs_guardrail_usage, + ) await process_spend_logs_guardrail_usage( prisma_client=prisma_client, logs_to_process=logs_to_process, @@ -4653,8 +4683,10 @@ async def _monitor_spend_logs_queue( db_writer_client: Optional HTTP handler for external spend logs endpoint proxy_logging_obj: Proxy logging object """ - from litellm.constants import (SPEND_LOG_QUEUE_POLL_INTERVAL, - SPEND_LOG_QUEUE_SIZE_THRESHOLD) + from litellm.constants import ( + SPEND_LOG_QUEUE_POLL_INTERVAL, + SPEND_LOG_QUEUE_SIZE_THRESHOLD, + ) threshold = SPEND_LOG_QUEUE_SIZE_THRESHOLD base_interval = SPEND_LOG_QUEUE_POLL_INTERVAL @@ -5175,11 +5207,12 @@ async def get_available_models_for_user( List of model names available to the user """ from litellm.proxy.auth.auth_checks import get_team_object - from litellm.proxy.auth.model_checks import (get_complete_model_list, - get_key_models, - get_team_models) - from litellm.proxy.management_endpoints.team_endpoints import \ - validate_membership + from litellm.proxy.auth.model_checks import ( + get_complete_model_list, + get_key_models, + get_team_models, + ) + from litellm.proxy.management_endpoints.team_endpoints import validate_membership # Get proxy model list and access groups if llm_router is None: diff --git a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py index e5be61690ba..836f9858744 100644 --- a/litellm/router_utils/pre_call_checks/model_rate_limit_check.py +++ b/litellm/router_utils/pre_call_checks/model_rate_limit_check.py @@ -129,27 +129,8 @@ class ModelRateLimitingCheck(CustomLogger): ), ) - # Check RPM limit + # Check RPM limit (atomic increment-first to avoid race conditions) if rpm_limit is not None: - # First check local cache - current_rpm = self.dual_cache.get_cache(key=rpm_key, local_only=True) - if current_rpm >= rpm_limit: - raise litellm.RateLimitError( - message=f"Model rate limit exceeded. RPM limit={rpm_limit}, current usage={current_rpm}", - llm_provider="", - model=model_name, - response=httpx.Response( - status_code=429, - content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={rpm_limit}. current usage={current_rpm}. id={model_id}, model_group={model_group}", - headers={"retry-after": str(60)}, - request=httpx.Request( - method="model_rate_limit_check", - url="https://github.com/BerriAI/litellm", - ), - ), - ) - - # Check redis cache and increment current_rpm = self.dual_cache.increment_cache( key=rpm_key, value=1, ttl=RoutingArgs.ttl ) @@ -226,30 +207,8 @@ class ModelRateLimitingCheck(CustomLogger): num_retries=0, # Don't retry - return 429 immediately ) - # Check RPM limit + # Check RPM limit (atomic increment-first to avoid race conditions) if rpm_limit is not None: - # First check local cache - current_rpm = await self.dual_cache.async_get_cache( - key=rpm_key, local_only=True - ) - if current_rpm is not None and current_rpm >= rpm_limit: - raise litellm.RateLimitError( - message=f"Model rate limit exceeded. RPM limit={rpm_limit}, current usage={current_rpm}", - llm_provider="", - model=model_name, - response=httpx.Response( - status_code=429, - content=f"{RouterErrors.user_defined_ratelimit_error.value} rpm limit={rpm_limit}. current usage={current_rpm}. id={model_id}, model_group={model_group}", - headers={"retry-after": str(60)}, - request=httpx.Request( - method="model_rate_limit_check", - url="https://github.com/BerriAI/litellm", - ), - ), - num_retries=0, # Don't retry - return 429 immediately - ) - - # Check redis cache and increment current_rpm = await self.dual_cache.async_increment_cache( key=rpm_key, value=1, diff --git a/litellm/types/integrations/datadog_metrics.py b/litellm/types/integrations/datadog_metrics.py new file mode 100644 index 00000000000..4c980cdee6d --- /dev/null +++ b/litellm/types/integrations/datadog_metrics.py @@ -0,0 +1,20 @@ +from typing import List, Optional + +from typing_extensions import TypedDict + + +class DatadogMetricPoint(TypedDict): + timestamp: int # Unix epoch seconds + value: float # The metric value + + +class DatadogMetricSeries(TypedDict, total=False): + metric: str + type: int # 0=unspecified, 1=count, 2=rate, 3=gauge + points: List[DatadogMetricPoint] + tags: List[str] + interval: Optional[int] # Required for count (type=1) and rate (type=2) metrics + + +class DatadogMetricsPayload(TypedDict): + series: List[DatadogMetricSeries] diff --git a/tests/litellm/test_no_hardcoded_secrets.py b/tests/litellm/test_no_hardcoded_secrets.py new file mode 100644 index 00000000000..f22eb1f3a72 --- /dev/null +++ b/tests/litellm/test_no_hardcoded_secrets.py @@ -0,0 +1,74 @@ +""" +Test to ensure no hardcoded secrets exist in the codebase. + +This catches Base64 Basic Authentication strings and other secret patterns +that would be flagged by secret scanners like GitGuardian/ggshield. +""" + +import base64 +import os +import re + +import pytest + +# Root of the litellm package +LITELLM_ROOT = os.path.join(os.path.dirname(__file__), "..", "..", "litellm") + +# Regex for Base64 Basic Auth patterns: 'Basic ' +# Matches strings like: Basic YW55dGhpbmc6YW55dGhpbmc= +BASIC_AUTH_PATTERN = re.compile( + r"""['"]Basic\s+([A-Za-z0-9+/]{16,}={0,2})['"]""" +) + +# Directories/files to skip +SKIP_DIRS = {"__pycache__", ".git", "node_modules", ".mypy_cache", ".ruff_cache"} + + +def _is_real_base64_credentials(match_str: str) -> bool: + """Check if a Base64 string decodes to something that looks like credentials (user:pass).""" + try: + # Add padding if needed - Base64 strings may omit trailing '=' + padded = match_str + "=" * (-len(match_str) % 4) + decoded = base64.b64decode(padded).decode("utf-8", errors="ignore") + return ":" in decoded + except Exception: + return False + + +def _collect_python_files(): + """Collect all Python files under the litellm package.""" + python_files = [] + for root, dirs, files in os.walk(LITELLM_ROOT): + dirs[:] = [d for d in dirs if d not in SKIP_DIRS] + for f in files: + if f.endswith(".py"): + python_files.append(os.path.join(root, f)) + return python_files + + +def test_no_hardcoded_basic_auth_secrets(): + """Ensure no hardcoded Base64 Basic Authentication credentials exist in source code. + + This test prevents regressions like the one caught by T-Mobile's GitGuardian + container scan, where a docstring contained a literal Base64-encoded + 'Basic YW55dGhpbmc6YW55dGhpbmc' string (anything:anything). + """ + violations = [] + + for filepath in _collect_python_files(): + with open(filepath, "r", errors="ignore") as f: + for line_num, line in enumerate(f, start=1): + for match in BASIC_AUTH_PATTERN.finditer(line): + b64_value = match.group(1) + if _is_real_base64_credentials(b64_value): + rel_path = os.path.relpath(filepath, LITELLM_ROOT) + violations.append( + f" {rel_path}:{line_num}: {match.group(0)}" + ) + + assert not violations, ( + "Found hardcoded Base64 Basic Auth credentials that will be flagged by " + "secret scanners (e.g. GitGuardian/ggshield):\n" + + "\n".join(violations) + + "\n\nUse placeholders like '' in comments/docs instead." + ) diff --git a/tests/test_litellm/integrations/datadog/test_datadog_metrics.py b/tests/test_litellm/integrations/datadog/test_datadog_metrics.py new file mode 100644 index 00000000000..757c558c298 --- /dev/null +++ b/tests/test_litellm/integrations/datadog/test_datadog_metrics.py @@ -0,0 +1,273 @@ +import os +import time +from datetime import datetime, timedelta +from unittest.mock import AsyncMock + +import pytest +from httpx import Request, Response + +from litellm.integrations.datadog.datadog_metrics import DatadogMetricsLogger +from litellm.types.utils import StandardLoggingPayload + + +@pytest.fixture +def clean_env(): + """Set test env vars and restore originals after test.""" + keys = ["DD_API_KEY", "DD_APP_KEY", "DD_SITE", "DD_ENV", "DD_SERVICE", "DD_VERSION"] + originals = {k: os.environ.get(k) for k in keys} + + os.environ["DD_API_KEY"] = "test_api_key" + os.environ["DD_APP_KEY"] = "test_app_key" + os.environ["DD_SITE"] = "test.datadoghq.com" + os.environ["DD_ENV"] = "test-env" + os.environ["DD_SERVICE"] = "test-service" + os.environ["DD_VERSION"] = "1.0.0" + + yield + + for k, v in originals.items(): + if v is not None: + os.environ[k] = v + elif k in os.environ: + del os.environ[k] + + +@pytest.mark.asyncio +async def test_init(clean_env): + """Test initialization sets up clients and url correctly.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + assert logger.upload_url == "https://api.test.datadoghq.com/api/v2/series" + + +@pytest.mark.asyncio +async def test_extract_tags(clean_env): + """Test tag extraction from a StandardLoggingPayload.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + + payload = StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + model_group="gpt-4", + metadata={"user_api_key_team_alias": "test-team"}, + ) + + tags = logger._extract_tags(log=payload, status_code="200") + + assert "env:test-env" in tags + assert "service:test-service" in tags + assert "version:1.0.0" in tags + assert "provider:openai" in tags + assert "model_name:gpt-4o" in tags + assert "model_group:gpt-4" in tags + assert "status_code:200" in tags + assert "team:test-team" in tags + + +@pytest.mark.asyncio +async def test_extract_tags_no_team(clean_env): + """Test tag extraction when no team info is present.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + + payload = StandardLoggingPayload( + custom_llm_provider="anthropic", + model="claude-3-sonnet", + ) + + tags = logger._extract_tags(log=payload, status_code="500") + + assert "provider:anthropic" in tags + assert "model_name:claude-3-sonnet" in tags + assert "status_code:500" in tags + assert not any(tag.startswith("team:") for tag in tags) + + +@pytest.mark.asyncio +async def test_add_metrics_from_log(clean_env): + """Test that _add_metrics_from_log appends the correct metric series to the queue.""" + logger = DatadogMetricsLogger(batch_size=100, start_periodic_flush=False) + + now = datetime.now() + start_time = now - timedelta(seconds=2) + api_call_start_time = now - timedelta(seconds=1) + + payload = StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + ) + + kwargs = { + "start_time": start_time, + "api_call_start_time": api_call_start_time, + "end_time": now, + } + + logger._add_metrics_from_log(log=payload, kwargs=kwargs, status_code="200") + + # Should have 3 series: total_latency, llm_api_latency, request_count + assert len(logger.log_queue) == 3 + + metrics = {s["metric"]: s for s in logger.log_queue} + + # Total latency ~2s + total = metrics["litellm.request.total_latency"] + assert total["type"] == 3 # gauge + assert abs(total["points"][0]["value"] - 2.0) < 0.1 + + # LLM API latency ~1s + llm = metrics["litellm.llm_api.latency"] + assert llm["type"] == 3 # gauge + assert abs(llm["points"][0]["value"] - 1.0) < 0.1 + + # Request count + count = metrics["litellm.llm_api.request_count"] + assert count["type"] == 1 # count + assert count["points"][0]["value"] == 1.0 + assert "status_code:200" in count["tags"] + + +@pytest.mark.asyncio +async def test_async_log_success_event(clean_env): + """Test that success events are added to the queue.""" + logger = DatadogMetricsLogger(batch_size=100, start_periodic_flush=False) + + now = datetime.now() + start_time = now - timedelta(seconds=1) + + await logger.async_log_success_event( + kwargs={ + "standard_logging_object": StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + ), + "start_time": start_time, + "end_time": now, + }, + response_obj=None, + start_time=start_time, + end_time=now, + ) + + # At least request_count and total_latency + assert len(logger.log_queue) >= 2 + + +@pytest.mark.asyncio +async def test_async_log_success_event_no_standard_logging_object(clean_env): + """Test that events without standard_logging_object are skipped.""" + logger = DatadogMetricsLogger(batch_size=100, start_periodic_flush=False) + + await logger.async_log_success_event( + kwargs={}, + response_obj=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + assert len(logger.log_queue) == 0 + + +@pytest.mark.asyncio +async def test_async_log_failure_event_extracts_status_code(clean_env): + """Test that failure events extract the error status code.""" + logger = DatadogMetricsLogger(batch_size=100, start_periodic_flush=False) + + now = datetime.now() + start_time = now - timedelta(seconds=1) + + await logger.async_log_failure_event( + kwargs={ + "standard_logging_object": StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + error_information={"error_code": "429"}, + ), + "start_time": start_time, + "end_time": now, + }, + response_obj=None, + start_time=start_time, + end_time=now, + ) + + count_series = next( + (s for s in logger.log_queue if s["metric"] == "litellm.llm_api.request_count"), + None, + ) + assert count_series is not None + assert "status_code:429" in count_series["tags"] + + +@pytest.mark.asyncio +async def test_async_log_failure_event_default_status_code(clean_env): + """Test that failure events default to 500 when no error_code is present.""" + logger = DatadogMetricsLogger(batch_size=100, start_periodic_flush=False) + + now = datetime.now() + + await logger.async_log_failure_event( + kwargs={ + "standard_logging_object": StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + ), + "start_time": now, + "end_time": now, + }, + response_obj=None, + start_time=now, + end_time=now, + ) + + count_series = next( + (s for s in logger.log_queue if s["metric"] == "litellm.llm_api.request_count"), + None, + ) + assert count_series is not None + assert "status_code:500" in count_series["tags"] + + +@pytest.mark.asyncio +async def test_async_send_batch(clean_env): + """Test that async_send_batch uploads metrics to Datadog.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + logger.async_client = AsyncMock() + mock_request = Request("POST", "https://api.test.datadoghq.com/api/v2/series") + logger.async_client.post.return_value = Response( + 202, json={"status": "ok"}, request=mock_request + ) + + # Manually add a metric series to the queue + logger.log_queue = [ + { + "metric": "litellm.request.total_latency", + "type": 3, + "points": [{"timestamp": int(time.time()), "value": 1.5}], + "tags": ["env:test"], + } + ] + + await logger.async_send_batch() + + assert logger.async_client.post.called + call_args = logger.async_client.post.call_args + assert call_args[0][0] == "https://api.test.datadoghq.com/api/v2/series" + + # Verify gzip + JSON payload + import gzip + import json + + compressed = call_args[1]["content"] + payload = json.loads(gzip.decompress(compressed).decode("utf-8")) + assert len(payload["series"]) == 1 + assert payload["series"][0]["metric"] == "litellm.request.total_latency" + + +@pytest.mark.asyncio +async def test_async_send_batch_empty_queue(clean_env): + """Test that async_send_batch does nothing when queue is empty.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + logger.async_client = AsyncMock() + + await logger.async_send_batch() + + assert not logger.async_client.post.called 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 f6d3d3c12f7..345f3ae7c5d 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -3493,3 +3493,262 @@ class TestBedrockMinThinkingBudgetTokens: drop_params=False, ) assert "thinking" not in result or result.get("thinking") is None + +def test_transform_response_with_both_json_tool_call_and_real_tool(): + """ + When Bedrock returns BOTH json_tool_call AND a real tool (get_weather), + only the real tool should remain in tool_calls. The json_tool_call should be filtered out. + Fixes https://github.com/BerriAI/litellm/issues/18381 + """ + from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig + from litellm.types.utils import ModelResponse + + response_json = { + "metrics": {"latencyMs": 200}, + "output": { + "message": { + "role": "assistant", + "content": [ + { + "toolUse": { + "toolUseId": "tooluse_json_001", + "name": "json_tool_call", + "input": { + "Current_Temperature": 62, + "Weather_Explanation": "Mild and cool.", + }, + } + }, + { + "toolUse": { + "toolUseId": "tooluse_weather_001", + "name": "get_weather", + "input": { + "location": "San Francisco, CA", + "unit": "fahrenheit", + }, + } + }, + ], + } + }, + "stopReason": "tool_use", + "usage": { + "inputTokens": 100, + "outputTokens": 50, + "totalTokens": 150, + "cacheReadInputTokenCount": 0, + "cacheReadInputTokens": 0, + "cacheWriteInputTokenCount": 0, + "cacheWriteInputTokens": 0, + }, + } + + class MockResponse: + def json(self): + return response_json + + @property + def text(self): + return json.dumps(response_json) + + config = AmazonConverseConfig() + model_response = ModelResponse() + optional_params = {"json_mode": True} + + result = config._transform_response( + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + response=MockResponse(), + model_response=model_response, + stream=False, + logging_obj=None, + optional_params=optional_params, + api_key=None, + data=None, + messages=[], + encoding=None, + ) + + # Only real tool should remain + assert result.choices[0].message.tool_calls is not None + assert len(result.choices[0].message.tool_calls) == 1 + assert result.choices[0].message.tool_calls[0].function.name == "get_weather" + assert ( + result.choices[0].message.tool_calls[0].function.arguments + == '{"location": "San Francisco, CA", "unit": "fahrenheit"}' + ) + + # json_tool_call content should be preserved as message text + content = result.choices[0].message.content + assert content is not None + parsed = json.loads(content) + assert parsed["Current_Temperature"] == 62 + assert parsed["Weather_Explanation"] == "Mild and cool." + + +def test_transform_response_does_not_mutate_optional_params(): + """ + Verify that optional_params still contains json_mode after _transform_response. + Previously, .pop() was used which mutated the caller's dict. + """ + from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig + from litellm.types.utils import ModelResponse + + response_json = { + "metrics": {"latencyMs": 50}, + "output": { + "message": { + "role": "assistant", + "content": [ + { + "toolUse": { + "toolUseId": "tooluse_001", + "name": "json_tool_call", + "input": {"result": "ok"}, + } + } + ], + } + }, + "stopReason": "tool_use", + "usage": { + "inputTokens": 10, + "outputTokens": 5, + "totalTokens": 15, + "cacheReadInputTokenCount": 0, + "cacheReadInputTokens": 0, + "cacheWriteInputTokenCount": 0, + "cacheWriteInputTokens": 0, + }, + } + + class MockResponse: + def json(self): + return response_json + + @property + def text(self): + return json.dumps(response_json) + + config = AmazonConverseConfig() + model_response = ModelResponse() + optional_params = {"json_mode": True, "other_key": "value"} + + config._transform_response( + model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", + response=MockResponse(), + model_response=model_response, + stream=False, + logging_obj=None, + optional_params=optional_params, + api_key=None, + data=None, + messages=[], + encoding=None, + ) + + # json_mode should still be in optional_params (not popped) + assert "json_mode" in optional_params + assert optional_params["json_mode"] is True + assert optional_params["other_key"] == "value" + + +def test_streaming_filters_json_tool_call_with_real_tools(): + """ + Simulate streaming chunks where both json_tool_call and a real tool arrive. + Verify json_tool_call chunks are converted to text content while real tool + chunks pass through normally. + """ + from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder + from litellm.types.llms.bedrock import ( + ContentBlockDeltaEvent, + ContentBlockStartEvent, + ) + + decoder = AWSEventStreamDecoder(model="test-model", json_mode=True) + + # Chunk 1: json_tool_call start + json_start = ContentBlockStartEvent( + toolUse={ + "toolUseId": "tooluse_json_001", + "name": "json_tool_call", + } + ) + tool_use_1, _, _ = decoder._handle_converse_start_event(json_start) + # json_tool_call start should be suppressed (return None tool_use) + assert tool_use_1 is None + # tool_calls_index should NOT have been incremented + assert decoder.tool_calls_index is None + + # Chunk 2: json_tool_call delta — should become text, not tool_use + json_delta = ContentBlockDeltaEvent(toolUse={"input": '{"temp": 62}'}) + text_2, tool_use_2, _, _, _ = decoder._handle_converse_delta_event( + json_delta, index=0 + ) + assert text_2 == '{"temp": 62}' + assert tool_use_2 is None + + # Chunk 3: json_tool_call stop + stop_tool = decoder._handle_converse_stop_event(index=0) + assert stop_tool is None + # _current_tool_name should be reset + assert decoder._current_tool_name is None + + # Chunk 4: real tool start + real_start = ContentBlockStartEvent( + toolUse={ + "toolUseId": "tooluse_weather_001", + "name": "get_weather", + } + ) + tool_use_4, _, _ = decoder._handle_converse_start_event(real_start) + assert tool_use_4 is not None + assert tool_use_4["function"]["name"] == "get_weather" + assert decoder.tool_calls_index == 0 + + # Chunk 5: real tool delta + real_delta = ContentBlockDeltaEvent( + toolUse={"input": '{"location": "SF"}'} + ) + text_5, tool_use_5, _, _, _ = decoder._handle_converse_delta_event( + real_delta, index=1 + ) + assert text_5 == "" + assert tool_use_5 is not None + assert tool_use_5["function"]["arguments"] == '{"location": "SF"}' + + +def test_streaming_without_json_mode_passes_all_tools(): + """ + Verify backward compatibility: when json_mode=False, all tools + (including json_tool_call if present) pass through unchanged. + """ + from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder + from litellm.types.llms.bedrock import ( + ContentBlockDeltaEvent, + ContentBlockStartEvent, + ) + + decoder = AWSEventStreamDecoder(model="test-model", json_mode=False) + + # json_tool_call start — should pass through when json_mode=False + json_start = ContentBlockStartEvent( + toolUse={ + "toolUseId": "tooluse_json_001", + "name": "json_tool_call", + } + ) + tool_use, _, _ = decoder._handle_converse_start_event(json_start) + assert tool_use is not None + assert tool_use["function"]["name"] == "json_tool_call" + assert decoder.tool_calls_index == 0 + + # json_tool_call delta — should be a tool_use, not text + json_delta = ContentBlockDeltaEvent(toolUse={"input": '{"data": 1}'}) + text, tool_use_delta, _, _, _ = decoder._handle_converse_delta_event( + json_delta, index=0 + ) + assert text == "" + assert tool_use_delta is not None + assert tool_use_delta["function"]["arguments"] == '{"data": 1}' + diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index b56d13bb932..8418dde5e9c 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1485,4 +1485,273 @@ async def test_get_objects_resolves_org_by_name(): ) +# --------------------------------------------------------------------------- +# Fix 1: OIDC discovery URL resolution +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_resolve_jwks_url_passthrough_for_direct_jwks_url(): + """Non-discovery URLs are returned unchanged.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.dual_cache import DualCache + + handler = JWTHandler() + handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + url = "https://login.microsoftonline.com/common/discovery/keys" + result = await handler._resolve_jwks_url(url) + assert result == url + + +@pytest.mark.asyncio +async def test_resolve_jwks_url_resolves_oidc_discovery_document(): + """ + A .well-known/openid-configuration URL should be fetched and its + jwks_uri returned. + """ + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.caching.dual_cache import DualCache + + handler = JWTHandler() + cache = DualCache() + handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" + jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = {"jwks_uri": jwks_url, "issuer": "https://..."} + + with patch.object(handler.http_handler, "get", new_callable=AsyncMock, return_value=mock_response) as mock_get: + result = await handler._resolve_jwks_url(discovery_url) + + assert result == jwks_url + mock_get.assert_called_once_with(discovery_url) + + +@pytest.mark.asyncio +async def test_resolve_jwks_url_caches_resolved_jwks_uri(): + """Resolved jwks_uri is cached — second call does not hit the network.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.caching.dual_cache import DualCache + + handler = JWTHandler() + cache = DualCache() + handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + + discovery_url = "https://login.microsoftonline.com/tenant/.well-known/openid-configuration" + jwks_url = "https://login.microsoftonline.com/tenant/discovery/keys" + + mock_response = MagicMock() + mock_response.json.return_value = {"jwks_uri": jwks_url} + + with patch.object(handler.http_handler, "get", new_callable=AsyncMock, return_value=mock_response) as mock_get: + first = await handler._resolve_jwks_url(discovery_url) + second = await handler._resolve_jwks_url(discovery_url) + + assert first == jwks_url + assert second == jwks_url + # Network should only be hit once + assert mock_get.call_count == 1 + + +@pytest.mark.asyncio +async def test_resolve_jwks_url_raises_if_no_jwks_uri_in_discovery_doc(): + """Raise a helpful error if the discovery document has no jwks_uri.""" + from unittest.mock import AsyncMock, MagicMock, patch + + from litellm.caching.dual_cache import DualCache + + handler = JWTHandler() + handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + + discovery_url = "https://example.com/.well-known/openid-configuration" + mock_response = MagicMock() + mock_response.json.return_value = {"issuer": "https://example.com"} # no jwks_uri + + with patch.object(handler.http_handler, "get", new_callable=AsyncMock, return_value=mock_response): + with pytest.raises(Exception, match="jwks_uri"): + await handler._resolve_jwks_url(discovery_url) + + +# --------------------------------------------------------------------------- +# Fix 2: handle array values in team_id_jwt_field (e.g. AAD "roles" claim) +# --------------------------------------------------------------------------- + + +def _make_jwt_handler(team_id_jwt_field: str) -> JWTHandler: + from litellm.caching.dual_cache import DualCache + + handler = JWTHandler() + handler.update_environment( + prisma_client=None, + user_api_key_cache=DualCache(), + litellm_jwtauth=LiteLLM_JWTAuth(team_id_jwt_field=team_id_jwt_field), + ) + return handler + + +def test_get_team_id_returns_first_element_when_roles_is_list(): + """ + AAD sends roles as a list. get_team_id() must return the first string + element rather than the raw list (which would later crash with + 'unhashable type: list'). + """ + handler = _make_jwt_handler("roles") + token = {"oid": "user-oid", "roles": ["team1"]} + result = handler.get_team_id(token=token, default_value=None) + assert result == "team1" + + +def test_get_team_id_returns_first_element_from_multi_value_roles_list(): + """When roles has multiple entries, the first one is used.""" + handler = _make_jwt_handler("roles") + token = {"roles": ["team2", "team1"]} + result = handler.get_team_id(token=token, default_value=None) + assert result == "team2" + + +def test_get_team_id_returns_default_when_roles_list_is_empty(): + """Empty list should fall back to default_value.""" + handler = _make_jwt_handler("roles") + token = {"roles": []} + result = handler.get_team_id(token=token, default_value="fallback") + assert result == "fallback" + + +def test_get_team_id_still_works_with_string_value(): + """String values (non-array) continue to work as before.""" + handler = _make_jwt_handler("appid") + token = {"appid": "my-team-id"} + result = handler.get_team_id(token=token, default_value=None) + assert result == "my-team-id" + + +def test_get_team_id_list_result_is_hashable(): + """ + The value returned by get_team_id() must be hashable so it can be + added to a set (the operation that previously crashed). + """ + handler = _make_jwt_handler("roles") + token = {"roles": ["team1"]} + result = handler.get_team_id(token=token, default_value=None) + # This must not raise TypeError + s: set = set() + s.add(result) + assert "team1" in s + + +# --------------------------------------------------------------------------- +# Fix 3: helpful error message for dot-notation array indexing (roles.0) +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_hints_bracket_notation(): + """ + When team_id_jwt_field is set to 'roles.0' (unsupported dot-notation for + array indexing) and no team is found, the exception message should suggest + using 'roles' instead (and explain LiteLLM auto-unwraps list values). + """ + from unittest.mock import MagicMock + + from litellm.caching.dual_cache import DualCache + + handler = _make_jwt_handler("roles.0") + # token has roles as a list — dot-notation won't find anything + token = {"roles": ["team1"]} + + with pytest.raises(Exception) as exc_info: + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + + error_msg = str(exc_info.value) + # Should mention the bad field name and suggest the fix + assert "roles.0" in error_msg, f"Expected field name in: {error_msg}" + assert "roles" in error_msg and "list" in error_msg, ( + f"Expected hint about using 'roles' instead: {error_msg}" + ) + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_hints_bracket_index_notation(): + """ + When team_id_jwt_field is set to 'roles[0]' (bracket indexing, also unsupported + in get_nested_value) the error message should suggest using 'roles' instead. + """ + from unittest.mock import MagicMock + + from litellm.caching.dual_cache import DualCache + + handler = _make_jwt_handler("roles[0]") + token = {"roles": ["team1"]} + + with pytest.raises(Exception) as exc_info: + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + + error_msg = str(exc_info.value) + assert "roles[0]" in error_msg, f"Expected field name in: {error_msg}" + assert "roles" in error_msg and "list" in error_msg, ( + f"Expected hint about using 'roles' instead: {error_msg}" + ) + + +@pytest.mark.asyncio +async def test_find_and_validate_specific_team_id_no_hint_for_valid_field(): + """ + When team_id_jwt_field is a normal field name (no dot-notation) the + error message should not contain a spurious bracket-notation hint. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.dual_cache import DualCache + + handler = _make_jwt_handler("appid") + token = {} # no appid — triggers the "no team found" path + + with pytest.raises(Exception) as exc_info: + await JWTAuthManager.find_and_validate_specific_team_id( + jwt_handler=handler, + jwt_valid_token=token, + prisma_client=None, + user_api_key_cache=DualCache(), + parent_otel_span=None, + proxy_logging_obj=MagicMock(), + ) + + error_msg = str(exc_info.value) + assert "Hint" not in error_msg diff --git a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py index 308c8cdbce1..234b83bcd95 100644 --- a/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py +++ b/tests/test_litellm/proxy/common_utils/test_key_rotation_integration.py @@ -142,3 +142,156 @@ class TestKeyRotationManagerPassesKeyAlias: assert ( captured_request.key_alias is None ), "key_alias should be None for keys without alias" + +class TestKeyRotationSecretNamingStability: + """ + Tests that the fallback secret name in the rotation hook remains stable + across rotations to prevent AWS secret sprawl. + + Couple this with the validation fix (Step 1-2) to ensure a stable + experience for secret management. + """ + + @pytest.mark.asyncio + async def test_rotation_hook_uses_initial_secret_name_fallback(self): + """ + GIVEN: A key WITHOUT an alias (has an initial_secret_name based on token ID) + WHEN: The key is rotated + THEN: The hook MUST reuse the existing secret name, NOT generate a new one + based on the new token ID. + """ + from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + + # 1. Existing key without alias + initial_token_hash = "hashed-initial-token" + existing_key = MagicMock(spec=LiteLLM_VerificationToken) + existing_key.token = initial_token_hash + existing_key.key_alias = None + initial_secret_name = f"virtual-key-{initial_token_hash}" + + # 2. Rotation response (new token ID) + new_token_id = "hashed-new-token" + response = GenerateKeyResponse( + key="sk-new-key", + token_id=new_token_id, + key_alias=None + ) + + # 3. Request data without alias + request_data = RegenerateKeyRequest( + key=initial_token_hash, + key_alias=None + ) + + with patch("litellm.proxy.hooks.key_management_event_hooks.KeyManagementEventHooks._rotate_virtual_key_in_secret_manager", new_callable=AsyncMock) as mock_rotate: + await KeyManagementEventHooks.async_key_rotated_hook( + data=request_data, + existing_key_row=existing_key, + response=response, + user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin", api_key="sk-1234", user_id="1234") + ) + + # ASSERT: The new_secret_name MUST be the same as initial_secret_name + # This ensures PutSecretValue instead of a new secret creation + mock_rotate.assert_called_once() + call_kwargs = mock_rotate.call_args.kwargs + assert call_kwargs["current_secret_name"] == initial_secret_name + assert call_kwargs["new_secret_name"] == initial_secret_name, \ + f"Secret name drift! Expected {initial_secret_name}, got {call_kwargs['new_secret_name']}. This causes secret sprawl." + + @pytest.mark.asyncio + async def test_rotation_hook_pre_rotation_alias_consistency(self): + """ + GIVEN: A key WITH an alias + WHEN: The key is rotated + THEN: The hook uses the alias for both current and new names. + """ + from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks + from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth + + test_alias = "tenant1/stable-key" + existing_key = MagicMock(spec=LiteLLM_VerificationToken) + existing_key.token = "old-hash" + existing_key.key_alias = test_alias + + response = GenerateKeyResponse(token_id="new-hash", key="sk-new", key_alias=test_alias) + request_data = RegenerateKeyRequest(key="old-hash", key_alias=test_alias) + + with patch("litellm.proxy.hooks.key_management_event_hooks.KeyManagementEventHooks._rotate_virtual_key_in_secret_manager", new_callable=AsyncMock) as mock_rotate: + await KeyManagementEventHooks.async_key_rotated_hook( + data=request_data, + existing_key_row=existing_key, + response=response, + user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin", api_key="sk-123", user_id="1") + ) + mock_rotate.assert_called_once() + assert mock_rotate.call_args.kwargs["current_secret_name"] == test_alias + assert mock_rotate.call_args.kwargs["new_secret_name"] == test_alias + + @pytest.mark.asyncio + async def test_set_key_rotation_fields_requires_alias(self): + """ + Tests that _set_key_rotation_fields enforces key_alias requirement + when secret storage is enabled. + """ + import litellm + from litellm.proxy.management_endpoints.key_management_endpoints import _set_key_rotation_fields + from litellm.proxy._types import ProxyException + # Create a mock for settings + mock_settings = MagicMock() + mock_settings.store_virtual_keys = True + + # Mock settings: store_virtual_keys = True + with patch("litellm._key_management_settings", mock_settings): + data = {"auto_rotate": True} # Missing key_alias + + # Should raise ProxyException 400 + with pytest.raises(ProxyException) as exc: + _set_key_rotation_fields(data, auto_rotate=True, rotation_interval="30d") + + assert str(exc.value.code) == "400" + assert "key_alias is required" in str(exc.value.message) + + # Adding key_alias should work + data["key_alias"] = "valid-alias" + _set_key_rotation_fields(data, auto_rotate=True, rotation_interval="30d") + assert data["auto_rotate"] is True + assert "key_rotation_at" in data + + @pytest.mark.asyncio + async def test_set_key_rotation_fields_with_existing_alias(self): + """ + Tests that _set_key_rotation_fields allows enabling rotation + if the key already has an alias in the database (even if not in current request). + """ + from litellm.proxy.management_endpoints.key_management_endpoints import _set_key_rotation_fields + from unittest.mock import MagicMock, patch + + mock_settings = MagicMock() + mock_settings.store_virtual_keys = True + + with patch("litellm._key_management_settings", mock_settings): + # 1. No alias in request, but HAS existing_key_alias + data = {"auto_rotate": True} + _set_key_rotation_fields( + data, + auto_rotate=True, + rotation_interval="30d", + existing_key_alias="already-exists-in-db" + ) + # Should NOT raise, and field should be set + assert data["auto_rotate"] is True + assert "key_rotation_at" in data + + # 2. Verify it still fails if NO alias AND NO existing_key_alias + from litellm.proxy._types import ProxyException + data_fail = {"auto_rotate": True} + with pytest.raises(ProxyException) as exc: + _set_key_rotation_fields( + data_fail, + auto_rotate=True, + rotation_interval="30d", + existing_key_alias=None + ) + assert str(exc.value.code) == "400" 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 7565e901ecd..11c80351839 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 @@ -6070,6 +6070,121 @@ async def test_build_key_filter_admin_all_member_overlap(): ) +@pytest.mark.asyncio +async def test_build_key_filter_project_id(): + """ + Test that project_id is applied as a global AND condition, narrowing all results + to keys that belong to the specified project. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_key_filter_conditions, + ) + + user_id = "user-123" + project_id = "proj-abc" + + where = _build_key_filter_conditions( + user_id=user_id, + team_id=None, + organization_id=None, + key_alias=None, + key_hash=None, + exclude_team_id=None, + admin_team_ids=None, + member_team_ids=None, + include_created_by_keys=False, + project_id=project_id, + ) + + # Should be wrapped in a top-level AND for the project_id filter + assert "AND" in where + and_parts = where["AND"] + assert len(and_parts) == 2 + + # Second part of AND should be the project_id filter + assert {"project_id": project_id} in and_parts + + +@pytest.mark.asyncio +async def test_build_key_filter_access_group_id(): + """ + Test that access_group_id is applied as a global AND condition using hasSome, + narrowing results to keys whose access_group_ids array contains the given ID. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_key_filter_conditions, + ) + + user_id = "user-123" + access_group_id = "ag-xyz" + + where = _build_key_filter_conditions( + user_id=user_id, + team_id=None, + organization_id=None, + key_alias=None, + key_hash=None, + exclude_team_id=None, + admin_team_ids=None, + member_team_ids=None, + include_created_by_keys=False, + access_group_id=access_group_id, + ) + + # Should be wrapped in a top-level AND for the access_group_id filter + assert "AND" in where + and_parts = where["AND"] + assert len(and_parts) == 2 + + # Second part of AND should use hasSome for the array field + assert {"access_group_ids": {"hasSome": [access_group_id]}} in and_parts + + +@pytest.mark.asyncio +async def test_build_key_filter_project_id_and_access_group_id(): + """ + Test that project_id and access_group_id stack correctly when both are provided. + Both should be applied as AND conditions, narrowing results to keys that match both. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_key_filter_conditions, + ) + + user_id = "user-123" + project_id = "proj-abc" + access_group_id = "ag-xyz" + + where = _build_key_filter_conditions( + user_id=user_id, + team_id=None, + organization_id=None, + key_alias=None, + key_hash=None, + exclude_team_id=None, + admin_team_ids=None, + member_team_ids=None, + include_created_by_keys=False, + project_id=project_id, + access_group_id=access_group_id, + ) + + # After project_id: {"AND": [visibility_where, {"project_id": ...}]} + # After access_group_id: {"AND": [above, {"access_group_ids": ...}]} + assert "AND" in where + outer_and = where["AND"] + assert len(outer_and) == 2 + + # The access_group_ids filter is the outermost AND + access_group_filter = outer_and[1] + assert access_group_filter == {"access_group_ids": {"hasSome": [access_group_id]}} + + # The project_id filter is nested one level in + inner = outer_and[0] + assert "AND" in inner + inner_and = inner["AND"] + assert {"project_id": project_id} in inner_and + + @pytest.mark.asyncio async def test_get_member_team_ids(): """ @@ -6305,3 +6420,38 @@ async def test_key_aliases_no_search_omits_ilike_filter(): assert "ILIKE" not in count_sql + +class TestValidateKeyAliasFormat: + def test_validate_key_alias_format_valid(self): + from litellm.proxy.management_endpoints.key_management_endpoints import _validate_key_alias_format + # Valid cases + _validate_key_alias_format(None) # OK + _validate_key_alias_format("valid-alias") + _validate_key_alias_format("valid_alias") + _validate_key_alias_format("valid.alias") + _validate_key_alias_format("valid/alias") + _validate_key_alias_format("a" * 255) + _validate_key_alias_format("my-key-123") + + def test_validate_key_alias_format_invalid(self): + from litellm.proxy.management_endpoints.key_management_endpoints import _validate_key_alias_format + from litellm.proxy._types import ProxyException + + invalid_aliases = [ + "", # empty + " ", # whitespace + "a", # too short (min 2) + "!", # special char + "-start", # non-alphanumeric start + "end-", # non-alphanumeric end + "invalid@char", # invalid char + "a" * 256, # too long + " leading", + "trailing ", + ] + + for alias in invalid_aliases: + with pytest.raises(ProxyException) as exc: + _validate_key_alias_format(alias) + assert str(exc.value.code) == "400" + assert "Invalid key_alias format" in str(exc.value.message) diff --git a/tests/test_litellm/proxy/test_audio_speech_prometheus_hooks.py b/tests/test_litellm/proxy/test_audio_speech_prometheus_hooks.py new file mode 100644 index 00000000000..5a13a4dc531 --- /dev/null +++ b/tests/test_litellm/proxy/test_audio_speech_prometheus_hooks.py @@ -0,0 +1,139 @@ +import asyncio +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi.testclient import TestClient + +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.proxy_server import app, initialize + + +def _mock_user_api_key_auth(): + """Bypass auth for tests so /v1/audio/speech doesn't require a real key.""" + return MagicMock() + + +def _make_mock_tts_response(): + """Mock response that simulates HttpxBinaryResponseContent with aiter_bytes.""" + + async def _chunks(): + yield b"\xff\xfb" + + def _aiter_bytes(chunk_size=8192): + async def _wrapper(): + return _chunks() + + return _wrapper() + + inner = MagicMock() + inner.aiter_bytes = _aiter_bytes + inner._hidden_params = {} + + async def _resolver(): + return inner + + return _resolver() + + +@pytest.fixture +def client_no_auth(): + from litellm.proxy.proxy_server import cleanup_router_config_variables + + cleanup_router_config_variables() + filepath = os.path.dirname(os.path.abspath(__file__)) + config_fp = os.path.join(filepath, "test_configs", "test_config_no_auth.yaml") + asyncio.run(initialize(config=config_fp, debug=True)) + return TestClient(app) + + +@pytest.mark.asyncio +@pytest.mark.retry(retries=0) +async def test_audio_speech_success_does_not_call_post_call_success_hook( + client_no_auth, +): + """TTS success path must NOT call post_call_success_hook. + + TTS returns a streaming binary response (HttpxBinaryResponseContent) which + is not in LLMResponseTypes. Prometheus metrics for successful requests are + tracked at the litellm level via async_log_success_event, not here. + """ + mock_success_hook = AsyncMock() + mock_failure_hook = AsyncMock() + mock_pre_call = AsyncMock(side_effect=lambda *, data, **kw: data) + mock_update_status = AsyncMock() + + mock_logging = MagicMock() + mock_logging.post_call_success_hook = mock_success_hook + mock_logging.post_call_failure_hook = mock_failure_hook + mock_logging.pre_call_hook = mock_pre_call + mock_logging.update_request_status = mock_update_status + + async def _mock_route_request(*, data, route_type, llm_router, user_model): + assert route_type == "aspeech" + return _make_mock_tts_response() + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = _mock_user_api_key_auth + try: + with ( + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging), + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_mock_route_request, + ), + ): + response = client_no_auth.post( + "/v1/audio/speech", + json={"model": "tts-1", "input": "hello"}, + headers={"Content-Type": "application/json"}, + ) + finally: + app.dependency_overrides = original_overrides + + assert response.status_code == 200 + mock_success_hook.assert_not_called() + mock_failure_hook.assert_not_called() + + +@pytest.mark.asyncio +@pytest.mark.retry(retries=0) +async def test_audio_speech_failure_calls_post_call_failure_hook(client_no_auth): + """TTS failure path must call proxy_logging_obj.post_call_failure_hook (Prometheus failed requests).""" + mock_success_hook = AsyncMock() + mock_failure_hook = AsyncMock() + mock_pre_call = AsyncMock(side_effect=lambda *, data, **kw: data) + + mock_logging = MagicMock() + mock_logging.post_call_success_hook = mock_success_hook + mock_logging.post_call_failure_hook = mock_failure_hook + mock_logging.pre_call_hook = mock_pre_call + + async def _mock_route_request_raise(*, data, route_type, llm_router, user_model): + raise ValueError("mock rate limit") + + original_overrides = app.dependency_overrides.copy() + app.dependency_overrides[user_api_key_auth] = _mock_user_api_key_auth + # Don't re-raise server exceptions so we get the 500 response instead of ValueError + client = TestClient(app, raise_server_exceptions=False) + try: + with ( + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging), + patch( + "litellm.proxy.proxy_server.route_request", + side_effect=_mock_route_request_raise, + ), + ): + response = client.post( + "/v1/audio/speech", + json={"model": "tts-1", "input": "hello"}, + headers={"Content-Type": "application/json"}, + ) + finally: + app.dependency_overrides = original_overrides + + assert response.status_code == 500 + mock_failure_hook.assert_awaited_once() + mock_success_hook.assert_not_called() + call_kw = mock_failure_hook.call_args.kwargs + assert "user_api_key_dict" in call_kw and "original_exception" in call_kw diff --git a/tests/test_litellm/proxy/test_health_check_max_tokens.py b/tests/test_litellm/proxy/test_health_check_max_tokens.py new file mode 100644 index 00000000000..e26f7fb9f20 --- /dev/null +++ b/tests/test_litellm/proxy/test_health_check_max_tokens.py @@ -0,0 +1,75 @@ +import pytest +from litellm.proxy.health_check import _update_litellm_params_for_health_check +from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers +from unittest.mock import AsyncMock, patch, MagicMock + + +@pytest.mark.asyncio +async def test_update_litellm_params_max_tokens_default(): + """ + Test that max_tokens defaults to 1 for non-wildcard models. + """ + model_info = {} + litellm_params = {"model": "gpt-4"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated_params["max_tokens"] == 1 + + +@pytest.mark.asyncio +async def test_update_litellm_params_max_tokens_custom(): + """ + Test that max_tokens respects health_check_max_tokens from model_info. + """ + model_info = {"health_check_max_tokens": 5} + litellm_params = {"model": "gpt-4"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + assert updated_params["max_tokens"] == 5 + + +@pytest.mark.asyncio +async def test_update_litellm_params_max_tokens_wildcard(): + """ + Test that max_tokens does NOT default to 1 for wildcard models. + """ + model_info = {} + litellm_params = {"model": "openai/*"} + + updated_params = _update_litellm_params_for_health_check(model_info, litellm_params) + + # Should not be set to 1 + assert "max_tokens" not in updated_params or updated_params["max_tokens"] != 1 + + +@pytest.mark.asyncio +async def test_ahealth_check_wildcard_models_respects_max_tokens(): + """ + Test that ahealth_check_wildcard_models respects max_tokens if passed, + otherwise defaults to 10. + """ + with patch( + "litellm.litellm_core_utils.llm_request_utils.pick_cheapest_chat_models_from_llm_provider", + return_value=["gpt-4o-mini"], + ), patch("litellm.acompletion", new_callable=AsyncMock): + # Test Case 1: No max_tokens passed, should default to 10 + model_params = {} + await HealthCheckHelpers.ahealth_check_wildcard_models( + model="openai/*", + custom_llm_provider="openai", + model_params=model_params, + litellm_logging_obj=MagicMock(), + ) + assert model_params["max_tokens"] == 10 + + # Test Case 2: Custom health_check_max_tokens passed via model_params, should be respected + model_params = {"max_tokens": 3} + await HealthCheckHelpers.ahealth_check_wildcard_models( + model="openai/*", + custom_llm_provider="openai", + model_params=model_params, + litellm_logging_obj=MagicMock(), + ) + assert model_params["max_tokens"] == 3 diff --git a/tests/test_litellm/proxy/test_update_llm_router_resilience.py b/tests/test_litellm/proxy/test_update_llm_router_resilience.py new file mode 100644 index 00000000000..0ee865ab48c --- /dev/null +++ b/tests/test_litellm/proxy/test_update_llm_router_resilience.py @@ -0,0 +1,173 @@ +""" +Test that _update_llm_router and _delete_deployment are resilient to +config loading failures (e.g. database timeouts). + +This addresses a bug where httpcore.ReadTimeout from the Prisma client +during get_config() would prevent ALL DB models from loading into the +router, because the exception propagated up and was caught by the +catch-all handler in _update_llm_router. +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + +from litellm.proxy.proxy_server import ProxyConfig + + +def _make_db_model(model_name: str, model_id: str): + """Helper to create a mock DB model record.""" + record = MagicMock() + record.model_id = model_id + record.model_name = model_name + record.litellm_params = {"model": model_name} + record.model_info = {"id": model_id} + record.created_by = "default_user_id" + record.created_at = None + record.updated_at = None + record.updated_by = None + return record + + +class TestUpdateLlmRouterResilience: + """Test _update_llm_router handles get_config failures gracefully.""" + + @pytest.mark.asyncio + async def test_models_loaded_when_get_config_times_out(self): + """DB models should still be added to the router when get_config() raises a timeout.""" + proxy_config = ProxyConfig() + + db_models = [_make_db_model("gpt-5.1", "db-id-1")] + + mock_router = MagicMock() + mock_router.get_model_list.return_value = [] + mock_router.get_model_ids.return_value = [] + + mock_proxy_logging = MagicMock() + + with ( + patch.object( + proxy_config, + "get_config", + new_callable=AsyncMock, + side_effect=Exception("httpcore.ReadTimeout"), + ), + patch.object(proxy_config, "_add_deployment", return_value=1) as mock_add, + patch.object( + proxy_config, + "_delete_deployment", + new_callable=AsyncMock, + return_value=0, + ), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.master_key", "sk-test"), + patch("litellm.proxy.proxy_server.llm_model_list", []), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + await proxy_config._update_llm_router( + new_models=db_models, + proxy_logging_obj=mock_proxy_logging, + ) + + # _add_deployment should still have been called despite get_config failure + mock_add.assert_called_once_with(db_models=db_models) + + @pytest.mark.asyncio + async def test_get_config_success_still_works(self): + """Normal flow should still work when get_config succeeds.""" + proxy_config = ProxyConfig() + + db_models = [_make_db_model("gpt-5.1", "db-id-1")] + + mock_router = MagicMock() + mock_router.get_model_list.return_value = [] + mock_router.get_model_ids.return_value = [] + + mock_proxy_logging = MagicMock() + + with ( + patch.object( + proxy_config, + "get_config", + new_callable=AsyncMock, + return_value={"model_list": []}, + ), + patch.object(proxy_config, "_add_deployment", return_value=1) as mock_add, + patch.object( + proxy_config, + "_delete_deployment", + new_callable=AsyncMock, + return_value=0, + ), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.master_key", "sk-test"), + patch("litellm.proxy.proxy_server.llm_model_list", []), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + await proxy_config._update_llm_router( + new_models=db_models, + proxy_logging_obj=mock_proxy_logging, + ) + + mock_add.assert_called_once_with(db_models=db_models) + + +class TestDeleteDeploymentResilience: + """Test _delete_deployment handles get_config failures gracefully.""" + + @pytest.mark.asyncio + async def test_returns_zero_when_get_config_times_out(self): + """Should return 0 (no deletions) when get_config fails, not raise.""" + proxy_config = ProxyConfig() + + db_models = [_make_db_model("gpt-5.1", "db-id-1")] + + mock_router = MagicMock() + mock_router.get_model_ids.return_value = ["db-id-1", "config-id-1"] + + with ( + patch.object( + proxy_config, + "get_config", + new_callable=AsyncMock, + side_effect=Exception("httpcore.ReadTimeout"), + ), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", False), + ): + result = await proxy_config._delete_deployment(db_models=db_models) + + # Should safely return 0 instead of raising + assert result == 0 + # Should NOT have deleted any deployments + mock_router.delete_deployment.assert_not_called() + + @pytest.mark.asyncio + async def test_normal_delete_still_works(self): + """Normal deletion should work when get_config succeeds.""" + proxy_config = ProxyConfig() + + db_models = [_make_db_model("gpt-5.1", "db-id-1")] + + mock_router = MagicMock() + # Router has a model ID that's not in DB or config -> should be deleted + mock_router.get_model_ids.return_value = ["db-id-1", "stale-id"] + mock_router.delete_deployment.return_value = True + mock_router._generate_model_id = MagicMock(return_value="config-id-1") + + with ( + patch.object( + proxy_config, + "get_config", + new_callable=AsyncMock, + return_value={"model_list": [ + {"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}, "model_info": {"id": "config-id-1"}} + ]}, + ), + patch("litellm.proxy.proxy_server.llm_router", mock_router), + patch("litellm.proxy.proxy_server.premium_user", False), + ): + result = await proxy_config._delete_deployment(db_models=db_models) + + # "stale-id" should have been deleted (not in db_models or config) + assert result == 1 + mock_router.delete_deployment.assert_called_once_with(id="stale-id") diff --git a/tests/test_litellm/test_router/test_enforce_model_rate_limits.py b/tests/test_litellm/test_router/test_enforce_model_rate_limits.py index 3bca3df4e1d..1def253ac93 100644 --- a/tests/test_litellm/test_router/test_enforce_model_rate_limits.py +++ b/tests/test_litellm/test_router/test_enforce_model_rate_limits.py @@ -5,12 +5,14 @@ This feature allows users to enforce TPM/RPM limits set on model deployments regardless of the routing strategy being used. """ +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest import litellm from litellm import Router +from litellm.caching.dual_cache import DualCache from litellm.router_utils.pre_call_checks.model_rate_limit_check import ( ModelRateLimitingCheck, ) @@ -88,7 +90,7 @@ class TestModelRateLimitingCheck: def test_pre_call_check_raises_rate_limit_error_when_over_rpm(self): """Test that RateLimitError is raised when RPM limit is exceeded.""" mock_cache = MagicMock() - mock_cache.get_cache.return_value = 10 # Already at limit + mock_cache.increment_cache.return_value = 11 # Over limit after increment check = ModelRateLimitingCheck(dual_cache=mock_cache) @@ -103,12 +105,11 @@ class TestModelRateLimitingCheck: check.pre_call_check(deployment) assert "RPM limit=10" in str(exc_info.value) - assert "current usage=10" in str(exc_info.value) + assert "current usage=11" in str(exc_info.value) def test_pre_call_check_allows_request_under_limit(self): """Test that requests are allowed when under the limit.""" mock_cache = MagicMock() - mock_cache.get_cache.return_value = 5 mock_cache.increment_cache.return_value = 6 check = ModelRateLimitingCheck(dual_cache=mock_cache) @@ -188,7 +189,8 @@ class TestModelRateLimitingCheckAsync: async def test_async_pre_call_check_raises_rate_limit_error_when_over_rpm(self): """Test that RateLimitError is raised when RPM limit is exceeded (async).""" mock_cache = MagicMock() - mock_cache.async_get_cache = AsyncMock(return_value=10) # Already at limit + mock_cache.async_get_cache = AsyncMock(return_value=None) + mock_cache.async_increment_cache = AsyncMock(return_value=11) # Over limit check = ModelRateLimitingCheck(dual_cache=mock_cache) @@ -208,7 +210,7 @@ class TestModelRateLimitingCheckAsync: async def test_async_pre_call_check_allows_request_under_limit(self): """Test that requests are allowed when under the limit (async).""" mock_cache = MagicMock() - mock_cache.async_get_cache = AsyncMock(return_value=5) + mock_cache.async_get_cache = AsyncMock(return_value=None) mock_cache.async_increment_cache = AsyncMock(return_value=6) check = ModelRateLimitingCheck(dual_cache=mock_cache) @@ -313,3 +315,41 @@ class TestRouterWithEnforceModelRateLimits: break assert found, "ModelRateLimitingCheck should be in litellm.callbacks" + + +class TestModelRateLimitConcurrency: + """Test that RPM rate limiting is atomic under concurrent requests.""" + + @pytest.mark.asyncio + async def test_concurrent_requests_respect_rpm_limit(self): + """ + Fire 4 concurrent async requests with RPM limit of 2. + Exactly 2 should succeed and 2 should raise RateLimitError. + + This test validates the atomic increment-first pattern: + the old check-then-increment pattern would let 3+ through + due to a race condition on the local cache read. + """ + dual_cache = DualCache() + check = ModelRateLimitingCheck(dual_cache=dual_cache) + + deployment = { + "rpm": 2, + "litellm_params": {"model": "gpt-4"}, + "model_info": {"id": "concurrent-test-id"}, + "model_name": "test-model", + } + + async def attempt_request(): + return await check.async_pre_call_check(deployment) + + results = await asyncio.gather( + *[attempt_request() for _ in range(4)], + return_exceptions=True, + ) + + successes = [r for r in results if not isinstance(r, Exception)] + failures = [r for r in results if isinstance(r, litellm.RateLimitError)] + + assert len(successes) == 2, f"Expected 2 successes, got {len(successes)}" + assert len(failures) == 2, f"Expected 2 rate limit errors, got {len(failures)}" diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts index 1643412d1e9..80cb69495da 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.test.ts @@ -50,6 +50,7 @@ const mockKeys: KeyResponse[] = [ config: {}, user_id: "user-1", team_id: null, + project_id: null, max_parallel_requests: 10, metadata: {}, tpm_limit: 1000, @@ -105,6 +106,7 @@ const mockKeys: KeyResponse[] = [ config: {}, user_id: "user-2", team_id: "team-1", + project_id: "project-1", max_parallel_requests: 5, metadata: {}, tpm_limit: 500, @@ -396,6 +398,76 @@ describe("useKeys", () => { }, ); }); + + it("should pass projectID filter to the API", async () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => ({ + keys: [mockKeys[1]], + total_count: 1, + current_page: 1, + total_pages: 1, + }), + }); + + const { result } = renderHook( + () => useKeys(1, 10, { projectID: "project-1" }), + { wrapper }, + ); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + }); + + const callUrl = mockFetch.mock.calls[0][0]; + expect(callUrl).toContain("project_id=project-1"); + expect(result.current.data?.keys).toHaveLength(1); + expect(result.current.data?.keys[0].project_id).toBe("project-1"); + }); + + it("should pass both projectID and teamID filters to the API", async () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => ({ + keys: [mockKeys[1]], + total_count: 1, + current_page: 1, + total_pages: 1, + }), + }); + + const { result } = renderHook( + () => useKeys(1, 10, { projectID: "project-1", teamID: "team-1" }), + { wrapper }, + ); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + }); + + const callUrl = mockFetch.mock.calls[0][0]; + expect(callUrl).toContain("project_id=project-1"); + expect(callUrl).toContain("team_id=team-1"); + }); + + it("should not include project_id param when projectID is null", async () => { + mockFetch.mockResolvedValueOnce({ + ok: true, + json: async () => mockKeysResponse, + }); + + const { result } = renderHook( + () => useKeys(1, 10, { projectID: null }), + { wrapper }, + ); + + await waitFor(() => { + expect(result.current.isLoading).toBe(false); + }); + + const callUrl = mockFetch.mock.calls[0][0]; + expect(callUrl).not.toContain("project_id"); + }); }); describe("useDeletedKeys", () => { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts index cf477a2e556..fbe5eccb75a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/keys/useKeys.ts @@ -33,6 +33,7 @@ export interface DeletedKeysResponse { export interface KeyListCallOptions { organizationID?: string | null; teamID?: string | null; + projectID?: string | null; selectedKeyAlias?: string | null; userID?: string | null; keyHash?: string | null; @@ -57,6 +58,7 @@ const keyListCall = async ( const params = new URLSearchParams( Object.entries({ team_id: options.teamID, + project_id: options.projectID, organization_id: options.organizationID, key_alias: options.selectedKeyAlias, key_hash: options.keyHash, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjectDetails.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjectDetails.ts new file mode 100644 index 00000000000..1d35ac1bf70 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useProjectDetails.ts @@ -0,0 +1,63 @@ +import { useQuery, useQueryClient } from "@tanstack/react-query"; +import { + getProxyBaseUrl, + getGlobalLitellmHeaderName, + deriveErrorMessage, + handleError, +} from "@/components/networking"; +import { all_admin_roles } from "@/utils/roles"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { ProjectResponse, projectKeys } from "./useProjects"; + +// ── Fetch function ─────────────────────────────────────────────────────────── + +const fetchProjectDetails = async ( + accessToken: string, + projectId: string, +): Promise => { + const baseUrl = getProxyBaseUrl(); + const url = `${baseUrl}/project/info?project_id=${encodeURIComponent(projectId)}`; + + const response = await fetch(url, { + method: "GET", + headers: { + [getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return response.json(); +}; + +// ── Hook ───────────────────────────────────────────────────────────────────── + +export const useProjectDetails = (projectId?: string) => { + const { accessToken, userRole } = useAuthorized(); + const queryClient = useQueryClient(); + + return useQuery({ + queryKey: projectKeys.detail(projectId!), + queryFn: async () => fetchProjectDetails(accessToken!, projectId!), + enabled: + Boolean(accessToken && projectId) && + all_admin_roles.includes(userRole || ""), + + // Seed from the list cache when available + initialData: () => { + if (!projectId) return undefined; + + const projects = queryClient.getQueryData( + projectKeys.list({}), + ); + + return projects?.find((p) => p.project_id === projectId); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useUpdateProject.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useUpdateProject.ts new file mode 100644 index 00000000000..e6cd3071f5f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/projects/useUpdateProject.ts @@ -0,0 +1,75 @@ +import { useMutation, useQueryClient } from "@tanstack/react-query"; +import { + getProxyBaseUrl, + getGlobalLitellmHeaderName, + deriveErrorMessage, + handleError, +} from "@/components/networking"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { ProjectResponse, projectKeys } from "./useProjects"; + +// ── Types ──────────────────────────────────────────────────────────────────── + +export interface ProjectUpdateParams { + project_alias?: string; + description?: string; + team_id?: string; + models?: string[]; + max_budget?: number; + blocked?: boolean; + metadata?: Record; + model_rpm_limit?: Record; + model_tpm_limit?: Record; +} + +// ── Fetch function ─────────────────────────────────────────────────────────── + +const updateProject = async ( + accessToken: string, + projectId: string, + params: ProjectUpdateParams, +): Promise => { + const baseUrl = getProxyBaseUrl(); + const url = `${baseUrl}/project/update`; + + const response = await fetch(url, { + method: "POST", + headers: { + [getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`, + "Content-Type": "application/json", + }, + body: JSON.stringify({ project_id: projectId, ...params }), + }); + + if (!response.ok) { + const errorData = await response.json(); + const errorMessage = deriveErrorMessage(errorData); + handleError(errorMessage); + throw new Error(errorMessage); + } + + return response.json(); +}; + +// ── Hook ───────────────────────────────────────────────────────────────────── + +export const useUpdateProject = () => { + const { accessToken } = useAuthorized(); + const queryClient = useQueryClient(); + + return useMutation< + ProjectResponse, + Error, + { projectId: string; params: ProjectUpdateParams } + >({ + mutationFn: async ({ projectId, params }) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return updateProject(accessToken, projectId, params); + }, + onSuccess: () => { + queryClient.invalidateQueries({ queryKey: projectKeys.all }); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/add_margin_form.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/add_margin_form.test.tsx new file mode 100644 index 00000000000..d61f0987ac2 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/add_margin_form.test.tsx @@ -0,0 +1,148 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../tests/test-utils"; +import AddMarginForm from "./add_margin_form"; +import { MarginConfig } from "./types"; + +vi.mock("../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, + marginType: "percentage" as const, + percentageValue: "", + fixedAmountValue: "", + onProviderChange: vi.fn(), + onMarginTypeChange: vi.fn(), + onPercentageChange: vi.fn(), + onFixedAmountChange: vi.fn(), + onAddProvider: vi.fn(), +}; + +describe("AddMarginForm", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render", () => { + renderWithProviders(); + expect(screen.getByRole("button", { name: /add provider margin/i })).toBeInTheDocument(); + }); + + it("should show the percentage input when marginType is percentage", () => { + renderWithProviders(); + expect(screen.getByPlaceholderText("10")).toBeInTheDocument(); + }); + + it("should show the fixed amount input when marginType is fixed", () => { + renderWithProviders(); + expect(screen.getByPlaceholderText("0.001")).toBeInTheDocument(); + }); + + it("should not show the fixed amount input when marginType is percentage", () => { + renderWithProviders(); + expect(screen.queryByPlaceholderText("0.001")).not.toBeInTheDocument(); + }); + + it("should not show the percentage input when marginType is fixed", () => { + renderWithProviders(); + expect(screen.queryByPlaceholderText("10")).not.toBeInTheDocument(); + }); + + it("should show the Percentage-based and Fixed Amount radio options", () => { + renderWithProviders(); + expect(screen.getByText("Percentage-based")).toBeInTheDocument(); + expect(screen.getByText("Fixed Amount")).toBeInTheDocument(); + }); + + it("should disable the submit button when no provider is selected (percentage mode)", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider margin/i })).toBeDisabled(); + }); + + it("should disable the submit button when provider is selected but no percentage value (percentage mode)", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider margin/i })).toBeDisabled(); + }); + + it("should enable the submit button when provider and percentage value are both provided", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider margin/i })).not.toBeDisabled(); + }); + + it("should disable the submit button in fixed mode when no fixed amount is provided", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider margin/i })).toBeDisabled(); + }); + + it("should enable the submit button in fixed mode when provider and fixed amount are provided", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider margin/i })).not.toBeDisabled(); + }); + + it("should call onAddProvider when the enabled submit button is clicked", async () => { + const onAddProvider = vi.fn(); + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /add provider margin/i })); + expect(onAddProvider).toHaveBeenCalledTimes(1); + }); + + it("should call onMarginTypeChange when the Fixed Amount radio is clicked", async () => { + const onMarginTypeChange = vi.fn(); + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByText("Fixed Amount")); + expect(onMarginTypeChange).toHaveBeenCalledWith("fixed"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/add_provider_form.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/add_provider_form.test.tsx new file mode 100644 index 00000000000..611c8609c36 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/add_provider_form.test.tsx @@ -0,0 +1,98 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen } from "@testing-library/react"; +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("../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 = { + discountConfig: {} as DiscountConfig, + selectedProvider: undefined, + newDiscount: "", + onProviderChange: vi.fn(), + onDiscountChange: vi.fn(), + onAddProvider: vi.fn(), +}; + +describe("AddProviderForm", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render", () => { + renderWithProviders(); + expect(screen.getByRole("button", { name: /add provider discount/i })).toBeInTheDocument(); + }); + + it("should render the discount percentage input field", () => { + renderWithProviders(); + expect(screen.getByPlaceholderText("5")).toBeInTheDocument(); + }); + + it("should disable the submit button when no provider is selected and no discount is entered", () => { + renderWithProviders(); + expect(screen.getByRole("button", { name: /add provider discount/i })).toBeDisabled(); + }); + + it("should disable the submit button when a provider is selected but no discount is entered", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider discount/i })).toBeDisabled(); + }); + + it("should disable the submit button when a discount is entered but no provider is selected", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider discount/i })).toBeDisabled(); + }); + + it("should enable the submit button when both a provider and a discount value are provided", () => { + renderWithProviders( + + ); + expect(screen.getByRole("button", { name: /add provider discount/i })).not.toBeDisabled(); + }); + + it("should call onAddProvider when the enabled submit button is clicked", async () => { + const onAddProvider = vi.fn(); + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /add provider discount/i })); + expect(onAddProvider).toHaveBeenCalledTimes(1); + }); + + it("should show the percent sign next to the discount input", () => { + renderWithProviders(); + expect(screen.getByText("%")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/cost_tracking_settings.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/cost_tracking_settings.test.tsx new file mode 100644 index 00000000000..db6899ba17f --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/cost_tracking_settings.test.tsx @@ -0,0 +1,201 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../tests/test-utils"; +import CostTrackingSettings from "./cost_tracking_settings"; + +// Mock sub-hooks so we can control their state without network calls +const mockDiscountConfig = vi.fn(() => ({})); +const mockMarginConfig = vi.fn(() => ({})); + +vi.mock("./use_discount_config", () => ({ + useDiscountConfig: () => ({ + discountConfig: mockDiscountConfig(), + fetchDiscountConfig: vi.fn().mockResolvedValue(undefined), + handleAddProvider: vi.fn().mockResolvedValue(true), + handleRemoveProvider: vi.fn().mockResolvedValue(undefined), + handleDiscountChange: vi.fn().mockResolvedValue(undefined), + }), +})); + +vi.mock("./use_margin_config", () => ({ + useMarginConfig: () => ({ + marginConfig: mockMarginConfig(), + fetchMarginConfig: vi.fn().mockResolvedValue(undefined), + handleAddMargin: vi.fn().mockResolvedValue(true), + handleRemoveMargin: vi.fn().mockResolvedValue(undefined), + handleMarginChange: vi.fn().mockResolvedValue(undefined), + }), +})); + +vi.mock("./pricing_calculator/index", () => ({ + default: () => Pricing Calculator, +})); + +vi.mock("../playground/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn().mockResolvedValue([]), +})); + +vi.mock("../HelpLink", () => ({ + DocsMenu: () => null, +})); + +vi.mock("./how_it_works", () => ({ + default: () => How It Works, +})); + +vi.mock("../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(), +})); + +const ADMIN_PROPS = { + userID: "user-1", + userRole: "proxy_admin", + accessToken: "test-token", +}; + +describe("CostTrackingSettings", () => { + beforeEach(() => { + vi.clearAllMocks(); + mockDiscountConfig.mockReturnValue({}); + mockMarginConfig.mockReturnValue({}); + }); + + it("should return nothing when accessToken is null", () => { + const { container } = renderWithProviders( + + ); + expect(container.firstChild).toBeNull(); + }); + + it("should render the page title", () => { + renderWithProviders(); + expect(screen.getByText("Cost Tracking Settings")).toBeInTheDocument(); + }); + + it("should show the Provider Discounts accordion header for proxy_admin", () => { + renderWithProviders(); + expect(screen.getByText("Provider Discounts")).toBeInTheDocument(); + }); + + it("should show the Fee/Price Margin accordion header for proxy_admin", () => { + renderWithProviders(); + expect(screen.getByText("Fee/Price Margin")).toBeInTheDocument(); + }); + + it("should always show the Pricing Calculator section", () => { + renderWithProviders(); + // The accordion header text appears in the DOM; getAllByText tolerates duplicates + expect(screen.getAllByText("Pricing Calculator").length).toBeGreaterThan(0); + }); + + it("should show the pricing calculator component", async () => { + renderWithProviders(); + expect(await screen.findByTestId("pricing-calculator")).toBeInTheDocument(); + }); + + it("should not show Provider Discounts section for a non-admin role", () => { + renderWithProviders( + + ); + expect(screen.queryByText("Provider Discounts")).not.toBeInTheDocument(); + }); + + it("should not show Fee/Price Margin section for a non-admin role", () => { + renderWithProviders( + + ); + expect(screen.queryByText("Fee/Price Margin")).not.toBeInTheDocument(); + }); + + it("should show Provider Discounts for the 'Admin' role as well", () => { + renderWithProviders( + + ); + expect(screen.getByText("Provider Discounts")).toBeInTheDocument(); + }); + + it("should show the subtitle describing discount/margin configuration", () => { + renderWithProviders(); + expect( + screen.getByText(/configure cost discounts and margins/i) + ).toBeInTheDocument(); + }); + + describe("Add Provider Discount modal", () => { + it("should open the Add Provider Discount modal when the button is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + // The button lives inside the Provider Discounts accordion — click the header to expand first + const accordionHeader = screen.getByText("Provider Discounts").closest("button"); + if (accordionHeader) { + await user.click(accordionHeader); + } + + const addButton = await screen.findByRole("button", { name: /add provider discount/i }); + await user.click(addButton); + + expect( + await screen.findByText("Add Provider Discount", { selector: "h2" }) + ).toBeInTheDocument(); + }); + }); + + describe("Add Provider Margin modal", () => { + it("should open the Add Provider Margin modal when the button is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const accordionHeader = screen.getByText("Fee/Price Margin").closest("button"); + if (accordionHeader) { + await user.click(accordionHeader); + } + + const addButton = await screen.findByRole("button", { name: /add provider margin/i }); + await user.click(addButton); + + expect( + await screen.findByText("Add Provider Margin", { selector: "h2" }) + ).toBeInTheDocument(); + }); + }); + + describe("empty state messages", () => { + it("should show the empty state message when no discount config is loaded", async () => { + mockDiscountConfig.mockReturnValue({}); + renderWithProviders(); + + const accordionHeader = screen.getByText("Provider Discounts").closest("button"); + if (accordionHeader) { + await userEvent.setup().click(accordionHeader); + } + + expect( + await screen.findByText(/no provider discounts configured/i) + ).toBeInTheDocument(); + }); + + it("should show the empty state message when no margin config is loaded", async () => { + mockMarginConfig.mockReturnValue({}); + renderWithProviders(); + + const accordionHeader = screen.getByText("Fee/Price Margin").closest("button"); + if (accordionHeader) { + await userEvent.setup().click(accordionHeader); + } + + expect( + await screen.findByText(/no provider margins configured/i) + ).toBeInTheDocument(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/how_it_works.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/how_it_works.test.tsx new file mode 100644 index 00000000000..fa608f555ce --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/how_it_works.test.tsx @@ -0,0 +1,95 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../tests/test-utils"; +import HowItWorks from "./how_it_works"; + +vi.mock("@/app/(dashboard)/api-reference/components/CodeBlock", () => ({ + default: ({ code }: { code: string }) => {code}, +})); + +describe("HowItWorks", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render", () => { + renderWithProviders(); + expect(screen.getByText("Cost Calculation")).toBeInTheDocument(); + }); + + it("should display the cost calculation formula", () => { + renderWithProviders(); + expect(screen.getByText(/final_cost = base_cost/i)).toBeInTheDocument(); + }); + + it("should display the valid range information", () => { + renderWithProviders(); + expect(screen.getByText(/0% and 100%/i)).toBeInTheDocument(); + }); + + it("should render the code block with a curl example", () => { + renderWithProviders(); + expect(screen.getByTestId("code-block")).toBeInTheDocument(); + expect(screen.getByTestId("code-block").textContent).toContain("curl"); + }); + + it("should show the response header names for discount verification", () => { + renderWithProviders(); + expect(screen.getByText("x-litellm-response-cost")).toBeInTheDocument(); + expect(screen.getByText("x-litellm-response-cost-original")).toBeInTheDocument(); + expect(screen.getByText("x-litellm-response-cost-discount-amount")).toBeInTheDocument(); + }); + + it("should not show calculated results initially when no input is provided", () => { + renderWithProviders(); + expect(screen.queryByText("Calculated Results")).not.toBeInTheDocument(); + }); + + it("should not show calculated results when only response cost is entered", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const responseCostInput = screen.getByPlaceholderText("0.0171938125"); + await user.type(responseCostInput, "0.01"); + + expect(screen.queryByText("Calculated Results")).not.toBeInTheDocument(); + }); + + it("should not show calculated results when only discount amount is entered", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const discountAmountInput = screen.getByPlaceholderText("0.0009049375"); + await user.type(discountAmountInput, "0.001"); + + expect(screen.queryByText("Calculated Results")).not.toBeInTheDocument(); + }); + + it("should show calculated results when both fields are filled", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + const responseCostInput = screen.getByPlaceholderText("0.0171938125"); + const discountAmountInput = screen.getByPlaceholderText("0.0009049375"); + + await user.type(responseCostInput, "0.0171938125"); + await user.type(discountAmountInput, "0.0009049375"); + + expect(await screen.findByText("Calculated Results")).toBeInTheDocument(); + }); + + it("should show original cost, final cost, and discount amount in results", async () => { + const user = userEvent.setup(); + renderWithProviders(); + + await user.type(screen.getByPlaceholderText("0.0171938125"), "0.0171938125"); + await user.type(screen.getByPlaceholderText("0.0009049375"), "0.0009049375"); + + expect(await screen.findByText("Original Cost:")).toBeInTheDocument(); + expect(screen.getByText("Final Cost:")).toBeInTheDocument(); + expect(screen.getByText("Discount Amount:")).toBeInTheDocument(); + expect(screen.getByText("Discount Applied:")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_discount_table.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_discount_table.test.tsx new file mode 100644 index 00000000000..7697c6e7686 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_discount_table.test.tsx @@ -0,0 +1,241 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { screen, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { renderWithProviders } from "../../../tests/test-utils"; +import ProviderDiscountTable from "./provider_discount_table"; + +vi.mock("@heroicons/react/outline", () => ({ + TrashIcon: function TrashIcon() { return null; }, + PencilAltIcon: function PencilAltIcon() { return null; }, + CheckIcon: function CheckIcon() { return null; }, + XIcon: function XIcon() { return null; }, +})); + +vi.mock("@tremor/react", () => ({ + Table: ({ children }: any) => {children}, + TableHead: ({ children }: any) => {children}, + TableRow: ({ children }: any) => {children}, + TableHeaderCell: ({ children }: any) => {children}, + TableBody: ({ children }: any) => {children}, + TableCell: ({ children }: any) => {children}, + Text: ({ children }: any) => {children}, + TextInput: ({ value, onValueChange, onKeyDown, placeholder, ...rest }: any) => ( + onValueChange?.(e.target.value)} + onKeyDown={onKeyDown} + placeholder={placeholder} + {...rest} + /> + ), + Icon: ({ icon: IconComponent, onClick }: any) => { + const name = IconComponent?.displayName ?? IconComponent?.name ?? "icon"; + return ; + }, +})); + +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, +}; + +describe("ProviderDiscountTable", () => { + const onDiscountChange = vi.fn(); + const onRemoveProvider = vi.fn(); + + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render", () => { + renderWithProviders( + + ); + expect(screen.getByRole("table")).toBeInTheDocument(); + }); + + it("should render the table headers", () => { + renderWithProviders( + + ); + expect(screen.getByText("Provider")).toBeInTheDocument(); + expect(screen.getByText("Discount Percentage")).toBeInTheDocument(); + expect(screen.getByText("Actions")).toBeInTheDocument(); + }); + + it("should display provider display names in the table", () => { + renderWithProviders( + + ); + expect(screen.getByText("OpenAI")).toBeInTheDocument(); + }); + + it("should display the formatted discount percentage", () => { + renderWithProviders( + + ); + expect(screen.getByText("5.0%")).toBeInTheDocument(); + }); + + it("should show a text input when the edit icon is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + const pencilButton = screen.getByRole("button", { name: /PencilAltIcon/i }); + await user.click(pencilButton); + + expect(screen.getByPlaceholderText("5")).toBeInTheDocument(); + }); + + it("should hide the formatted percentage when in edit mode", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + + expect(screen.queryByText("5.0%")).not.toBeInTheDocument(); + }); + + it("should call onDiscountChange with the new value when the save icon is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + + const input = screen.getByPlaceholderText("5"); + await user.clear(input); + await user.type(input, "10"); + + await user.click(screen.getByRole("button", { name: /CheckIcon/i })); + + expect(onDiscountChange).toHaveBeenCalledWith("openai", "0.1"); + }); + + it("should restore the display view after saving", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + await user.click(screen.getByRole("button", { name: /CheckIcon/i })); + + expect(screen.queryByPlaceholderText("5")).not.toBeInTheDocument(); + }); + + it("should cancel edit mode when the cancel icon is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + await user.click(screen.getByRole("button", { name: /XIcon/i })); + + expect(screen.queryByPlaceholderText("5")).not.toBeInTheDocument(); + expect(onDiscountChange).not.toHaveBeenCalled(); + expect(screen.getByText("5.0%")).toBeInTheDocument(); + }); + + it("should not call onDiscountChange when canceling edit", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + await user.click(screen.getByRole("button", { name: /XIcon/i })); + + expect(onDiscountChange).not.toHaveBeenCalled(); + }); + + it("should call onRemoveProvider with the provider key and display name when the trash icon is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /TrashIcon/i })); + + expect(onRemoveProvider).toHaveBeenCalledWith("openai", "OpenAI"); + }); + + it("should not call onDiscountChange when the entered value is out of range", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + const input = screen.getByPlaceholderText("5"); + await user.clear(input); + await user.type(input, "150"); + await user.click(screen.getByRole("button", { name: /CheckIcon/i })); + + expect(onDiscountChange).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_display_helpers.test.ts b/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_display_helpers.test.ts new file mode 100644 index 00000000000..9668f07c2c5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_display_helpers.test.ts @@ -0,0 +1,91 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { getProviderDisplayInfo, getProviderBackendValue, handleImageError } from "./provider_display_helpers"; + +vi.mock("../provider_info_helpers", () => ({ + 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"); + }); + + it("should return the backend value for another known provider", () => { + expect(getProviderBackendValue("Anthropic")).toBe("anthropic"); + }); + + it("should return null for an unknown enum key", () => { + 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/components/CostTrackingSettings/provider_margin_table.test.tsx b/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_margin_table.test.tsx new file mode 100644 index 00000000000..1ad937e6cb4 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/provider_margin_table.test.tsx @@ -0,0 +1,246 @@ +import React from "react"; +import { describe, it, expect, vi, beforeEach } from "vitest"; +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"; + +vi.mock("@heroicons/react/outline", () => ({ + TrashIcon: function TrashIcon() { return null; }, + PencilAltIcon: function PencilAltIcon() { return null; }, + CheckIcon: function CheckIcon() { return null; }, + XIcon: function XIcon() { return null; }, +})); + +vi.mock("@tremor/react", () => ({ + Table: ({ children }: any) => {children}, + TableHead: ({ children }: any) => {children}, + TableRow: ({ children }: any) => {children}, + TableHeaderCell: ({ children }: any) => {children}, + TableBody: ({ children }: any) => {children}, + TableCell: ({ children }: any) => {children}, + Text: ({ children }: any) => {children}, + TextInput: ({ value, onValueChange, placeholder, autoFocus, className }: any) => ( + onValueChange?.(e.target.value)} + placeholder={placeholder} + autoFocus={autoFocus} + className={className} + /> + ), + Icon: ({ icon: IconComponent, onClick }: any) => { + const name = IconComponent?.displayName ?? IconComponent?.name ?? "icon"; + return ; + }, +})); + +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(); + + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("should render", () => { + renderWithProviders( + + ); + expect(screen.getByRole("table")).toBeInTheDocument(); + }); + + it("should render the table headers", () => { + renderWithProviders( + + ); + expect(screen.getByText("Provider")).toBeInTheDocument(); + expect(screen.getByText("Margin")).toBeInTheDocument(); + expect(screen.getByText("Actions")).toBeInTheDocument(); + }); + + it("should display the provider display name", () => { + renderWithProviders( + + ); + expect(screen.getByText("OpenAI")).toBeInTheDocument(); + }); + + it("should display the global provider as 'Global (All Providers)'", () => { + renderWithProviders( + + ); + expect(screen.getByText("Global (All Providers)")).toBeInTheDocument(); + }); + + it("should display a numeric margin as a percentage", () => { + renderWithProviders( + + ); + expect(screen.getByText("10.0%")).toBeInTheDocument(); + }); + + it("should display a fixed amount margin with dollar sign", () => { + renderWithProviders( + + ); + expect(screen.getByText("$0.001000")).toBeInTheDocument(); + }); + + it("should display a combined percentage and fixed margin", () => { + renderWithProviders( + + ); + expect(screen.getByText(/10\.0%.*\$0\.001000/)).toBeInTheDocument(); + }); + + it("should show edit inputs when the pencil icon is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + + expect(screen.getByPlaceholderText("10")).toBeInTheDocument(); + expect(screen.getByPlaceholderText("0.001")).toBeInTheDocument(); + }); + + it("should call onMarginChange with a percentage value when save is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + + const percentInput = screen.getByPlaceholderText("10"); + await user.clear(percentInput); + await user.type(percentInput, "20"); + + await user.click(screen.getByRole("button", { name: /CheckIcon/i })); + + expect(onMarginChange).toHaveBeenCalledWith("openai", 0.2); + }); + + it("should cancel edit mode without calling onMarginChange when X is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + await user.click(screen.getByRole("button", { name: /XIcon/i })); + + expect(onMarginChange).not.toHaveBeenCalled(); + expect(screen.queryByPlaceholderText("10")).not.toBeInTheDocument(); + }); + + it("should call onRemoveProvider with provider key and display name when trash is clicked", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /TrashIcon/i })); + + expect(onRemoveProvider).toHaveBeenCalledWith("openai", "OpenAI"); + }); + + it("should call onRemoveProvider with 'Global' display name for the global provider", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /TrashIcon/i })); + + expect(onRemoveProvider).toHaveBeenCalledWith("global", "Global"); + }); + + describe("when both percentage and fixed amount are entered", () => { + it("should call onMarginChange with an object containing both values", async () => { + const user = userEvent.setup(); + renderWithProviders( + + ); + + await user.click(screen.getByRole("button", { name: /PencilAltIcon/i })); + + const percentInput = screen.getByPlaceholderText("10"); + await user.clear(percentInput); + await user.type(percentInput, "5"); + + const fixedInput = screen.getByPlaceholderText("0.001"); + await user.type(fixedInput, "0.002"); + + await user.click(screen.getByRole("button", { name: /CheckIcon/i })); + + expect(onMarginChange).toHaveBeenCalledWith("openai", { + percentage: 0.05, + fixed_amount: 0.002, + }); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/use_discount_config.test.ts b/ui/litellm-dashboard/src/components/CostTrackingSettings/use_discount_config.test.ts new file mode 100644 index 00000000000..9161d78e56e --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/use_discount_config.test.ts @@ -0,0 +1,244 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { renderHook, act } from "@testing-library/react"; +import { useDiscountConfig } from "./use_discount_config"; +import NotificationsManager from "@/components/molecules/notifications_manager"; + +vi.mock("@/components/networking", () => ({ + getProxyBaseUrl: vi.fn(() => ""), + getGlobalLitellmHeaderName: vi.fn(() => "Authorization"), +})); + +vi.mock("./provider_display_helpers", () => ({ + getProviderBackendValue: vi.fn((enumKey: string) => { + const map: Record = { + OpenAI: "openai", + Anthropic: "anthropic", + }; + return map[enumKey] ?? null; + }), +})); + +vi.mock("../provider_info_helpers", () => ({ + Providers: { + OpenAI: "OpenAI", + Anthropic: "Anthropic", + }, +})); + +describe("useDiscountConfig", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + describe("fetchDiscountConfig", () => { + it("should populate discountConfig with fetched values on success", async () => { + vi.spyOn(global, "fetch").mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.05, anthropic: 0.1 } }), + } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + expect(result.current.discountConfig).toEqual({ openai: 0.05, anthropic: 0.1 }); + }); + + it("should set an empty config when the response has no values", async () => { + vi.spyOn(global, "fetch").mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: {} }), + } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + expect(result.current.discountConfig).toEqual({}); + }); + + it("should show an error notification when the fetch throws", async () => { + vi.spyOn(global, "fetch").mockRejectedValueOnce(new Error("Network error")); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith( + expect.stringMatching(/failed to fetch/i) + ); + }); + }); + + describe("handleAddProvider", () => { + it("should return false and notify when no provider is selected", async () => { + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddProvider(undefined, "5"); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalled(); + }); + + it("should return false and notify when no discount is provided", async () => { + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddProvider("OpenAI", ""); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalled(); + }); + + it("should return false and notify when the discount exceeds 100", async () => { + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddProvider("OpenAI", "150"); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith( + expect.stringMatching(/0%.*100%/i) + ); + }); + + it("should return false and notify when the provider already exists in the config", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.05 } }), + } as Response) + .mockResolvedValue({ ok: true, json: async () => ({}) } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddProvider("OpenAI", "10"); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith( + expect.stringMatching(/already exists/i) + ); + }); + + it("should save the config and return true on a valid new provider", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ ok: true, json: async () => ({ values: {} }) } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({ values: { openai: 0.05 } }) } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddProvider("OpenAI", "5"); + }); + + expect(success!).toBe(true); + expect(NotificationsManager.success).toHaveBeenCalled(); + }); + }); + + describe("handleRemoveProvider", () => { + it("should remove the provider from the config and save", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.05, anthropic: 0.1 } }), + } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { anthropic: 0.1 } }), + } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + expect(result.current.discountConfig).toHaveProperty("openai"); + + await act(async () => { + await result.current.handleRemoveProvider("openai"); + }); + + // The optimistic update removes openai immediately + expect(result.current.discountConfig).not.toHaveProperty("openai"); + }); + }); + + describe("handleDiscountChange", () => { + it("should update the discount value and save when the value is valid", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.05 } }), + } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.1 } }), + } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + await act(async () => { + await result.current.handleDiscountChange("openai", "0.1"); + }); + + // Optimistic update applied immediately + expect(result.current.discountConfig["openai"]).toBe(0.1); + }); + + it("should not save when the value is greater than 1 (invalid fraction)", async () => { + vi.spyOn(global, "fetch").mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.05 } }), + } as Response); + + const { result } = renderHook(() => useDiscountConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchDiscountConfig(); + }); + + // Clear mocks after the initial fetch so we can check that no PATCH was made + vi.clearAllMocks(); + + await act(async () => { + await result.current.handleDiscountChange("openai", "1.5"); + }); + + expect(global.fetch).not.toHaveBeenCalled(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/CostTrackingSettings/use_margin_config.test.ts b/ui/litellm-dashboard/src/components/CostTrackingSettings/use_margin_config.test.ts new file mode 100644 index 00000000000..3363bc58d93 --- /dev/null +++ b/ui/litellm-dashboard/src/components/CostTrackingSettings/use_margin_config.test.ts @@ -0,0 +1,316 @@ +import { describe, it, expect, vi, beforeEach } from "vitest"; +import { renderHook, act } from "@testing-library/react"; +import { useMarginConfig } from "./use_margin_config"; +import NotificationsManager from "@/components/molecules/notifications_manager"; + +vi.mock("@/components/networking", () => ({ + getProxyBaseUrl: vi.fn(() => ""), + getGlobalLitellmHeaderName: vi.fn(() => "Authorization"), +})); + +vi.mock("./provider_display_helpers", () => ({ + getProviderBackendValue: vi.fn((enumKey: string) => { + const map: Record = { + OpenAI: "openai", + Anthropic: "anthropic", + }; + return map[enumKey] ?? null; + }), +})); + +vi.mock("../provider_info_helpers", () => ({ + Providers: { + OpenAI: "OpenAI", + Anthropic: "Anthropic", + }, +})); + +describe("useMarginConfig", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + describe("fetchMarginConfig", () => { + it("should populate marginConfig with fetched values on success", async () => { + vi.spyOn(global, "fetch").mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.1, global: 0.05 } }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + expect(result.current.marginConfig).toEqual({ openai: 0.1, global: 0.05 }); + }); + + it("should set an empty config when the response has no values", async () => { + vi.spyOn(global, "fetch").mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: {} }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + expect(result.current.marginConfig).toEqual({}); + }); + + it("should show an error notification when fetch throws", async () => { + vi.spyOn(global, "fetch").mockRejectedValueOnce(new Error("Network error")); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith( + expect.stringMatching(/failed to fetch/i) + ); + }); + }); + + describe("handleAddMargin", () => { + it("should return false and notify when no provider is selected", async () => { + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddMargin({ + selectedProvider: undefined, + marginType: "percentage", + percentageValue: "10", + fixedAmountValue: "", + }); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalled(); + }); + + it("should return false and notify when percentage is out of range", async () => { + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddMargin({ + selectedProvider: "OpenAI", + marginType: "percentage", + percentageValue: "2000", + fixedAmountValue: "", + }); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith( + expect.stringMatching(/0%.*1000%/i) + ); + }); + + it("should return false when the provider already has a margin configured", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.05 } }), + } as Response) + .mockResolvedValue({ ok: true, json: async () => ({}) } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddMargin({ + selectedProvider: "OpenAI", + marginType: "percentage", + percentageValue: "10", + fixedAmountValue: "", + }); + }); + + expect(success!).toBe(false); + expect(NotificationsManager.fromBackend).toHaveBeenCalledWith( + expect.stringMatching(/already exists/i) + ); + }); + + it("should save a percentage margin and return true for a valid new provider", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ ok: true, json: async () => ({ values: {} }) } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({ values: { openai: 0.1 } }) } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddMargin({ + selectedProvider: "OpenAI", + marginType: "percentage", + percentageValue: "10", + fixedAmountValue: "", + }); + }); + + expect(success!).toBe(true); + expect(NotificationsManager.success).toHaveBeenCalled(); + }); + + it("should save a fixed amount margin and return true for a valid new provider", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ ok: true, json: async () => ({ values: {} }) } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: { fixed_amount: 0.001 } } }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddMargin({ + selectedProvider: "OpenAI", + marginType: "fixed", + percentageValue: "", + fixedAmountValue: "0.001", + }); + }); + + expect(success!).toBe(true); + expect(NotificationsManager.success).toHaveBeenCalled(); + }); + + it("should accept the global provider without provider_map lookup", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ ok: true, json: async () => ({ values: {} }) } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { global: 0.05 } }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + let success: boolean; + await act(async () => { + success = await result.current.handleAddMargin({ + selectedProvider: "global", + marginType: "percentage", + percentageValue: "5", + fixedAmountValue: "", + }); + }); + + expect(success!).toBe(true); + }); + }); + + describe("handleRemoveMargin", () => { + it("should remove the provider from the config and save", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.1, anthropic: 0.05 } }), + } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { anthropic: 0.05 } }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + expect(result.current.marginConfig).toHaveProperty("openai"); + + await act(async () => { + await result.current.handleRemoveMargin("openai"); + }); + + expect(result.current.marginConfig).not.toHaveProperty("openai"); + }); + }); + + describe("handleMarginChange", () => { + it("should update the margin value for a provider and save", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.1 } }), + } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.2 } }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + await act(async () => { + await result.current.handleMarginChange("openai", 0.2); + }); + + // Optimistic update applied immediately + expect(result.current.marginConfig["openai"]).toBe(0.2); + }); + + it("should update the margin with a complex value (percentage + fixed)", async () => { + vi.spyOn(global, "fetch") + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ values: { openai: 0.1 } }), + } as Response) + .mockResolvedValueOnce({ ok: true, json: async () => ({}) } as Response) + .mockResolvedValueOnce({ + ok: true, + json: async () => ({ + values: { openai: { percentage: 0.05, fixed_amount: 0.001 } }, + }), + } as Response); + + const { result } = renderHook(() => useMarginConfig({ accessToken: "test-token" })); + + await act(async () => { + await result.current.fetchMarginConfig(); + }); + + await act(async () => { + await result.current.handleMarginChange("openai", { percentage: 0.05, fixed_amount: 0.001 }); + }); + + expect(result.current.marginConfig["openai"]).toEqual({ + percentage: 0.05, + fixed_amount: 0.001, + }); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectDetailsPage.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectDetailsPage.tsx new file mode 100644 index 00000000000..77beac65ad7 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Projects/ProjectDetailsPage.tsx @@ -0,0 +1,317 @@ +import { useProjectDetails } from "@/app/(dashboard)/hooks/projects/useProjectDetails"; +import { useTeam } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { + Button, + Card, + Col, + Descriptions, + Empty, + Flex, + Layout, + Progress, + Row, + Spin, + Tag, + theme, + Typography, +} from "antd"; +import { LoadingOutlined } from "@ant-design/icons"; +import { BarChart } from "@tremor/react"; +import { ArrowLeftIcon, DollarSignIcon, EditIcon, UsersIcon } from "lucide-react"; +import { useMemo, useState } from "react"; +import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; +import { EditProjectModal } from "./ProjectModals/EditProjectModal"; +import { ProjectKeysSection } from "./ProjectKeysSection"; + +const { Title, Text } = Typography; +const { Content } = Layout; + +interface TeamInfoShape { + team_id: string; + team_alias?: string; + models?: string[]; + max_budget?: number | null; + budget_duration?: string | null; + spend?: number; + members_with_roles?: { user_id: string; role: string }[]; +} + +interface ProjectDetailProps { + projectId: string; + onBack: () => void; +} + +export function ProjectDetail({ projectId, onBack }: ProjectDetailProps) { + const { data: project, isLoading } = useProjectDetails(projectId); + const { data: teamData } = useTeam(project?.team_id ?? undefined); + // teamInfoCall returns { team_id, team_info: {...}, keys, team_memberships } + const teamInfo: TeamInfoShape | undefined = ((teamData as unknown as { team_info?: TeamInfoShape })?.team_info ?? + teamData) as TeamInfoShape | undefined; + const { token } = theme.useToken(); + const [isEditModalVisible, setIsEditModalVisible] = useState(false); + + const spend = project?.spend ?? 0; + const maxBudget = project?.litellm_budget_table?.max_budget ?? null; + const hasLimit = maxBudget != null && maxBudget > 0; + const spendPercent = hasLimit ? Math.min((spend / maxBudget) * 100, 100) : 0; + const spendColor = spendPercent >= 90 ? "#f5222d" : spendPercent >= 70 ? "#faad14" : "#52c41a"; + + const modelSpendData = useMemo(() => { + const raw = (project?.model_spend ?? {}) as Record; + return Object.entries(raw) + .map(([model, value]) => ({ model, spend: value })) + .sort((a, b) => b.spend - a.spend); + }, [project?.model_spend]); + + if (isLoading) { + return ( + + + } size="large" /> + + + ); + } + + if (!project) { + return ( + + } onClick={onBack} type="text" style={{ marginBottom: 16 }} /> + + + ); + } + + return ( + + {/* Header */} + + + } onClick={onBack} type="text" /> + + + + {project.project_alias ?? project.project_id} + + {project.blocked ? "Blocked" : "Active"} + + + ID: {project.project_id} + + + + } onClick={() => setIsEditModalVisible(true)}> + Edit Project + + + + {/* Project Details */} + + + + {project.description || "\u2014"} + + {new Date(project.created_at).toLocaleString()} + {project.created_by && ( + + {"by"} + + + )} + + + {new Date(project.updated_at).toLocaleString()} + {project.updated_by && ( + + {"by"} + + + )} + + + + + + {/* Spend / Budget */} + + + + + Budget + + } + style={{ height: "100%" }} + > + + + + ${spend.toFixed(2)} + + + {hasLimit ? `of $${maxBudget.toFixed(2)} budget` : "No budget limit"} + + {hasLimit && ( + + + + {(Math.round(spendPercent * 10) / 10).toFixed(1)}% utilized + + + )} + + + + + + {modelSpendData.length > 0 ? ( + `$${value.toFixed(4)}`} + yAxisWidth={140} + showLegend={false} + style={{ height: Math.max(modelSpendData.length * 40, 120) }} + /> + ) : ( + + )} + + + + + {/* Keys & Team */} + + + + + + + + Team + + } + style={{ height: "100%" }} + > + {teamInfo ? ( + (() => { + const teamBudget = teamInfo.max_budget ?? null; + const teamSpend = teamInfo.spend ?? 0; + const teamHasLimit = teamBudget != null && teamBudget > 0; + const teamPercent = teamHasLimit ? Math.min((teamSpend / teamBudget) * 100, 100) : 0; + const teamColor = teamPercent >= 90 ? "#f5222d" : teamPercent >= 70 ? "#faad14" : "#52c41a"; + + return ( + + {/* Team name + ID */} + + + {teamInfo.team_alias || teamInfo.team_id} + + + + ID:{" "} + + {teamInfo.team_id} + + + + + {/* Models */} + + + Models + + {(teamInfo.models?.length ?? 0) > 0 ? ( + + {teamInfo.models?.map((m: string) => ( + + {m} + + ))} + + ) : ( + All models + )} + + + {/* Budget + Spend compact */} + + + + Spend + + + ${teamSpend.toFixed(2)} + {teamHasLimit ? ( + + {" "} + / ${teamBudget.toFixed(2)} + + ) : ( + + {" "} + (Unlimited) + + )} + + + {teamHasLimit && ( + + )} + + + {/* Members */} + + + Members + + {teamInfo.members_with_roles?.length ?? 0} + + + ); + })() + ) : project.team_id ? ( + + } size="small" /> + + ) : ( + + )} + + + + + {/* Edit Modal */} + setIsEditModalVisible(false)} /> + + ); +} diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectKeysSection.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectKeysSection.tsx new file mode 100644 index 00000000000..9f60db596e8 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Projects/ProjectKeysSection.tsx @@ -0,0 +1,67 @@ +import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys"; +import { LoadingOutlined } from "@ant-design/icons"; +import { Card, Flex, Input, Pagination, Spin } from "antd"; +import { KeyIcon, SearchIcon } from "lucide-react"; +import { useEffect, useState } from "react"; +import { ProjectKeysTable } from "./ProjectKeysTable"; + +interface ProjectKeysSectionProps { + projectId: string; +} + +const PAGE_SIZE = 5; + +export function ProjectKeysSection({ projectId }: ProjectKeysSectionProps) { + const [page, setPage] = useState(1); + const [keyAlias, setKeyAlias] = useState(""); + + const { data, isLoading } = useKeys(page, PAGE_SIZE, { + projectID: projectId, + selectedKeyAlias: keyAlias || null, + }); + + // Reset to page 1 when filter changes + useEffect(() => { + setPage(1); + }, [keyAlias]); + + const keys = data?.keys ?? []; + const totalCount = data?.total_count ?? 0; + + return ( + + + Keys + + } + style={{ height: "100%" }} + > + + } + placeholder="Filter by key name..." + style={{ maxWidth: 220 }} + value={keyAlias} + onChange={(e) => setKeyAlias(e.target.value)} + allowClear + size="small" + /> + `${total} keys`} + /> + + } /> } : false} + /> + + ); +} diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectKeysTable.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectKeysTable.tsx new file mode 100644 index 00000000000..cb80d0e27a5 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Projects/ProjectKeysTable.tsx @@ -0,0 +1,58 @@ +import { KeyResponse } from "@/components/key_team_helpers/key_list"; +import { Empty, Table, Tooltip } from "antd"; +import type { ColumnsType } from "antd/es/table"; +import type { SpinProps } from "antd"; +import DefaultProxyAdminTag from "../common_components/DefaultProxyAdminTag"; + +interface ProjectKeysTableProps { + keys: KeyResponse[]; + loading?: boolean | SpinProps; +} + +const columns: ColumnsType = [ + { + title: "Key Name", + dataIndex: "key_alias", + key: "key_alias", + render: (alias: string | null) => alias || "—", + }, + { + title: "Owner", + key: "owner", + render: (_: unknown, record: KeyResponse) => { + const email = record.user?.user_email ?? record.user_id ?? null; + if (!email) return "—"; + return ( + + + + ); + }, + }, + { + title: "Created", + dataIndex: "created_at", + key: "created_at", + render: (date: string) => (date ? new Date(date).toLocaleDateString() : "—"), + }, + { + title: "Last Active", + dataIndex: "last_active", + key: "last_active", + render: (date: string | null) => (date ? new Date(date).toLocaleDateString() : "Never"), + }, +]; + +export function ProjectKeysTable({ keys, loading }: ProjectKeysTableProps) { + return ( + }} + /> + ); +} diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectModals/CreateProjectModal.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectModals/CreateProjectModal.tsx index 14b4d70b743..e490f89303f 100644 --- a/ui/litellm-dashboard/src/components/Projects/ProjectModals/CreateProjectModal.tsx +++ b/ui/litellm-dashboard/src/components/Projects/ProjectModals/CreateProjectModal.tsx @@ -1,95 +1,39 @@ -import { useEffect, useState } from "react"; +import { Modal, Form, Button, Typography, message } from "antd"; +import { FolderAddOutlined } from "@ant-design/icons"; import { - Alert, - Modal, - Form, - Input, - Select, - Switch, - InputNumber, - Collapse, - Button, - Col, - Flex, - Row, - Space, - Divider, - Typography, - message, -} from "antd"; -import { FolderAddOutlined, PlusOutlined, MinusCircleOutlined } from "@ant-design/icons"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; -import { useCreateProject, ProjectCreateParams } from "@/app/(dashboard)/hooks/projects/useCreateProject"; -import { Team } from "../../key_team_helpers/key_list"; -import { fetchTeamModels } from "../../organisms/create_key_button"; -import { getModelDisplayName } from "../../key_team_helpers/fetch_available_models_team_key"; + useCreateProject, + ProjectCreateParams, +} from "@/app/(dashboard)/hooks/projects/useCreateProject"; +import { + ProjectBaseForm, + ProjectFormValues, +} from "./ProjectBaseForm"; +import { buildProjectApiParams } from "./projectFormUtils"; interface CreateProjectModalProps { isOpen: boolean; onClose: () => void; } -export function CreateProjectModal({ isOpen, onClose }: CreateProjectModalProps) { - const [form] = Form.useForm(); - const { accessToken, userId, userRole } = useAuthorized(); - const { data: teams } = useTeams(); +export function CreateProjectModal({ + isOpen, + onClose, +}: CreateProjectModalProps) { + const [form] = Form.useForm(); const createMutation = useCreateProject(); - const [selectedTeam, setSelectedTeam] = useState(null); - const [modelsToPick, setModelsToPick] = useState([]); - - // Fetch team-scoped models when team selection changes - useEffect(() => { - if (userId && userRole && accessToken && selectedTeam) { - fetchTeamModels(userId, userRole, accessToken, selectedTeam.team_id).then((models) => { - const allModels = Array.from(new Set([...(selectedTeam.models ?? []), ...models])); - setModelsToPick(allModels); - }); - } else { - setModelsToPick([]); - } - form.setFieldValue("models", []); - }, [selectedTeam, accessToken, userId, userRole, form]); - const handleSubmit = async () => { try { const values = await form.validateFields(); - - // Build model-specific limits from the dynamic form list - const modelRpmLimit: Record = {}; - const modelTpmLimit: Record = {}; - for (const entry of values.modelLimits ?? []) { - if (entry.model) { - if (entry.rpm != null) modelRpmLimit[entry.model] = entry.rpm; - if (entry.tpm != null) modelTpmLimit[entry.model] = entry.tpm; - } - } - - // Build metadata from the dynamic form list - const metadata: Record = {}; - for (const entry of values.metadata ?? []) { - if (entry.key) metadata[entry.key] = entry.value; - } - const params: ProjectCreateParams = { - project_alias: values.project_alias, - description: values.description, + ...buildProjectApiParams(values), team_id: values.team_id, - models: values.models ?? [], - max_budget: values.max_budget, - blocked: values.isBlocked ?? false, - ...(Object.keys(modelRpmLimit).length > 0 && { model_rpm_limit: modelRpmLimit }), - ...(Object.keys(modelTpmLimit).length > 0 && { model_tpm_limit: modelTpmLimit }), - ...(Object.keys(metadata).length > 0 && { metadata }), }; createMutation.mutate(params, { onSuccess: () => { message.success("Project created successfully"); form.resetFields(); - setSelectedTeam(null); - setModelsToPick([]); onClose(); }, onError: (error) => { @@ -103,16 +47,9 @@ export function CreateProjectModal({ isOpen, onClose }: CreateProjectModalProps) const handleCancel = () => { form.resetFields(); - setSelectedTeam(null); - setModelsToPick([]); onClose(); }; - const handleTeamChange = (teamId: string) => { - const team = teams?.find((t) => t.team_id === teamId) ?? null; - setSelectedTeam(team); - }; - return ( Cancel , - } loading={createMutation.isPending} onClick={handleSubmit}> + } + loading={createMutation.isPending} + onClick={handleSubmit} + > Create Project , ]} > - - {/* Basic Info */} - - Basic Information - - - - - - - - - - - - { - const team = teams?.find((t) => t.team_id === option?.value); - if (!team) return false; - const search = input.toLowerCase().trim(); - return ( - (team.team_alias || "").toLowerCase().includes(search) || - team.team_id.toLowerCase().includes(search) - ); - }} - > - {teams?.map((team) => ( - - {team.team_alias}{" "} - ({team.team_id}) - - ))} - - - - - - - - - - - - - - - - - { - if (values.includes("all-team-models")) { - form.setFieldsValue({ models: ["all-team-models"] }); - } - }} - > - - All Team Models - - {modelsToPick.map((model) => ( - - {getModelDisplayName(model)} - - ))} - - - - - - - - - - - - - - {/* Advanced Settings */} - - - - - Advanced Settings - - } - key="1" - > - - Block Project - - - - - prev.isBlocked !== cur.isBlocked}> - {({ getFieldValue }) => - getFieldValue("isBlocked") ? ( - - ) : null - } - - - - - - Model-Specific Limits - - - {(fields, { add, remove }) => ( - <> - {fields.map(({ key, name, ...restField }) => ( - - - - - - - - - - - remove(name)} style={{ color: "#ef4444" }} /> - - ))} - - add()} block icon={}> - Add Model Limit - - - > - )} - - - - - - Metadata - - - {(fields, { add, remove }) => ( - <> - {fields.map(({ key, name, ...restField }) => ( - - - - - - - - remove(name)} style={{ color: "#ef4444" }} /> - - ))} - - add()} block icon={}> - Add Key-Value Pair - - - > - )} - - - - - - + ); } diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectModals/EditProjectModal.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectModals/EditProjectModal.tsx new file mode 100644 index 00000000000..75f56b1373f --- /dev/null +++ b/ui/litellm-dashboard/src/components/Projects/ProjectModals/EditProjectModal.tsx @@ -0,0 +1,126 @@ +import { useEffect } from "react"; +import { Modal, Form, Button, Typography, message } from "antd"; +import { SaveOutlined } from "@ant-design/icons"; +import { ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; +import { + useUpdateProject, + ProjectUpdateParams, +} from "@/app/(dashboard)/hooks/projects/useUpdateProject"; +import { ProjectBaseForm, ProjectFormValues } from "./ProjectBaseForm"; +import { buildProjectApiParams } from "./projectFormUtils"; + +interface EditProjectModalProps { + isOpen: boolean; + project: ProjectResponse; + onClose: () => void; + onSuccess?: () => void; +} + +export function EditProjectModal({ + isOpen, + project, + onClose, + onSuccess, +}: EditProjectModalProps) { + const [form] = Form.useForm(); + const updateMutation = useUpdateProject(); + + // Populate form with existing project data when modal opens + useEffect(() => { + if (isOpen && project) { + // Model limits are stored inside metadata by the backend + const metadataObj = (project.metadata ?? {}) as Record; + const rpmLimits = (metadataObj.model_rpm_limit ?? {}) as Record; + const tpmLimits = (metadataObj.model_tpm_limit ?? {}) as Record; + + const modelLimits: ProjectFormValues["modelLimits"] = []; + const allLimitModels = new Set([ + ...Object.keys(rpmLimits), + ...Object.keys(tpmLimits), + ]); + for (const model of allLimitModels) { + modelLimits.push({ + model, + rpm: rpmLimits[model], + tpm: tpmLimits[model], + }); + } + + // Filter out internal keys from user-facing metadata + const internalKeys = new Set(["model_rpm_limit", "model_tpm_limit"]); + const metadata: ProjectFormValues["metadata"] = []; + for (const [key, value] of Object.entries(metadataObj)) { + if (!internalKeys.has(key)) { + metadata.push({ key, value: String(value) }); + } + } + + form.setFieldsValue({ + project_alias: project.project_alias ?? "", + team_id: project.team_id ?? "", + description: project.description ?? "", + models: project.models ?? [], + max_budget: project.litellm_budget_table?.max_budget ?? undefined, + isBlocked: project.blocked, + modelLimits: modelLimits.length > 0 ? modelLimits : undefined, + metadata: metadata.length > 0 ? metadata : undefined, + }); + } + }, [isOpen, project, form]); + + const handleSubmit = async () => { + try { + const values = await form.validateFields(); + const params: ProjectUpdateParams = { + ...buildProjectApiParams(values), + team_id: values.team_id, + }; + + updateMutation.mutate( + { projectId: project.project_id, params }, + { + onSuccess: () => { + message.success("Project updated successfully"); + onSuccess?.(); + onClose(); + }, + onError: (error) => { + message.error(error.message || "Failed to update project"); + }, + }, + ); + } catch (error) { + console.error("Validation failed:", error); + } + }; + + return ( + + Edit Project + + } + open={isOpen} + onCancel={onClose} + width={720} + destroyOnHidden + footer={[ + + Cancel + , + } + loading={updateMutation.isPending} + onClick={handleSubmit} + > + Save Changes + , + ]} + > + + + ); +} diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectModals/ProjectBaseForm.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectModals/ProjectBaseForm.tsx new file mode 100644 index 00000000000..bf1eca882c3 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Projects/ProjectModals/ProjectBaseForm.tsx @@ -0,0 +1,401 @@ +import { useEffect, useState } from "react"; +import { + Alert, + Col, + Collapse, + Divider, + Flex, + Form, + Input, + InputNumber, + Row, + Select, + Space, + Switch, + Typography, + Button, +} from "antd"; +import type { FormInstance } from "antd"; +import { PlusOutlined, MinusCircleOutlined } from "@ant-design/icons"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; +import { Team } from "../../key_team_helpers/key_list"; +import { fetchTeamModels } from "../../organisms/create_key_button"; +import { getModelDisplayName } from "../../key_team_helpers/fetch_available_models_team_key"; + +export interface ProjectFormValues { + project_alias: string; + team_id: string; + description?: string; + models: string[]; + max_budget?: number; + isBlocked: boolean; + modelLimits?: { model: string; tpm?: number; rpm?: number }[]; + metadata?: { key: string; value: string }[]; +} + +interface ProjectBaseFormProps { + form: FormInstance; +} + +export function ProjectBaseForm({ + form, +}: ProjectBaseFormProps) { + const { accessToken, userId, userRole } = useAuthorized(); + const { data: teams } = useTeams(); + + const [selectedTeam, setSelectedTeam] = useState(null); + const [modelsToPick, setModelsToPick] = useState([]); + + // Sync selectedTeam from form value (needed for edit mode pre-fill) + const teamIdValue = Form.useWatch("team_id", form); + useEffect(() => { + if (teamIdValue && teams) { + const team = teams.find((t) => t.team_id === teamIdValue) ?? null; + if (team && team.team_id !== selectedTeam?.team_id) { + setSelectedTeam(team); + } + } + }, [teamIdValue, teams, selectedTeam?.team_id]); + + // Fetch team-scoped models when team selection changes + useEffect(() => { + if (userId && userRole && accessToken && selectedTeam) { + fetchTeamModels(userId, userRole, accessToken, selectedTeam.team_id).then( + (models) => { + const allModels = Array.from( + new Set([...(selectedTeam.models ?? []), ...models]), + ); + setModelsToPick(allModels); + }, + ); + } else { + setModelsToPick([]); + } + }, [selectedTeam, accessToken, userId, userRole]); + + const handleTeamChange = (teamId: string) => { + const team = teams?.find((t) => t.team_id === teamId) ?? null; + setSelectedTeam(team); + form.setFieldValue("models", []); + }; + + return ( + + {/* Basic Info */} + + Basic Information + + + + + + + + + + + + { + const team = teams?.find((t) => t.team_id === option?.value); + if (!team) return false; + const search = input.toLowerCase().trim(); + return ( + (team.team_alias || "").toLowerCase().includes(search) || + team.team_id.toLowerCase().includes(search) + ); + }} + > + {teams?.map((team) => ( + + {team.team_alias}{" "} + ({team.team_id}) + + ))} + + + + + + + + + + + + + + + + + { + if (values.includes("all-team-models")) { + form.setFieldsValue({ models: ["all-team-models"] }); + } + }} + > + + All Team Models + + {modelsToPick.map((model) => ( + + {getModelDisplayName(model)} + + ))} + + + + + + + + + + + + + + {/* Advanced Settings */} + + + + Advanced Settings + + ), + children: ( + <> + + Block Project + + + + + prev.isBlocked !== cur.isBlocked} + > + {({ getFieldValue }) => + getFieldValue("isBlocked") ? ( + + ) : null + } + + + + + + Model-Specific Limits + + + {(fields, { add, remove }) => ( + <> + {fields.map(({ key, name, ...restField }) => ( + + { + if (!value) return Promise.resolve(); + const all = form.getFieldValue("modelLimits") ?? []; + const dupes = all.filter( + (entry: { model?: string }) => entry?.model === value, + ); + if (dupes.length > 1) { + return Promise.reject(new Error("Duplicate model")); + } + return Promise.resolve(); + }, + }, + ]} + > + + + + + + + + + remove(name)} + style={{ color: "#ef4444" }} + /> + + ))} + + add()} + block + icon={} + > + Add Model Limit + + + > + )} + + + + + + Metadata + + + {(fields, { add, remove }) => ( + <> + {fields.map(({ key, name, ...restField }) => ( + + { + if (!value) return Promise.resolve(); + const all = form.getFieldValue("metadata") ?? []; + const dupes = all.filter( + (entry: { key?: string }) => entry?.key === value, + ); + if (dupes.length > 1) { + return Promise.reject(new Error("Duplicate key")); + } + return Promise.resolve(); + }, + }, + ]} + > + + + + + + remove(name)} + style={{ color: "#ef4444" }} + /> + + ))} + + add()} + block + icon={} + > + Add Key-Value Pair + + + > + )} + + > + ), + }, + ]} + /> + + + + ); +} diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectModals/projectFormUtils.ts b/ui/litellm-dashboard/src/components/Projects/ProjectModals/projectFormUtils.ts new file mode 100644 index 00000000000..97c093b57d9 --- /dev/null +++ b/ui/litellm-dashboard/src/components/Projects/ProjectModals/projectFormUtils.ts @@ -0,0 +1,36 @@ +import { ProjectFormValues } from "./ProjectBaseForm"; + +/** + * Transforms ProjectFormValues into the flat API param shape + * shared by both create and update endpoints. + */ +export function buildProjectApiParams(values: ProjectFormValues) { + const modelRpmLimit: Record = {}; + const modelTpmLimit: Record = {}; + for (const entry of values.modelLimits ?? []) { + if (entry.model) { + if (entry.rpm != null) modelRpmLimit[entry.model] = entry.rpm; + if (entry.tpm != null) modelTpmLimit[entry.model] = entry.tpm; + } + } + + const metadata: Record = {}; + for (const entry of values.metadata ?? []) { + if (entry.key) metadata[entry.key] = entry.value; + } + + return { + project_alias: values.project_alias, + description: values.description, + models: values.models ?? [], + max_budget: values.max_budget, + blocked: values.isBlocked ?? false, + ...(Object.keys(modelRpmLimit).length > 0 && { + model_rpm_limit: modelRpmLimit, + }), + ...(Object.keys(modelTpmLimit).length > 0 && { + model_tpm_limit: modelTpmLimit, + }), + ...(Object.keys(metadata).length > 0 && { metadata }), + }; +} diff --git a/ui/litellm-dashboard/src/components/Projects/ProjectsPage.tsx b/ui/litellm-dashboard/src/components/Projects/ProjectsPage.tsx index 40ab5045703..9c75e19ac4e 100644 --- a/ui/litellm-dashboard/src/components/Projects/ProjectsPage.tsx +++ b/ui/litellm-dashboard/src/components/Projects/ProjectsPage.tsx @@ -1,13 +1,16 @@ import { useProjects, ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams"; -import { PlusOutlined } from "@ant-design/icons"; +import { LoadingOutlined, PlusOutlined } from "@ant-design/icons"; import { + Alert, Button, Card, Flex, Input, Layout, + Pagination, Space, + Spin, Table, Tag, theme, @@ -18,6 +21,7 @@ import type { ColumnsType } from "antd/es/table"; import { LayersIcon, SearchIcon } from "lucide-react"; import { useEffect, useMemo, useState } from "react"; import { CreateProjectModal } from "./ProjectModals/CreateProjectModal"; +import { ProjectDetail } from "./ProjectDetailsPage"; const { Title, Text } = Typography; const { Content } = Layout; @@ -25,8 +29,9 @@ const { Content } = Layout; export function ProjectsPage() { const { token } = theme.useToken(); const { data: projects, isLoading } = useProjects(); - const { data: teams } = useTeams(); + const { data: teams, isLoading: isTeamsLoading } = useTeams(); + const [selectedProjectId, setSelectedProjectId] = useState(null); const [isCreateModalVisible, setIsCreateModalVisible] = useState(false); const [searchText, setSearchText] = useState(""); const [currentPage, setCurrentPage] = useState(1); @@ -74,6 +79,7 @@ export function ProjectsPage() { ellipsis className="text-blue-500 bg-blue-50 hover:bg-blue-100 text-xs cursor-pointer" style={{ fontSize: 14, padding: "1px 8px" }} + onClick={() => setSelectedProjectId(id)} > {id} @@ -96,8 +102,11 @@ export function ProjectsPage() { return aAlias.localeCompare(bAlias); }, render: (_: unknown, record: ProjectResponse) => { - const alias = teamAliasMap.get(record.team_id ?? ""); - return alias ?? record.team_id ?? "—"; + if (!record.team_id) return "—"; + const alias = teamAliasMap.get(record.team_id); + if (alias) return alias; + if (isTeamsLoading) return } size="small" />; + return record.team_id; }, }, { @@ -144,10 +153,25 @@ export function ProjectsPage() { }, ]; + if (selectedProjectId) { + return ( + setSelectedProjectId(null)} + /> + ); + } + return ( + setSearchText(e.target.value)} allowClear /> + setCurrentPage(page)} + size="small" + showTotal={(total) => `${total} projects`} + showSizeChanger={false} + /> setCurrentPage(page), - size: "small", - showTotal: (total) => `${total} projects`, - showSizeChanger: false, - }} + pagination={false} /> diff --git a/ui/litellm-dashboard/src/components/common_components/ProjectDropdown.tsx b/ui/litellm-dashboard/src/components/common_components/ProjectDropdown.tsx new file mode 100644 index 00000000000..f91cbd4a14b --- /dev/null +++ b/ui/litellm-dashboard/src/components/common_components/ProjectDropdown.tsx @@ -0,0 +1,62 @@ +import React from "react"; +import { Select, Spin } from "antd"; +import { LoadingOutlined } from "@ant-design/icons"; +import { ProjectResponse } from "@/app/(dashboard)/hooks/projects/useProjects"; + +interface ProjectDropdownProps { + projects?: ProjectResponse[] | null; + value?: string; + onChange?: (value: string) => void; + disabled?: boolean; + loading?: boolean; + /** When set, only show projects belonging to this team */ + teamId?: string | null; +} + +const ProjectDropdown: React.FC = ({ + projects, + value, + onChange, + disabled, + loading, + teamId, +}) => { + const filtered = teamId + ? projects?.filter((p) => p.team_id === teamId) + : projects; + + return ( + } size="small" /> : undefined} + filterOption={(input, option) => { + if (!option) return false; + const project = filtered?.find((p) => p.project_id === option.key); + if (!project) return false; + + const searchTerm = input.toLowerCase().trim(); + const alias = (project.project_alias || "").toLowerCase(); + const id = (project.project_id || "").toLowerCase(); + + return alias.includes(searchTerm) || id.includes(searchTerm); + }} + optionFilterProp="children" + > + {!loading && + filtered?.map((project) => ( + + {project.project_alias || project.project_id}{" "} + ({project.project_id}) + + ))} + + ); +}; + +export default ProjectDropdown; diff --git a/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx b/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx index d54724da2a7..9e79ea2950a 100644 --- a/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx +++ b/ui/litellm-dashboard/src/components/common_components/team_dropdown.tsx @@ -7,10 +7,10 @@ interface TeamDropdownProps { value?: string; onChange?: (value: string) => void; disabled?: boolean; + loading?: boolean; } -const TeamDropdown: React.FC = ({ teams, value, onChange, disabled }) => { - console.log("disabled", disabled); +const TeamDropdown: React.FC = ({ teams, value, onChange, disabled, loading }) => { return ( = ({ teams, value, onChange, dis value={value} onChange={onChange} disabled={disabled} + loading={loading} allowClear filterOption={(input, option) => { if (!option) return false; diff --git a/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx b/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx index 5512809ba3f..08bccda7749 100644 --- a/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx +++ b/ui/litellm-dashboard/src/components/key_team_helpers/key_list.tsx @@ -30,6 +30,7 @@ export interface KeyResponse { config: Record; user_id: string; team_id: string | null; + project_id: string | null; max_parallel_requests: number; metadata: Record; tpm_limit: number; diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx index 12118a6aaa0..c72cb56ffa5 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx @@ -15,6 +15,7 @@ vi.mock("../networking", () => ({ keyCreateCall: mockKeyCreateCall, modelAvailableCall: vi.fn().mockResolvedValue({ data: [{ id: "gpt-4" }, { id: "gpt-3.5-turbo" }] }), getGuardrailsList: vi.fn().mockResolvedValue({ guardrails: [] }), + getPoliciesList: vi.fn().mockResolvedValue({ policies: [] }), getPromptsList: vi.fn().mockResolvedValue({ prompts: [] }), proxyBaseUrl: "http://localhost:4000", getPossibleUserRoles: vi.fn().mockResolvedValue({ @@ -27,7 +28,7 @@ vi.mock("../networking", () => ({ soft_budget: null, }), fetchMCPAccessGroups: vi.fn().mockResolvedValue([]), - getAgentsList: vi.fn().mockResolvedValue([]), + getAgentsList: vi.fn().mockResolvedValue({ agents: [] }), })); vi.mock("../molecules/notifications_manager", () => ({ @@ -41,6 +42,20 @@ vi.mock("../molecules/notifications_manager", () => ({ }, })); +vi.mock("@/app/(dashboard)/hooks/projects/useProjects", () => ({ + useProjects: vi.fn().mockReturnValue({ data: [], isLoading: false }), +})); + +vi.mock("../common_components/ProjectDropdown", () => ({ + default: ({ value, onChange }: { value?: string; onChange?: (v: string) => void }) => ( + onChange?.(e.target.value)} + /> + ), +})); + vi.mock("../common_components/AccessGroupSelector", () => ({ default: ({ value = [], onChange }: { value?: string[]; onChange?: (v: string[]) => void }) => ( = ({ team, teams, data, addKey }) => { const { accessToken, userId: userID, userRole, premiumUser } = useAuthorized(); + const { data: projects, isLoading: isProjectsLoading } = useProjects(); const queryClient = useQueryClient(); const [form] = Form.useForm(); const [isModalVisible, setIsModalVisible] = useState(false); @@ -157,6 +160,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { const [promptsList, setPromptsList] = useState([]); const [loggingSettings, setLoggingSettings] = useState([]); const [selectedCreateKeyTeam, setSelectedCreateKeyTeam] = useState(team); + const [selectedProjectId, setSelectedProjectId] = useState(null); const [isCreateUserModalVisible, setIsCreateUserModalVisible] = useState(false); const [newlyCreatedUserId, setNewlyCreatedUserId] = useState(null); const [possibleUIRoles, setPossibleUIRoles] = useState>>({}); @@ -184,6 +188,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { setRouterSettings(null); setRouterSettingsKey((prev) => prev + 1); setSelectedAgentId(null); + setSelectedProjectId(null); }; const handleCancel = () => { @@ -200,6 +205,7 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { setRouterSettings(null); setRouterSettingsKey((prev) => prev + 1); setSelectedAgentId(null); + setSelectedProjectId(null); }; useEffect(() => { @@ -468,6 +474,14 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { }; useEffect(() => { + if (selectedProjectId) { + // When a project is selected, use the project's models + const project = projects?.find((p) => p.project_id === selectedProjectId); + const projectModels = project?.models ?? []; + setModelsToPick(projectModels); + form.setFieldValue("models", []); + return; + } if (userID && userRole && accessToken) { fetchTeamModels(userID, userRole, accessToken, selectedCreateKeyTeam?.team_id ?? null).then((models) => { let allModels = Array.from(new Set([...(selectedCreateKeyTeam?.models ?? []), ...models])); @@ -475,7 +489,22 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { }); } form.setFieldValue("models", []); - }, [selectedCreateKeyTeam, accessToken, userID, userRole]); + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [selectedCreateKeyTeam, selectedProjectId, accessToken, userID, userRole]); + + // Sync team when project is selected but teams loaded later (race condition) + useEffect(() => { + if (!selectedProjectId || !teams) return; + const project = projects?.find((p) => p.project_id === selectedProjectId); + if (!project?.team_id) return; + // If team is already set correctly, skip + if (selectedCreateKeyTeam?.team_id === project.team_id) return; + const projectTeam = teams.find((t) => t.team_id === project.team_id) || null; + if (projectTeam) { + setSelectedCreateKeyTeam(projectTeam); + form.setFieldValue("team_id", projectTeam.team_id); + } + }, [teams, selectedProjectId, projects]); // Add a callback function to handle user creation const handleUserCreated = (userId: string) => { @@ -653,9 +682,40 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { > { const selectedTeam = teams?.find((t) => t.team_id === teamId) || null; setSelectedCreateKeyTeam(selectedTeam); + setSelectedProjectId(null); + form.setFieldValue("project_id", undefined); + }} + /> + + + Project{" "} + + + + + } + name="project_id" + className="mt-4" + > + { + if (!projectId) { + setSelectedProjectId(null); + setSelectedCreateKeyTeam(null); + form.setFieldValue("team_id", undefined); + return; + } + setSelectedProjectId(projectId); }} /> @@ -735,9 +795,11 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { } }} > - - All Team Models - + {!selectedProjectId && ( + + All Team Models + + )} {modelsToPick.map((model: string) => ( {getModelDisplayName(model)} diff --git a/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx b/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx index c2c32390cd8..00cf2343935 100644 --- a/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/KeyInfoView.handleKeyUpdate.test.tsx @@ -249,6 +249,11 @@ vi.mock("@/app/(dashboard)/hooks/useTeams", () => ({ })), })); +// Mock useProjects hook +vi.mock("@/app/(dashboard)/hooks/projects/useProjects", () => ({ + useProjects: vi.fn().mockReturnValue({ data: [], isLoading: false }), +})); + // KeyEditView mock: triggers onSubmit with our injected form values vi.mock("./key_edit_view", async () => { const React = await import("react"); diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index 71e00e16542..4923b71f760 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -1,4 +1,5 @@ import GuardrailSelector from "@/components/guardrails/GuardrailSelector"; +import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects"; import PolicySelector from "@/components/policies/PolicySelector"; import { InfoCircleOutlined } from "@ant-design/icons"; import { TextInput, Button as TremorButton } from "@tremor/react"; @@ -95,6 +96,15 @@ export function KeyEditView({ const [autoRotationEnabled, setAutoRotationEnabled] = useState(keyData.auto_rotate || false); const [rotationInterval, setRotationInterval] = useState(keyData.rotation_interval || ""); const [isKeySaving, setIsKeySaving] = useState(false); + const { data: projects } = useProjects(); + const hasProject = Boolean(keyData.project_id); + const projectDisplay = (() => { + if (!keyData.project_id) return null; + const project = projects?.find((p) => p.project_id === keyData.project_id); + return project?.project_alias + ? `${project.project_alias} (${keyData.project_id})` + : keyData.project_id; + })(); useEffect(() => { const fetchModels = async () => { @@ -590,10 +600,15 @@ export function KeyEditView({ /> - + { const team = teams?.find((t) => t.team_id === option?.value); @@ -601,7 +616,6 @@ export function KeyEditView({ return team.team_alias?.toLowerCase().includes(input.toLowerCase()) ?? false; }} > - {/* Only show All Team Models if team has models */} {teams?.map((team) => ( {`${team.team_alias} (${team.team_id})`} @@ -609,6 +623,11 @@ export function KeyEditView({ ))} + {hasProject && ( + + + + )} ({ default: vi.fn(), })); +vi.mock("@/app/(dashboard)/hooks/projects/useProjects", () => ({ + useProjects: vi.fn().mockReturnValue({ data: [], isLoading: false }), +})); + vi.mock("../networking", () => ({ keyDeleteCall: vi.fn().mockResolvedValue({}), keyUpdateCall: vi.fn().mockResolvedValue({}), @@ -49,6 +55,7 @@ describe("KeyInfoView", () => { config: {}, user_id: "default_user_id", team_id: null, + project_id: null, max_parallel_requests: 10, metadata: { logging: [], @@ -121,7 +128,7 @@ describe("KeyInfoView", () => { it("should render tags", async () => { vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock); - render( + renderWithProviders( { }} @@ -138,7 +145,7 @@ describe("KeyInfoView", () => { it("should not render tags in metadata textarea", async () => { vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock); - const { container } = render( + const { container } = renderWithProviders( { }} @@ -168,7 +175,7 @@ describe("KeyInfoView", () => { }); const keyData = { ...MOCK_KEY_DATA, user_id: "other-user-id" }; - render( + renderWithProviders( { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />, ); @@ -213,7 +220,7 @@ describe("KeyInfoView", () => { }); const keyData = { ...MOCK_KEY_DATA, team_id: teamId, user_id: "other-user-id" }; - render( + renderWithProviders( { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />, ); @@ -237,7 +244,7 @@ describe("KeyInfoView", () => { const ownerUserId = "owner-user-id"; const keyData = { ...MOCK_KEY_DATA, user_id: ownerUserId }; - render( + renderWithProviders( { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />, ); @@ -260,7 +267,7 @@ describe("KeyInfoView", () => { }); const keyData = { ...MOCK_KEY_DATA, user_id: "owner-user-id" }; - render( + renderWithProviders( { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />, ); @@ -284,7 +291,7 @@ describe("KeyInfoView", () => { const ownerUserId = "internal-viewer-user-id"; const keyData = { ...MOCK_KEY_DATA, user_id: ownerUserId }; - render( + renderWithProviders( { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />, ); @@ -328,7 +335,7 @@ describe("KeyInfoView", () => { }); const keyData = { ...MOCK_KEY_DATA, team_id: "non-matching-team-id", user_id: "other-user-id" }; - render( + renderWithProviders( { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />, ); @@ -342,7 +349,7 @@ describe("KeyInfoView", () => { vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock); const onCloseMock = vi.fn(); - render( + renderWithProviders( { describe("'Edit Settings' button visibility in the Settings tab", () => { const renderAndOpenSettingsTab = async (keyData = MOCK_KEY_DATA) => { - render( + renderWithProviders( {}} @@ -474,7 +481,7 @@ describe("KeyInfoView", () => { }, }; - render( + renderWithProviders( { }} @@ -500,7 +507,7 @@ describe("KeyInfoView", () => { }, }; - render( + renderWithProviders( { }} @@ -518,7 +525,7 @@ describe("KeyInfoView", () => { it("should display no key found message when keyData is undefined", async () => { vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock); - render( + renderWithProviders( { }} diff --git a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx index 378f5b3872b..9aac652a58d 100644 --- a/ui/litellm-dashboard/src/components/templates/key_info_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_info_view.tsx @@ -1,4 +1,5 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects"; import useTeams from "@/app/(dashboard)/hooks/useTeams"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { mapEmptyStringToNull } from "@/utils/keyUpdateUtils"; @@ -48,6 +49,7 @@ export default function KeyInfoView({ }: KeyInfoViewProps) { const { accessToken, userId: userID, userRole, premiumUser } = useAuthorized(); const { teams: teamsData } = useTeams(); + const { data: projects } = useProjects(); const [isEditing, setIsEditing] = useState(false); const [form] = Form.useForm(); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); @@ -571,6 +573,20 @@ export default function KeyInfoView({ {currentKeyData.team_id || "Not Set"} + + Project + + {currentKeyData.project_id + ? (() => { + const project = projects?.find((p) => p.project_id === currentKeyData.project_id); + return project?.project_alias + ? `${project.project_alias} (${currentKeyData.project_id})` + : currentKeyData.project_id; + })() + : "Not Set"} + + + Organization {currentKeyData.organization_id || "Not Set"}
{code}