mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
92f7789f10
commit
1b8708fccc
16 changed files with 552 additions and 72 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -14,3 +14,4 @@ class ArizeConfig(BaseModel):
|
|||
api_key: Optional[str] = None
|
||||
protocol: Protocol
|
||||
endpoint: str
|
||||
project_name: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue