diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 5d8052493a1..b7b30d36f99 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -1,8 +1,32 @@ model_list: - model_name: openai/gpt-4o litellm_params: - model: openai/gpt-4o - api_key: os.environ/OPENAI_API_KEY + model: openai/gpt-4o-mini + api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5" + api_key: dummy + - model_name: "byok-wildcard/*" + litellm_params: + model: openai/* + - model_name: xai-grok-3 + litellm_params: + model: xai/grok-3 + - model_name: hosted_vllm/whisper-v3 + litellm_params: + model: hosted_vllm/whisper-v3 + api_base: "https://webhook.site/2f385e05-00aa-402b-86d1-efc9261471a5" + api_key: dummy + +# mcp_servers: +# github_mcp: +# url: "https://api.githubcopilot.com/mcp" +# auth_type: oauth2 +# authorization_url: https://github.com/login/oauth/authorize +# token_url: https://github.com/login/oauth/access_token +# client_id: os.environ/GITHUB_OAUTH_CLIENT_ID +# client_secret: os.environ/GITHUB_OAUTH_CLIENT_SECRET +# scopes: ["public_repo", "user:email"] +# allowed_tools: ["list_tools"] +# # disallowed_tools: ["repo_delete"] litellm_settings: callbacks: ["prometheus"] diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 7ce38c2d683..0eb2d27463f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -90,10 +90,10 @@ def _get_user_in_team( def _calculate_key_rotation_time(rotation_interval: str) -> datetime: """ Helper function to calculate the next rotation time for a key based on the rotation interval. - + Args: rotation_interval: String representing the rotation interval (e.g., '30d', '90d', '1h') - + Returns: datetime: The calculated next rotation time in UTC """ @@ -102,28 +102,34 @@ def _calculate_key_rotation_time(rotation_interval: str) -> datetime: return now + timedelta(seconds=interval_seconds) -def _set_key_rotation_fields(data: dict, auto_rotate: bool, rotation_interval: Optional[str]) -> None: +def _set_key_rotation_fields( + data: dict, auto_rotate: bool, rotation_interval: Optional[str] +) -> None: """ Helper function to set rotation fields in key data if auto_rotate is enabled. - + Args: data: Dictionary to update with rotation fields auto_rotate: Whether auto rotation is enabled rotation_interval: The rotation interval string (required if auto_rotate is True) """ if auto_rotate and rotation_interval: - data.update({ - "auto_rotate": auto_rotate, - "rotation_interval": rotation_interval, - "key_rotation_at": _calculate_key_rotation_time(rotation_interval) - }) + data.update( + { + "auto_rotate": auto_rotate, + "rotation_interval": rotation_interval, + "key_rotation_at": _calculate_key_rotation_time(rotation_interval), + } + ) def _is_allowed_to_make_key_request( - user_api_key_dict: UserAPIKeyAuth, user_id: Optional[str], team_id: Optional[str] + user_api_key_dict: UserAPIKeyAuth, + user_id: Optional[str], + team_id: Optional[str], ) -> bool: """ - Assert user only creates keys for themselves + Assert user only creates/updates keys for themselves Relevant issue: https://github.com/BerriAI/litellm/issues/7336 """ @@ -332,6 +338,7 @@ def common_key_access_checks( data: Union[GenerateKeyRequest, UpdateKeyRequest], llm_router: Optional[Router], premium_user: bool, + user_id: Optional[str] = None, ) -> Literal[True]: """ Check if user is allowed to make a key request, for this key @@ -339,7 +346,7 @@ def common_key_access_checks( try: _is_allowed_to_make_key_request( user_api_key_dict=user_api_key_dict, - user_id=data.user_id, + user_id=user_id or data.user_id, team_id=data.team_id, ) except AssertionError as e: @@ -1136,13 +1143,6 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") - common_key_access_checks( - user_api_key_dict=user_api_key_dict, - data=data, - llm_router=llm_router, - premium_user=premium_user, - ) - existing_key_row = await prisma_client.get_data( token=data.key, table_name="key", query_type="find_unique" ) @@ -1153,6 +1153,14 @@ async def update_key_fn( detail={"error": f"Team not found, passed team_id={data.team_id}"}, ) + common_key_access_checks( + user_api_key_dict=user_api_key_dict, + data=data, + user_id=existing_key_row.user_id, + llm_router=llm_router, + premium_user=premium_user, + ) + # check if user has permission to update key await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( user_api_key_dict=user_api_key_dict, @@ -1198,9 +1206,9 @@ async def update_key_fn( # Handle rotation fields if auto_rotate is being enabled _set_key_rotation_fields( - non_default_values, - non_default_values.get("auto_rotate", False), - non_default_values.get("rotation_interval") + non_default_values, + non_default_values.get("auto_rotate", False), + non_default_values.get("rotation_interval"), ) _data = {**non_default_values, "token": key} @@ -1602,8 +1610,6 @@ def _check_model_access_group( return True - - async def generate_key_helper_fn( # noqa: PLR0915 request_type: Literal[ "user", "key" @@ -1766,12 +1772,12 @@ async def generate_key_helper_fn( # noqa: PLR0915 "allowed_routes": allowed_routes or [], "object_permission_id": object_permission_id, } - + # Add rotation fields if auto_rotate is enabled _set_key_rotation_fields( data=key_data, auto_rotate=auto_rotate or False, - rotation_interval=rotation_interval + rotation_interval=rotation_interval, ) if ( diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index e2abf9c4c75..c364c311b14 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -6,7 +6,6 @@ sys.path.insert(0, os.path.abspath("../..")) import asyncio import logging -from litellm._uuid import uuid from datetime import datetime, timedelta, timezone from unittest.mock import MagicMock, call, patch @@ -16,6 +15,7 @@ from prometheus_client import REGISTRY, CollectorRegistry import litellm from litellm import completion from litellm._logging import verbose_logger +from litellm._uuid import uuid from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.utils import ( StandardLoggingHiddenParams, @@ -1033,10 +1033,10 @@ def test_deployment_state_management(prometheus_logger): # Test set_deployment_healthy (state=0) prometheus_logger.set_deployment_healthy(**test_params) prometheus_logger.litellm_deployment_state.labels.assert_called_with( - test_params["litellm_model_name"], - test_params["model_id"], - test_params["api_base"], - test_params["api_provider"], + litellm_model_name=test_params["litellm_model_name"], + model_id=test_params["model_id"], + api_base=test_params["api_base"], + api_provider=test_params["api_provider"], ) prometheus_logger.litellm_deployment_state.labels().set.assert_called_with(0) @@ -1153,22 +1153,28 @@ def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch): # Configure tags with wildcard patterns monkeypatch.setattr( - "litellm.custom_prometheus_tags", - ["User-Agent: curl/*", "User-Agent: python-requests/*", "Environment: prod*", "Service: api-gateway*", "exact-match"] + "litellm.custom_prometheus_tags", + [ + "User-Agent: curl/*", + "User-Agent: python-requests/*", + "Environment: prod*", + "Service: api-gateway*", + "exact-match", + ], ) - + # Test tags that should match the wildcard patterns tags = [ - "User-Agent: curl/7.68.0", - "User-Agent: python-requests/2.28.1", + "User-Agent: curl/7.68.0", + "User-Agent: python-requests/2.28.1", "Environment: production", "Service: api-gateway-v2", "exact-match", - "other-tag" + "other-tag", ] - + result = get_custom_labels_from_tags(tags) - + expected = { "tag_User_Agent__curl__": "true", # matches "User-Agent: curl/*" "tag_User_Agent__python_requests__": "true", # matches "User-Agent: python-requests/*" @@ -1176,7 +1182,7 @@ def test_get_custom_labels_from_tags_wildcard_patterns(monkeypatch): "tag_Service__api_gateway_": "true", # matches "Service: api-gateway*" "tag_exact_match": "true", # exact match } - + assert result == expected @@ -1186,26 +1192,26 @@ def test_get_custom_labels_from_tags_wildcard_no_matches(monkeypatch): # Configure tags with wildcard patterns monkeypatch.setattr( - "litellm.custom_prometheus_tags", - ["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"] + "litellm.custom_prometheus_tags", + ["User-Agent: firefox/*", "Environment: dev*", "Service: web-app*"], ) - + # Test tags that should NOT match the wildcard patterns tags = [ "User-Agent: curl/7.68.0", # doesn't match "User-Agent: firefox/*" - "Environment: production", # doesn't match "Environment: dev*" + "Environment: production", # doesn't match "Environment: dev*" "Service: api-gateway-v2", # doesn't match "Service: web-app*" - "other-tag" + "other-tag", ] - + result = get_custom_labels_from_tags(tags) - + expected = { "tag_User_Agent__firefox__": "false", # no match for "User-Agent: firefox/*" "tag_Environment__dev_": "false", # no match for "Environment: dev*" "tag_Service__web_app_": "false", # no match for "Service: web-app*" } - + assert result == expected @@ -1216,48 +1222,69 @@ def test_tag_matches_wildcard_configured_pattern(): ) # Test cases that should match - assert _tag_matches_wildcard_configured_pattern( - tags=["User-Agent: curl/7.68.0", "prod", "other"], - configured_tag="User-Agent: curl/*" - ) is True - - assert _tag_matches_wildcard_configured_pattern( - tags=["User-Agent: python-requests/2.28.1", "test"], - configured_tag="User-Agent: python-requests/*" - ) is True - - assert _tag_matches_wildcard_configured_pattern( - tags=["Environment: production", "debug"], - configured_tag="Environment: prod*" - ) is True - + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["User-Agent: curl/7.68.0", "prod", "other"], + configured_tag="User-Agent: curl/*", + ) + is True + ) + + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["User-Agent: python-requests/2.28.1", "test"], + configured_tag="User-Agent: python-requests/*", + ) + is True + ) + + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["Environment: production", "debug"], + configured_tag="Environment: prod*", + ) + is True + ) + # Test exact match (no wildcard) - assert _tag_matches_wildcard_configured_pattern( - tags=["prod", "test"], - configured_tag="prod" - ) is True - + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["prod", "test"], configured_tag="prod" + ) + is True + ) + # Test cases that should NOT match - assert _tag_matches_wildcard_configured_pattern( - tags=["User-Agent: firefox/98.0", "prod"], - configured_tag="User-Agent: curl/*" - ) is False - - assert _tag_matches_wildcard_configured_pattern( - tags=["Environment: development", "test"], - configured_tag="Environment: prod*" - ) is False - - assert _tag_matches_wildcard_configured_pattern( - tags=["staging", "test"], - configured_tag="prod" - ) is False - + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["User-Agent: firefox/98.0", "prod"], + configured_tag="User-Agent: curl/*", + ) + is False + ) + + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["Environment: development", "test"], + configured_tag="Environment: prod*", + ) + is False + ) + + assert ( + _tag_matches_wildcard_configured_pattern( + tags=["staging", "test"], configured_tag="prod" + ) + is False + ) + # Test with empty tags - assert _tag_matches_wildcard_configured_pattern( - tags=[], - configured_tag="User-Agent: curl/*" - ) is False + assert ( + _tag_matches_wildcard_configured_pattern( + tags=[], configured_tag="User-Agent: curl/*" + ) + is False + ) @pytest.mark.asyncio(scope="session") @@ -1920,12 +1947,12 @@ def test_set_llm_deployment_success_metrics_with_label_filtering(): async def test_prometheus_token_metrics_with_prometheus_config(): """ Test that validates the renamed token metrics are incremented correctly with a prometheus config. - + This test ensures that after the metric renaming (git diff): - litellm_total_tokens -> litellm_total_tokens_metric - - litellm_input_tokens -> litellm_input_tokens_metric + - litellm_input_tokens -> litellm_input_tokens_metric - litellm_output_tokens -> litellm_output_tokens_metric - + All three metrics should be properly incremented when making a successful completion request. """ from prometheus_client import CollectorRegistry, Counter @@ -1937,39 +1964,39 @@ async def test_prometheus_token_metrics_with_prometheus_config(): collectors = list(REGISTRY._collector_to_names.keys()) for collector in collectors: REGISTRY.unregister(collector) - + # Set up prometheus configuration that includes the token metrics config = [ PrometheusMetricsConfig( group="token_metrics_test", metrics=[ "litellm_total_tokens_metric", - "litellm_input_tokens_metric", + "litellm_input_tokens_metric", "litellm_output_tokens_metric", - "litellm_requests_metric" + "litellm_requests_metric", ], include_labels=[ "model", - "hashed_api_key", + "hashed_api_key", "api_key_alias", "team", - "team_alias" + "team_alias", ], ) ] - + # Mock litellm.prometheus_metrics_config with patch("litellm.prometheus_metrics_config", config): # Create PrometheusLogger with the configuration prometheus_logger = PrometheusLogger() - + # Test data with specific token counts standard_logging_payload = create_standard_logging_payload() standard_logging_payload["total_tokens"] = 1500 standard_logging_payload["prompt_tokens"] = 900 standard_logging_payload["completion_tokens"] = 600 standard_logging_payload["response_cost"] = 0.075 - + kwargs = { "model": "gpt-3.5-turbo", "stream": False, @@ -1983,7 +2010,7 @@ async def test_prometheus_token_metrics_with_prometheus_config(): } }, "start_time": datetime.now() - timedelta(seconds=2), - "completion_start_time": datetime.now() - timedelta(seconds=1), + "completion_start_time": datetime.now() - timedelta(seconds=1), "api_call_start_time": datetime.now() - timedelta(seconds=1.5), "end_time": datetime.now(), "standard_logging_object": standard_logging_payload, @@ -1999,69 +2026,75 @@ async def test_prometheus_token_metrics_with_prometheus_config(): print("final registry values", REGISTRY._collector_to_names) - # Get metric collectors directly from registry + # Get metric collectors directly from registry metric_collectors = {} for collector, names in REGISTRY._collector_to_names.items(): metric_name = names[0] # First name is the base metric name metric_collectors[metric_name] = collector print("=== Final Metric Values (Direct Access) ===") - - # Expected values + + # Expected values expected_values = { "litellm_total_tokens_metric": 1500.0, "litellm_input_tokens_metric": 900.0, "litellm_output_tokens_metric": 600.0, - "litellm_requests_metric": 1.0 + "litellm_requests_metric": 1.0, } - + expected_label_values = { - 'api_key_alias': 'test_alias', - 'hashed_api_key': 'test_hash', - 'model': 'gpt-3.5-turbo', - 'team': 'test_team', - 'team_alias': 'test_team_alias' + "api_key_alias": "test_alias", + "hashed_api_key": "test_hash", + "model": "gpt-3.5-turbo", + "team": "test_team", + "team_alias": "test_team_alias", } # Validate each metric directly for metric_name, expected_value in expected_values.items(): if metric_name in metric_collectors: collector = metric_collectors[metric_name] - + # Get all samples for this metric samples = list(collector.collect())[0].samples - + # Find the _total sample (the actual counter value) total_sample = None for sample in samples: - if sample.name.endswith('_total'): + if sample.name.endswith("_total"): total_sample = sample break - + if total_sample: actual_value = total_sample.value actual_labels = total_sample.labels - - print(f"✓ {metric_name}: expected={expected_value}, actual={actual_value}") + + print( + f"✓ {metric_name}: expected={expected_value}, actual={actual_value}" + ) print(f" Labels: {actual_labels}") - + # Validate the value - assert actual_value == expected_value, f"Expected {expected_value}, got {actual_value} for {metric_name}" - + assert ( + actual_value == expected_value + ), f"Expected {expected_value}, got {actual_value} for {metric_name}" + # Validate the labels - for label_key, expected_label_value in expected_label_values.items(): + for ( + label_key, + expected_label_value, + ) in expected_label_values.items(): actual_label_value = actual_labels.get(label_key) - assert actual_label_value == expected_label_value, f"Expected label {label_key}={expected_label_value}, got {actual_label_value}" - + assert ( + actual_label_value == expected_label_value + ), f"Expected label {label_key}={expected_label_value}, got {actual_label_value}" + print(f" ✓ {metric_name} VALIDATED") else: raise AssertionError(f"No _total sample found for {metric_name}") else: raise AssertionError(f"Metric {metric_name} not found in registry") - + print("✓ All token metrics validated successfully!") # check final value of metrics in registry - - -