diff --git a/docker/docker-compose.tracing.yml b/docker/docker-compose.tracing.yml new file mode 100644 index 00000000000..b39fc8f4561 --- /dev/null +++ b/docker/docker-compose.tracing.yml @@ -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: diff --git a/docker/tracing-config.yaml b/docker/tracing-config.yaml new file mode 100644 index 00000000000..03637cfa9fb --- /dev/null +++ b/docker/tracing-config.yaml @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 6a18273ed4c..226e88b3430 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -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> { + 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, ¶meters).await }, + async move { + litellm_traces::execute_named_read(&client, &connection, query, ¶meters).await + }, map_error, ) } diff --git a/litellm-rust/crates/traces/query/list_traces.sql b/litellm-rust/crates/traces/query/list_traces.sql new file mode 100644 index 00000000000..c0c1b28aa7f --- /dev/null +++ b/litellm-rust/crates/traces/query/list_traces.sql @@ -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} diff --git a/litellm-rust/crates/traces/query/span_detail.sql b/litellm-rust/crates/traces/query/span_detail.sql new file mode 100644 index 00000000000..37bb4e8a87e --- /dev/null +++ b/litellm-rust/crates/traces/query/span_detail.sql @@ -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 diff --git a/litellm-rust/crates/traces/query/spend_by_response_ids.sql b/litellm-rust/crates/traces/query/spend_by_response_ids.sql new file mode 100644 index 00000000000..285e9235629 --- /dev/null +++ b/litellm-rust/crates/traces/query/spend_by_response_ids.sql @@ -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 diff --git a/litellm-rust/crates/traces/query/trace_spans.sql b/litellm-rust/crates/traces/query/trace_spans.sql new file mode 100644 index 00000000000..409e6328198 --- /dev/null +++ b/litellm-rust/crates/traces/query/trace_spans.sql @@ -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 diff --git a/litellm-rust/crates/traces/src/error.rs b/litellm-rust/crates/traces/src/error.rs index 125edc35422..4a4fdaa00f7 100644 --- a/litellm-rust/crates/traces/src/error.rs +++ b/litellm-rust/crates/traces/src/error.rs @@ -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}")] diff --git a/litellm-rust/crates/traces/src/insert.rs b/litellm-rust/crates/traces/src/insert.rs index 6f2d6acb023..5c5ed7e3c9c 100644 --- a/litellm-rust/crates/traces/src/insert.rs +++ b/litellm-rust/crates/traces/src/insert.rs @@ -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) diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index 279afb20e9b..5402b54385e 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -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)] diff --git a/litellm-rust/crates/traces/src/sql.rs b/litellm-rust/crates/traces/src/sql.rs index 1aa21a59caa..8925b942a7f 100644 --- a/litellm-rust/crates/traces/src/sql.rs +++ b/litellm-rust/crates/traces/src/sql.rs @@ -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 { + 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, +) -> Result { + execute_read(client, connection, query.sql(), parameters).await +} diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index ac0266409fa..bea5016322a 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -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, +) -> 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( diff --git a/litellm/integrations/clickhouse/clickhouse_batch_logger.py b/litellm/integrations/clickhouse/clickhouse_batch_logger.py index fb9088f44ff..81601ea2a78 100644 --- a/litellm/integrations/clickhouse/clickhouse_batch_logger.py +++ b/litellm/integrations/clickhouse/clickhouse_batch_logger.py @@ -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()) diff --git a/litellm/integrations/clickhouse/clickhouse_spend_logger.py b/litellm/integrations/clickhouse/clickhouse_spend_logger.py new file mode 100644 index 00000000000..237d297c38a --- /dev/null +++ b/litellm/integrations/clickhouse/clickhouse_spend_logger.py @@ -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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f89dde04e06..22688631b96 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index a0b010caea0..e1a3d1b6acc 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -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 diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index 244eddd3def..806757306c0 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -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: diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index b8b6f646111..f7e75538951 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -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): diff --git a/scripts/run_tracing_proxy_local.sh b/scripts/run_tracing_proxy_local.sh new file mode 100755 index 00000000000..fd48590bf93 --- /dev/null +++ b/scripts/run_tracing_proxy_local.sh @@ -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 diff --git a/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py b/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py index e5b00b3c783..bae94ba6100 100644 --- a/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py +++ b/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py @@ -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()) diff --git a/tests/test_litellm/integrations/clickhouse/test_clickhouse_spend_logger.py b/tests/test_litellm/integrations/clickhouse/test_clickhouse_spend_logger.py new file mode 100644 index 00000000000..1fc10813b8d --- /dev/null +++ b/tests/test_litellm/integrations/clickhouse/test_clickhouse_spend_logger.py @@ -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() diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 0157200ed5c..7096bc7c632 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -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 diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 39ce3162073..7ee772e078c 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -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 diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx index 6d28e8a53e9..126542e848f 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesSection.test.tsx @@ -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 () => { diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx index d3912b7be5c..0eeb6b76837 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/AgentTracesTable.tsx @@ -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 (
@@ -74,7 +70,7 @@ export function AgentTracesTable({ Agents Steps Duration - {showCost && Cost} + Cost Failed @@ -110,11 +106,9 @@ export function AgentTracesTable({ {run.agent_count.toLocaleString()} {run.span_count.toLocaleString()} {fmtMs(run.duration_ms)} - {showCost && ( - - {formatCost((run as SummaryWithSpend).spend ?? 0)} - - )} + + {run.spend == null ? "—" : formatCost(run.spend)} + {run.error_count > 0 ? ( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx index 156393e6b37..f6088fc8640 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/RequestDetail.tsx @@ -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)], diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx index cddba8ffbea..21f580fdb1b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx @@ -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 }) {
+ {failed && }
diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts index 18f705676c3..f733bae1a6f 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts @@ -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 {