✨ feat(health): add database-backed test email resolution for email health checks

- 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
This commit is contained in:
yangdx 2026-04-01 14:39:49 +08:00
parent aa19b7467c
commit 8d4b8ce23f
2 changed files with 271 additions and 18 deletions

View file

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

View file

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