Litellm embeddings calltype fix for guardrail precallhook (#18740)

* adding signoz integration to observability docs

* Fixing build

* Adding timeout for flaky test

* Fixing e2e

* add team member budget duration in team/update

* Reusable Duration Select and update team member budget UI

* feat: allow configuring project name for OpenTelemetry service name

* docs: sets ARIZE_PROJECT_NAME

* added valid callType for bedrock guardrail pre hook

This is to resolve the error when bedrock guardrails are enabled and invoke the embedding models.   {"error":{"message":"'embeddings' is not a valid CallTypes","type":"None","param":"None","code":"500"}}*

* updated the test case to reflect valid callType

---------

Co-authored-by: Goutham Karthi <goutham@signoz.io>
Co-authored-by: yuneng-jiang <yuneng.jiang@gmail.com>
Co-authored-by: YutaSaito <36355491+uc4w6c@users.noreply.github.com>
Co-authored-by: Yuta Saito <uc4w6c@bma.biglobe.ne.jp>
This commit is contained in:
kothamah 2026-01-07 11:10:36 -05:00 • committed by GitHub
parent 92f7789f10
commit 1b8708fccc
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
16 changed files with 552 additions and 72 deletions

View file

@ -68,6 +68,7 @@ environment_variables:
ARIZE_API_KEY: "141a****"
ARIZE_ENDPOINT: "https://otlp.arize.com/v1" # OPTIONAL - your custom arize GRPC api endpoint
ARIZE_HTTP_ENDPOINT: "https://otlp.arize.com/v1" # OPTIONAL - your custom arize HTTP api endpoint. Set either this or ARIZE_ENDPOINT or Neither (defaults to https://otlp.arize.com/v1 on grpc)
ARIZE_PROJECT_NAME: "my-litellm-project" # OPTIONAL - sets the arize project name
```
2. Start the proxy

View file

@ -51,6 +51,7 @@ class ArizeLogger(OpenTelemetry):
space_id = os.environ.get("ARIZE_SPACE_ID")
space_key = os.environ.get("ARIZE_SPACE_KEY")
api_key = os.environ.get("ARIZE_API_KEY")
project_name = os.environ.get("ARIZE_PROJECT_NAME")
grpc_endpoint = os.environ.get("ARIZE_ENDPOINT")
http_endpoint = os.environ.get("ARIZE_HTTP_ENDPOINT")
@ -74,6 +75,7 @@ class ArizeLogger(OpenTelemetry):
api_key=api_key,
protocol=protocol,
endpoint=endpoint,
project_name=project_name,
)
async def async_service_success_hook(

View file

@ -54,38 +54,6 @@ RAW_REQUEST_SPAN_NAME = "raw_gen_ai_request"
LITELLM_REQUEST_SPAN_NAME = "litellm_request"
def _get_litellm_resource():
"""
Create a proper OpenTelemetry Resource that respects OTEL_RESOURCE_ATTRIBUTES
while maintaining backward compatibility with LiteLLM-specific environment variables.
"""
from opentelemetry.sdk.resources import OTELResourceDetector, Resource
# Create base resource attributes with LiteLLM-specific defaults
# These will be overridden by OTEL_RESOURCE_ATTRIBUTES if present
base_attributes: Dict[str, Optional[str]] = {
"service.name": os.getenv("OTEL_SERVICE_NAME", "litellm"),
"deployment.environment": os.getenv("OTEL_ENVIRONMENT_NAME", "production"),
# Fix the model_id to use proper environment variable or default to service name
"model_id": os.getenv(
"OTEL_MODEL_ID", os.getenv("OTEL_SERVICE_NAME", "litellm")
),
}
# Create base resource with LiteLLM-specific defaults
base_resource = Resource.create(base_attributes) # type: ignore
# Create resource from OTEL_RESOURCE_ATTRIBUTES using the detector
otel_resource_detector = OTELResourceDetector()
env_resource = otel_resource_detector.detect()
# Merge the resources: env_resource takes precedence over base_resource
# This ensures OTEL_RESOURCE_ATTRIBUTES overrides LiteLLM defaults
merged_resource = base_resource.merge(env_resource)
return merged_resource
@dataclass
class OpenTelemetryConfig:
exporter: Union[str, SpanExporter] = "console"
@ -93,6 +61,19 @@ class OpenTelemetryConfig:
headers: Optional[str] = None
enable_metrics: bool = False
enable_events: bool = False
service_name: Optional[str] = None
deployment_environment: Optional[str] = None
model_id: Optional[str] = None
def __post_init__(self) -> None:
if not self.service_name:
self.service_name = os.getenv("OTEL_SERVICE_NAME", "litellm")
if not self.deployment_environment:
self.deployment_environment = os.getenv(
"OTEL_ENVIRONMENT_NAME", "production"
)
if not self.model_id:
self.model_id = os.getenv("OTEL_MODEL_ID", self.service_name)
@classmethod
def from_env(cls):
@ -122,6 +103,9 @@ class OpenTelemetryConfig:
os.getenv("LITELLM_OTEL_INTEGRATION_ENABLE_EVENTS", "false").lower()
== "true"
)
service_name = os.getenv("OTEL_SERVICE_NAME", "litellm")
deployment_environment = os.getenv("OTEL_ENVIRONMENT_NAME", "production")
model_id = os.getenv("OTEL_MODEL_ID", service_name)
if exporter == "in_memory":
return cls(exporter=InMemorySpanExporter())
@ -131,6 +115,9 @@ class OpenTelemetryConfig:
headers=headers, # example: OTEL_HEADERS=x-honeycomb-team=B85YgLm96***"
enable_metrics=enable_metrics,
enable_events=enable_events,
service_name=service_name,
deployment_environment=deployment_environment,
model_id=model_id,
)
@ -174,6 +161,22 @@ class OpenTelemetry(CustomLogger):
self._init_logs(logger_provider)
self._init_otel_logger_on_litellm_proxy()
@staticmethod
def _get_litellm_resource(config: OpenTelemetryConfig):
"""Create an OpenTelemetry Resource using config-driven defaults."""
from opentelemetry.sdk.resources import OTELResourceDetector, Resource
base_attributes: Dict[str, Optional[str]] = {
"service.name": config.service_name,
"deployment.environment": config.deployment_environment,
"model_id": config.model_id or config.service_name,
}
base_resource = Resource.create(base_attributes) # type: ignore[arg-type]
otel_resource_detector = OTELResourceDetector()
env_resource = otel_resource_detector.detect()
return base_resource.merge(env_resource)
def _init_otel_logger_on_litellm_proxy(self):
"""
Initializes OpenTelemetry for litellm proxy server
@ -266,7 +269,7 @@ class OpenTelemetry(CustomLogger):
from opentelemetry.trace import SpanKind
def create_tracer_provider():
provider = TracerProvider(resource=_get_litellm_resource())
provider = TracerProvider(resource=self._get_litellm_resource(self.config))
provider.add_span_processor(self._get_span_processor())
return provider
@ -300,7 +303,8 @@ class OpenTelemetry(CustomLogger):
def create_meter_provider():
metric_reader = self._get_metric_reader()
return MeterProvider(
metric_readers=[metric_reader], resource=_get_litellm_resource()
metric_readers=[metric_reader],
resource=self._get_litellm_resource(self.config),
)
meter_provider = self._get_or_create_provider(
@ -355,7 +359,9 @@ class OpenTelemetry(CustomLogger):
from opentelemetry.sdk._logs.export import BatchLogRecordProcessor
def create_logger_provider():
provider = OTLoggerProvider(resource=_get_litellm_resource())
provider = OTLoggerProvider(
resource=self._get_litellm_resource(self.config)
)
log_exporter = self._get_log_exporter()
provider.add_log_record_processor(
BatchLogRecordProcessor(log_exporter) # type: ignore[arg-type]
@ -606,7 +612,7 @@ class OpenTelemetry(CustomLogger):
from opentelemetry.sdk.trace import TracerProvider
# Create a temporary tracer provider with dynamic headers
temp_provider = TracerProvider(resource=_get_litellm_resource())
temp_provider = TracerProvider(resource=self._get_litellm_resource(self.config))
temp_provider.add_span_processor(
self._get_span_processor(dynamic_headers=dynamic_headers)
)
@ -987,9 +993,9 @@ class OpenTelemetry(CustomLogger):
# Get the resource from the logger provider
logger_provider = get_logger_provider()
resource = (
getattr(logger_provider, "_resource", None) or _get_litellm_resource()
)
resource = getattr(
logger_provider, "_resource", None
) or self._get_litellm_resource(self.config)
parent_ctx = span.get_span_context()
provider = (kwargs.get("litellm_params") or {}).get(
@ -1910,7 +1916,9 @@ class OpenTelemetry(CustomLogger):
)
_split_otel_headers = OpenTelemetry._get_headers_dictionary(self.OTEL_HEADERS)
normalized_endpoint = self._normalize_otel_endpoint(self.OTEL_ENDPOINT, "metrics")
normalized_endpoint = self._normalize_otel_endpoint(
self.OTEL_ENDPOINT, "metrics"
)
if self.OTEL_EXPORTER == "console":
exporter = ConsoleMetricExporter()

View file

@ -3630,6 +3630,7 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
otel_config = OpenTelemetryConfig(
exporter=arize_config.protocol,
endpoint=arize_config.endpoint,
service_name=arize_config.project_name,
)
os.environ[

View file

@ -1532,6 +1532,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
guardrails: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
team_member_budget: Optional[float] = None
team_member_budget_duration: Optional[str] = None
team_member_rpm_limit: Optional[int] = None
team_member_tpm_limit: Optional[int] = None
team_member_key_duration: Optional[str] = None

View file

@ -957,11 +957,11 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
def _get_pre_call_type(
route_type: Literal["acompletion", "aembedding", "aresponses", "allm_passthrough_route"],
) -> Literal["completion", "embeddings", "responses", "allm_passthrough_route"]:
) -> Literal["completion", "embedding", "responses", "allm_passthrough_route"]:
if route_type == "acompletion":
return "completion"
elif route_type == "aembedding":
return "embeddings"
return "embedding"
elif route_type == "aresponses":
return "responses"
elif route_type == "allm_passthrough_route":

View file

@ -112,6 +112,7 @@ class TeamMemberBudgetHandler:
team_member_budget: Optional[float] = None,
team_member_rpm_limit: Optional[int] = None,
team_member_tpm_limit: Optional[int] = None,
team_member_budget_duration: Optional[str] = None,
) -> bool:
"""Check if any team member limits are provided"""
return any(
@ -119,6 +120,7 @@ class TeamMemberBudgetHandler:
team_member_budget is not None,
team_member_rpm_limit is not None,
team_member_tpm_limit is not None,
team_member_budget_duration is not None,
]
)
@ -130,6 +132,7 @@ class TeamMemberBudgetHandler:
team_member_budget: Optional[float] = None,
team_member_rpm_limit: Optional[int] = None,
team_member_tpm_limit: Optional[int] = None,
team_member_budget_duration: Optional[str] = None,
) -> dict:
"""Create team member budget table with provided limits"""
from litellm.proxy._types import BudgetNewRequest
@ -147,7 +150,7 @@ class TeamMemberBudgetHandler:
# Create budget request with all provided limits
budget_request = BudgetNewRequest(
budget_id=budget_id,
budget_duration=data.budget_duration,
budget_duration=data.budget_duration or team_member_budget_duration,
)
if team_member_budget is not None:
@ -156,6 +159,8 @@ class TeamMemberBudgetHandler:
budget_request.rpm_limit = team_member_rpm_limit
if team_member_tpm_limit is not None:
budget_request.tpm_limit = team_member_tpm_limit
if team_member_budget_duration is not None:
budget_request.budget_duration = team_member_budget_duration
team_member_budget_table = await new_budget(
budget_obj=budget_request,
@ -182,6 +187,7 @@ class TeamMemberBudgetHandler:
team_member_budget: Optional[float] = None,
team_member_rpm_limit: Optional[int] = None,
team_member_tpm_limit: Optional[int] = None,
team_member_budget_duration: Optional[str] = None,
) -> dict:
"""Upsert team member budget table with provided limits"""
from litellm.proxy._types import BudgetNewRequest
@ -203,6 +209,8 @@ class TeamMemberBudgetHandler:
budget_request.rpm_limit = team_member_rpm_limit
if team_member_tpm_limit is not None:
budget_request.tpm_limit = team_member_tpm_limit
if team_member_budget_duration is not None:
budget_request.budget_duration = team_member_budget_duration
budget_row = await update_budget(
budget_obj=budget_request,
@ -223,6 +231,7 @@ class TeamMemberBudgetHandler:
team_member_budget=team_member_budget,
team_member_rpm_limit=team_member_rpm_limit,
team_member_tpm_limit=team_member_tpm_limit,
team_member_budget_duration=team_member_budget_duration,
)
# Remove team member fields from updated_kv
@ -233,6 +242,7 @@ class TeamMemberBudgetHandler:
def _clean_team_member_fields(data_dict: dict) -> None:
"""Remove team member fields from data dictionary"""
data_dict.pop("team_member_budget", None)
data_dict.pop("team_member_budget_duration", None)
data_dict.pop("team_member_rpm_limit", None)
data_dict.pop("team_member_tpm_limit", None)
@ -1214,6 +1224,7 @@ async def update_team( # noqa: PLR0915
- disable_global_guardrails: Optional[bool] - Whether to disable global guardrails for the key.
- object_permission: Optional[LiteLLM_ObjectPermissionBase] - team-specific object permission. Example - {"vector_stores": ["vector_store_1", "vector_store_2"], "agents": ["agent_1", "agent_2"], "agent_access_groups": ["dev_group"]}. IF null or {} then no object permission.
- team_member_budget: Optional[float] - The maximum budget allocated to an individual team member.
- team_member_budget_duration: Optional[str] - The duration of the budget for the team member. Doc [here](https://docs.litellm.ai/docs/proxy/team_budgets)
- team_member_rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for individual team members.
- team_member_tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for individual team members.
- team_member_key_duration: Optional[str] - The duration for a team member's key. e.g. "1d", "1w", "1mo"
@ -1349,6 +1360,7 @@ async def update_team( # noqa: PLR0915
team_member_budget=data.team_member_budget,
team_member_rpm_limit=data.team_member_rpm_limit,
team_member_tpm_limit=data.team_member_tpm_limit,
team_member_budget_duration=data.team_member_budget_duration,
):
updated_kv = await TeamMemberBudgetHandler.upsert_team_member_budget_table(
team_table=existing_team_row,
@ -1357,6 +1369,7 @@ async def update_team( # noqa: PLR0915
team_member_budget=data.team_member_budget,
team_member_rpm_limit=data.team_member_rpm_limit,
team_member_tpm_limit=data.team_member_tpm_limit,
team_member_budget_duration=data.team_member_budget_duration,
)
else:
TeamMemberBudgetHandler._clean_team_member_fields(updated_kv)

View file

@ -14,3 +14,4 @@ class ArizeConfig(BaseModel):
api_key: Optional[str] = None
protocol: Protocol
endpoint: str
project_name: Optional[str] = None

View file

@ -71,6 +71,7 @@ def test_get_arize_config(mock_env_vars):
assert config.api_key == "test_api_key"
assert config.endpoint == "https://otlp.arize.com/v1"
assert config.protocol == "otlp_grpc"
assert config.project_name is None
def test_get_arize_config_with_endpoints(mock_env_vars, monkeypatch):
@ -79,10 +80,12 @@ def test_get_arize_config_with_endpoints(mock_env_vars, monkeypatch):
"""
monkeypatch.setenv("ARIZE_ENDPOINT", "grpc://test.endpoint")
monkeypatch.setenv("ARIZE_HTTP_ENDPOINT", "http://test.endpoint")
monkeypatch.setenv("ARIZE_PROJECT_NAME", "custom-project")
config = ArizeLogger.get_arize_config()
assert config.endpoint == "grpc://test.endpoint"
assert config.protocol == "otlp_grpc"
assert config.project_name == "custom-project"
@pytest.mark.skip(

View file

@ -123,7 +123,8 @@ class TestArizeIntegrationWithProxy:
with patch.dict(os.environ, {
"ARIZE_SPACE_KEY": "test-space-123",
"ARIZE_API_KEY": "test-api-456",
"ARIZE_ENDPOINT": "https://custom.arize.com/v1"
"ARIZE_ENDPOINT": "https://custom.arize.com/v1",
"ARIZE_PROJECT_NAME": "custom-project",
}):
config = ArizeLogger.get_arize_config()
@ -131,13 +132,15 @@ class TestArizeIntegrationWithProxy:
assert config.api_key == "test-api-456"
assert config.endpoint == "https://custom.arize.com/v1"
assert config.protocol == "otlp_grpc"
assert config.project_name == "custom-project"
def test_arize_get_config_defaults(self):
"""Test ArizeLogger.get_arize_config() with default endpoint."""
with patch.dict(os.environ, {
"ARIZE_SPACE_KEY": "test-space-default",
"ARIZE_API_KEY": "test-api-default"
"ARIZE_API_KEY": "test-api-default",
"ARIZE_PROJECT_NAME": "default-project",
}, clear=True):
config = ArizeLogger.get_arize_config()
@ -145,6 +148,7 @@ class TestArizeIntegrationWithProxy:
assert config.api_key == "test-api-default"
assert config.endpoint == "https://otlp.arize.com/v1" # Default endpoint
assert config.protocol == "otlp_grpc" # Default protocol
assert config.project_name == "default-project"
def test_arize_construct_dynamic_headers(self):
"""Test dynamic OTEL headers construction for team/key logging."""
@ -180,4 +184,4 @@ class TestArizeIntegrationWithProxy:
if __name__ == "__main__":
pytest.main([__file__, "-v"])
pytest.main([__file__, "-v"])

View file

@ -530,7 +530,7 @@ class TestPassthroughCallTypeHandling:
)
assert (
ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="aembedding")
== "embeddings"
== "embedding"
)
assert (
ProxyBaseLLMRequestProcessing._get_pre_call_type(route_type="aresponses")

View file

@ -258,6 +258,22 @@ class TestOpenTelemetry(unittest.TestCase):
MODEL = "arn:aws:bedrock:us-west-2:1234567890123:inference-profile/us.anthropic.claude-3-7-sonnet-20250219-v1:0"
HERE = os.path.dirname(__file__)
@patch.dict(os.environ, {}, clear=True)
def test_open_telemetry_config_manual_defaults(self):
"""Manual OpenTelemetryConfig creation should populate default identifiers."""
config = OpenTelemetryConfig(exporter="console", endpoint="http://collector")
self.assertEqual(config.service_name, "litellm")
self.assertEqual(config.deployment_environment, "production")
self.assertEqual(config.model_id, "litellm")
@patch.dict(os.environ, {}, clear=True)
def test_open_telemetry_config_custom_service_name(self):
"""Model ID should inherit provided service name when not explicitly set."""
config = OpenTelemetryConfig(service_name="custom-service", exporter="console")
self.assertEqual(config.service_name, "custom-service")
self.assertEqual(config.deployment_environment, "production")
self.assertEqual(config.model_id, "custom-service")
def wait_for_spans(self, exporter: InMemorySpanExporter, prefix: str):
"""Poll until we see at least one span with an attribute key starting with `prefix`."""
deadline = time.time() + self.POLL_TIMEOUT
@ -504,8 +520,6 @@ class TestOpenTelemetry(unittest.TestCase):
self, mock_detector_cls, mock_resource_create
):
"""Test _get_litellm_resource with default values when no environment variables are set."""
from litellm.integrations.opentelemetry import _get_litellm_resource
# Mock the Resource.create method
mock_base_resource = MagicMock()
mock_resource_create.return_value = mock_base_resource
@ -520,8 +534,8 @@ class TestOpenTelemetry(unittest.TestCase):
mock_merged_resource = MagicMock()
mock_base_resource.merge.return_value = mock_merged_resource
# Call the function
result = _get_litellm_resource()
config = OpenTelemetryConfig()
result = OpenTelemetry._get_litellm_resource(config)
# Verify Resource.create was called with correct default attributes
expected_attributes = {
@ -549,8 +563,6 @@ class TestOpenTelemetry(unittest.TestCase):
self, mock_detector_cls, mock_resource_create
):
"""Test _get_litellm_resource with LiteLLM-specific environment variables."""
from litellm.integrations.opentelemetry import _get_litellm_resource
# Mock the Resource.create method
mock_base_resource = MagicMock()
mock_resource_create.return_value = mock_base_resource
@ -565,8 +577,8 @@ class TestOpenTelemetry(unittest.TestCase):
mock_merged_resource = MagicMock()
mock_base_resource.merge.return_value = mock_merged_resource
# Call the function
result = _get_litellm_resource()
config = OpenTelemetryConfig.from_env()
result = OpenTelemetry._get_litellm_resource(config)
# Verify Resource.create was called with environment variable values
expected_attributes = {
@ -593,8 +605,6 @@ class TestOpenTelemetry(unittest.TestCase):
self, mock_detector_cls, mock_resource_create
):
"""Test _get_litellm_resource with OTEL_RESOURCE_ATTRIBUTES environment variable."""
from litellm.integrations.opentelemetry import _get_litellm_resource
# Mock the Resource.create method to simulate the actual behavior
# In reality, Resource.create() would parse OTEL_RESOURCE_ATTRIBUTES and merge it
mock_base_resource = MagicMock()
@ -610,8 +620,8 @@ class TestOpenTelemetry(unittest.TestCase):
mock_merged_resource = MagicMock()
mock_base_resource.merge.return_value = mock_merged_resource
# Call the function
result = _get_litellm_resource()
config = OpenTelemetryConfig.from_env()
result = OpenTelemetry._get_litellm_resource(config)
# Verify Resource.create was called with the base attributes
# The actual OTEL_RESOURCE_ATTRIBUTES parsing is handled by OpenTelemetry SDK
@ -628,10 +638,8 @@ class TestOpenTelemetry(unittest.TestCase):
@patch.dict(os.environ, {}, clear=True)
def test_get_litellm_resource_integration_with_real_resource(self):
"""Integration test to verify _get_litellm_resource works with actual OpenTelemetry Resource."""
from litellm.integrations.opentelemetry import _get_litellm_resource
# This test uses the real OpenTelemetry Resource.create() method
result = _get_litellm_resource()
config = OpenTelemetryConfig()
result = OpenTelemetry._get_litellm_resource(config)
# Verify the result is a Resource instance
from opentelemetry.sdk.resources import Resource
@ -653,10 +661,8 @@ class TestOpenTelemetry(unittest.TestCase):
)
def test_get_litellm_resource_real_otel_resource_attributes(self):
"""Integration test to verify OTEL_RESOURCE_ATTRIBUTES is properly handled."""
from litellm.integrations.opentelemetry import _get_litellm_resource
# This test uses the real OpenTelemetry Resource.create() method
result = _get_litellm_resource()
config = OpenTelemetryConfig.from_env()
result = OpenTelemetry._get_litellm_resource(config)
print("RESULT", result)
@ -683,10 +689,8 @@ class TestOpenTelemetry(unittest.TestCase):
)
def test_get_litellm_resource_precedence(self):
"""Test that OTEL_SERVICE_NAME takes precedence over OTEL_RESOURCE_ATTRIBUTES according to OpenTelemetry spec."""
from litellm.integrations.opentelemetry import _get_litellm_resource
# This test verifies the OpenTelemetry standard behavior
result = _get_litellm_resource()
config = OpenTelemetryConfig.from_env()
result = OpenTelemetry._get_litellm_resource(config)
# Verify the result is a Resource instance
from opentelemetry.sdk.resources import Resource

View file

@ -1279,7 +1279,7 @@ async def test_update_team_team_member_budget_not_passed_to_db():
# Mock budget upsert to return updated_kv without team_member_budget
def mock_upsert_side_effect(
team_table, user_api_key_dict, updated_kv, team_member_budget=None, team_member_rpm_limit=None, team_member_tpm_limit=None
team_table, user_api_key_dict, updated_kv, team_member_budget=None, team_member_rpm_limit=None, team_member_tpm_limit=None, team_member_budget_duration=None
):
# Remove team_member_budget from updated_kv as the real function does
result_kv = updated_kv.copy()
@ -1376,6 +1376,370 @@ async def test_update_team_team_member_budget_not_passed_to_db():
)
def test_clean_team_member_fields():
"""
Test that _clean_team_member_fields removes all team member fields from a dictionary.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
TeamMemberBudgetHandler,
)
data_dict = {
"team_id": "test_team",
"team_alias": "Test Team",
"team_member_budget": 100.0,
"team_member_budget_duration": "30d",
"team_member_rpm_limit": 50,
"team_member_tpm_limit": 1000,
"other_field": "should_remain",
}
TeamMemberBudgetHandler._clean_team_member_fields(data_dict)
assert "team_member_budget" not in data_dict
assert "team_member_budget_duration" not in data_dict
assert "team_member_rpm_limit" not in data_dict
assert "team_member_tpm_limit" not in data_dict
assert data_dict["team_id"] == "test_team"
assert data_dict["team_alias"] == "Test Team"
assert data_dict["other_field"] == "should_remain"
def test_clean_team_member_fields_with_missing_fields():
"""
Test that _clean_team_member_fields handles dictionaries without team member fields gracefully.
"""
from litellm.proxy.management_endpoints.team_endpoints import (
TeamMemberBudgetHandler,
)
data_dict = {
"team_id": "test_team",
"team_alias": "Test Team",
}
TeamMemberBudgetHandler._clean_team_member_fields(data_dict)
assert data_dict["team_id"] == "test_team"
assert data_dict["team_alias"] == "Test Team"
@pytest.mark.asyncio
async def test_create_team_member_budget_table():
"""
Test that create_team_member_budget_table creates a budget and adds it to metadata.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LitellmUserRoles, NewTeamRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import (
TeamMemberBudgetHandler,
)
mock_user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id"
)
data = NewTeamRequest(
team_id="test_team_id",
team_alias="Test Team",
budget_duration="1mo",
)
new_team_data_json = {
"team_id": "test_team_id",
"team_alias": "Test Team",
"team_member_budget": 100.0,
"team_member_budget_duration": "30d",
"team_member_rpm_limit": 50,
"team_member_tpm_limit": 1000,
}
mock_budget_response = MagicMock()
mock_budget_response.budget_id = "budget_123"
with patch(
"litellm.proxy.management_endpoints.budget_management_endpoints.new_budget",
new_callable=AsyncMock
) as mock_new_budget:
mock_new_budget.return_value = mock_budget_response
result = await TeamMemberBudgetHandler.create_team_member_budget_table(
data=data,
new_team_data_json=new_team_data_json,
user_api_key_dict=mock_user_api_key_dict,
team_member_budget=100.0,
team_member_rpm_limit=50,
team_member_tpm_limit=1000,
team_member_budget_duration="30d",
)
assert mock_new_budget.called
call_args = mock_new_budget.call_args
budget_request = call_args[1]["budget_obj"]
assert budget_request.max_budget == 100.0
assert budget_request.rpm_limit == 50
assert budget_request.tpm_limit == 1000
assert budget_request.budget_duration == "30d"
assert budget_request.budget_id is not None
assert "team-" in budget_request.budget_id
assert "team_member_budget_id" in result["metadata"]
assert result["metadata"]["team_member_budget_id"] == "budget_123"
assert "team_member_budget" not in result
assert "team_member_budget_duration" not in result
assert "team_member_rpm_limit" not in result
assert "team_member_tpm_limit" not in result
@pytest.mark.asyncio
async def test_create_team_member_budget_table_without_team_alias():
"""
Test that create_team_member_budget_table generates budget_id correctly when team_alias is None.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LitellmUserRoles, NewTeamRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import (
TeamMemberBudgetHandler,
)
mock_user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id"
)
data = NewTeamRequest(team_id="test_team_id")
new_team_data_json = {
"team_id": "test_team_id",
"team_member_budget": 100.0,
}
mock_budget_response = MagicMock()
mock_budget_response.budget_id = "budget_123"
with patch(
"litellm.proxy.management_endpoints.budget_management_endpoints.new_budget",
new_callable=AsyncMock
) as mock_new_budget:
mock_new_budget.return_value = mock_budget_response
result = await TeamMemberBudgetHandler.create_team_member_budget_table(
data=data,
new_team_data_json=new_team_data_json,
user_api_key_dict=mock_user_api_key_dict,
team_member_budget=100.0,
)
assert mock_new_budget.called
call_args = mock_new_budget.call_args
budget_request = call_args[1]["budget_obj"]
assert budget_request.budget_id is not None
assert budget_request.budget_id.startswith("team-budget-")
@pytest.mark.asyncio
async def test_upsert_team_member_budget_table_existing_budget():
"""
Test that upsert_team_member_budget_table updates an existing budget when team_member_budget_id exists.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LitellmUserRoles, LiteLLM_TeamTable, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import (
TeamMemberBudgetHandler,
)
mock_user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id"
)
team_table = MagicMock(spec=LiteLLM_TeamTable)
team_table.metadata = {"team_member_budget_id": "existing_budget_123"}
updated_kv = {
"team_id": "test_team_id",
"team_member_budget": 200.0,
"team_member_budget_duration": "60d",
"team_member_rpm_limit": 100,
}
mock_budget_response = MagicMock()
mock_budget_response.budget_id = "existing_budget_123"
with patch(
"litellm.proxy.management_endpoints.budget_management_endpoints.update_budget",
new_callable=AsyncMock
) as mock_update_budget:
mock_update_budget.return_value = mock_budget_response
result = await TeamMemberBudgetHandler.upsert_team_member_budget_table(
team_table=team_table,
user_api_key_dict=mock_user_api_key_dict,
updated_kv=updated_kv,
team_member_budget=200.0,
team_member_budget_duration="60d",
team_member_rpm_limit=100,
)
assert mock_update_budget.called
call_args = mock_update_budget.call_args
budget_request = call_args[1]["budget_obj"]
assert budget_request.budget_id == "existing_budget_123"
assert budget_request.max_budget == 200.0
assert budget_request.budget_duration == "60d"
assert budget_request.rpm_limit == 100
assert "team_member_budget_id" in result["metadata"]
assert result["metadata"]["team_member_budget_id"] == "existing_budget_123"
assert "team_member_budget" not in result
assert "team_member_budget_duration" not in result
assert "team_member_rpm_limit" not in result
@pytest.mark.asyncio
async def test_upsert_team_member_budget_table_no_existing_budget():
"""
Test that upsert_team_member_budget_table creates a new budget when team_member_budget_id does not exist.
"""
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import LitellmUserRoles, LiteLLM_TeamTable, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import (
TeamMemberBudgetHandler,
)
mock_user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id"
)
team_table = MagicMock(spec=LiteLLM_TeamTable)
team_table.metadata = {}
team_table.team_alias = "Test Team"
team_table.budget_duration = None
updated_kv = {
"team_id": "test_team_id",
"team_member_budget": 150.0,
"team_member_budget_duration": "45d",
}
mock_budget_response = MagicMock()
mock_budget_response.budget_id = "new_budget_456"
with patch(
"litellm.proxy.management_endpoints.budget_management_endpoints.new_budget",
new_callable=AsyncMock
) as mock_new_budget:
mock_new_budget.return_value = mock_budget_response
result = await TeamMemberBudgetHandler.upsert_team_member_budget_table(
team_table=team_table,
user_api_key_dict=mock_user_api_key_dict,
updated_kv=updated_kv,
team_member_budget=150.0,
team_member_budget_duration="45d",
)
assert mock_new_budget.called
assert "team_member_budget_id" in result["metadata"]
assert result["metadata"]["team_member_budget_id"] == "new_budget_456"
assert "team_member_budget" not in result
assert "team_member_budget_duration" not in result
@pytest.mark.asyncio
async def test_update_team_with_team_member_budget_duration():
"""
Test that team/update endpoint properly handles team_member_budget_duration.
"""
from unittest.mock import AsyncMock, MagicMock, Mock, patch
from fastapi import Request
from litellm.proxy._types import LitellmUserRoles, UpdateTeamRequest, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import update_team
mock_request = Mock(spec=Request)
mock_user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id"
)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch(
"litellm.proxy.proxy_server.llm_router"
) as mock_llm_router, patch(
"litellm.proxy.proxy_server.user_api_key_cache"
) as mock_cache, patch(
"litellm.proxy.proxy_server.proxy_logging_obj"
) as mock_logging, patch(
"litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin"
), patch(
"litellm.proxy.auth.auth_checks._cache_team_object"
) as mock_cache_team, patch(
"litellm.proxy.management_endpoints.team_endpoints.TeamMemberBudgetHandler.upsert_team_member_budget_table"
) as mock_upsert_budget:
mock_existing_team = MagicMock()
mock_existing_team.model_dump.return_value = {
"team_id": "test_team_id",
"team_alias": "test_team",
"metadata": {"team_member_budget_id": "budget_123"},
}
mock_existing_team.metadata = {"team_member_budget_id": "budget_123"}
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(
return_value=mock_existing_team
)
mock_updated_team = MagicMock()
mock_updated_team.team_id = "test_team_id"
mock_updated_team.model_dump.return_value = {"team_id": "test_team_id"}
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(
return_value=mock_updated_team
)
mock_prisma_client.jsonify_team_object = MagicMock(
side_effect=lambda db_data: db_data
)
def mock_upsert_side_effect(
team_table, user_api_key_dict, updated_kv, team_member_budget=None, team_member_rpm_limit=None, team_member_tpm_limit=None, team_member_budget_duration=None
):
result_kv = updated_kv.copy()
result_kv.pop("team_member_budget", None)
result_kv.pop("team_member_budget_duration", None)
return result_kv
mock_upsert_budget.side_effect = mock_upsert_side_effect
update_request = UpdateTeamRequest(
team_id="test_team_id",
team_alias="updated_alias",
team_member_budget=100.0,
team_member_budget_duration="30d",
)
result = await update_team(
data=update_request,
http_request=mock_request,
user_api_key_dict=mock_user_api_key_dict,
)
assert mock_upsert_budget.called
call_args = mock_upsert_budget.call_args
assert call_args[1]["team_member_budget"] == 100.0
assert call_args[1]["team_member_budget_duration"] == "30d"
assert mock_prisma_client.db.litellm_teamtable.update.called
update_call_args = mock_prisma_client.db.litellm_teamtable.update.call_args
update_data = update_call_args[1]["data"]
assert "team_member_budget" not in update_data
assert "team_member_budget_duration" not in update_data
@pytest.mark.asyncio
async def test_bulk_team_member_add_success():
"""

View file

@ -0,0 +1,49 @@
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi } from "vitest";
import DurationSelect from "./DurationSelect";
describe("DurationSelect", () => {
it("should render", () => {
render(<DurationSelect />);
expect(screen.getByRole("combobox")).toBeInTheDocument();
});
it("should render all three duration options", async () => {
const user = userEvent.setup();
render(<DurationSelect />);
const select = screen.getByRole("combobox");
await user.click(select);
expect(screen.getByText("Daily")).toBeInTheDocument();
expect(screen.getByText("Weekly")).toBeInTheDocument();
expect(screen.getByText("Monthly")).toBeInTheDocument();
});
it("should apply className prop", () => {
render(<DurationSelect className="test-class" />);
const select = screen.getByRole("combobox");
expect(select.closest(".test-class")).toBeInTheDocument();
});
it("should call onChange when an option is selected", async () => {
const user = userEvent.setup();
const onChange = vi.fn();
render(<DurationSelect onChange={onChange} />);
const select = screen.getByRole("combobox");
await user.click(select);
const dailyOption = screen.getByText("Daily");
await user.click(dailyOption);
expect(onChange).toHaveBeenCalledWith("24h", expect.any(Object));
});
it("should accept and pass value prop to Select", () => {
render(<DurationSelect value="7d" />);
const select = screen.getByRole("combobox");
expect(select).toBeInTheDocument();
});
});

View file

@ -0,0 +1,17 @@
import { Select } from "antd";
interface DurationSelectProps {
className?: string;
value?: string;
onChange?: (value: string) => void;
}
export default function DurationSelect({ className, value, onChange }: DurationSelectProps) {
return (
<Select className={className} value={value} onChange={onChange}>
<Select.Option value="24h">Daily</Select.Option>
<Select.Option value="7d">Weekly</Select.Option>
<Select.Option value="30d">Monthly</Select.Option>
</Select>
);
}

View file

@ -48,6 +48,7 @@ import EditLoggingSettings from "./EditLoggingSettings";
import MemberModal from "./EditMembership";
import MemberPermissions from "./member_permissions";
import TeamMembersComponent from "./team_member_view";
import DurationSelect from "../common_components/DurationSelect";
export interface TeamMembership {
user_id: string;
@ -413,6 +414,7 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
};
updateData.max_budget = mapEmptyStringToNull(updateData.max_budget);
updateData.team_member_budget_duration = values.team_member_budget_duration;
if (values.team_member_budget !== undefined) {
updateData.team_member_budget = Number(values.team_member_budget);
@ -650,6 +652,8 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
budget_duration: info.budget_duration,
team_member_tpm_limit: info.team_member_budget_table?.tpm_limit,
team_member_rpm_limit: info.team_member_budget_table?.rpm_limit,
team_member_budget: info.team_member_budget_table?.max_budget,
team_member_budget_duration: info.team_member_budget_table?.budget_duration,
guardrails: info.metadata?.guardrails || [],
disable_global_guardrails: info.metadata?.disable_global_guardrails || false,
metadata: info.metadata
@ -747,6 +751,13 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
<NumericalInput step={0.01} precision={2} style={{ width: "100%" }} />
</Form.Item>
<Form.Item label="Team Member Budget Duration" name="team_member_budget_duration">
<DurationSelect
onChange={(value) => form.setFieldValue("team_member_budget_duration", value)}
value={form.getFieldValue("team_member_budget_duration")}
/>
</Form.Item>
<Form.Item
label="Team Member Key Duration (eg: 1d, 1mo)"
name="team_member_key_duration"
@ -991,6 +1002,7 @@ const TeamInfoView: React.FC<TeamInfoProps> = ({
</Tooltip>
</Text>
<div>Max Budget: {info.team_member_budget_table?.max_budget || "No Limit"}</div>
<div>Budget Duration: {info.team_member_budget_table?.budget_duration || "No Limit"}</div>
<div>Key Duration: {info.metadata?.team_member_key_duration || "No Limit"}</div>
<div>TPM Limit: {info.team_member_budget_table?.tpm_limit || "No Limit"}</div>
<div>RPM Limit: {info.team_member_budget_table?.rpm_limit || "No Limit"}</div>