mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(prometheus): preserve failed deployment labels across fallbacks
This commit is contained in:
parent
f5e4aa38ba
commit
25a68ffcb6
4 changed files with 98 additions and 11 deletions
|
|
@ -2295,6 +2295,34 @@ class PrometheusLogger(CustomLogger):
|
|||
_labels,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_deployment_failure_model_id(
|
||||
request_kwargs: Mapping[str, object], standard_logging_payload: StandardLoggingPayload
|
||||
) -> str | None:
|
||||
exception: Final = request_kwargs.get("exception")
|
||||
failed_deployment_id: Final = getattr(exception, "failed_deployment_id", None)
|
||||
if isinstance(failed_deployment_id, str) and failed_deployment_id:
|
||||
return failed_deployment_id
|
||||
|
||||
standard_model_id: Final = standard_logging_payload.get("model_id")
|
||||
if standard_model_id:
|
||||
return standard_model_id
|
||||
|
||||
litellm_params: Final = request_kwargs.get("litellm_params")
|
||||
if not isinstance(litellm_params, Mapping):
|
||||
return None
|
||||
for metadata_key in ("litellm_metadata", "metadata"):
|
||||
metadata = litellm_params.get(metadata_key)
|
||||
if not isinstance(metadata, Mapping):
|
||||
continue
|
||||
model_info = metadata.get("model_info")
|
||||
if not isinstance(model_info, Mapping):
|
||||
continue
|
||||
model_id = model_info.get("id")
|
||||
if isinstance(model_id, str) and model_id:
|
||||
return model_id
|
||||
return None
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
verbose_logger.debug(
|
||||
"prometheus Logging - Enters failure logging function (kwargs keys: %s)",
|
||||
|
|
@ -2318,6 +2346,13 @@ class PrometheusLogger(CustomLogger):
|
|||
user_api_team: Final = standard_logging_payload["metadata"]["user_api_key_team_id"]
|
||||
user_api_team_alias: Final = standard_logging_payload["metadata"]["user_api_key_team_alias"]
|
||||
user_api_key_org_id: Final = standard_logging_payload["metadata"].get("user_api_key_org_id")
|
||||
model_id: Final = (
|
||||
self._get_deployment_failure_model_id(
|
||||
request_kwargs=kwargs,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
)
|
||||
or ""
|
||||
)
|
||||
|
||||
try:
|
||||
enum_values: Final = UserAPIKeyLabelValues(
|
||||
|
|
@ -2328,7 +2363,7 @@ class PrometheusLogger(CustomLogger):
|
|||
team=user_api_team,
|
||||
team_alias=user_api_team_alias,
|
||||
user=user_id,
|
||||
model_id=standard_logging_payload.get("model_id", ""),
|
||||
model_id=model_id,
|
||||
custom_metadata_labels=get_custom_labels_from_metadata(
|
||||
metadata=_get_combined_custom_metadata_from_standard_logging_payload(
|
||||
standard_logging_payload=standard_logging_payload
|
||||
|
|
@ -2728,17 +2763,11 @@ class PrometheusLogger(CustomLogger):
|
|||
litellm_model_name: Final = request_kwargs.get("model", None)
|
||||
model_group = standard_logging_payload.get("model_group", None)
|
||||
api_base: Final = standard_logging_payload.get("api_base", None)
|
||||
model_id = standard_logging_payload.get("model_id", None)
|
||||
exception: Final = request_kwargs.get("exception", None)
|
||||
|
||||
# Fallback: model_id from litellm_metadata.model_info
|
||||
if model_id is None:
|
||||
_model_info: Final = (
|
||||
(_litellm_params.get("litellm_metadata") or {}).get("model_info")
|
||||
or (_litellm_params.get("metadata") or {}).get("model_info")
|
||||
or {}
|
||||
)
|
||||
model_id = _model_info.get("id")
|
||||
model_id: Final = self._get_deployment_failure_model_id(
|
||||
request_kwargs=request_kwargs,
|
||||
standard_logging_payload=standard_logging_payload,
|
||||
)
|
||||
|
||||
# Fallback: model_group from litellm_metadata
|
||||
if model_group is None:
|
||||
|
|
|
|||
|
|
@ -8555,6 +8555,7 @@ class Router:
|
|||
try:
|
||||
await _callback.async_pre_call_check(deployment, parent_otel_span)
|
||||
except litellm.RateLimitError as e:
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
|
|
@ -8573,6 +8574,7 @@ class Router:
|
|||
)
|
||||
raise e
|
||||
except Exception as e:
|
||||
self._set_failed_deployment_id_on_exception(e, deployment)
|
||||
## LOG FAILURE EVENT
|
||||
if logging_obj is not None:
|
||||
asyncio.create_task(
|
||||
|
|
|
|||
|
|
@ -83,6 +83,11 @@ def _requested_model_values(metric) -> set[str]:
|
|||
return {sample_key[index] for sample_key in metric._metrics}
|
||||
|
||||
|
||||
def _model_id_values(metric) -> set[str]:
|
||||
index = metric._labelnames.index("model_id")
|
||||
return {sample_key[index] for sample_key in metric._metrics}
|
||||
|
||||
|
||||
def _series_count(metric) -> int:
|
||||
return len(metric._metrics)
|
||||
|
||||
|
|
@ -194,6 +199,55 @@ async def test_sdk_router_originated_metrics_keep_labels_without_proxy_router():
|
|||
assert _requested_model_values(logger.litellm_deployment_failed_fallbacks) == {"sdk-fallback-group"}
|
||||
|
||||
|
||||
def test_deployment_failure_prefers_stamped_failed_deployment_id_over_mutated_metadata():
|
||||
logger = PrometheusLogger()
|
||||
exception = _ClientSideError("deployment-a exceeded its TPM limit")
|
||||
exception.failed_deployment_id = "deployment-a"
|
||||
|
||||
logger.set_llm_deployment_failure_metrics(
|
||||
request_kwargs={
|
||||
"model": "model-group",
|
||||
"litellm_params": {"metadata": {"model_info": {"id": "deployment-b"}}},
|
||||
"standard_logging_object": {"model_id": "deployment-b"},
|
||||
"exception": exception,
|
||||
}
|
||||
)
|
||||
|
||||
assert _model_id_values(logger.litellm_deployment_failure_responses) == {"deployment-a"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_failure_metrics_prefer_stamped_failed_deployment_id():
|
||||
logger = PrometheusLogger()
|
||||
exception = _ClientSideError("deployment-a exceeded its TPM limit")
|
||||
exception.failed_deployment_id = "deployment-a"
|
||||
|
||||
await logger.async_log_failure_event(
|
||||
kwargs={
|
||||
"model": "model-group",
|
||||
"litellm_params": {"metadata": {"model_info": {"id": "deployment-b"}}},
|
||||
"standard_logging_object": {
|
||||
"model_id": "deployment-b",
|
||||
"model_group": "model-group",
|
||||
"metadata": {
|
||||
"user_api_key_user_id": "user",
|
||||
"user_api_key_hash": "hash",
|
||||
"user_api_key_alias": "alias",
|
||||
"user_api_key_team_id": "team",
|
||||
"user_api_key_team_alias": "team-alias",
|
||||
},
|
||||
},
|
||||
"exception": exception,
|
||||
},
|
||||
response_obj=None,
|
||||
start_time=None,
|
||||
end_time=None,
|
||||
)
|
||||
|
||||
assert _model_id_values(logger.litellm_llm_api_failed_requests_metric) == {"deployment-a"}
|
||||
assert _model_id_values(logger.litellm_deployment_failure_responses) == {"deployment-a"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sdk_fallback_labels_survive_non_import_errors_from_proxy_module(monkeypatch):
|
||||
logger = PrometheusLogger()
|
||||
|
|
|
|||
|
|
@ -14460,6 +14460,8 @@ async def test_async_routing_strategy_pre_call_checks_failure_logging_is_coordin
|
|||
type(hook_error),
|
||||
)
|
||||
|
||||
assert hook_error.failed_deployment_id == deployment["model_info"]["id"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_callback_filter_deployments_failure_logging_is_coordinated():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue