Merge branch 'BerriAI:main' into main

This commit is contained in:
Esteban Zeller 2026-02-25 18:35:11 -03:00 • committed by GitHub
commit 7a85ee8c0a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
72 changed files with 4414 additions and 478 deletions

View file

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

View file

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

View 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

View file

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

Binary file not shown.

View file

@ -0,0 +1,3 @@
-- AlterTable
ALTER TABLE "LiteLLM_SpendLogs" ADD COLUMN "request_duration_ms" INTEGER;

View file

@ -0,0 +1,40 @@
-- AlterTable
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN "object_permission_id" TEXT;
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" DROP COLUMN "spec_path";
-- AlterTable
ALTER TABLE "LiteLLM_VerificationToken" ADD COLUMN "agent_id" TEXT;
-- CreateTable
CREATE TABLE "LiteLLM_ToolTable" (
"tool_id" TEXT NOT NULL,
"tool_name" TEXT NOT NULL,
"origin" TEXT,
"call_policy" TEXT NOT NULL DEFAULT 'untrusted',
"call_count" INTEGER NOT NULL DEFAULT 0,
"assignments" JSONB DEFAULT '{}',
"key_hash" TEXT,
"team_id" TEXT,
"key_alias" TEXT,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"created_by" TEXT,
"updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_by" TEXT,
CONSTRAINT "LiteLLM_ToolTable_pkey" PRIMARY KEY ("tool_id")
);
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_ToolTable_tool_name_key" ON "LiteLLM_ToolTable"("tool_name");
-- CreateIndex
CREATE INDEX "LiteLLM_ToolTable_call_policy_idx" ON "LiteLLM_ToolTable"("call_policy");
-- CreateIndex
CREATE INDEX "LiteLLM_ToolTable_team_id_idx" ON "LiteLLM_ToolTable"("team_id");
-- AddForeignKey
ALTER TABLE "LiteLLM_AgentsTable" ADD CONSTRAINT "LiteLLM_AgentsTable_object_permission_id_fkey" FOREIGN KEY ("object_permission_id") REFERENCES "LiteLLM_ObjectPermissionTable"("object_permission_id") ON DELETE SET NULL ON UPDATE CASCADE;

View file

@ -64,6 +64,8 @@ model LiteLLM_AgentsTable {
litellm_params Json?
agent_card_params Json
agent_access_groups String[] @default([])
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@ -264,6 +266,7 @@ model LiteLLM_ObjectPermissionTable {
organizations LiteLLM_OrganizationTable[]
users LiteLLM_UserTable[]
end_users LiteLLM_EndUserTable[]
agents_table LiteLLM_AgentsTable[]
}
// Holds the MCP server configuration
@ -273,7 +276,6 @@ model LiteLLM_MCPServerTable {
alias String?
description String?
url String?
spec_path String?
transport String @default("sse")
auth_type String?
credentials Json? @default("{}")
@ -315,6 +317,7 @@ model LiteLLM_VerificationToken {
router_settings Json? @default("{}")
user_id String?
team_id String?
agent_id String?
project_id String?
permissions Json @default("{}")
max_parallel_requests Int?
@ -477,6 +480,7 @@ model LiteLLM_SpendLogs {
completion_tokens Int @default(0)
startTime DateTime // Assuming start_time is a DateTime field
endTime DateTime // Assuming end_time is a DateTime field
request_duration_ms Int?
completionStartTime DateTime? // Assuming completionStartTime is a DateTime field
model String @default("")
model_id String? @default("") // the model id stored in proxy model db
@ -1052,6 +1056,26 @@ model LiteLLM_PolicyAttachmentTable {
updated_by String?
}
// Global tool registry - auto-discovered from LLM responses; admins set call_policy here
model LiteLLM_ToolTable {
tool_id String @id @default(uuid())
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
origin String? // MCP server name or "user_defined"
call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked"
call_count Int @default(0) // cumulative number of times this tool was seen
assignments Json? @default("{}")
key_hash String? // hash of the virtual key that first called this tool
team_id String? // team that first called this tool
key_alias String? // human-readable alias of the virtual key
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
@@index([call_policy])
@@index([team_id])
}
//Unified Access Groups table for storing unified access groups
model LiteLLM_AccessGroupTable {
access_group_id String @id @default(uuid())

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm-proxy-extras"
version = "0.4.47"
version = "0.4.48"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
authors = ["BerriAI"]
readme = "README.md"
@ -22,7 +22,7 @@ requires = ["poetry-core"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "0.4.47"
version = "0.4.48"
version_files = [
"pyproject.toml:version",
"../requirements.txt:litellm-proxy-extras==",

View file

@ -22,7 +22,11 @@ from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_REDIS_MAJOR_VERSION
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.caching import (
RedisPipelineIncrementOperation,
RedisPipelineLpopOperation,
RedisPipelineRpushOperation,
)
from litellm.types.services import ServiceTypes
from .base_cache import BaseCache
@ -1320,6 +1324,75 @@ class RedisCache(BaseCache):
)
raise e
async def _pipeline_rpush_helper(
self,
pipe: pipeline,
rpush_list: List[RedisPipelineRpushOperation],
) -> List[int]:
"""Helper function for pipeline rpush operations"""
for rpush_op in rpush_list:
pipe.rpush(rpush_op["key"], *rpush_op["values"])
results = await pipe.execute()
# Preserve positional correspondence — raise on per-command errors
for r in results:
if isinstance(r, Exception):
raise r
return results
async def async_rpush_pipeline(
self,
rpush_list: List[RedisPipelineRpushOperation],
) -> List[int]:
"""
Use Redis Pipelines for bulk RPUSH operations
Args:
rpush_list: List of RedisPipelineRpushOperation dicts containing:
- key: str
- values: List[Any]
Returns:
List[int]: List lengths after each push
"""
if len(rpush_list) == 0:
return []
_redis_client: Any = self.init_async_client()
start_time = time.time()
try:
async with _redis_client.pipeline(transaction=False) as pipe:
results = await self._pipeline_rpush_helper(pipe, rpush_list)
## LOGGING ##
end_time = time.time()
_duration = end_time - start_time
asyncio.create_task(
self.service_logger_obj.async_service_success_hook(
service=ServiceTypes.REDIS,
duration=_duration,
call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}",
)
)
return results
except Exception as e:
## LOGGING ##
end_time = time.time()
_duration = end_time - start_time
asyncio.create_task(
self.service_logger_obj.async_service_failure_hook(
service=ServiceTypes.REDIS,
duration=_duration,
error=e,
call_type=f"async_rpush_pipeline <- {_get_call_stack_info()}",
)
)
verbose_logger.error(
"LiteLLM Redis Caching: async_rpush_pipeline() - Got exception from REDIS %s",
str(e),
)
raise e
async def handle_lpop_count_for_older_redis_versions(
self, pipe: pipeline, key: str, count: int
) -> List[bytes]:
@ -1400,3 +1473,120 @@ class RedisCache(BaseCache):
f"LiteLLM Redis Cache LPOP: - Got exception from REDIS : {str(e)}"
)
raise e
async def _pipeline_lpop_helper(
self,
pipe: pipeline,
lpop_list: List[RedisPipelineLpopOperation],
) -> List[Optional[List[str]]]:
"""Helper function for pipeline lpop operations.
For Redis >= 7, queues one LPOP(key, count) per operation.
For Redis < 7, queues `count` individual LPOP(key) commands per operation.
"""
major_version = self._parse_redis_major_version()
if major_version >= 7:
for lpop_op in lpop_list:
pipe.lpop(lpop_op["key"], lpop_op["count"])
raw_results = await pipe.execute()
else:
# For Redis < 7, LPOP doesn't support count param.
# Issue `count` individual LPOP commands per key, all in one pipeline.
counts: List[int] = []
for lpop_op in lpop_list:
count = lpop_op["count"] or 1
counts.append(count)
for _ in range(count):
pipe.lpop(lpop_op["key"])
flat_results = await pipe.execute()
# Re-group the flat results back into per-key lists
raw_results = []
offset = 0
for count in counts:
key_results = [
r for r in flat_results[offset : offset + count] if r is not None
]
raw_results.append(key_results if key_results else None)
offset += count
# Raise on per-command errors (matches _pipeline_rpush_helper behavior)
for r in raw_results:
if isinstance(r, Exception):
raise r
# Decode bytes -> str for each result set
decoded_results: List[Optional[List[str]]] = []
for r in raw_results:
if r is None:
decoded_results.append(None)
elif isinstance(r, list):
try:
decoded_results.append(
[
item.decode("utf-8") if isinstance(item, bytes) else item
for item in r
if item is not None
]
or None
)
except Exception:
decoded_results.append(r) # type: ignore
else:
decoded_results.append(None)
return decoded_results
async def async_lpop_pipeline(
self,
lpop_list: List[RedisPipelineLpopOperation],
) -> List[Optional[List[str]]]:
"""
Use Redis Pipelines for bulk LPOP operations
Args:
lpop_list: List of RedisPipelineLpopOperation dicts containing:
- key: str
- count: Optional[int]
Returns:
List[Optional[List[str]]]: Decoded results per key, None if key was empty
"""
if len(lpop_list) == 0:
return []
_redis_client: Any = self.init_async_client()
start_time = time.time()
try:
async with _redis_client.pipeline(transaction=False) as pipe:
results = await self._pipeline_lpop_helper(pipe, lpop_list)
## LOGGING ##
end_time = time.time()
_duration = end_time - start_time
asyncio.create_task(
self.service_logger_obj.async_service_success_hook(
service=ServiceTypes.REDIS,
duration=_duration,
call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}",
)
)
return results
except Exception as e:
## LOGGING ##
end_time = time.time()
_duration = end_time - start_time
asyncio.create_task(
self.service_logger_obj.async_service_failure_hook(
service=ServiceTypes.REDIS,
duration=_duration,
error=e,
call_type=f"async_lpop_pipeline <- {_get_call_stack_info()}",
)
)
verbose_logger.error(
"LiteLLM Redis Caching: async_lpop_pipeline() - Got exception from REDIS %s",
str(e),
)
raise e

View file

@ -412,6 +412,26 @@ class MCPRequestHandler:
)
return []
#########################################################
# Check agent permissions if agent_id is set on the key
#########################################################
if user_api_key_auth and user_api_key_auth.agent_id:
allowed_mcp_servers_for_agent = (
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(
user_api_key_auth
)
)
if len(allowed_mcp_servers_for_agent) > 0:
# Intersect: agent can only use servers allowed by BOTH key/team AND agent config
allowed_mcp_servers = [
s
for s in allowed_mcp_servers
if s in allowed_mcp_servers_for_agent
]
verbose_logger.debug(
f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}"
)
return list(set(allowed_mcp_servers))
except Exception as e:
verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}")
@ -513,13 +533,33 @@ class MCPRequestHandler:
if team_tools:
if key_tools:
# Both have restrictions → intersection
return list(set(team_tools) & set(key_tools))
allowed_tools = list(set(team_tools) & set(key_tools))
else:
# Only team has restrictions → inherit from team
return team_tools
allowed_tools = team_tools
else:
# No team restrictions → use key restrictions
return key_tools
allowed_tools = key_tools
# Intersect with agent's tool permissions if agent_id is set
if user_api_key_auth.agent_id:
# Pre-fetch agent object_permission once to avoid duplicate DB query
agent_obj_perm = await MCPRequestHandler._get_agent_object_permission(
user_api_key_auth
)
agent_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server(
server_id=server_id,
user_api_key_auth=user_api_key_auth,
agent_object_permission=agent_obj_perm,
)
if agent_tools is not None:
if allowed_tools is not None:
allowed_tools = list(
set(allowed_tools) & set(agent_tools)
)
else:
allowed_tools = agent_tools
return allowed_tools
except Exception as e:
verbose_logger.warning(f"Failed to get allowed tools for server: {str(e)}")
@ -715,6 +755,131 @@ class MCPRequestHandler:
)
return []
@staticmethod
async def _get_agent_object_permission(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
):
"""
Fetch the agent's object_permission from the DB (single query).
Returns the object_permission object or None.
"""
from litellm.proxy.proxy_server import prisma_client
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
if prisma_client is None:
verbose_logger.debug("prisma_client is None")
return None
try:
agent_row = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": user_api_key_auth.agent_id},
include={"object_permission": True},
)
if agent_row is None or agent_row.object_permission is None:
return None
return agent_row.object_permission
except Exception as e:
verbose_logger.warning(
f"Failed to get agent object permission: {str(e)}"
)
return None
@staticmethod
async def _get_allowed_mcp_servers_for_agent(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
agent_object_permission=None,
) -> List[str]:
"""
Get allowed MCP servers for an agent (from the agent's object_permission).
Returns the MCP servers from the agent's object_permission.
If agent has no object_permission, returns [] (no extra restriction).
Args:
user_api_key_auth: User auth with agent_id
agent_object_permission: Pre-fetched object_permission to avoid duplicate DB query.
If None, will be fetched from DB.
"""
if not user_api_key_auth or not user_api_key_auth.agent_id:
return []
try:
obj_perm = agent_object_permission
if obj_perm is None:
obj_perm = await MCPRequestHandler._get_agent_object_permission(
user_api_key_auth
)
if obj_perm is None:
return []
direct_mcp_servers = getattr(obj_perm, "mcp_servers", None) or []
if isinstance(direct_mcp_servers, str):
direct_mcp_servers = []
mcp_access_groups = getattr(obj_perm, "mcp_access_groups", None) or []
if isinstance(mcp_access_groups, str):
mcp_access_groups = []
access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
mcp_access_groups
)
)
all_servers = list(direct_mcp_servers) + access_group_servers
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(
f"Failed to get allowed MCP servers for agent: {str(e)}"
)
return []
@staticmethod
async def _get_agent_tool_permissions_for_server(
server_id: str,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
agent_object_permission=None,
) -> Optional[List[str]]:
"""
Get allowed tool names for a server from the agent's object_permission.
Returns None if agent has no tool restrictions for this server.
Args:
server_id: Server ID to check permissions for
user_api_key_auth: User auth with agent_id
agent_object_permission: Pre-fetched object_permission to avoid duplicate DB query.
If None, will be fetched from DB.
"""
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
try:
obj_perm = agent_object_permission
if obj_perm is None:
obj_perm = await MCPRequestHandler._get_agent_object_permission(
user_api_key_auth
)
if obj_perm is None:
return None
mcp_tool_permissions = getattr(
obj_perm, "mcp_tool_permissions", None
)
if not mcp_tool_permissions:
return None
if isinstance(mcp_tool_permissions, dict):
tools = mcp_tool_permissions.get(server_id)
else:
tools = None
return list(tools) if tools else None
except Exception as e:
verbose_logger.warning(
f"Failed to get agent tool permissions for server: {str(e)}"
)
return None
@staticmethod
def _get_config_server_ids_for_access_groups(
config_mcp_servers, access_groups: List[str]

View file

@ -1,60 +1,40 @@
import enum
import json
from datetime import datetime
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Literal, Optional, Union
from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Literal,
Optional, Union)
import httpx
from pydantic import (
BaseModel,
ConfigDict,
Field,
Json,
field_validator,
model_validator,
)
from pydantic import (BaseModel, ConfigDict, Field, Json, field_validator,
model_validator)
from typing_extensions import Required, TypedDict
from litellm._uuid import uuid
from litellm.types.integrations.slack_alerting import AlertType
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIFileObject,
ResponsesAPIResponse,
)
from litellm.types.mcp import (
MCPAuth,
MCPAuthType,
MCPCredentials,
MCPTransport,
MCPTransportType,
)
from litellm.types.llms.openai import (AllMessageValues, OpenAIFileObject,
ResponsesAPIResponse)
from litellm.types.mcp import (MCPAuth, MCPAuthType, MCPCredentials,
MCPTransport, MCPTransportType)
from litellm.types.mcp_server.mcp_server_manager import MCPInfo
from litellm.types.router import RouterErrors, UpdateRouterConfig
from litellm.types.secret_managers.main import KeyManagementSystem
from litellm.types.utils import (
CallTypes,
CostBreakdown,
EmbeddingResponse,
GenericBudgetConfigType,
ImageResponse,
LiteLLMBatch,
LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase,
ModelResponse,
ProviderField,
StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse,
)
from litellm.types.utils import (CallTypes, CostBreakdown, EmbeddingResponse,
GenericBudgetConfigType, ImageResponse,
LiteLLMBatch, LiteLLMFineTuningJob,
LiteLLMPydanticObjectBase, ModelResponse,
ProviderField, StandardCallbackDynamicParams,
StandardLoggingGuardrailInformation,
StandardLoggingMCPToolCall,
StandardLoggingModelInformation,
StandardLoggingPayloadErrorInformation,
StandardLoggingPayloadStatus,
StandardLoggingVectorStoreRequest,
StandardPassThroughResponseObject,
TextCompletionResponse)
from litellm.types.videos.main import VideoObject
from .types_utils.utils import get_instance_fn, validate_custom_validate_return_type
from .types_utils.utils import (get_instance_fn,
validate_custom_validate_return_type)
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -2202,6 +2182,7 @@ class LiteLLM_VerificationToken(LiteLLMPydanticObjectBase):
config: Dict = {}
user_id: Optional[str] = None
team_id: Optional[str] = None
agent_id: Optional[str] = None
project_id: Optional[str] = None
max_parallel_requests: Optional[int] = None
metadata: Dict = {}
@ -2379,7 +2360,8 @@ class UserAPIKeyAuth(
This is used to track number of requests/spend for health check calls.
"""
from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
from litellm.constants import \
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME,
@ -2411,7 +2393,8 @@ class UserAPIKeyAuth(
This is used to track actions performed by automated system jobs.
"""
from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
from litellm.constants import \
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME
return cls(
api_key=LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME,
@ -2802,7 +2785,8 @@ class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):
@model_validator(mode="after")
def mask_api_keys(self):
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.litellm_core_utils.sensitive_data_masker import \
SensitiveDataMasker
masker = SensitiveDataMasker(sensitive_patterns={"key"})
@ -3105,6 +3089,7 @@ class SpendLogsPayload(TypedDict):
response: Optional[Union[str, list, dict]]
proxy_server_request: Optional[str]
session_id: Optional[str]
request_duration_ms: Optional[int]
status: Literal["success", "failure"]

View file

@ -5,6 +5,9 @@ from typing import Any, Dict, List, Optional
import litellm
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy.management_helpers.object_permission_utils import (
handle_update_object_permission_common,
)
from litellm.proxy.utils import PrismaClient
from litellm.types.agents import AgentConfig, AgentResponse, PatchAgentRequest
@ -117,20 +120,39 @@ class AgentRegistry:
)
agent_card_params: str = safe_dumps(agent_card_params_dict)
# Handle object_permission (MCP tool access for agent)
object_permission_id: Optional[str] = None
if agent.get("object_permission") is not None:
agent_copy = dict(agent)
object_permission_id = await handle_update_object_permission_common(
agent_copy, None, prisma_client
)
create_data: Dict[str, Any] = {
"agent_name": agent_name,
"litellm_params": litellm_params,
"agent_card_params": agent_card_params,
"created_by": created_by,
"updated_by": created_by,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
if object_permission_id is not None:
create_data["object_permission_id"] = object_permission_id
# Create agent in DB
created_agent = await prisma_client.db.litellm_agentstable.create(
data={
"agent_name": agent_name,
"litellm_params": litellm_params,
"agent_card_params": agent_card_params,
"created_by": created_by,
"updated_by": created_by,
"created_at": datetime.now(timezone.utc),
"updated_at": datetime.now(timezone.utc),
}
data=create_data,
include={"object_permission": True},
)
return AgentResponse(**created_agent.model_dump()) # type: ignore
created_agent_dict = created_agent.model_dump()
if created_agent.object_permission is not None:
try:
created_agent_dict["object_permission"] = created_agent.object_permission.model_dump()
except Exception:
created_agent_dict["object_permission"] = created_agent.object_permission.dict()
return AgentResponse(**created_agent_dict) # type: ignore
except Exception as e:
raise Exception(f"Error adding agent to DB: {str(e)}")
@ -181,7 +203,7 @@ class AgentRegistry:
raise Exception(f"Agent with ID {agent_id} not found")
augment_agent = {**existing_agent, **agent}
update_data = {}
update_data: Dict[str, Any] = {}
if augment_agent.get("agent_name"):
update_data["agent_name"] = augment_agent.get("agent_name")
if augment_agent.get("litellm_params"):
@ -192,6 +214,20 @@ class AgentRegistry:
update_data["agent_card_params"] = safe_dumps(
augment_agent.get("agent_card_params")
)
if agent.get("object_permission") is not None:
agent_copy = dict(augment_agent)
existing_object_permission_id = existing_agent.get(
"object_permission_id"
)
object_permission_id = (
await handle_update_object_permission_common(
agent_copy,
existing_object_permission_id,
prisma_client,
)
)
if object_permission_id is not None:
update_data["object_permission_id"] = object_permission_id
# Patch agent in DB
patched_agent = await prisma_client.db.litellm_agentstable.update(
where={"agent_id": agent_id},
@ -200,8 +236,15 @@ class AgentRegistry:
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),
},
include={"object_permission": True},
)
return AgentResponse(**patched_agent.model_dump()) # type: ignore
patched_agent_dict = patched_agent.model_dump()
if patched_agent.object_permission is not None:
try:
patched_agent_dict["object_permission"] = patched_agent.object_permission.model_dump()
except Exception:
patched_agent_dict["object_permission"] = patched_agent.object_permission.dict()
return AgentResponse(**patched_agent_dict) # type: ignore
except Exception as e:
raise Exception(f"Error patching agent in DB: {str(e)}")
@ -238,19 +281,47 @@ class AgentRegistry:
)
agent_card_params: str = safe_dumps(agent_card_params_dict)
update_data: Dict[str, Any] = {
"agent_name": agent_name,
"litellm_params": litellm_params,
"agent_card_params": agent_card_params,
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),
}
if agent.get("object_permission") is not None:
existing_agent = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": agent_id}
)
existing_object_permission_id = (
existing_agent.object_permission_id
if existing_agent is not None
else None
)
agent_copy = dict(agent)
object_permission_id = (
await handle_update_object_permission_common(
agent_copy,
existing_object_permission_id,
prisma_client,
)
)
if object_permission_id is not None:
update_data["object_permission_id"] = object_permission_id
# Update agent in DB
updated_agent = await prisma_client.db.litellm_agentstable.update(
where={"agent_id": agent_id},
data={
"agent_name": agent_name,
"litellm_params": litellm_params,
"agent_card_params": agent_card_params,
"updated_by": updated_by,
"updated_at": datetime.now(timezone.utc),
},
data=update_data,
include={"object_permission": True},
)
return AgentResponse(**updated_agent.model_dump()) # type: ignore
updated_agent_dict = updated_agent.model_dump()
if updated_agent.object_permission is not None:
try:
updated_agent_dict["object_permission"] = updated_agent.object_permission.model_dump()
except Exception:
updated_agent_dict["object_permission"] = updated_agent.object_permission.dict()
return AgentResponse(**updated_agent_dict) # type: ignore
except Exception as e:
raise Exception(f"Error updating agent in DB: {str(e)}")
@ -264,11 +335,19 @@ class AgentRegistry:
try:
agents_from_db = await prisma_client.db.litellm_agentstable.find_many(
order={"created_at": "desc"},
include={"object_permission": True},
)
agents: List[Dict[str, Any]] = []
for agent in agents_from_db:
agents.append(dict(agent))
agent_dict = dict(agent)
# object_permission is eagerly loaded via include above
if agent.object_permission is not None:
try:
agent_dict["object_permission"] = agent.object_permission.model_dump()
except Exception:
agent_dict["object_permission"] = agent.object_permission.dict()
agents.append(agent_dict)
return agents
except Exception as e:

View file

@ -16,6 +16,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
from litellm.types.agents import (
AgentConfig,
AgentMakePublicResponse,
@ -23,8 +24,6 @@ from litellm.types.agents import (
MakeAgentsPublicRequest,
PatchAgentRequest,
)
from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity
from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
)
@ -233,11 +232,18 @@ async def get_agent_by_id(agent_id: str):
try:
agent = AGENT_REGISTRY.get_agent_by_id(agent_id=agent_id)
if agent is None:
agent = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": agent_id}
agent_row = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": agent_id},
include={"object_permission": True},
)
if agent is not None:
agent = AgentResponse(**agent.model_dump()) # type: ignore
if agent_row is not None:
agent_dict = agent_row.model_dump()
if agent_row.object_permission is not None:
try:
agent_dict["object_permission"] = agent_row.object_permission.model_dump()
except Exception:
agent_dict["object_permission"] = agent_row.object_permission.dict()
agent = AgentResponse(**agent_dict) # type: ignore
if agent is None:
raise HTTPException(

View file

@ -139,53 +139,20 @@ class DBSpendUpdateWriter:
payload["startTime"] = payload["startTime"].isoformat()
if isinstance(payload["endTime"], datetime):
payload["endTime"] = payload["endTime"].isoformat()
if org_id is not None and org_id != "":
payload["organization_id"] = org_id
if team_id is not None and team_id != "":
payload["team_id"] = team_id
asyncio.create_task(
self._update_user_db(
response_cost=response_cost,
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
end_user_id=end_user_id,
)
)
asyncio.create_task(
self._update_key_db(
response_cost=response_cost,
hashed_token=hashed_token,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self._update_team_db(
response_cost=response_cost,
team_id=team_id,
user_id=user_id,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self._update_org_db(
response_cost=response_cost,
org_id=org_id,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self._update_tag_db(
response_cost=response_cost,
request_tags=copy.deepcopy(payload.get("request_tags")),
prisma_client=prisma_client,
)
)
# One deepcopy shared by all 6 daily spend helpers (was 5, fixes agent bug)
payload_copy = copy.deepcopy(payload)
# Deepcopy request_tags for _update_tag_db
request_tags = copy.deepcopy(payload.get("request_tags"))
# Keep _insert_spend_log_to_db awaited inline (not a task, preserve current behavior)
if disable_spend_logs is False:
await self._insert_spend_log_to_db(
payload=copy.deepcopy(payload),
@ -196,44 +163,20 @@ class DBSpendUpdateWriter:
"disable_spend_logs=True. Skipping writing spend logs to db. Other spend updates - Key/User/Team table will still occur."
)
# Single task replaces 11 create_task() calls
asyncio.create_task(
self.add_spend_log_transaction_to_daily_user_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_end_user_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_agent_transaction(
payload=payload,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_team_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_org_transaction(
payload=copy.deepcopy(payload),
self._batch_database_updates(
response_cost=response_cost,
user_id=user_id,
hashed_token=hashed_token,
team_id=team_id,
org_id=org_id,
end_user_id=end_user_id,
prisma_client=prisma_client,
)
)
asyncio.create_task(
self.add_spend_log_transaction_to_daily_tag_transaction(
payload=copy.deepcopy(payload),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
payload_copy=payload_copy,
request_tags=request_tags,
)
)
@ -357,6 +300,157 @@ class DBSpendUpdateWriter:
"_enqueue_tool_registry_upsert error (non-blocking): %s", e
)
async def _batch_database_updates(
self,
*,
response_cost: Optional[float],
user_id: Optional[str],
hashed_token: Optional[str],
team_id: Optional[str],
org_id: Optional[str],
end_user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
litellm_proxy_budget_name: Optional[str],
payload_copy: dict,
request_tags: Optional[Any],
):
"""
Runs all 11 spend-update helpers sequentially inside a single asyncio task.
Each helper is wrapped in try/except so one failure doesn't prevent the others.
"""
try:
await self._update_user_db(
response_cost=response_cost,
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
litellm_proxy_budget_name=litellm_proxy_budget_name,
end_user_id=end_user_id,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_user_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_key_db(
response_cost=response_cost,
hashed_token=hashed_token,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_key_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_team_db(
response_cost=response_cost,
team_id=team_id,
user_id=user_id,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_team_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_org_db(
response_cost=response_cost,
org_id=org_id,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_org_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_tag_db(
response_cost=response_cost,
request_tags=request_tags,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: _update_tag_db failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_user_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_user_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_end_user_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_end_user_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_agent_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_agent_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_team_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_team_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_org_transaction(
payload=payload_copy,
org_id=org_id,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_org_transaction failed: %s",
traceback.format_exc(),
)
try:
await self.add_spend_log_transaction_to_daily_tag_transaction(
payload=payload_copy,
prisma_client=prisma_client,
)
except Exception:
verbose_proxy_logger.debug(
"_batch_database_updates: add_spend_log_transaction_to_daily_tag_transaction failed: %s",
traceback.format_exc(),
)
async def _update_key_db(
self,
response_cost: Optional[float],
@ -666,9 +760,16 @@ class DBSpendUpdateWriter:
verbose_proxy_logger.debug("acquired lock for spend updates")
try:
db_spend_update_transactions = (
await self.redis_update_buffer.get_all_update_transactions_from_redis_buffer()
)
(
db_spend_update_transactions,
daily_spend_update_transactions,
daily_team_spend_update_transactions,
daily_org_spend_update_transactions,
daily_end_user_spend_update_transactions,
daily_agent_spend_update_transactions,
daily_tag_spend_update_transactions,
) = await self.redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
if db_spend_update_transactions is not None:
verbose_proxy_logger.info(
"Spend tracking - committing spend updates from Redis to DB: "
@ -688,9 +789,6 @@ class DBSpendUpdateWriter:
db_spend_update_transactions=db_spend_update_transactions,
)
daily_spend_update_transactions = (
await self.redis_update_buffer.get_all_daily_spend_update_transactions_from_redis_buffer()
)
if daily_spend_update_transactions is not None:
await DBSpendUpdateWriter.update_daily_user_spend(
n_retry_times=n_retry_times,
@ -698,9 +796,6 @@ class DBSpendUpdateWriter:
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_spend_update_transactions,
)
daily_team_spend_update_transactions = (
await self.redis_update_buffer.get_all_daily_team_spend_update_transactions_from_redis_buffer()
)
if daily_team_spend_update_transactions is not None:
await DBSpendUpdateWriter.update_daily_team_spend(
n_retry_times=n_retry_times,
@ -709,9 +804,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_team_spend_update_transactions,
)
daily_org_spend_update_transactions = (
await self.redis_update_buffer.get_all_daily_org_spend_update_transactions_from_redis_buffer()
)
if daily_org_spend_update_transactions is not None:
await DBSpendUpdateWriter.update_daily_org_spend(
n_retry_times=n_retry_times,
@ -720,9 +812,6 @@ class DBSpendUpdateWriter:
daily_spend_transactions=daily_org_spend_update_transactions,
)
daily_tag_spend_update_transactions = (
await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer()
)
if daily_tag_spend_update_transactions is not None:
await DBSpendUpdateWriter.update_daily_tag_spend(
n_retry_times=n_retry_times,
@ -730,9 +819,6 @@ class DBSpendUpdateWriter:
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_tag_spend_update_transactions,
)
daily_end_user_spend_update_transactions = (
await self.redis_update_buffer.get_all_daily_end_user_spend_update_transactions_from_redis_buffer()
)
if daily_end_user_spend_update_transactions is not None:
await DBSpendUpdateWriter.update_daily_end_user_spend(
n_retry_times=n_retry_times,
@ -740,9 +826,6 @@ class DBSpendUpdateWriter:
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_end_user_spend_update_transactions,
)
daily_agent_spend_update_transactions = (
await self.redis_update_buffer.get_all_daily_agent_spend_update_transactions_from_redis_buffer()
)
if daily_agent_spend_update_transactions is not None:
await DBSpendUpdateWriter.update_daily_agent_spend(
n_retry_times=n_retry_times,
@ -1072,7 +1155,7 @@ class DBSpendUpdateWriter:
team_id = key.split("::")[1]
user_id = key.split("::")[3]
team_memberships_to_invalidate.append((user_id, team_id))
for i in range(n_retry_times + 1):
start_time = time.time()
try:
@ -1109,11 +1192,13 @@ class DBSpendUpdateWriter:
_raise_failed_update_spend_exception(
e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj
)
# Invalidate cache for updated team memberships
# This ensures budget checks read fresh spend data from the database
if team_memberships_to_invalidate and proxy_logging_obj is not None:
user_api_key_cache = proxy_logging_obj.call_details.get("user_api_key_cache")
user_api_key_cache = proxy_logging_obj.call_details.get(
"user_api_key_cache"
)
if user_api_key_cache is not None:
for user_id, team_id in team_memberships_to_invalidate:
cache_key = "team_membership:{}:{}".format(user_id, team_id)
@ -1425,7 +1510,9 @@ class DBSpendUpdateWriter:
),
"endpoint": transaction.get("endpoint") or "",
"prompt_tokens": transaction["prompt_tokens"],
"completion_tokens": transaction["completion_tokens"],
"completion_tokens": transaction[
"completion_tokens"
],
"spend": transaction["spend"],
"api_requests": transaction["api_requests"],
"successful_requests": transaction[
@ -1436,12 +1523,14 @@ class DBSpendUpdateWriter:
# Add cache-related fields if they exist
if "cache_read_input_tokens" in transaction:
common_data["cache_read_input_tokens"] = (
transaction.get("cache_read_input_tokens", 0)
)
common_data[
"cache_read_input_tokens"
] = transaction.get("cache_read_input_tokens", 0)
if "cache_creation_input_tokens" in transaction:
common_data["cache_creation_input_tokens"] = (
transaction.get("cache_creation_input_tokens", 0)
common_data[
"cache_creation_input_tokens"
] = transaction.get(
"cache_creation_input_tokens", 0
)
if entity_type == "tag" and "request_id" in transaction:
@ -1484,10 +1573,14 @@ class DBSpendUpdateWriter:
}
if entity_type == "tag" and "request_id" in transaction:
update_data["request_id"] = transaction.get("request_id")
update_data["request_id"] = transaction.get(
"request_id"
)
# Add endpoint to update_data so existing rows get their endpoint field updated
update_data["endpoint"] = transaction.get("endpoint") or ""
update_data["endpoint"] = (
transaction.get("endpoint") or ""
)
table.upsert(
where=where_clause,
@ -1671,7 +1764,9 @@ class DBSpendUpdateWriter:
self,
payload: Union[dict, SpendLogsPayload],
prisma_client: PrismaClient,
type: Literal["user", "team", "org", "request_tags", "end_user", "agent"] = "user",
type: Literal[
"user", "team", "org", "request_tags", "end_user", "agent"
] = "user",
) -> Optional[BaseDailySpendTransaction]:
common_expected_keys = ["startTime", "api_key"]
if type == "user":
@ -1730,7 +1825,7 @@ class DBSpendUpdateWriter:
endpoint = None
if call_type:
endpoint = ROUTE_ENDPOINT_MAPPING.get(call_type, None)
daily_transaction = BaseDailySpendTransaction(
date=date,
api_key=payload["api_key"],
@ -1942,7 +2037,7 @@ class DBSpendUpdateWriter:
endpoint_str = base_daily_transaction.get("endpoint") or ""
daily_transaction_key = f"{payload['agent_id']}_{base_daily_transaction['date']}_{payload_with_agent_id['api_key']}_{payload_with_agent_id['model']}_{payload_with_agent_id['custom_llm_provider']}_{endpoint_str}"
daily_transaction = DailyAgentSpendTransaction(
agent_id=payload['agent_id'], **base_daily_transaction
agent_id=payload["agent_id"], **base_daily_transaction
)
await self.daily_agent_spend_update_queue.add_update(
update={daily_transaction_key: daily_transaction}

View file

@ -6,7 +6,7 @@ This is to prevent deadlocks and improve reliability
import asyncio
import json
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
from litellm._logging import verbose_proxy_logger
from litellm.caching import RedisCache
@ -36,6 +36,7 @@ from litellm.proxy.db.db_transaction_queue.daily_spend_update_queue import (
)
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
from litellm.secret_managers.main import str_to_bool
from litellm.types.caching import RedisPipelineLpopOperation, RedisPipelineRpushOperation
from litellm.types.services import ServiceTypes
if TYPE_CHECKING:
@ -209,47 +210,44 @@ class RedisUpdateBuffer:
"ALL DAILY SPEND UPDATE TRANSACTIONS: %s", daily_spend_update_transactions
)
await self._store_transactions_in_redis(
transactions=db_spend_update_transactions,
redis_key=REDIS_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_SPEND_UPDATE_QUEUE,
# Build a list of rpush operations, skipping empty/None transaction sets
_queue_configs: List[Tuple[Any, str, ServiceTypes]] = [
(db_spend_update_transactions, REDIS_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_SPEND_UPDATE_QUEUE),
(daily_spend_update_transactions, REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_SPEND_UPDATE_QUEUE),
(daily_team_spend_update_transactions, REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_TEAM_SPEND_UPDATE_QUEUE),
(daily_org_spend_update_transactions, REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_ORG_SPEND_UPDATE_QUEUE),
(daily_end_user_spend_update_transactions, REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_END_USER_SPEND_UPDATE_QUEUE),
(daily_agent_spend_update_transactions, REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_AGENT_SPEND_UPDATE_QUEUE),
(daily_tag_spend_update_transactions, REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY, ServiceTypes.REDIS_DAILY_TAG_SPEND_UPDATE_QUEUE),
]
rpush_list: List[RedisPipelineRpushOperation] = []
service_types: List[ServiceTypes] = []
for transactions, redis_key, service_type in _queue_configs:
if transactions is None or len(transactions) == 0:
continue
rpush_list.append(
RedisPipelineRpushOperation(
key=redis_key,
values=[safe_dumps(transactions)],
)
)
service_types.append(service_type)
if len(rpush_list) == 0:
return
result_lengths = await self.redis_cache.async_rpush_pipeline(
rpush_list=rpush_list,
)
await self._store_transactions_in_redis(
transactions=daily_spend_update_transactions,
redis_key=REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_DAILY_SPEND_UPDATE_QUEUE,
)
await self._store_transactions_in_redis(
transactions=daily_team_spend_update_transactions,
redis_key=REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_DAILY_TEAM_SPEND_UPDATE_QUEUE,
)
await self._store_transactions_in_redis(
transactions=daily_org_spend_update_transactions,
redis_key=REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_DAILY_SPEND_UPDATE_QUEUE,
)
await self._store_transactions_in_redis(
transactions=daily_end_user_spend_update_transactions,
redis_key=REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_DAILY_END_USER_SPEND_UPDATE_QUEUE,
)
await self._store_transactions_in_redis(
transactions=daily_agent_spend_update_transactions,
redis_key=REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_DAILY_AGENT_SPEND_UPDATE_QUEUE,
)
await self._store_transactions_in_redis(
transactions=daily_tag_spend_update_transactions,
redis_key=REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY,
service_type=ServiceTypes.REDIS_DAILY_TAG_SPEND_UPDATE_QUEUE,
)
# Emit gauge events for each queue
for i, queue_size in enumerate(result_lengths):
if i < len(service_types):
await self._emit_new_item_added_to_redis_buffer_event(
queue_size=queue_size,
service=service_types[i],
)
@staticmethod
def _number_of_transactions_to_store_in_redis(
@ -338,6 +336,77 @@ class RedisUpdateBuffer:
return combined_transaction
async def get_all_transactions_from_redis_buffer_pipeline(
self,
) -> Tuple[
Optional[DBSpendUpdateTransactions],
Optional[Dict[str, DailyUserSpendTransaction]],
Optional[Dict[str, DailyTeamSpendTransaction]],
Optional[Dict[str, DailyOrganizationSpendTransaction]],
Optional[Dict[str, DailyEndUserSpendTransaction]],
Optional[Dict[str, DailyAgentSpendTransaction]],
Optional[Dict[str, DailyTagSpendTransaction]],
]:
"""
Drains all 7 Redis buffer queues in a single pipeline round-trip.
Returns a 7-tuple of parsed results in this order:
0: DBSpendUpdateTransactions
1: daily user spend
2: daily team spend
3: daily org spend
4: daily end-user spend
5: daily agent spend
6: daily tag spend
"""
if self.redis_cache is None:
return None, None, None, None, None, None, None
lpop_list: List[RedisPipelineLpopOperation] = [
RedisPipelineLpopOperation(key=REDIS_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
RedisPipelineLpopOperation(key=REDIS_DAILY_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
RedisPipelineLpopOperation(key=REDIS_DAILY_TEAM_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
RedisPipelineLpopOperation(key=REDIS_DAILY_ORG_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
RedisPipelineLpopOperation(key=REDIS_DAILY_END_USER_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
RedisPipelineLpopOperation(key=REDIS_DAILY_AGENT_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
RedisPipelineLpopOperation(key=REDIS_DAILY_TAG_SPEND_UPDATE_BUFFER_KEY, count=MAX_REDIS_BUFFER_DEQUEUE_COUNT),
]
raw_results = await self.redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
# Pad with None if pipeline returned fewer results than expected
while len(raw_results) < 7:
raw_results.append(None)
# Slot 0: DBSpendUpdateTransactions
db_spend: Optional[DBSpendUpdateTransactions] = None
if raw_results[0] is not None:
parsed = self._parse_list_of_transactions(raw_results[0])
if len(parsed) > 0:
db_spend = self._combine_list_of_transactions(parsed)
# Slots 1-6: daily spend categories
daily_results: List[Optional[Dict[str, Any]]] = []
for slot in range(1, 7):
if raw_results[slot] is None:
daily_results.append(None)
else:
list_of_daily = [json.loads(t) for t in raw_results[slot]] # type: ignore
aggregated = DailySpendUpdateQueue.get_aggregated_daily_spend_update_transactions(
list_of_daily
)
daily_results.append(aggregated)
return (
db_spend,
cast(Optional[Dict[str, DailyUserSpendTransaction]], daily_results[0]),
cast(Optional[Dict[str, DailyTeamSpendTransaction]], daily_results[1]),
cast(Optional[Dict[str, DailyOrganizationSpendTransaction]], daily_results[2]),
cast(Optional[Dict[str, DailyEndUserSpendTransaction]], daily_results[3]),
cast(Optional[Dict[str, DailyAgentSpendTransaction]], daily_results[4]),
cast(Optional[Dict[str, DailyTagSpendTransaction]], daily_results[5]),
)
async def get_all_daily_spend_update_transactions_from_redis_buffer(
self,
) -> Optional[Dict[str, DailyUserSpendTransaction]]:

View file

@ -0,0 +1,208 @@
"""
Max Iterations Limiter for LiteLLM Proxy.
Enforces a per-session cap on the number of LLM calls an agentic loop can make.
Callers send a `session_id` with each request (via `x-litellm-session-id` header
or `metadata.session_id`), and this hook counts calls per session. When the count
exceeds `max_iterations` (configured in key/team metadata), returns 429.
Works across multiple proxy instances via DualCache (in-memory + Redis).
Follows the same pattern as parallel_request_limiter_v3.py.
"""
import os
from typing import TYPE_CHECKING, Any, Optional, Union
from fastapi import HTTPException
from litellm import DualCache
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
if TYPE_CHECKING:
from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache
InternalUsageCache = _InternalUsageCache
else:
InternalUsageCache = Any
# Redis Lua script for atomic increment with TTL.
# Returns the new count after increment.
# Only sets EXPIRE on first increment (when count becomes 1).
MAX_ITERATIONS_INCREMENT_SCRIPT = """
local key = KEYS[1]
local ttl = tonumber(ARGV[1])
local current = redis.call('INCR', key)
if current == 1 then
redis.call('EXPIRE', key, ttl)
end
return current
"""
# Default TTL for session iteration counters (1 hour)
DEFAULT_MAX_ITERATIONS_TTL = 3600
class _PROXY_MaxIterationsHandler(CustomLogger):
"""
Pre-call hook that enforces max_iterations per session.
Configuration:
- max_iterations: set in key metadata via /key/generate or /key/update
e.g. metadata={"max_iterations": 25}
- session_id: sent by caller via x-litellm-session-id header or
metadata.session_id in request body
Cache key pattern:
{session_iterations:<session_id>}:count
Multi-instance support:
Uses Redis Lua script for atomic increment (same pattern as
parallel_request_limiter_v3). Falls back to in-memory cache
when Redis is unavailable.
"""
def __init__(self, internal_usage_cache: InternalUsageCache):
self.internal_usage_cache = internal_usage_cache
self.ttl = int(
os.getenv("LITELLM_MAX_ITERATIONS_TTL", DEFAULT_MAX_ITERATIONS_TTL)
)
# Register Lua script with Redis if available (same pattern as v3 limiter)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
self.increment_script = (
self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
MAX_ITERATIONS_INCREMENT_SCRIPT
)
)
else:
self.increment_script = None
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: str,
) -> Optional[Union[Exception, str, dict]]:
"""
Check session iteration count before making the API call.
Extracts session_id from request metadata and max_iterations from
key metadata. If the session has exceeded max_iterations, raises 429.
"""
# Extract session_id from request data
session_id = self._get_session_id(data)
if session_id is None:
return None
# Extract max_iterations from key metadata
max_iterations = self._get_max_iterations(user_api_key_dict)
if max_iterations is None:
return None
verbose_proxy_logger.debug(
"MaxIterationsHandler: session_id=%s, max_iterations=%s",
session_id,
max_iterations,
)
# Increment and check
cache_key = self._make_cache_key(session_id)
current_count = await self._increment_and_get(cache_key)
if current_count > max_iterations:
raise HTTPException(
status_code=429,
detail=(
f"Max iterations exceeded for session {session_id}. "
f"Current count: {current_count}, max_iterations: {max_iterations}."
),
)
verbose_proxy_logger.debug(
"MaxIterationsHandler: session_id=%s, count=%s/%s",
session_id,
current_count,
max_iterations,
)
return None
def _get_session_id(self, data: dict) -> Optional[str]:
"""Extract session_id from request metadata."""
metadata = data.get("metadata") or {}
session_id = metadata.get("session_id")
if session_id is not None:
return str(session_id)
# Also check litellm_metadata (used for /thread and /assistant endpoints)
litellm_metadata = data.get("litellm_metadata") or {}
session_id = litellm_metadata.get("session_id")
if session_id is not None:
return str(session_id)
return None
def _get_max_iterations(
self, user_api_key_dict: UserAPIKeyAuth
) -> Optional[int]:
"""Extract max_iterations from key metadata."""
metadata = user_api_key_dict.metadata or {}
max_iterations = metadata.get("max_iterations")
if max_iterations is not None:
return int(max_iterations)
return None
def _make_cache_key(self, session_id: str) -> str:
"""
Create cache key for session iteration counter.
Uses Redis hash-tag pattern {session_iterations:<session_id>} so all
keys for a session land on the same Redis Cluster slot.
"""
return f"{{session_iterations:{session_id}}}:count"
async def _increment_and_get(self, cache_key: str) -> int:
"""
Atomically increment the session counter and return the new value.
Tries Redis first (via registered Lua script for atomicity across
instances), falls back to in-memory cache.
"""
if self.increment_script is not None:
try:
result = await self.increment_script(
keys=[cache_key],
args=[self.ttl],
)
return int(result)
except Exception as e:
verbose_proxy_logger.warning(
"MaxIterationsHandler: Redis failed, falling back to in-memory: %s",
str(e),
)
# Fallback: in-memory cache
return await self._in_memory_increment(cache_key)
async def _in_memory_increment(self, cache_key: str) -> int:
"""Increment counter in in-memory cache with TTL."""
current = await self.internal_usage_cache.async_get_cache(
key=cache_key,
litellm_parent_otel_span=None,
local_only=True,
)
new_value = (int(current) if current is not None else 0) + 1
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=new_value,
ttl=self.ttl,
litellm_parent_otel_span=None,
local_only=True,
)
return new_value

View file

@ -480,6 +480,7 @@ model LiteLLM_SpendLogs {
completion_tokens Int @default(0)
startTime DateTime // Assuming start_time is a DateTime field
endTime DateTime // Assuming end_time is a DateTime field
request_duration_ms Int?
completionStartTime DateTime? // Assuming completionStartTime is a DateTime field
model String @default("")
model_id String? @default("") // the model id stored in proxy model db

View file

@ -1726,7 +1726,7 @@ async def ui_view_spend_logs( # noqa: PLR0915
)
# Validate sort_by and sort_order
valid_sort_fields = {"spend", "total_tokens", "startTime", "endTime"}
valid_sort_fields = {"spend", "total_tokens", "startTime", "endTime", "request_duration_ms"}
if sort_by not in valid_sort_fields:
raise ProxyException(
message=f"Invalid sort_by: {sort_by}. Must be one of: {', '.join(sorted(valid_sort_fields))}",
@ -1939,7 +1939,8 @@ async def ui_view_spend_logs( # noqa: PLR0915
custom_llm_provider, api_base, "user", metadata,
cache_hit, cache_key, request_tags, team_id,
organization_id, end_user, requester_ip_address,
session_id, status, mcp_namespaced_tool_name, agent_id
session_id, status, mcp_namespaced_tool_name, agent_id,
COALESCE(request_duration_ms, (EXTRACT(EPOCH FROM ("endTime" - "startTime")) * 1000)::INTEGER) AS request_duration_ms
FROM "LiteLLM_SpendLogs"
WHERE {" AND ".join(sql_conditions)}
ORDER BY {_sql_col} {_sql_dir}

View file

@ -447,6 +447,7 @@ def get_logging_payload( # noqa: PLR0915
kwargs=kwargs,
standard_logging_payload=standard_logging_payload,
),
request_duration_ms=_get_request_duration_ms(start_time, end_time),
status=_get_status_for_spend_log(
metadata=metadata,
),
@ -496,6 +497,16 @@ def _get_session_id_for_spend_log(
return str(uuid.uuid4())
def _get_request_duration_ms(
start_time: datetime, end_time: datetime
) -> Optional[int]:
"""Compute request duration in milliseconds from start and end times."""
try:
return int((end_time - start_time).total_seconds() * 1000)
except Exception:
return None
def _ensure_datetime_utc(timestamp: datetime) -> datetime:
"""Helper to ensure datetime is in UTC"""
timestamp = timestamp.astimezone(timezone.utc)

View file

@ -167,16 +167,25 @@ class AugmentedAgentCard(AgentCard):
is_public: bool
# Object permission shape for agent MCP tool access (mirrors LiteLLM_ObjectPermissionBase)
class AgentObjectPermission(TypedDict, total=False):
mcp_servers: Optional[List[str]]
mcp_access_groups: Optional[List[str]]
mcp_tool_permissions: Optional[Dict[str, List[str]]]
class AgentConfig(TypedDict, total=False):
agent_name: Required[str]
agent_card_params: Required[AgentCard]
litellm_params: Dict[str, Any] # allow for any future litellm params
object_permission: AgentObjectPermission
class PatchAgentRequest(TypedDict, total=False):
agent_name: str
agent_card_params: AgentCard
litellm_params: Dict[str, Any]
object_permission: AgentObjectPermission
# Request/Response models for CRUD endpoints
@ -187,6 +196,7 @@ class AgentResponse(BaseModel):
agent_name: str
litellm_params: Optional[Dict[str, Any]] = None
agent_card_params: Dict[str, Any]
object_permission: Optional[Dict[str, Any]] = None
created_at: Optional[datetime] = None
updated_at: Optional[datetime] = None
created_by: Optional[str] = None

View file

@ -52,6 +52,24 @@ class RedisPipelineSetOperation(TypedDict):
ttl: Optional[int]
class RedisPipelineRpushOperation(TypedDict):
"""
TypedDict for 1 Redis Pipeline RPUSH Operation
"""
key: str
values: List[Any]
class RedisPipelineLpopOperation(TypedDict):
"""
TypedDict for 1 Redis Pipeline LPOP Operation
"""
key: str
count: Optional[int]
DynamicCacheControl = TypedDict(
"DynamicCacheControl",
{

8
poetry.lock generated
View file

@ -3222,15 +3222,15 @@ files = [
[[package]]
name = "litellm-proxy-extras"
version = "0.4.47"
version = "0.4.48"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
optional = true
python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8"
groups = ["main"]
markers = "extra == \"proxy\""
files = [
{file = "litellm_proxy_extras-0.4.47-py3-none-any.whl", hash = "sha256:2e900ae3edfbc20d27556f092d914974d37bac213efe88b8fd5287f77b2b7ca7"},
{file = "litellm_proxy_extras-0.4.47.tar.gz", hash = "sha256:42d88929f9eaf0b827046d3712095354db843c1716ccabb2a40c806ea5f809b9"},
{file = "litellm_proxy_extras-0.4.48-py3-none-any.whl", hash = "sha256:097001fccec5dbf4cffd902114898a9cfeba62673202447d55d2d0286cf93126"},
{file = "litellm_proxy_extras-0.4.48.tar.gz", hash = "sha256:5d5d8acf31b92d0cd6738555fb4a2411819755155438de9fb23c724c356400a2"},
]
[[package]]
@ -7989,4 +7989,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "6f39f8c731e625f37460e1f5f9ba3cce63956540dc04baf6ce1b7c18b88b8322"
content-hash = "b9b1e47b3b84748c0053be6a544c2399bf2601746a4f88dcb1be7c5e4eeab359"

View file

@ -61,7 +61,7 @@ boto3 = { version = "1.40.76", optional = true }
redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"}
mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"}
a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"}
litellm-proxy-extras = {version = "0.4.47", optional = true}
litellm-proxy-extras = {version = "0.4.48", optional = true}
rich = {version = "13.7.1", optional = true}
litellm-enterprise = {version = "0.1.32", optional = true}
diskcache = {version = "^5.6.1", optional = true}

View file

@ -57,7 +57,7 @@ grpcio>=1.75.0; python_version >= "3.14"
sentry_sdk==2.21.0 # for sentry error handling
detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests
tzdata==2025.1 # IANA time zone database
litellm-proxy-extras==0.4.47 # for proxy extras - e.g. prisma migrations
litellm-proxy-extras==0.4.48 # for proxy extras - e.g. prisma migrations
llm-sandbox==0.3.31 # for skill execution in sandbox
### LITELLM PACKAGE DEPENDENCIES
python-dotenv==1.0.1 # for env

View file

@ -64,6 +64,8 @@ model LiteLLM_AgentsTable {
litellm_params Json?
agent_card_params Json
agent_access_groups String[] @default([])
object_permission_id String?
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
created_at DateTime @default(now()) @map("created_at")
created_by String
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
@ -264,6 +266,7 @@ model LiteLLM_ObjectPermissionTable {
organizations LiteLLM_OrganizationTable[]
users LiteLLM_UserTable[]
end_users LiteLLM_EndUserTable[]
agents_table LiteLLM_AgentsTable[]
}
// Holds the MCP server configuration
@ -273,7 +276,6 @@ model LiteLLM_MCPServerTable {
alias String?
description String?
url String?
spec_path String?
transport String @default("sse")
auth_type String?
credentials Json? @default("{}")
@ -315,6 +317,7 @@ model LiteLLM_VerificationToken {
router_settings Json? @default("{}")
user_id String?
team_id String?
agent_id String?
project_id String?
permissions Json @default("{}")
max_parallel_requests Int?
@ -477,6 +480,7 @@ model LiteLLM_SpendLogs {
completion_tokens Int @default(0)
startTime DateTime // Assuming start_time is a DateTime field
endTime DateTime // Assuming end_time is a DateTime field
request_duration_ms Int?
completionStartTime DateTime? // Assuming completionStartTime is a DateTime field
model String @default("")
model_id String? @default("") // the model id stored in proxy model db
@ -1052,6 +1056,26 @@ model LiteLLM_PolicyAttachmentTable {
updated_by String?
}
// Global tool registry - auto-discovered from LLM responses; admins set call_policy here
model LiteLLM_ToolTable {
tool_id String @id @default(uuid())
tool_name String @unique // e.g. "huggingface_remote-mcp__dynamic_space"
origin String? // MCP server name or "user_defined"
call_policy String @default("untrusted") // "trusted" | "untrusted" | "dual_llm" | "blocked"
call_count Int @default(0) // cumulative number of times this tool was seen
assignments Json? @default("{}")
key_hash String? // hash of the virtual key that first called this tool
team_id String? // team that first called this tool
key_alias String? // human-readable alias of the virtual key
created_at DateTime @default(now())
created_by String?
updated_at DateTime @default(now()) @updatedAt
updated_by String?
@@index([call_policy])
@@index([team_id])
}
//Unified Access Groups table for storing unified access groups
model LiteLLM_AccessGroupTable {
access_group_id String @id @default(uuid())

View file

@ -0,0 +1,186 @@
#!/usr/bin/env bash
#
# Test agent endpoint-level changes for MCP tool permissions (object_permission).
# Requires: proxy running, valid admin API key, curl, jq.
#
# Usage:
# export LITELLM_PROXY_BASE_URL="http://localhost:4000" # optional, default below
# export LITELLM_API_KEY="sk-..." # required
# ./scripts/test_agent_mcp_endpoints.sh
#
set -euo pipefail
BASE_URL="${LITELLM_PROXY_BASE_URL:-http://localhost:4000}"
API_KEY="${LITELLM_API_KEY:-}"
if ! command -v jq &>/dev/null; then
echo "Error: jq is required. Install with: brew install jq (macOS) or apt install jq (Linux)"
exit 1
fi
if [[ -z "$API_KEY" ]]; then
echo "Error: LITELLM_API_KEY is not set. Export it or pass via env."
exit 1
fi
AUTH_HEADER="Authorization: Bearer $API_KEY"
AGENT_NAME="test-agent-mcp-$(date +%s)"
# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
NC='\033[0m'
pass() { echo -e "${GREEN}PASS${NC}: $*"; }
fail() { echo -e "${RED}FAIL${NC}: $*"; exit 1; }
info() { echo -e "${YELLOW}INFO${NC}: $*"; }
# --- 1. Create agent with object_permission ---
info "Creating agent with object_permission (mcp_servers, mcp_tool_permissions)..."
CREATE_RESP=$(curl -s -w "\n%{http_code}" -X POST "$BASE_URL/v1/agents" \
-H "$AUTH_HEADER" \
-H "Content-Type: application/json" \
-d '{
"agent_name": "'"$AGENT_NAME"'",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Test MCP Agent",
"description": "Agent for endpoint tests",
"url": "http://localhost:9999/",
"version": "1.0.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {"streaming": true},
"skills": []
},
"object_permission": {
"mcp_servers": ["server_1", "server_2"],
"mcp_access_groups": ["group_a"],
"mcp_tool_permissions": {"server_1": ["tool_a", "tool_b"], "server_2": ["tool_c"]}
}
}')
HTTP_CODE=$(echo "$CREATE_RESP" | tail -n1)
BODY=$(echo "$CREATE_RESP" | sed '$d')
if [[ "$HTTP_CODE" != "200" ]]; then
fail "POST /v1/agents returned $HTTP_CODE. Body: $BODY"
fi
AGENT_ID=$(echo "$BODY" | jq -r '.agent_id')
if [[ -z "$AGENT_ID" || "$AGENT_ID" == "null" ]]; then
fail "POST /v1/agents did not return agent_id. Body: $BODY"
fi
pass "Created agent $AGENT_ID"
# Check create response includes object_permission
OP=$(echo "$BODY" | jq '.object_permission')
if [[ "$OP" == "null" || -z "$OP" ]]; then
fail "POST /v1/agents response missing object_permission. Body: $BODY"
fi
SERVERS=$(echo "$OP" | jq -r '.mcp_servers | join(",")')
if [[ "$SERVERS" != "server_1,server_2" ]]; then
fail "object_permission.mcp_servers unexpected: $SERVERS"
fi
pass "Create response includes object_permission with mcp_servers and mcp_tool_permissions"
# --- 2. GET /v1/agents (list) includes object_permission for our agent ---
info "GET /v1/agents and check one agent has object_permission..."
LIST_RESP=$(curl -s -w "\n%{http_code}" -X GET "$BASE_URL/v1/agents" -H "$AUTH_HEADER")
LIST_CODE=$(echo "$LIST_RESP" | tail -n1)
LIST_BODY=$(echo "$LIST_RESP" | sed '$d')
if [[ "$LIST_CODE" != "200" ]]; then
fail "GET /v1/agents returned $LIST_CODE"
fi
AGENT_IN_LIST=$(echo "$LIST_BODY" | jq --arg id "$AGENT_ID" '.[] | select(.agent_id == $id)')
if [[ -z "$AGENT_IN_LIST" ]]; then
fail "GET /v1/agents did not return agent $AGENT_ID (list might be key-scoped)"
fi
OP_LIST=$(echo "$AGENT_IN_LIST" | jq '.object_permission')
if [[ "$OP_LIST" == "null" || -z "$OP_LIST" ]]; then
fail "GET /v1/agents list entry for agent missing object_permission"
fi
pass "GET /v1/agents list includes object_permission for agent"
# --- 3. GET /v1/agents/{agent_id} returns object_permission ---
info "GET /v1/agents/{agent_id}..."
GET_RESP=$(curl -s -w "\n%{http_code}" -X GET "$BASE_URL/v1/agents/$AGENT_ID" -H "$AUTH_HEADER")
GET_CODE=$(echo "$GET_RESP" | tail -n1)
GET_BODY=$(echo "$GET_RESP" | sed '$d')
if [[ "$GET_CODE" != "200" ]]; then
fail "GET /v1/agents/$AGENT_ID returned $GET_CODE. Body: $GET_BODY"
fi
OP_GET=$(echo "$GET_BODY" | jq '.object_permission')
if [[ "$OP_GET" == "null" || -z "$OP_GET" ]]; then
fail "GET /v1/agents/$AGENT_ID response missing object_permission"
fi
TOOL_PERMS=$(echo "$OP_GET" | jq -r '.mcp_tool_permissions.server_1 | join(",")')
if [[ "$TOOL_PERMS" != "tool_a,tool_b" ]]; then
fail "object_permission.mcp_tool_permissions.server_1 unexpected: $TOOL_PERMS"
fi
pass "GET /v1/agents/{agent_id} returns object_permission with mcp_tool_permissions"
# --- 4. PATCH /v1/agents/{agent_id} with new object_permission ---
info "PATCH /v1/agents/{agent_id} with updated object_permission..."
PATCH_RESP=$(curl -s -w "\n%{http_code}" -X PATCH "$BASE_URL/v1/agents/$AGENT_ID" \
-H "$AUTH_HEADER" \
-H "Content-Type: application/json" \
-d '{
"object_permission": {
"mcp_servers": ["server_3"],
"mcp_tool_permissions": {"server_3": ["tool_x"]}
}
}')
PATCH_CODE=$(echo "$PATCH_RESP" | tail -n1)
PATCH_BODY=$(echo "$PATCH_RESP" | sed '$d')
if [[ "$PATCH_CODE" != "200" ]]; then
fail "PATCH /v1/agents/$AGENT_ID returned $PATCH_CODE. Body: $PATCH_BODY"
fi
OP_PATCH=$(echo "$PATCH_BODY" | jq '.object_permission')
if [[ "$OP_PATCH" == "null" || -z "$OP_PATCH" ]]; then
fail "PATCH response missing object_permission"
fi
PATCH_SERVERS=$(echo "$OP_PATCH" | jq -r '.mcp_servers | join(",")')
if [[ "$PATCH_SERVERS" != "server_3" ]]; then
fail "PATCH object_permission.mcp_servers unexpected: $PATCH_SERVERS"
fi
pass "PATCH /v1/agents/{agent_id} updates and returns object_permission"
# --- 5. Create agent without object_permission; GET should still work ---
info "Creating agent without object_permission..."
AGENT_NAME_2="test-agent-no-mcp-$(date +%s)"
CREATE2_RESP=$(curl -s -w "\n%{http_code}" -X POST "$BASE_URL/v1/agents" \
-H "$AUTH_HEADER" \
-H "Content-Type: application/json" \
-d '{
"agent_name": "'"$AGENT_NAME_2"'",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "No MCP Agent",
"description": "No object_permission",
"url": "http://localhost:9999/",
"version": "1.0.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {},
"skills": []
}
}')
CODE2=$(echo "$CREATE2_RESP" | tail -n1)
BODY2=$(echo "$CREATE2_RESP" | sed '$d')
if [[ "$CODE2" != "200" ]]; then
fail "POST /v1/agents (no object_permission) returned $CODE2. Body: $BODY2"
fi
AGENT_ID_2=$(echo "$BODY2" | jq -r '.agent_id')
# object_permission may be null or absent
pass "Created agent without object_permission: $AGENT_ID_2"
# --- 6. Cleanup: delete both agents ---
info "Deleting test agents..."
for AID in "$AGENT_ID" "$AGENT_ID_2"; do
DEL_CODE=$(curl -s -o /dev/null -w "%{http_code}" -X DELETE "$BASE_URL/v1/agents/$AID" -H "$AUTH_HEADER")
if [[ "$DEL_CODE" != "200" ]]; then
info "DELETE /v1/agents/$AID returned $DEL_CODE (non-fatal)"
fi
done
pass "Cleanup done"
echo ""
echo -e "${GREEN}All endpoint checks passed.${NC}"

View file

@ -122,6 +122,249 @@ async def test_handle_lpop_count_for_older_redis_versions(monkeypatch):
assert mock_pipeline.execute.call_count == 2
@pytest.mark.asyncio
async def test_async_rpush_pipeline_executes_all_operations(monkeypatch, redis_no_ping):
"""Verify that multiple rpush ops are batched into a single pipeline execute"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.rpush = MagicMock()
mock_pipeline.execute = AsyncMock(return_value=[3, 5, 1])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineRpushOperation
rpush_list = [
RedisPipelineRpushOperation(key="key1", values=["a", "b"]),
RedisPipelineRpushOperation(key="key2", values=["c"]),
RedisPipelineRpushOperation(key="key3", values=["d", "e", "f"]),
]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
result = await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
assert result == [3, 5, 1]
assert mock_pipeline.rpush.call_count == 3
mock_pipeline.rpush.assert_any_call("key1", "a", "b")
mock_pipeline.rpush.assert_any_call("key2", "c")
mock_pipeline.rpush.assert_any_call("key3", "d", "e", "f")
mock_pipeline.execute.assert_called_once()
@pytest.mark.asyncio
async def test_async_rpush_pipeline_empty_list_returns_empty(monkeypatch, redis_no_ping):
"""Empty rpush_list should return empty list without touching Redis"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
result = await redis_cache.async_rpush_pipeline(rpush_list=[])
assert result == []
mock_redis_instance.pipeline.assert_not_called()
@pytest.mark.asyncio
async def test_async_rpush_pipeline_raises_on_redis_error(monkeypatch, redis_no_ping):
"""Pipeline errors should propagate"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.rpush = MagicMock()
mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down"))
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineRpushOperation
rpush_list = [RedisPipelineRpushOperation(key="key1", values=["a"])]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
with pytest.raises(ConnectionError, match="Redis down"):
await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
@pytest.mark.asyncio
async def test_async_lpop_pipeline_single_round_trip(monkeypatch, redis_no_ping):
"""Verify that multiple lpop ops are batched into a single pipeline execute"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "7.0.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
mock_pipeline.execute = AsyncMock(return_value=[
[b"val1", b"val2"], # key1 results
None, # key2 empty
[b"val3"], # key3 results
])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [
RedisPipelineLpopOperation(key="key1", count=10),
RedisPipelineLpopOperation(key="key2", count=10),
RedisPipelineLpopOperation(key="key3", count=5),
]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
assert len(results) == 3
assert results[0] == ["val1", "val2"]
assert results[1] is None
assert results[2] == ["val3"]
mock_pipeline.execute.assert_called_once()
@pytest.mark.asyncio
async def test_async_lpop_pipeline_redis_lt7_regroups_flat_results(monkeypatch, redis_no_ping):
"""Verify Redis < 7 fallback issues individual LPOPs and regroups correctly"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "6.2.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
# With count=3 for key1 and count=2 for key2, we get 5 individual LPOP commands
# Simulate: key1 has 2 values then None, key2 has 1 value then None
mock_pipeline.execute = AsyncMock(return_value=[
b"val1", b"val2", None, # 3 LPOPs for key1
b"val3", None, # 2 LPOPs for key2
])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [
RedisPipelineLpopOperation(key="key1", count=3),
RedisPipelineLpopOperation(key="key2", count=2),
]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
results = await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
assert len(results) == 2
assert results[0] == ["val1", "val2"] # 2 values, None filtered out
assert results[1] == ["val3"] # 1 value, None filtered out
# All 5 individual LPOPs should be queued, but only 1 execute() call
assert mock_pipeline.lpop.call_count == 5
mock_pipeline.execute.assert_called_once()
@pytest.mark.asyncio
async def test_async_rpush_pipeline_raises_on_per_command_error(monkeypatch, redis_no_ping):
"""Verify that per-command errors in pipeline results are raised, not silently dropped"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.rpush = MagicMock()
# Simulate: first RPUSH succeeds, second returns a per-command error
mock_pipeline.execute = AsyncMock(return_value=[3, Exception("WRONGTYPE")])
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineRpushOperation
rpush_list = [
RedisPipelineRpushOperation(key="key1", values=["a"]),
RedisPipelineRpushOperation(key="key2", values=["b"]),
]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
with pytest.raises(Exception, match="WRONGTYPE"):
await redis_cache.async_rpush_pipeline(rpush_list=rpush_list)
@pytest.mark.asyncio
async def test_async_lpop_pipeline_raises_on_per_command_error(monkeypatch, redis_no_ping):
"""Verify that per-command errors in LPOP pipeline results are raised, not silently dropped"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "7.0.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
# Simulate: first LPOP succeeds, second returns a per-command error
mock_pipeline.execute = AsyncMock(
return_value=[[b"val1"], Exception("WRONGTYPE")]
)
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [
RedisPipelineLpopOperation(key="key1", count=10),
RedisPipelineLpopOperation(key="key2", count=10),
]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
with pytest.raises(Exception, match="WRONGTYPE"):
await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
@pytest.mark.asyncio
async def test_async_lpop_pipeline_empty_list(monkeypatch, redis_no_ping):
"""Empty lpop_list should return empty list without touching Redis"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
result = await redis_cache.async_lpop_pipeline(lpop_list=[])
assert result == []
mock_redis_instance.pipeline.assert_not_called()
@pytest.mark.asyncio
async def test_async_lpop_pipeline_propagates_redis_exception(monkeypatch, redis_no_ping):
"""Pipeline errors should propagate"""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
redis_cache.redis_version = "7.0.0"
mock_redis_instance = AsyncMock()
mock_pipeline = MagicMock()
mock_pipeline.__aenter__ = AsyncMock(return_value=mock_pipeline)
mock_pipeline.__aexit__ = AsyncMock(return_value=None)
mock_pipeline.lpop = MagicMock()
mock_pipeline.execute = AsyncMock(side_effect=ConnectionError("Redis down"))
mock_redis_instance.pipeline = MagicMock(return_value=mock_pipeline)
from litellm.types.caching import RedisPipelineLpopOperation
lpop_list = [RedisPipelineLpopOperation(key="key1", count=10)]
with patch.object(redis_cache, "init_async_client", return_value=mock_redis_instance):
with pytest.raises(ConnectionError, match="Redis down"):
await redis_cache.async_lpop_pipeline(lpop_list=lpop_list)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"redis_version",

View file

@ -1595,3 +1595,146 @@ async def test_get_allowed_mcp_servers_for_key_prefers_in_memory_permission():
assert set(result) == {"direct-server", "group-server"}
mock_get_perm.assert_not_called()
mock_access_groups.assert_called_once_with(["grp-alpha"])
@pytest.mark.asyncio
class TestAgentMCPPermissions:
"""Test agent-level MCP server and tool permission intersection."""
async def test_get_allowed_mcp_servers_agent_intersection(self):
"""Key/team allow [server_1, server_2]; agent allows [server_1]. Result = [server_1]."""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="test-team",
agent_id="agent-123",
)
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
) as mock_key:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
) as mock_team:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent"
) as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = ["server_1"]
result = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth
)
assert sorted(result) == ["server_1"]
mock_agent.assert_called_once_with(user_api_key_auth)
async def test_get_allowed_mcp_servers_agent_no_restriction(self):
"""Agent with no object_permission returns []; no intersection applied (inherit key/team)."""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
agent_id="agent-456",
)
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
) as mock_key:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
) as mock_team:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent"
) as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = [] # no agent-level restriction
result = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth
)
assert sorted(result) == ["server_1", "server_2"]
mock_agent.assert_called_once_with(user_api_key_auth)
async def test_get_allowed_mcp_servers_key_team_agent_intersection(self):
"""Key allows [1, 2], agent allows [2, 3]. Result = [2]."""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
agent_id="agent-789",
)
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_key"
) as mock_key:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_team"
) as mock_team:
with patch.object(
MCPRequestHandler, "_get_allowed_mcp_servers_for_agent"
) as mock_agent:
mock_key.return_value = ["server_1", "server_2"]
mock_team.return_value = []
mock_agent.return_value = ["server_2", "server_3"]
result = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth
)
assert sorted(result) == ["server_2"]
async def test_get_allowed_tools_for_server_agent_intersection(self):
"""Key allows [tool_a, tool_b], agent allows [tool_a]. Result = [tool_a]."""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
agent_id="agent-tools",
)
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
team_perm = None
with patch.object(
MCPRequestHandler, "_get_key_object_permission", return_value=key_perm
):
with patch.object(
MCPRequestHandler, "_get_team_object_permission",
new_callable=AsyncMock,
return_value=team_perm,
):
with patch.object(
MCPRequestHandler,
"_get_agent_tool_permissions_for_server",
new_callable=AsyncMock,
return_value=["tool_a"],
) as mock_agent_tools:
result = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server_1",
user_api_key_auth=user_api_key_auth,
)
assert result == ["tool_a"]
mock_agent_tools.assert_called_once()
call_kwargs = mock_agent_tools.call_args.kwargs
assert call_kwargs["server_id"] == "server_1"
assert call_kwargs["user_api_key_auth"] == user_api_key_auth
async def test_get_allowed_tools_for_server_agent_no_restriction(self):
"""Agent has no tool permissions for server; key/team result is unchanged."""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
agent_id="agent-no-tools",
)
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]}
with patch.object(
MCPRequestHandler, "_get_key_object_permission", return_value=key_perm
):
with patch.object(
MCPRequestHandler, "_get_team_object_permission",
new_callable=AsyncMock,
return_value=None,
):
with patch.object(
MCPRequestHandler,
"_get_agent_tool_permissions_for_server",
new_callable=AsyncMock,
return_value=None,
):
result = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server_1",
user_api_key_auth=user_api_key_auth,
)
assert sorted(result) == ["tool_a", "tool_b"]

View file

@ -0,0 +1,194 @@
import json
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer
from litellm.types.caching import RedisPipelineRpushOperation
@pytest.fixture
def mock_redis_cache():
"""Create a mock RedisCache instance"""
mock = AsyncMock()
return mock
@pytest.fixture
def redis_update_buffer(mock_redis_cache):
"""Create a RedisUpdateBuffer with a mock RedisCache"""
return RedisUpdateBuffer(redis_cache=mock_redis_cache)
@pytest.mark.asyncio
async def test_store_in_memory_spend_updates_uses_pipeline(redis_update_buffer, mock_redis_cache):
"""
Verify store_in_memory_spend_updates_in_redis calls async_rpush_pipeline once
with the correct operations and skips empty queues.
"""
mock_redis_cache.async_rpush_pipeline = AsyncMock(return_value=[3, 5, 2])
# Create mock queues - only 3 of 7 have data
spend_update_queue = AsyncMock()
spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions = AsyncMock(
return_value={"key_list_transactions": {"key1": 1.0}}
)
daily_spend_queue = AsyncMock()
daily_spend_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value={"user_key1": {"spend": 1.0}}
)
daily_team_queue = AsyncMock()
daily_team_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value={"team_key1": {"spend": 2.0}}
)
# Empty queues
daily_org_queue = AsyncMock()
daily_org_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value={}
)
daily_end_user_queue = AsyncMock()
daily_end_user_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value=None
)
daily_agent_queue = AsyncMock()
daily_agent_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value={}
)
daily_tag_queue = AsyncMock()
daily_tag_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value={}
)
await redis_update_buffer.store_in_memory_spend_updates_in_redis(
spend_update_queue=spend_update_queue,
daily_spend_update_queue=daily_spend_queue,
daily_team_spend_update_queue=daily_team_queue,
daily_org_spend_update_queue=daily_org_queue,
daily_end_user_spend_update_queue=daily_end_user_queue,
daily_agent_spend_update_queue=daily_agent_queue,
daily_tag_spend_update_queue=daily_tag_queue,
)
# Should be called exactly once (pipeline)
mock_redis_cache.async_rpush_pipeline.assert_called_once()
# Verify only 3 operations were included (empty ones skipped)
call_args = mock_redis_cache.async_rpush_pipeline.call_args
rpush_list = call_args.kwargs["rpush_list"]
assert len(rpush_list) == 3
@pytest.mark.asyncio
async def test_store_in_memory_spend_updates_all_empty_returns_early(
redis_update_buffer, mock_redis_cache
):
"""
When all queues are empty, pipeline should never be called.
"""
mock_redis_cache.async_rpush_pipeline = AsyncMock()
# All queues return empty
empty_queue = AsyncMock()
empty_queue.flush_and_get_aggregated_db_spend_update_transactions = AsyncMock(
return_value={}
)
empty_daily_queue = AsyncMock()
empty_daily_queue.flush_and_get_aggregated_daily_spend_update_transactions = AsyncMock(
return_value={}
)
await redis_update_buffer.store_in_memory_spend_updates_in_redis(
spend_update_queue=empty_queue,
daily_spend_update_queue=empty_daily_queue,
daily_team_spend_update_queue=empty_daily_queue,
daily_org_spend_update_queue=empty_daily_queue,
daily_end_user_spend_update_queue=empty_daily_queue,
daily_agent_spend_update_queue=empty_daily_queue,
daily_tag_spend_update_queue=empty_daily_queue,
)
mock_redis_cache.async_rpush_pipeline.assert_not_called()
@pytest.mark.asyncio
async def test_get_all_transactions_from_redis_buffer_pipeline(
redis_update_buffer, mock_redis_cache
):
"""
Verify get_all_transactions_from_redis_buffer_pipeline correctly parses
and aggregates results from async_lpop_pipeline.
"""
# Simulate pipeline results: slot 0 = spend updates, slots 1-6 = daily categories
db_spend_json = json.dumps(
{
"key_list_transactions": {"key1": 1.0, "key2": 2.0},
"user_list_transactions": {"user1": 0.5},
"end_user_list_transactions": {},
"team_list_transactions": {},
"team_member_list_transactions": {},
"org_list_transactions": {},
"tag_list_transactions": {},
}
)
daily_user_json = json.dumps({"user_key1": {"spend": 1.0, "api_requests": 1}})
daily_team_json = json.dumps({"team_key1": {"spend": 2.0, "api_requests": 2}})
mock_redis_cache.async_lpop_pipeline = AsyncMock(
return_value=[
[db_spend_json], # slot 0: db spend updates
[daily_user_json], # slot 1: daily user
[daily_team_json], # slot 2: daily team
None, # slot 3: daily org (empty)
None, # slot 4: daily end-user (empty)
None, # slot 5: daily agent (empty)
None, # slot 6: daily tag (empty)
]
)
result = await redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline()
assert len(result) == 7
db_spend, daily_user, daily_team, daily_org, daily_end_user, daily_agent, daily_tag = result
# Verify db spend was parsed correctly
assert db_spend is not None
assert db_spend["key_list_transactions"]["key1"] == 1.0
assert db_spend["key_list_transactions"]["key2"] == 2.0
assert db_spend["user_list_transactions"]["user1"] == 0.5
# Verify daily user was parsed
assert daily_user is not None
assert daily_user["user_key1"]["spend"] == 1.0
# Verify daily team was parsed
assert daily_team is not None
assert daily_team["team_key1"]["spend"] == 2.0
# Verify empty slots
assert daily_org is None
assert daily_end_user is None
assert daily_agent is None
assert daily_tag is None
# Verify pipeline was called once with correct keys
mock_redis_cache.async_lpop_pipeline.assert_called_once()
@pytest.mark.asyncio
async def test_get_all_transactions_from_redis_buffer_pipeline_no_redis():
"""When redis_cache is None, should return all Nones"""
buffer = RedisUpdateBuffer(redis_cache=None)
result = await buffer.get_all_transactions_from_redis_buffer_pipeline()
assert result == (None, None, None, None, None, None, None)

View file

@ -1,3 +1,4 @@
import asyncio
import json
import os
import sys
@ -51,6 +52,9 @@ async def test_daily_spend_tracking_with_disabled_spend_logs():
# Call the method
await db_writer.update_database(**test_data)
# Let the single batched task run
await asyncio.sleep(0)
# Verify that _insert_spend_log_to_db was NOT called (since disable_spend_logs is True)
db_writer._insert_spend_log_to_db.assert_not_called()
@ -115,7 +119,9 @@ async def test_update_daily_spend_with_null_entity_id():
# Verify the where clause contains null entity_id
call_args = mock_table.upsert.call_args[1]
where_clause = call_args["where"]["user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint"]
where_clause = call_args["where"][
"user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint"
]
assert where_clause["user_id"] is None
assert where_clause["date"] == "2024-01-01"
assert where_clause["api_key"] == "test-api-key"
@ -161,7 +167,7 @@ async def test_update_daily_spend_sorting():
upsert_calls = []
for i in range(50):
daily_spend_transactions[f"test_key_{i}"] = {
"user_id": f"user{60-i}", # user60 ... user11, reverse order
"user_id": f"user{60-i}", # user60 ... user11, reverse order
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
@ -173,46 +179,48 @@ async def test_update_daily_spend_sorting():
"successful_requests": 1,
"failed_requests": 0,
}
upsert_calls.append(call(
where={
"user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint": {
"user_id": f"user{i+11}", # user11 ... user60, sorted order
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
"custom_llm_provider": "openai",
"mcp_namespaced_tool_name": "",
"endpoint": "",
}
},
data={
"create": {
"user_id": f"user{i+11}",
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
"model_group": None,
"mcp_namespaced_tool_name": "",
"custom_llm_provider": "openai",
"endpoint": "",
"prompt_tokens": 10,
"completion_tokens": 20,
"spend": 0.1,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
upsert_calls.append(
call(
where={
"user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint": {
"user_id": f"user{i+11}", # user11 ... user60, sorted order
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
"custom_llm_provider": "openai",
"mcp_namespaced_tool_name": "",
"endpoint": "",
}
},
"update": {
"prompt_tokens": {"increment": 10},
"completion_tokens": {"increment": 20},
"spend": {"increment": 0.1},
"api_requests": {"increment": 1},
"successful_requests": {"increment": 1},
"failed_requests": {"increment": 0},
"endpoint": "",
data={
"create": {
"user_id": f"user{i+11}",
"date": "2024-01-01",
"api_key": "test-api-key",
"model": "gpt-4",
"model_group": None,
"mcp_namespaced_tool_name": "",
"custom_llm_provider": "openai",
"endpoint": "",
"prompt_tokens": 10,
"completion_tokens": 20,
"spend": 0.1,
"api_requests": 1,
"successful_requests": 1,
"failed_requests": 0,
},
"update": {
"prompt_tokens": {"increment": 10},
"completion_tokens": {"increment": 20},
"spend": {"increment": 0.1},
"api_requests": {"increment": 1},
"successful_requests": {"increment": 1},
"failed_requests": {"increment": 0},
"endpoint": "",
},
},
},
))
)
)
# Call the method
await DBSpendUpdateWriter._update_daily_spend(
@ -275,7 +283,7 @@ async def test_update_daily_spend_tag_with_request_id():
# Verify that table.upsert was called
mock_table.upsert.assert_called_once()
# Verify request_id is in update_data
call_args = mock_table.upsert.call_args[1]
update_data = call_args["data"]["update"]
@ -283,15 +291,13 @@ async def test_update_daily_spend_tag_with_request_id():
assert update_data["request_id"] == "test-request-id-123"
@pytest.mark.asyncio
async def test_update_daily_spend_with_none_values_in_sorting_fields():
"""
Test that _update_daily_spend handles None values in sorting fields correctly.
This test ensures that when fields like date, api_key, model, or custom_llm_provider
are None, the sorting doesn't crash with TypeError: '<' not supported between
are None, the sorting doesn't crash with TypeError: '<' not supported between
instances of 'NoneType' and 'str'.
"""
# Setup
@ -509,6 +515,7 @@ async def test_update_tag_db_without_prisma_client():
assert writer.spend_update_queue.add_update.call_count == 0
@pytest.mark.asyncio
async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_id():
"""
@ -518,7 +525,7 @@ async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_i
writer = DBSpendUpdateWriter()
mock_prisma = MagicMock()
mock_prisma.get_request_status = MagicMock(return_value="success")
request_id = "test-request-id-123"
payload = {
"request_id": request_id,
@ -546,13 +553,15 @@ async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_i
# Should be called twice (once for each tag)
assert writer.daily_tag_spend_update_queue.add_update.call_count == 2
# Check that request_id is included in both transactions
for call in writer.daily_tag_spend_update_queue.add_update.call_args_list:
transaction_dict = call[1]["update"]
# Each transaction should have one key with the format tag_date_api_key_model_provider
for key, transaction in transaction_dict.items():
assert transaction["request_id"] == request_id, f"request_id should be {request_id} but got {transaction.get('request_id')}"
assert (
transaction["request_id"] == request_id
), f"request_id should be {request_id} but got {transaction.get('request_id')}"
@pytest.mark.asyncio
@ -866,11 +875,11 @@ async def test_endpoint_field_is_correctly_mapped_from_call_type():
call_args = writer.daily_spend_update_queue.add_update.call_args[1]
update_dict = call_args["update"]
assert len(update_dict) == 1
for key, transaction in update_dict.items():
# Verify endpoint is included in the key
assert key == f"test-user_2024-01-01_test-key_gpt-4_openai_/chat/completions"
# Verify endpoint is set in the transaction
assert transaction["endpoint"] == "/chat/completions"
assert transaction["user_id"] == "test-user"
@ -887,7 +896,7 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure():
This ensures proper debugging information is available for issues like unique constraint violations.
"""
from litellm._logging import verbose_proxy_logger
# Setup
mock_prisma_client = MagicMock()
mock_batcher = MagicMock()
@ -895,13 +904,13 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure():
mock_batch_context = MagicMock()
mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher)
mock_batcher.litellm_dailyuserspend = mock_table
# Make the batch context manager's exit raise an exception
# This simulates a batch commit failure (e.g., unique constraint violation)
test_exception = Exception("Unique constraint violation")
mock_batch_context.__aexit__ = AsyncMock(side_effect=test_exception)
mock_prisma_client.db.batch_.return_value = mock_batch_context
# Create a transaction
daily_spend_transactions = {
"test_key": {
@ -918,13 +927,13 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure():
"failed_requests": 0,
}
}
# Create a mock proxy_logging_obj with failure_handler as AsyncMock
mock_proxy_logging = MagicMock()
mock_proxy_logging.failure_handler = AsyncMock()
# Mock the logger to capture exception calls
with patch.object(verbose_proxy_logger, 'exception') as mock_exception_logger:
with patch.object(verbose_proxy_logger, "exception") as mock_exception_logger:
# Call the method and expect it to raise the exception
with pytest.raises(Exception, match="Unique constraint violation"):
await DBSpendUpdateWriter._update_daily_spend(
@ -937,13 +946,16 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure():
table_name="litellm_dailyuserspend",
unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint",
)
# Verify that exception was logged with detailed information
assert mock_exception_logger.called
call_args = mock_exception_logger.call_args[0][0]
assert "Daily user spend batch upsert failed" in call_args
assert "Table: litellm_dailyuserspend" in call_args
assert "Constraint: user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint" in call_args
assert (
"Constraint: user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint"
in call_args
)
assert "Batch size: 1" in call_args
assert "Unique constraint violation" in call_args
@ -961,7 +973,7 @@ async def test_update_daily_spend_re_raises_exception_after_logging():
mock_batch_context = MagicMock()
mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher)
mock_batcher.litellm_dailyuserspend = mock_table
# Create a transaction
daily_spend_transactions = {
"test_key": {
@ -978,16 +990,16 @@ async def test_update_daily_spend_re_raises_exception_after_logging():
"failed_requests": 0,
}
}
# Create a custom exception to verify it's re-raised
custom_exception = ValueError("Database connection lost")
mock_batch_context.__aexit__ = AsyncMock(side_effect=custom_exception)
mock_prisma_client.db.batch_.return_value = mock_batch_context
# Create a mock proxy_logging_obj with failure_handler as AsyncMock
mock_proxy_logging = MagicMock()
mock_proxy_logging.failure_handler = AsyncMock()
# Verify the exception is re-raised
with pytest.raises(ValueError, match="Database connection lost"):
await DBSpendUpdateWriter._update_daily_spend(
@ -1018,10 +1030,12 @@ async def test_commit_key_spend_updates_includes_last_active():
mock_transaction = AsyncMock()
mock_transaction.__aenter__ = AsyncMock(return_value=mock_transaction)
mock_transaction.__aexit__ = AsyncMock(return_value=False)
mock_transaction.batch_ = MagicMock(return_value=AsyncMock(
__aenter__=AsyncMock(return_value=mock_batcher),
__aexit__=AsyncMock(return_value=False),
))
mock_transaction.batch_ = MagicMock(
return_value=AsyncMock(
__aenter__=AsyncMock(return_value=mock_batcher),
__aexit__=AsyncMock(return_value=False),
)
)
mock_prisma_client = MagicMock()
mock_prisma_client.db = MagicMock()
@ -1049,9 +1063,7 @@ async def test_commit_key_spend_updates_includes_last_active():
before_call = datetime.now(timezone.utc)
with patch(
"litellm.proxy.utils._raise_failed_update_spend_exception"
):
with patch("litellm.proxy.utils._raise_failed_update_spend_exception"):
await db_writer._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=0,
@ -1076,3 +1088,212 @@ async def test_commit_key_spend_updates_includes_last_active():
last_active = call_kwargs["data"]["last_active"]
assert isinstance(last_active, datetime)
assert before_call <= last_active <= after_call
@pytest.mark.asyncio
async def test_update_database_creates_single_task():
"""
Test that update_database() fires exactly 1 asyncio.create_task() call
(the batched task) instead of the previous 11.
"""
db_writer = DBSpendUpdateWriter()
# Mock all helpers so nothing real runs
db_writer._insert_spend_log_to_db = AsyncMock()
db_writer._batch_database_updates = AsyncMock()
with patch("litellm.proxy.proxy_server.disable_spend_logs", False), patch(
"litellm.proxy.proxy_server.prisma_client", MagicMock()
), patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch(
"litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"
), patch(
"litellm.proxy.db.db_spend_update_writer.asyncio.create_task"
) as mock_create_task:
await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id="test-end-user",
start_time=datetime.now(),
end_time=datetime.now(),
team_id="test-team",
org_id="test-org",
completion_response=MagicMock(),
response_cost=0.1,
kwargs={"model": "gpt-4", "custom_llm_provider": "openai"},
)
# Exactly 1 create_task call (the batch), not 11
assert mock_create_task.call_count == 1
@pytest.mark.asyncio
async def test_batch_database_updates_isolation_on_failure():
"""
Test that if one helper inside _batch_database_updates raises,
all other helpers still execute.
"""
db_writer = DBSpendUpdateWriter()
# Make _update_key_db raise
db_writer._update_key_db = AsyncMock(side_effect=RuntimeError("key db boom"))
# All other helpers are normal mocks
db_writer._update_user_db = AsyncMock()
db_writer._update_team_db = AsyncMock()
db_writer._update_org_db = AsyncMock()
db_writer._update_tag_db = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_user_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_end_user_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_agent_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_team_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_org_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_tag_transaction = AsyncMock()
await db_writer._batch_database_updates(
response_cost=0.1,
user_id="u1",
hashed_token="t1",
team_id="team1",
org_id="org1",
end_user_id="eu1",
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
litellm_proxy_budget_name="budget",
payload_copy={"key": "value"},
request_tags=None,
)
# _update_key_db raised, but all others should still have been called
db_writer._update_user_db.assert_awaited_once()
db_writer._update_key_db.assert_awaited_once()
db_writer._update_team_db.assert_awaited_once()
db_writer._update_org_db.assert_awaited_once()
db_writer._update_tag_db.assert_awaited_once()
db_writer.add_spend_log_transaction_to_daily_user_transaction.assert_awaited_once()
db_writer.add_spend_log_transaction_to_daily_end_user_transaction.assert_awaited_once()
db_writer.add_spend_log_transaction_to_daily_agent_transaction.assert_awaited_once()
db_writer.add_spend_log_transaction_to_daily_team_transaction.assert_awaited_once()
db_writer.add_spend_log_transaction_to_daily_org_transaction.assert_awaited_once()
db_writer.add_spend_log_transaction_to_daily_tag_transaction.assert_awaited_once()
@pytest.mark.asyncio
async def test_daily_agent_receives_deepcopied_payload():
"""
Test that the daily agent handler receives a deepcopied payload (not the original).
Previously, add_spend_log_transaction_to_daily_agent_transaction received the raw
payload without a deepcopy, which was a mutation bug. This test goes through
update_database() to verify the production deepcopy path.
"""
db_writer = DBSpendUpdateWriter()
# Capture the payload object that get_logging_payload returns (the "original")
# and the payload the agent handler receives (should be a deepcopy)
original_payload_ref = {}
captured_agent_payloads = []
async def capture_agent_payload(**kwargs):
captured_agent_payloads.append(kwargs.get("payload"))
# Mock all helpers
db_writer._insert_spend_log_to_db = AsyncMock()
db_writer._update_user_db = AsyncMock()
db_writer._update_key_db = AsyncMock()
db_writer._update_team_db = AsyncMock()
db_writer._update_org_db = AsyncMock()
db_writer._update_tag_db = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_user_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_end_user_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_agent_transaction = AsyncMock(
side_effect=capture_agent_payload
)
db_writer.add_spend_log_transaction_to_daily_team_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_org_transaction = AsyncMock()
db_writer.add_spend_log_transaction_to_daily_tag_transaction = AsyncMock()
# Mock get_logging_payload to return a known dict and capture its identity
fake_payload = {
"startTime": "2024-01-01T00:00:00",
"endTime": "2024-01-01T00:01:00",
"model": "gpt-4",
"custom_llm_provider": "openai",
"spend": 0.0,
"nested": {"a": 1},
}
original_payload_ref["obj"] = fake_payload # store reference to the original
with patch("litellm.proxy.proxy_server.disable_spend_logs", True), patch(
"litellm.proxy.proxy_server.prisma_client", MagicMock()
), patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), patch(
"litellm.proxy.proxy_server.litellm_proxy_budget_name", "test-budget"
), patch(
"litellm.proxy.spend_tracking.spend_tracking_utils.get_logging_payload",
return_value=fake_payload,
):
await db_writer.update_database(
token="test-token",
user_id="test-user",
end_user_id="test-end-user",
team_id="test-team",
org_id="test-org",
kwargs={"model": "gpt-4", "custom_llm_provider": "openai"},
completion_response=MagicMock(),
start_time=datetime.now(),
end_time=datetime.now(),
response_cost=0.1,
)
# Let the single batched task run
await asyncio.sleep(0)
# The agent handler should have been called
assert len(captured_agent_payloads) == 1
# The payload must NOT be the same object as the original (deepcopy occurred)
assert captured_agent_payloads[0] is not original_payload_ref["obj"]
# But it should have equivalent content
assert captured_agent_payloads[0]["model"] == "gpt-4"
assert captured_agent_payloads[0]["spend"] == 0.1
@pytest.mark.asyncio
async def test_commit_spend_updates_uses_pipeline():
"""
Verify that _commit_spend_updates_to_db_with_redis uses
get_all_transactions_from_redis_buffer_pipeline instead of 7 individual calls.
"""
db_writer = DBSpendUpdateWriter()
mock_redis_update_buffer = AsyncMock()
mock_redis_update_buffer.store_in_memory_spend_updates_in_redis = AsyncMock()
# Return all-None tuple (no data to commit)
mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline = AsyncMock(
return_value=(None, None, None, None, None, None, None)
)
db_writer.redis_update_buffer = mock_redis_update_buffer
mock_pod_lock_manager = AsyncMock()
mock_pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
mock_pod_lock_manager.release_lock = AsyncMock()
db_writer.pod_lock_manager = mock_pod_lock_manager
mock_prisma_client = MagicMock()
mock_proxy_logging = MagicMock()
await db_writer._commit_spend_updates_to_db_with_redis(
prisma_client=mock_prisma_client,
n_retry_times=1,
proxy_logging_obj=mock_proxy_logging,
)
# Pipeline method should be called once
mock_redis_update_buffer.get_all_transactions_from_redis_buffer_pipeline.assert_called_once()
# Individual methods should NOT be called
mock_redis_update_buffer.get_all_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_team_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_org_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_end_user_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_agent_spend_update_transactions_from_redis_buffer.assert_not_called()
mock_redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer.assert_not_called()

View file

@ -0,0 +1,106 @@
"""
Unit Tests for the max iterations limiter for the proxy.
Tests that session-scoped iteration counting works correctly:
- Enforces max_iterations per session_id
- Different sessions have independent counters
"""
import pytest
from fastapi import HTTPException
from litellm.caching.caching import DualCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.hooks.max_iterations_limiter import _PROXY_MaxIterationsHandler
from litellm.proxy.utils import InternalUsageCache
@pytest.mark.asyncio
async def test_max_iterations_basic_enforcement():
"""
Test that max_iterations is enforced per session_id.
- 3 requests with the same session_id should succeed when max_iterations=3
- 4th request should raise 429
"""
local_cache = DualCache()
handler = _PROXY_MaxIterationsHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-1234", metadata={"max_iterations": 3}
)
# First 3 requests should succeed
for i in range(3):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-abc"}},
call_type="",
)
# 4th request should fail with 429
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-abc"}},
call_type="",
)
assert exc_info.value.status_code == 429
assert "max_iterations" in str(exc_info.value.detail).lower()
@pytest.mark.asyncio
async def test_max_iterations_different_sessions_independent():
"""
Test that different session_ids have independent iteration counters.
- Session A and Session B each get their own max_iterations budget
- Exhausting Session A does not affect Session B
"""
local_cache = DualCache()
handler = _PROXY_MaxIterationsHandler(
internal_usage_cache=InternalUsageCache(local_cache),
)
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-test-key-5678", metadata={"max_iterations": 2}
)
# Session A: 2 calls succeed
for _ in range(2):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
# Session B: 2 calls succeed (independent counter)
for _ in range(2):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)
# Session A: 3rd call fails
with pytest.raises(HTTPException) as exc_info:
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-A"}},
call_type="",
)
assert exc_info.value.status_code == 429
# Session B: 3rd call also fails
with pytest.raises(HTTPException):
await handler.async_pre_call_hook(
user_api_key_dict=user_api_key_dict,
cache=local_cache,
data={"metadata": {"session_id": "session-B"}},
call_type="",
)

View file

@ -274,6 +274,7 @@ ignored_keys = [
"endTime",
"completionStartTime",
"endTime",
"request_duration_ms",
"organization_id",
"metadata.model_map_information",
"metadata.usage_object",
@ -606,6 +607,82 @@ async def test_ui_view_spend_logs_sort_validation_errors(
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_sort_by_request_duration_ms(client, monkeypatch):
"""Test that request_duration_ms is accepted as a valid sort_by field."""
base_logs = [
{
"request_id": "req_fast",
"api_key": "sk-test-key",
"user": "user1",
"spend": 0.10,
"total_tokens": 100,
"request_duration_ms": 100,
"startTime": "2025-01-01T00:00:00+00:00",
"endTime": "2025-01-01T00:00:00.100000+00:00",
"model": "gpt-4",
},
{
"request_id": "req_slow",
"api_key": "sk-test-key",
"user": "user1",
"spend": 0.05,
"total_tokens": 50,
"request_duration_ms": 5000,
"startTime": "2025-01-01T00:00:01+00:00",
"endTime": "2025-01-01T00:00:06+00:00",
"model": "gpt-4",
},
]
async def mock_count(*args, **kwargs):
return len(base_logs)
async def mock_query_raw(sql_query, *params):
reverse = "DESC" in sql_query
sorted_logs = sorted(
base_logs, key=lambda x: x.get("request_duration_ms", 0), reverse=reverse
)
page_size = params[-2] if len(params) >= 2 else 50
skip = params[-1] if len(params) >= 1 else 0
return sorted_logs[skip : skip + page_size]
class MockPrismaClient:
def __init__(self):
self.db = MagicMock()
self.db.litellm_spendlogs = MagicMock()
self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count)
self.db.query_raw = AsyncMock(side_effect=mock_query_raw)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient())
monkeypatch.setattr(
"litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe",
lambda user_api_key_dict: True,
)
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user"
)
try:
response = client.get(
"/spend/logs/ui",
params={
"start_date": "2024-12-25 00:00:00",
"end_date": "2025-01-02 23:59:59",
"sort_by": "request_duration_ms",
"sort_order": "asc",
},
headers={"Authorization": "Bearer sk-test"},
)
assert response.status_code == 200, response.text
data = response.json()
actual_ids = [log["request_id"] for log in data["data"]]
assert actual_ids == ["req_fast", "req_slow"]
finally:
app.dependency_overrides.pop(ps.user_api_key_auth, None)
@pytest.mark.asyncio
async def test_ui_view_spend_logs_with_team_id(client, monkeypatch):
mock_spend_logs = [

View file

@ -20,6 +20,7 @@ from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD, REDACTED_BY_LITEL
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy.spend_tracking.spend_tracking_utils import (
_get_proxy_server_request_for_spend_logs_payload,
_get_request_duration_ms,
_get_response_for_spend_logs_payload,
_get_spend_logs_metadata,
_get_vector_store_request_for_spend_logs_payload,
@ -1232,3 +1233,50 @@ def test_get_logging_payload_handles_missing_retry_info_gracefully():
metadata.get("max_retries") is None
), "max_retries should be None when not provided"
def test_get_request_duration_ms_normal():
"""Test that request duration is correctly computed in milliseconds."""
start = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
end = datetime.datetime(2025, 1, 1, 0, 0, 2, 500000, tzinfo=timezone.utc) # 2.5s later
result = _get_request_duration_ms(start, end)
assert result == 2500
def test_get_request_duration_ms_sub_millisecond():
"""Test that sub-millisecond durations are truncated to int."""
start = datetime.datetime(2025, 1, 1, 0, 0, 0, 0, tzinfo=timezone.utc)
end = datetime.datetime(2025, 1, 1, 0, 0, 0, 500, tzinfo=timezone.utc) # 0.5ms
result = _get_request_duration_ms(start, end)
assert result == 0
def test_get_request_duration_ms_zero():
"""Test that identical start and end times produce 0."""
t = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
result = _get_request_duration_ms(t, t)
assert result == 0
def test_get_logging_payload_includes_request_duration_ms():
"""Test that get_logging_payload populates request_duration_ms."""
start_time = datetime.datetime(2025, 1, 1, 0, 0, 0, tzinfo=timezone.utc)
end_time = datetime.datetime(2025, 1, 1, 0, 0, 3, tzinfo=timezone.utc) # 3s later
kwargs = {
"model": "gpt-4",
"litellm_params": {"api_base": "https://api.openai.com"},
"standard_logging_object": None,
}
response_obj = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
with patch("litellm.proxy.proxy_server.master_key", None), \
patch("litellm.proxy.proxy_server.general_settings", {}):
payload = get_logging_payload(
kwargs=kwargs,
response_obj=response_obj,
start_time=start_time,
end_time=end_time,
)
assert payload["request_duration_ms"] == 3000

View file

@ -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",
@ -12995,6 +13056,21 @@
"type": "github",
"url": "https://github.com/sponsors/wooorm"
}
},
"node_modules/@next/swc-win32-ia32-msvc": {
"version": "14.2.33",
"resolved": "https://registry.npmjs.org/@next/swc-win32-ia32-msvc/-/swc-win32-ia32-msvc-14.2.33.tgz",
"integrity": "sha512-pc9LpGNKhJ0dXQhZ5QMmYxtARwwmWLpeocFmVG5Z0DzWq5Uf0izcI8tLc+qOpqxO1PWqZ5A7J1blrUIKrIFc7Q==",
"cpu": [
"ia32"
],
"optional": true,
"os": [
"win32"
],
"engines": {
"node": ">= 10"
}
}
}
}

View file

@ -1,7 +1,7 @@
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import React from "react";
import { describe, beforeEach, expect, it, vi } from "vitest";
import { describe, expect, it, vi } from "vitest";
import ModelRetrySettingsTab from "./ModelRetrySettingsTab";
// TabPanel requires a parent Tabs context in Tremor. We stub it to render children

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -2,15 +2,7 @@
import React, { useCallback, useDeferredValue, useEffect, useState } from "react";
import { Select, Switch, Tooltip } from "antd";
import { Select, 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";
@ -18,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";
@ -57,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,
},
}}
popupMatchSelectWidth={false}
options={POLICY_OPTIONS.map((o) => ({
value: o.value,
@ -134,7 +113,9 @@ export const ToolPolicies: React.FC<ToolPoliciesProps> = ({ accessToken }) => {
}
}, [accessToken]);
useEffect(() => { load(); }, [load]);
useEffect(() => {
load();
}, [load]);
useEffect(() => {
if (!isLiveTail) return;
@ -147,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 {
@ -179,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[] = [
{
@ -246,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">
@ -257,16 +237,29 @@ 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>
<div className="flex items-center gap-2">
<span className="text-sm font-medium text-gray-900">Live Tail</span>
<Switch color="green" checked={isLiveTail} onChange={setIsLiveTail} />
<Switch checked={isLiveTail} onChange={setIsLiveTail} />
</div>
<button
@ -274,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>
@ -283,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>
@ -310,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>
)}
@ -322,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>
@ -398,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>
)}

View file

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

View file

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

View file

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

View file

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

View file

@ -79,6 +79,7 @@ vi.mock("@ant-design/icons", async () => {
ShoppingCartOutlined: Icon,
TagsOutlined: Icon,
RobotOutlined: Icon,
UserOutlined: Icon,
LineChartOutlined: Icon,
BarChartOutlined: Icon,
};

View file

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

View file

@ -1,13 +1,13 @@
import React, { useState, useEffect } from "react";
import { Button } from "@tremor/react";
import { Modal } from "antd";
import { getAgentsList, deleteAgentCall } from "./networking";
import { Modal, Alert } from "antd";
import { getAgentsList, deleteAgentCall, keyListCall } from "./networking";
import AddAgentForm from "./agents/add_agent_form";
import AgentTable from "./agents/agent_table";
import AgentCardGrid from "./agents/agent_card_grid";
import { isAdminRole } from "@/utils/roles";
import AgentInfoView from "./agents/agent_info";
import NotificationsManager from "./molecules/notifications_manager";
import { Agent } from "./agents/types";
import { Agent, AgentKeyInfo } from "./agents/types";
interface AgentsPanelProps {
accessToken: string | null;
@ -20,6 +20,7 @@ interface AgentsResponse {
const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole }) => {
const [agentsList, setAgentsList] = useState<Agent[]>([]);
const [keyInfoMap, setKeyInfoMap] = useState<Record<string, AgentKeyInfo>>({});
const [isAddModalVisible, setIsAddModalVisible] = useState(false);
const [isLoading, setIsLoading] = useState(false);
const [isDeleting, setIsDeleting] = useState(false);
@ -36,8 +37,7 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole }) => {
setIsLoading(true);
try {
const response: AgentsResponse = await getAgentsList(accessToken);
console.log(`agents: ${JSON.stringify(response)}`);
setAgentsList(response.agents);
setAgentsList(response.agents || []);
} catch (error) {
console.error("Error fetching agents:", error);
} finally {
@ -45,10 +45,50 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole }) => {
}
};
const fetchKeysForAgents = async () => {
if (!accessToken) return;
try {
const { keys = [] } = await keyListCall(
accessToken,
null,
null,
null,
null,
null,
1,
500
);
const map: Record<string, AgentKeyInfo> = {};
for (const key of keys) {
const agentId = (key as { agent_id?: string }).agent_id;
if (agentId && !map[agentId]) {
map[agentId] = {
has_key: true,
key_alias: (key as { key_alias?: string }).key_alias,
token_prefix: (key as { token?: string }).token
? `${(key as { token: string }).token.slice(0, 8)}…`
: undefined,
};
}
}
setKeyInfoMap(map);
} catch (error) {
console.error("Error fetching keys for agents:", error);
}
};
useEffect(() => {
fetchAgents();
}, [accessToken]);
useEffect(() => {
if (accessToken && agentsList.length > 0) {
fetchKeysForAgents();
} else if (agentsList.length === 0) {
setKeyInfoMap({});
}
}, [accessToken, agentsList.length]);
const handleAddAgent = () => {
if (selectedAgentId) {
setSelectedAgentId(null);
@ -94,6 +134,13 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole }) => {
<div className="flex flex-col gap-2 mb-4">
<h1 className="text-2xl font-bold">Agents</h1>
<p className="text-sm text-gray-600">List of A2A-spec agents that are available to be used in your organization. Go to AI Hub, to make agents public.</p>
<Alert
message="Why do agents need keys?"
description="Keys scope access to an agent and allow it to call MCP tools. Assign a key when creating an agent or from the Virtual Keys page."
type="info"
showIcon
className="mb-3"
/>
<div className="mt-2">
<Button onClick={handleAddAgent} disabled={!accessToken}>
+ Add New Agent
@ -109,8 +156,9 @@ const AgentsPanel: React.FC<AgentsPanelProps> = ({ accessToken, userRole }) => {
isAdmin={isAdmin}
/>
) : (
<AgentTable
<AgentCardGrid
agentsList={agentsList}
keyInfoMap={keyInfoMap}
isLoading={isLoading}
onDeleteClick={handleDeleteClick}
accessToken={accessToken}

View file

@ -1,7 +1,7 @@
import React, { useState, useEffect } from "react";
import { Modal, Form, message, Select, Input, Steps, Radio, Tag, Divider } from "antd";
import { Button } from "@tremor/react";
import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined } from "@ant-design/icons";
import { CheckCircleFilled, KeyOutlined, RobotOutlined, AppstoreOutlined, InfoCircleOutlined } from "@ant-design/icons";
import CreatedKeyDisplay from "../shared/CreatedKeyDisplay";
import {
createAgentCall,
@ -9,11 +9,16 @@ import {
keyCreateForAgentCall,
keyListCall,
keyUpdateCall,
modelAvailableCall,
AgentCreateInfo,
} from "../networking";
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
import { getModelDisplayName } from "../key_team_helpers/fetch_available_models_team_key";
import AgentFormFields from "./agent_form_fields";
import DynamicAgentFormFields, { buildDynamicAgentData } from "./dynamic_agent_form_fields";
import { getDefaultFormValues, buildAgentDataFromForm } from "./agent_config";
import MCPServerSelector from "../mcp_server_management/MCPServerSelector";
import MCPToolPermissions from "../mcp_server_management/MCPToolPermissions";
const { Step } = Steps;
@ -32,6 +37,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
accessToken,
onSuccess,
}) => {
const { userId, userRole } = useAuthorized();
const [form] = Form.useForm();
const [currentStep, setCurrentStep] = useState(0);
const [isSubmitting, setIsSubmitting] = useState(false);
@ -46,6 +52,8 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
const [existingKeys, setExistingKeys] = useState<any[]>([]);
const [selectedExistingKey, setSelectedExistingKey] = useState<string | null>(null);
const [loadingKeys, setLoadingKeys] = useState(false);
const [availableModels, setAvailableModels] = useState<string[]>([]);
const [loadingModels, setLoadingModels] = useState(false);
// Step 2: results
const [createdAgentName, setCreatedAgentName] = useState<string>("");
@ -68,9 +76,9 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
fetchMetadata();
}, []);
// Fetch existing keys when assign key step becomes active
// Fetch existing keys when assign key step becomes active (step 2)
useEffect(() => {
if (currentStep === 1 && accessToken && existingKeys.length === 0) {
if (currentStep === 2 && accessToken && existingKeys.length === 0) {
const fetchKeys = async () => {
setLoadingKeys(true);
try {
@ -86,6 +94,31 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
}
}, [currentStep, accessToken]);
// Fetch available models when Assign Key step is active (same list as key generation)
useEffect(() => {
if (currentStep !== 2 || !accessToken || !userId || !userRole) return;
let cancelled = false;
setLoadingModels(true);
modelAvailableCall(accessToken, userId, userRole)
.then((response) => {
if (cancelled) return;
const modelsArray = response?.data ?? (Array.isArray(response) ? response : []);
const ids = modelsArray
.map((m: { id?: string; model_name?: string }) => m.id ?? m.model_name)
.filter(Boolean) as string[];
setAvailableModels(ids);
})
.catch((error) => {
if (!cancelled) console.error("Error fetching models:", error);
})
.finally(() => {
if (!cancelled) setLoadingModels(false);
});
return () => {
cancelled = true;
};
}, [currentStep, accessToken, userId, userRole]);
const selectedAgentTypeInfo = agentTypeMetadata.find(
(info) => info.agent_type === agentType
);
@ -156,8 +189,6 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
setIsSubmitting(true);
try {
// getFieldsValue(true) returns ALL preserved values including fields from
// unmounted steps; merge with any currently-mounted validated fields.
await form.validateFields();
const values = { ...form.getFieldsValue(true) };
const agentData = buildAgentData(values);
@ -167,6 +198,26 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
return;
}
// Build object_permission from MCP Tools step (allowed_mcp_servers_and_groups, mcp_tool_permissions)
const mcpServersAndGroups = values.allowed_mcp_servers_and_groups;
const mcpToolPermissions = values.mcp_tool_permissions || {};
if (
mcpServersAndGroups &&
(mcpServersAndGroups.servers?.length > 0 || mcpServersAndGroups.accessGroups?.length > 0) ||
Object.keys(mcpToolPermissions).length > 0
) {
agentData.object_permission = {};
if (mcpServersAndGroups?.servers?.length > 0) {
agentData.object_permission.mcp_servers = mcpServersAndGroups.servers;
}
if (mcpServersAndGroups?.accessGroups?.length > 0) {
agentData.object_permission.mcp_access_groups = mcpServersAndGroups.accessGroups;
}
if (Object.keys(mcpToolPermissions).length > 0) {
agentData.object_permission.mcp_tool_permissions = mcpToolPermissions;
}
}
const agentResponse = await createAgentCall(accessToken, agentData);
const agentId: string = agentResponse.agent_id;
const agentName: string = agentResponse.agent_name || values.agent_name || agentId;
@ -194,11 +245,12 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
setAssignedKeyAlias(keyInfo?.key_alias || selectedExistingKey.slice(0, 12) + "…");
}
setCurrentStep(2);
setCurrentStep(3);
onSuccess();
} catch (error) {
console.error("Error creating agent:", error);
message.error("Failed to create agent");
const errorMessage = error instanceof Error ? error.message : String(error);
message.error(errorMessage ? `Failed to create agent: ${errorMessage}` : "Failed to create agent");
} finally {
setIsSubmitting(false);
}
@ -218,6 +270,54 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
onClose();
};
const renderMCPToolsStep = () => (
<div className="space-y-4">
<p className="text-sm text-gray-600">
Optionally restrict which MCP servers and tools this agent can use. Leave empty to allow all (subject to key/team permissions).
</p>
<Form.Item
label={
<span>
Allowed MCP Servers{" "}
<InfoCircleOutlined title="Select which MCP servers or access groups this agent can access" style={{ marginLeft: "4px" }} />
</span>
}
name="allowed_mcp_servers_and_groups"
initialValue={{ servers: [], accessGroups: [] }}
>
<MCPServerSelector
onChange={(val: { servers?: string[]; accessGroups?: string[] }) =>
form.setFieldValue("allowed_mcp_servers_and_groups", val)
}
value={form.getFieldValue("allowed_mcp_servers_and_groups") || { servers: [], accessGroups: [] }}
accessToken={accessToken ?? ""}
placeholder="Select MCP servers or access groups (optional)"
/>
</Form.Item>
<Form.Item name="mcp_tool_permissions" initialValue={{}} hidden>
<Input type="hidden" />
</Form.Item>
<Form.Item
noStyle
shouldUpdate={(prev, curr) =>
prev.allowed_mcp_servers_and_groups !== curr.allowed_mcp_servers_and_groups ||
prev.mcp_tool_permissions !== curr.mcp_tool_permissions
}
>
{() => (
<div className="mt-4">
<MCPToolPermissions
accessToken={accessToken ?? ""}
selectedServers={form.getFieldValue("allowed_mcp_servers_and_groups")?.servers ?? []}
toolPermissions={form.getFieldValue("mcp_tool_permissions") ?? {}}
onChange={(toolPerms: Record<string, string[]>) => form.setFieldsValue({ mcp_tool_permissions: toolPerms })}
/>
</div>
)}
</Form.Item>
</div>
);
const handleAgentTypeChange = (value: string) => {
setAgentType(value);
form.resetFields();
@ -412,10 +512,16 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
<Select
mode="tags"
style={{ width: "100%" }}
placeholder="e.g. gpt-4o, claude-3-5-sonnet"
placeholder={loadingModels ? "Loading models..." : "e.g. gpt-4o, claude-3-5-sonnet"}
value={newKeyModels}
onChange={setNewKeyModels}
tokenSeparators={[","]}
loading={loadingModels}
showSearch
options={availableModels.map((m) => ({
label: getModelDisplayName(m),
value: m,
}))}
/>
</div>
</div>
@ -537,6 +643,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
{/* Step indicator */}
<Steps current={currentStep} size="small" className="mb-8">
<Step title="Configure" />
<Step title="MCP Tools" />
<Step title="Assign Key" />
<Step title="Ready" />
</Steps>
@ -544,18 +651,23 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
<Form
form={form}
layout="vertical"
initialValues={agentType === "a2a" ? getDefaultFormValues() : {}}
initialValues={
agentType === "a2a"
? { ...getDefaultFormValues(), allowed_mcp_servers_and_groups: { servers: [], accessGroups: [] }, mcp_tool_permissions: {} }
: { allowed_mcp_servers_and_groups: { servers: [], accessGroups: [] }, mcp_tool_permissions: {} }
}
className="space-y-4"
>
{currentStep === 0 && renderConfigureStep()}
{currentStep === 1 && renderAssignKeyStep()}
{currentStep === 2 && renderReadyStep()}
{currentStep === 1 && renderMCPToolsStep()}
{currentStep === 2 && renderAssignKeyStep()}
{currentStep === 3 && renderReadyStep()}
</Form>
{/* Footer navigation */}
<div className="flex items-center justify-between pt-6 border-t border-gray-100 mt-6">
<div>
{currentStep > 0 && currentStep < 2 && (
{currentStep > 0 && currentStep < 3 && (
<button
type="button"
onClick={handleBack}
@ -566,7 +678,7 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
)}
</div>
<div className="flex gap-3">
{currentStep < 2 && (
{currentStep < 3 && (
<Button variant="secondary" onClick={handleClose}>
Cancel
</Button>
@ -577,11 +689,16 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
</Button>
)}
{currentStep === 1 && (
<Button variant="primary" onClick={handleNext}>
Next →
</Button>
)}
{currentStep === 2 && (
<Button variant="primary" loading={isSubmitting} onClick={handleCreateAgent}>
{isSubmitting ? "Creating..." : "Create Agent →"}
</Button>
)}
{currentStep === 2 && (
{currentStep === 3 && (
<Button variant="primary" onClick={handleClose}>
Done
</Button>

View file

@ -0,0 +1,103 @@
import React from "react";
import { Card, Badge, Tooltip, Button } from "antd";
import { CopyOutlined, KeyOutlined, WarningOutlined, DeleteOutlined } from "@ant-design/icons";
import { Agent, AgentKeyInfo } from "./types";
interface AgentCardProps {
agent: Agent;
keyInfo?: AgentKeyInfo;
onAgentClick: (agentId: string) => void;
onDeleteClick?: (agentId: string, agentName: string) => void;
accessToken: string | null;
isAdmin: boolean;
onAgentUpdated: () => void;
}
const AgentCard: React.FC<AgentCardProps> = ({
agent,
keyInfo,
onAgentClick,
onDeleteClick,
isAdmin,
}) => {
const description =
agent.agent_card_params?.description || "No description";
const url = agent.agent_card_params?.url;
const hasKey = keyInfo?.has_key ?? false;
const statusBadge = hasKey ? (
<Badge status="success" text="Active" />
) : (
<Badge status="warning" text="Needs Setup" />
);
const copyToClipboard = (e: React.MouseEvent, text: string) => {
e.stopPropagation();
navigator.clipboard.writeText(text);
};
return (
<Card
hoverable
className="h-full flex flex-col"
styles={{
body: { flex: 1, display: "flex", flexDirection: "column" },
}}
onClick={() => onAgentClick(agent.agent_id)}
>
<div className="flex items-start justify-between gap-2 mb-2">
<div className="flex-1 min-w-0">
<div className="flex items-center gap-2 flex-wrap">
<span className="font-medium text-gray-900 truncate">
{agent.agent_name}
</span>
<Tooltip title="Copy Agent ID">
<CopyOutlined
onClick={(e) => copyToClipboard(e, agent.agent_id)}
className="cursor-pointer text-gray-400 hover:text-blue-500 text-xs shrink-0"
/>
</Tooltip>
</div>
<div className="mt-1">{statusBadge}</div>
</div>
{isAdmin && onDeleteClick && (
<Tooltip title="Delete agent">
<Button
type="text"
size="small"
danger
icon={<DeleteOutlined />}
onClick={(e) => {
e.stopPropagation();
onDeleteClick(agent.agent_id, agent.agent_name);
}}
className="shrink-0 -mr-1"
/>
</Tooltip>
)}
</div>
<p className="text-sm text-gray-600 line-clamp-2 flex-1 mb-3">
{description}
</p>
{url && (
<p className="text-xs text-gray-500 truncate mb-2" title={url}>
{url}
</p>
)}
<div className="mt-auto pt-3 border-t border-gray-100 text-xs">
{hasKey ? (
<div className="flex items-center gap-1.5 text-gray-600">
<KeyOutlined />
<span>{keyInfo?.key_alias || keyInfo?.token_prefix || "Key assigned"}</span>
</div>
) : (
<div className="flex items-center gap-1.5 text-amber-600">
<WarningOutlined />
<span>No key assigned</span>
</div>
)}
</div>
</Card>
);
};
export default AgentCard;

View file

@ -0,0 +1,63 @@
import React from "react";
import { Skeleton } from "antd";
import AgentCard from "./agent_card";
import { Agent, AgentKeyInfo } from "./types";
interface AgentCardGridProps {
agentsList: Agent[];
keyInfoMap: Record<string, AgentKeyInfo>;
isLoading: boolean;
onDeleteClick: (agentId: string, agentName: string) => void;
accessToken: string | null;
onAgentUpdated: () => void;
isAdmin: boolean;
onAgentClick: (agentId: string) => void;
}
const AgentCardGrid: React.FC<AgentCardGridProps> = ({
agentsList,
keyInfoMap,
isLoading,
onDeleteClick,
accessToken,
onAgentUpdated,
isAdmin,
onAgentClick,
}) => {
if (isLoading) {
return (
<div className="grid grid-cols-1 md:grid-cols-2 lg:grid-cols-3 gap-6">
{[1, 2, 3].map((i) => (
<Skeleton key={i} active paragraph={{ rows: 3 }} />
))}
</div>
);
}
if (!agentsList || agentsList.length === 0) {
return (
<div className="rounded-lg border border-gray-200 bg-gray-50/50 py-12 text-center">
<p className="text-gray-500">No agents found. Create one to get started.</p>
</div>
);
}
return (
<div className="grid grid-cols-1 md:grid-cols-2 lg:grid-cols-3 gap-6">
{agentsList.map((agent) => (
<AgentCard
key={agent.agent_id}
agent={agent}
keyInfo={keyInfoMap[agent.agent_id]}
onAgentClick={onAgentClick}
onDeleteClick={isAdmin ? onDeleteClick : undefined}
accessToken={accessToken}
isAdmin={isAdmin}
onAgentUpdated={onAgentUpdated}
/>
))}
</div>
);
};
export default AgentCardGrid;

View file

@ -205,6 +205,44 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({
<Descriptions.Item label="Updated At">{formatDate(agent.updated_at)}</Descriptions.Item>
</Descriptions>
{agent.object_permission &&
(agent.object_permission.mcp_servers?.length ||
agent.object_permission.mcp_access_groups?.length ||
(agent.object_permission.mcp_tool_permissions &&
Object.keys(agent.object_permission.mcp_tool_permissions).length > 0)) && (
<div style={{ marginTop: 24 }}>
<Title>MCP Tool Permissions</Title>
<Descriptions bordered column={1} style={{ marginTop: 16 }}>
{agent.object_permission.mcp_servers && agent.object_permission.mcp_servers.length > 0 && (
<Descriptions.Item label="MCP Servers">
{agent.object_permission.mcp_servers.join(", ")}
</Descriptions.Item>
)}
{agent.object_permission.mcp_access_groups &&
agent.object_permission.mcp_access_groups.length > 0 && (
<Descriptions.Item label="MCP Access Groups">
{agent.object_permission.mcp_access_groups.join(", ")}
</Descriptions.Item>
)}
{agent.object_permission.mcp_tool_permissions &&
Object.keys(agent.object_permission.mcp_tool_permissions).length > 0 && (
<Descriptions.Item label="Tool permissions per server">
<div className="space-y-1">
{Object.entries(agent.object_permission.mcp_tool_permissions).map(
([serverId, tools]) => (
<div key={serverId}>
<span className="font-medium">{serverId}:</span>{" "}
{Array.isArray(tools) ? tools.join(", ") : String(tools)}
</div>
)
)}
</div>
</Descriptions.Item>
)}
</Descriptions>
</div>
)}
<AgentCostView agent={agent} />
{agent.agent_card_params?.skills && agent.agent_card_params.skills.length > 0 && (

View file

@ -1,3 +1,15 @@
export interface AgentKeyInfo {
key_alias?: string;
token_prefix?: string;
has_key: boolean;
}
export interface AgentObjectPermission {
mcp_servers?: string[];
mcp_access_groups?: string[];
mcp_tool_permissions?: Record<string, string[]>;
}
export interface Agent {
agent_id: string;
agent_name: string;
@ -7,8 +19,10 @@ export interface Agent {
};
agent_card_params?: {
description?: string;
url?: string;
[key: string]: any;
};
object_permission?: AgentObjectPermission;
created_at?: string;
updated_at?: string;
created_by?: string;

View file

@ -1,7 +1,7 @@
import React, { useState } from "react";
import { Input } from "antd";
import { SearchOutlined, ArrowRightOutlined } from "@ant-design/icons";
import { GuardrailCardInfo, LITELLM_CONTENT_FILTER_CARDS, PARTNER_GUARDRAIL_CARDS, ALL_CARDS } from "./guardrail_garden_data";
import { GuardrailCardInfo, ALL_CARDS } from "./guardrail_garden_data";
import GuardrailCard from "./guardrail_garden_card";
import GuardrailDetailView from "./guardrail_garden_detail";

View file

@ -194,7 +194,7 @@ const MCPToolConfiguration: React.FC<MCPToolConfigurationProps> = ({
{filteredTools.length === 0 ? (
<div className="text-center py-6 text-gray-400 border rounded-lg border-dashed">
<SearchOutlined className="text-2xl mb-2" />
<Text>No tools found matching "{toolSearchTerm}"</Text>
<Text>No tools found matching &quot;{toolSearchTerm}&quot;</Text>
</div>
) : (
filteredTools.map((tool, index) => (

View file

@ -1,12 +1,12 @@
import React, { useState } from "react";
import { useQuery, useMutation } from "@tanstack/react-query";
import { ToolTestPanel } from "./ToolTestPanel";
import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse, AUTH_TYPE } from "./types";
import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse } from "./types";
import { listMCPTools, callMCPTool } from "../networking";
import { Card, Title, Text } from "@tremor/react";
import { RobotOutlined, ToolOutlined, SearchOutlined, LockOutlined, KeyOutlined } from "@ant-design/icons";
import { Input, Alert, Button as AntdButton } from "antd";
import { RobotOutlined, ToolOutlined, SearchOutlined, KeyOutlined } from "@ant-design/icons";
import { Input, Button as AntdButton } from "antd";
const MCPToolsViewer = ({
serverId,
@ -134,7 +134,7 @@ const MCPToolsViewer = ({
{!showHeaderInput && Object.keys(passthroughHeaders).length === 0 && (
<Text className="text-xs text-blue-700">
This server requires additional headers. Click "Configure" to provide values.
This server requires additional headers. Click &quot;Configure&quot; to provide values.
</Text>
)}
@ -255,7 +255,7 @@ const MCPToolsViewer = ({
<div className="p-4 text-center bg-white border border-gray-200 rounded-lg">
<SearchOutlined className="text-2xl text-gray-400 mb-2" />
<p className="text-xs font-medium text-gray-700 mb-1">No tools found</p>
<p className="text-xs text-gray-500">No tools match "{toolSearchTerm}"</p>
<p className="text-xs text-gray-500">No tools match &quot;{toolSearchTerm}&quot;</p>
</div>
) : (
<div

View file

@ -5908,6 +5908,7 @@ export const enrichPolicyTemplateStream = async (
const decoder = new TextDecoder();
let buffer = "";
// eslint-disable-next-line no-constant-condition -- stream read loop
while (true) {
const { done, value } = await reader.read();
if (done) break;
@ -5982,6 +5983,7 @@ export const usageAiChatStream = async (
const decoder = new TextDecoder();
let buffer = "";
// eslint-disable-next-line no-constant-condition -- stream read loop
while (true) {
const { done, value } = await reader.read();
if (done) break;

View file

@ -27,6 +27,7 @@ vi.mock("../networking", () => ({
soft_budget: null,
}),
fetchMCPAccessGroups: vi.fn().mockResolvedValue([]),
getAgentsList: vi.fn().mockResolvedValue([]),
}));
vi.mock("../molecules/notifications_manager", () => ({

View file

@ -5,13 +5,13 @@ import { formatNumberWithCommas } from "@/utils/dataUtils";
import { InfoCircleOutlined } from "@ant-design/icons";
import { useQueryClient } from "@tanstack/react-query";
import { Accordion, AccordionBody, AccordionHeader, Button, Col, Grid, Text, TextInput, Title } from "@tremor/react";
import { Button as Button2, Form, Input, Modal, Radio, Select, Switch, Tag, Tooltip } from "antd";
import { Button as Button2, Form, Input, message, Modal, Radio, Select, Switch, Tag, Tooltip } from "antd";
import debounce from "lodash/debounce";
import React, { useCallback, useEffect, useState } from "react";
import { CopyToClipboard } from "react-copy-to-clipboard";
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";
@ -20,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";
@ -40,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;
@ -299,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;
@ -559,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>
@ -1005,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 }))}
/>
@ -1059,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={
@ -1297,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>

View file

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

View file

@ -998,7 +998,7 @@ const PipelineTestPanel: React.FC<PipelineTestPanelProps> = ({
{!result && !error && complianceResults.length === 0 && (
<div style={{ textAlign: "center", color: "#9ca3af", fontSize: 13, marginTop: 24 }}>
Choose a test source above (quick chat or a compliance dataset) and click "Run Test"
Choose a test source above (quick chat or a compliance dataset) and click &quot;Run Test&quot;
</div>
)}
</div>

View file

@ -1,5 +1,5 @@
import React, { useState, useEffect, useMemo } from "react";
import { Card, Button, Spin, message, Checkbox, Badge } from "antd";
import { Card, Button, Spin, message, Checkbox } from "antd";
import {
ShieldCheckIcon,
ShieldExclamationIcon,

View file

@ -1,4 +1,4 @@
import { render, screen, waitFor } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import { KeyInfoHeader, KeyInfoData } from "./KeyInfoHeader";

View file

@ -4,7 +4,6 @@ import moment from "moment";
import { LogEntry } from "../columns";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import GuardrailViewer from "../GuardrailViewer/GuardrailViewer";
import CompliancePanel from "../GuardrailViewer/CompliancePanel";
import { CostBreakdownViewer } from "../CostBreakdownViewer";
import { ConfigInfoMessage } from "../ConfigInfoMessage";
import { VectorStoreViewer } from "../VectorStoreViewer";

View file

@ -15,7 +15,7 @@ import { fetchAllKeyAliases } from "../key_team_helpers/filter_helpers";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import { PaginatedModelSelect } from "../ModelSelect/PaginatedModelSelect/PaginatedModelSelect";
import FilterComponent, { FilterOption } from "../molecules/filter";
import { allEndUsersCall, keyInfoV1Call, keyListCall, uiSpendLogsCall } from "../networking";
import { allEndUsersCall, keyInfoV1Call, uiSpendLogsCall } from "../networking";
import KeyInfoView from "../templates/key_info_view";
import AuditLogs from "./audit_logs";
import { createColumns, LogEntry, type LogsSortField } from "./columns";

View file

@ -1,7 +1,7 @@
// Auto-generated from block_investment.csv — do not edit manually.
// Regenerate: python scripts/generate_compliance_prompts.py --csv ... --output ...
import type { CompliancePrompt, ComplianceFramework } from "./compliancePrompts";
import type { CompliancePrompt } from "./compliancePrompts";
export const financialCompliancePrompts: CompliancePrompt[] = [
{

View file

@ -1,7 +1,7 @@
// Auto-generated from block_insults.csv — do not edit manually.
// Regenerate: python scripts/generate_compliance_prompts.py --csv ... --output ...
import type { CompliancePrompt, ComplianceFramework } from "./compliancePrompts";
import type { CompliancePrompt } from "./compliancePrompts";
export const insultsCompliancePrompts: CompliancePrompt[] = [
{

View file

@ -14,7 +14,7 @@
"moduleResolution": "bundler",
"resolveJsonModule": true,
"isolatedModules": true,
"jsx": "preserve",
"jsx": "react-jsx",
"incremental": true,
"plugins": [
{