mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(clickhouse): add clickhouse spend-log callback
This commit is contained in:
parent
f512c161f0
commit
4d350c3787
1 changed files with 149 additions and 0 deletions
149
litellm/integrations/clickhouse/clickhouse_spend_logger.py
Normal file
149
litellm/integrations/clickhouse/clickhouse_spend_logger.py
Normal file
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue