mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
✨ 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:
parent
aa19b7467c
commit
8d4b8ce23f
2 changed files with 271 additions and 18 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue