mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge branch 'main' into litellm_dev_02_24_2026_p2
This commit is contained in:
commit
2278d0220b
49 changed files with 3045 additions and 342 deletions
|
|
@ -174,6 +174,8 @@ When opening issues or pull requests, follow these templates:
|
|||
3. **Rate Limits**: Respect provider rate limits in tests
|
||||
4. **Memory Usage**: Be mindful of memory usage in streaming scenarios
|
||||
5. **Dependencies**: Keep dependencies minimal and well-justified
|
||||
6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections
|
||||
7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks
|
||||
|
||||
## HELPFUL RESOURCES
|
||||
|
||||
|
|
|
|||
|
|
@ -97,6 +97,10 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components:
|
|||
- Integration tests for each provider in `tests/llm_translation/`
|
||||
- Proxy tests in `tests/proxy_unit_tests/`
|
||||
- Load tests in `tests/load_tests/`
|
||||
- **Always add tests when adding new entity types or features** — if the existing test file covers other entity types, add corresponding tests for the new one
|
||||
|
||||
### UI / Backend Consistency
|
||||
- When wiring a new UI entity type to an existing backend endpoint, verify the backend API contract (single value vs. array, required vs. optional params) and ensure the UI controls match — e.g., use a single-select dropdown when the backend accepts a single value, not a multi-select
|
||||
|
||||
### Database Migrations
|
||||
- Prisma handles schema migrations
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ Use LiteLLM to call Google AI's generateContent endpoints for text generation, m
|
|||
| Streaming | ✅ | |
|
||||
| Fallbacks | ✅ | between supported models |
|
||||
| Loadbalancing | ✅ | between supported models |
|
||||
| Metadata Tracking | ✅ | passes trace ID, metadata to observability callbacks (e.g. S3, Langfuse) |
|
||||
|
||||
## Usage
|
||||
---
|
||||
|
|
|
|||
19
docs/my-website/docs/proxy/credential_usage_tracking.md
Normal file
19
docs/my-website/docs/proxy/credential_usage_tracking.md
Normal file
|
|
@ -0,0 +1,19 @@
|
|||
# Credential Usage Tracking
|
||||
|
||||
When a model is attached to a [reusable credential](./ui_credentials.md), LiteLLM automatically injects the credential name as a tag on every request that uses that model. This means credential-level spend and usage are tracked with zero extra configuration.
|
||||
|
||||
## How It Works
|
||||
|
||||
When you attach a model to a reusable credential via `litellm_credential_name`, each request routed through that model is tagged `Credential: <name>` (for example, `Credential: xAI`). This tag flows into `DailyTagSpend` and appears in the **Tag** view on the Usage page, where you can filter spend and usage by credential.
|
||||
|
||||
If a model has no credential attached, behavior is unchanged—no credential tag is added.
|
||||
|
||||
## Viewing Credential Usage
|
||||
|
||||
In the Admin UI, go to **Usage → Tag** and look for tags with the `Credential: ` prefix. These represent aggregated spend and token usage across all requests that used that credential.
|
||||
|
||||
## Related Documentation
|
||||
|
||||
- [Adding LLM Credentials](./ui_credentials.md) - How to create and attach reusable credentials to models
|
||||
- [Tag Budgets](./tag_budgets.md) - Setting spend limits on tags
|
||||
- [Tag Routing](./tag_routing.md) - Routing requests based on tags
|
||||
|
|
@ -37,11 +37,11 @@ The following rules determine which headers are forwarded (see [`_get_forwardabl
|
|||
|
||||
| Rule | Example | Forwarded? |
|
||||
|---|---|---|
|
||||
| Headers starting with `x-` | `x-trace-id`, `x-custom-header`, `x-request-source` | ✅ Yes |
|
||||
| `anthropic-beta` header | `anthropic-beta: prompt-caching-2024-07-31` | ✅ Yes |
|
||||
| Headers starting with `x-stainless-*` | `x-stainless-lang`, `x-stainless-arch` | ❌ No (causes OpenAI SDK issues) |
|
||||
| Standard HTTP headers | `Authorization`, `Content-Type`, `Host` | ❌ No |
|
||||
| Other provider headers | `Accept`, `User-Agent` | ❌ No |
|
||||
| Headers starting with `x-` | `x-trace-id`, `x-custom-header`, `x-request-source` | Yes |
|
||||
| `anthropic-beta` header | `anthropic-beta: prompt-caching-2024-07-31` | Yes |
|
||||
| Headers starting with `x-stainless-*` | `x-stainless-lang`, `x-stainless-arch` | No (causes OpenAI SDK issues) |
|
||||
| Standard HTTP headers | `Authorization`, `Content-Type`, `Host` | No |
|
||||
| Other provider headers | `Accept`, `User-Agent` | No |
|
||||
|
||||
### Additional Header Mechanisms
|
||||
|
||||
|
|
@ -61,6 +61,125 @@ general_settings:
|
|||
forward_client_headers_to_llm_api: true
|
||||
```
|
||||
|
||||
## Forward LLM Provider Authentication Headers
|
||||
|
||||
**New in v1.82+**: By default, LiteLLM strips authentication headers like `x-api-key`, `x-goog-api-key`, and `api-key` from client requests for security (these are typically used to authenticate with the proxy itself). However, you can enable forwarding of these LLM provider authentication headers to allow **Bring Your Own Key (BYOK)** scenarios where clients send their own API keys to the LLM provider.
|
||||
|
||||
### Configuration
|
||||
|
||||
Add `forward_llm_provider_auth_headers: true` to your `general_settings`:
|
||||
|
||||
```yaml
|
||||
general_settings:
|
||||
forward_client_headers_to_llm_api: true
|
||||
forward_llm_provider_auth_headers: true # 👈 Enable BYOK
|
||||
```
|
||||
|
||||
### Which Headers Are Forwarded
|
||||
|
||||
When `forward_llm_provider_auth_headers: true`, the following LLM provider authentication headers are preserved and forwarded:
|
||||
|
||||
| Header | Provider | Example |
|
||||
|--------|----------|---------|
|
||||
| `x-api-key` | Anthropic, Azure AI, Databricks | `x-api-key: sk-ant-api03-...` |
|
||||
| `x-goog-api-key` | Google AI Studio | `x-goog-api-key: AIza...` |
|
||||
| `api-key` | Azure OpenAI | `api-key: your-azure-key` |
|
||||
| `ocp-apim-subscription-key` | Azure APIM | `ocp-apim-subscription-key: your-key` |
|
||||
|
||||
:::warning Important Security Note
|
||||
The proxy's `Authorization` header (used for proxy authentication) is **never** forwarded to LLM providers, even with this setting enabled. This ensures your proxy authentication remains secure.
|
||||
:::
|
||||
|
||||
### Use Case: Client-Side API Keys (BYOK)
|
||||
|
||||
This feature enables scenarios where:
|
||||
1. **Clients bring their own LLM provider API keys** instead of using keys configured in the proxy
|
||||
2. **Multi-tenant applications** where each tenant has their own Anthropic/OpenAI account
|
||||
3. **Development environments** where developers use their personal API keys through a shared proxy
|
||||
|
||||
#### Example: Anthropic BYOK
|
||||
|
||||
```yaml
|
||||
# proxy_config.yaml
|
||||
model_list:
|
||||
- model_name: claude-sonnet-4
|
||||
litellm_params:
|
||||
model: anthropic/claude-sonnet-4-20250514
|
||||
# No api_key configured! Will use client's key
|
||||
|
||||
general_settings:
|
||||
forward_client_headers_to_llm_api: true
|
||||
forward_llm_provider_auth_headers: true # Enable BYOK
|
||||
```
|
||||
|
||||
Client request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/messages" \
|
||||
-H "Authorization: Bearer sk-proxy-auth-123" \ # Proxy authentication (stripped)
|
||||
-H "x-api-key: sk-ant-api03-YOUR-KEY..." \ # Client's Anthropic key (forwarded!)
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "claude-sonnet-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"max_tokens": 100
|
||||
}'
|
||||
```
|
||||
|
||||
#### Example: Google AI Studio BYOK
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gemini-pro
|
||||
litellm_params:
|
||||
model: gemini/gemini-1.5-pro
|
||||
# No api_key configured
|
||||
|
||||
general_settings:
|
||||
forward_client_headers_to_llm_api: true
|
||||
forward_llm_provider_auth_headers: true
|
||||
```
|
||||
|
||||
Client request:
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/chat/completions" \
|
||||
-H "Authorization: Bearer sk-proxy-auth-123" \
|
||||
-H "x-goog-api-key: AIza..." \
|
||||
-d '{
|
||||
"model": "gemini-pro",
|
||||
"messages": [{"role": "user", "content": "Hello"}]
|
||||
}'
|
||||
```
|
||||
|
||||
### Security Considerations
|
||||
|
||||
**When to Use This Feature:**
|
||||
- Internal tools where you trust all clients
|
||||
- Development/testing environments
|
||||
- Multi-tenant apps with proper client authentication
|
||||
- Scenarios where you want clients to use their own API keys
|
||||
|
||||
**When NOT to Use:**
|
||||
- Public APIs where you don't trust all clients
|
||||
- When you want centralized billing/cost control
|
||||
- When you need to enforce rate limits at the proxy level
|
||||
|
||||
### Backward Compatibility
|
||||
|
||||
For backward compatibility, if you have `forward_client_headers_to_llm_api: true` but don't explicitly set `forward_llm_provider_auth_headers`, the behavior is:
|
||||
- **Default**: LLM provider auth headers are **NOT** forwarded (safe default)
|
||||
- **Explicit `true`**: LLM provider auth headers **ARE** forwarded (BYOK enabled)
|
||||
|
||||
```yaml
|
||||
# Safe default - auth headers NOT forwarded
|
||||
general_settings:
|
||||
forward_client_headers_to_llm_api: true
|
||||
|
||||
# BYOK enabled - auth headers ARE forwarded
|
||||
general_settings:
|
||||
forward_client_headers_to_llm_api: true
|
||||
forward_llm_provider_auth_headers: true # 👈 Opt-in required
|
||||
```
|
||||
|
||||
## Enable for a Model Group
|
||||
|
||||
Add the `forward_client_headers_to_llm_api` setting under `model_group_settings` in your configuration:
|
||||
|
|
|
|||
|
|
@ -46,6 +46,10 @@ Go to Add Model -> Existing Credentials -> Select your credential in the dropdow
|
|||
|
||||
<Image img={require('../../img/use_model_cred.png')} />
|
||||
|
||||
## Usage Tracking
|
||||
|
||||
Models attached to a reusable credential are automatically tracked in the Usage page. Each request is tagged `Credential: <name>` and appears in the **Tag** view, so you can filter spend and usage by credential without any extra configuration. See [Credential Usage Tracking](./credential_usage_tracking.md) for details.
|
||||
|
||||
## Frequently Asked Questions
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -592,6 +592,21 @@ def test_pii_masking_allows_normal_text():
|
|||
|
||||
## Part 7: Troubleshooting
|
||||
|
||||
### Issue: Guardrail failure: non-JSON response from Presidio
|
||||
|
||||
**Symptom:** You receive an error indicating `expected application/json Content-Type but received text/html` or similar.
|
||||
|
||||
**Root cause:** Your ingress controller or reverse proxy might be routing the `/analyze` or `/anonymize` POST request to a health endpoint (like `/health` or `/presidio-analyzer/health`) which returns plain text instead of JSON.
|
||||
|
||||
**Fix:** Ensure your `PRESIDIO_ANALYZER_API_BASE` and `PRESIDIO_ANONYMIZER_API_BASE` are correctly pointing directly to the Presidio API endpoints, or that your ingress routes the path correctly without stripping it and inadvertently forwarding to a plain-text health check endpoint.
|
||||
|
||||
**Verification:** You can verify your endpoints using `curl`. It should return a JSON array, not `text/html`:
|
||||
```bash
|
||||
curl -sv -X POST http://your-analyzer-endpoint/analyze \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"text":"test","language":"en"}'
|
||||
```
|
||||
|
||||
### Issue: Presidio Not Detecting PII
|
||||
|
||||
**Check 1: Language Configuration**
|
||||
|
|
|
|||
|
|
@ -37,7 +37,9 @@ from litellm.types.llms.openai import (
|
|||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import (
|
||||
LIST_BATCHES_SUPPORTED_PROVIDERS,
|
||||
OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS,
|
||||
ListBatchesSupportedProvider,
|
||||
LiteLLMBatch,
|
||||
LlmProviders,
|
||||
)
|
||||
|
|
@ -674,7 +676,7 @@ def retrieve_batch(
|
|||
async def alist_batches(
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "hosted_vllm", "vertex_ai"] = "openai",
|
||||
custom_llm_provider: ListBatchesSupportedProvider = "openai",
|
||||
metadata: Optional[Dict[str, str]] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
|
|
@ -717,7 +719,7 @@ async def alist_batches(
|
|||
def list_batches(
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "hosted_vllm", "vertex_ai"] = "openai",
|
||||
custom_llm_provider: ListBatchesSupportedProvider = "openai",
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
**kwargs,
|
||||
|
|
@ -843,8 +845,9 @@ def list_batches(
|
|||
)
|
||||
else:
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message="LiteLLM doesn't support {} for 'list_batch'. Supported providers: openai, azure, vertex_ai.".format(
|
||||
custom_llm_provider
|
||||
message="LiteLLM doesn't support {} for 'list_batch'. Supported providers: {}.".format(
|
||||
custom_llm_provider,
|
||||
", ".join(sorted(LIST_BATCHES_SUPPORTED_PROVIDERS)),
|
||||
),
|
||||
model="n/a",
|
||||
llm_provider=custom_llm_provider,
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ Helper functions for health check calls.
|
|||
|
||||
from typing import TYPE_CHECKING, Callable, Dict, Literal, Optional
|
||||
|
||||
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
|
||||
|
|
@ -82,6 +84,27 @@ class HealthCheckHelpers:
|
|||
"tags": [LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME],
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
async def _batch_health_check(
|
||||
custom_llm_provider: str,
|
||||
model_params: dict,
|
||||
filtered_model_params: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Health check for batch mode.
|
||||
|
||||
Calls list_batches for providers that support it (openai, hosted_vllm, azure,
|
||||
vertex_ai). For all other providers (e.g. bedrock) the batch API surface doesn't
|
||||
include list_batches, so we fall back to acompletion to verify connectivity and
|
||||
credential validity instead.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS:
|
||||
return await litellm.alist_batches(**filtered_model_params)
|
||||
else:
|
||||
return await litellm.acompletion(**model_params)
|
||||
|
||||
@staticmethod
|
||||
def get_mode_handlers(
|
||||
model: str,
|
||||
|
|
@ -176,8 +199,10 @@ class HealthCheckHelpers:
|
|||
api_key=model_params.get("api_key", None),
|
||||
api_version=model_params.get("api_version", None),
|
||||
),
|
||||
"batch": lambda: litellm.alist_batches(
|
||||
**_filter_model_params(model_params=model_params),
|
||||
"batch": lambda: HealthCheckHelpers._batch_health_check(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_params=model_params,
|
||||
filtered_model_params=_filter_model_params(model_params=model_params),
|
||||
),
|
||||
"responses": lambda: litellm.aresponses(
|
||||
**_filter_model_params(model_params=model_params),
|
||||
|
|
|
|||
|
|
@ -593,9 +593,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
user_id=user_id,
|
||||
team_id=team_id,
|
||||
team_alias=(
|
||||
team_object.team_alias
|
||||
if team_object is not None
|
||||
else None
|
||||
team_object.team_alias if team_object is not None else None
|
||||
),
|
||||
team_metadata=team_object.metadata
|
||||
if team_object is not None
|
||||
|
|
@ -709,12 +707,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
if isinstance(api_key, str):
|
||||
return UserAPIKeyAuth(
|
||||
api_key=api_key,
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
else:
|
||||
return UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
parent_otel_span=parent_otel_span,
|
||||
)
|
||||
elif api_key is None: # only require api key if master key is set
|
||||
|
|
@ -846,7 +844,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
)
|
||||
valid_token.parent_otel_span = parent_otel_span
|
||||
if _end_user_object is not None:
|
||||
valid_token.end_user_object_permission = _end_user_object.object_permission
|
||||
valid_token.end_user_object_permission = (
|
||||
_end_user_object.object_permission
|
||||
)
|
||||
|
||||
return valid_token
|
||||
|
||||
|
|
@ -954,7 +954,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
if isinstance(
|
||||
api_key, str
|
||||
): # if generated token, make sure it starts with sk-.
|
||||
_masked_key = "{}****{}".format(api_key[:4], api_key[-4:]) if len(api_key) > 8 else "****"
|
||||
_masked_key = (
|
||||
"{}****{}".format(api_key[:4], api_key[-4:])
|
||||
if len(api_key) > 8
|
||||
else "****"
|
||||
)
|
||||
assert api_key.startswith(
|
||||
"sk-"
|
||||
), "LiteLLM Virtual Key expected. Received={}, expected to start with 'sk-'.".format(
|
||||
|
|
@ -1304,9 +1308,9 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
|
|||
|
||||
if _end_user_object is not None:
|
||||
valid_token_dict.update(end_user_params)
|
||||
valid_token_dict["end_user_object_permission"] = (
|
||||
_end_user_object.object_permission
|
||||
)
|
||||
valid_token_dict[
|
||||
"end_user_object_permission"
|
||||
] = _end_user_object.object_permission
|
||||
|
||||
# check if token is from litellm-ui, litellm ui makes keys to allow users to login with sso. These keys can only be used for LiteLLM UI functions
|
||||
# sso/login, ui/login, /key functions and /user functions
|
||||
|
|
|
|||
|
|
@ -1,6 +1,10 @@
|
|||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from fastapi.responses import ORJSONResponse, StreamingResponse
|
||||
|
||||
import litellm
|
||||
from litellm._uuid import uuid
|
||||
from litellm.proxy._types import *
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
|
|
@ -17,7 +21,8 @@ router = APIRouter(
|
|||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
@router.post(
|
||||
"/models/{model_name:path}:generateContent", dependencies=[Depends(user_api_key_auth)]
|
||||
"/models/{model_name:path}:generateContent",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def google_generate_content(
|
||||
request: Request,
|
||||
|
|
@ -36,12 +41,12 @@ async def google_generate_content(
|
|||
data = await _read_request_body(request=request)
|
||||
if "model" not in data:
|
||||
data["model"] = model_name
|
||||
|
||||
|
||||
# Extract generationConfig and pass it as config parameter
|
||||
generation_config = data.pop("generationConfig", None)
|
||||
if generation_config:
|
||||
data["config"] = generation_config
|
||||
|
||||
|
||||
# Add user authentication metadata for cost tracking
|
||||
data = await add_litellm_data_to_request(
|
||||
data=data,
|
||||
|
|
@ -51,7 +56,19 @@ async def google_generate_content(
|
|||
general_settings=general_settings,
|
||||
version=version,
|
||||
)
|
||||
|
||||
|
||||
# Create logging object with full request metadata so callbacks (e.g. S3) get user/trace_id
|
||||
data["litellm_call_id"] = request.headers.get(
|
||||
"x-litellm-call-id", str(uuid.uuid4())
|
||||
)
|
||||
logging_obj, data = litellm.utils.function_setup(
|
||||
original_function="agenerate_content",
|
||||
rules_obj=litellm.utils.Rules(),
|
||||
start_time=datetime.now(),
|
||||
**data,
|
||||
)
|
||||
data["litellm_logging_obj"] = logging_obj
|
||||
|
||||
# call router
|
||||
if llm_router is None:
|
||||
raise HTTPException(status_code=500, detail="Router not initialized")
|
||||
|
|
@ -103,6 +120,18 @@ async def google_stream_generate_content(
|
|||
version=version,
|
||||
)
|
||||
|
||||
# Create logging object with full request metadata so streaming END callbacks (e.g. S3) get user/trace_id
|
||||
data["litellm_call_id"] = request.headers.get(
|
||||
"x-litellm-call-id", str(uuid.uuid4())
|
||||
)
|
||||
logging_obj, data = litellm.utils.function_setup(
|
||||
original_function="agenerate_content_stream",
|
||||
rules_obj=litellm.utils.Rules(),
|
||||
start_time=datetime.now(),
|
||||
**data,
|
||||
)
|
||||
data["litellm_logging_obj"] = logging_obj
|
||||
|
||||
# call router
|
||||
if llm_router is None:
|
||||
raise HTTPException(status_code=500, detail="Router not initialized")
|
||||
|
|
@ -247,11 +276,11 @@ async def create_interaction(
|
|||
)
|
||||
|
||||
data = await _read_request_body(request=request)
|
||||
|
||||
|
||||
# Default to gemini provider for interactions
|
||||
if "custom_llm_provider" not in data:
|
||||
data["custom_llm_provider"] = "gemini"
|
||||
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
|
|
@ -301,7 +330,7 @@ async def get_interaction(
|
|||
):
|
||||
"""
|
||||
Get an interaction by ID.
|
||||
|
||||
|
||||
Per OpenAPI spec: GET /{api_version}/interactions/{interaction_id}
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
|
|
@ -319,7 +348,7 @@ async def get_interaction(
|
|||
)
|
||||
|
||||
data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"}
|
||||
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
|
|
@ -369,7 +398,7 @@ async def delete_interaction(
|
|||
):
|
||||
"""
|
||||
Delete an interaction by ID.
|
||||
|
||||
|
||||
Per OpenAPI spec: DELETE /{api_version}/interactions/{interaction_id}
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
|
|
@ -387,7 +416,7 @@ async def delete_interaction(
|
|||
)
|
||||
|
||||
data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"}
|
||||
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
|
|
@ -437,7 +466,7 @@ async def cancel_interaction(
|
|||
):
|
||||
"""
|
||||
Cancel an interaction by ID.
|
||||
|
||||
|
||||
Per OpenAPI spec: POST /{api_version}/interactions/{interaction_id}:cancel
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
|
|
@ -455,7 +484,7 @@ async def cancel_interaction(
|
|||
)
|
||||
|
||||
data = {"interaction_id": interaction_id, "custom_llm_provider": "gemini"}
|
||||
|
||||
|
||||
processor = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
return await processor.base_process_llm_request(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
CRUD ENDPOINTS FOR GUARDRAILS
|
||||
"""
|
||||
|
||||
import concurrent.futures
|
||||
import inspect
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union, cast
|
||||
|
||||
|
|
@ -11,9 +12,16 @@ from pydantic import BaseModel
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.guardrails.guardrail_registry import GuardrailRegistry
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.code_validator import (
|
||||
CustomCodeValidationError,
|
||||
validate_custom_code,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import (
|
||||
get_custom_code_primitives,
|
||||
)
|
||||
from litellm.proxy.guardrails.usage_endpoints import router as guardrails_usage_router
|
||||
from litellm.types.guardrails import (
|
||||
PII_ENTITY_CATEGORIES_MAP,
|
||||
|
|
@ -243,9 +251,11 @@ class CreateGuardrailRequest(BaseModel):
|
|||
@router.post(
|
||||
"/guardrails",
|
||||
tags=["Guardrails"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def create_guardrail(request: CreateGuardrailRequest):
|
||||
async def create_guardrail(
|
||||
request: CreateGuardrailRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a new guardrail
|
||||
|
||||
|
|
@ -296,6 +306,12 @@ async def create_guardrail(request: CreateGuardrailRequest):
|
|||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required to manage guardrails",
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
|
|
@ -332,9 +348,12 @@ class UpdateGuardrailRequest(BaseModel):
|
|||
@router.put(
|
||||
"/guardrails/{guardrail_id}",
|
||||
tags=["Guardrails"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def update_guardrail(guardrail_id: str, request: UpdateGuardrailRequest):
|
||||
async def update_guardrail(
|
||||
guardrail_id: str,
|
||||
request: UpdateGuardrailRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Update an existing guardrail
|
||||
|
||||
|
|
@ -385,6 +404,12 @@ async def update_guardrail(guardrail_id: str, request: UpdateGuardrailRequest):
|
|||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required to manage guardrails",
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
|
|
@ -429,9 +454,11 @@ async def update_guardrail(guardrail_id: str, request: UpdateGuardrailRequest):
|
|||
@router.delete(
|
||||
"/guardrails/{guardrail_id}",
|
||||
tags=["Guardrails"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_guardrail(guardrail_id: str):
|
||||
async def delete_guardrail(
|
||||
guardrail_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Delete a guardrail
|
||||
|
||||
|
|
@ -453,6 +480,12 @@ async def delete_guardrail(guardrail_id: str):
|
|||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required to manage guardrails",
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
|
|
@ -495,9 +528,12 @@ async def delete_guardrail(guardrail_id: str):
|
|||
@router.patch(
|
||||
"/guardrails/{guardrail_id}",
|
||||
tags=["Guardrails"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def patch_guardrail(guardrail_id: str, request: PatchGuardrailRequest):
|
||||
async def patch_guardrail(
|
||||
guardrail_id: str,
|
||||
request: PatchGuardrailRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Partially update an existing guardrail
|
||||
|
||||
|
|
@ -546,6 +582,12 @@ async def patch_guardrail(guardrail_id: str, request: PatchGuardrailRequest):
|
|||
from litellm.proxy.guardrails.guardrail_registry import IN_MEMORY_GUARDRAIL_HANDLER
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required to manage guardrails",
|
||||
)
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Prisma client not initialized")
|
||||
|
||||
|
|
@ -1302,9 +1344,9 @@ async def get_provider_specific_params():
|
|||
lakera_v2_fields = _get_fields_from_model(LakeraV2GuardrailConfigModel)
|
||||
tool_permission_fields = _get_fields_from_model(ToolPermissionGuardrailConfigModel)
|
||||
|
||||
tool_permission_fields["ui_friendly_name"] = (
|
||||
ToolPermissionGuardrailConfigModel.ui_friendly_name()
|
||||
)
|
||||
tool_permission_fields[
|
||||
"ui_friendly_name"
|
||||
] = ToolPermissionGuardrailConfigModel.ui_friendly_name()
|
||||
|
||||
# Return the provider-specific parameters
|
||||
provider_params = {
|
||||
|
|
@ -1364,10 +1406,12 @@ class TestCustomCodeGuardrailResponse(BaseModel):
|
|||
@router.post(
|
||||
"/guardrails/test_custom_code",
|
||||
tags=["Guardrails"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=TestCustomCodeGuardrailResponse,
|
||||
)
|
||||
async def test_custom_code_guardrail(request: TestCustomCodeGuardrailRequest):
|
||||
async def test_custom_code_guardrail(
|
||||
request: TestCustomCodeGuardrailRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Test custom code guardrail logic without creating a guardrail.
|
||||
|
||||
|
|
@ -1440,63 +1484,26 @@ async def test_custom_code_guardrail(request: TestCustomCodeGuardrailRequest):
|
|||
}
|
||||
```
|
||||
"""
|
||||
import concurrent.futures
|
||||
import re
|
||||
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import (
|
||||
get_custom_code_primitives,
|
||||
)
|
||||
|
||||
# Security validation patterns
|
||||
FORBIDDEN_PATTERNS = [
|
||||
# Import statements
|
||||
(r"\bimport\s+", "import statements are not allowed"),
|
||||
(r"\bfrom\s+\w+\s+import\b", "from...import statements are not allowed"),
|
||||
(r"__import__\s*\(", "__import__() is not allowed"),
|
||||
# Dangerous builtins
|
||||
(r"\bexec\s*\(", "exec() is not allowed"),
|
||||
(r"\beval\s*\(", "eval() is not allowed"),
|
||||
(r"\bcompile\s*\(", "compile() is not allowed"),
|
||||
(r"\bopen\s*\(", "open() is not allowed"),
|
||||
(r"\bgetattr\s*\(", "getattr() is not allowed"),
|
||||
(r"\bsetattr\s*\(", "setattr() is not allowed"),
|
||||
(r"\bdelattr\s*\(", "delattr() is not allowed"),
|
||||
(r"\bglobals\s*\(", "globals() is not allowed"),
|
||||
(r"\blocals\s*\(", "locals() is not allowed"),
|
||||
(r"\bvars\s*\(", "vars() is not allowed"),
|
||||
(r"\bdir\s*\(", "dir() is not allowed"),
|
||||
(r"\bbreakpoint\s*\(", "breakpoint() is not allowed"),
|
||||
(r"\binput\s*\(", "input() is not allowed"),
|
||||
# Dangerous dunder access
|
||||
(r"__builtins__", "__builtins__ access is not allowed"),
|
||||
(r"__globals__", "__globals__ access is not allowed"),
|
||||
(r"__code__", "__code__ access is not allowed"),
|
||||
(r"__subclasses__", "__subclasses__ access is not allowed"),
|
||||
(r"__bases__", "__bases__ access is not allowed"),
|
||||
(r"__mro__", "__mro__ access is not allowed"),
|
||||
(r"__class__", "__class__ access is not allowed"),
|
||||
(r"__dict__", "__dict__ access is not allowed"),
|
||||
(r"__getattribute__", "__getattribute__ access is not allowed"),
|
||||
(r"__reduce__", "__reduce__ access is not allowed"),
|
||||
(r"__reduce_ex__", "__reduce_ex__ access is not allowed"),
|
||||
# OS/system access
|
||||
(r"\bos\.", "os module access is not allowed"),
|
||||
(r"\bsys\.", "sys module access is not allowed"),
|
||||
(r"\bsubprocess\.", "subprocess module access is not allowed"),
|
||||
]
|
||||
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail="Admin access required to test custom code guardrails",
|
||||
)
|
||||
|
||||
EXECUTION_TIMEOUT_SECONDS = 5
|
||||
|
||||
try:
|
||||
# Step 0: Security validation - check for forbidden patterns
|
||||
code = request.custom_code
|
||||
for pattern, error_msg in FORBIDDEN_PATTERNS:
|
||||
if re.search(pattern, code):
|
||||
return TestCustomCodeGuardrailResponse(
|
||||
success=False,
|
||||
error=f"Security violation: {error_msg}",
|
||||
error_type="compilation",
|
||||
)
|
||||
|
||||
try:
|
||||
validate_custom_code(request.custom_code)
|
||||
except CustomCodeValidationError as e:
|
||||
return TestCustomCodeGuardrailResponse(
|
||||
success=False,
|
||||
error=str(e),
|
||||
error_type="compilation",
|
||||
)
|
||||
|
||||
# Step 1: Compile the custom code with restricted environment
|
||||
exec_globals = get_custom_code_primitives().copy()
|
||||
|
|
@ -1612,10 +1619,10 @@ async def apply_guardrail(
|
|||
from litellm.proxy.utils import handle_exception_on_proxy
|
||||
|
||||
try:
|
||||
active_guardrail: Optional[CustomGuardrail] = (
|
||||
GUARDRAIL_REGISTRY.get_initialized_guardrail_callback(
|
||||
guardrail_name=request.guardrail_name
|
||||
)
|
||||
active_guardrail: Optional[
|
||||
CustomGuardrail
|
||||
] = GUARDRAIL_REGISTRY.get_initialized_guardrail_callback(
|
||||
guardrail_name=request.guardrail_name
|
||||
)
|
||||
if active_guardrail is None:
|
||||
raise HTTPException(
|
||||
|
|
|
|||
|
|
@ -0,0 +1,63 @@
|
|||
import re
|
||||
from typing import List, Tuple
|
||||
|
||||
# Security validation patterns
|
||||
FORBIDDEN_PATTERNS: List[Tuple[str, str]] = [
|
||||
# Import statements
|
||||
(r"\bimport\s+", "import statements are not allowed"),
|
||||
(r"\bfrom\s+\w+\s+import\b", "from...import statements are not allowed"),
|
||||
(r"__import__\s*\(", "__import__() is not allowed"),
|
||||
# Dangerous builtins
|
||||
(r"\bexec\s*\(", "exec() is not allowed"),
|
||||
(r"\beval\s*\(", "eval() is not allowed"),
|
||||
(r"\bcompile\s*\(", "compile() is not allowed"),
|
||||
(r"\bopen\s*\(", "open() is not allowed"),
|
||||
(r"\bgetattr\s*\(", "getattr() is not allowed"),
|
||||
(r"\bsetattr\s*\(", "setattr() is not allowed"),
|
||||
(r"\bdelattr\s*\(", "delattr() is not allowed"),
|
||||
(r"\bglobals\s*\(", "globals() is not allowed"),
|
||||
(r"\blocals\s*\(", "locals() is not allowed"),
|
||||
(r"\bvars\s*\(", "vars() is not allowed"),
|
||||
(r"\bdir\s*\(", "dir() is not allowed"),
|
||||
(r"\bbreakpoint\s*\(", "breakpoint() is not allowed"),
|
||||
(r"\binput\s*\(", "input() is not allowed"),
|
||||
# Dangerous dunder access
|
||||
(r"__builtins__", "__builtins__ access is not allowed"),
|
||||
(r"__globals__", "__globals__ access is not allowed"),
|
||||
(r"__code__", "__code__ access is not allowed"),
|
||||
(r"__subclasses__", "__subclasses__ access is not allowed"),
|
||||
(r"__bases__", "__bases__ access is not allowed"),
|
||||
(r"__mro__", "__mro__ access is not allowed"),
|
||||
(r"__class__", "__class__ access is not allowed"),
|
||||
(r"__dict__", "__dict__ access is not allowed"),
|
||||
(r"__getattribute__", "__getattribute__ access is not allowed"),
|
||||
(r"__reduce__", "__reduce__ access is not allowed"),
|
||||
(r"__reduce_ex__", "__reduce_ex__ access is not allowed"),
|
||||
# OS/system access
|
||||
(r"\bos\.", "os module access is not allowed"),
|
||||
(r"\bsys\.", "sys module access is not allowed"),
|
||||
(r"\bsubprocess\.", "subprocess module access is not allowed"),
|
||||
(r"\bshutil\.", "shutil module access is not allowed"),
|
||||
(r"\bctypes\.", "ctypes module access is not allowed"),
|
||||
(r"\bsocket\.", "socket module access is not allowed"),
|
||||
(r"\bpickle\.", "pickle module access is not allowed"),
|
||||
]
|
||||
|
||||
|
||||
class CustomCodeValidationError(Exception):
|
||||
"""Raised when custom code fails security validation."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
def validate_custom_code(code: str) -> None:
|
||||
"""
|
||||
Validate custom code against forbidden patterns.
|
||||
|
||||
Raises CustomCodeValidationError if any forbidden pattern is found.
|
||||
"""
|
||||
if not code:
|
||||
return
|
||||
for pattern, error_msg in FORBIDDEN_PATTERNS:
|
||||
if re.search(pattern, code):
|
||||
raise CustomCodeValidationError(f"Security violation: {error_msg}")
|
||||
|
|
@ -41,18 +41,19 @@ from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast
|
|||
from fastapi import HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import (CustomGuardrail,
|
||||
log_guardrail_information)
|
||||
from litellm.integrations.custom_guardrail import (
|
||||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import \
|
||||
GuardrailConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
from .code_validator import CustomCodeValidationError, validate_custom_code
|
||||
from .primitives import get_custom_code_primitives
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import \
|
||||
Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
|
||||
class CustomCodeGuardrailError(Exception):
|
||||
|
|
@ -143,6 +144,33 @@ class CustomCodeGuardrail(CustomGuardrail):
|
|||
"""Returns the config model for the UI."""
|
||||
return CustomCodeGuardrailConfigModel
|
||||
|
||||
def _do_compile(self) -> None:
|
||||
"""Internal compilation method without lock. Expected to run inside _compile_lock."""
|
||||
# Create a restricted execution environment
|
||||
# Only include our safe primitives
|
||||
exec_globals = get_custom_code_primitives().copy()
|
||||
|
||||
# CRITICAL: Restrict __builtins__ to prevent sandbox escape
|
||||
exec_globals["__builtins__"] = {}
|
||||
|
||||
# Execute the user code in the restricted environment
|
||||
exec(compile(self.custom_code, "<guardrail>", "exec"), exec_globals)
|
||||
|
||||
# Extract the apply_guardrail function
|
||||
if "apply_guardrail" not in exec_globals:
|
||||
raise CustomCodeCompilationError(
|
||||
"Custom code must define an 'apply_guardrail' function. "
|
||||
"Expected signature: apply_guardrail(inputs, request_data, input_type)"
|
||||
)
|
||||
|
||||
apply_fn = exec_globals["apply_guardrail"]
|
||||
if not callable(apply_fn):
|
||||
raise CustomCodeCompilationError(
|
||||
"'apply_guardrail' must be a callable function"
|
||||
)
|
||||
|
||||
self._compiled_function = apply_fn
|
||||
|
||||
def _compile_custom_code(self) -> None:
|
||||
"""
|
||||
Compile the custom code and extract the apply_guardrail function.
|
||||
|
|
@ -154,27 +182,14 @@ class CustomCodeGuardrail(CustomGuardrail):
|
|||
return
|
||||
|
||||
try:
|
||||
# Create a restricted execution environment
|
||||
# Only include our safe primitives
|
||||
exec_globals = get_custom_code_primitives().copy()
|
||||
# Step 1: Security validation — forbidden pattern check
|
||||
try:
|
||||
validate_custom_code(self.custom_code)
|
||||
except CustomCodeValidationError as e:
|
||||
raise CustomCodeCompilationError(str(e)) from e
|
||||
|
||||
# Execute the user code in the restricted environment
|
||||
exec(compile(self.custom_code, "<guardrail>", "exec"), exec_globals)
|
||||
|
||||
# Extract the apply_guardrail function
|
||||
if "apply_guardrail" not in exec_globals:
|
||||
raise CustomCodeCompilationError(
|
||||
"Custom code must define an 'apply_guardrail' function. "
|
||||
"Expected signature: apply_guardrail(inputs, request_data, input_type)"
|
||||
)
|
||||
|
||||
apply_fn = exec_globals["apply_guardrail"]
|
||||
if not callable(apply_fn):
|
||||
raise CustomCodeCompilationError(
|
||||
"'apply_guardrail' must be a callable function"
|
||||
)
|
||||
|
||||
self._compiled_function = apply_fn
|
||||
# Step 2: Compile logic
|
||||
self._do_compile()
|
||||
verbose_proxy_logger.debug(
|
||||
f"Custom code guardrail '{self.guardrail_name}' compiled successfully"
|
||||
)
|
||||
|
|
@ -390,6 +405,12 @@ class CustomCodeGuardrail(CustomGuardrail):
|
|||
Raises:
|
||||
CustomCodeCompilationError: If the new code fails to compile
|
||||
"""
|
||||
# Validate BEFORE acquiring lock / resetting state
|
||||
try:
|
||||
validate_custom_code(new_code)
|
||||
except CustomCodeValidationError as e:
|
||||
raise CustomCodeCompilationError(str(e)) from e
|
||||
|
||||
with self._compile_lock:
|
||||
# Reset state
|
||||
old_function = self._compiled_function
|
||||
|
|
@ -399,12 +420,24 @@ class CustomCodeGuardrail(CustomGuardrail):
|
|||
|
||||
try:
|
||||
self.custom_code = new_code
|
||||
self._compile_custom_code()
|
||||
self._do_compile()
|
||||
verbose_proxy_logger.info(
|
||||
f"Custom code guardrail '{self.guardrail_name}': Code updated successfully"
|
||||
)
|
||||
except SyntaxError as e:
|
||||
# Rollback on failure
|
||||
self.custom_code = old_code
|
||||
self._compiled_function = old_function
|
||||
self._compile_error = f"Syntax error in custom code: {e}"
|
||||
raise CustomCodeCompilationError(self._compile_error) from e
|
||||
except CustomCodeCompilationError:
|
||||
# Rollback on failure
|
||||
self.custom_code = old_code
|
||||
self._compiled_function = old_function
|
||||
raise
|
||||
except Exception as e:
|
||||
# Rollback on failure
|
||||
self.custom_code = old_code
|
||||
self._compiled_function = old_function
|
||||
self._compile_error = f"Failed to compile custom code: {e}"
|
||||
raise CustomCodeCompilationError(self._compile_error) from e
|
||||
|
|
|
|||
|
|
@ -322,12 +322,6 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
analyze_payload,
|
||||
)
|
||||
|
||||
async with session.post(analyze_url, json=analyze_payload) as response:
|
||||
analyze_results = await response.json()
|
||||
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
|
||||
|
||||
# Handle error responses from Presidio (e.g., {'error': 'No text provided'})
|
||||
# Presidio may return a dict instead of a list when errors occur
|
||||
def _fail_on_invalid_response(
|
||||
reason: str,
|
||||
) -> List[PresidioAnalyzeResponseItem]:
|
||||
|
|
@ -347,6 +341,36 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
)
|
||||
return []
|
||||
|
||||
async with session.post(
|
||||
analyze_url,
|
||||
json=analyze_payload,
|
||||
headers={"Accept": "application/json"},
|
||||
) as response:
|
||||
# Validate HTTP status
|
||||
if response.status >= 400:
|
||||
error_body = await response.text()
|
||||
return _fail_on_invalid_response(
|
||||
f"HTTP {response.status} from Presidio analyzer: {error_body[:200]}"
|
||||
)
|
||||
|
||||
# Validate Content-Type is JSON
|
||||
content_type = getattr(
|
||||
response,
|
||||
"content_type",
|
||||
response.headers.get("Content-Type", ""),
|
||||
)
|
||||
if "application/json" not in content_type:
|
||||
error_body = await response.text()
|
||||
return _fail_on_invalid_response(
|
||||
f"expected application/json Content-Type but received '{content_type}'; body: '{error_body[:200]}'"
|
||||
)
|
||||
|
||||
analyze_results = await response.json()
|
||||
verbose_proxy_logger.debug("analyze_results: %s", analyze_results)
|
||||
|
||||
# Handle error responses from Presidio (e.g., {'error': 'No text provided'})
|
||||
# Presidio may return a dict instead of a list when errors occur
|
||||
|
||||
if isinstance(analyze_results, dict):
|
||||
if "error" in analyze_results:
|
||||
return _fail_on_invalid_response(
|
||||
|
|
@ -423,8 +447,29 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
}
|
||||
|
||||
async with session.post(
|
||||
anonymize_url, json=anonymize_payload
|
||||
anonymize_url,
|
||||
json=anonymize_payload,
|
||||
headers={"Accept": "application/json"},
|
||||
) as response:
|
||||
# Validate HTTP status
|
||||
if response.status >= 400:
|
||||
error_body = await response.text()
|
||||
raise Exception(
|
||||
f"Presidio anonymizer returned HTTP {response.status}: {error_body[:200]}"
|
||||
)
|
||||
|
||||
# Validate Content-Type is JSON
|
||||
content_type = getattr(
|
||||
response,
|
||||
"content_type",
|
||||
response.headers.get("Content-Type", ""),
|
||||
)
|
||||
if "application/json" not in content_type:
|
||||
error_body = await response.text()
|
||||
raise Exception(
|
||||
f"Presidio anonymizer returned non-JSON Content-Type '{content_type}'; body: '{error_body[:200]}'"
|
||||
)
|
||||
|
||||
redacted_text = await response.json()
|
||||
|
||||
new_text = text
|
||||
|
|
@ -456,7 +501,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
|
|||
except Exception as e:
|
||||
# Sanitize exception to avoid leaking the original text (which may
|
||||
# contain API keys or other secrets) in error responses.
|
||||
if "Invalid anonymizer response" in str(e):
|
||||
error_str = str(e)
|
||||
if (
|
||||
"Invalid anonymizer response" in error_str
|
||||
or "Presidio anonymizer returned" in error_str
|
||||
):
|
||||
raise
|
||||
raise Exception(
|
||||
f"Presidio PII anonymization failed: {type(e).__name__}"
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
)
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_checks import log_db_metrics
|
||||
from litellm.proxy.auth.auth_checks import get_key_object, get_team_object, log_db_metrics
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.utils import ProxyUpdateSpend
|
||||
from litellm.types.utils import (
|
||||
|
|
@ -76,6 +76,10 @@ class _ProxyDBLogger(CustomLogger):
|
|||
traceback_str=traceback_str,
|
||||
)
|
||||
|
||||
_metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(
|
||||
metadata=_metadata,
|
||||
)
|
||||
|
||||
existing_metadata: dict = request_data.get("metadata", None) or {}
|
||||
existing_metadata.update(_metadata)
|
||||
|
||||
|
|
@ -255,6 +259,72 @@ class _ProxyDBLogger(CustomLogger):
|
|||
"Error in tracking cost callback - %s", str(e)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict:
|
||||
"""
|
||||
Enriches failure spend log metadata by looking up the key object (and team object)
|
||||
from cache/DB when key fields are missing.
|
||||
|
||||
This handles two scenarios:
|
||||
1. Auth errors (401): UserAPIKeyAuth is created with only api_key set, all other
|
||||
fields are null. We look up the full key object to fill in alias, user_id,
|
||||
team_id, etc.
|
||||
2. Post-auth failures (provider errors, rate limits): key fields are populated
|
||||
but team_alias is missing because LiteLLM_VerificationTokenView SQL view
|
||||
doesn't include it. We look up the team object to fill in team_alias.
|
||||
"""
|
||||
api_key_hash = metadata.get("user_api_key")
|
||||
if not api_key_hash:
|
||||
return metadata
|
||||
|
||||
from litellm.proxy.proxy_server import (
|
||||
prisma_client,
|
||||
proxy_logging_obj,
|
||||
user_api_key_cache,
|
||||
)
|
||||
|
||||
# Step 1: If key fields are missing, look up the full key object
|
||||
if metadata.get("user_api_key_alias") is None:
|
||||
try:
|
||||
key_obj = await get_key_object(
|
||||
hashed_token=api_key_hash,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if metadata.get("user_api_key_alias") is None:
|
||||
metadata["user_api_key_alias"] = key_obj.key_alias
|
||||
if metadata.get("user_api_key_user_id") is None:
|
||||
metadata["user_api_key_user_id"] = key_obj.user_id
|
||||
if metadata.get("user_api_key_team_id") is None:
|
||||
metadata["user_api_key_team_id"] = key_obj.team_id
|
||||
if metadata.get("user_api_key_org_id") is None:
|
||||
metadata["user_api_key_org_id"] = key_obj.org_id
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to enrich failure metadata with key info for api_key=%s",
|
||||
api_key_hash,
|
||||
)
|
||||
|
||||
# Step 2: If team_id is known but team_alias is missing, look up the team object
|
||||
team_id = metadata.get("user_api_key_team_id")
|
||||
if team_id and metadata.get("user_api_key_team_alias") is None:
|
||||
try:
|
||||
team_obj = await get_team_object(
|
||||
team_id=team_id,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
if team_obj.team_alias is not None:
|
||||
metadata["user_api_key_team_alias"] = team_obj.team_alias
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Failed to enrich failure metadata with team_alias for team_id=%s",
|
||||
team_id,
|
||||
)
|
||||
return metadata
|
||||
|
||||
@staticmethod
|
||||
def _should_track_errors_in_db():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -10,10 +10,15 @@ import litellm
|
|||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._service_logger import ServiceLogging
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.proxy._types import (AddTeamCallback, CommonProxyErrors,
|
||||
LitellmDataForBackendLLMCall,
|
||||
LitellmUserRoles, SpecialHeaders,
|
||||
TeamCallbackMetadata, UserAPIKeyAuth)
|
||||
from litellm.proxy._types import (
|
||||
AddTeamCallback,
|
||||
CommonProxyErrors,
|
||||
LitellmDataForBackendLLMCall,
|
||||
LitellmUserRoles,
|
||||
SpecialHeaders,
|
||||
TeamCallbackMetadata,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers
|
||||
|
||||
# Cache special headers as a frozenset for O(1) lookup performance
|
||||
|
|
@ -23,9 +28,12 @@ _SPECIAL_HEADERS_CACHE = frozenset(
|
|||
from litellm.router import Router
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_API_HEADERS
|
||||
from litellm.types.services import ServiceTypes
|
||||
from litellm.types.utils import (LlmProviders, ProviderSpecificHeader,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
SupportedCacheControls)
|
||||
from litellm.types.utils import (
|
||||
LlmProviders,
|
||||
ProviderSpecificHeader,
|
||||
StandardLoggingUserAPIKeyMetadata,
|
||||
SupportedCacheControls,
|
||||
)
|
||||
|
||||
service_logger_obj = ServiceLogging() # used for tracking latency on OTEL
|
||||
|
||||
|
|
@ -228,7 +236,9 @@ def _get_dynamic_logging_metadata(
|
|||
|
||||
|
||||
def clean_headers(
|
||||
headers: Headers, litellm_key_header_name: Optional[str] = None
|
||||
headers: Headers,
|
||||
litellm_key_header_name: Optional[str] = None,
|
||||
forward_llm_provider_auth_headers: bool = False,
|
||||
) -> dict:
|
||||
"""
|
||||
Removes litellm api key from headers
|
||||
|
|
@ -238,15 +248,18 @@ def clean_headers(
|
|||
clean_headers = {}
|
||||
litellm_key_lower = (
|
||||
litellm_key_header_name.lower() if litellm_key_header_name is not None else None
|
||||
)
|
||||
|
||||
)
|
||||
for header, value in headers.items():
|
||||
header_lower = header.lower()
|
||||
# Preserve Authorization header if it contains Anthropic OAuth token (sk-ant-oat*)
|
||||
# This allows OAuth tokens to be forwarded to Anthropic-compatible providers
|
||||
# via add_provider_specific_headers_to_request()
|
||||
|
||||
if header_lower == "authorization" and is_anthropic_oauth_key(value):
|
||||
clean_headers[header] = value
|
||||
elif forward_llm_provider_auth_headers and header_lower in _SPECIAL_HEADERS_CACHE:
|
||||
if litellm_key_lower and header_lower == litellm_key_lower:
|
||||
continue
|
||||
if header_lower == "authorization":
|
||||
continue
|
||||
clean_headers[header] = value
|
||||
# Check if header should be excluded: either in special headers cache or matches custom litellm key
|
||||
elif header_lower not in _SPECIAL_HEADERS_CACHE and (
|
||||
litellm_key_lower is None or header_lower != litellm_key_lower
|
||||
|
|
@ -654,7 +667,8 @@ class LiteLLMProxyRequestSetup:
|
|||
return data
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ManagementEndpoint_MetadataFields,
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium)
|
||||
LiteLLM_ManagementEndpoint_MetadataFields_Premium,
|
||||
)
|
||||
|
||||
# ignore any special fields
|
||||
added_metadata = {}
|
||||
|
|
@ -826,6 +840,11 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
|
||||
|
||||
_raw_headers: Dict[str, str] = _safe_get_request_headers(request)
|
||||
|
||||
forward_llm_auth = False
|
||||
if general_settings:
|
||||
forward_llm_auth = general_settings.get("forward_llm_provider_auth_headers", False)
|
||||
|
||||
_headers: Dict[str, str] = clean_headers(
|
||||
request.headers,
|
||||
litellm_key_header_name=(
|
||||
|
|
@ -833,7 +852,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915
|
|||
if general_settings is not None
|
||||
else None
|
||||
),
|
||||
forward_llm_provider_auth_headers=forward_llm_auth,
|
||||
)
|
||||
verbose_proxy_logger.debug(f"Request Headers: {_headers}")
|
||||
verbose_proxy_logger.debug(f"Raw Headers: {_raw_headers}")
|
||||
|
||||
##########################################################
|
||||
# Init - Proxy Server Request
|
||||
|
|
@ -1479,8 +1501,7 @@ async def move_guardrails_to_metadata(
|
|||
|
||||
# Only check policy engine if no local config (avoid import + registry lookup)
|
||||
if not (has_key_config or has_team_config or has_request_config):
|
||||
from litellm.proxy.policy_engine.policy_registry import \
|
||||
get_policy_registry
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
|
||||
if not get_policy_registry().is_initialized():
|
||||
# Nothing configured anywhere - clean up request body fields and return
|
||||
|
|
@ -1544,16 +1565,14 @@ async def move_guardrails_to_metadata(
|
|||
|
||||
def _is_policy_version_id(s: str) -> bool:
|
||||
"""Return True if string is a policy version ID (starts with policy_<uuid> prefix)."""
|
||||
from litellm.proxy.policy_engine.policy_registry import \
|
||||
POLICY_VERSION_ID_PREFIX
|
||||
from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX
|
||||
|
||||
return isinstance(s, str) and s.startswith(POLICY_VERSION_ID_PREFIX)
|
||||
|
||||
|
||||
def _extract_policy_id(s: str) -> Optional[str]:
|
||||
"""Extract raw UUID from policy_<uuid> string, or None if not a valid version ID."""
|
||||
from litellm.proxy.policy_engine.policy_registry import \
|
||||
POLICY_VERSION_ID_PREFIX
|
||||
from litellm.proxy.policy_engine.policy_registry import POLICY_VERSION_ID_PREFIX
|
||||
|
||||
if not _is_policy_version_id(s):
|
||||
return None
|
||||
|
|
@ -1574,9 +1593,10 @@ def _match_and_track_policies(
|
|||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_policy_sources_to_metadata, add_policy_to_applied_policies_header)
|
||||
from litellm.proxy.policy_engine.attachment_registry import \
|
||||
get_attachment_registry
|
||||
add_policy_sources_to_metadata,
|
||||
add_policy_to_applied_policies_header,
|
||||
)
|
||||
from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry
|
||||
from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher
|
||||
|
||||
# Get matching policies via attachments (with match reasons for attribution)
|
||||
|
|
@ -1721,8 +1741,7 @@ async def add_guardrails_from_policy_engine(
|
|||
user_api_key_dict: The user's API key authentication info
|
||||
"""
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.common_utils.http_parsing_utils import \
|
||||
get_tags_from_request_body
|
||||
from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body
|
||||
from litellm.proxy.policy_engine.policy_registry import get_policy_registry
|
||||
from litellm.types.proxy.policy_engine import PolicyMatchContext
|
||||
|
||||
|
|
|
|||
|
|
@ -1216,9 +1216,7 @@ try:
|
|||
|
||||
# Case 2: Runtime UI exists and is ready
|
||||
if has_content and is_pre_restructured:
|
||||
verbose_proxy_logger.info(
|
||||
f"Using pre-restructured UI at {runtime_ui_path}"
|
||||
)
|
||||
verbose_proxy_logger.info(f"Using pre-restructured UI at {runtime_ui_path}")
|
||||
ui_path = runtime_ui_path
|
||||
|
||||
# Case 3: Runtime UI exists but needs restructuring
|
||||
|
|
@ -2994,6 +2992,10 @@ class ProxyConfig:
|
|||
|
||||
if master_key is not None and isinstance(master_key, str):
|
||||
litellm_master_key_hash = hash_token(master_key)
|
||||
else:
|
||||
verbose_proxy_logger.critical(
|
||||
"LITELLM_MASTER_KEY is not set! All requests will be treated as INTERNAL_USER with no admin access. Set LITELLM_MASTER_KEY for production use."
|
||||
)
|
||||
### USER API KEY CACHE IN-MEMORY TTL ###
|
||||
user_api_key_cache_ttl = general_settings.get(
|
||||
"user_api_key_cache_ttl", None
|
||||
|
|
@ -3796,6 +3798,7 @@ class ProxyConfig:
|
|||
parsed = value
|
||||
elif isinstance(value, str):
|
||||
import json
|
||||
|
||||
try:
|
||||
parsed = yaml.safe_load(value)
|
||||
except (yaml.YAMLError, json.JSONDecodeError):
|
||||
|
|
@ -4381,10 +4384,12 @@ class ProxyConfig:
|
|||
|
||||
if self._should_load_db_object(object_type="model_cost_map"):
|
||||
await self._check_and_reload_model_cost_map(prisma_client=prisma_client)
|
||||
|
||||
|
||||
if self._should_load_db_object(object_type="anthropic_beta_headers"):
|
||||
await self._check_and_reload_anthropic_beta_headers(prisma_client=prisma_client)
|
||||
|
||||
await self._check_and_reload_anthropic_beta_headers(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
if self._should_load_db_object(object_type="sso_settings"):
|
||||
await self._init_sso_settings_in_db(prisma_client=prisma_client)
|
||||
if self._should_load_db_object(object_type="cache_settings"):
|
||||
|
|
@ -4614,7 +4619,9 @@ class ProxyConfig:
|
|||
f"Error in _check_and_reload_model_cost_map: {str(e)}"
|
||||
)
|
||||
|
||||
async def _check_and_reload_anthropic_beta_headers(self, prisma_client: PrismaClient):
|
||||
async def _check_and_reload_anthropic_beta_headers(
|
||||
self, prisma_client: PrismaClient
|
||||
):
|
||||
"""
|
||||
Check if anthropic beta headers config needs to be reloaded based on database configuration.
|
||||
This function runs every 10 seconds as part of _init_non_llm_objects_in_db.
|
||||
|
|
@ -4705,7 +4712,11 @@ class ProxyConfig:
|
|||
)
|
||||
|
||||
# Count providers in config
|
||||
provider_count = sum(1 for k in new_config.keys() if k != "provider_aliases" and k != "description")
|
||||
provider_count = sum(
|
||||
1
|
||||
for k in new_config.keys()
|
||||
if k != "provider_aliases" and k != "description"
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
f"Anthropic beta headers config reloaded successfully. Providers: {provider_count}"
|
||||
)
|
||||
|
|
@ -5687,8 +5698,7 @@ class ProxyStartupEvent:
|
|||
):
|
||||
_db_val = _db_gs_record.param_value.get("store_model_in_db")
|
||||
if _db_val is True or (
|
||||
isinstance(_db_val, str)
|
||||
and _db_val.lower() == "true"
|
||||
isinstance(_db_val, str) and _db_val.lower() == "true"
|
||||
):
|
||||
store_model_in_db = True
|
||||
verbose_proxy_logger.info(
|
||||
|
|
@ -6155,6 +6165,7 @@ class ProxyStartupEvent:
|
|||
"Pyroscope profiling will not run. Install with: pip install pyroscope-io"
|
||||
)
|
||||
|
||||
|
||||
#### API ENDPOINTS ####
|
||||
@router.get(
|
||||
"/v1/models", dependencies=[Depends(user_api_key_auth)], tags=["model management"]
|
||||
|
|
@ -10993,18 +11004,14 @@ async def get_favicon():
|
|||
from fastapi.responses import Response
|
||||
|
||||
current_dir = os.path.dirname(os.path.abspath(__file__))
|
||||
default_favicon = os.path.join(
|
||||
current_dir, "_experimental", "out", "favicon.ico"
|
||||
)
|
||||
default_favicon = os.path.join(current_dir, "_experimental", "out", "favicon.ico")
|
||||
|
||||
favicon_url = os.getenv("LITELLM_FAVICON_URL", "")
|
||||
|
||||
if not favicon_url:
|
||||
if os.path.exists(default_favicon):
|
||||
return FileResponse(default_favicon, media_type="image/x-icon")
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Default favicon not found"
|
||||
)
|
||||
raise HTTPException(status_code=404, detail="Default favicon not found")
|
||||
|
||||
if favicon_url.startswith(("http://", "https://")):
|
||||
try:
|
||||
|
|
@ -11019,9 +11026,7 @@ async def get_favicon():
|
|||
)
|
||||
response = await async_client.get(favicon_url)
|
||||
if response.status_code == 200:
|
||||
content_type = response.headers.get(
|
||||
"content-type", "image/x-icon"
|
||||
)
|
||||
content_type = response.headers.get("content-type", "image/x-icon")
|
||||
return Response(
|
||||
content=response.content,
|
||||
media_type=content_type,
|
||||
|
|
@ -11033,12 +11038,8 @@ async def get_favicon():
|
|||
response.status_code,
|
||||
)
|
||||
if os.path.exists(default_favicon):
|
||||
return FileResponse(
|
||||
default_favicon, media_type="image/x-icon"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Favicon not found"
|
||||
)
|
||||
return FileResponse(default_favicon, media_type="image/x-icon")
|
||||
raise HTTPException(status_code=404, detail="Favicon not found")
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -11046,20 +11047,14 @@ async def get_favicon():
|
|||
"Error downloading favicon from %s: %s", favicon_url, e
|
||||
)
|
||||
if os.path.exists(default_favicon):
|
||||
return FileResponse(
|
||||
default_favicon, media_type="image/x-icon"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Favicon not found"
|
||||
)
|
||||
return FileResponse(default_favicon, media_type="image/x-icon")
|
||||
raise HTTPException(status_code=404, detail="Favicon not found")
|
||||
else:
|
||||
if os.path.exists(favicon_url):
|
||||
return FileResponse(favicon_url, media_type="image/x-icon")
|
||||
if os.path.exists(default_favicon):
|
||||
return FileResponse(default_favicon, media_type="image/x-icon")
|
||||
raise HTTPException(
|
||||
status_code=404, detail="Favicon not found"
|
||||
)
|
||||
raise HTTPException(status_code=404, detail="Favicon not found")
|
||||
|
||||
|
||||
#### INVITATION MANAGEMENT ####
|
||||
|
|
@ -12545,7 +12540,9 @@ async def reload_anthropic_beta_headers(
|
|||
},
|
||||
)
|
||||
|
||||
provider_count = sum(1 for k in new_config.keys() if k not in ["provider_aliases", "description"])
|
||||
provider_count = sum(
|
||||
1 for k in new_config.keys() if k not in ["provider_aliases", "description"]
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
f"Anthropic beta headers config reloaded successfully in current pod. Providers: {provider_count}"
|
||||
)
|
||||
|
|
@ -12557,7 +12554,9 @@ async def reload_anthropic_beta_headers(
|
|||
"timestamp": current_time.isoformat(),
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Failed to reload anthropic beta headers: {str(e)}")
|
||||
verbose_proxy_logger.exception(
|
||||
f"Failed to reload anthropic beta headers: {str(e)}"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Failed to reload anthropic beta headers: {str(e)}"
|
||||
)
|
||||
|
|
@ -12679,7 +12678,8 @@ async def cancel_anthropic_beta_headers_reload(
|
|||
f"Failed to cancel anthropic beta headers reload: {str(e)}"
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=500, detail=f"Failed to cancel anthropic beta headers reload: {str(e)}"
|
||||
status_code=500,
|
||||
detail=f"Failed to cancel anthropic beta headers reload: {str(e)}",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -12726,7 +12726,9 @@ async def get_anthropic_beta_headers_reload_status(
|
|||
)
|
||||
|
||||
if config_record is None or config_record.param_value is None:
|
||||
verbose_proxy_logger.info("No anthropic beta headers reload configuration found")
|
||||
verbose_proxy_logger.info(
|
||||
"No anthropic beta headers reload configuration found"
|
||||
)
|
||||
return {
|
||||
"scheduled": False,
|
||||
"interval_hours": None,
|
||||
|
|
@ -12752,7 +12754,9 @@ async def get_anthropic_beta_headers_reload_status(
|
|||
# Use pod's in-memory last reload time
|
||||
if last_anthropic_beta_headers_reload is not None:
|
||||
try:
|
||||
last_reload_time = datetime.fromisoformat(last_anthropic_beta_headers_reload)
|
||||
last_reload_time = datetime.fromisoformat(
|
||||
last_anthropic_beta_headers_reload
|
||||
)
|
||||
time_since_last_reload = current_time - last_reload_time
|
||||
hours_since_last_reload = time_since_last_reload.total_seconds() / 3600
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from typing_extensions import Any, List, Optional, TypedDict
|
|||
|
||||
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
|
||||
|
||||
Phase = Optional[Literal["commentary", "final_answer"]] # TODO: Once openai sdk has updated, we can remove this and use the openai sdk type
|
||||
Phase = Optional[Literal["commentary", "final_answer"]]
|
||||
|
||||
class GenericResponseOutputItemContentAnnotation(BaseLiteLLMOpenAIResponseObject):
|
||||
"""Annotation for content in a message"""
|
||||
|
|
|
|||
|
|
@ -1,7 +1,17 @@
|
|||
import json
|
||||
import time
|
||||
from enum import Enum
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Mapping, Optional, Union
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Mapping,
|
||||
Optional,
|
||||
Union,
|
||||
get_args,
|
||||
)
|
||||
|
||||
from openai._models import BaseModel as OpenAIObject
|
||||
from openai.types.audio.transcription_create_params import (
|
||||
|
|
@ -3186,6 +3196,12 @@ OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS: set[str] = {
|
|||
LlmProviders.HOSTED_VLLM.value,
|
||||
}
|
||||
|
||||
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "vertex_ai"]
|
||||
|
||||
LIST_BATCHES_SUPPORTED_PROVIDERS: frozenset[str] = frozenset(
|
||||
get_args(ListBatchesSupportedProvider)
|
||||
)
|
||||
|
||||
|
||||
class SearchProviders(str, Enum):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -4644,11 +4644,12 @@ def add_provider_specific_params_to_optional_params(
|
|||
)
|
||||
is False
|
||||
):
|
||||
extra_body = passed_params.pop("extra_body", {})
|
||||
extra_body = passed_params.pop("extra_body", None) or {}
|
||||
for k in passed_params.keys():
|
||||
if k not in openai_params and passed_params[k] is not None:
|
||||
extra_body[k] = passed_params[k]
|
||||
optional_params.setdefault("extra_body", {})
|
||||
if not isinstance(optional_params.get("extra_body"), dict):
|
||||
optional_params["extra_body"] = {}
|
||||
initial_extra_body = {
|
||||
**optional_params["extra_body"],
|
||||
**extra_body,
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ cryptography==46.0.5 #GHSA-r6ph-v2qm-q3c2
|
|||
|
||||
anyio==4.8.0 # openai + http req.
|
||||
httpx==0.28.1
|
||||
openai==2.9.0 # openai req.
|
||||
openai==2.24.0 # openai req.
|
||||
fastapi==0.120.1 # server dep
|
||||
starlette==0.49.1 # starlette fastapi dep
|
||||
backoff==2.2.1 # server dep
|
||||
|
|
|
|||
94
tests/litellm/proxy/guardrails/test_custom_code_security.py
Normal file
94
tests/litellm/proxy/guardrails/test_custom_code_security.py
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
import pytest
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.code_validator import (
|
||||
validate_custom_code,
|
||||
CustomCodeValidationError,
|
||||
)
|
||||
from litellm.proxy.guardrails.guardrail_hooks.custom_code.custom_code_guardrail import (
|
||||
CustomCodeGuardrail,
|
||||
)
|
||||
|
||||
# Phase 4.1: Test forbidden pattern validation
|
||||
|
||||
|
||||
def test_validate_custom_code_import_os():
|
||||
code = "import os\ndef apply_guardrail(inputs, req, ty):\n return allow()"
|
||||
with pytest.raises(CustomCodeValidationError, match="import statements are not"):
|
||||
validate_custom_code(code)
|
||||
|
||||
|
||||
def test_validate_custom_code_from_subprocess():
|
||||
code = (
|
||||
"from subprocess import call\ndef apply_guardrail(i, r, t):\n return allow()"
|
||||
)
|
||||
with pytest.raises(
|
||||
CustomCodeValidationError, match="import statements are not allowed"
|
||||
):
|
||||
validate_custom_code(code)
|
||||
|
||||
|
||||
def test_validate_custom_code_exec():
|
||||
code = "def apply_guardrail(i, r, t):\n exec('print(1)')\n return allow()"
|
||||
with pytest.raises(CustomCodeValidationError, match=r"exec\(\) is not allowed"):
|
||||
validate_custom_code(code)
|
||||
|
||||
|
||||
def test_validate_custom_code_builtins():
|
||||
code = "def apply_guardrail(i, r, t):\n print(__builtins__)\n return allow()"
|
||||
with pytest.raises(
|
||||
CustomCodeValidationError, match="__builtins__ access is not allowed"
|
||||
):
|
||||
validate_custom_code(code)
|
||||
|
||||
|
||||
def test_validate_custom_code_subclasses():
|
||||
code = "def apply_guardrail(i, r, t):\n print(''.__class__.__mro__[1].__subclasses__())\n return allow()"
|
||||
with pytest.raises(
|
||||
CustomCodeValidationError, match="__subclasses__ access is not allowed"
|
||||
):
|
||||
validate_custom_code(code)
|
||||
|
||||
|
||||
def test_validate_custom_code_clean():
|
||||
code = (
|
||||
"def apply_guardrail(inputs, request_data, input_type):\n return allow()\n"
|
||||
)
|
||||
# Should not raise any exception
|
||||
validate_custom_code(code)
|
||||
|
||||
|
||||
# Phase 4.2: Test __builtins__ restriction in execution
|
||||
|
||||
|
||||
def test_custom_code_compile_valid():
|
||||
code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()"
|
||||
guardrail = CustomCodeGuardrail(custom_code=code, guardrail_name="test")
|
||||
# if it doesn't fail, we successfully compiled
|
||||
assert guardrail._compiled_function is not None
|
||||
|
||||
|
||||
def test_custom_code_override_builtins():
|
||||
# Verify that even if pattern validation is bypassed, __builtins__ = {} blocks dangerous builtins.
|
||||
# We test this by compiling safe code and verifying builtins are not accessible in the sandbox.
|
||||
code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()"
|
||||
guardrail = CustomCodeGuardrail(custom_code=code, guardrail_name="test")
|
||||
# The compiled function's globals should have empty __builtins__
|
||||
fn_globals = guardrail._compiled_function.__globals__
|
||||
assert fn_globals.get("__builtins__") == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_custom_code_guardrail_apply():
|
||||
code = "def apply_guardrail(inputs, request_data, input_type):\n return allow()"
|
||||
guardrail = CustomCodeGuardrail(custom_code=code, guardrail_name="test")
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=GenericGuardrailAPIInputs(texts=["test"]),
|
||||
request_data={},
|
||||
input_type="request",
|
||||
)
|
||||
assert result["texts"][0] == "test"
|
||||
|
||||
|
||||
# The RBAC endpoint tests are harder to write right here, but the core security
|
||||
# validations are fully covered by the simple tests above.
|
||||
|
|
@ -334,6 +334,81 @@ def test_chat_completion_forward_headers(
|
|||
pytest.fail(f"LiteLLM Proxy test failed. Exception - {str(e)}")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("forward_llm_auth_headers", [True, False])
|
||||
@mock_patch_acompletion()
|
||||
def test_chat_completion_forward_llm_provider_auth_headers(
|
||||
mock_acompletion, client_no_auth, forward_llm_auth_headers
|
||||
):
|
||||
"""
|
||||
Test that LLM provider auth headers (x-api-key, x-goog-api-key) are forwarded
|
||||
when forward_llm_provider_auth_headers=True.
|
||||
|
||||
This allows clients to send their own LLM provider API keys through the proxy.
|
||||
"""
|
||||
try:
|
||||
# Configure general settings
|
||||
gs = getattr(litellm.proxy.proxy_server, "general_settings")
|
||||
gs["forward_client_headers_to_llm_api"] = True
|
||||
gs["forward_llm_provider_auth_headers"] = forward_llm_auth_headers
|
||||
setattr(litellm.proxy.proxy_server, "general_settings", gs)
|
||||
|
||||
# Test data
|
||||
test_data = {
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"},
|
||||
],
|
||||
"max_tokens": 10,
|
||||
}
|
||||
|
||||
# Headers including LLM provider auth
|
||||
request_headers = {
|
||||
"Authorization": "Bearer sk-proxy-auth-123", # Proxy auth (should be stripped)
|
||||
"x-api-key": "sk-ant-api03-test-anthropic-key", # Anthropic API key
|
||||
"x-goog-api-key": "google-api-key-123", # Google API key
|
||||
"X-Custom-Header": "custom-value", # Custom header (should be forwarded)
|
||||
}
|
||||
|
||||
# Make request
|
||||
response = client_no_auth.post(
|
||||
"/v1/chat/completions", json=test_data, headers=request_headers
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
# Check forwarded headers
|
||||
forwarded_headers = mock_acompletion.call_args.kwargs.get("headers", {})
|
||||
|
||||
if forward_llm_auth_headers:
|
||||
# LLM provider auth headers should be forwarded
|
||||
assert "x-api-key" in forwarded_headers
|
||||
assert forwarded_headers["x-api-key"] == "sk-ant-api03-test-anthropic-key"
|
||||
assert "x-goog-api-key" in forwarded_headers
|
||||
assert forwarded_headers["x-goog-api-key"] == "google-api-key-123"
|
||||
else:
|
||||
# LLM provider auth headers should be stripped
|
||||
assert "x-api-key" not in forwarded_headers
|
||||
assert "x-goog-api-key" not in forwarded_headers
|
||||
|
||||
# Custom headers should always be forwarded (when forward_client_headers_to_llm_api=True)
|
||||
assert "x-custom-header" in forwarded_headers
|
||||
assert forwarded_headers["x-custom-header"] == "custom-value"
|
||||
|
||||
# Proxy Authorization should never be forwarded
|
||||
assert "authorization" not in forwarded_headers
|
||||
|
||||
print(f"✓ Test passed with forward_llm_provider_auth_headers={forward_llm_auth_headers}")
|
||||
print(f" Forwarded headers: {list(forwarded_headers.keys())}")
|
||||
|
||||
except Exception as e:
|
||||
pytest.fail(f"Test failed with forward_llm_auth_headers={forward_llm_auth_headers}: {str(e)}")
|
||||
finally:
|
||||
# Clean up
|
||||
gs = getattr(litellm.proxy.proxy_server, "general_settings")
|
||||
gs.pop("forward_llm_provider_auth_headers", None)
|
||||
setattr(litellm.proxy.proxy_server, "general_settings", gs)
|
||||
|
||||
|
||||
@mock_patch_acompletion()
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_disable_guardrails(mock_acompletion, client_no_auth):
|
||||
|
|
|
|||
|
|
@ -147,7 +147,8 @@ class TestResponseCompliance:
|
|||
"""Verify status enum values match spec."""
|
||||
schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"]
|
||||
status_prop = schema["properties"]["status"]
|
||||
expected_statuses = ["UNSPECIFIED", "IN_PROGRESS", "REQUIRES_ACTION", "COMPLETED", "FAILED", "CANCELLED", "INCOMPLETE"]
|
||||
# Google Interactions API uses lowercase status values (updated Feb 2026)
|
||||
expected_statuses = ["in_progress", "requires_action", "completed", "failed", "cancelled", "incomplete"]
|
||||
assert status_prop["enum"] == expected_statuses
|
||||
print(f"✓ Status enum values: {expected_statuses}")
|
||||
|
||||
|
|
|
|||
|
|
@ -363,6 +363,87 @@ class TestProxyOAuthHeaderForwarding:
|
|||
assert "authorization" not in cleaned
|
||||
assert cleaned["content-type"] == "application/json"
|
||||
|
||||
def test_clean_headers_forwards_anthropic_api_key_when_enabled(self):
|
||||
"""clean_headers should preserve x-api-key when forward_llm_provider_auth_headers=True."""
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy.litellm_pre_call_utils import clean_headers
|
||||
|
||||
raw_headers = Headers(
|
||||
raw=[
|
||||
(b"authorization", b"Bearer sk-proxy-auth"),
|
||||
(b"x-api-key", b"sk-ant-api03-test-key"),
|
||||
(b"content-type", b"application/json"),
|
||||
]
|
||||
)
|
||||
cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=True)
|
||||
|
||||
# x-api-key should be preserved when flag is True
|
||||
assert "x-api-key" in cleaned
|
||||
assert cleaned["x-api-key"] == "sk-ant-api03-test-key"
|
||||
# Authorization (proxy auth) should still be stripped
|
||||
assert "authorization" not in cleaned
|
||||
assert cleaned["content-type"] == "application/json"
|
||||
|
||||
def test_clean_headers_strips_anthropic_api_key_when_disabled(self):
|
||||
"""clean_headers should strip x-api-key when forward_llm_provider_auth_headers=False (default)."""
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy.litellm_pre_call_utils import clean_headers
|
||||
|
||||
raw_headers = Headers(
|
||||
raw=[
|
||||
(b"x-api-key", b"sk-ant-api03-test-key"),
|
||||
(b"content-type", b"application/json"),
|
||||
]
|
||||
)
|
||||
cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=False)
|
||||
|
||||
# x-api-key should be stripped by default
|
||||
assert "x-api-key" not in cleaned
|
||||
assert cleaned["content-type"] == "application/json"
|
||||
|
||||
def test_clean_headers_forwards_google_api_key_when_enabled(self):
|
||||
"""clean_headers should preserve x-goog-api-key when forward_llm_provider_auth_headers=True."""
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy.litellm_pre_call_utils import clean_headers
|
||||
|
||||
raw_headers = Headers(
|
||||
raw=[
|
||||
(b"x-goog-api-key", b"google-api-key-123"),
|
||||
(b"content-type", b"application/json"),
|
||||
]
|
||||
)
|
||||
cleaned = clean_headers(raw_headers, forward_llm_provider_auth_headers=True)
|
||||
|
||||
assert "x-goog-api-key" in cleaned
|
||||
assert cleaned["x-goog-api-key"] == "google-api-key-123"
|
||||
assert cleaned["content-type"] == "application/json"
|
||||
|
||||
def test_clean_headers_preserves_oauth_regardless_of_forward_flag(self):
|
||||
"""clean_headers should always preserve OAuth tokens regardless of forward_llm_provider_auth_headers."""
|
||||
from starlette.datastructures import Headers
|
||||
|
||||
from litellm.proxy.litellm_pre_call_utils import clean_headers
|
||||
|
||||
raw_headers = Headers(
|
||||
raw=[
|
||||
(b"authorization", f"Bearer {FAKE_OAUTH_TOKEN}".encode()),
|
||||
(b"content-type", b"application/json"),
|
||||
]
|
||||
)
|
||||
|
||||
# Should preserve OAuth even with flag=False
|
||||
cleaned_without_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=False)
|
||||
assert "authorization" in cleaned_without_flag
|
||||
assert cleaned_without_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
|
||||
# Should also preserve OAuth with flag=True
|
||||
cleaned_with_flag = clean_headers(raw_headers, forward_llm_provider_auth_headers=True)
|
||||
assert "authorization" in cleaned_with_flag
|
||||
assert cleaned_with_flag["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}"
|
||||
|
||||
def test_add_provider_specific_headers_forwards_oauth(self):
|
||||
"""add_provider_specific_headers_to_request should forward OAuth Authorization
|
||||
as a ProviderSpecificHeader scoped to Anthropic-compatible providers."""
|
||||
|
|
|
|||
|
|
@ -0,0 +1,103 @@
|
|||
"""
|
||||
Tests for hosted_vllm responses API support.
|
||||
|
||||
Regression test for: https://github.com/BerriAI/litellm/issues
|
||||
Bug: client.responses.create() raised TypeError: 'NoneType' object is not a mapping
|
||||
when extra_body=None was passed through the responses→completion pipeline for
|
||||
hosted_vllm (and any OpenAI-compatible provider using add_provider_specific_params_to_optional_params).
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
def _make_mock_chat_completion_response(content: str = "Hello! I'm doing well.") -> dict:
|
||||
return {
|
||||
"id": "chatcmpl-test123",
|
||||
"object": "chat.completion",
|
||||
"created": 1234567890,
|
||||
"model": "Qwen/Qwen3-8B",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": content},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
||||
}
|
||||
|
||||
|
||||
def _make_mock_http_client(response_body: dict) -> MagicMock:
|
||||
mock_client = MagicMock()
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "application/json"}
|
||||
mock_response.json.return_value = response_body
|
||||
mock_response.text = json.dumps(response_body)
|
||||
mock_client.post.return_value = mock_response
|
||||
return mock_client
|
||||
|
||||
|
||||
def test_hosted_vllm_responses_create_with_string_input():
|
||||
"""
|
||||
Regression test: responses.create() with string input must not raise
|
||||
TypeError: 'NoneType' object is not a mapping.
|
||||
|
||||
Root cause: extra_body=None was passed explicitly through the
|
||||
responses→completion pipeline. In add_provider_specific_params_to_optional_params(),
|
||||
passed_params.pop("extra_body", {}) returned None (key existed with value None),
|
||||
and **None raised TypeError at dict unpacking.
|
||||
|
||||
Fix: normalize None to {} for both extra_body and optional_params["extra_body"].
|
||||
"""
|
||||
mock_client = _make_mock_http_client(
|
||||
_make_mock_chat_completion_response("I'm doing well, thanks!")
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.llm_http_handler._get_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
response = litellm.responses(
|
||||
model="hosted_vllm/Qwen/Qwen3-8B",
|
||||
input="Hello, how are you?",
|
||||
api_base="https://test-vllm.example.com/v1",
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
assert response is not None
|
||||
assert isinstance(response, ResponsesAPIResponse)
|
||||
assert len(response.output) > 0
|
||||
output_message = response.output[0]
|
||||
assert output_message.role == "assistant" # type: ignore[union-attr]
|
||||
assert len(output_message.content) > 0 # type: ignore[union-attr]
|
||||
assert "well" in output_message.content[0].text # type: ignore[union-attr]
|
||||
|
||||
|
||||
def test_hosted_vllm_responses_create_with_explicit_none_extra_body():
|
||||
"""
|
||||
Directly verify the fix in add_provider_specific_params_to_optional_params:
|
||||
extra_body=None must not crash when building optional_params.
|
||||
"""
|
||||
from litellm.utils import get_optional_params
|
||||
|
||||
# This should not raise TypeError: 'NoneType' object is not a mapping
|
||||
optional_params = get_optional_params(
|
||||
model="Qwen/Qwen3-8B",
|
||||
custom_llm_provider="hosted_vllm",
|
||||
extra_body=None,
|
||||
)
|
||||
|
||||
# extra_body=None should be normalized to an empty dict (or absent)
|
||||
assert optional_params.get("extra_body") is not None or "extra_body" not in optional_params
|
||||
|
|
@ -2,10 +2,9 @@
|
|||
"""
|
||||
Test to verify the Google GenAI proxy API endpoints
|
||||
"""
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -13,7 +12,6 @@ sys.path.insert(
|
|||
0, os.path.abspath("../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
def test_google_generate_content_endpoint():
|
||||
|
|
@ -401,3 +399,123 @@ def test_google_generate_content_with_image_config():
|
|||
assert "contents" in called_data
|
||||
assert len(called_data["contents"]) == 1
|
||||
assert called_data["contents"][0]["role"] == "user"
|
||||
|
||||
|
||||
def test_google_generate_content_metadata_and_trace_id_callbacks():
|
||||
"""Test that google_generate_content sets litellm_call_id and logging_obj for callbacks (e.g. S3, Langfuse)"""
|
||||
try:
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy.google_endpoints.endpoints import router as google_router
|
||||
except ImportError as e:
|
||||
pytest.skip(f"Skipping test due to missing dependency: {e}")
|
||||
|
||||
# Create a FastAPI app and include the router
|
||||
app = FastAPI()
|
||||
app.include_router(google_router)
|
||||
|
||||
# Create a test client
|
||||
client = TestClient(app)
|
||||
|
||||
# Mock all required proxy server dependencies
|
||||
with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {}
|
||||
), patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config, patch(
|
||||
"litellm.proxy.proxy_server.version", "1.0.0"
|
||||
), patch(
|
||||
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
|
||||
) as mock_add_data:
|
||||
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
|
||||
|
||||
# Mock add_litellm_data_to_request to return data with metadata
|
||||
async def mock_add_litellm_data(
|
||||
data, request, user_api_key_dict, proxy_config, general_settings, version
|
||||
):
|
||||
# Simulate adding user metadata
|
||||
data["litellm_metadata"] = {
|
||||
"user_api_key_user_id": "test-user-id",
|
||||
}
|
||||
return data
|
||||
|
||||
mock_add_data.side_effect = mock_add_litellm_data
|
||||
|
||||
# Send a request to the endpoint with x-litellm-call-id header
|
||||
test_call_id = "test-custom-call-id"
|
||||
response = client.post(
|
||||
"/v1beta/models/test-model:generateContent",
|
||||
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
|
||||
headers={
|
||||
"Authorization": "Bearer sk-test-key",
|
||||
"x-litellm-call-id": test_call_id,
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
mock_router.agenerate_content.assert_called_once()
|
||||
call_args = mock_router.agenerate_content.call_args
|
||||
called_data = call_args[1]
|
||||
|
||||
# Verify that the litellm_logging_obj got assigned in the final called_data to router
|
||||
assert "litellm_logging_obj" in called_data
|
||||
assert "litellm_call_id" in called_data
|
||||
assert called_data["litellm_call_id"] == test_call_id
|
||||
|
||||
|
||||
def test_google_stream_generate_content_metadata_and_trace_id_callbacks():
|
||||
"""Test that google_stream_generate_content sets litellm_call_id and logging_obj for callbacks"""
|
||||
try:
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy.google_endpoints.endpoints import router as google_router
|
||||
except ImportError as e:
|
||||
pytest.skip(f"Skipping test due to missing dependency: {e}")
|
||||
|
||||
app = FastAPI()
|
||||
app.include_router(google_router)
|
||||
client = TestClient(app)
|
||||
|
||||
mock_stream = AsyncMock()
|
||||
mock_stream.__aiter__ = lambda self: mock_stream
|
||||
mock_stream.__anext__.side_effect = StopAsyncIteration
|
||||
|
||||
with patch("litellm.proxy.proxy_server.llm_router") as mock_router, patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {}
|
||||
), patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config, patch(
|
||||
"litellm.proxy.proxy_server.version", "1.0.0"
|
||||
), patch(
|
||||
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
|
||||
) as mock_add_data:
|
||||
mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream)
|
||||
|
||||
async def mock_add_litellm_data(
|
||||
data, request, user_api_key_dict, proxy_config, general_settings, version
|
||||
):
|
||||
data["litellm_metadata"] = {
|
||||
"user_api_key_user_id": "test-user-id",
|
||||
}
|
||||
return data
|
||||
|
||||
mock_add_data.side_effect = mock_add_litellm_data
|
||||
|
||||
test_call_id = "test-custom-stream-call-id"
|
||||
response = client.post(
|
||||
"/v1beta/models/test-model:streamGenerateContent",
|
||||
json={"contents": [{"role": "user", "parts": [{"text": "Hello stream"}]}]},
|
||||
headers={
|
||||
"Authorization": "Bearer sk-test-key",
|
||||
"x-litellm-call-id": test_call_id,
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
|
||||
mock_router.agenerate_content_stream.assert_called_once()
|
||||
call_args = mock_router.agenerate_content_stream.call_args
|
||||
called_data = call_args[1]
|
||||
|
||||
assert "litellm_logging_obj" in called_data
|
||||
assert "litellm_call_id" in called_data
|
||||
assert called_data["litellm_call_id"] == test_call_id
|
||||
|
|
|
|||
|
|
@ -24,12 +24,29 @@ from litellm.types.guardrails import LitellmParams, PiiAction, PiiEntityType
|
|||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
|
||||
def _make_mock_session_iterator(json_response):
|
||||
def _make_mock_session_iterator(
|
||||
json_response, status=200, content_type="application/json", text_response=""
|
||||
):
|
||||
"""Create a mock _get_session_iterator that yields a session returning json_response."""
|
||||
|
||||
@asynccontextmanager
|
||||
async def mock_iterator():
|
||||
class MockResponse:
|
||||
def __init__(self):
|
||||
self.status = status
|
||||
self.content_type = content_type
|
||||
self.headers = {"Content-Type": content_type}
|
||||
|
||||
async def text(self):
|
||||
if text_response:
|
||||
return text_response
|
||||
import json
|
||||
|
||||
try:
|
||||
return json.dumps(json_response)
|
||||
except Exception:
|
||||
return str(json_response)
|
||||
|
||||
async def json(self):
|
||||
return json_response
|
||||
|
||||
|
|
@ -41,6 +58,7 @@ def _make_mock_session_iterator(json_response):
|
|||
|
||||
class MockSession:
|
||||
def post(self, *args, **kwargs):
|
||||
self.last_kwargs = kwargs
|
||||
return MockResponse()
|
||||
|
||||
async def __aenter__(self):
|
||||
|
|
@ -1444,3 +1462,149 @@ def test_deny_list_and_score_threshold_combined():
|
|||
# EMAIL_ADDRESS passes both filters
|
||||
assert len(filtered) == 1
|
||||
assert filtered[0]["entity_type"] == "EMAIL_ADDRESS"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_analyze_text_non_json_content_type_fail_closed():
|
||||
"""
|
||||
Test that analyze_text raises GuardrailRaisedException when Presidio health
|
||||
endpoint returns text/html and fail-closed is enabled.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
presidio_analyzer_api_base="http://test-analyzer/",
|
||||
presidio_anonymizer_api_base="http://test-anonymizer/",
|
||||
pii_entities_config={"PERSON": PiiAction.BLOCK},
|
||||
mock_testing=False,
|
||||
)
|
||||
|
||||
mock_iterator = _make_mock_session_iterator(
|
||||
json_response=None,
|
||||
status=200,
|
||||
content_type="text/html; charset=utf-8",
|
||||
text_response="Presidio Analyzer service is up.",
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
|
||||
with pytest.raises(GuardrailRaisedException) as exc_info:
|
||||
await guardrail.analyze_text(
|
||||
text="Hello world",
|
||||
presidio_config=None,
|
||||
request_data={},
|
||||
)
|
||||
assert "expected application/json Content-Type" in str(exc_info.value)
|
||||
assert "text/html" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_analyze_text_non_json_content_type_fail_open():
|
||||
"""
|
||||
Test that analyze_text returns empty list when Presidio returns text/html
|
||||
and fail-closed is NOT enabled.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
presidio_analyzer_api_base="http://test-analyzer/",
|
||||
presidio_anonymizer_api_base="http://test-anonymizer/",
|
||||
mock_testing=False,
|
||||
)
|
||||
|
||||
mock_iterator = _make_mock_session_iterator(
|
||||
json_response=None,
|
||||
status=200,
|
||||
content_type="text/html; charset=utf-8",
|
||||
text_response="Presidio Analyzer service is up.",
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
|
||||
results = await guardrail.analyze_text(
|
||||
text="Hello world",
|
||||
presidio_config=None,
|
||||
request_data={},
|
||||
)
|
||||
assert results == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_analyze_text_http_error_status():
|
||||
"""
|
||||
Test that analyze_text handles 5xx HTTP errors properly.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
presidio_analyzer_api_base="http://test-analyzer/",
|
||||
presidio_anonymizer_api_base="http://test-anonymizer/",
|
||||
pii_entities_config={"PERSON": PiiAction.BLOCK},
|
||||
mock_testing=False,
|
||||
)
|
||||
|
||||
mock_iterator = _make_mock_session_iterator(
|
||||
json_response=None,
|
||||
status=500,
|
||||
content_type="text/plain",
|
||||
text_response="Internal Server Error",
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
|
||||
with pytest.raises(GuardrailRaisedException) as exc_info:
|
||||
await guardrail.analyze_text(
|
||||
text="Hello world",
|
||||
presidio_config=None,
|
||||
request_data={},
|
||||
)
|
||||
assert "HTTP 500" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anonymize_text_non_json_content_type():
|
||||
"""
|
||||
Test that anonymize_text raises Exception for non-JSON responses.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
presidio_analyzer_api_base="http://test-analyzer/",
|
||||
presidio_anonymizer_api_base="http://test-anonymizer/",
|
||||
mock_testing=False,
|
||||
)
|
||||
|
||||
mock_iterator = _make_mock_session_iterator(
|
||||
json_response=None,
|
||||
status=200,
|
||||
content_type="text/html",
|
||||
text_response="Presidio Anonymizer service is up.",
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
|
||||
with pytest.raises(
|
||||
Exception, match="Presidio anonymizer returned non-JSON Content-Type"
|
||||
):
|
||||
await guardrail.anonymize_text(
|
||||
text="Hello world",
|
||||
analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}],
|
||||
output_parse_pii=False,
|
||||
masked_entity_count={},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anonymize_text_http_error_status():
|
||||
"""
|
||||
Test that anonymize_text raises Exception on HTTP error.
|
||||
"""
|
||||
guardrail = _OPTIONAL_PresidioPIIMasking(
|
||||
presidio_analyzer_api_base="http://test-analyzer/",
|
||||
presidio_anonymizer_api_base="http://test-anonymizer/",
|
||||
mock_testing=False,
|
||||
)
|
||||
|
||||
mock_iterator = _make_mock_session_iterator(
|
||||
json_response=None,
|
||||
status=502,
|
||||
content_type="text/plain",
|
||||
text_response="Bad Gateway",
|
||||
)
|
||||
|
||||
with patch.object(guardrail, "_get_session_iterator", mock_iterator):
|
||||
with pytest.raises(Exception, match="Presidio anonymizer returned HTTP 502"):
|
||||
await guardrail.anonymize_text(
|
||||
text="Hello world",
|
||||
analyze_results=[{"start": 0, "end": 5, "entity_type": "PERSON"}],
|
||||
output_parse_pii=False,
|
||||
masked_entity_count={},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -169,6 +169,223 @@ async def test_track_cost_callback_skips_when_no_standard_logging_object():
|
|||
mock_proxy_logging.failed_tracking_alert.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrich_failure_metadata_with_team_alias():
|
||||
"""
|
||||
When team_id is set but team_alias is missing (and key_alias is present),
|
||||
_enrich_failure_metadata_with_key_info should look up the team from cache
|
||||
and populate user_api_key_team_alias.
|
||||
"""
|
||||
mock_team_obj = MagicMock()
|
||||
mock_team_obj.team_alias = "my-team-alias"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_team_obj,
|
||||
):
|
||||
metadata = {
|
||||
"user_api_key": "hashed_key",
|
||||
"user_api_key_alias": "my-key-alias", # already set
|
||||
"user_api_key_team_id": "test_team_id",
|
||||
"user_api_key_team_alias": None,
|
||||
}
|
||||
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
|
||||
assert result["user_api_key_team_alias"] == "my-team-alias"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrich_failure_metadata_with_full_key_lookup():
|
||||
"""
|
||||
When all key fields are null (auth error 401 scenario), _enrich_failure_metadata_with_key_info
|
||||
should look up the key object from cache/DB and populate alias, user_id, team_id,
|
||||
then look up the team to get team_alias.
|
||||
"""
|
||||
mock_key_obj = MagicMock()
|
||||
mock_key_obj.key_alias = "fetched-key-alias"
|
||||
mock_key_obj.user_id = "fetched-user-id"
|
||||
mock_key_obj.team_id = "fetched-team-id"
|
||||
mock_key_obj.org_id = "fetched-org-id"
|
||||
|
||||
mock_team_obj = MagicMock()
|
||||
mock_team_obj.team_alias = "fetched-team-alias"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_key_obj,
|
||||
), patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_team_obj,
|
||||
):
|
||||
metadata = {
|
||||
"user_api_key": "hashed_key",
|
||||
"user_api_key_alias": None, # all null - simulates auth error path
|
||||
"user_api_key_user_id": None,
|
||||
"user_api_key_team_id": None,
|
||||
"user_api_key_team_alias": None,
|
||||
"user_api_key_org_id": None,
|
||||
}
|
||||
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
|
||||
assert result["user_api_key_alias"] == "fetched-key-alias"
|
||||
assert result["user_api_key_user_id"] == "fetched-user-id"
|
||||
assert result["user_api_key_team_id"] == "fetched-team-id"
|
||||
assert result["user_api_key_org_id"] == "fetched-org-id"
|
||||
assert result["user_api_key_team_alias"] == "fetched-team-alias"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrich_failure_metadata_skips_when_team_alias_present():
|
||||
"""
|
||||
When team_alias is already populated, _enrich_failure_metadata_with_key_info
|
||||
should not perform a team cache lookup.
|
||||
"""
|
||||
with patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key, patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_team:
|
||||
metadata = {
|
||||
"user_api_key": "hashed_key",
|
||||
"user_api_key_alias": "existing-alias",
|
||||
"user_api_key_team_id": "test_team_id",
|
||||
"user_api_key_team_alias": "already-set",
|
||||
}
|
||||
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
|
||||
assert result["user_api_key_team_alias"] == "already-set"
|
||||
mock_get_key.assert_not_called()
|
||||
mock_get_team.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enrich_failure_metadata_skips_when_no_api_key():
|
||||
"""
|
||||
When api_key hash is absent, _enrich_failure_metadata_with_key_info should
|
||||
not perform any lookups.
|
||||
"""
|
||||
with patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_get_key:
|
||||
metadata = {
|
||||
"user_api_key": None,
|
||||
"user_api_key_alias": None,
|
||||
"user_api_key_team_id": None,
|
||||
"user_api_key_team_alias": None,
|
||||
}
|
||||
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
|
||||
mock_get_key.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_failure_hook_enriches_auth_error_metadata():
|
||||
"""
|
||||
Simulates a 401 ProxyException (e.g. can_key_call_model). In this case
|
||||
UserAPIKeyAuth is created with only api_key set. The failure hook should
|
||||
look up the key and team from cache/DB to populate all missing fields.
|
||||
"""
|
||||
logger = _ProxyDBLogger()
|
||||
|
||||
# This is what auth_exception_handler creates for 401 errors
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="hashed_key",
|
||||
# key_alias, user_id, team_id, team_alias are all None
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"model": "claude-haiku-4-5",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"metadata": {},
|
||||
"litellm_params": {},
|
||||
}
|
||||
|
||||
mock_key_obj = MagicMock()
|
||||
mock_key_obj.key_alias = "my-key-alias"
|
||||
mock_key_obj.user_id = "my-user-id"
|
||||
mock_key_obj.team_id = "my-team-id"
|
||||
mock_key_obj.org_id = None
|
||||
|
||||
mock_team_obj = MagicMock()
|
||||
mock_team_obj.team_alias = "my-team-alias"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_update_database, patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_key_obj,
|
||||
), patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_team_obj,
|
||||
):
|
||||
await logger.async_post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=Exception("401 - model not allowed"),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
mock_update_database.assert_called_once()
|
||||
call_args = mock_update_database.call_args[1]
|
||||
metadata = call_args["kwargs"]["litellm_params"]["metadata"]
|
||||
assert metadata["user_api_key_alias"] == "my-key-alias"
|
||||
assert metadata["user_api_key_user_id"] == "my-user-id"
|
||||
assert metadata["user_api_key_team_id"] == "my-team-id"
|
||||
assert metadata["user_api_key_team_alias"] == "my-team-alias"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_post_call_failure_hook_enriches_missing_team_alias():
|
||||
"""
|
||||
When user_api_key_dict has a team_id but no team_alias, async_post_call_failure_hook
|
||||
should look up the team from cache and populate user_api_key_team_alias in the
|
||||
spend log metadata written to the DB.
|
||||
"""
|
||||
logger = _ProxyDBLogger()
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test_api_key",
|
||||
key_alias="test_alias",
|
||||
user_id="test_user_id",
|
||||
team_id="test_team_id",
|
||||
team_alias=None, # Missing - simulates regular key auth where SQL view omits team_alias
|
||||
)
|
||||
|
||||
request_data = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"metadata": {},
|
||||
"litellm_params": {},
|
||||
}
|
||||
|
||||
mock_team_obj = MagicMock()
|
||||
mock_team_obj.team_alias = "enriched-team-alias"
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_update_database, patch(
|
||||
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_team_obj,
|
||||
):
|
||||
await logger.async_post_call_failure_hook(
|
||||
request_data=request_data,
|
||||
original_exception=Exception("Provider rate limit"),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
mock_update_database.assert_called_once()
|
||||
call_args = mock_update_database.call_args[1]
|
||||
metadata = call_args["kwargs"]["litellm_params"]["metadata"]
|
||||
assert metadata["user_api_key_team_alias"] == "enriched-team-alias"
|
||||
assert metadata["user_api_key_team_id"] == "test_team_id"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model_value", [None, ""])
|
||||
async def test_track_cost_callback_skips_for_falsy_model_and_no_slo(model_value):
|
||||
|
|
|
|||
93
ui/litellm-dashboard/package-lock.json
generated
93
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -90,6 +90,7 @@
|
|||
"version": "5.2.0",
|
||||
"resolved": "https://registry.npmjs.org/@alloc/quick-lru/-/quick-lru-5.2.0.tgz",
|
||||
"integrity": "sha512-UrcABB+4bUrFABwbluTIBErXwvbsU/V7TZWfmbgJfbkwiBuziS9gxdODUyuiecfdGQ85jglMW6juS3+z5TsKLw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=10"
|
||||
|
|
@ -1771,6 +1772,7 @@
|
|||
"version": "0.3.13",
|
||||
"resolved": "https://registry.npmjs.org/@jridgewell/gen-mapping/-/gen-mapping-0.3.13.tgz",
|
||||
"integrity": "sha512-2kkt/7niJ6MgEPxF0bYdQ6etZaA+fQvDcLKckhy1yIQOzaoKjBBjSj63/aLVjYE3qhRt5dvM+uUyfCg6UKCBbA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@jridgewell/sourcemap-codec": "^1.5.0",
|
||||
|
|
@ -1781,6 +1783,7 @@
|
|||
"version": "3.1.2",
|
||||
"resolved": "https://registry.npmjs.org/@jridgewell/resolve-uri/-/resolve-uri-3.1.2.tgz",
|
||||
"integrity": "sha512-bRISgCIjP20/tbWSPWMEi54QVPRZExkuD9lJL+UIxUKtwVJA8wW1Trb1jMs1RFXo1CBTNZ/5hpC9QvmKWdopKw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=6.0.0"
|
||||
|
|
@ -1790,12 +1793,14 @@
|
|||
"version": "1.5.5",
|
||||
"resolved": "https://registry.npmjs.org/@jridgewell/sourcemap-codec/-/sourcemap-codec-1.5.5.tgz",
|
||||
"integrity": "sha512-cYQ9310grqxueWbl+WuIUIaiUaDcj7WOq5fVhEljNVgRfOUhY9fy2zTvfoqWsnebh8Sl70VScFbICvJnLKB0Og==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@jridgewell/trace-mapping": {
|
||||
"version": "0.3.31",
|
||||
"resolved": "https://registry.npmjs.org/@jridgewell/trace-mapping/-/trace-mapping-0.3.31.tgz",
|
||||
"integrity": "sha512-zzNR+SdQSDJzc8joaeP8QQoCQr8NuYx2dIIytl1QeBEZHJ9uW6hebsrYgbz8hJwUQao3TWCMtmfV8Nu1twOLAw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@jridgewell/resolve-uri": "^3.1.0",
|
||||
|
|
@ -1973,6 +1978,7 @@
|
|||
"version": "2.1.5",
|
||||
"resolved": "https://registry.npmjs.org/@nodelib/fs.scandir/-/fs.scandir-2.1.5.tgz",
|
||||
"integrity": "sha512-vq24Bq3ym5HEQm2NKCr3yXDwjc7vTsEThRDnkp2DK9p1uqLR+DHurm/NOTo0KG7HYHU7eppKZj3MyqYuMBf62g==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@nodelib/fs.stat": "2.0.5",
|
||||
|
|
@ -1986,6 +1992,7 @@
|
|||
"version": "2.0.5",
|
||||
"resolved": "https://registry.npmjs.org/@nodelib/fs.stat/-/fs.stat-2.0.5.tgz",
|
||||
"integrity": "sha512-RkhPPp2zrqDAQA/2jNhnztcPAlv64XdhIp7a7454A5ovI7Bukxgt7MX7udwAu3zg1DcpPU0rz3VV1SeaqvY4+A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 8"
|
||||
|
|
@ -1995,6 +2002,7 @@
|
|||
"version": "1.2.8",
|
||||
"resolved": "https://registry.npmjs.org/@nodelib/fs.walk/-/fs.walk-1.2.8.tgz",
|
||||
"integrity": "sha512-oGB+UxlgWcgQkgwo8GcEGwemoTFt3FIO9ababBmaGwXIoBKZ+GTy0pP185beGg7Llih/NSHSV2XAs1lnznocSg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@nodelib/fs.scandir": "2.1.5",
|
||||
|
|
@ -2318,7 +2326,7 @@
|
|||
"version": "1.58.1",
|
||||
"resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.58.1.tgz",
|
||||
"integrity": "sha512-6LdVIUERWxQMmUSSQi0I53GgCBYgM2RpGngCPY7hSeju+VrKjq3lvs7HpJoPbDiY5QM5EYRtRX5fvrinnMAz3w==",
|
||||
"devOptional": true,
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"dependencies": {
|
||||
"playwright": "1.58.1"
|
||||
|
|
@ -3423,12 +3431,14 @@
|
|||
"version": "15.7.15",
|
||||
"resolved": "https://registry.npmjs.org/@types/prop-types/-/prop-types-15.7.15.tgz",
|
||||
"integrity": "sha512-F6bEyamV9jKGAFBEmlQnesRPGOQqS2+Uwi0Em15xenOxHaf2hv6L8YCVn3rPdPJOiJfPiCnLIRyvwVaqMY3MIw==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@types/react": {
|
||||
"version": "18.2.48",
|
||||
"resolved": "https://registry.npmjs.org/@types/react/-/react-18.2.48.tgz",
|
||||
"integrity": "sha512-qboRCl6Ie70DQQG9hhNREz81jqC1cs9EVNcjQ1AU+jH6NFfSAhVVbrrY/+nSF+Bsk4AOwm9Qa61InvMCyV+H3w==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@types/prop-types": "*",
|
||||
|
|
@ -3470,6 +3480,7 @@
|
|||
"version": "0.26.0",
|
||||
"resolved": "https://registry.npmjs.org/@types/scheduler/-/scheduler-0.26.0.tgz",
|
||||
"integrity": "sha512-WFHp9YUJQ6CKshqoC37iOlHnQSmxNc795UhB26CyBBttrN9svdIrUjl/NjnNmfcwtncN0h/0PPAFWv9ovP8mLA==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/@types/unist": {
|
||||
|
|
@ -4330,12 +4341,14 @@
|
|||
"version": "1.3.0",
|
||||
"resolved": "https://registry.npmjs.org/any-promise/-/any-promise-1.3.0.tgz",
|
||||
"integrity": "sha512-7UvmKalWRt1wgjL1RrGxoSJW/0QZFIegpeGvZG9kjp8vrRu55XTHbwnqq2GpXm9uLbcuhxm3IqX9OB4MZR1b2A==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/anymatch": {
|
||||
"version": "3.1.3",
|
||||
"resolved": "https://registry.npmjs.org/anymatch/-/anymatch-3.1.3.tgz",
|
||||
"integrity": "sha512-KMReFUr0B4t+D+OBkjR3KYqvocp2XaSzO55UcB6mgQMd3KbcE+mWTyvVV7D/zsdEbNnV6acZUutkiHQXvTr1Rw==",
|
||||
"dev": true,
|
||||
"license": "ISC",
|
||||
"dependencies": {
|
||||
"normalize-path": "^3.0.0",
|
||||
|
|
@ -4349,6 +4362,7 @@
|
|||
"version": "2.3.1",
|
||||
"resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.1.tgz",
|
||||
"integrity": "sha512-JU3teHTNjmE2VCGFzuY8EXzCDVwEqB2a8fsIvwaStHhAWJEeVd1o1QD80CU6+ZdEXXSLbSsuLwJjkCBWqRQUVA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=8.6"
|
||||
|
|
@ -4361,6 +4375,7 @@
|
|||
"version": "5.0.2",
|
||||
"resolved": "https://registry.npmjs.org/arg/-/arg-5.0.2.tgz",
|
||||
"integrity": "sha512-PYjyFOLKQ9y57JvQ6QLo8dAgNqswh8M1RMJYdQduT6xbWSgK36P/Z/v+p888pM69jMMfS8Xd8F6I1kQ/I9HUGg==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/argparse": {
|
||||
|
|
@ -4732,6 +4747,7 @@
|
|||
"version": "2.3.0",
|
||||
"resolved": "https://registry.npmjs.org/binary-extensions/-/binary-extensions-2.3.0.tgz",
|
||||
"integrity": "sha512-Ceh+7ox5qe7LJuLHoY0feh3pHuUDHAcRUeyL2VYghZwfpkNIy/+8Ocg0a3UuSoYzavmylwuLWQOf3hl0jjMMIw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=8"
|
||||
|
|
@ -4757,6 +4773,7 @@
|
|||
"version": "3.0.3",
|
||||
"resolved": "https://registry.npmjs.org/braces/-/braces-3.0.3.tgz",
|
||||
"integrity": "sha512-yQbXgO/OSZVD2IsiLlro+7Hf6Q18EJrKSEsdoMzKePKXct3gvD8oLcOQdIzGupr5Fj+EDe8gO/lxc1BzfMpxvA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"fill-range": "^7.1.1"
|
||||
|
|
@ -4872,6 +4889,7 @@
|
|||
"version": "2.0.1",
|
||||
"resolved": "https://registry.npmjs.org/camelcase-css/-/camelcase-css-2.0.1.tgz",
|
||||
"integrity": "sha512-QOSvevhslijgYwRx6Rv7zKdMF8lbRmx+uQGx2+vDc+KI/eBnsy9kit5aj23AgGu3pa4t9AgwbnXWqS+iOY+2aA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 6"
|
||||
|
|
@ -4995,6 +5013,7 @@
|
|||
"version": "3.6.0",
|
||||
"resolved": "https://registry.npmjs.org/chokidar/-/chokidar-3.6.0.tgz",
|
||||
"integrity": "sha512-7VT13fmjotKpGipCW9JEQAusEPE+Ei8nl6/g4FBAmIm0GOOLMua9NDDo/DWp0ZAxCr3cPq5ZpBqmPAQgDda2Pw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"anymatch": "~3.1.2",
|
||||
|
|
@ -5019,6 +5038,7 @@
|
|||
"version": "5.1.2",
|
||||
"resolved": "https://registry.npmjs.org/glob-parent/-/glob-parent-5.1.2.tgz",
|
||||
"integrity": "sha512-AOIgSQCepiJYwP3ARnGx+5VnTu2HBYdzbGP45eLw1vr3zB3vZLeyed1sC9hnbcOc9/SrMyM5RPQrkGz4aS9Zow==",
|
||||
"dev": true,
|
||||
"license": "ISC",
|
||||
"dependencies": {
|
||||
"is-glob": "^4.0.1"
|
||||
|
|
@ -5094,6 +5114,7 @@
|
|||
"version": "4.1.1",
|
||||
"resolved": "https://registry.npmjs.org/commander/-/commander-4.1.1.tgz",
|
||||
"integrity": "sha512-NOKm8xhkzAjzFx8B2v5OAHT+u5pRQc2UCa2Vq9jYL/31o2wi9mxBA7LIFs3sV5VSC49z6pEhfbMULvShKj26WA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 6"
|
||||
|
|
@ -5154,6 +5175,7 @@
|
|||
"version": "3.0.0",
|
||||
"resolved": "https://registry.npmjs.org/cssesc/-/cssesc-3.0.0.tgz",
|
||||
"integrity": "sha512-/Tb/JcjK111nNScGob5MNtsntNM1aCNUDipB/TkwZFhyDrrE47SOx/18wF2bbjgc3ZzCSKW1T5nt5EbFoAz/Vg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"bin": {
|
||||
"cssesc": "bin/cssesc"
|
||||
|
|
@ -5567,12 +5589,14 @@
|
|||
"version": "1.2.2",
|
||||
"resolved": "https://registry.npmjs.org/didyoumean/-/didyoumean-1.2.2.tgz",
|
||||
"integrity": "sha512-gxtyfqMg7GKyhQmb056K7M3xszy/myH8w+B4RT+QXBQsvAOdc3XymqDDPHx1BgPgsdAA5SIifona89YtRATDzw==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0"
|
||||
},
|
||||
"node_modules/dlv": {
|
||||
"version": "1.1.3",
|
||||
"resolved": "https://registry.npmjs.org/dlv/-/dlv-1.1.3.tgz",
|
||||
"integrity": "sha512-+HlytyjlPKnIG8XuRG8WvmBP8xs8P71y+SKKS6ZXWoEgLuePxtDoUEiH7WkdePWrQ5JBpE6aoVqfZfJUQkjXwA==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/doctrine": {
|
||||
|
|
@ -6486,6 +6510,7 @@
|
|||
"version": "1.20.1",
|
||||
"resolved": "https://registry.npmjs.org/fastq/-/fastq-1.20.1.tgz",
|
||||
"integrity": "sha512-GGToxJ/w1x32s/D2EKND7kTil4n8OVk/9mycTc4VDza13lOvpUZTGX3mFSCtV9ksdGBVzvsyAVLM6mHFThxXxw==",
|
||||
"dev": true,
|
||||
"license": "ISC",
|
||||
"dependencies": {
|
||||
"reusify": "^1.0.4"
|
||||
|
|
@ -6518,6 +6543,7 @@
|
|||
"version": "6.5.0",
|
||||
"resolved": "https://registry.npmjs.org/fdir/-/fdir-6.5.0.tgz",
|
||||
"integrity": "sha512-tIbYtZbucOs0BRGqPJkshJUYdL+SDH7dVM8gjy+ERp3WAUjLEFJE+02kanyHtwjWOnwrKYBiwAmM0p4kLJAnXg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=12.0.0"
|
||||
|
|
@ -6555,6 +6581,7 @@
|
|||
"version": "7.1.1",
|
||||
"resolved": "https://registry.npmjs.org/fill-range/-/fill-range-7.1.1.tgz",
|
||||
"integrity": "sha512-YsGpe3WHLK8ZYi4tWDg2Jy3ebRz2rXowDxnld4bkQB00cc/1Zw9AWnC0i9ztDJitivtQvaI9KaLyKrc+hBW0yg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"to-regex-range": "^5.0.1"
|
||||
|
|
@ -6715,6 +6742,7 @@
|
|||
"version": "2.3.2",
|
||||
"resolved": "https://registry.npmjs.org/fsevents/-/fsevents-2.3.2.tgz",
|
||||
"integrity": "sha512-xiqMQR4xAeHTuB9uWm+fFRcIOgKBMiOBP+eXiyT7jsgVCq1bkVygt00oASowB7EdtpOHaaPgKt812P9ab+DDKA==",
|
||||
"dev": true,
|
||||
"hasInstallScript": true,
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
|
|
@ -6865,6 +6893,7 @@
|
|||
"version": "6.0.2",
|
||||
"resolved": "https://registry.npmjs.org/glob-parent/-/glob-parent-6.0.2.tgz",
|
||||
"integrity": "sha512-XxwI8EOhVQgWp6iDL+3b0r86f4d6AX6zSU55HfB4ydCEuXLXc5FcYeOu+nnGftS4TEju/11rt4KJPTMgbfmv4A==",
|
||||
"dev": true,
|
||||
"license": "ISC",
|
||||
"dependencies": {
|
||||
"is-glob": "^4.0.3"
|
||||
|
|
@ -7362,6 +7391,7 @@
|
|||
"version": "2.1.0",
|
||||
"resolved": "https://registry.npmjs.org/is-binary-path/-/is-binary-path-2.1.0.tgz",
|
||||
"integrity": "sha512-ZMERYes6pDydyuGidse7OsHxtbI7WVeUEozgR/g7rd0xUimYNlvZRE/K2MgZTjWy725IfelLeVcEM97mmtRGXw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"binary-extensions": "^2.0.0"
|
||||
|
|
@ -7414,6 +7444,7 @@
|
|||
"version": "2.16.1",
|
||||
"resolved": "https://registry.npmjs.org/is-core-module/-/is-core-module-2.16.1.tgz",
|
||||
"integrity": "sha512-UfoeMA6fIJ8wTYFEUjelnaGI67v6+N7qXJEvQuIGa99l4xsCruSYOVSQ0uPANn4dAzm8lkYPaKLrrijLq7x23w==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"hasown": "^2.0.2"
|
||||
|
|
@ -7474,6 +7505,7 @@
|
|||
"version": "2.1.1",
|
||||
"resolved": "https://registry.npmjs.org/is-extglob/-/is-extglob-2.1.1.tgz",
|
||||
"integrity": "sha512-SbKbANkN603Vi4jEZv49LeVJMn4yGwsbzZworEoyEiutsN3nJYdbO36zfhGJ6QEDpOZIFkDtnq5JRxmvl3jsoQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=0.10.0"
|
||||
|
|
@ -7519,6 +7551,7 @@
|
|||
"version": "4.0.3",
|
||||
"resolved": "https://registry.npmjs.org/is-glob/-/is-glob-4.0.3.tgz",
|
||||
"integrity": "sha512-xelSayHH36ZgE7ZWhli7pW34hNbNl8Ojv5KVmkJD4hBdD3th8Tfk9vYasLM+mXWOZhFkgZfxhLSnrwRr4elSSg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"is-extglob": "^2.1.1"
|
||||
|
|
@ -7567,6 +7600,7 @@
|
|||
"version": "7.0.0",
|
||||
"resolved": "https://registry.npmjs.org/is-number/-/is-number-7.0.0.tgz",
|
||||
"integrity": "sha512-41Cifkg6e8TylSpdtTpeLVMqvSBEVzTttHvERD741+pnZ8ANv0004MRL43QKPDlK9cGvNp6NZWZUBlbGXYxxng==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=0.12.0"
|
||||
|
|
@ -7843,6 +7877,7 @@
|
|||
"version": "1.21.7",
|
||||
"resolved": "https://registry.npmjs.org/jiti/-/jiti-1.21.7.tgz",
|
||||
"integrity": "sha512-/imKNG4EbWNrVjoNC/1H5/9GFy+tqjGBHCaSsN+P2RnPqjsLmv6UD3Ej+Kj8nBWaRAwyk7kK5ZUc+OEatnTR3A==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"bin": {
|
||||
"jiti": "bin/jiti.js"
|
||||
|
|
@ -8128,6 +8163,7 @@
|
|||
"version": "3.1.3",
|
||||
"resolved": "https://registry.npmjs.org/lilconfig/-/lilconfig-3.1.3.tgz",
|
||||
"integrity": "sha512-/vlFKAoH5Cgt3Ie+JLhRbwOsCQePABiU3tJ1egGvyQ+33R/vcwM2Zl2QR/LzjsBeItPt3oSVXapn+m4nQDvpzw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=14"
|
||||
|
|
@ -8140,6 +8176,7 @@
|
|||
"version": "1.2.4",
|
||||
"resolved": "https://registry.npmjs.org/lines-and-columns/-/lines-and-columns-1.2.4.tgz",
|
||||
"integrity": "sha512-7ylylesZQ/PV29jhEDl3Ufjo6ZX7gCqJr5F7PKrqc93v7fzSymt1BpwEU8nAUXs8qzzvqhbjhK5QZg6Mt/HkBg==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/locate-path": {
|
||||
|
|
@ -8454,6 +8491,7 @@
|
|||
"version": "1.4.1",
|
||||
"resolved": "https://registry.npmjs.org/merge2/-/merge2-1.4.1.tgz",
|
||||
"integrity": "sha512-8q7VEgMJW4J8tcfVPy8g09NcQwZdbwFEqhe/WZkoIzjn/3TGDwtOCYtXGxA3O8tPzpczCCDgv+P2P5y00ZJOOg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 8"
|
||||
|
|
@ -8905,6 +8943,7 @@
|
|||
"version": "4.0.8",
|
||||
"resolved": "https://registry.npmjs.org/micromatch/-/micromatch-4.0.8.tgz",
|
||||
"integrity": "sha512-PXwfBhYu0hBCPw8Dn0E+WDYb7af3dSLVWKi3HGv84IdF4TyFoC0ysxFd0Goxw7nSv4T/PzEJQxsYsEiFCKo2BA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"braces": "^3.0.3",
|
||||
|
|
@ -8918,6 +8957,7 @@
|
|||
"version": "2.3.1",
|
||||
"resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.1.tgz",
|
||||
"integrity": "sha512-JU3teHTNjmE2VCGFzuY8EXzCDVwEqB2a8fsIvwaStHhAWJEeVd1o1QD80CU6+ZdEXXSLbSsuLwJjkCBWqRQUVA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=8.6"
|
||||
|
|
@ -9032,6 +9072,7 @@
|
|||
"version": "2.7.0",
|
||||
"resolved": "https://registry.npmjs.org/mz/-/mz-2.7.0.tgz",
|
||||
"integrity": "sha512-z81GNO7nnYMEhrGh9LeymoE4+Yr0Wn5McHIZMK5cfQCl+NDX08sCZgUc9/6MHni9IWuFLm1Z3HTCXu2z9fN62Q==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"any-promise": "^1.0.0",
|
||||
|
|
@ -9243,6 +9284,7 @@
|
|||
"version": "3.0.0",
|
||||
"resolved": "https://registry.npmjs.org/normalize-path/-/normalize-path-3.0.0.tgz",
|
||||
"integrity": "sha512-6eZs5Ls3WtCisHWp9S2GUy8dqkpGi4BVSz3GaqiE6ezub0512ESztXUwUB6C6IKbQkY2Pnb/mD4WYojCRwcwLA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=0.10.0"
|
||||
|
|
@ -9261,6 +9303,7 @@
|
|||
"version": "3.0.0",
|
||||
"resolved": "https://registry.npmjs.org/object-hash/-/object-hash-3.0.0.tgz",
|
||||
"integrity": "sha512-RSn9F68PjH9HqtltsSnqYC1XXoWe9Bju5+213R98cNGttag9q9yAOTzdbsqvIa7aNm5WffBZFpWYr2aWrklWAw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 6"
|
||||
|
|
@ -9605,6 +9648,7 @@
|
|||
"version": "1.0.7",
|
||||
"resolved": "https://registry.npmjs.org/path-parse/-/path-parse-1.0.7.tgz",
|
||||
"integrity": "sha512-LDJzPVEEEPR+y48z93A0Ed0yXb8pAByGWo/k5YYdYgpY2/2EsOsksJrq7lOHxryrVOn1ejG6oAp8ahvOIQD8sw==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/path-scurry": {
|
||||
|
|
@ -9651,6 +9695,7 @@
|
|||
"version": "4.0.3",
|
||||
"resolved": "https://registry.npmjs.org/picomatch/-/picomatch-4.0.3.tgz",
|
||||
"integrity": "sha512-5gTmgEY/sqK6gFXLIsQNH19lWb4ebPDLA4SdLP7dsWkIXHWlG66oPuVvXSGFPppYZz8ZDZq0dYYrbHfBCVUb1Q==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=12"
|
||||
|
|
@ -9663,6 +9708,7 @@
|
|||
"version": "2.3.0",
|
||||
"resolved": "https://registry.npmjs.org/pify/-/pify-2.3.0.tgz",
|
||||
"integrity": "sha512-udgsAY+fTnvv7kI7aaxbqwWNb0AHiB0qBO89PZKPkoTmGOgdbrHDKD+0B2X4uTfJ/FT1R09r9gTsjUjNJotuog==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=0.10.0"
|
||||
|
|
@ -9672,6 +9718,7 @@
|
|||
"version": "4.0.7",
|
||||
"resolved": "https://registry.npmjs.org/pirates/-/pirates-4.0.7.tgz",
|
||||
"integrity": "sha512-TfySrs/5nm8fQJDcBDuUng3VOUKsd7S+zqvbOTiGXHfxX4wK31ard+hoNuvkicM/2YFzlpDgABOevKSsB4G/FA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 6"
|
||||
|
|
@ -9681,7 +9728,7 @@
|
|||
"version": "1.58.1",
|
||||
"resolved": "https://registry.npmjs.org/playwright/-/playwright-1.58.1.tgz",
|
||||
"integrity": "sha512-+2uTZHxSCcxjvGc5C891LrS1/NlxglGxzrC4seZiVjcYVQfUa87wBL6rTDqzGjuoWNjnBzRqKmF6zRYGMvQUaQ==",
|
||||
"devOptional": true,
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"dependencies": {
|
||||
"playwright-core": "1.58.1"
|
||||
|
|
@ -9700,7 +9747,7 @@
|
|||
"version": "1.58.1",
|
||||
"resolved": "https://registry.npmjs.org/playwright-core/-/playwright-core-1.58.1.tgz",
|
||||
"integrity": "sha512-bcWzOaTxcW+VOOGBCQgnaKToLJ65d6AqfLVKEWvexyS3AS6rbXl+xdpYRMGSRBClPvyj44njOWoxjNdL/H9UNg==",
|
||||
"devOptional": true,
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"bin": {
|
||||
"playwright-core": "cli.js"
|
||||
|
|
@ -9723,6 +9770,7 @@
|
|||
"version": "8.5.6",
|
||||
"resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.6.tgz",
|
||||
"integrity": "sha512-3Ybi1tAuwAP9s0r1UQ2J4n5Y0G05bJkpUIO0/bI9MhwmD70S5aTWbXGBwxHrelT+XM1k6dM0pk+SwNkpTRN7Pg==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "opencollective",
|
||||
|
|
@ -9751,6 +9799,7 @@
|
|||
"version": "15.1.0",
|
||||
"resolved": "https://registry.npmjs.org/postcss-import/-/postcss-import-15.1.0.tgz",
|
||||
"integrity": "sha512-hpr+J05B2FVYUAXHeK1YyI267J/dDDhMU6B6civm8hSY1jYJnBXxzKDKDswzJmtLHryrjhnDjqqp/49t8FALew==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"postcss-value-parser": "^4.0.0",
|
||||
|
|
@ -9768,6 +9817,7 @@
|
|||
"version": "4.1.0",
|
||||
"resolved": "https://registry.npmjs.org/postcss-js/-/postcss-js-4.1.0.tgz",
|
||||
"integrity": "sha512-oIAOTqgIo7q2EOwbhb8UalYePMvYoIeRY2YKntdpFQXNosSu3vLrniGgmH9OKs/qAkfoj5oB3le/7mINW1LCfw==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "opencollective",
|
||||
|
|
@ -9793,6 +9843,7 @@
|
|||
"version": "6.0.1",
|
||||
"resolved": "https://registry.npmjs.org/postcss-load-config/-/postcss-load-config-6.0.1.tgz",
|
||||
"integrity": "sha512-oPtTM4oerL+UXmx+93ytZVN82RrlY/wPUV8IeDxFrzIjXOLF1pN+EmKPLbubvKHT2HC20xXsCAH2Z+CKV6Oz/g==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "opencollective",
|
||||
|
|
@ -9835,6 +9886,7 @@
|
|||
"version": "6.2.0",
|
||||
"resolved": "https://registry.npmjs.org/postcss-nested/-/postcss-nested-6.2.0.tgz",
|
||||
"integrity": "sha512-HQbt28KulC5AJzG+cZtj9kvKB93CFCdLvog1WFLf1D+xmMvPGlBstkpTEZfK5+AN9hfJocyBFCNiqyS48bpgzQ==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "opencollective",
|
||||
|
|
@ -9860,6 +9912,7 @@
|
|||
"version": "6.1.2",
|
||||
"resolved": "https://registry.npmjs.org/postcss-selector-parser/-/postcss-selector-parser-6.1.2.tgz",
|
||||
"integrity": "sha512-Q8qQfPiZ+THO/3ZrOrO0cJJKfpYCagtMUkXbnEfmgUjwXg6z/WBeOyS9APBBPCTSiDV+s4SwQGu8yFsiMRIudg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"cssesc": "^3.0.0",
|
||||
|
|
@ -9873,6 +9926,7 @@
|
|||
"version": "4.2.0",
|
||||
"resolved": "https://registry.npmjs.org/postcss-value-parser/-/postcss-value-parser-4.2.0.tgz",
|
||||
"integrity": "sha512-1NNCs6uurfkVbeXG4S8JFT9t19m45ICnif8zWLd5oPSZ50QnwMfK+H3jv408d4jw/7Bttv5axS5IiHoLaVNHeQ==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/prelude-ls": {
|
||||
|
|
@ -9986,6 +10040,7 @@
|
|||
"version": "1.2.3",
|
||||
"resolved": "https://registry.npmjs.org/queue-microtask/-/queue-microtask-1.2.3.tgz",
|
||||
"integrity": "sha512-NuaNSa6flKT5JaSYQzJok04JzTL1CA6aGhv5rfLW3PgqA+M2ChpZQnAC8h8i4ZFkBS8X5RqkDBHA7r4hej3K9A==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "github",
|
||||
|
|
@ -10774,6 +10829,7 @@
|
|||
"version": "1.0.0",
|
||||
"resolved": "https://registry.npmjs.org/read-cache/-/read-cache-1.0.0.tgz",
|
||||
"integrity": "sha512-Owdv/Ft7IjOgm/i0xvNDZ1LrRANRfew4b2prF3OWMQLxLfu3bS8FVhCsrSCMK4lR56Y9ya+AThoTpDCTxCmpRA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"pify": "^2.3.0"
|
||||
|
|
@ -10783,6 +10839,7 @@
|
|||
"version": "3.6.0",
|
||||
"resolved": "https://registry.npmjs.org/readdirp/-/readdirp-3.6.0.tgz",
|
||||
"integrity": "sha512-hOS089on8RduqdbhvQ5Z37A0ESjsqz6qnRcffsMU3495FuTdqSm+7bhJ29JvIOsBDEEnan5DPu9t3To9VRlMzA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"picomatch": "^2.2.1"
|
||||
|
|
@ -10795,6 +10852,7 @@
|
|||
"version": "2.3.1",
|
||||
"resolved": "https://registry.npmjs.org/picomatch/-/picomatch-2.3.1.tgz",
|
||||
"integrity": "sha512-JU3teHTNjmE2VCGFzuY8EXzCDVwEqB2a8fsIvwaStHhAWJEeVd1o1QD80CU6+ZdEXXSLbSsuLwJjkCBWqRQUVA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=8.6"
|
||||
|
|
@ -11059,6 +11117,7 @@
|
|||
"version": "1.22.11",
|
||||
"resolved": "https://registry.npmjs.org/resolve/-/resolve-1.22.11.tgz",
|
||||
"integrity": "sha512-RfqAvLnMl313r7c9oclB1HhUEAezcpLjz95wFH4LVuhk9JF/r22qmVP9AMmOU4vMX7Q8pN8jwNg/CSpdFnMjTQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"is-core-module": "^2.16.1",
|
||||
|
|
@ -11099,6 +11158,7 @@
|
|||
"version": "1.1.0",
|
||||
"resolved": "https://registry.npmjs.org/reusify/-/reusify-1.1.0.tgz",
|
||||
"integrity": "sha512-g6QUff04oZpHs0eG5p83rFLhHeV00ug/Yf9nZM6fLeUrPguBTkTQOdpAWWspMh55TZfVQDPaN3NQJfbVRAxdIw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"iojs": ">=1.0.0",
|
||||
|
|
@ -11154,6 +11214,7 @@
|
|||
"version": "1.2.0",
|
||||
"resolved": "https://registry.npmjs.org/run-parallel/-/run-parallel-1.2.0.tgz",
|
||||
"integrity": "sha512-5l4VyZR86LZ/lDxZTR6jqL8AFE2S0IFLMP26AbjsLVADxHdhB/c0GUsH+y39UfCi3dzz8OlQuPmnaJOMoDHQBA==",
|
||||
"dev": true,
|
||||
"funding": [
|
||||
{
|
||||
"type": "github",
|
||||
|
|
@ -11794,6 +11855,7 @@
|
|||
"version": "3.35.1",
|
||||
"resolved": "https://registry.npmjs.org/sucrase/-/sucrase-3.35.1.tgz",
|
||||
"integrity": "sha512-DhuTmvZWux4H1UOnWMB3sk0sbaCVOoQZjv8u1rDoTV0HTdGem9hkAZtl4JZy8P2z4Bg0nT+YMeOFyVr4zcG5Tw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@jridgewell/gen-mapping": "^0.3.2",
|
||||
|
|
@ -11829,6 +11891,7 @@
|
|||
"version": "1.0.0",
|
||||
"resolved": "https://registry.npmjs.org/supports-preserve-symlinks-flag/-/supports-preserve-symlinks-flag-1.0.0.tgz",
|
||||
"integrity": "sha512-ot0WnXS9fgdkgIcePe6RHNk1WA8+muPa6cSjeR3V8K27q9BB1rTE3R1p7Hv0z1ZyAc8s6Vvv8DIyWf681MAt0w==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">= 0.4"
|
||||
|
|
@ -11864,6 +11927,7 @@
|
|||
"version": "3.4.19",
|
||||
"resolved": "https://registry.npmjs.org/tailwindcss/-/tailwindcss-3.4.19.tgz",
|
||||
"integrity": "sha512-3ofp+LL8E+pK/JuPLPggVAIaEuhvIz4qNcf3nA1Xn2o/7fb7s/TYpHhwGDv1ZU3PkBluUVaF8PyCHcm48cKLWQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@alloc/quick-lru": "^5.2.0",
|
||||
|
|
@ -11901,6 +11965,7 @@
|
|||
"version": "3.3.3",
|
||||
"resolved": "https://registry.npmjs.org/fast-glob/-/fast-glob-3.3.3.tgz",
|
||||
"integrity": "sha512-7MptL8U0cqcFdzIzwOTHoilX9x5BrNqye7Z/LuC7kCMRio1EMSyqRK3BEAUD7sXRq4iT4AzTVuZdhgQ2TCvYLg==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"@nodelib/fs.stat": "^2.0.2",
|
||||
|
|
@ -11917,6 +11982,7 @@
|
|||
"version": "5.1.2",
|
||||
"resolved": "https://registry.npmjs.org/glob-parent/-/glob-parent-5.1.2.tgz",
|
||||
"integrity": "sha512-AOIgSQCepiJYwP3ARnGx+5VnTu2HBYdzbGP45eLw1vr3zB3vZLeyed1sC9hnbcOc9/SrMyM5RPQrkGz4aS9Zow==",
|
||||
"dev": true,
|
||||
"license": "ISC",
|
||||
"dependencies": {
|
||||
"is-glob": "^4.0.1"
|
||||
|
|
@ -11944,6 +12010,7 @@
|
|||
"version": "3.3.1",
|
||||
"resolved": "https://registry.npmjs.org/thenify/-/thenify-3.3.1.tgz",
|
||||
"integrity": "sha512-RVZSIV5IG10Hk3enotrhvz0T9em6cyHBLkH/YAZuKqd8hRkKhSfCGIcP2KUY0EPxndzANBmNllzWPwak+bheSw==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"any-promise": "^1.0.0"
|
||||
|
|
@ -11953,6 +12020,7 @@
|
|||
"version": "1.6.0",
|
||||
"resolved": "https://registry.npmjs.org/thenify-all/-/thenify-all-1.6.0.tgz",
|
||||
"integrity": "sha512-RNxQH/qI8/t3thXJDwcstUO4zeqo64+Uy/+sNVRBx4Xn2OX+OZ9oP+iJnNFqplFra2ZUVeKCSa2oVWi3T4uVmA==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"thenify": ">= 3.1.0 < 4"
|
||||
|
|
@ -11994,6 +12062,7 @@
|
|||
"version": "0.2.15",
|
||||
"resolved": "https://registry.npmjs.org/tinyglobby/-/tinyglobby-0.2.15.tgz",
|
||||
"integrity": "sha512-j2Zq4NyQYG5XMST4cbs02Ak8iJUdxRM0XI5QyxXuZOzKOINmWurp3smXu3y5wDcJrptwpSjgXHzIQxR0omXljQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"fdir": "^6.5.0",
|
||||
|
|
@ -12060,6 +12129,7 @@
|
|||
"version": "5.0.1",
|
||||
"resolved": "https://registry.npmjs.org/to-regex-range/-/to-regex-range-5.0.1.tgz",
|
||||
"integrity": "sha512-65P7iz6X5yEr1cwcgvQxbbIw7Uk3gOy5dIdtZ4rDveLqhrdJP+Li/Hx6tyK0NEb+2GCyneCMJiGqrADCSNk8sQ==",
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"dependencies": {
|
||||
"is-number": "^7.0.0"
|
||||
|
|
@ -12147,6 +12217,7 @@
|
|||
"version": "0.1.13",
|
||||
"resolved": "https://registry.npmjs.org/ts-interface-checker/-/ts-interface-checker-0.1.13.tgz",
|
||||
"integrity": "sha512-Y/arvbn+rrz3JCKl9C4kVNfTfSm2/mEp5FSz5EsZSANGPSlQrpRI5M4PKF+mJnE52jOO90PnPSc3Ur3bTQw0gA==",
|
||||
"dev": true,
|
||||
"license": "Apache-2.0"
|
||||
},
|
||||
"node_modules/tsconfig-paths": {
|
||||
|
|
@ -12263,7 +12334,7 @@
|
|||
"version": "5.3.3",
|
||||
"resolved": "https://registry.npmjs.org/typescript/-/typescript-5.3.3.tgz",
|
||||
"integrity": "sha512-pXWcraxM0uxAS+tN0AG/BF2TyqmHO014Z070UsJ+pFvYuRSq8KH8DmWpnbXe0pEPDHXZV3FcAbJkijJ5oNEnWw==",
|
||||
"devOptional": true,
|
||||
"dev": true,
|
||||
"license": "Apache-2.0",
|
||||
"bin": {
|
||||
"tsc": "bin/tsc",
|
||||
|
|
@ -12465,6 +12536,7 @@
|
|||
"version": "1.0.2",
|
||||
"resolved": "https://registry.npmjs.org/util-deprecate/-/util-deprecate-1.0.2.tgz",
|
||||
"integrity": "sha512-EPD5q1uXyFxJpCrLnCc1nHnq3gOa6DZBocAIiI2TaSCA7VCJ1UJDMagCzIkXNsUYfD1daK//LTEQ8xiIbrHtcw==",
|
||||
"dev": true,
|
||||
"license": "MIT"
|
||||
},
|
||||
"node_modules/uuid": {
|
||||
|
|
@ -12918,7 +12990,7 @@
|
|||
"version": "8.19.0",
|
||||
"resolved": "https://registry.npmjs.org/ws/-/ws-8.19.0.tgz",
|
||||
"integrity": "sha512-blAT2mjOEIi0ZzruJfIhb3nps74PRWTCz1IjglWEEpQl5XS/UNama6u2/rjFkDDouqr4L67ry+1aGIALViWjDg==",
|
||||
"devOptional": true,
|
||||
"dev": true,
|
||||
"license": "MIT",
|
||||
"engines": {
|
||||
"node": ">=10.0.0"
|
||||
|
|
@ -12975,17 +13047,6 @@
|
|||
"url": "https://github.com/sponsors/sindresorhus"
|
||||
}
|
||||
},
|
||||
"node_modules/zod": {
|
||||
"version": "3.25.76",
|
||||
"resolved": "https://registry.npmjs.org/zod/-/zod-3.25.76.tgz",
|
||||
"integrity": "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ==",
|
||||
"license": "MIT",
|
||||
"optional": true,
|
||||
"peer": true,
|
||||
"funding": {
|
||||
"url": "https://github.com/sponsors/colinhacks"
|
||||
}
|
||||
},
|
||||
"node_modules/zwitch": {
|
||||
"version": "2.0.4",
|
||||
"resolved": "https://registry.npmjs.org/zwitch/-/zwitch-2.0.4.tgz",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,131 @@
|
|||
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 PricingCalculator from "./index";
|
||||
import type { ModelEntry } from "./types";
|
||||
import type { MultiModelResult } from "./types";
|
||||
|
||||
vi.mock("./use_multi_cost_estimate", () => ({
|
||||
useMultiCostEstimate: vi.fn(() => ({
|
||||
debouncedFetchForEntry: vi.fn(),
|
||||
removeEntry: vi.fn(),
|
||||
getMultiModelResult: vi.fn((entries: ModelEntry[]): MultiModelResult => ({
|
||||
entries: entries.map((e) => ({ entry: e, result: null, loading: false, error: null })),
|
||||
totals: {
|
||||
cost_per_request: 0,
|
||||
daily_cost: null,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
})),
|
||||
})),
|
||||
}));
|
||||
|
||||
vi.mock("./multi_export_utils", () => ({
|
||||
exportMultiToPDF: vi.fn(),
|
||||
exportMultiToCSV: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/utils/dataUtils", () => ({
|
||||
formatNumberWithCommas: vi.fn((v: number, d: number = 0) =>
|
||||
Number.isFinite(v) ? v.toFixed(d) : "-"
|
||||
),
|
||||
}));
|
||||
|
||||
const DEFAULT_PROPS = {
|
||||
accessToken: "test-token",
|
||||
models: ["gpt-4", "gpt-3.5-turbo", "claude-3-sonnet"],
|
||||
};
|
||||
|
||||
describe("PricingCalculator", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should render the calculator with an initial model row", () => {
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
expect(screen.getByRole("table")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render the time period toggle with Per Day and Per Month options", () => {
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
expect(screen.getByText("Per Day")).toBeInTheDocument();
|
||||
expect(screen.getByText("Per Month")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should render an Add Another Model button", () => {
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
expect(screen.getByRole("button", { name: /add another model/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the Requests/Month column header by default", () => {
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
expect(screen.getByText("Requests/Month")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should add a new row when Add Another Model is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
|
||||
const table = screen.getByRole("table");
|
||||
const initialRows = within(table).getAllByRole("row");
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /add another model/i }));
|
||||
|
||||
const updatedRows = within(table).getAllByRole("row");
|
||||
// One new data row added (header row + data rows)
|
||||
expect(updatedRows.length).toBeGreaterThan(initialRows.length);
|
||||
});
|
||||
|
||||
it("should have the delete button disabled when there is only one row", () => {
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
const allButtons = screen.getAllByRole("button");
|
||||
const disabledButtons = allButtons.filter((btn) => btn.hasAttribute("disabled"));
|
||||
expect(disabledButtons.length).toBeGreaterThan(0);
|
||||
});
|
||||
|
||||
it("should have no disabled buttons after adding a second row", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /add another model/i }));
|
||||
|
||||
// With two rows, no delete buttons should be disabled
|
||||
const allButtons = screen.getAllByRole("button");
|
||||
const disabledButtons = allButtons.filter((btn) => btn.hasAttribute("disabled"));
|
||||
expect(disabledButtons.length).toBe(0);
|
||||
});
|
||||
|
||||
describe("time period toggle", () => {
|
||||
it("should switch the column header to Requests/Day when Per Day is selected", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
|
||||
await user.click(screen.getByText("Per Day"));
|
||||
|
||||
expect(screen.getByText("Requests/Day")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should switch the column header back to Requests/Month when Per Month is selected", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
|
||||
await user.click(screen.getByText("Per Day"));
|
||||
expect(screen.getByText("Requests/Day")).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByText("Per Month"));
|
||||
expect(screen.getByText("Requests/Month")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should render column headers for Model, Input Tokens, and Output Tokens", () => {
|
||||
renderWithProviders(<PricingCalculator {...DEFAULT_PROPS} />);
|
||||
expect(screen.getByText("Model")).toBeInTheDocument();
|
||||
expect(screen.getByText("Input Tokens")).toBeInTheDocument();
|
||||
expect(screen.getByText("Output Tokens")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,305 @@
|
|||
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 MultiCostResults from "./multi_cost_results";
|
||||
import type { MultiModelResult } from "./types";
|
||||
import type { CostEstimateResponse } from "../types";
|
||||
|
||||
vi.mock("./multi_export_utils", () => ({
|
||||
exportMultiToPDF: vi.fn(),
|
||||
exportMultiToCSV: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("@/utils/dataUtils", () => ({
|
||||
formatNumberWithCommas: vi.fn((v: number, d: number = 0) =>
|
||||
Number.isFinite(v) ? v.toFixed(d) : "-"
|
||||
),
|
||||
}));
|
||||
|
||||
function makeCostResponse(overrides: Partial<CostEstimateResponse> = {}): CostEstimateResponse {
|
||||
return {
|
||||
model: "gpt-4",
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
num_requests_per_day: 100,
|
||||
num_requests_per_month: null,
|
||||
cost_per_request: 0.05,
|
||||
input_cost_per_request: 0.03,
|
||||
output_cost_per_request: 0.02,
|
||||
margin_cost_per_request: 0,
|
||||
daily_cost: 5.0,
|
||||
daily_input_cost: 3.0,
|
||||
daily_output_cost: 2.0,
|
||||
daily_margin_cost: 0,
|
||||
monthly_cost: null,
|
||||
monthly_input_cost: null,
|
||||
monthly_output_cost: null,
|
||||
monthly_margin_cost: null,
|
||||
input_cost_per_token: null,
|
||||
output_cost_per_token: null,
|
||||
provider: "openai",
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function makeMultiResult(overrides: Partial<MultiModelResult> = {}): MultiModelResult {
|
||||
return {
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 },
|
||||
result: makeCostResponse(),
|
||||
loading: false,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: 0.05,
|
||||
daily_cost: 5.0,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function emptyMultiResult(): MultiModelResult {
|
||||
return {
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "", input_tokens: 1000, output_tokens: 500 },
|
||||
result: null,
|
||||
loading: false,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: 0,
|
||||
daily_cost: null,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
describe("MultiCostResults", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
describe("when no model has been selected", () => {
|
||||
it("should show a prompt to select models", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={emptyMultiResult()} timePeriod="month" />
|
||||
);
|
||||
expect(screen.getByText(/select models above to see cost estimates/i)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("when results are loading and no data has arrived yet", () => {
|
||||
it("should show a calculating costs spinner", () => {
|
||||
const multiResult: MultiModelResult = {
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 },
|
||||
result: null,
|
||||
loading: true,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: 0,
|
||||
daily_cost: null,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
};
|
||||
|
||||
renderWithProviders(<MultiCostResults multiResult={multiResult} timePeriod="month" />);
|
||||
expect(screen.getByText(/calculating costs/i)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("when there are errors but no valid results", () => {
|
||||
it("should display the error message with the model name", () => {
|
||||
const multiResult: MultiModelResult = {
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "bad-model", input_tokens: 0, output_tokens: 0 },
|
||||
result: null,
|
||||
loading: false,
|
||||
error: "Pricing not found",
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: 0,
|
||||
daily_cost: null,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
};
|
||||
|
||||
renderWithProviders(<MultiCostResults multiResult={multiResult} timePeriod="month" />);
|
||||
expect(screen.getByText(/bad-model/i)).toBeInTheDocument();
|
||||
expect(screen.getByText(/Pricing not found/i)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("when valid results are available", () => {
|
||||
it("should show the Cost Estimates heading", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.getByText("Cost Estimates")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display the Total Per Request statistic", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.getByText("Total Per Request")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display Total Daily statistic when timePeriod is day", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.getByText("Total Daily")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should display Total Monthly statistic when timePeriod is month", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults
|
||||
multiResult={makeMultiResult({
|
||||
totals: { cost_per_request: 0.05, daily_cost: null, monthly_cost: 150.0, margin_per_request: 0, daily_margin: null, monthly_margin: null },
|
||||
})}
|
||||
timePeriod="month"
|
||||
/>
|
||||
);
|
||||
expect(screen.getByText("Total Monthly")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the model name in the summary table", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.getByText("gpt-4")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the provider tag next to the model name", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.getByText("openai")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the Export button when results are available", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.getByRole("button", { name: /export/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should expand the model breakdown row when the expand button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
|
||||
// The expand column renders a button (RightOutlined icon) for rows without errors
|
||||
const expandButtons = screen.getAllByRole("button");
|
||||
// Find the small expand button (not the Export button)
|
||||
const expandButton = expandButtons.find(
|
||||
(btn) => !btn.textContent?.toLowerCase().includes("export")
|
||||
);
|
||||
expect(expandButton).toBeDefined();
|
||||
|
||||
await user.click(expandButton!);
|
||||
|
||||
// After expanding, the SingleModelBreakdown should be visible
|
||||
expect(screen.getByText("Total/Request")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the collapse icon after expanding a row", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
|
||||
const getExpandButton = () => {
|
||||
const allButtons = screen.getAllByRole("button");
|
||||
return allButtons.find((btn) => !btn.textContent?.toLowerCase().includes("export"));
|
||||
};
|
||||
|
||||
// Before expand: button has the "down" aria-label (RightOutlined renders as down in ant icons)
|
||||
// Just verify clicking works and the breakdown content appears
|
||||
await user.click(getExpandButton()!);
|
||||
expect(screen.getByText("Total/Request")).toBeInTheDocument();
|
||||
|
||||
// After a second click, the row collapses — content may be hidden or removed
|
||||
await user.click(getExpandButton()!);
|
||||
// The expanded content should no longer be visible
|
||||
expect(screen.queryByText("Total/Request")).not.toBeVisible();
|
||||
});
|
||||
});
|
||||
|
||||
describe("margin section", () => {
|
||||
it("should show margin fee details when margin per request is greater than zero", () => {
|
||||
const multiResult = makeMultiResult({
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 },
|
||||
result: makeCostResponse({ margin_cost_per_request: 0.01, daily_margin_cost: 1.0 }),
|
||||
loading: false,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: 0.06,
|
||||
daily_cost: 6.0,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0.01,
|
||||
daily_margin: 1.0,
|
||||
monthly_margin: null,
|
||||
},
|
||||
});
|
||||
|
||||
renderWithProviders(<MultiCostResults multiResult={multiResult} timePeriod="day" />);
|
||||
expect(screen.getByText("Margin Fee/Request")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should not show margin fee details when margin per request is zero", () => {
|
||||
renderWithProviders(
|
||||
<MultiCostResults multiResult={makeMultiResult()} timePeriod="day" />
|
||||
);
|
||||
expect(screen.queryByText("Margin Fee/Request")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
describe("when a model has zero cost", () => {
|
||||
it("should show a warning about missing pricing data", () => {
|
||||
const multiResult = makeMultiResult({
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "custom-model", input_tokens: 1000, output_tokens: 500 },
|
||||
result: makeCostResponse({ model: "custom-model", cost_per_request: 0 }),
|
||||
loading: false,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
});
|
||||
|
||||
renderWithProviders(<MultiCostResults multiResult={multiResult} timePeriod="day" />);
|
||||
expect(screen.getByText(/no pricing data found/i)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,146 @@
|
|||
import React from "react";
|
||||
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
|
||||
import { screen, fireEvent } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { renderWithProviders } from "../../../../tests/test-utils";
|
||||
import MultiExportDropdown from "./multi_export_dropdown";
|
||||
import type { MultiModelResult } from "./types";
|
||||
|
||||
vi.mock("./multi_export_utils", () => ({
|
||||
exportMultiToPDF: vi.fn(),
|
||||
exportMultiToCSV: vi.fn(),
|
||||
}));
|
||||
|
||||
import { exportMultiToPDF, exportMultiToCSV } from "./multi_export_utils";
|
||||
|
||||
function makeMultiResult(hasResult: boolean): MultiModelResult {
|
||||
return {
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 },
|
||||
result: hasResult
|
||||
? {
|
||||
model: "gpt-4",
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
num_requests_per_day: null,
|
||||
num_requests_per_month: null,
|
||||
cost_per_request: 0.05,
|
||||
input_cost_per_request: 0.03,
|
||||
output_cost_per_request: 0.02,
|
||||
margin_cost_per_request: 0,
|
||||
daily_cost: null,
|
||||
daily_input_cost: null,
|
||||
daily_output_cost: null,
|
||||
daily_margin_cost: null,
|
||||
monthly_cost: null,
|
||||
monthly_input_cost: null,
|
||||
monthly_output_cost: null,
|
||||
monthly_margin_cost: null,
|
||||
input_cost_per_token: null,
|
||||
output_cost_per_token: null,
|
||||
provider: "openai",
|
||||
}
|
||||
: null,
|
||||
loading: false,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: hasResult ? 0.05 : 0,
|
||||
daily_cost: null,
|
||||
monthly_cost: null,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
};
|
||||
}
|
||||
|
||||
describe("MultiExportDropdown", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should not render anything when no entries have results", () => {
|
||||
const { container } = renderWithProviders(
|
||||
<MultiExportDropdown multiResult={makeMultiResult(false)} />
|
||||
);
|
||||
expect(container.firstChild).toBeNull();
|
||||
});
|
||||
|
||||
it("should render the Export button when at least one entry has a result", () => {
|
||||
renderWithProviders(<MultiExportDropdown multiResult={makeMultiResult(true)} />);
|
||||
expect(screen.getByRole("button", { name: /^export$/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should show the export menu when the Export button is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<MultiExportDropdown multiResult={makeMultiResult(true)} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
|
||||
expect(screen.getByText("Export as PDF")).toBeInTheDocument();
|
||||
expect(screen.getByText("Export as CSV")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should hide the export menu when the Export button is clicked again", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<MultiExportDropdown multiResult={makeMultiResult(true)} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
expect(screen.getByText("Export as PDF")).toBeInTheDocument();
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
expect(screen.queryByText("Export as PDF")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call exportMultiToPDF and close the menu when Export as PDF is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<MultiExportDropdown multiResult={makeMultiResult(true)} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
await user.click(screen.getByText("Export as PDF"));
|
||||
|
||||
expect(exportMultiToPDF).toHaveBeenCalledTimes(1);
|
||||
expect(screen.queryByText("Export as PDF")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should call exportMultiToCSV and close the menu when Export as CSV is clicked", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(<MultiExportDropdown multiResult={makeMultiResult(true)} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
await user.click(screen.getByText("Export as CSV"));
|
||||
|
||||
expect(exportMultiToCSV).toHaveBeenCalledTimes(1);
|
||||
expect(screen.queryByText("Export as CSV")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should pass the multiResult to the export functions", async () => {
|
||||
const user = userEvent.setup();
|
||||
const multiResult = makeMultiResult(true);
|
||||
renderWithProviders(<MultiExportDropdown multiResult={multiResult} />);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
await user.click(screen.getByText("Export as PDF"));
|
||||
|
||||
expect(exportMultiToPDF).toHaveBeenCalledWith(multiResult);
|
||||
});
|
||||
|
||||
it("should close the menu when clicking outside", async () => {
|
||||
const user = userEvent.setup();
|
||||
renderWithProviders(
|
||||
<div>
|
||||
<MultiExportDropdown multiResult={makeMultiResult(true)} />
|
||||
<div data-testid="outside">Outside</div>
|
||||
</div>
|
||||
);
|
||||
|
||||
await user.click(screen.getByRole("button", { name: /^export$/i }));
|
||||
expect(screen.getByText("Export as PDF")).toBeInTheDocument();
|
||||
|
||||
fireEvent.mouseDown(screen.getByTestId("outside"));
|
||||
expect(screen.queryByText("Export as PDF")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,274 @@
|
|||
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
|
||||
import { exportMultiToPDF, exportMultiToCSV } from "./multi_export_utils";
|
||||
import type { MultiModelResult } from "./types";
|
||||
import type { CostEstimateResponse } from "../types";
|
||||
|
||||
vi.mock("@/utils/dataUtils", () => ({
|
||||
formatNumberWithCommas: vi.fn((v: number, d: number = 0) =>
|
||||
Number.isFinite(v) ? v.toFixed(d) : "-"
|
||||
),
|
||||
}));
|
||||
|
||||
function makeCostResponse(overrides: Partial<CostEstimateResponse> = {}): CostEstimateResponse {
|
||||
return {
|
||||
model: "gpt-4",
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
num_requests_per_day: 100,
|
||||
num_requests_per_month: 3000,
|
||||
cost_per_request: 0.05,
|
||||
input_cost_per_request: 0.03,
|
||||
output_cost_per_request: 0.02,
|
||||
margin_cost_per_request: 0,
|
||||
daily_cost: 5.0,
|
||||
daily_input_cost: 3.0,
|
||||
daily_output_cost: 2.0,
|
||||
daily_margin_cost: 0,
|
||||
monthly_cost: 150.0,
|
||||
monthly_input_cost: 90.0,
|
||||
monthly_output_cost: 60.0,
|
||||
monthly_margin_cost: 0,
|
||||
input_cost_per_token: 0.00003,
|
||||
output_cost_per_token: 0.00004,
|
||||
provider: "openai",
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function makeMultiResult(overrides: Partial<MultiModelResult> = {}): MultiModelResult {
|
||||
return {
|
||||
entries: [
|
||||
{
|
||||
entry: { id: "entry-1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 },
|
||||
result: makeCostResponse(),
|
||||
loading: false,
|
||||
error: null,
|
||||
},
|
||||
],
|
||||
totals: {
|
||||
cost_per_request: 0.05,
|
||||
daily_cost: 5.0,
|
||||
monthly_cost: 150.0,
|
||||
margin_per_request: 0,
|
||||
daily_margin: null,
|
||||
monthly_margin: null,
|
||||
},
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
describe("exportMultiToPDF", () => {
|
||||
let mockPrintWindow: {
|
||||
document: { write: ReturnType<typeof vi.fn>; close: ReturnType<typeof vi.fn> };
|
||||
print: ReturnType<typeof vi.fn>;
|
||||
onload: (() => void) | null;
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
mockPrintWindow = {
|
||||
document: { write: vi.fn(), close: vi.fn() },
|
||||
print: vi.fn(),
|
||||
onload: null,
|
||||
};
|
||||
vi.spyOn(window, "open").mockReturnValue(mockPrintWindow as unknown as Window);
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
});
|
||||
|
||||
it("should open a new popup window", () => {
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
expect(window.open).toHaveBeenCalledWith("", "_blank");
|
||||
});
|
||||
|
||||
it("should write HTML containing the report title", () => {
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
const html = mockPrintWindow.document.write.mock.calls[0][0] as string;
|
||||
expect(html).toContain("LLM Cost Estimate Report");
|
||||
});
|
||||
|
||||
it("should include model name and provider in the generated HTML", () => {
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
const html = mockPrintWindow.document.write.mock.calls[0][0] as string;
|
||||
expect(html).toContain("gpt-4");
|
||||
expect(html).toContain("openai");
|
||||
});
|
||||
|
||||
it("should close the document after writing", () => {
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
expect(mockPrintWindow.document.close).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should call print after the window finishes loading", () => {
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
expect(mockPrintWindow.print).not.toHaveBeenCalled();
|
||||
mockPrintWindow.onload!();
|
||||
expect(mockPrintWindow.print).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should show the margin section when margin per request is greater than zero", () => {
|
||||
const multiResult = makeMultiResult({
|
||||
totals: {
|
||||
cost_per_request: 0.06,
|
||||
daily_cost: 5.0,
|
||||
monthly_cost: 150.0,
|
||||
margin_per_request: 0.01,
|
||||
daily_margin: 1.0,
|
||||
monthly_margin: 30.0,
|
||||
},
|
||||
});
|
||||
exportMultiToPDF(multiResult);
|
||||
const html = mockPrintWindow.document.write.mock.calls[0][0] as string;
|
||||
expect(html).toContain("Margin/Request");
|
||||
});
|
||||
|
||||
it("should not show the margin section when margin per request is zero", () => {
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
const html = mockPrintWindow.document.write.mock.calls[0][0] as string;
|
||||
expect(html).not.toContain("Margin/Request");
|
||||
});
|
||||
|
||||
it("should alert when popup is blocked", () => {
|
||||
vi.spyOn(window, "open").mockReturnValue(null);
|
||||
const alertSpy = vi.spyOn(window, "alert").mockImplementation(() => {});
|
||||
exportMultiToPDF(makeMultiResult());
|
||||
expect(alertSpy).toHaveBeenCalledWith("Please allow popups to export PDF");
|
||||
});
|
||||
|
||||
it("should only include entries that have a result", () => {
|
||||
const multiResult: MultiModelResult = {
|
||||
entries: [
|
||||
{ entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, result: null, loading: false, error: null },
|
||||
{ entry: { id: "e2", model: "claude-3", input_tokens: 500, output_tokens: 250 }, result: makeCostResponse({ model: "claude-3", provider: "anthropic" }), loading: false, error: null },
|
||||
],
|
||||
totals: { cost_per_request: 0.05, daily_cost: 5.0, monthly_cost: 150.0, margin_per_request: 0, daily_margin: null, monthly_margin: null },
|
||||
};
|
||||
exportMultiToPDF(multiResult);
|
||||
const html = mockPrintWindow.document.write.mock.calls[0][0] as string;
|
||||
expect(html).toContain("1 model configured");
|
||||
expect(html).toContain("claude-3");
|
||||
});
|
||||
|
||||
it("should show plural 'models' when multiple results are present", () => {
|
||||
const multiResult: MultiModelResult = {
|
||||
entries: [
|
||||
{ entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, result: makeCostResponse(), loading: false, error: null },
|
||||
{ entry: { id: "e2", model: "claude-3", input_tokens: 500, output_tokens: 250 }, result: makeCostResponse({ model: "claude-3" }), loading: false, error: null },
|
||||
],
|
||||
totals: { cost_per_request: 0.10, daily_cost: 10.0, monthly_cost: 300.0, margin_per_request: 0, daily_margin: null, monthly_margin: null },
|
||||
};
|
||||
exportMultiToPDF(multiResult);
|
||||
const html = mockPrintWindow.document.write.mock.calls[0][0] as string;
|
||||
expect(html).toContain("2 models configured");
|
||||
});
|
||||
});
|
||||
|
||||
describe("exportMultiToCSV", () => {
|
||||
beforeEach(() => {
|
||||
document.body.innerHTML = "";
|
||||
window.URL.createObjectURL = vi.fn(() => "blob:mock-url");
|
||||
window.URL.revokeObjectURL = vi.fn();
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.restoreAllMocks();
|
||||
});
|
||||
|
||||
it("should create an object URL and revoke it after download", () => {
|
||||
exportMultiToCSV(makeMultiResult());
|
||||
expect(window.URL.createObjectURL).toHaveBeenCalledTimes(1);
|
||||
expect(window.URL.revokeObjectURL).toHaveBeenCalledWith("blob:mock-url");
|
||||
});
|
||||
|
||||
it("should set the download filename to include today's date", () => {
|
||||
const createdAnchors: HTMLAnchorElement[] = [];
|
||||
const originalCreate = document.createElement.bind(document);
|
||||
vi.spyOn(document, "createElement").mockImplementation((tag: string) => {
|
||||
const el = originalCreate(tag);
|
||||
if (tag === "a") createdAnchors.push(el as HTMLAnchorElement);
|
||||
return el;
|
||||
});
|
||||
|
||||
const today = new Date().toISOString().split("T")[0];
|
||||
exportMultiToCSV(makeMultiResult());
|
||||
|
||||
expect(createdAnchors[0].download).toBe(`cost_estimate_multi_model_${today}.csv`);
|
||||
});
|
||||
|
||||
it("should generate CSV content containing a header row and model data", () => {
|
||||
let csvContent = "";
|
||||
const OriginalBlob = globalThis.Blob;
|
||||
globalThis.Blob = class extends OriginalBlob {
|
||||
constructor(parts?: BlobPart[], options?: BlobPropertyBag) {
|
||||
super(parts, options);
|
||||
if (typeof parts?.[0] === "string") csvContent = parts[0];
|
||||
}
|
||||
} as unknown as typeof Blob;
|
||||
|
||||
exportMultiToCSV(makeMultiResult());
|
||||
globalThis.Blob = OriginalBlob;
|
||||
|
||||
expect(csvContent).toContain("Model");
|
||||
expect(csvContent).toContain("Cost/Request");
|
||||
expect(csvContent).toContain("gpt-4");
|
||||
expect(csvContent).toContain("openai");
|
||||
});
|
||||
|
||||
it("should include the combined totals section in CSV", () => {
|
||||
let csvContent = "";
|
||||
const OriginalBlob = globalThis.Blob;
|
||||
globalThis.Blob = class extends OriginalBlob {
|
||||
constructor(parts?: BlobPart[], options?: BlobPropertyBag) {
|
||||
super(parts, options);
|
||||
if (typeof parts?.[0] === "string") csvContent = parts[0];
|
||||
}
|
||||
} as unknown as typeof Blob;
|
||||
|
||||
exportMultiToCSV(makeMultiResult());
|
||||
globalThis.Blob = OriginalBlob;
|
||||
|
||||
expect(csvContent).toContain("COMBINED TOTALS");
|
||||
});
|
||||
|
||||
it("should create a blob with the correct CSV mime type", () => {
|
||||
let capturedType = "";
|
||||
const OriginalBlob = globalThis.Blob;
|
||||
globalThis.Blob = class extends OriginalBlob {
|
||||
constructor(parts?: BlobPart[], options?: BlobPropertyBag) {
|
||||
super(parts, options);
|
||||
if (options?.type) capturedType = options.type;
|
||||
}
|
||||
} as unknown as typeof Blob;
|
||||
|
||||
exportMultiToCSV(makeMultiResult());
|
||||
globalThis.Blob = OriginalBlob;
|
||||
|
||||
expect(capturedType).toBe("text/csv;charset=utf-8;");
|
||||
});
|
||||
|
||||
it("should skip entries with null results", () => {
|
||||
const multiResult: MultiModelResult = {
|
||||
entries: [
|
||||
{ entry: { id: "e1", model: "gpt-4", input_tokens: 1000, output_tokens: 500 }, result: null, loading: false, error: null },
|
||||
],
|
||||
totals: { cost_per_request: 0, daily_cost: null, monthly_cost: null, margin_per_request: 0, daily_margin: null, monthly_margin: null },
|
||||
};
|
||||
|
||||
let csvContent = "";
|
||||
const OriginalBlob = globalThis.Blob;
|
||||
globalThis.Blob = class extends OriginalBlob {
|
||||
constructor(parts?: BlobPart[], options?: BlobPropertyBag) {
|
||||
super(parts, options);
|
||||
if (typeof parts?.[0] === "string") csvContent = parts[0];
|
||||
}
|
||||
} as unknown as typeof Blob;
|
||||
|
||||
exportMultiToCSV(multiResult);
|
||||
globalThis.Blob = OriginalBlob;
|
||||
|
||||
// CSV should have metadata rows but no model data row for gpt-4
|
||||
const lines = csvContent.split("\n").filter((l) => l.includes('"gpt-4"'));
|
||||
expect(lines).toHaveLength(0);
|
||||
});
|
||||
});
|
||||
|
|
@ -0,0 +1,342 @@
|
|||
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
|
||||
import { renderHook, act } from "@testing-library/react";
|
||||
import { useMultiCostEstimate } from "./use_multi_cost_estimate";
|
||||
import type { ModelEntry } from "./types";
|
||||
import type { CostEstimateResponse } from "../types";
|
||||
|
||||
vi.mock("@/components/networking", () => ({
|
||||
getProxyBaseUrl: vi.fn(() => ""),
|
||||
getGlobalLitellmHeaderName: vi.fn(() => "Authorization"),
|
||||
}));
|
||||
|
||||
function makeEntry(overrides: Partial<ModelEntry> = {}): ModelEntry {
|
||||
return {
|
||||
id: "entry-1",
|
||||
model: "gpt-4",
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
function makeApiResponse(overrides: Partial<CostEstimateResponse> = {}): CostEstimateResponse {
|
||||
return {
|
||||
model: "gpt-4",
|
||||
input_tokens: 1000,
|
||||
output_tokens: 500,
|
||||
num_requests_per_day: null,
|
||||
num_requests_per_month: null,
|
||||
cost_per_request: 0.05,
|
||||
input_cost_per_request: 0.03,
|
||||
output_cost_per_request: 0.02,
|
||||
margin_cost_per_request: 0,
|
||||
daily_cost: null,
|
||||
daily_input_cost: null,
|
||||
daily_output_cost: null,
|
||||
daily_margin_cost: null,
|
||||
monthly_cost: null,
|
||||
monthly_input_cost: null,
|
||||
monthly_output_cost: null,
|
||||
monthly_margin_cost: null,
|
||||
input_cost_per_token: 0.00003,
|
||||
output_cost_per_token: 0.00004,
|
||||
provider: "openai",
|
||||
...overrides,
|
||||
};
|
||||
}
|
||||
|
||||
describe("useMultiCostEstimate", () => {
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
vi.useFakeTimers();
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
vi.useRealTimers();
|
||||
});
|
||||
|
||||
describe("debouncedFetchForEntry", () => {
|
||||
it("should not fetch when access token is null", async () => {
|
||||
const fetchSpy = vi.spyOn(global, "fetch");
|
||||
const { result } = renderHook(() => useMultiCostEstimate(null));
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(makeEntry());
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
expect(fetchSpy).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not fetch when the model field is empty", async () => {
|
||||
const fetchSpy = vi.spyOn(global, "fetch");
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(makeEntry({ model: "" }));
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
expect(fetchSpy).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("should not fetch immediately — only after the debounce delay", async () => {
|
||||
const fetchSpy = vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => makeApiResponse(),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
|
||||
act(() => {
|
||||
result.current.debouncedFetchForEntry(makeEntry());
|
||||
});
|
||||
|
||||
expect(fetchSpy).not.toHaveBeenCalled();
|
||||
|
||||
await act(async () => {
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
expect(fetchSpy).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should cancel an in-flight debounce when called again for the same entry", async () => {
|
||||
const fetchSpy = vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => makeApiResponse(),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(makeEntry());
|
||||
vi.advanceTimersByTime(200);
|
||||
result.current.debouncedFetchForEntry(makeEntry());
|
||||
vi.advanceTimersByTime(200);
|
||||
result.current.debouncedFetchForEntry(makeEntry());
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
expect(fetchSpy).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should store the API result after a successful fetch", async () => {
|
||||
vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => makeApiResponse(),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry]);
|
||||
expect(multiResult.entries[0].result).not.toBeNull();
|
||||
expect(multiResult.entries[0].result?.cost_per_request).toBe(0.05);
|
||||
expect(multiResult.entries[0].loading).toBe(false);
|
||||
expect(multiResult.entries[0].error).toBeNull();
|
||||
});
|
||||
|
||||
it("should set an error message when the API returns a non-ok response", async () => {
|
||||
vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: false,
|
||||
json: async () => ({ detail: { error: "Model not found" } }),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry]);
|
||||
expect(multiResult.entries[0].result).toBeNull();
|
||||
expect(multiResult.entries[0].error).toBe("Model not found");
|
||||
});
|
||||
|
||||
it("should fall back to detail string when error has no nested error field", async () => {
|
||||
vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: false,
|
||||
json: async () => ({ detail: "Bad request" }),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry]);
|
||||
expect(multiResult.entries[0].error).toBe("Bad request");
|
||||
});
|
||||
|
||||
it("should set 'Network error' when fetch throws", async () => {
|
||||
vi.spyOn(global, "fetch").mockRejectedValue(new Error("connection refused"));
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry]);
|
||||
expect(multiResult.entries[0].error).toBe("Network error");
|
||||
expect(multiResult.entries[0].result).toBeNull();
|
||||
});
|
||||
});
|
||||
|
||||
describe("removeEntry", () => {
|
||||
it("should remove an entry's cached result", async () => {
|
||||
vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => makeApiResponse(),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
// Confirm result was stored
|
||||
expect(result.current.getMultiModelResult([entry]).entries[0].result).not.toBeNull();
|
||||
|
||||
act(() => {
|
||||
result.current.removeEntry(entry.id);
|
||||
});
|
||||
|
||||
// After removal, the entry should return as if it never fetched
|
||||
const multiResult = result.current.getMultiModelResult([entry]);
|
||||
expect(multiResult.entries[0].result).toBeNull();
|
||||
});
|
||||
|
||||
it("should cancel a pending debounce for the removed entry", async () => {
|
||||
const fetchSpy = vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => makeApiResponse(),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
act(() => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
result.current.removeEntry(entry.id);
|
||||
});
|
||||
|
||||
await act(async () => {
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
expect(fetchSpy).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
describe("getMultiModelResult", () => {
|
||||
it("should return zero totals when no entries have results", () => {
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const multiResult = result.current.getMultiModelResult([makeEntry()]);
|
||||
|
||||
expect(multiResult.totals.cost_per_request).toBe(0);
|
||||
expect(multiResult.totals.margin_per_request).toBe(0);
|
||||
expect(multiResult.totals.daily_cost).toBeNull();
|
||||
expect(multiResult.totals.monthly_cost).toBeNull();
|
||||
});
|
||||
|
||||
it("should return an empty entries array for an empty input list", () => {
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const multiResult = result.current.getMultiModelResult([]);
|
||||
|
||||
expect(multiResult.entries).toHaveLength(0);
|
||||
expect(multiResult.totals.daily_cost).toBeNull();
|
||||
expect(multiResult.totals.monthly_cost).toBeNull();
|
||||
});
|
||||
|
||||
it("should sum cost_per_request across multiple loaded entries", async () => {
|
||||
const entry1 = makeEntry({ id: "e1", model: "gpt-4" });
|
||||
const entry2 = makeEntry({ id: "e2", model: "claude-3" });
|
||||
|
||||
let callIndex = 0;
|
||||
const responses = [
|
||||
makeApiResponse({ cost_per_request: 0.05, margin_cost_per_request: 0 }),
|
||||
makeApiResponse({ model: "claude-3", cost_per_request: 0.10, margin_cost_per_request: 0 }),
|
||||
];
|
||||
|
||||
vi.spyOn(global, "fetch").mockImplementation(async () => ({
|
||||
ok: true,
|
||||
json: async () => responses[callIndex++],
|
||||
} as Response));
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry1);
|
||||
result.current.debouncedFetchForEntry(entry2);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry1, entry2]);
|
||||
expect(multiResult.totals.cost_per_request).toBeCloseTo(0.15);
|
||||
});
|
||||
|
||||
it("should accumulate daily cost only when entries have a daily cost", async () => {
|
||||
const entry1 = makeEntry({ id: "e1", model: "gpt-4" });
|
||||
const entry2 = makeEntry({ id: "e2", model: "claude-3" });
|
||||
|
||||
let callIndex = 0;
|
||||
const responses = [
|
||||
makeApiResponse({ daily_cost: 5.0, daily_margin_cost: 0, monthly_cost: null, monthly_margin_cost: null }),
|
||||
makeApiResponse({ model: "claude-3", daily_cost: 10.0, daily_margin_cost: 0, monthly_cost: null, monthly_margin_cost: null }),
|
||||
];
|
||||
|
||||
vi.spyOn(global, "fetch").mockImplementation(async () => ({
|
||||
ok: true,
|
||||
json: async () => responses[callIndex++],
|
||||
} as Response));
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry1);
|
||||
result.current.debouncedFetchForEntry(entry2);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry1, entry2]);
|
||||
expect(multiResult.totals.daily_cost).toBeCloseTo(15.0);
|
||||
expect(multiResult.totals.monthly_cost).toBeNull();
|
||||
});
|
||||
|
||||
it("should mark each entry's loading and error state from cached data", async () => {
|
||||
vi.spyOn(global, "fetch").mockResolvedValue({
|
||||
ok: false,
|
||||
json: async () => ({ detail: "Not found" }),
|
||||
} as Response);
|
||||
|
||||
const { result } = renderHook(() => useMultiCostEstimate("token123"));
|
||||
const entry = makeEntry();
|
||||
|
||||
await act(async () => {
|
||||
result.current.debouncedFetchForEntry(entry);
|
||||
await vi.runAllTimersAsync();
|
||||
});
|
||||
|
||||
const multiResult = result.current.getMultiModelResult([entry]);
|
||||
expect(multiResult.entries[0].error).toBe("Not found");
|
||||
expect(multiResult.entries[0].loading).toBe(false);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -17,6 +17,7 @@ interface UsageExportHeaderProps {
|
|||
selectedFilters?: string[];
|
||||
onFiltersChange?: (filters: string[]) => void;
|
||||
filterOptions?: Array<{ label: string; value: string }>;
|
||||
filterMode?: "multiple" | "single";
|
||||
customTitle?: string;
|
||||
compactLayout?: boolean;
|
||||
teams?: Team[];
|
||||
|
|
@ -32,6 +33,7 @@ const UsageExportHeader: React.FC<UsageExportHeaderProps> = ({
|
|||
selectedFilters = [],
|
||||
onFiltersChange,
|
||||
filterOptions = [],
|
||||
filterMode = "multiple",
|
||||
customTitle,
|
||||
compactLayout = false,
|
||||
teams = [],
|
||||
|
|
@ -59,11 +61,17 @@ const UsageExportHeader: React.FC<UsageExportHeaderProps> = ({
|
|||
<div>
|
||||
{filterLabel && <Text className="mb-2">{filterLabel}</Text>}
|
||||
<Select
|
||||
mode="multiple"
|
||||
mode={filterMode === "single" ? undefined : "multiple"}
|
||||
style={{ width: "100%" }}
|
||||
placeholder={filterPlaceholder}
|
||||
value={selectedFilters}
|
||||
onChange={onFiltersChange}
|
||||
value={filterMode === "single" ? (selectedFilters[0] ?? undefined) : selectedFilters}
|
||||
onChange={(value: any) => {
|
||||
if (filterMode === "single") {
|
||||
onFiltersChange?.(value ? [value] : []);
|
||||
} else {
|
||||
onFiltersChange?.(value);
|
||||
}
|
||||
}}
|
||||
options={filterOptions}
|
||||
allowClear
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import type { Team } from "@/components/key_team_helpers/key_list";
|
|||
|
||||
export type ExportFormat = "csv" | "json";
|
||||
export type ExportScope = "daily" | "daily_with_keys" | "daily_with_models";
|
||||
export type EntityType = "tag" | "team" | "organization" | "customer" | "agent";
|
||||
export type EntityType = "tag" | "team" | "organization" | "customer" | "agent" | "user";
|
||||
|
||||
export interface EntitySpendData {
|
||||
results: any[];
|
||||
|
|
|
|||
|
|
@ -2,14 +2,7 @@
|
|||
|
||||
import React, { useCallback, useDeferredValue, useEffect, useState } from "react";
|
||||
import { Select, Switch, Tooltip } from "antd";
|
||||
import {
|
||||
Table,
|
||||
TableHead,
|
||||
TableHeaderCell,
|
||||
TableBody,
|
||||
TableRow,
|
||||
TableCell,
|
||||
} from "@tremor/react";
|
||||
import { Table, TableHead, TableHeaderCell, TableBody, TableRow, TableCell } from "@tremor/react";
|
||||
import { TimeCell } from "./view_logs/time_cell";
|
||||
import { TableHeaderSortDropdown } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
|
||||
import type { SortState } from "./common_components/TableHeaderSortDropdown/TableHeaderSortDropdown";
|
||||
|
|
@ -17,14 +10,13 @@ import FilterComponent, { FilterOption } from "./molecules/filter";
|
|||
import { fetchToolsList, updateToolPolicy, ToolRow } from "./networking";
|
||||
|
||||
const POLICY_OPTIONS = [
|
||||
{ value: "trusted", label: "trusted", color: "#065f46", bg: "#d1fae5", border: "#6ee7b7" },
|
||||
{ value: "blocked", label: "blocked", color: "#991b1b", bg: "#fee2e2", border: "#fca5a5" },
|
||||
{ value: "trusted", label: "trusted", color: "#065f46", bg: "#d1fae5", border: "#6ee7b7" },
|
||||
{ value: "blocked", label: "blocked", color: "#991b1b", bg: "#fee2e2", border: "#fca5a5" },
|
||||
] as const;
|
||||
|
||||
type PolicyValue = "trusted" | "blocked";
|
||||
|
||||
const policyStyle = (p: string) =>
|
||||
POLICY_OPTIONS.find((o) => o.value === p) ?? POLICY_OPTIONS[1];
|
||||
const policyStyle = (p: string) => POLICY_OPTIONS.find((o) => o.value === p) ?? POLICY_OPTIONS[1];
|
||||
|
||||
type SortField = "tool_name" | "call_policy" | "team_id" | "key_alias" | "created_at" | "call_count";
|
||||
|
||||
|
|
@ -56,18 +48,6 @@ const PolicySelect: React.FC<{
|
|||
minWidth: 110,
|
||||
fontWeight: 500,
|
||||
}}
|
||||
styles={{
|
||||
selector: {
|
||||
backgroundColor: style.bg,
|
||||
borderColor: style.border,
|
||||
color: style.color,
|
||||
borderRadius: 999,
|
||||
fontSize: 11,
|
||||
fontWeight: 600,
|
||||
paddingLeft: 8,
|
||||
paddingRight: 4,
|
||||
},
|
||||
} as Record<string, React.CSSProperties>}
|
||||
popupMatchSelectWidth={false}
|
||||
options={POLICY_OPTIONS.map((o) => ({
|
||||
value: o.value,
|
||||
|
|
@ -133,7 +113,9 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
}
|
||||
}, [accessToken]);
|
||||
|
||||
useEffect(() => { load(); }, [load]);
|
||||
useEffect(() => {
|
||||
load();
|
||||
}, [load]);
|
||||
|
||||
useEffect(() => {
|
||||
if (!isLiveTail) return;
|
||||
|
|
@ -146,9 +128,7 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
setSaving(toolName);
|
||||
try {
|
||||
await updateToolPolicy(accessToken, toolName, newPolicy);
|
||||
setTools((prev) =>
|
||||
prev.map((t) => (t.tool_name === toolName ? { ...t, call_policy: newPolicy } : t))
|
||||
);
|
||||
setTools((prev) => prev.map((t) => (t.tool_name === toolName ? { ...t, call_policy: newPolicy } : t)));
|
||||
} catch (e: any) {
|
||||
alert(`Failed to update policy: ${e.message}`);
|
||||
} finally {
|
||||
|
|
@ -178,12 +158,14 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
};
|
||||
|
||||
// Build unique team/key options from loaded data
|
||||
const teamOptions = Array.from(new Set(tools.map((t) => t.team_id).filter(Boolean))).map(
|
||||
(v) => ({ label: v as string, value: v as string })
|
||||
);
|
||||
const keyAliasOptions = Array.from(new Set(tools.map((t) => t.key_alias).filter(Boolean))).map(
|
||||
(v) => ({ label: v as string, value: v as string })
|
||||
);
|
||||
const teamOptions = Array.from(new Set(tools.map((t) => t.team_id).filter(Boolean))).map((v) => ({
|
||||
label: v as string,
|
||||
value: v as string,
|
||||
}));
|
||||
const keyAliasOptions = Array.from(new Set(tools.map((t) => t.key_alias).filter(Boolean))).map((v) => ({
|
||||
label: v as string,
|
||||
value: v as string,
|
||||
}));
|
||||
|
||||
const filterOptions: FilterOption[] = [
|
||||
{
|
||||
|
|
@ -245,7 +227,6 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
<div className="p-6 w-full">
|
||||
<h1 className="text-2xl font-semibold text-gray-900 mb-6">Tool Policies</h1>
|
||||
<div className="bg-white rounded-lg shadow w-full max-w-full box-border">
|
||||
|
||||
{/* Toolbar */}
|
||||
<div className="border-b px-6 py-4 w-full max-w-full box-border">
|
||||
<div className="flex flex-col md:flex-row items-start md:items-center justify-between space-y-4 md:space-y-0 w-full max-w-full box-border">
|
||||
|
|
@ -256,10 +237,23 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
placeholder="Search by Tool Name"
|
||||
className="w-full px-3 py-2 pl-8 border rounded-md text-sm focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-blue-500"
|
||||
value={searchTerm}
|
||||
onChange={(e) => { setSearchTerm(e.target.value); setCurrentPage(1); }}
|
||||
onChange={(e) => {
|
||||
setSearchTerm(e.target.value);
|
||||
setCurrentPage(1);
|
||||
}}
|
||||
/>
|
||||
<svg className="absolute left-2.5 top-2.5 h-4 w-4 text-gray-500" fill="none" stroke="currentColor" viewBox="0 0 24 24">
|
||||
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M21 21l-6-6m2-5a7 7 0 11-14 0 7 7 0 0114 0z" />
|
||||
<svg
|
||||
className="absolute left-2.5 top-2.5 h-4 w-4 text-gray-500"
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
d="M21 21l-6-6m2-5a7 7 0 11-14 0 7 7 0 0114 0z"
|
||||
/>
|
||||
</svg>
|
||||
</div>
|
||||
|
||||
|
|
@ -273,8 +267,18 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
disabled={isButtonLoading}
|
||||
className="flex items-center gap-1.5 px-3 py-2 text-sm border rounded-md hover:bg-gray-50 disabled:opacity-60"
|
||||
>
|
||||
<svg className={`w-4 h-4 ${isButtonLoading ? "animate-spin" : ""}`} fill="none" stroke="currentColor" viewBox="0 0 24 24">
|
||||
<path strokeLinecap="round" strokeLinejoin="round" strokeWidth={2} d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15" />
|
||||
<svg
|
||||
className={`w-4 h-4 ${isButtonLoading ? "animate-spin" : ""}`}
|
||||
fill="none"
|
||||
stroke="currentColor"
|
||||
viewBox="0 0 24 24"
|
||||
>
|
||||
<path
|
||||
strokeLinecap="round"
|
||||
strokeLinejoin="round"
|
||||
strokeWidth={2}
|
||||
d="M4 4v5h.582m15.356 2A8.001 8.001 0 004.582 9m0 0H9m11 11v-5h-.581m0 0a8.003 8.003 0 01-15.357-2m15.357 2H15"
|
||||
/>
|
||||
</svg>
|
||||
{isButtonLoading ? "Fetching" : "Fetch"}
|
||||
</button>
|
||||
|
|
@ -282,14 +286,27 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
|
||||
<div className="flex items-center gap-4 text-sm text-gray-600 whitespace-nowrap">
|
||||
<span>
|
||||
Showing {filtered.length === 0 ? 0 : (currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, filtered.length)} of {filtered.length} results
|
||||
Showing {filtered.length === 0 ? 0 : (currentPage - 1) * pageSize + 1} -{" "}
|
||||
{Math.min(currentPage * pageSize, filtered.length)} of {filtered.length} results
|
||||
</span>
|
||||
<span>
|
||||
Page {currentPage} of {totalPages}
|
||||
</span>
|
||||
<span>Page {currentPage} of {totalPages}</span>
|
||||
<div className="flex gap-1">
|
||||
<button onClick={() => setCurrentPage((p) => Math.max(1, p - 1))} disabled={currentPage === 1}
|
||||
className="px-3 py-1.5 border rounded-md text-sm hover:bg-gray-50 disabled:opacity-40">Previous</button>
|
||||
<button onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))} disabled={currentPage === totalPages}
|
||||
className="px-3 py-1.5 border rounded-md text-sm hover:bg-gray-50 disabled:opacity-40">Next</button>
|
||||
<button
|
||||
onClick={() => setCurrentPage((p) => Math.max(1, p - 1))}
|
||||
disabled={currentPage === 1}
|
||||
className="px-3 py-1.5 border rounded-md text-sm hover:bg-gray-50 disabled:opacity-40"
|
||||
>
|
||||
Previous
|
||||
</button>
|
||||
<button
|
||||
onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))}
|
||||
disabled={currentPage === totalPages}
|
||||
className="px-3 py-1.5 border rounded-md text-sm hover:bg-gray-50 disabled:opacity-40"
|
||||
>
|
||||
Next
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
|
@ -309,7 +326,9 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
{isLiveTail && (
|
||||
<div className="bg-green-50 border-b border-green-100 px-6 py-2 flex items-center justify-between">
|
||||
<span className="text-sm text-green-700">Auto-refreshing every 15 seconds</span>
|
||||
<button onClick={() => setIsLiveTail(false)} className="text-xs text-green-600 underline">Stop</button>
|
||||
<button onClick={() => setIsLiveTail(false)} className="text-xs text-green-600 underline">
|
||||
Stop
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
|
||||
|
|
@ -321,20 +340,34 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
<Table className="[&_td]:py-0.5 [&_th]:py-1 w-full">
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeaderCell className="py-1 h-8"><SortHeader label="Discovered" field="created_at" /></TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8"><SortHeader label="Tool Name" field="tool_name" /></TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8"><SortHeader label="Policy" field="call_policy" /></TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8"><SortHeader label="# Calls" field="call_count" /></TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8"><SortHeader label="Team Name" field="team_id" /></TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">
|
||||
<SortHeader label="Discovered" field="created_at" />
|
||||
</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">
|
||||
<SortHeader label="Tool Name" field="tool_name" />
|
||||
</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">
|
||||
<SortHeader label="Policy" field="call_policy" />
|
||||
</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">
|
||||
<SortHeader label="# Calls" field="call_count" />
|
||||
</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">
|
||||
<SortHeader label="Team Name" field="team_id" />
|
||||
</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">Key Hash</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8"><SortHeader label="Key Name" field="key_alias" /></TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">
|
||||
<SortHeader label="Key Name" field="key_alias" />
|
||||
</TableHeaderCell>
|
||||
<TableHeaderCell className="py-1 h-8">Origin</TableHeaderCell>
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{loading ? (
|
||||
<TableRow>
|
||||
<TableCell colSpan={8} className="h-8 text-center text-gray-500">Loading tools…</TableCell>
|
||||
<TableCell colSpan={8} className="h-8 text-center text-gray-500">
|
||||
Loading tools…
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
) : paginated.length === 0 ? (
|
||||
<TableRow>
|
||||
|
|
@ -397,12 +430,25 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
|
|||
{/* Bottom pagination (only when > 1 page) */}
|
||||
{totalPages > 1 && (
|
||||
<div className="border-t px-6 py-3 flex items-center justify-between text-sm text-gray-600">
|
||||
<span>Showing {(currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, sorted.length)} of {sorted.length}</span>
|
||||
<span>
|
||||
Showing {(currentPage - 1) * pageSize + 1} - {Math.min(currentPage * pageSize, sorted.length)} of{" "}
|
||||
{sorted.length}
|
||||
</span>
|
||||
<div className="flex gap-1">
|
||||
<button onClick={() => setCurrentPage((p) => Math.max(1, p - 1))} disabled={currentPage === 1}
|
||||
className="px-3 py-1.5 border rounded-md hover:bg-gray-50 disabled:opacity-40">Previous</button>
|
||||
<button onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))} disabled={currentPage === totalPages}
|
||||
className="px-3 py-1.5 border rounded-md hover:bg-gray-50 disabled:opacity-40">Next</button>
|
||||
<button
|
||||
onClick={() => setCurrentPage((p) => Math.max(1, p - 1))}
|
||||
disabled={currentPage === 1}
|
||||
className="px-3 py-1.5 border rounded-md hover:bg-gray-50 disabled:opacity-40"
|
||||
>
|
||||
Previous
|
||||
</button>
|
||||
<button
|
||||
onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))}
|
||||
disabled={currentPage === totalPages}
|
||||
className="px-3 py-1.5 border rounded-md hover:bg-gray-50 disabled:opacity-40"
|
||||
>
|
||||
Next
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ vi.mock("../../../networking", () => ({
|
|||
organizationDailyActivityCall: vi.fn(),
|
||||
customerDailyActivityCall: vi.fn(),
|
||||
agentDailyActivityCall: vi.fn(),
|
||||
userDailyActivityCall: vi.fn(),
|
||||
}));
|
||||
|
||||
// Mock the child components to simplify testing
|
||||
|
|
@ -58,6 +59,7 @@ describe("EntityUsage", () => {
|
|||
const mockOrganizationDailyActivityCall = vi.mocked(networking.organizationDailyActivityCall);
|
||||
const mockCustomerDailyActivityCall = vi.mocked(networking.customerDailyActivityCall);
|
||||
const mockAgentDailyActivityCall = vi.mocked(networking.agentDailyActivityCall);
|
||||
const mockUserDailyActivityCall = vi.mocked(networking.userDailyActivityCall);
|
||||
|
||||
const mockSpendData = {
|
||||
results: [
|
||||
|
|
@ -146,11 +148,13 @@ describe("EntityUsage", () => {
|
|||
mockOrganizationDailyActivityCall.mockClear();
|
||||
mockCustomerDailyActivityCall.mockClear();
|
||||
mockAgentDailyActivityCall.mockClear();
|
||||
mockUserDailyActivityCall.mockClear();
|
||||
mockTagDailyActivityCall.mockResolvedValue(mockSpendData);
|
||||
mockTeamDailyActivityCall.mockResolvedValue(mockSpendData);
|
||||
mockOrganizationDailyActivityCall.mockResolvedValue(mockSpendData);
|
||||
mockCustomerDailyActivityCall.mockResolvedValue(mockSpendData);
|
||||
mockAgentDailyActivityCall.mockResolvedValue(mockSpendData);
|
||||
mockUserDailyActivityCall.mockResolvedValue(mockSpendData);
|
||||
});
|
||||
|
||||
it("should render with tag entity type and display spend metrics", async () => {
|
||||
|
|
@ -232,6 +236,21 @@ describe("EntityUsage", () => {
|
|||
});
|
||||
});
|
||||
|
||||
it("should render with user entity type and call user API", async () => {
|
||||
render(<EntityUsage {...defaultProps} entityType="user" />);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(mockUserDailyActivityCall).toHaveBeenCalled();
|
||||
});
|
||||
|
||||
expect(screen.getByText("User Spend Overview")).toBeInTheDocument();
|
||||
|
||||
await waitFor(() => {
|
||||
const spendElements = screen.getAllByText("$100.50");
|
||||
expect(spendElements.length).toBeGreaterThan(0);
|
||||
});
|
||||
});
|
||||
|
||||
it("should switch between tabs", async () => {
|
||||
render(<EntityUsage {...defaultProps} />);
|
||||
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ import {
|
|||
organizationDailyActivityCall,
|
||||
tagDailyActivityCall,
|
||||
teamDailyActivityCall,
|
||||
userDailyActivityCall,
|
||||
} from "../../../networking";
|
||||
import { getProviderLogoAndName } from "../../../provider_info_helpers";
|
||||
import { BreakdownMetrics, DailyData, EntityMetricWithMetadata, KeyMetricWithMetadata, TagUsage } from "../../types";
|
||||
|
|
@ -156,6 +157,15 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
selectedTags.length > 0 ? selectedTags : null,
|
||||
);
|
||||
setSpendData(data);
|
||||
} else if (entityType === "user") {
|
||||
const data = await userDailyActivityCall(
|
||||
accessToken,
|
||||
startTime,
|
||||
endTime,
|
||||
1,
|
||||
selectedTags.length > 0 ? selectedTags[0] : null,
|
||||
);
|
||||
setSpendData(data);
|
||||
} else {
|
||||
throw new Error("Invalid entity type");
|
||||
}
|
||||
|
|
@ -391,6 +401,7 @@ const EntityUsage: React.FC<EntityUsageProps> = ({ accessToken, entityType, enti
|
|||
selectedFilters={selectedTags}
|
||||
onFiltersChange={setSelectedTags}
|
||||
filterOptions={getAllTags() || undefined}
|
||||
filterMode={entityType === "user" ? "single" : "multiple"}
|
||||
teams={teams || []}
|
||||
/>
|
||||
<TabGroup>
|
||||
|
|
|
|||
|
|
@ -713,7 +713,7 @@ describe("UsagePage", () => {
|
|||
// Admin should see the user selector select element with the placeholder attribute
|
||||
const userSelects = screen.getAllByRole("combobox");
|
||||
const userSelect = userSelects.find(
|
||||
(el) => el.getAttribute("placeholder") === "All Users (Global View)",
|
||||
(el) => el.getAttribute("placeholder") === "Select user to filter...",
|
||||
);
|
||||
expect(userSelect).toBeDefined();
|
||||
});
|
||||
|
|
@ -828,7 +828,7 @@ describe("UsagePage", () => {
|
|||
// Non-admin should not see the user selector
|
||||
const userSelects = screen.getAllByRole("combobox");
|
||||
const userSelect = userSelects.find(
|
||||
(el) => el.getAttribute("placeholder") === "All Users (Global View)",
|
||||
(el) => el.getAttribute("placeholder") === "Select user to filter...",
|
||||
);
|
||||
expect(userSelect).toBeUndefined();
|
||||
});
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@
|
|||
* Works at 1m+ spend logs, by querying an aggregate table instead.
|
||||
*/
|
||||
|
||||
import { InfoCircleOutlined, LoadingOutlined, UserOutlined } from "@ant-design/icons";
|
||||
import { InfoCircleOutlined, LoadingOutlined } from "@ant-design/icons";
|
||||
import {
|
||||
BarChart,
|
||||
Card,
|
||||
|
|
@ -498,6 +498,36 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
{/* Your Usage Panel */}
|
||||
{usageView === "global" && (
|
||||
<>
|
||||
{isAdmin && (
|
||||
<div className="mb-4">
|
||||
<Text className="mb-2">Filter by user</Text>
|
||||
<Select
|
||||
showSearch
|
||||
allowClear
|
||||
style={{ width: "100%" }}
|
||||
placeholder="Select user to filter..."
|
||||
value={selectedUserId}
|
||||
onChange={(value) => setSelectedUserId(value ?? null)}
|
||||
filterOption={false}
|
||||
onSearch={handleUserSearchChange}
|
||||
searchValue={userSearchInput}
|
||||
onPopupScroll={handleUserPopupScroll}
|
||||
loading={isLoadingUsers}
|
||||
notFoundContent={isLoadingUsers ? <LoadingOutlined spin /> : "No users found"}
|
||||
options={userOptions}
|
||||
popupRender={(menu) => (
|
||||
<>
|
||||
{menu}
|
||||
{isFetchingNextUsersPage && (
|
||||
<div style={{ textAlign: "center", padding: 8 }}>
|
||||
<LoadingOutlined spin />
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
<TabGroup>
|
||||
<div className="flex justify-between items-center">
|
||||
<TabList variant="solid" className="mt-1">
|
||||
|
|
@ -560,41 +590,6 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
</>
|
||||
)}
|
||||
</Text>
|
||||
{isAdmin && (
|
||||
<div className="flex items-center gap-2">
|
||||
<UserOutlined style={{ fontSize: "14px", color: "#6b7280" }} />
|
||||
<Select
|
||||
showSearch
|
||||
allowClear
|
||||
style={{ width: 300 }}
|
||||
placeholder="All Users (Global View)"
|
||||
value={selectedUserId}
|
||||
onChange={(value) => setSelectedUserId(value ?? null)}
|
||||
filterOption={false}
|
||||
onSearch={handleUserSearchChange}
|
||||
searchValue={userSearchInput}
|
||||
onPopupScroll={handleUserPopupScroll}
|
||||
loading={isLoadingUsers}
|
||||
notFoundContent={isLoadingUsers ? <LoadingOutlined spin /> : "No users found"}
|
||||
options={userOptions}
|
||||
popupRender={(menu) => (
|
||||
<>
|
||||
{menu}
|
||||
{isFetchingNextUsersPage && (
|
||||
<div style={{ textAlign: "center", padding: 8 }}>
|
||||
<LoadingOutlined spin />
|
||||
</div>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
/>
|
||||
{selectedUserId && (
|
||||
<span className="text-xs text-gray-500">
|
||||
Filtering by user
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<ViewUserSpend
|
||||
|
|
@ -912,6 +907,18 @@ const UsagePage: React.FC<UsagePageProps> = ({ teams, organizations }) => {
|
|||
dateValue={dateValue}
|
||||
/>
|
||||
)}
|
||||
{/* User Usage Panel */}
|
||||
{usageView === "user" && (
|
||||
<EntityUsage
|
||||
accessToken={accessToken}
|
||||
entityType="user"
|
||||
userID={userID}
|
||||
userRole={userRole}
|
||||
entityList={userOptions.length > 0 ? userOptions : null}
|
||||
premiumUser={premiumUser}
|
||||
dateValue={dateValue}
|
||||
/>
|
||||
)}
|
||||
{/* User Agent Activity Panel */}
|
||||
{usageView === "user-agent-activity" && (
|
||||
<UserAgentActivity accessToken={accessToken} userRole={userRole} dateValue={dateValue} />
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ vi.mock("@ant-design/icons", async () => {
|
|||
ShoppingCartOutlined: Icon,
|
||||
TagsOutlined: Icon,
|
||||
RobotOutlined: Icon,
|
||||
UserOutlined: Icon,
|
||||
LineChartOutlined: Icon,
|
||||
BarChartOutlined: Icon,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -7,10 +7,11 @@ import {
|
|||
ShoppingCartOutlined,
|
||||
TagsOutlined,
|
||||
TeamOutlined,
|
||||
UserOutlined,
|
||||
} from "@ant-design/icons";
|
||||
import { Badge, Select } from "antd";
|
||||
import React from "react";
|
||||
export type UsageOption = "global" | "organization" | "team" | "customer" | "tag" | "agent" | "user-agent-activity";
|
||||
export type UsageOption = "global" | "organization" | "team" | "customer" | "tag" | "agent" | "user" | "user-agent-activity";
|
||||
export interface UsageViewSelectProps {
|
||||
value: UsageOption;
|
||||
onChange: (value: UsageOption) => void;
|
||||
|
|
@ -79,6 +80,13 @@ const OPTIONS: OptionConfig[] = [
|
|||
icon: <RobotOutlined style={{ fontSize: "16px" }} />,
|
||||
adminOnly: true,
|
||||
},
|
||||
{
|
||||
value: "user",
|
||||
label: "User Usage",
|
||||
description: "View usage by individual users",
|
||||
icon: <UserOutlined style={{ fontSize: "16px" }} />,
|
||||
adminOnly: true,
|
||||
},
|
||||
{
|
||||
value: "user-agent-activity",
|
||||
label: "User Agent Activity",
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ vi.mock("../networking", () => ({
|
|||
soft_budget: null,
|
||||
}),
|
||||
fetchMCPAccessGroups: vi.fn().mockResolvedValue([]),
|
||||
getAgentsList: vi.fn().mockResolvedValue([]),
|
||||
}));
|
||||
|
||||
vi.mock("../molecules/notifications_manager", () => ({
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import React, { useCallback, useEffect, useState } from "react";
|
|||
import { rolesWithWriteAccess } from "../../utils/roles";
|
||||
import AgentSelector from "../agent_management/AgentSelector";
|
||||
import { mapDisplayToInternalNames } from "../callback_info_helpers";
|
||||
import AccessGroupSelector from "../common_components/AccessGroupSelector";
|
||||
import BudgetDurationDropdown from "../common_components/budget_duration_dropdown";
|
||||
import SchemaFormFields from "../common_components/check_openapi_schema";
|
||||
import KeyLifecycleSettings from "../common_components/KeyLifecycleSettings";
|
||||
|
|
@ -19,7 +20,6 @@ import PassThroughRoutesSelector from "../common_components/PassThroughRoutesSel
|
|||
import PremiumLoggingSettings from "../common_components/PremiumLoggingSettings";
|
||||
import RateLimitTypeFormItem from "../common_components/RateLimitTypeFormItem";
|
||||
import RouterSettingsAccordion, { RouterSettingsAccordionValue } from "../common_components/RouterSettingsAccordion";
|
||||
import AccessGroupSelector from "../common_components/AccessGroupSelector";
|
||||
import TeamDropdown from "../common_components/team_dropdown";
|
||||
import { CreateUserButton } from "../CreateUserButton";
|
||||
import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key";
|
||||
|
|
@ -39,10 +39,10 @@ import {
|
|||
proxyBaseUrl,
|
||||
userFilterUICall,
|
||||
} from "../networking";
|
||||
import CreatedKeyDisplay from "../shared/CreatedKeyDisplay";
|
||||
import NumericalInput from "../shared/numerical_input";
|
||||
import VectorStoreSelector from "../vector_store_management/VectorStoreSelector";
|
||||
import { simplifyKeyGenerateError } from "./utils";
|
||||
import CreatedKeyDisplay from "../shared/CreatedKeyDisplay";
|
||||
|
||||
const { Option } = Select;
|
||||
|
||||
|
|
@ -298,7 +298,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
|
|||
formValues.user_id = userID;
|
||||
} else if (keyOwner === "agent") {
|
||||
if (!selectedAgentId) {
|
||||
message.error("Please select an agent");
|
||||
NotificationsManager.error("Please select an agent");
|
||||
return;
|
||||
}
|
||||
formValues.agent_id = selectedAgentId;
|
||||
|
|
@ -558,7 +558,9 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
|
|||
<Radio value="you">You</Radio>
|
||||
<Radio value="service_account">Service Account</Radio>
|
||||
{userRole === "Admin" && <Radio value="another_user">Another User</Radio>}
|
||||
<Radio value="agent">Agent <Tag color="purple">New</Tag></Radio>
|
||||
<Radio value="agent">
|
||||
Agent <Tag color="purple">New</Tag>
|
||||
</Radio>
|
||||
</Radio.Group>
|
||||
</Form.Item>
|
||||
|
||||
|
|
@ -1004,9 +1006,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
|
|||
style={{ width: "100%" }}
|
||||
disabled={!premiumUser}
|
||||
placeholder={
|
||||
!premiumUser
|
||||
? "Premium feature - Upgrade to set policies by key"
|
||||
: "Select or enter policies"
|
||||
!premiumUser ? "Premium feature - Upgrade to set policies by key" : "Select or enter policies"
|
||||
}
|
||||
options={policiesList.map((name) => ({ value: name, label: name }))}
|
||||
/>
|
||||
|
|
@ -1058,9 +1058,7 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
|
|||
className="mt-4"
|
||||
help="Select access groups to assign to this key"
|
||||
>
|
||||
<AccessGroupSelector
|
||||
placeholder="Select access groups (optional)"
|
||||
/>
|
||||
<AccessGroupSelector placeholder="Select access groups (optional)" />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
label={
|
||||
|
|
@ -1296,7 +1294,11 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey }) => {
|
|||
accessToken={accessToken || ""}
|
||||
value={routerSettings || undefined}
|
||||
onChange={setRouterSettings}
|
||||
modelData={userModels.length > 0 ? { data: userModels.map((model) => ({ model_name: model })) } : undefined}
|
||||
modelData={
|
||||
userModels.length > 0
|
||||
? { data: userModels.map((model) => ({ model_name: model })) }
|
||||
: undefined
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
</AccordionBody>
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ export const pageDescriptions: Record<string, string> = {
|
|||
guardrails: "Set up content moderation and safety guardrails",
|
||||
policies: "Define access control and usage policies",
|
||||
"search-tools": "Configure RAG search and retrieval tools",
|
||||
"tool-policies": "Configure tool use policies and permissions",
|
||||
"vector-stores": "Manage vector databases for embeddings",
|
||||
new_usage: "View usage analytics and metrics",
|
||||
logs: "Access request and response logs",
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@
|
|||
"moduleResolution": "bundler",
|
||||
"resolveJsonModule": true,
|
||||
"isolatedModules": true,
|
||||
"jsx": "preserve",
|
||||
"jsx": "react-jsx",
|
||||
"incremental": true,
|
||||
"plugins": [
|
||||
{
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue