Merge remote-tracking branch 'origin/main' into litellm_main_stability_fix

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-02-28 18:04:35 +00:00
commit 93e08a6509
72 changed files with 6259 additions and 621 deletions

View file

@ -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

208
.github/scripts/close_duplicate_issues.py vendored Executable file
View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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 .

View file

@ -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
<Image img={require('../../img/dd_llm_obs.png')} />
## 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 |

View file

@ -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

View file

@ -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)

View file

@ -105,6 +105,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
"prometheus",
"otel",
"datadog",
"datadog_metrics",
"datadog_llm_observability",
"galileo",
"braintrust",

View file

@ -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))

View file

@ -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"
}
]
]

View file

@ -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

View file

@ -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,

View file

@ -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,
},
),
}
}

View file

@ -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):

View file

@ -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:

View file

@ -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

View file

@ -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"])

View file

@ -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 = {

View file

@ -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

View file

@ -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.

View file

@ -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)}")

View file

@ -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)

View file

@ -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:

View file

@ -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

View file

@ -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",

View file

@ -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,

View file

@ -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)

View file

@ -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:

View file

@ -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,

View file

@ -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]

View file

@ -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 <base64string>'
# 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 '<base64(username:password)>' in comments/docs instead."
)

View file

@ -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

View file

@ -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}'

View file

@ -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

View file

@ -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"

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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)}"

View file

@ -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", () => {

View file

@ -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,

View file

@ -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<ProjectResponse> => {
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<ProjectResponse>({
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<ProjectResponse[]>(
projectKeys.list({}),
);
return projects?.find((p) => p.project_id === projectId);
},
});
};

View file

@ -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<string, unknown>;
model_rpm_limit?: Record<string, number>;
model_tpm_limit?: Record<string, number>;
}
// ── Fetch function ───────────────────────────────────────────────────────────
const updateProject = async (
accessToken: string,
projectId: string,
params: ProjectUpdateParams,
): Promise<ProjectResponse> => {
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 });
},
});
};

View file

@ -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(<AddMarginForm {...DEFAULT_PROPS} />);
expect(screen.getByRole("button", { name: /add provider margin/i })).toBeInTheDocument();
});
it("should show the percentage input when marginType is percentage", () => {
renderWithProviders(<AddMarginForm {...DEFAULT_PROPS} marginType="percentage" />);
expect(screen.getByPlaceholderText("10")).toBeInTheDocument();
});
it("should show the fixed amount input when marginType is fixed", () => {
renderWithProviders(<AddMarginForm {...DEFAULT_PROPS} marginType="fixed" />);
expect(screen.getByPlaceholderText("0.001")).toBeInTheDocument();
});
it("should not show the fixed amount input when marginType is percentage", () => {
renderWithProviders(<AddMarginForm {...DEFAULT_PROPS} marginType="percentage" />);
expect(screen.queryByPlaceholderText("0.001")).not.toBeInTheDocument();
});
it("should not show the percentage input when marginType is fixed", () => {
renderWithProviders(<AddMarginForm {...DEFAULT_PROPS} marginType="fixed" />);
expect(screen.queryByPlaceholderText("10")).not.toBeInTheDocument();
});
it("should show the Percentage-based and Fixed Amount radio options", () => {
renderWithProviders(<AddMarginForm {...DEFAULT_PROPS} />);
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(
<AddMarginForm {...DEFAULT_PROPS} selectedProvider={undefined} percentageValue="10" />
);
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(
<AddMarginForm {...DEFAULT_PROPS} selectedProvider="OpenAI" percentageValue="" />
);
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(
<AddMarginForm {...DEFAULT_PROPS} selectedProvider="OpenAI" percentageValue="10" />
);
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(
<AddMarginForm
{...DEFAULT_PROPS}
selectedProvider="OpenAI"
marginType="fixed"
fixedAmountValue=""
/>
);
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(
<AddMarginForm
{...DEFAULT_PROPS}
selectedProvider="OpenAI"
marginType="fixed"
fixedAmountValue="0.001"
/>
);
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(
<AddMarginForm
{...DEFAULT_PROPS}
selectedProvider="OpenAI"
percentageValue="10"
onAddProvider={onAddProvider}
/>
);
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(
<AddMarginForm {...DEFAULT_PROPS} onMarginTypeChange={onMarginTypeChange} />
);
await user.click(screen.getByText("Fixed Amount"));
expect(onMarginTypeChange).toHaveBeenCalledWith("fixed");
});
});

View file

@ -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(<AddProviderForm {...DEFAULT_PROPS} />);
expect(screen.getByRole("button", { name: /add provider discount/i })).toBeInTheDocument();
});
it("should render the discount percentage input field", () => {
renderWithProviders(<AddProviderForm {...DEFAULT_PROPS} />);
expect(screen.getByPlaceholderText("5")).toBeInTheDocument();
});
it("should disable the submit button when no provider is selected and no discount is entered", () => {
renderWithProviders(<AddProviderForm {...DEFAULT_PROPS} />);
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(
<AddProviderForm {...DEFAULT_PROPS} selectedProvider="OpenAI" newDiscount="" />
);
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(
<AddProviderForm {...DEFAULT_PROPS} selectedProvider={undefined} newDiscount="5" />
);
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(
<AddProviderForm {...DEFAULT_PROPS} selectedProvider="OpenAI" newDiscount="5" />
);
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(
<AddProviderForm
{...DEFAULT_PROPS}
selectedProvider="OpenAI"
newDiscount="5"
onAddProvider={onAddProvider}
/>
);
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(<AddProviderForm {...DEFAULT_PROPS} />);
expect(screen.getByText("%")).toBeInTheDocument();
});
});

View file

@ -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: () => <div data-testid="pricing-calculator">Pricing Calculator</div>,
}));
vi.mock("../playground/llm_calls/fetch_models", () => ({
fetchAvailableModels: vi.fn().mockResolvedValue([]),
}));
vi.mock("../HelpLink", () => ({
DocsMenu: () => null,
}));
vi.mock("./how_it_works", () => ({
default: () => <div data-testid="how-it-works">How It Works</div>,
}));
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(
<CostTrackingSettings userID="user-1" userRole="proxy_admin" accessToken={null} />
);
expect(container.firstChild).toBeNull();
});
it("should render the page title", () => {
renderWithProviders(<CostTrackingSettings {...ADMIN_PROPS} />);
expect(screen.getByText("Cost Tracking Settings")).toBeInTheDocument();
});
it("should show the Provider Discounts accordion header for proxy_admin", () => {
renderWithProviders(<CostTrackingSettings {...ADMIN_PROPS} />);
expect(screen.getByText("Provider Discounts")).toBeInTheDocument();
});
it("should show the Fee/Price Margin accordion header for proxy_admin", () => {
renderWithProviders(<CostTrackingSettings {...ADMIN_PROPS} />);
expect(screen.getByText("Fee/Price Margin")).toBeInTheDocument();
});
it("should always show the Pricing Calculator section", () => {
renderWithProviders(<CostTrackingSettings {...ADMIN_PROPS} />);
// 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(<CostTrackingSettings {...ADMIN_PROPS} />);
expect(await screen.findByTestId("pricing-calculator")).toBeInTheDocument();
});
it("should not show Provider Discounts section for a non-admin role", () => {
renderWithProviders(
<CostTrackingSettings userID="user-1" userRole="internal_user" accessToken="test-token" />
);
expect(screen.queryByText("Provider Discounts")).not.toBeInTheDocument();
});
it("should not show Fee/Price Margin section for a non-admin role", () => {
renderWithProviders(
<CostTrackingSettings userID="user-1" userRole="internal_user" accessToken="test-token" />
);
expect(screen.queryByText("Fee/Price Margin")).not.toBeInTheDocument();
});
it("should show Provider Discounts for the 'Admin' role as well", () => {
renderWithProviders(
<CostTrackingSettings userID="user-1" userRole="Admin" accessToken="test-token" />
);
expect(screen.getByText("Provider Discounts")).toBeInTheDocument();
});
it("should show the subtitle describing discount/margin configuration", () => {
renderWithProviders(<CostTrackingSettings {...ADMIN_PROPS} />);
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(<CostTrackingSettings {...ADMIN_PROPS} />);
// 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(<CostTrackingSettings {...ADMIN_PROPS} />);
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(<CostTrackingSettings {...ADMIN_PROPS} />);
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(<CostTrackingSettings {...ADMIN_PROPS} />);
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();
});
});
});

View file

@ -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 }) => <pre data-testid="code-block">{code}</pre>,
}));
describe("HowItWorks", () => {
beforeEach(() => {
vi.clearAllMocks();
});
it("should render", () => {
renderWithProviders(<HowItWorks />);
expect(screen.getByText("Cost Calculation")).toBeInTheDocument();
});
it("should display the cost calculation formula", () => {
renderWithProviders(<HowItWorks />);
expect(screen.getByText(/final_cost = base_cost/i)).toBeInTheDocument();
});
it("should display the valid range information", () => {
renderWithProviders(<HowItWorks />);
expect(screen.getByText(/0% and 100%/i)).toBeInTheDocument();
});
it("should render the code block with a curl example", () => {
renderWithProviders(<HowItWorks />);
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(<HowItWorks />);
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(<HowItWorks />);
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(<HowItWorks />);
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(<HowItWorks />);
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(<HowItWorks />);
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(<HowItWorks />);
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();
});
});

View file

@ -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) => <table>{children}</table>,
TableHead: ({ children }: any) => <thead>{children}</thead>,
TableRow: ({ children }: any) => <tr>{children}</tr>,
TableHeaderCell: ({ children }: any) => <th>{children}</th>,
TableBody: ({ children }: any) => <tbody>{children}</tbody>,
TableCell: ({ children }: any) => <td>{children}</td>,
Text: ({ children }: any) => <span>{children}</span>,
TextInput: ({ value, onValueChange, onKeyDown, placeholder, ...rest }: any) => (
<input
value={value}
onChange={(e) => onValueChange?.(e.target.value)}
onKeyDown={onKeyDown}
placeholder={placeholder}
{...rest}
/>
),
Icon: ({ icon: IconComponent, onClick }: any) => {
const name = IconComponent?.displayName ?? IconComponent?.name ?? "icon";
return <button onClick={onClick} aria-label={name} />;
},
}));
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(
<ProviderDiscountTable
discountConfig={DEFAULT_DISCOUNT_CONFIG}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByRole("table")).toBeInTheDocument();
});
it("should render the table headers", () => {
renderWithProviders(
<ProviderDiscountTable
discountConfig={DEFAULT_DISCOUNT_CONFIG}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={DEFAULT_DISCOUNT_CONFIG}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("OpenAI")).toBeInTheDocument();
});
it("should display the formatted discount percentage", () => {
renderWithProviders(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("5.0%")).toBeInTheDocument();
});
it("should show a text input when the edit icon is clicked", async () => {
const user = userEvent.setup();
renderWithProviders(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderDiscountTable
discountConfig={{ openai: 0.05 }}
onDiscountChange={onDiscountChange}
onRemoveProvider={onRemoveProvider}
/>
);
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();
});
});

View file

@ -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();
});
});

View file

@ -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) => <table>{children}</table>,
TableHead: ({ children }: any) => <thead>{children}</thead>,
TableRow: ({ children }: any) => <tr>{children}</tr>,
TableHeaderCell: ({ children }: any) => <th>{children}</th>,
TableBody: ({ children }: any) => <tbody>{children}</tbody>,
TableCell: ({ children }: any) => <td>{children}</td>,
Text: ({ children }: any) => <span>{children}</span>,
TextInput: ({ value, onValueChange, placeholder, autoFocus, className }: any) => (
<input
value={value}
onChange={(e) => onValueChange?.(e.target.value)}
placeholder={placeholder}
autoFocus={autoFocus}
className={className}
/>
),
Icon: ({ icon: IconComponent, onClick }: any) => {
const name = IconComponent?.displayName ?? IconComponent?.name ?? "icon";
return <button onClick={onClick} aria-label={name} />;
},
}));
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(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByRole("table")).toBeInTheDocument();
});
it("should render the table headers", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("Provider")).toBeInTheDocument();
expect(screen.getByText("Margin")).toBeInTheDocument();
expect(screen.getByText("Actions")).toBeInTheDocument();
});
it("should display the provider display name", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("OpenAI")).toBeInTheDocument();
});
it("should display the global provider as 'Global (All Providers)'", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ global: 0.05 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("Global (All Providers)")).toBeInTheDocument();
});
it("should display a numeric margin as a percentage", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("10.0%")).toBeInTheDocument();
});
it("should display a fixed amount margin with dollar sign", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ openai: { fixed_amount: 0.001 } }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
expect(screen.getByText("$0.001000")).toBeInTheDocument();
});
it("should display a combined percentage and fixed margin", () => {
renderWithProviders(
<ProviderMarginTable
marginConfig={{ openai: { percentage: 0.1, fixed_amount: 0.001 } }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderMarginTable
marginConfig={{ global: 0.05 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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(
<ProviderMarginTable
marginConfig={{ openai: 0.1 }}
onMarginChange={onMarginChange}
onRemoveProvider={onRemoveProvider}
/>
);
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,
});
});
});
});

View file

@ -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<string, string> = {
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();
});
});
});

View file

@ -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<string, string> = {
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,
});
});
});
});

View file

@ -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<string, number>;
return Object.entries(raw)
.map(([model, value]) => ({ model, spend: value }))
.sort((a, b) => b.spend - a.spend);
}, [project?.model_spend]);
if (isLoading) {
return (
<Content
style={{
padding: token.paddingLG,
paddingInline: token.paddingLG * 2,
}}
>
<Flex justify="center" align="center" style={{ minHeight: 300 }}>
<Spin indicator={<LoadingOutlined spin />} size="large" />
</Flex>
</Content>
);
}
if (!project) {
return (
<Content
style={{
padding: token.paddingLG,
paddingInline: token.paddingLG * 2,
}}
>
<Button icon={<ArrowLeftIcon size={16} />} onClick={onBack} type="text" style={{ marginBottom: 16 }} />
<Empty description="Project not found" />
</Content>
);
}
return (
<Content style={{ padding: token.paddingLG, paddingInline: token.paddingLG * 2 }}>
{/* Header */}
<div
style={{
display: "flex",
justifyContent: "space-between",
alignItems: "center",
marginBottom: 24,
}}
>
<div style={{ display: "flex", alignItems: "center", gap: 16 }}>
<Button icon={<ArrowLeftIcon size={16} />} onClick={onBack} type="text" />
<div>
<Flex align="center" gap={8}>
<Title level={2} style={{ margin: 0 }}>
{project.project_alias ?? project.project_id}
</Title>
<Tag color={project.blocked ? "red" : "green"}>{project.blocked ? "Blocked" : "Active"}</Tag>
</Flex>
<Text type="secondary">
ID: <Text copyable>{project.project_id}</Text>
</Text>
</div>
</div>
<Button type="primary" icon={<EditIcon size={16} />} onClick={() => setIsEditModalVisible(true)}>
Edit Project
</Button>
</div>
{/* Project Details */}
<Row style={{ marginBottom: 24 }}>
<Card>
<Descriptions title="Project Details" column={1}>
<Descriptions.Item label="Description">{project.description || "\u2014"}</Descriptions.Item>
<Descriptions.Item label="Created">
{new Date(project.created_at).toLocaleString()}
{project.created_by && (
<Text>
&nbsp;{"by"}&nbsp;
<DefaultProxyAdminTag userId={project.created_by} />
</Text>
)}
</Descriptions.Item>
<Descriptions.Item label="Last Updated">
{new Date(project.updated_at).toLocaleString()}
{project.updated_by && (
<Text>
&nbsp;{"by"}&nbsp;
<DefaultProxyAdminTag userId={project.updated_by} />
</Text>
)}
</Descriptions.Item>
</Descriptions>
</Card>
</Row>
{/* Spend / Budget */}
<Row gutter={[16, 16]} style={{ marginBottom: 24 }}>
<Col xs={24} lg={8}>
<Card
title={
<Flex align="center" gap={8}>
<DollarSignIcon size={16} />
Budget
</Flex>
}
style={{ height: "100%" }}
>
<Flex vertical gap={16}>
<div>
<Text strong style={{ fontSize: 28, lineHeight: 1 }}>
${spend.toFixed(2)}
</Text>
<br />
<Text type="secondary">{hasLimit ? `of $${maxBudget.toFixed(2)} budget` : "No budget limit"}</Text>
</div>
{hasLimit && (
<div>
<Progress percent={Math.round(spendPercent * 10) / 10} strokeColor={spendColor} showInfo={false} />
<Text type="secondary" style={{ fontSize: 12 }}>
{(Math.round(spendPercent * 10) / 10).toFixed(1)}% utilized
</Text>
</div>
)}
</Flex>
</Card>
</Col>
<Col xs={24} lg={16}>
<Card title="Spend by Model" style={{ height: "100%" }}>
{modelSpendData.length > 0 ? (
<BarChart
data={modelSpendData}
index="model"
categories={["spend"]}
colors={["cyan"]}
layout="vertical"
valueFormatter={(value) => `$${value.toFixed(4)}`}
yAxisWidth={140}
showLegend={false}
style={{ height: Math.max(modelSpendData.length * 40, 120) }}
/>
) : (
<Empty description="No model spend recorded yet" image={Empty.PRESENTED_IMAGE_SIMPLE} />
)}
</Card>
</Col>
</Row>
{/* Keys & Team */}
<Row gutter={[16, 16]} style={{ marginBottom: 24 }}>
<Col xs={24} lg={12}>
<ProjectKeysSection projectId={project.project_id} />
</Col>
<Col xs={24} lg={12}>
<Card
title={
<Flex align="center" gap={8}>
<UsersIcon size={16} />
Team
</Flex>
}
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 (
<Flex vertical gap={12}>
{/* Team name + ID */}
<div>
<Text strong style={{ fontSize: 16 }}>
{teamInfo.team_alias || teamInfo.team_id}
</Text>
<br />
<Text type="secondary" style={{ fontSize: 12 }}>
ID:{" "}
<Text copyable style={{ fontSize: 12 }}>
{teamInfo.team_id}
</Text>
</Text>
</div>
{/* Models */}
<div>
<Text type="secondary" style={{ fontSize: 12, display: "block", marginBottom: 4 }}>
Models
</Text>
{(teamInfo.models?.length ?? 0) > 0 ? (
<Flex wrap="wrap" gap={4} style={{ maxHeight: 60, overflow: "hidden" }}>
{teamInfo.models?.map((m: string) => (
<Tag key={m} style={{ margin: 0 }}>
{m}
</Tag>
))}
</Flex>
) : (
<Text type="secondary">All models</Text>
)}
</div>
{/* Budget + Spend compact */}
<div>
<Flex justify="space-between" align="center" style={{ marginBottom: 2 }}>
<Text type="secondary" style={{ fontSize: 12 }}>
Spend
</Text>
<Text style={{ fontSize: 12 }}>
${teamSpend.toFixed(2)}
{teamHasLimit ? (
<Text type="secondary" style={{ fontSize: 12 }}>
{" "}
/ ${teamBudget.toFixed(2)}
</Text>
) : (
<Text type="secondary" style={{ fontSize: 12 }}>
{" "}
(Unlimited)
</Text>
)}
</Text>
</Flex>
{teamHasLimit && (
<Progress
percent={Math.round(teamPercent * 10) / 10}
strokeColor={teamColor}
size="small"
showInfo={false}
/>
)}
</div>
{/* Members */}
<Flex justify="space-between">
<Text type="secondary" style={{ fontSize: 12 }}>
Members
</Text>
<Text style={{ fontSize: 12 }}>{teamInfo.members_with_roles?.length ?? 0}</Text>
</Flex>
</Flex>
);
})()
) : project.team_id ? (
<Flex justify="center" align="center" style={{ padding: 16 }}>
<Spin indicator={<LoadingOutlined spin />} size="small" />
</Flex>
) : (
<Empty description="No team assigned" image={Empty.PRESENTED_IMAGE_SIMPLE} />
)}
</Card>
</Col>
</Row>
{/* Edit Modal */}
<EditProjectModal isOpen={isEditModalVisible} project={project} onClose={() => setIsEditModalVisible(false)} />
</Content>
);
}

View file

@ -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<string>("");
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 (
<Card
title={
<Flex align="center" gap={8}>
<KeyIcon size={16} />
Keys
</Flex>
}
style={{ height: "100%" }}
>
<Flex justify="space-between" align="center" style={{ marginBottom: 12 }}>
<Input
prefix={<SearchIcon size={14} />}
placeholder="Filter by key name..."
style={{ maxWidth: 220 }}
value={keyAlias}
onChange={(e) => setKeyAlias(e.target.value)}
allowClear
size="small"
/>
<Pagination
current={page}
total={totalCount}
pageSize={PAGE_SIZE}
onChange={setPage}
size="small"
showSizeChanger={false}
showTotal={(total) => `${total} keys`}
/>
</Flex>
<ProjectKeysTable
keys={keys}
loading={isLoading ? { indicator: <Spin indicator={<LoadingOutlined spin />} /> } : false}
/>
</Card>
);
}

View file

@ -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<KeyResponse> = [
{
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 (
<Tooltip title={email}>
<DefaultProxyAdminTag userId={email} />
</Tooltip>
);
},
},
{
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 (
<Table
columns={columns}
dataSource={keys}
rowKey="token"
loading={loading}
pagination={false}
size="small"
locale={{ emptyText: <Empty description="No keys found" image={Empty.PRESENTED_IMAGE_SIMPLE} /> }}
/>
);
}

View file

@ -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<ProjectFormValues>();
const createMutation = useCreateProject();
const [selectedTeam, setSelectedTeam] = useState<Team | null>(null);
const [modelsToPick, setModelsToPick] = useState<string[]>([]);
// 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<string, number> = {};
const modelTpmLimit: Record<string, number> = {};
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<string, unknown> = {};
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 (
<Modal
title={
@ -123,226 +60,23 @@ export function CreateProjectModal({ isOpen, onClose }: CreateProjectModalProps)
open={isOpen}
onCancel={handleCancel}
width={720}
destroyOnHidden
footer={[
<Button key="cancel" onClick={handleCancel}>
Cancel
</Button>,
<Button key="submit" type="primary" icon={<FolderAddOutlined />} loading={createMutation.isPending} onClick={handleSubmit}>
<Button
key="submit"
type="primary"
icon={<FolderAddOutlined />}
loading={createMutation.isPending}
onClick={handleSubmit}
>
Create Project
</Button>,
]}
>
<Form
form={form}
layout="vertical"
initialValues={{
isBlocked: false,
}}
style={{ marginTop: 24 }}
>
{/* Basic Info */}
<Typography.Text
strong
style={{ fontSize: 13, color: "#374151", textTransform: "uppercase", letterSpacing: "0.05em" }}
>
Basic Information
</Typography.Text>
<Divider style={{ marginTop: 8, marginBottom: 16 }} />
<Row gutter={24}>
<Col span={12}>
<Form.Item
name="project_alias"
label="Project Name"
rules={[{ required: true, message: "Please enter a project name" }]}
>
<Input placeholder="e.g. Customer Support Bot" />
</Form.Item>
</Col>
<Col span={12}>
<Form.Item name="team_id" label="Team" rules={[{ required: true, message: "Please select a team" }]}>
<Select
showSearch
placeholder="Search or select a team"
onChange={handleTeamChange}
allowClear
optionLabelProp="label"
filterOption={(input, option) => {
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) => (
<Select.Option key={team.team_id} value={team.team_id} label={team.team_alias || team.team_id}>
<span style={{ fontWeight: 500 }}>{team.team_alias}</span>{" "}
<span style={{ color: "#9ca3af" }}>({team.team_id})</span>
</Select.Option>
))}
</Select>
</Form.Item>
</Col>
</Row>
<Row>
<Col span={24}>
<Form.Item name="description" label="Description">
<Input.TextArea placeholder="Describe the purpose of this project" rows={3} />
</Form.Item>
</Col>
</Row>
<Row>
<Col span={24}>
<Form.Item
name="models"
label="Allowed Models (scoped to selected team's models)"
help={!selectedTeam ? "Select a team first to see available models" : undefined}
>
<Select
mode="multiple"
placeholder={selectedTeam ? "Select models" : "Select a team first"}
disabled={!selectedTeam}
allowClear
maxTagCount="responsive"
onChange={(values) => {
if (values.includes("all-team-models")) {
form.setFieldsValue({ models: ["all-team-models"] });
}
}}
>
<Select.Option key="all-team-models" value="all-team-models">
All Team Models
</Select.Option>
{modelsToPick.map((model) => (
<Select.Option key={model} value={model}>
{getModelDisplayName(model)}
</Select.Option>
))}
</Select>
</Form.Item>
</Col>
</Row>
<Row gutter={24}>
<Col span={12}>
<Form.Item name="max_budget" label="Max Budget (USD)">
<InputNumber prefix="$" style={{ width: "100%" }} placeholder="0.00" min={0} precision={2} />
</Form.Item>
</Col>
</Row>
{/* Advanced Settings */}
<Row>
<Col span={24}>
<Collapse ghost style={{ background: "#f9fafb", borderRadius: 8, border: "1px solid #e5e7eb" }}>
<Collapse.Panel
header={
<Typography.Text strong style={{ color: "#374151" }}>
Advanced Settings
</Typography.Text>
}
key="1"
>
<Flex align="center" gap={12}>
<Typography.Text strong>Block Project</Typography.Text>
<Form.Item name="isBlocked" valuePropName="checked" noStyle>
<Switch />
</Form.Item>
</Flex>
<Form.Item noStyle shouldUpdate={(prev, cur) => prev.isBlocked !== cur.isBlocked}>
{({ getFieldValue }) =>
getFieldValue("isBlocked") ? (
<Alert
banner
type="warning"
showIcon
message="All API requests using keys under this project will be rejected."
style={{ marginTop: 12 }}
/>
) : null
}
</Form.Item>
<Divider />
<Typography.Text strong style={{ display: "block", marginBottom: 12 }}>
Model-Specific Limits
</Typography.Text>
<Form.List name="modelLimits">
{(fields, { add, remove }) => (
<>
{fields.map(({ key, name, ...restField }) => (
<Space key={key} style={{ display: "flex", marginBottom: 8 }} align="baseline">
<Form.Item
{...restField}
name={[name, "model"]}
rules={[{ required: true, message: "Missing model" }]}
>
<Input placeholder="Model name (e.g. gpt-4)" />
</Form.Item>
<Form.Item {...restField} name={[name, "tpm"]}>
<InputNumber placeholder="TPM Limit" min={0} />
</Form.Item>
<Form.Item {...restField} name={[name, "rpm"]}>
<InputNumber placeholder="RPM Limit" min={0} />
</Form.Item>
<MinusCircleOutlined onClick={() => remove(name)} style={{ color: "#ef4444" }} />
</Space>
))}
<Form.Item>
<Button type="dashed" onClick={() => add()} block icon={<PlusOutlined />}>
Add Model Limit
</Button>
</Form.Item>
</>
)}
</Form.List>
<Divider />
<Typography.Text strong style={{ display: "block", marginBottom: 12 }}>
Metadata
</Typography.Text>
<Form.List name="metadata">
{(fields, { add, remove }) => (
<>
{fields.map(({ key, name, ...restField }) => (
<Space key={key} style={{ display: "flex", marginBottom: 8 }} align="baseline">
<Form.Item
{...restField}
name={[name, "key"]}
rules={[{ required: true, message: "Missing key" }]}
>
<Input placeholder="Key" />
</Form.Item>
<Form.Item
{...restField}
name={[name, "value"]}
rules={[{ required: true, message: "Missing value" }]}
>
<Input placeholder="Value" />
</Form.Item>
<MinusCircleOutlined onClick={() => remove(name)} style={{ color: "#ef4444" }} />
</Space>
))}
<Form.Item>
<Button type="dashed" onClick={() => add()} block icon={<PlusOutlined />}>
Add Key-Value Pair
</Button>
</Form.Item>
</>
)}
</Form.List>
</Collapse.Panel>
</Collapse>
</Col>
</Row>
</Form>
<ProjectBaseForm form={form} />
</Modal>
);
}

View file

@ -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<ProjectFormValues>();
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<string, unknown>;
const rpmLimits = (metadataObj.model_rpm_limit ?? {}) as Record<string, number>;
const tpmLimits = (metadataObj.model_tpm_limit ?? {}) as Record<string, number>;
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 (
<Modal
title={
<Typography.Text strong style={{ fontSize: 18 }}>
Edit Project
</Typography.Text>
}
open={isOpen}
onCancel={onClose}
width={720}
destroyOnHidden
footer={[
<Button key="cancel" onClick={onClose}>
Cancel
</Button>,
<Button
key="submit"
type="primary"
icon={<SaveOutlined />}
loading={updateMutation.isPending}
onClick={handleSubmit}
>
Save Changes
</Button>,
]}
>
<ProjectBaseForm form={form} />
</Modal>
);
}

View file

@ -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<ProjectFormValues>;
}
export function ProjectBaseForm({
form,
}: ProjectBaseFormProps) {
const { accessToken, userId, userRole } = useAuthorized();
const { data: teams } = useTeams();
const [selectedTeam, setSelectedTeam] = useState<Team | null>(null);
const [modelsToPick, setModelsToPick] = useState<string[]>([]);
// 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 (
<Form
form={form}
layout="vertical"
name="project_form"
initialValues={{ isBlocked: false }}
style={{ marginTop: 24 }}
>
{/* Basic Info */}
<Typography.Text
strong
style={{
fontSize: 13,
color: "#374151",
textTransform: "uppercase",
letterSpacing: "0.05em",
}}
>
Basic Information
</Typography.Text>
<Divider style={{ marginTop: 8, marginBottom: 16 }} />
<Row gutter={24}>
<Col span={12}>
<Form.Item
name="project_alias"
label="Project Name"
rules={[
{ required: true, message: "Please enter a project name" },
]}
>
<Input placeholder="e.g. Customer Support Bot" />
</Form.Item>
</Col>
<Col span={12}>
<Form.Item
name="team_id"
label="Team"
rules={[{ required: true, message: "Please select a team" }]}
>
<Select
showSearch
placeholder="Search or select a team"
onChange={handleTeamChange}
allowClear
optionLabelProp="label"
filterOption={(input, option) => {
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) => (
<Select.Option
key={team.team_id}
value={team.team_id}
label={team.team_alias || team.team_id}
>
<span style={{ fontWeight: 500 }}>{team.team_alias}</span>{" "}
<span style={{ color: "#9ca3af" }}>({team.team_id})</span>
</Select.Option>
))}
</Select>
</Form.Item>
</Col>
</Row>
<Row>
<Col span={24}>
<Form.Item name="description" label="Description">
<Input.TextArea
placeholder="Describe the purpose of this project"
rows={3}
/>
</Form.Item>
</Col>
</Row>
<Row>
<Col span={24}>
<Form.Item
name="models"
label="Allowed Models (scoped to selected team's models)"
help={
!selectedTeam
? "Select a team first to see available models"
: undefined
}
>
<Select
mode="multiple"
placeholder={
selectedTeam ? "Select models" : "Select a team first"
}
disabled={!selectedTeam}
allowClear
maxTagCount="responsive"
onChange={(values) => {
if (values.includes("all-team-models")) {
form.setFieldsValue({ models: ["all-team-models"] });
}
}}
>
<Select.Option key="all-team-models" value="all-team-models">
All Team Models
</Select.Option>
{modelsToPick.map((model) => (
<Select.Option key={model} value={model}>
{getModelDisplayName(model)}
</Select.Option>
))}
</Select>
</Form.Item>
</Col>
</Row>
<Row gutter={24}>
<Col span={12}>
<Form.Item name="max_budget" label="Max Budget (USD)">
<InputNumber
prefix="$"
style={{ width: "100%" }}
placeholder="0.00"
min={0}
precision={2}
/>
</Form.Item>
</Col>
</Row>
{/* Advanced Settings */}
<Row>
<Col span={24}>
<Collapse
ghost
style={{
background: "#f9fafb",
borderRadius: 8,
border: "1px solid #e5e7eb",
}}
items={[
{
key: "1",
label: (
<Typography.Text strong style={{ color: "#374151" }}>
Advanced Settings
</Typography.Text>
),
children: (
<>
<Flex align="center" gap={12}>
<Typography.Text strong>Block Project</Typography.Text>
<Form.Item name="isBlocked" valuePropName="checked" noStyle>
<Switch />
</Form.Item>
</Flex>
<Form.Item
noStyle
shouldUpdate={(prev, cur) => prev.isBlocked !== cur.isBlocked}
>
{({ getFieldValue }) =>
getFieldValue("isBlocked") ? (
<Alert
banner
type="warning"
showIcon
message="All API requests using keys under this project will be rejected."
style={{ marginTop: 12 }}
/>
) : null
}
</Form.Item>
<Divider />
<Typography.Text
strong
style={{ display: "block", marginBottom: 12 }}
>
Model-Specific Limits
</Typography.Text>
<Form.List name="modelLimits">
{(fields, { add, remove }) => (
<>
{fields.map(({ key, name, ...restField }) => (
<Space
key={key}
style={{ display: "flex", marginBottom: 8 }}
align="baseline"
>
<Form.Item
{...restField}
name={[name, "model"]}
rules={[
{ required: true, message: "Missing model" },
{
validator: (_, value) => {
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();
},
},
]}
>
<Input placeholder="Model name (e.g. gpt-4)" />
</Form.Item>
<Form.Item {...restField} name={[name, "tpm"]}>
<InputNumber placeholder="TPM Limit" min={0} />
</Form.Item>
<Form.Item {...restField} name={[name, "rpm"]}>
<InputNumber placeholder="RPM Limit" min={0} />
</Form.Item>
<MinusCircleOutlined
onClick={() => remove(name)}
style={{ color: "#ef4444" }}
/>
</Space>
))}
<Form.Item>
<Button
type="dashed"
onClick={() => add()}
block
icon={<PlusOutlined />}
>
Add Model Limit
</Button>
</Form.Item>
</>
)}
</Form.List>
<Divider />
<Typography.Text
strong
style={{ display: "block", marginBottom: 12 }}
>
Metadata
</Typography.Text>
<Form.List name="metadata">
{(fields, { add, remove }) => (
<>
{fields.map(({ key, name, ...restField }) => (
<Space
key={key}
style={{ display: "flex", marginBottom: 8 }}
align="baseline"
>
<Form.Item
{...restField}
name={[name, "key"]}
rules={[
{ required: true, message: "Missing key" },
{
validator: (_, value) => {
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();
},
},
]}
>
<Input placeholder="Key" />
</Form.Item>
<Form.Item
{...restField}
name={[name, "value"]}
rules={[
{ required: true, message: "Missing value" },
]}
>
<Input placeholder="Value" />
</Form.Item>
<MinusCircleOutlined
onClick={() => remove(name)}
style={{ color: "#ef4444" }}
/>
</Space>
))}
<Form.Item>
<Button
type="dashed"
onClick={() => add()}
block
icon={<PlusOutlined />}
>
Add Key-Value Pair
</Button>
</Form.Item>
</>
)}
</Form.List>
</>
),
},
]}
/>
</Col>
</Row>
</Form>
);
}

View file

@ -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<string, number> = {};
const modelTpmLimit: Record<string, number> = {};
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<string, unknown> = {};
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 }),
};
}

View file

@ -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<string | null>(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}
</Text>
@ -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 <Spin indicator={<LoadingOutlined spin />} size="small" />;
return record.team_id;
},
},
{
@ -144,10 +153,25 @@ export function ProjectsPage() {
},
];
if (selectedProjectId) {
return (
<ProjectDetail
projectId={selectedProjectId}
onBack={() => setSelectedProjectId(null)}
/>
);
}
return (
<Content
style={{ padding: token.paddingLG, paddingInline: token.paddingLG * 2 }}
>
<Alert
message="Projects is currently in beta. Features and behavior may change without notice."
type="warning"
showIcon
style={{ marginBottom: 16 }}
/>
<Flex
justify="space-between"
align="center"
@ -184,21 +208,22 @@ export function ProjectsPage() {
onChange={(e) => setSearchText(e.target.value)}
allowClear
/>
<Pagination
current={currentPage}
total={filteredProjects.length}
pageSize={pageSize}
onChange={(page) => setCurrentPage(page)}
size="small"
showTotal={(total) => `${total} projects`}
showSizeChanger={false}
/>
</Flex>
<Table
columns={columns}
dataSource={filteredProjects}
dataSource={filteredProjects.slice((currentPage - 1) * pageSize, currentPage * pageSize)}
rowKey="project_id"
loading={isLoading}
pagination={{
current: currentPage,
pageSize,
total: filteredProjects.length,
onChange: (page) => setCurrentPage(page),
size: "small",
showTotal: (total) => `${total} projects`,
showSizeChanger: false,
}}
pagination={false}
/>
</Card>

View file

@ -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<ProjectDropdownProps> = ({
projects,
value,
onChange,
disabled,
loading,
teamId,
}) => {
const filtered = teamId
? projects?.filter((p) => p.team_id === teamId)
: projects;
return (
<Select
showSearch
placeholder="Search or select a project"
value={value}
onChange={onChange}
disabled={disabled}
loading={loading}
allowClear
notFoundContent={loading ? <Spin indicator={<LoadingOutlined spin />} 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) => (
<Select.Option key={project.project_id} value={project.project_id}>
<span className="font-medium">{project.project_alias || project.project_id}</span>{" "}
<span className="text-gray-500">({project.project_id})</span>
</Select.Option>
))}
</Select>
);
};
export default ProjectDropdown;

View file

@ -7,10 +7,10 @@ interface TeamDropdownProps {
value?: string;
onChange?: (value: string) => void;
disabled?: boolean;
loading?: boolean;
}
const TeamDropdown: React.FC<TeamDropdownProps> = ({ teams, value, onChange, disabled }) => {
console.log("disabled", disabled);
const TeamDropdown: React.FC<TeamDropdownProps> = ({ teams, value, onChange, disabled, loading }) => {
return (
<Select
showSearch
@ -18,6 +18,7 @@ const TeamDropdown: React.FC<TeamDropdownProps> = ({ teams, value, onChange, dis
value={value}
onChange={onChange}
disabled={disabled}
loading={loading}
allowClear
filterOption={(input, option) => {
if (!option) return false;

View file

@ -30,6 +30,7 @@ export interface KeyResponse {
config: Record<string, unknown>;
user_id: string;
team_id: string | null;
project_id: string | null;
max_parallel_requests: number;
metadata: Record<string, unknown>;
tpm_limit: number;

View file

@ -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 }) => (
<input
data-testid="project-dropdown"
value={value || ""}
onChange={(e) => onChange?.(e.target.value)}
/>
),
}));
vi.mock("../common_components/AccessGroupSelector", () => ({
default: ({ value = [], onChange }: { value?: string[]; onChange?: (v: string[]) => void }) => (
<input

View file

@ -1,5 +1,6 @@
"use client";
import { keyKeys } from "@/app/(dashboard)/hooks/keys/useKeys";
import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { InfoCircleOutlined } from "@ant-design/icons";
@ -21,6 +22,7 @@ import PremiumLoggingSettings from "../common_components/PremiumLoggingSettings"
import RateLimitTypeFormItem from "../common_components/RateLimitTypeFormItem";
import RouterSettingsAccordion, { RouterSettingsAccordionValue } from "../common_components/RouterSettingsAccordion";
import TeamDropdown from "../common_components/team_dropdown";
import ProjectDropdown from "../common_components/ProjectDropdown";
import { CreateUserButton } from "../CreateUserButton";
import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key";
import { Team } from "../key_team_helpers/key_list";
@ -143,6 +145,7 @@ export const fetchUserModels = async (
*/
const CreateKey: React.FC<CreateKeyProps> = ({ 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<CreateKeyProps> = ({ team, teams, data, addKey }) => {
const [promptsList, setPromptsList] = useState<string[]>([]);
const [loggingSettings, setLoggingSettings] = useState<any[]>([]);
const [selectedCreateKeyTeam, setSelectedCreateKeyTeam] = useState<Team | null>(team);
const [selectedProjectId, setSelectedProjectId] = useState<string | null>(null);
const [isCreateUserModalVisible, setIsCreateUserModalVisible] = useState(false);
const [newlyCreatedUserId, setNewlyCreatedUserId] = useState<string | null>(null);
const [possibleUIRoles, setPossibleUIRoles] = useState<Record<string, Record<string, string>>>({});
@ -184,6 +188,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
setRouterSettings(null);
setRouterSettingsKey((prev) => prev + 1);
setSelectedAgentId(null);
setSelectedProjectId(null);
};
const handleCancel = () => {
@ -200,6 +205,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
setRouterSettings(null);
setRouterSettingsKey((prev) => prev + 1);
setSelectedAgentId(null);
setSelectedProjectId(null);
};
useEffect(() => {
@ -468,6 +474,14 @@ const CreateKey: React.FC<CreateKeyProps> = ({ 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<CreateKeyProps> = ({ 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<CreateKeyProps> = ({ team, teams, data, addKey }) => {
>
<TeamDropdown
teams={teams}
disabled={selectedProjectId !== null}
loading={!teams}
onChange={(teamId) => {
const selectedTeam = teams?.find((t) => t.team_id === teamId) || null;
setSelectedCreateKeyTeam(selectedTeam);
setSelectedProjectId(null);
form.setFieldValue("project_id", undefined);
}}
/>
</Form.Item>
<Form.Item
label={
<span>
Project{" "}
<Tooltip title="Assign this key to a project. Selecting a project will lock the team to the project's team.">
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
</Tooltip>
</span>
}
name="project_id"
className="mt-4"
>
<ProjectDropdown
projects={projects}
teamId={selectedCreateKeyTeam?.team_id}
loading={isProjectsLoading || !teams}
onChange={(projectId) => {
if (!projectId) {
setSelectedProjectId(null);
setSelectedCreateKeyTeam(null);
form.setFieldValue("team_id", undefined);
return;
}
setSelectedProjectId(projectId);
}}
/>
</Form.Item>
@ -735,9 +795,11 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
}
}}
>
<Option key="all-team-models" value="all-team-models">
All Team Models
</Option>
{!selectedProjectId && (
<Option key="all-team-models" value="all-team-models">
All Team Models
</Option>
)}
{modelsToPick.map((model: string) => (
<Option key={model} value={model}>
{getModelDisplayName(model)}

View file

@ -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");

View file

@ -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<boolean>(keyData.auto_rotate || false);
const [rotationInterval, setRotationInterval] = useState<string>(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({
/>
</Form.Item>
<Form.Item label="Team ID" name="team_id">
<Form.Item
label="Team ID"
name="team_id"
help={hasProject ? "Team is locked because this key belongs to a project" : undefined}
>
<Select
placeholder="Select team"
showSearch
disabled={hasProject}
style={{ width: "100%" }}
filterOption={(input, option) => {
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) => (
<Select.Option key={team.team_id} value={team.team_id}>
{`${team.team_alias} (${team.team_id})`}
@ -609,6 +623,11 @@ export function KeyEditView({
))}
</Select>
</Form.Item>
{hasProject && (
<Form.Item label="Project">
<Input value={projectDisplay ?? ""} disabled />
</Form.Item>
)}
<Form.Item label="Logging Settings" name="logging_settings">
<EditLoggingSettings
value={form.getFieldValue("logging_settings")}

View file

@ -1,6 +1,8 @@
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { useProjects } from "@/app/(dashboard)/hooks/projects/useProjects";
import useTeams from "@/app/(dashboard)/hooks/useTeams";
import { render, screen, waitFor } from "@testing-library/react";
import { renderWithProviders } from "../../../tests/test-utils";
import { screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
@ -14,6 +16,10 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
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(
<KeyInfoView
keyData={MOCK_KEY_DATA}
onClose={() => { }}
@ -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(
<KeyInfoView
keyData={MOCK_KEY_DATA}
onClose={() => { }}
@ -168,7 +175,7 @@ describe("KeyInfoView", () => {
});
const keyData = { ...MOCK_KEY_DATA, user_id: "other-user-id" };
render(
renderWithProviders(
<KeyInfoView keyData={keyData} onClose={() => { }} 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(
<KeyInfoView keyData={keyData} onClose={() => { }} 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(
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
);
@ -260,7 +267,7 @@ describe("KeyInfoView", () => {
});
const keyData = { ...MOCK_KEY_DATA, user_id: "owner-user-id" };
render(
renderWithProviders(
<KeyInfoView keyData={keyData} onClose={() => { }} 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(
<KeyInfoView keyData={keyData} onClose={() => { }} 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(
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
);
@ -342,7 +349,7 @@ describe("KeyInfoView", () => {
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
const onCloseMock = vi.fn();
render(
renderWithProviders(
<KeyInfoView
keyData={MOCK_KEY_DATA}
onClose={onCloseMock}
@ -365,7 +372,7 @@ describe("KeyInfoView", () => {
describe("'Edit Settings' button visibility in the Settings tab", () => {
const renderAndOpenSettingsTab = async (keyData = MOCK_KEY_DATA) => {
render(
renderWithProviders(
<KeyInfoView
keyData={keyData}
onClose={() => {}}
@ -474,7 +481,7 @@ describe("KeyInfoView", () => {
},
};
render(
renderWithProviders(
<KeyInfoView
keyData={keyDataWithGuardrails}
onClose={() => { }}
@ -500,7 +507,7 @@ describe("KeyInfoView", () => {
},
};
render(
renderWithProviders(
<KeyInfoView
keyData={keyDataWithPolicies}
onClose={() => { }}
@ -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(
<KeyInfoView
keyData={undefined}
onClose={() => { }}

View file

@ -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({
<Text>{currentKeyData.team_id || "Not Set"}</Text>
</div>
<div>
<Text className="font-medium">Project</Text>
<Text>
{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"}
</Text>
</div>
<div>
<Text className="font-medium">Organization</Text>
<Text>{currentKeyData.organization_id || "Not Set"}</Text>