feat(tracing): store spend in ClickHouse automatically (#43928)

This commit is contained in:
yujonglee 2026-09-30 15:17:43 -07:00 • committed by GitHub
parent b41715b0c2
commit 629c2b5808
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 883 additions and 103 deletions

View file

@ -0,0 +1,62 @@
name: litellm-tracing
services:
litellm:
build:
context: ..
target: runtime
command: ["--config", "/app/tracing-config.yaml", "--port", "4000"]
environment:
LITELLM_MASTER_KEY: local-tracing-master-key
LITELLM_SALT_KEY: sk-local-tracing-salt-key
DATABASE_URL: postgresql://litellm:litellm@db:5432/litellm
STORE_MODEL_IN_DB: "True"
CLICKHOUSE_URL: http://default:local-tracing@clickhouse:8123
CLICKHOUSE_READER_URL: http://default:local-tracing@clickhouse:8123
CLICKHOUSE_DATABASE: litellm
OPENAI_API_KEY: ${OPENAI_API_KEY:-}
volumes:
- ./tracing-config.yaml:/app/tracing-config.yaml:ro
ports:
- "127.0.0.1:4002:4000"
depends_on:
db:
condition: service_healthy
clickhouse:
condition: service_healthy
db:
image: postgres:16
environment:
POSTGRES_DB: litellm
POSTGRES_USER: litellm
POSTGRES_PASSWORD: litellm
volumes:
- postgres_data:/var/lib/postgresql/data
ports:
- "127.0.0.1:15432:5432"
healthcheck:
test: ["CMD-SHELL", "pg_isready -U litellm -d litellm"]
interval: 5s
timeout: 5s
retries: 10
clickhouse:
image: clickhouse/clickhouse-server:26.9.6.6
environment:
CLICKHOUSE_USER: default
CLICKHOUSE_PASSWORD: local-tracing
CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: "1"
volumes:
- clickhouse_data:/var/lib/clickhouse
ports:
- "127.0.0.1:18123:8123"
healthcheck:
test: ["CMD", "clickhouse-client", "--user", "default", "--password", "local-tracing", "--query", "SELECT 1"]
interval: 5s
timeout: 5s
retries: 20
volumes:
postgres_data:
clickhouse_data:

View file

@ -0,0 +1,10 @@
model_list:
- model_name: gpt-6.1-sol
litellm_params:
model: openai/gpt-6.1-sol
api_key: os.environ/OPENAI_API_KEY
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
tracing:
store: clickhouse

View file

@ -1,7 +1,7 @@
use std::collections::BTreeMap;
use litellm_http::ClientVariant;
use litellm_traces::{Connection, Error, InsertTable, Parameter};
use litellm_traces::{Connection, Error, InsertTable, Parameter, ReadQuery};
use pyo3::{
exceptions::{PyOverflowError, PyRuntimeError, PyValueError},
prelude::*,
@ -9,9 +9,11 @@ use pyo3::{
fn map_error(error: Error) -> PyErr {
match error {
Error::InvalidRow | Error::InvalidTable | Error::InvalidSchema | Error::EmptySql => {
PyValueError::new_err(error.to_string())
}
Error::InvalidRow
| Error::InvalidTable
| Error::InvalidSchema
| Error::EmptySql
| Error::InvalidQuery => PyValueError::new_err(error.to_string()),
Error::InsertTooLarge => PyOverflowError::new_err(error.to_string()),
Error::InvalidUrl
| Error::QueryFailed(_)
@ -94,19 +96,22 @@ impl NativeTraceStorage {
fn query<'py>(
&self,
py: Python<'py>,
sql: String,
query: &str,
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] parameters: BTreeMap<
String,
Parameter,
>,
) -> PyResult<Bound<'py, PyAny>> {
let query = ReadQuery::parse(query).map_err(map_error)?;
let connection = self.reader.clone().ok_or_else(|| {
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
})?;
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
crate::execution::run_async(
py,
async move { litellm_traces::execute_read(&client, &connection, &sql, &parameters).await },
async move {
litellm_traces::execute_named_read(&client, &connection, query, &parameters).await
},
map_error,
)
}

View file

@ -0,0 +1,23 @@
SELECT TraceId AS trace_id,
hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref,
TeamId AS team_id, ApiKeyHash AS api_key_hash,
ifNull(any(RootName), '') AS name, any(ServiceName) AS service,
ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status,
toUnixTimestamp64Milli(min(StartTs)) AS start_ms,
dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms,
sum(SpanCount) AS span_count, length(groupUniqArrayArray(AgentNames)) AS agent_count,
sum(AgentCount) AS agent_invocations,
sum(LlmCount) AS llm_calls, sum(ToolCount) AS tool_calls,
sum(InputTokens) AS input_tokens, sum(OutputTokens) AS output_tokens,
groupUniqArrayArray(Models) AS models, sum(ErrorCount) AS error_count,
arrayDistinct(groupArrayArray(RequestIds)) AS request_ids
FROM agent_traces_by_key
WHERE (empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)})
AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String})
GROUP BY TeamId, ApiKeyHash, TraceId
HAVING min(StartTs) >= fromUnixTimestamp64Milli({start_ms:Int64})
AND min(StartTs) < fromUnixTimestamp64Milli({end_ms:Int64})
AND ({cursor_ms:Int64} = 0 OR (toUnixTimestamp64Milli(min(StartTs)), trace_ref)
< ({cursor_ms:Int64}, {cursor_trace_id:String}))
ORDER BY start_ms DESC, trace_ref DESC
LIMIT {limit:UInt32}

View file

@ -0,0 +1,8 @@
SELECT SpanId AS span_id, Input AS input, Output AS output, SpanAttributes AS attributes
FROM otel_traces
WHERE TraceId = {trace_id:String} AND SpanId = {span_id:String}
AND (empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)})
AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String})
AND ({trace_ref:String} = '' OR
hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String})
LIMIT 1

View file

@ -0,0 +1,9 @@
SELECT request_id, response_id, team_id, api_key, spend,
toUnixTimestamp64Milli(start_time) AS start_ms
FROM spend_logs FINAL
WHERE response_id IN {response_ids:Array(String)}
AND start_time >= fromUnixTimestamp64Milli({start_ms:Int64})
AND start_time < fromUnixTimestamp64Milli({end_ms:Int64})
AND (empty({team_ids:Array(String)}) OR team_id IN {team_ids:Array(String)})
AND ({api_key_hash:String} = '' OR api_key = {api_key_hash:String})
ORDER BY start_time DESC

View file

@ -0,0 +1,16 @@
SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name,
o.ObservationType AS type, o.AgentName AS agent, o.StatusCode AS status,
o.StatusMessage AS status_message,
toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns,
o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model,
o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens,
o.LiteLLMRequestId AS litellm_request_id,
o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash
FROM otel_traces AS o
WHERE o.TraceId = {trace_id:String}
AND (empty({team_ids:Array(String)}) OR o.TeamId IN {team_ids:Array(String)})
AND ({api_key_hash:String} = '' OR o.ApiKeyHash = {api_key_hash:String})
AND ({trace_ref:String} = '' OR
hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String})
ORDER BY o.Timestamp
LIMIT 1 BY o.SpanId

View file

@ -10,6 +10,8 @@ pub enum Error {
InvalidSchema,
#[error("SQL query must not be empty")]
EmptySql,
#[error("unknown ClickHouse read query")]
InvalidQuery,
#[error("ClickHouse query failed with HTTP status {0}")]
QueryFailed(u16),
#[error("ClickHouse insert failed with HTTP status {0}")]

View file

@ -49,7 +49,24 @@ pub async fn insert_rows(
.map_err(|_| Error::InvalidRow)?;
let body = encoder.finish().map_err(|_| Error::InvalidRow)?;
let mut url = connection.url().clone();
let existing_pairs: Vec<(String, String)> = url
.query_pairs()
.filter(|(key, _)| {
!matches!(
key.as_ref(),
"query"
| "async_insert"
| "async_insert_deduplicate"
| "wait_for_async_insert"
| "input_format_skip_unknown_fields"
| "date_time_input_format"
)
})
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
url.query_pairs_mut()
.clear()
.extend_pairs(existing_pairs)
.append_pair(
"query",
&format!(
@ -60,6 +77,7 @@ pub async fn insert_rows(
.append_pair("async_insert", "1")
.append_pair("async_insert_deduplicate", "1")
.append_pair("wait_for_async_insert", "1")
.append_pair("input_format_skip_unknown_fields", "0")
.append_pair("date_time_input_format", "best_effort");
let response = client
.post(url)

View file

@ -8,7 +8,7 @@ pub use error::{DecodeError, Error};
pub use insert::{InsertTable, encode_rows, insert_rows};
pub use otlp::{DecodedSpan, decode_otlp};
pub use schema::{ensure_schema, schema_statements};
pub use sql::{Parameter, execute_read};
pub use sql::{Parameter, ReadQuery, execute_named_read, execute_read};
use url::Url;
#[derive(Clone)]

View file

@ -8,6 +8,34 @@ use crate::{Connection, Error};
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
pub enum ReadQuery {
ListTraces,
TraceSpans,
SpanDetail,
SpendByResponseIds,
}
impl ReadQuery {
pub fn parse(value: &str) -> Result<Self, Error> {
match value {
"list_traces" => Ok(Self::ListTraces),
"trace_spans" => Ok(Self::TraceSpans),
"span_detail" => Ok(Self::SpanDetail),
"spend_by_response_ids" => Ok(Self::SpendByResponseIds),
_ => Err(Error::InvalidQuery),
}
}
fn sql(&self) -> &'static str {
match self {
Self::ListTraces => include_str!("../query/list_traces.sql"),
Self::TraceSpans => include_str!("../query/trace_spans.sql"),
Self::SpanDetail => include_str!("../query/span_detail.sql"),
Self::SpendByResponseIds => include_str!("../query/spend_by_response_ids.sql"),
}
}
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum Parameter {
@ -112,3 +140,12 @@ pub async fn execute_read(
}
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
}
pub async fn execute_named_read(
client: &Client,
connection: &Connection,
query: ReadQuery,
parameters: &BTreeMap<String, Parameter>,
) -> Result<String, Error> {
execute_read(client, connection, query.sql(), parameters).await
}

View file

@ -2,7 +2,8 @@ use std::{collections::BTreeMap, time::Duration};
use litellm_http::Client;
use litellm_traces::{
Connection, Error, InsertTable, encode_rows, ensure_schema, execute_read, schema_statements,
Connection, Error, InsertTable, Parameter, ReadQuery, encode_rows, ensure_schema,
execute_named_read, execute_read, schema_statements,
};
use rstest::{fixture, rstest};
use testcontainers_modules::{
@ -124,6 +125,61 @@ async fn schema_supports_span_rollups_and_spend_joins(
}))?;
insert_rows(&database, "otel_traces", vec![span]).await?;
insert_rows(&database, "spend_logs", vec![spend]).await?;
let reader = Connection::reader(&database.url, "trace_test")?;
let list_parameters = BTreeMap::from([
("team_ids".into(), Parameter::Strings(vec!["team-1".into()])),
("api_key_hash".into(), Parameter::Text(String::new())),
(
"start_ms".into(),
Parameter::Integer(timestamp / 1_000_000 - 1000),
),
(
"end_ms".into(),
Parameter::Integer(timestamp / 1_000_000 + 1000),
),
("cursor_ms".into(), Parameter::Integer(0)),
("cursor_trace_id".into(), Parameter::Text(String::new())),
("limit".into(), Parameter::Integer(10)),
]);
let listed: serde_json::Value = serde_json::from_str(
&execute_named_read(
&database.client,
&reader,
ReadQuery::ListTraces,
&list_parameters,
)
.await?,
)?;
assert_eq!(
listed["data"][0]["request_ids"],
serde_json::json!(["response-1"])
);
let spend_parameters = BTreeMap::from([
(
"response_ids".into(),
Parameter::Strings(vec!["response-1".into()]),
),
("team_ids".into(), Parameter::Strings(vec!["team-1".into()])),
("api_key_hash".into(), Parameter::Text(String::new())),
(
"start_ms".into(),
Parameter::Integer(timestamp / 1_000_000 - 1000),
),
(
"end_ms".into(),
Parameter::Integer(timestamp / 1_000_000 + 1000),
),
]);
let matched: serde_json::Value = serde_json::from_str(
&execute_named_read(
&database.client,
&reader,
ReadQuery::SpendByResponseIds,
&spend_parameters,
)
.await?,
)?;
assert_eq!(matched["data"][0]["spend"], 0.125);
let body = read_json(
&database,
"SELECT o.TeamId, o.ApiKeyHash, o.ObservationType, o.InputPreview, s.spend, \
@ -154,6 +210,43 @@ async fn schema_supports_span_rollups_and_spend_joins(
Ok(())
}
#[rstest]
#[tokio::test]
async fn insert_rejects_unknown_columns_even_if_url_requests_skipping_them(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
) -> TestResult {
let database = database?;
let writer = Connection::writer(&format!(
"{}?input_format_skip_unknown_fields=1",
database.url
))?;
ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?;
let row = BTreeMap::from([
(
"Timestamp".to_owned(),
serde_json::json!(1_700_000_000_000_000_000_i64),
),
(
"unexpected".to_owned(),
serde_json::json!("dropped silently"),
),
]);
assert!(matches!(
litellm_traces::insert_rows(
&database.client,
&writer,
"trace_test",
InsertTable::OtelTraces,
vec![row]
)
.await,
Err(Error::InsertFailed(_))
));
assert_eq!(table_rows(&database, "otel_traces").await?, 0);
Ok(())
}
#[rstest]
#[tokio::test]
async fn retried_trace_insert_does_not_inflate_rollup(

View file

@ -5,11 +5,12 @@ Built on `CustomBatchLogger`: rows accumulate in `log_queue` and are flushed as
gzip JSONEachRow insert, either every `CLICKHOUSE_FLUSH_INTERVAL_SECONDS` or as soon as
`batch_size` rows are queued. Subclasses only pick the table and build rows:
- `ClickHouseSpendLogger` -> spend_logs (LiteLLM requests, via the `clickhouse` callback)
- `ClickHouseSpendLogger` -> spend_logs (LiteLLM requests when tracing is enabled)
"""
import asyncio
import os
from collections.abc import Mapping, Sequence
from typing import Any, ClassVar
from litellm._logging import verbose_logger
@ -43,20 +44,19 @@ class ClickHouseBatchLogger(CustomBatchLogger):
batch_size=CLICKHOUSE_BATCH_SIZE,
flush_interval=CLICKHOUSE_FLUSH_INTERVAL_SECONDS,
)
try:
asyncio.get_running_loop().create_task(self.periodic_flush())
except RuntimeError: # no loop yet (e.g. sync config load); proxy startup calls start()
pass
self._flush_task: asyncio.Task[None] | None = None
def start(self) -> None:
asyncio.get_running_loop().create_task(self.periodic_flush())
if self._flush_task is None or self._flush_task.done():
self._flush_task = asyncio.get_running_loop().create_task(self.periodic_flush())
def is_full(self) -> bool:
"""Backpressure signal: producers should reject (429) instead of enqueueing."""
return len(self.log_queue) >= CLICKHOUSE_MAX_BUFFERED_ROWS
def enqueue(self, rows: list[dict[str, Any]]) -> None:
def enqueue(self, rows: Sequence[Mapping[str, object]]) -> None:
"""Never awaits ClickHouse. Kicks off an early flush once a full batch is queued."""
self.start()
self.log_queue.extend(rows)
if len(self.log_queue) >= self.batch_size:
asyncio.get_running_loop().create_task(self.flush_queue())

View file

@ -0,0 +1,85 @@
import re
from collections.abc import Mapping
from datetime import datetime
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict, ValidationError
from litellm._logging import verbose_logger
from litellm.integrations.clickhouse.clickhouse_batch_logger import ClickHouseBatchLogger
from litellm.integrations.clickhouse.schema import SPEND_LOGS_TABLE
from litellm.rust_bridge.traces import TraceStorage
_CACHE_HIT_SUFFIX: Final = re.compile(r"_cache_hit[0-9.]+$")
class _SpendMetadata(BaseModel):
model_config = ConfigDict(frozen=True)
user_api_key_hash: str | None = None
user_api_key_team_id: str | None = None
class _SpendPayload(BaseModel):
model_config = ConfigDict(frozen=True)
id: str
call_type: str = ""
response_cost: float | None = None
prompt_tokens: int = 0
completion_tokens: int = 0
total_tokens: int = 0
startTime: float
endTime: float
metadata: _SpendMetadata = _SpendMetadata()
model: str | None = None
status: str = ""
cache_hit: bool | None = None
def spend_log_row_from_payload(payload: _SpendPayload) -> Mapping[str, object]:
return MappingProxyType(
{
"request_id": payload.id,
"response_id": _CACHE_HIT_SUFFIX.sub("", payload.id),
"call_type": payload.call_type,
"api_key": payload.metadata.user_api_key_hash or "",
"team_id": payload.metadata.user_api_key_team_id or "",
"model": payload.model or "",
"spend": payload.response_cost or 0.0,
"prompt_tokens": payload.prompt_tokens,
"completion_tokens": payload.completion_tokens,
"total_tokens": payload.total_tokens,
"start_time": int(payload.startTime * 1000),
"end_time": int(payload.endTime * 1000),
"status": payload.status,
"cache_hit": payload.cache_hit is True,
}
)
class ClickHouseSpendLogger(ClickHouseBatchLogger):
table = SPEND_LOGS_TABLE
def __init__(self, storage: TraceStorage) -> None:
super().__init__(storage=storage)
async def _log(self, kwargs: Mapping[str, object]) -> None:
try:
payload: Final = _SpendPayload.model_validate(kwargs.get("standard_logging_object"))
if payload.call_type.startswith("/v1/traces"):
return
self.enqueue((spend_log_row_from_payload(payload),))
except (ValidationError, RuntimeError, ValueError) as error:
verbose_logger.warning("ClickHouse spend logging failed: %s", error)
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
await self._log(kwargs)
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
) -> None:
await self._log(kwargs)

View file

@ -11321,25 +11321,36 @@ class ProxyStartupEvent:
return connected_client
@classmethod
async def init_tracing(cls, general_settings: dict) -> None:
async def init_tracing(cls, general_settings: dict, receiver: TraceReceiver | None = None) -> None:
"""
Enable agent tracing (`POST/GET /v1/traces`) when configured:
general_settings:
tracing:
store: clickhouse # CLICKHOUSE_URL / _USER / _PASSWORD / _DATABASE
store: clickhouse
"""
from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger
manager: Final = litellm.logging_callback_manager
for callback in manager.get_custom_loggers_for_type(ClickHouseSpendLogger):
manager.remove_callback_from_all_lists(callback)
tracing_endpoints.receiver = None
settings: Final = general_settings.get("tracing")
if not isinstance(settings, dict) or settings.get("store") != "clickhouse":
return
try:
tracing: Final = TraceReceiver.from_env()
tracing: Final = receiver if receiver is not None else TraceReceiver.from_env()
await tracing.start()
except (KeyError, OSError, RuntimeError, ValueError) as error:
tracing_endpoints.receiver = None
verbose_proxy_logger.warning("Agent tracing unavailable: %s", error)
return
tracing_endpoints.receiver = tracing
spend_logger: Final = ClickHouseSpendLogger(storage=tracing.store.storage)
manager.add_litellm_callback(spend_logger)
manager.add_litellm_success_callback(spend_logger)
manager.add_litellm_failure_callback(spend_logger)
manager.add_litellm_async_success_callback(spend_logger)
manager.add_litellm_async_failure_callback(spend_logger)
verbose_proxy_logger.info("Agent tracing enabled (store=clickhouse)")
@classmethod

View file

@ -1,6 +1,6 @@
from collections.abc import Awaitable, Mapping, Sequence
from types import MappingProxyType
from typing import Final, Protocol, TypedDict, cast
from typing import Final, Literal, Protocol, TypedDict, cast
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
from typing_extensions import ReadOnly
@ -31,6 +31,9 @@ class DecodedSpan(TypedDict):
events: ReadOnly[list[DecodedEvent]]
ReadQueryName = Literal["list_traces", "trace_spans", "span_detail", "spend_by_response_ids"]
class NativeStore(Protocol):
def __init__(self, database: str, url: str, reader_url: str | None = None) -> None: ...
@ -38,7 +41,7 @@ class NativeStore(Protocol):
def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Awaitable[None]: ...
def query(self, sql: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Awaitable[str]: ...
def query(self, name: ReadQueryName, parameters: Mapping[str, str | int | Sequence[str]]) -> Awaitable[str]: ...
class NativeTraces(Protocol):
@ -85,8 +88,10 @@ class TraceStorage:
async def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> None:
await self._native.insert_rows(table, INSERT_ROWS.validate_python(rows))
async def query(self, sql: str, parameters: Mapping[str, object] | None = None) -> list[dict[str, JsonValue]]:
async def query(
self, name: ReadQueryName, parameters: Mapping[str, object] | None = None
) -> list[dict[str, JsonValue]]:
result: Final = await self._native.query(
sql, QUERY_PARAMETERS.validate_python(parameters or MappingProxyType({}))
name, QUERY_PARAMETERS.validate_python(parameters or MappingProxyType({}))
)
return QueryResponse.model_validate_json(result).data

View file

@ -5,12 +5,15 @@ import binascii
import json
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from itertools import chain
from types import MappingProxyType
from typing import Any, Final
from pydantic import BaseModel, ConfigDict, TypeAdapter
from litellm._logging import verbose_logger
from litellm.constants import AGENT_TRACING_LIST_PAGE_SIZE
from litellm.integrations.clickhouse.schema import (
AGENT_TRACES_BY_KEY_TABLE,
OTEL_TRACES_TABLE,
)
from litellm.rust_bridge.traces import TraceStorage
@ -27,58 +30,39 @@ from litellm.tracing.types import (
)
NANOS_PER_MS: Final = 1_000_000
SPEND_WINDOW_MS: Final = 30 * 60 * 1000
_STATUS: Final = MappingProxyType({"STATUS_CODE_OK": "ok", "STATUS_CODE_ERROR": "error"})
_SCOPE_OTEL: Final = (
"(empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)})"
" AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String})"
)
_TRACE_REF_SQL: Final = "hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))"
LIST_TRACES_SQL: Final = f"""
SELECT TraceId AS trace_id, {_TRACE_REF_SQL} AS trace_ref,
ifNull(any(RootName), '') AS name, any(ServiceName) AS service,
ifNull(any(RootInput), '') AS input_preview, ifNull(any(RootStatus), '') AS status,
toUnixTimestamp64Milli(min(StartTs)) AS start_ms,
dateDiff('millisecond', min(StartTs), max(EndTs)) AS duration_ms,
sum(SpanCount) AS span_count, length(groupUniqArrayArray(AgentNames)) AS agent_count,
sum(AgentCount) AS agent_invocations,
sum(LlmCount) AS llm_calls, sum(ToolCount) AS tool_calls,
sum(InputTokens) AS input_tokens, sum(OutputTokens) AS output_tokens,
groupUniqArrayArray(Models) AS models, sum(ErrorCount) AS error_count
FROM {AGENT_TRACES_BY_KEY_TABLE}
WHERE (empty({{team_ids:Array(String)}}) OR TeamId IN {{team_ids:Array(String)}})
AND ({{api_key_hash:String}} = '' OR ApiKeyHash = {{api_key_hash:String}})
GROUP BY TeamId, ApiKeyHash, TraceId
HAVING min(StartTs) >= fromUnixTimestamp64Milli({{start_ms:Int64}})
AND min(StartTs) < fromUnixTimestamp64Milli({{end_ms:Int64}})
AND ({{cursor_ms:Int64}} = 0 OR (toUnixTimestamp64Milli(min(StartTs)), trace_ref)
< ({{cursor_ms:Int64}}, {{cursor_trace_id:String}}))
ORDER BY start_ms DESC, trace_ref DESC
LIMIT {{limit:UInt32}}
"""
TRACE_SPANS_SQL: Final = f"""
SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name,
o.ObservationType AS type, o.AgentName AS agent, o.StatusCode AS status,
o.StatusMessage AS status_message,
toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns,
o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model,
o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens,
o.LiteLLMRequestId AS litellm_request_id
FROM {OTEL_TRACES_TABLE} AS o
WHERE o.TraceId = {{trace_id:String}} AND {_SCOPE_OTEL}
AND ({{trace_ref:String}} = '' OR {_TRACE_REF_SQL} = {{trace_ref:String}})
ORDER BY o.Timestamp
LIMIT 1 BY o.SpanId
"""
class _SpendRow(BaseModel):
model_config = ConfigDict(frozen=True)
SPAN_DETAIL_SQL: Final = f"""
SELECT SpanId AS span_id, Input AS input, Output AS output, SpanAttributes AS attributes
FROM {OTEL_TRACES_TABLE}
WHERE TraceId = {{trace_id:String}} AND SpanId = {{span_id:String}} AND {_SCOPE_OTEL}
AND ({{trace_ref:String}} = '' OR {_TRACE_REF_SQL} = {{trace_ref:String}})
LIMIT 1
"""
request_id: str
response_id: str
team_id: str
api_key: str
spend: float
start_ms: int
_SPEND_ROWS: Final = TypeAdapter(tuple[_SpendRow, ...])
def _spend_for(request_id: str, team_id: str, api_key_hash: str, rows: Sequence[_SpendRow]) -> float | None:
matches: Final = tuple(
row for row in rows if row.response_id == request_id and row.team_id == team_id and row.api_key == api_key_hash
)
return matches[0].spend if len(matches) == 1 else None
def _trace_spend(
request_ids: Sequence[str], team_id: str, api_key_hash: str, rows: Sequence[_SpendRow]
) -> float | None:
ids: Final = frozenset(request_id for request_id in request_ids if request_id)
costs: Final = tuple(_spend_for(request_id, team_id, api_key_hash, rows) for request_id in ids)
return (
sum(cost for cost in costs if cost is not None) if costs and all(cost is not None for cost in costs) else None
)
def encode_cursor(start_ms: int, trace_id: str) -> str:
@ -113,7 +97,7 @@ def _status(code: str) -> SpanStatus:
return _STATUS.get(code, "unset")
def trace_summary_from_row(row: dict[str, Any]) -> TraceSummary:
def trace_summary_from_row(row: dict[str, Any], spend_rows: Sequence[_SpendRow] = ()) -> TraceSummary:
return TraceSummary(
trace_id=row["trace_id"],
trace_ref=row.get("trace_ref", ""),
@ -132,10 +116,13 @@ def trace_summary_from_row(row: dict[str, Any]) -> TraceSummary:
input_tokens=int(row["input_tokens"]),
output_tokens=int(row["output_tokens"]),
models=tuple(row["models"]),
spend=_trace_spend(
row.get("request_ids") or (), row.get("team_id") or "", row.get("api_key_hash") or "", spend_rows
),
)
def span_from_row(row: dict[str, Any], trace_start_ns: int) -> Span:
def span_from_row(row: dict[str, Any], trace_start_ns: int, spend_rows: Sequence[_SpendRow] = ()) -> Span:
return Span(
span_id=row["span_id"],
parent_span_id=row["parent_span_id"] or None,
@ -151,6 +138,11 @@ def span_from_row(row: dict[str, Any], trace_start_ns: int) -> Span:
input_tokens=int(row["input_tokens"]),
output_tokens=int(row["output_tokens"]),
litellm_request_id=row["litellm_request_id"] or None,
spend=(
_spend_for(row["litellm_request_id"], row.get("team_id") or "", row.get("api_key_hash") or "", spend_rows)
if row["litellm_request_id"]
else None
),
)
@ -182,6 +174,7 @@ def agent_nodes(spans: Sequence[Span]) -> tuple[AgentNode, ...]:
llm_calls=0,
tool_calls=0,
duration_ms=0.0,
spend=None,
),
)
node["invocations"] += 1
@ -194,15 +187,43 @@ def agent_nodes(spans: Sequence[Span]) -> tuple[AgentNode, ...]:
owner["llm_calls"] += 1
elif span["type"] == "tool":
owner["tool_calls"] += 1
return tuple(agents.values())
return tuple(
AgentNode(
name=agent["name"],
parent_agent=agent["parent_agent"],
invocations=agent["invocations"],
llm_calls=agent["llm_calls"],
tool_calls=agent["tool_calls"],
duration_ms=agent["duration_ms"],
spend=_agent_spend(spans, agent["name"]),
)
for agent in agents.values()
)
def trace_from_rows(trace_id: str, rows: list[dict[str, Any]], trace_ref: str = "") -> Trace | None:
def _agent_spend(spans: Sequence[Span], agent_name: str) -> float | None:
by_request: Final = MappingProxyType(
{
span["litellm_request_id"]: span["spend"]
for span in spans
if span["type"] == "llm" and span["agent"] == agent_name and span["litellm_request_id"]
}
)
return (
sum(cost for cost in by_request.values() if cost is not None)
if by_request and all(cost is not None for cost in by_request.values())
else None
)
def trace_from_rows(
trace_id: str, rows: list[dict[str, Any]], trace_ref: str = "", spend_rows: Sequence[_SpendRow] = ()
) -> Trace | None:
if not rows:
return None
trace_start_ns: Final = min(int(r["start_ns"]) for r in rows)
trace_end_ns: Final = max(int(r["start_ns"]) + int(r["duration_ns"]) for r in rows)
spans: Final = tuple(span_from_row(r, trace_start_ns) for r in rows)
spans: Final = tuple(span_from_row(r, trace_start_ns, spend_rows) for r in rows)
root: Final = next((s for s in spans if s["parent_span_id"] is None), spans[0])
agents: Final = agent_nodes(spans)
llm_spans: Final = tuple(s for s in spans if s["type"] == "llm")
@ -225,6 +246,12 @@ def trace_from_rows(trace_id: str, rows: list[dict[str, Any]], trace_ref: str =
input_tokens=sum(s["input_tokens"] for s in spans),
output_tokens=sum(s["output_tokens"] for s in spans),
models=tuple(sorted(frozenset(s["model"] for s in llm_spans if s["model"]))),
spend=_trace_spend(
tuple(row["litellm_request_id"] for row in rows),
rows[0].get("team_id") or "",
rows[0].get("api_key_hash") or "",
spend_rows,
),
),
agents=agents,
spans=spans,
@ -240,6 +267,29 @@ class ClickHouseTraceStore:
async def insert_spans(self, rows: Sequence[SpanRow]) -> None:
await self.storage.insert_rows(OTEL_TRACES_TABLE, tuple(rows))
async def _spend_rows(
self, scope: TraceScope, request_ids: Sequence[str], start_ms: int, end_ms: int
) -> tuple[_SpendRow, ...]:
ids: Final = tuple(sorted(frozenset(request_id for request_id in request_ids if request_id)))
if not ids:
return ()
try:
rows: Final = await self.storage.query(
"spend_by_response_ids",
MappingProxyType(
{
**scope,
"response_ids": ids,
"start_ms": start_ms - SPEND_WINDOW_MS,
"end_ms": end_ms + SPEND_WINDOW_MS,
}
),
)
except RuntimeError as error:
verbose_logger.warning("Trace spend lookup unavailable: %s", error)
return ()
return _SPEND_ROWS.validate_python(rows)
async def list_traces(
self,
scope: TraceScope,
@ -250,7 +300,7 @@ class ClickHouseTraceStore:
) -> TracePage:
cursor_ms, cursor_trace_id = decode_cursor(cursor)
rows = await self.storage.query(
LIST_TRACES_SQL,
"list_traces",
MappingProxyType(
{
**scope,
@ -262,18 +312,30 @@ class ClickHouseTraceStore:
}
),
)
spend_rows: Final = await self._spend_rows(
scope,
tuple(chain.from_iterable(row.get("request_ids") or () for row in rows)),
min((int(row["start_ms"]) for row in rows), default=start_ms),
max((int(row["start_ms"]) + int(row["duration_ms"]) for row in rows), default=end_ms),
)
next_cursor = encode_cursor(int(rows[-1]["start_ms"]), rows[-1]["trace_ref"]) if len(rows) == limit else None
return TracePage(data=tuple(trace_summary_from_row(r) for r in rows), next_cursor=next_cursor)
return TracePage(data=tuple(trace_summary_from_row(r, spend_rows) for r in rows), next_cursor=next_cursor)
async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None:
rows = await self.storage.query(
TRACE_SPANS_SQL, MappingProxyType({**scope, "trace_id": trace_id, "trace_ref": trace_ref})
"trace_spans", MappingProxyType({**scope, "trace_id": trace_id, "trace_ref": trace_ref})
)
return trace_from_rows(trace_id, rows, trace_ref)
spend_rows: Final = await self._spend_rows(
scope,
tuple(row["litellm_request_id"] for row in rows),
min((int(row["start_ns"]) // NANOS_PER_MS for row in rows), default=0),
max(((int(row["start_ns"]) + int(row["duration_ns"])) // NANOS_PER_MS for row in rows), default=0),
)
return trace_from_rows(trace_id, rows, trace_ref, spend_rows)
async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None:
rows = await self.storage.query(
SPAN_DETAIL_SQL,
"span_detail",
MappingProxyType({**scope, "trace_id": trace_id, "span_id": span_id, "trace_ref": trace_ref}),
)
if not rows:

View file

@ -32,6 +32,7 @@ class Span(TypedDict):
input_tokens: ReadOnly[int]
output_tokens: ReadOnly[int]
litellm_request_id: ReadOnly[str | None]
spend: ReadOnly[float | None]
class AgentNode(TypedDict):
@ -43,6 +44,7 @@ class AgentNode(TypedDict):
llm_calls: int
tool_calls: int
duration_ms: float
spend: ReadOnly[float | None]
class TraceSummary(TypedDict):
@ -63,6 +65,7 @@ class TraceSummary(TypedDict):
input_tokens: ReadOnly[int]
output_tokens: ReadOnly[int]
models: ReadOnly[tuple[str, ...]]
spend: ReadOnly[float | None]
class Trace(TypedDict):

View file

@ -0,0 +1,34 @@
#!/usr/bin/env bash
set -euo pipefail
repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
cd "$repo_root"
docker compose -f docker/docker-compose.tracing.yml up -d --wait db clickhouse
uv sync --inexact --frozen --extra proxy --group proxy-dev --no-install-project
"$repo_root/.venv/bin/python" scripts/prisma_generate_if_needed.py
VIRTUAL_ENV="$repo_root/.venv" uvx --from maturin==1.15.0 maturin develop \
--release --manifest-path litellm-rust/crates/python-bridge/Cargo.toml --features extension-module
config_file="$(mktemp "${TMPDIR:-/tmp}/litellm-tracing-local.XXXXXX.yaml")"
trap 'rm -f "$config_file"' EXIT
cat > "$config_file" <<'EOF'
model_list: []
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
tracing:
store: clickhouse
EOF
export LITELLM_MASTER_KEY=sk-local-tracing
export LITELLM_SALT_KEY=sk-local-tracing-salt-key
export DATABASE_URL=postgresql://litellm:litellm@127.0.0.1:15432/litellm
export STORE_MODEL_IN_DB=True
export CLICKHOUSE_URL=http://default:local-tracing@127.0.0.1:18123
export CLICKHOUSE_READER_URL="$CLICKHOUSE_URL"
export CLICKHOUSE_DATABASE=litellm
export LITELLM_LOCAL_MODEL_COST_MAP=True
printf 'Proxy: http://127.0.0.1:4002/ui\nMaster key: %s\n' "$LITELLM_MASTER_KEY"
"$repo_root/.venv/bin/python" litellm/proxy/proxy_cli.py \
--config "$config_file" --host 127.0.0.1 --port 4002

View file

@ -2,13 +2,13 @@
Tests for the CustomBatchLogger-based ClickHouse base logger.
"""
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.integrations.clickhouse import clickhouse_batch_logger as module
from litellm.integrations.clickhouse.clickhouse_batch_logger import ClickHouseBatchLogger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
class _TestLogger(ClickHouseBatchLogger):
@ -21,10 +21,6 @@ def _logger(insert: AsyncMock) -> _TestLogger:
return _TestLogger(storage=storage)
def test_is_a_custom_batch_logger():
assert issubclass(ClickHouseBatchLogger, CustomBatchLogger)
@pytest.mark.asyncio
async def test_flush_splits_into_batches_and_empties_queue():
insert = AsyncMock()
@ -40,6 +36,24 @@ async def test_flush_splits_into_batches_and_empties_queue():
assert logger.rows_written == 5
@pytest.mark.asyncio
async def test_first_enqueued_row_flushes_after_synchronous_construction():
flushed = asyncio.Event()
async def insert_rows(table: str, rows: list[dict[str, int]]) -> None:
assert table == "test_table"
assert rows == [{"i": 1}]
flushed.set()
logger = _logger(AsyncMock(side_effect=insert_rows))
logger.flush_interval = 0.01
logger.enqueue([{"i": 1}])
await asyncio.wait_for(flushed.wait(), timeout=1)
if logger._flush_task is not None:
logger._flush_task.cancel()
@pytest.mark.asyncio
async def test_is_full_signals_backpressure():
logger = _logger(AsyncMock())

View file

@ -0,0 +1,102 @@
from datetime import datetime, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger
def _payload(request_id: str, *, status: str, cost: float) -> dict[str, object]:
return {
"id": request_id,
"call_type": "acompletion",
"response_cost": cost,
"prompt_tokens": 7,
"completion_tokens": 3,
"total_tokens": 10,
"startTime": 1_700_000_000.123,
"endTime": 1_700_000_001.456,
"metadata": {"user_api_key_hash": "key-a", "user_api_key_team_id": "team-a"},
"model": "test-model",
"status": status,
}
@pytest.mark.asyncio
async def test_success_and_failure_events_write_scoped_spend_rows():
storage = MagicMock()
storage.ensure_schema = AsyncMock()
storage.insert_rows = AsyncMock()
logger = ClickHouseSpendLogger(storage=storage)
now = datetime.now(timezone.utc)
await logger.async_log_success_event(
{"standard_logging_object": _payload("response-1", status="success", cost=0.25)}, None, now, now
)
await logger.async_log_failure_event(
{"standard_logging_object": _payload("response-2_cache_hit123", status="failure", cost=0.0)},
None,
now,
now,
)
await logger.flush_queue()
if logger._flush_task is not None:
logger._flush_task.cancel()
storage.ensure_schema.assert_not_awaited()
assert storage.insert_rows.await_count == 1
table, rows = storage.insert_rows.await_args.args
assert table == "spend_logs"
assert rows == [
{
"request_id": "response-1",
"response_id": "response-1",
"call_type": "acompletion",
"api_key": "key-a",
"team_id": "team-a",
"model": "test-model",
"spend": 0.25,
"prompt_tokens": 7,
"completion_tokens": 3,
"total_tokens": 10,
"start_time": 1_700_000_000_123,
"end_time": 1_700_000_001_456,
"status": "success",
"cache_hit": False,
},
{
"request_id": "response-2_cache_hit123",
"response_id": "response-2",
"call_type": "acompletion",
"api_key": "key-a",
"team_id": "team-a",
"model": "test-model",
"spend": 0.0,
"prompt_tokens": 7,
"completion_tokens": 3,
"total_tokens": 10,
"start_time": 1_700_000_000_123,
"end_time": 1_700_000_001_456,
"status": "failure",
"cache_hit": False,
},
]
@pytest.mark.asyncio
async def test_trace_ingest_and_invalid_payload_do_not_write_spend():
storage = MagicMock()
storage.ensure_schema = AsyncMock()
logger = ClickHouseSpendLogger(storage=storage)
now = datetime.now(timezone.utc)
await logger.async_log_success_event(
{"standard_logging_object": {**_payload("trace", status="success", cost=0), "call_type": "/v1/traces"}},
None,
now,
now,
)
await logger.async_log_success_event({"standard_logging_object": "invalid"}, None, now, now)
assert logger.log_queue == []
storage.ensure_schema.assert_not_awaited()

View file

@ -22,6 +22,7 @@ from typing import Any, Dict, Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from pydantic import JsonValue, TypeAdapter, ValidationError
import litellm
from litellm.proxy._types import CommonProxyErrors
@ -33,14 +34,59 @@ from litellm.proxy.proxy_server import (
_scrub_guardrail_inner,
resolve_complexity_router_plugins,
resolve_routing_plugins,
validate_auto_router_capability_limits,
validate_deployment_access_windows,
validate_deployment_complexity_router_placement,
validate_deployment_max_agentic_loops,
validate_auto_router_capability_limits,
)
from .conftest import normalize
from pydantic import JsonValue, TypeAdapter, ValidationError
@pytest.mark.asyncio
async def test_tracing_config_automatically_logs_spend_without_callback_setting():
from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger
from litellm.proxy import tracing_endpoints
from litellm.proxy.proxy_server import ProxyStartupEvent
from litellm.tracing import TraceReceiver
from litellm.tracing.store import ClickHouseTraceStore
storage = MagicMock()
storage.ensure_schema = AsyncMock()
storage.insert_rows = AsyncMock()
receiver = TraceReceiver(ClickHouseTraceStore(storage))
prior_receiver = tracing_endpoints.receiver
try:
await ProxyStartupEvent.init_tracing({"tracing": {"store": "clickhouse"}}, receiver=receiver)
storage.ensure_schema.assert_awaited_once()
logger = next(
callback for callback in litellm._async_success_callback if isinstance(callback, ClickHouseSpendLogger)
)
now = datetime.now()
await logger.async_log_success_event(
{
"standard_logging_object": {
"id": "response-1",
"startTime": now.timestamp(),
"endTime": now.timestamp(),
"response_cost": 0.25,
}
},
None,
now,
now,
)
await logger.flush_queue()
assert storage.insert_rows.await_args.args[0] == "spend_logs"
assert storage.insert_rows.await_args.args[1][0]["spend"] == 0.25
await ProxyStartupEvent.init_tracing({})
assert all(not isinstance(callback, ClickHouseSpendLogger) for callback in litellm._async_success_callback)
finally:
await ProxyStartupEvent.init_tracing({})
tracing_endpoints.receiver = prior_receiver
# ---------------------------------------------------------------------------
# _is_remote_module_url

View file

@ -52,7 +52,7 @@ def _row(
}
def _llm_row(span_id: str, parent: str, agent: str, request_id: str, start_ms: float = 1) -> dict:
def _llm_row(span_id: str, parent: str, agent: str, request_id: str, start_ms: float = 1, **extra: Any) -> dict:
return _row(
span_id,
parent,
@ -65,6 +65,7 @@ def _llm_row(span_id: str, parent: str, agent: str, request_id: str, start_ms: f
input_tokens=100,
output_tokens=20,
litellm_request_id=request_id,
**extra,
)
@ -92,13 +93,14 @@ def test_empty_rows_is_none():
assert trace_from_rows("abc", []) is None
def test_llm_response_id_is_preserved_without_spend_enrichment():
def test_llm_response_id_is_preserved_when_spend_is_unavailable():
trace = trace_from_rows("t1", _deep_agent_rows())
assert trace is not None
spans = {span["span_id"]: span for span in trace["spans"]}
assert spans["llm-root"]["litellm_request_id"] == "chatcmpl-root"
assert spans["task"]["litellm_request_id"] is None
assert "spend" not in trace["summary"]
assert trace["summary"]["spend"] is None
assert spans["llm-root"]["spend"] is None
def test_summary_totals():
@ -163,6 +165,7 @@ def test_agent_nodes_parent_and_per_agent_counts():
"llm_calls": 1,
"tool_calls": 1,
"duration_ms": 1000,
"spend": None,
},
{
"name": "researcher",
@ -171,6 +174,7 @@ def test_agent_nodes_parent_and_per_agent_counts():
"llm_calls": 1,
"tool_calls": 1,
"duration_ms": 5,
"spend": None,
},
)
@ -309,3 +313,121 @@ async def test_get_span_not_found_and_found():
"output": "o",
"attributes": {"k": "v"},
}
@pytest.mark.asyncio
async def test_trace_cost_is_scoped_and_counts_repeated_request_once():
client = MagicMock()
spans = [
_row("root", "", "agent", "agent", "agent", team_id="team-a", api_key_hash="key-a"),
_llm_row("llm-1", "root", "agent", "response-1", team_id="team-a", api_key_hash="key-a"),
_llm_row("llm-2", "root", "agent", "response-1", team_id="team-a", api_key_hash="key-a"),
]
spend = [
{
"request_id": "request-other",
"response_id": "response-1",
"team_id": "team-b",
"api_key": "key-b",
"spend": 99.0,
"start_ms": T0 // MS,
},
{
"request_id": "request-1",
"response_id": "response-1",
"team_id": "team-a",
"api_key": "key-a",
"spend": 0.25,
"start_ms": T0 // MS,
},
{
"request_id": "request-other-key",
"response_id": "response-1",
"team_id": "team-a",
"api_key": "key-c",
"spend": 50.0,
"start_ms": T0 // MS,
},
]
client.query = AsyncMock(side_effect=[spans, spend])
store = ClickHouseTraceStore(client)
scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""}
trace = await store.get_trace("trace-1", scope)
assert trace is not None
assert trace["summary"]["spend"] == 0.25
assert trace["agents"][0]["spend"] == 0.25
assert [span["spend"] for span in trace["spans"]] == [None, 0.25, 0.25]
assert [call.args[0] for call in client.query.await_args_list] == ["trace_spans", "spend_by_response_ids"]
@pytest.mark.asyncio
async def test_run_list_uses_matching_spend_and_leaves_missing_cost_unavailable():
client = MagicMock()
rows = [
{
"trace_id": trace_id,
"trace_ref": trace_id,
"team_id": "team-a",
"api_key_hash": "key-a",
"request_ids": [request_id],
"name": "agent",
"service": "service",
"input_preview": "",
"start_ms": 1000,
"duration_ms": 100,
"status": "STATUS_CODE_OK",
"span_count": 1,
"agent_count": 1,
"llm_calls": 1,
"tool_calls": 0,
"input_tokens": 1,
"output_tokens": 1,
"models": [],
}
for trace_id, request_id in (("trace-1", "response-1"), ("trace-2", "response-2"))
]
spend = [
{
"request_id": "request-1",
"response_id": "response-1",
"team_id": "team-a",
"api_key": "key-a",
"spend": 0.25,
"start_ms": 1000,
}
]
client.query = AsyncMock(side_effect=[rows, spend])
scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""}
page = await ClickHouseTraceStore(client).list_traces(scope, 0, 2000)
assert [run["spend"] for run in page["data"]] == [0.25, None]
assert [call.args[0] for call in client.query.await_args_list] == ["list_traces", "spend_by_response_ids"]
@pytest.mark.asyncio
async def test_ambiguous_cache_response_id_keeps_cost_unavailable():
client = MagicMock()
span = _llm_row("llm-1", "", "agent", "response-1", team_id="", api_key_hash="key-a")
spend = [
{
"request_id": request_id,
"response_id": "response-1",
"team_id": "",
"api_key": "key-a",
"spend": cost,
"start_ms": T0 // MS,
}
for request_id, cost in (("response-1", 0.25), ("response-1_cache_hit123", 0.0))
]
client.query = AsyncMock(side_effect=[[span], spend])
store = ClickHouseTraceStore(client)
scope: TraceScope = {"team_ids": ("",), "api_key_hash": "key-a"}
trace = await store.get_trace("trace-1", scope)
assert trace is not None
assert trace["summary"]["spend"] is None
assert trace["spans"][0]["spend"] is None

View file

@ -123,7 +123,19 @@ describe("AgentTracesSection", () => {
const failed = rows.find((row) => row.textContent?.includes("acme-404")) as HTMLElement;
expect(within(failed).getByLabelText("2 errors")).toBeInTheDocument();
expect(screen.getByText(`${runs.length} runs`)).toBeInTheDocument();
expect(screen.queryByRole("columnheader", { name: "Cost" })).not.toBeInTheDocument();
expect(screen.getByRole("columnheader", { name: "Cost" })).toBeInTheDocument();
expect(within(failed).getByText("—")).toBeInTheDocument();
});
it("shows the spend returned for a run", async () => {
vi.mocked(agentTraceListCall).mockResolvedValue({
...(traceList as TracePage),
data: [{ ...runs[0], spend: 0.025 }],
});
renderSection();
const row = await screen.findByTestId("agent-trace-row");
expect(within(row).getByText("$0.03")).toBeInTheDocument();
});
it("filters by input text and by trace id", async () => {

View file

@ -17,9 +17,6 @@ interface AgentTracesTableProps {
onOpenTrace: (trace: TraceSummary) => void;
}
/** Spend is only on summaries once the spend-enrichment PR lands; show Cost when it's there. */
type SummaryWithSpend = TraceSummary & { spend?: number };
const SECOND_MS = 1000;
const MINUTE_S = 60;
const HOUR_M = 60;
@ -57,7 +54,6 @@ export function AgentTracesTable({
onLoadMore,
onOpenTrace,
}: AgentTracesTableProps) {
const showCost = traces.some((t) => typeof (t as SummaryWithSpend).spend === "number");
const isEmpty = !isLoading && !error && traces.length === 0;
return (
<div className="min-h-0 flex-1 overflow-auto" data-testid="runs-table">
@ -74,7 +70,7 @@ export function AgentTracesTable({
<th className={`w-[72px] ${TH_NUM}`}>Agents</th>
<th className={`w-[74px] ${TH_NUM}`}>Steps</th>
<th className={`w-[86px] ${TH_NUM}`}>Duration</th>
{showCost && <th className={`w-[80px] ${TH_NUM}`}>Cost</th>}
<th className={`w-[80px] ${TH_NUM}`}>Cost</th>
<th className={`w-[72px] ${TH_NUM}`}>Failed</th>
<th className="w-8" />
</tr>
@ -110,11 +106,9 @@ export function AgentTracesTable({
<td className={TD_NUM}>{run.agent_count.toLocaleString()}</td>
<td className={TD_NUM}>{run.span_count.toLocaleString()}</td>
<td className="px-3 text-right font-mono tabular-nums text-foreground">{fmtMs(run.duration_ms)}</td>
{showCost && (
<td className="px-3 text-right font-mono tabular-nums text-foreground">
{formatCost((run as SummaryWithSpend).spend ?? 0)}
</td>
)}
<td className="px-3 text-right font-mono tabular-nums text-foreground">
{run.spend == null ? "—" : formatCost(run.spend)}
</td>
<td className="px-3 text-right">
{run.error_count > 0 ? (
<StatusMark status="error" count={run.error_count} />

View file

@ -7,6 +7,7 @@ import { Button } from "@/components/ui/button";
import { LogDetailsDrawer } from "../LogDetailsDrawer";
import { CopyButton } from "./CopyButton";
import { formatCost } from "./AgentTracesTable";
import type { Span } from "./traceTypes";
import { fmtMs, fmtTok } from "./traceUtils";
import { useSpanRequestLog } from "./useSpanRequestLog";
@ -36,6 +37,7 @@ export function RequestDetail({ span, accessToken, traceStartMs }: RequestDetail
const rows: [string, string][] = [
["Model", span.model ?? "—"],
["Cost", span.spend == null ? "—" : formatCost(span.spend)],
["Input tokens", fmtTok(span.input_tokens)],
["Output tokens", fmtTok(span.output_tokens)],
["Total tokens", fmtTok(span.input_tokens + span.output_tokens)],

View file

@ -11,6 +11,7 @@ import { copyToClipboard } from "@/utils/dataUtils";
import { agentTraceCall, getProxyBaseUrl } from "../../networking";
import { DetailPane } from "./DetailPane";
import { formatCost } from "./AgentTracesTable";
import { SpanTree } from "./SpanTree";
import type { SpanTreeState, TreeRow } from "./traceTree";
import type { Trace } from "./traceTypes";
@ -120,6 +121,7 @@ function RunHeader({ trace, onBack }: { trace: Trace; onBack: () => void }) {
<div className="flex min-w-0 items-center gap-4 font-mono text-[11px] text-foreground tabular-nums">
<Stat label="duration" value={fmtMs(summary.duration_ms)} />
<Stat label="steps" value={summary.span_count.toLocaleString()} />
<Stat label="cost" value={summary.spend == null ? "—" : formatCost(summary.spend)} />
{failed && <Stat label="failed" value={summary.error_count.toLocaleString()} error />}
</div>
<div className="ml-auto">

View file

@ -25,6 +25,7 @@ export interface Span {
input_tokens: number;
output_tokens: number;
litellm_request_id: string | null;
spend?: number | null;
}
/** One distinct agent in a trace. 200 invocations of `researcher` = one node. */
@ -35,6 +36,7 @@ export interface AgentNode {
llm_calls: number;
tool_calls: number;
duration_ms: number;
spend?: number | null;
}
export interface TraceSummary {
@ -56,6 +58,7 @@ export interface TraceSummary {
input_tokens: number;
output_tokens: number;
models: string[];
spend?: number | null;
}
export interface Trace {