mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
commit
93e08a6509
72 changed files with 6259 additions and 621 deletions
|
|
@ -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
208
.github/scripts/close_duplicate_issues.py
vendored
Executable 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
23
.github/workflows/check_duplicate_issues.yml
vendored
23
.github/workflows/check_duplicate_issues.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
47
.github/workflows/scan_duplicate_issues.yml
vendored
Normal file
47
.github/workflows/scan_duplicate_issues.yml
vendored
Normal 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
|
||||
29
.github/workflows/test-linting.yml
vendored
29
.github/workflows/test-linting.yml
vendored
|
|
@ -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 .
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -105,6 +105,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"prometheus",
|
||||
"otel",
|
||||
"datadog",
|
||||
"datadog_metrics",
|
||||
"datadog_llm_observability",
|
||||
"galileo",
|
||||
"braintrust",
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
}
|
||||
]
|
||||
]
|
||||
|
|
|
|||
286
litellm/integrations/datadog/datadog_metrics.py
Normal file
286
litellm/integrations/datadog/datadog_metrics.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
20
litellm/types/integrations/datadog_metrics.py
Normal file
20
litellm/types/integrations/datadog_metrics.py
Normal 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]
|
||||
74
tests/litellm/test_no_hardcoded_secrets.py
Normal file
74
tests/litellm/test_no_hardcoded_secrets.py
Normal 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."
|
||||
)
|
||||
273
tests/test_litellm/integrations/datadog/test_datadog_metrics.py
Normal file
273
tests/test_litellm/integrations/datadog/test_datadog_metrics.py
Normal 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
|
||||
|
|
@ -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}'
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
139
tests/test_litellm/proxy/test_audio_speech_prometheus_hooks.py
Normal file
139
tests/test_litellm/proxy/test_audio_speech_prometheus_hooks.py
Normal 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
|
||||
75
tests/test_litellm/proxy/test_health_check_max_tokens.py
Normal file
75
tests/test_litellm/proxy/test_health_check_max_tokens.py
Normal 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
|
||||
173
tests/test_litellm/proxy/test_update_llm_router_resilience.py
Normal file
173
tests/test_litellm/proxy/test_update_llm_router_resilience.py
Normal 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")
|
||||
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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", () => {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -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 });
|
||||
},
|
||||
});
|
||||
};
|
||||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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,
|
||||
});
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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,
|
||||
});
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
{"by"}
|
||||
<DefaultProxyAdminTag userId={project.created_by} />
|
||||
</Text>
|
||||
)}
|
||||
</Descriptions.Item>
|
||||
<Descriptions.Item label="Last Updated">
|
||||
{new Date(project.updated_at).toLocaleString()}
|
||||
{project.updated_by && (
|
||||
<Text>
|
||||
{"by"}
|
||||
<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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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} /> }}
|
||||
/>
|
||||
);
|
||||
}
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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 }),
|
||||
};
|
||||
}
|
||||
|
|
@ -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>
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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")}
|
||||
|
|
|
|||
|
|
@ -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={() => { }}
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue