From b7c45991d8776e514ce088e06cbede75d613d49c Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 31 Jan 2026 12:25:04 -0800 Subject: [PATCH] Fix health endpoints --- .../health_endpoints/_health_endpoints.py | 45 +++++- .../health_endpoints/test_health_endpoints.py | 133 ++++++++++++++++++ 2 files changed, 177 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index eddd64c36c7..da90696ec2d 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -32,6 +32,7 @@ from litellm.proxy.health_check import ( run_with_timeout, ) from litellm.secret_managers.main import get_secret +from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry #### Health ENDPOINTS #### @@ -106,6 +107,35 @@ def _resolve_os_environ_variables(params: dict) -> dict: return resolved_root +def get_callback_identifier(callback): + """ + Get the callback identifier string, handling both strings and objects. + + This function extracts a string identifier from a callback, which can be: + - A string (returned as-is) + - An object with a callback_name attribute + - An object registered in CustomLoggerRegistry + - Falls back to callback_name() helper function + + Args: + callback: The callback to identify (can be str or object) + + Returns: + str: The callback identifier string + """ + if isinstance(callback, str): + return callback + if hasattr(callback, 'callback_name') and callback.callback_name: + return callback.callback_name + if hasattr(callback, '__class__'): + callback_strs = CustomLoggerRegistry.get_all_callback_strs_from_class_type(callback.__class__) + if hasattr(callback, 'callback_name') and callback.callback_name in callback_strs: + return callback.callback_name + if callback_strs: + return callback_strs[0] + return callback_name(callback) + + router = APIRouter() services = Union[ Literal[ @@ -203,11 +233,24 @@ async def health_services_endpoint( # noqa: PLR0915 }, ) + service_in_success_callbacks = False + if service in litellm.success_callback: + service_in_success_callbacks = True + else: + for cb in litellm.success_callback: + if hasattr(cb, 'callback_name') and cb.callback_name == service: + service_in_success_callbacks = True + break + cb_id = get_callback_identifier(cb) + if cb_id == service: + service_in_success_callbacks = True + break + if ( service == "openmeter" or service == "braintrust" or service == "generic_api" - or (service in litellm.success_callback and service != "langfuse") + or (service_in_success_callbacks and service != "langfuse") ): _ = await litellm.acompletion( model="openai/litellm-mock-response-model", 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 d6393bc6414..97ab8355343 100644 --- a/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/test_litellm/proxy/health_endpoints/test_health_endpoints.py @@ -15,6 +15,7 @@ from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, Prisma from litellm.proxy.health_endpoints._health_endpoints import ( _db_health_readiness_check, db_health_cache, + get_callback_identifier, health_license_endpoint, health_services_endpoint, ) @@ -478,3 +479,135 @@ def test_health_readiness(proxy_client): f"Unexpected db status: {db_status}" print("="*60 + "\n") + + +def test_get_callback_identifier_string_and_object_with_callback_name(): + """ + Test get_callback_identifier with string callbacks and objects with callback_name attribute. + + Covers: + - String callback (returned as-is) + - Object with callback_name attribute + - Object with empty/None callback_name (should fall through to other checks) + """ + from litellm.proxy.health_endpoints._health_endpoints import get_callback_identifier + + # Test 1: String callback should be returned as-is + assert get_callback_identifier("datadog") == "datadog" + assert get_callback_identifier("langfuse") == "langfuse" + + # Test 2: Object with callback_name attribute + class MockCallbackWithName: + def __init__(self, name): + self.callback_name = name + + callback_obj = MockCallbackWithName("custom_callback") + assert get_callback_identifier(callback_obj) == "custom_callback" + + # Test 3: Object with empty callback_name should fall through + callback_obj_empty = MockCallbackWithName("") + # This should fall through to CustomLoggerRegistry or callback_name() fallback + # We'll verify it doesn't return empty string + result = get_callback_identifier(callback_obj_empty) + assert result != "" # Should not return empty string + assert isinstance(result, str) # Should still return a string + + +def test_get_callback_identifier_custom_logger_registry_and_fallback(): + """ + Test get_callback_identifier with CustomLoggerRegistry lookup and fallback scenarios. + + Covers: + - Object registered in CustomLoggerRegistry + - Object with callback_name that matches registry entry + - Fallback to callback_name() helper function + """ + from litellm.proxy.health_endpoints._health_endpoints import get_callback_identifier + from litellm.litellm_core_utils.custom_logger_registry import CustomLoggerRegistry + + # Test 1: Object registered in CustomLoggerRegistry (without callback_name attribute) + # Mock a class that's registered in the registry + class MockRegisteredLogger: + pass + + # Mock the registry to return callback strings for our mock class + with patch.object( + CustomLoggerRegistry, + 'get_all_callback_strs_from_class_type', + return_value=['mock_logger'] + ): + mock_instance = MockRegisteredLogger() + result = get_callback_identifier(mock_instance) + assert result == "mock_logger" + + # Test 2: Object with callback_name that matches registry entry + class MockCallbackWithMatchingName: + def __init__(self): + self.callback_name = "matched_name" + + callback_with_matching = MockCallbackWithMatchingName() + # Mock registry to return list containing the matching name + with patch.object( + CustomLoggerRegistry, + 'get_all_callback_strs_from_class_type', + return_value=['matched_name', 'other_name'] + ): + result = get_callback_identifier(callback_with_matching) + assert result == "matched_name" + + # Test 3: Object with falsy callback_name (empty string), should use registry + class MockCallbackWithEmptyName: + def __init__(self): + self.callback_name = "" # Empty string is falsy + + callback_empty = MockCallbackWithEmptyName() + # Mock registry to return list - should use first registry entry since callback_name is falsy + with patch.object( + CustomLoggerRegistry, + 'get_all_callback_strs_from_class_type', + return_value=['registry_name'] + ): + result = get_callback_identifier(callback_empty) + assert result == "registry_name" + + # Test 3b: Object with truthy callback_name not in registry - returns callback_name immediately + # (This tests that truthy callback_name takes precedence over registry) + class MockCallbackWithNonMatchingName: + def __init__(self): + self.callback_name = "non_matching" + + callback_non_matching = MockCallbackWithNonMatchingName() + # Even if registry has different values, truthy callback_name is returned first + with patch.object( + CustomLoggerRegistry, + 'get_all_callback_strs_from_class_type', + return_value=['registry_name'] + ): + result = get_callback_identifier(callback_non_matching) + # Should return callback_name because it's truthy (checked before registry) + assert result == "non_matching" + + # Test 4: Object not in registry, falls back to callback_name() helper + class UnregisteredCallback: + def __init__(self): + pass + + unregistered = UnregisteredCallback() + # Mock registry to return empty list (not registered) + with patch.object( + CustomLoggerRegistry, + 'get_all_callback_strs_from_class_type', + return_value=[] + ): + result = get_callback_identifier(unregistered) + # Should fall back to callback_name() which returns __class__.__name__ + assert result == "UnregisteredCallback" + + # Test 5: Function callback (not a class instance) + def my_callback_function(): + pass + + # Function won't have __class__, so it will skip registry check and go to callback_name() + result = get_callback_identifier(my_callback_function) + # Should fall back to callback_name() which returns __name__ + assert result == "my_callback_function"