diff --git a/litellm/integrations/clickhouse/clickhouse_spend_logger.py b/litellm/integrations/clickhouse/clickhouse_spend_logger.py new file mode 100644 index 00000000000..b55b074ef18 --- /dev/null +++ b/litellm/integrations/clickhouse/clickhouse_spend_logger.py @@ -0,0 +1,149 @@ +""" +`clickhouse` logging callback: one `spend_logs` row per LiteLLM request. + +Agent LLM spans join to these rows on `otel_traces.LiteLLMRequestId = spend_logs.response_id`, +so `response_id` is always the raw provider response id (cache-hit suffix stripped). +""" + +import json +import re +from collections.abc import Mapping +from typing import Any, Final + +import litellm +from litellm._logging import verbose_logger +from litellm.constants import MILLISECONDS_PER_SECOND +from litellm.integrations.clickhouse.clickhouse_batch_logger import ClickHouseBatchLogger +from litellm.integrations.clickhouse.schema import SPEND_LOGS_TABLE +from litellm.tracing.types import SpendLogRecord +from litellm.types.utils import StandardLoggingPayload + +# litellm_logging.py rewrites cache-hit ids as f"{id}_cache_hit{time.time()}" +_CACHE_HIT_SUFFIX: Final = re.compile(r"_cache_hit[0-9.]*$") +# W3C trace context: version-traceid-parentid-flags +_TRACEPARENT: Final = re.compile(r"^[0-9a-f]{2}-([0-9a-f]{32})-([0-9a-f]{16})-[0-9a-f]{2}$") +_INVALID_TRACE_ID: Final = "0" * 32 +_INVALID_SPAN_ID: Final = "0" * 16 + + +def strip_cache_hit_suffix(request_id: str) -> str: + return _CACHE_HIT_SUFFIX.sub("", request_id) + + +def parse_traceparent(value: object) -> tuple[str, str]: + """(trace_id, span_id) from a W3C `traceparent` header, or ("", "") if absent/invalid.""" + if not isinstance(value, str): + return "", "" + match = _TRACEPARENT.match(value.strip().lower()) + if match is None or match.group(1) == _INVALID_TRACE_ID or match.group(2) == _INVALID_SPAN_ID: + return "", "" + return match.group(1), match.group(2) + + +def _to_ms(seconds: object) -> int | None: + return int(float(seconds) * MILLISECONDS_PER_SECOND) if isinstance(seconds, (int, float)) else None + + +def _int(value: object) -> int: + return value if isinstance(value, int) and not isinstance(value, bool) else 0 + + +def _json(value: object) -> str: + if value is None or value == "": + return "" + return value if isinstance(value, str) else json.dumps(value, default=str) + + +def _find_traceparent(metadata: Mapping[str, Any], kwargs: Mapping[str, Any]) -> tuple[str, str]: + custom_headers = metadata.get("requester_custom_headers") or {} + proxy_request = (kwargs.get("litellm_params") or {}).get("proxy_server_request") or {} + request_headers = proxy_request.get("headers") or {} + for headers in (custom_headers, request_headers): + for name, value in headers.items(): + if str(name).lower() == "traceparent": + return parse_traceparent(value) + return "", "" + + +def _cache_tokens(usage: Mapping[str, Any]) -> tuple[int, int]: + """(cache_read, cache_write) from a Usage dict: OpenAI prompt_tokens_details first, Anthropic fields as fallback.""" + details = usage.get("prompt_tokens_details") or {} + cache_read = _int(details.get("cached_tokens")) or _int(usage.get("cache_read_input_tokens")) + cache_write = ( + _int(details.get("cache_write_tokens")) + or _int(details.get("cache_creation_tokens")) + or _int(usage.get("cache_creation_input_tokens")) + ) + return cache_read, cache_write + + +def _session_id(payload: StandardLoggingPayload, kwargs: Mapping[str, Any]) -> str: + """Mirrors proxy `_get_session_id_for_spend_log`: explicit session id, else the payload trace id.""" + request_metadata = (kwargs.get("litellm_params") or {}).get("metadata") or {} + return str(payload.get("session_id") or request_metadata.get("session_id") or payload.get("trace_id") or "") + + +def spend_log_row_from_payload(payload: StandardLoggingPayload, kwargs: Mapping[str, Any]) -> SpendLogRecord: + metadata: Mapping[str, Any] = payload.get("metadata") or {} + hidden_params: Mapping[str, Any] = payload.get("hidden_params") or {} + usage: Mapping[str, Any] = metadata.get("usage_object") or hidden_params.get("usage_object") or {} + cache_read_tokens, cache_write_tokens = _cache_tokens(usage) + trace_id, span_id = _find_traceparent(metadata, kwargs) + request_id = str(payload.get("id") or "") + redact = litellm.turn_off_message_logging is True + completion_start_ms = _to_ms(payload.get("completionStartTime")) + return SpendLogRecord( + request_id=request_id, + response_id=strip_cache_hit_suffix(request_id), + call_type=payload.get("call_type") or "", + api_key=metadata.get("user_api_key_hash") or "", + key_alias=metadata.get("user_api_key_alias") or "", + team_id=metadata.get("user_api_key_team_id") or metadata.get("team_id") or "", + team_alias=metadata.get("user_api_key_team_alias") or metadata.get("team_alias") or "", + organization_id=metadata.get("user_api_key_org_id") or "", + user=metadata.get("user_api_key_user_id") or "", + end_user=payload.get("end_user") or metadata.get("user_api_key_end_user_id") or "", + model=payload.get("model") or "", + model_group=payload.get("model_group") or "", + model_id=payload.get("model_id") or "", + custom_llm_provider=payload.get("custom_llm_provider") or "", + api_base=payload.get("api_base") or "", + spend=float(payload.get("response_cost") or 0.0), + prompt_tokens=_int(payload.get("prompt_tokens")), + completion_tokens=_int(payload.get("completion_tokens")), + total_tokens=_int(payload.get("total_tokens")), + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + start_time=_to_ms(payload.get("startTime")) or 0, + end_time=_to_ms(payload.get("endTime")) or 0, + completion_start_time=completion_start_ms or None, + status=payload.get("status") or "", + error_str=payload.get("error_str") or "", + cache_hit=payload.get("cache_hit") is True, + session_id=_session_id(payload, kwargs), + trace_id=trace_id, + span_id=span_id, + request_tags=[str(tag) for tag in payload.get("request_tags") or []], + metadata=_json(dict(metadata)), + messages="" if redact else _json(payload.get("messages")), + response="" if redact else _json(payload.get("response")), + ) + + +class ClickHouseSpendLogger(ClickHouseBatchLogger): + table = SPEND_LOGS_TABLE + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time) -> None: + self._log(kwargs) + + async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time) -> None: + self._log(kwargs) + + def _log(self, kwargs: Mapping[str, Any]) -> None: + try: + payload = kwargs.get("standard_logging_object") + if payload is None: + return + self.enqueue([dict(spend_log_row_from_payload(payload, kwargs))]) + except Exception as e: + verbose_logger.exception("ClickHouseSpendLogger: failed to log request: %s", e)