From 8d4b8ce23fe383f3ae945368c047a217bc48dd9a Mon Sep 17 00:00:00 2001 From: yangdx Date: Wed, 1 Apr 2026 14:39:49 +0800 Subject: [PATCH] =?UTF-8?q?=E2=9C=A8=20feat(health):=20add=20database-back?= =?UTF-8?q?ed=20test=20email=20resolution=20for=20email=20health=20checks?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - add _resolve_test_email_address() to fetch TEST_EMAIL_ADDRESS from DB when store_model_in_db is enabled - add helper functions _parse_config_row_param_value() and _is_truthy_config_flag() for config parsing - support encrypted values with fallback to environment variable - refactor email health check to use new resolver instead of direct config access ✅ test(health): add comprehensive tests for db-backed email resolution - test db email takes precedence when store_model_in_db is enabled - test fallback to env var when db is disabled or config missing - test json string parsing for environment_variables row --- .../health_endpoints/_health_endpoints.py | 106 ++++++++-- .../health_endpoints/test_health_endpoints.py | 183 ++++++++++++++++++ 2 files changed, 271 insertions(+), 18 deletions(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 4fa24143c9d..3fb04306200 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -1,5 +1,6 @@ import asyncio import copy +import json import logging import os import time @@ -36,7 +37,8 @@ from litellm.proxy.health_check import ( from litellm.proxy.middleware.in_flight_requests_middleware import ( get_in_flight_requests, ) -from litellm.secret_managers.main import get_secret +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper +from litellm.secret_managers.main import get_secret, get_secret_bool #### Health ENDPOINTS #### @@ -145,6 +147,90 @@ def get_callback_identifier(callback): return callback_name(callback) +def _parse_config_row_param_value(param_value: Any) -> dict: + if param_value is None: + return {} + + if isinstance(param_value, str): + try: + parsed_value = json.loads(param_value) + except json.JSONDecodeError: + return {} + return parsed_value if isinstance(parsed_value, dict) else {} + + if isinstance(param_value, dict): + return dict(param_value) + + try: + parsed_value = dict(param_value) + except (TypeError, ValueError): + return {} + + return parsed_value if isinstance(parsed_value, dict) else {} + + +def _is_truthy_config_flag(value: Any) -> bool: + if isinstance(value, bool): + return value + + if isinstance(value, str): + return value.strip().lower() == "true" + + if value is None: + return False + + return bool(value) + + +async def _resolve_test_email_address(prisma_client: Any) -> Optional[str]: + test_email_address = os.getenv("TEST_EMAIL_ADDRESS") + + try: + store_model_in_db = ( + get_secret_bool("STORE_MODEL_IN_DB", default_value=False) is True + ) + + if not store_model_in_db and prisma_client is not None: + general_settings_row = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "general_settings"} + ) + general_settings = _parse_config_row_param_value( + getattr(general_settings_row, "param_value", None) + ) + store_model_in_db = _is_truthy_config_flag( + general_settings.get("store_model_in_db") + ) + + if not store_model_in_db or prisma_client is None: + return test_email_address + + environment_variables_row = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "environment_variables"} + ) + environment_variables = _parse_config_row_param_value( + getattr(environment_variables_row, "param_value", None) + ) + db_test_email_address = environment_variables.get("TEST_EMAIL_ADDRESS") + + if db_test_email_address is None: + return test_email_address + + decrypted_test_email_address = decrypt_value_helper( + value=db_test_email_address, + key="TEST_EMAIL_ADDRESS", + exception_type="debug", + return_original_value=True, + ) + + return decrypted_test_email_address or test_email_address + except Exception as e: + verbose_proxy_logger.debug( + "Falling back to TEST_EMAIL_ADDRESS from env after DB lookup failed: %s", + str(e), + ) + return test_email_address + + router = APIRouter() services = Union[ Literal[ @@ -438,22 +524,6 @@ async def health_services_endpoint( # noqa: PLR0915 }, ) if service == "email": - from litellm.proxy.proxy_server import proxy_config, store_model_in_db - - # TEST_EMAIL_ADDRESS is stored encrypted in the DB when the proxy - # is running in DB mode. Calling get_config() ensures the value is - # freshly decrypted and available both in the returned config dict - # and in os.environ. For YAML / env-var deployments the call is a - # cheap no-op and os.getenv() still works as the fallback. - if store_model_in_db and prisma_client is not None: - _fresh_config = await proxy_config.get_config() - _env_vars = _fresh_config.get("environment_variables", {}) - _test_email_address = _env_vars.get( - "TEST_EMAIL_ADDRESS" - ) or os.getenv("TEST_EMAIL_ADDRESS") - else: - _test_email_address = os.getenv("TEST_EMAIL_ADDRESS") - webhook_event = WebhookEvent( event="key_created", event_group=Litellm_EntityType.KEY, @@ -463,7 +533,7 @@ async def health_services_endpoint( # noqa: PLR0915 spend=0, max_budget=0, user_id=user_api_key_dict.user_id, - user_email=_test_email_address, + user_email=await _resolve_test_email_address(prisma_client), team_id=user_api_key_dict.team_id, ) diff --git a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py index bc3aec58991..f9127210a3b 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -1,3 +1,4 @@ +import json import os import sys import time @@ -337,6 +338,188 @@ async def test_health_services_endpoint_sqs(status, error_message): mock_instance.async_health_check.assert_awaited_once() +@pytest.mark.asyncio +async def test_health_services_endpoint_email_should_use_test_email_address_from_db_when_store_model_in_db_enabled(): + mock_prisma = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock( + side_effect=[ + SimpleNamespace(param_value={"store_model_in_db": True}), + SimpleNamespace(param_value={"TEST_EMAIL_ADDRESS": "encrypted-db-value"}), + ] + ) + mock_slack_alerting = SimpleNamespace( + send_key_created_or_user_invited_email=AsyncMock() + ) + mock_proxy_logging_obj = SimpleNamespace( + slack_alerting_instance=mock_slack_alerting + ) + mock_user_api_key_dict = SimpleNamespace( + token="test-token", + user_id="test-user", + team_id="test-team", + ) + + with patch.dict(os.environ, {"TEST_EMAIL_ADDRESS": "env@example.com"}), patch( + "litellm.proxy.proxy_server.general_settings", + {}, + ), patch( + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging_obj, + ), patch( + "litellm.proxy.health_endpoints._health_endpoints.get_secret_bool", + return_value=False, + ), patch( + "litellm.proxy.health_endpoints._health_endpoints.decrypt_value_helper", + return_value="db@example.com", + ): + result = await health_services_endpoint( + service="email", + user_api_key_dict=mock_user_api_key_dict, + ) + + assert result["status"] == "success" + assert ( + mock_prisma.db.litellm_config.find_unique.await_args_list[0].kwargs["where"] + == {"param_name": "general_settings"} + ) + assert ( + mock_prisma.db.litellm_config.find_unique.await_args_list[1].kwargs["where"] + == {"param_name": "environment_variables"} + ) + webhook_event = ( + mock_slack_alerting.send_key_created_or_user_invited_email.await_args.kwargs[ + "webhook_event" + ] + ) + assert webhook_event.user_email == "db@example.com" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "store_model_in_db_secret,general_settings_row,environment_variables_row,expected_db_calls", + [ + (False, SimpleNamespace(param_value={"store_model_in_db": False}), None, 1), + (True, None, None, 1), + ], + ids=["db-disabled", "config-row-missing"], +) +async def test_health_services_endpoint_email_should_fall_back_to_env_test_email_address_when_db_disabled_or_missing( + store_model_in_db_secret, + general_settings_row, + environment_variables_row, + expected_db_calls, +): + mock_prisma = MagicMock() + db_rows = [] + if general_settings_row is not None: + db_rows.append(general_settings_row) + if store_model_in_db_secret: + db_rows.append(environment_variables_row) + mock_prisma.db.litellm_config.find_unique = AsyncMock(side_effect=db_rows) + mock_slack_alerting = SimpleNamespace( + send_key_created_or_user_invited_email=AsyncMock() + ) + mock_proxy_logging_obj = SimpleNamespace( + slack_alerting_instance=mock_slack_alerting + ) + mock_user_api_key_dict = SimpleNamespace( + token="test-token", + user_id="test-user", + team_id="test-team", + ) + + with patch.dict(os.environ, {"TEST_EMAIL_ADDRESS": "env@example.com"}), patch( + "litellm.proxy.proxy_server.general_settings", + {}, + ), patch( + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging_obj, + ), patch( + "litellm.proxy.health_endpoints._health_endpoints.get_secret_bool", + return_value=store_model_in_db_secret, + ), patch( + "litellm.proxy.health_endpoints._health_endpoints.decrypt_value_helper" + ) as decrypt_mock: + result = await health_services_endpoint( + service="email", + user_api_key_dict=mock_user_api_key_dict, + ) + + assert result["status"] == "success" + webhook_event = ( + mock_slack_alerting.send_key_created_or_user_invited_email.await_args.kwargs[ + "webhook_event" + ] + ) + assert webhook_event.user_email == "env@example.com" + assert mock_prisma.db.litellm_config.find_unique.await_count == expected_db_calls + decrypt_mock.assert_not_called() + + +@pytest.mark.asyncio +async def test_health_services_endpoint_email_should_accept_json_string_environment_variables(): + mock_prisma = MagicMock() + mock_prisma.db.litellm_config.find_unique = AsyncMock( + return_value=SimpleNamespace( + param_value=json.dumps( + {"TEST_EMAIL_ADDRESS": "json-string-db-value"} + ) + ) + ) + mock_slack_alerting = SimpleNamespace( + send_key_created_or_user_invited_email=AsyncMock() + ) + mock_proxy_logging_obj = SimpleNamespace( + slack_alerting_instance=mock_slack_alerting + ) + mock_user_api_key_dict = SimpleNamespace( + token="test-token", + user_id="test-user", + team_id="test-team", + ) + + with patch.dict(os.environ, {"TEST_EMAIL_ADDRESS": "env@example.com"}), patch( + "litellm.proxy.proxy_server.general_settings", + {}, + ), patch( + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging_obj, + ), patch( + "litellm.proxy.health_endpoints._health_endpoints.get_secret_bool", + return_value=True, + ), patch( + "litellm.proxy.health_endpoints._health_endpoints.decrypt_value_helper", + return_value="json@example.com", + ) as decrypt_mock: + result = await health_services_endpoint( + service="email", + user_api_key_dict=mock_user_api_key_dict, + ) + + assert result["status"] == "success" + decrypt_mock.assert_called_once_with( + value="json-string-db-value", + key="TEST_EMAIL_ADDRESS", + exception_type="debug", + return_original_value=True, + ) + webhook_event = ( + mock_slack_alerting.send_key_created_or_user_invited_email.await_args.kwargs[ + "webhook_event" + ] + ) + assert webhook_event.user_email == "json@example.com" + + @pytest.mark.asyncio async def test_health_license_endpoint_with_active_license(): license_data = {