refactor(proxy): inject tracing receiver and access context (#44035)

* refactor(proxy): inject tracing receiver and access context

* refactor(proxy): own tracing resources through FastAPI lifespan

* test(proxy): pass tracing dependency in Lens lifecycle

* refactor(proxy): stop tracing logger cooperatively

* refactor(proxy): derive tracing permissions in one place

* refactor(proxy): compose application lifespan state

* refactor(proxy): give Lens tracing storage directly

* refactor(tracing): name shared ClickHouse storage explicitly

* refactor(tracing): extract shared ClickHouse storage crate

* test(proxy): isolate db push timeout from Lens safety check

* fix(tracing): drain spend retries during shutdown
This commit is contained in:
yujonglee 2026-10-01 13:45:32 -07:00 • committed by GitHub
parent 163ebccad5
commit be67fce26a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
36 changed files with 1152 additions and 559 deletions

View file

@ -4086,6 +4086,7 @@ dependencies = [
"litellm-secrets",
"litellm-secrets-aws",
"litellm-secrets-types",
"litellm-storage-clickhouse",
"litellm-token-counter",
"litellm-traces",
"litellm-tracing",
@ -4288,6 +4289,20 @@ dependencies = [
"veil",
]
[[package]]
name = "litellm-storage-clickhouse"
version = "0.1.0"
dependencies = [
"flate2",
"litellm-http",
"rstest",
"serde",
"serde_json",
"thiserror 2.0.19",
"tokio",
"url",
]
[[package]]
name = "litellm-testkit"
version = "0.1.0"
@ -4371,6 +4386,7 @@ dependencies = [
"base64 0.22.1",
"flate2",
"litellm-http",
"litellm-storage-clickhouse",
"opentelemetry-proto",
"prost",
"rstest",
@ -4381,7 +4397,6 @@ dependencies = [
"thiserror 2.0.19",
"time",
"tokio",
"url",
]
[[package]]

View file

@ -13,6 +13,7 @@ litellm-config = { path = "crates/config" }
litellm-router = { path = "crates/router" }
litellm-tracing = { path = "crates/tracing" }
litellm-traces = { path = "crates/traces" }
litellm-storage-clickhouse = { path = "crates/storage-clickhouse" }
litellm-core = { path = "crates/core" }
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
litellm-gateway = { path = "crates/gateway" }

View file

@ -22,6 +22,7 @@ tiktoken = ["litellm-token-counter/tiktoken"]
fancy-regex.workspace = true
litellm-tracing.workspace = true
litellm-traces.workspace = true
litellm-storage-clickhouse.workspace = true
litellm-host.workspace = true
bytes.workspace = true
futures-util.workspace = true

View file

@ -1,7 +1,8 @@
use std::collections::BTreeMap;
use litellm_http::ClientVariant;
use litellm_traces::{Connection, Error, InsertTable, Parameter, ReadQuery};
use litellm_storage_clickhouse::Storage;
use litellm_traces::{Error, InsertTable, Parameter, ReadQuery};
use pyo3::{
exceptions::{PyOverflowError, PyRuntimeError, PyValueError},
prelude::*,
@ -27,9 +28,7 @@ fn map_error(error: Error) -> PyErr {
#[pyclass]
pub struct NativeTraceStorage {
database: String,
writer: Connection,
reader: Option<Connection>,
storage: Storage,
}
#[pymethods]
@ -39,12 +38,7 @@ impl NativeTraceStorage {
fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult<Self> {
litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?;
Ok(Self {
writer: Connection::writer(url).map_err(map_error)?,
reader: reader_url
.map(|value| Connection::reader(value, &database))
.transpose()
.map_err(map_error)?,
database,
storage: Storage::new(database, url, reader_url).map_err(map_error)?,
})
}
@ -55,8 +49,8 @@ impl NativeTraceStorage {
spend_log_retention_days: u32,
) -> PyResult<Bound<'py, PyAny>> {
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
let connection = self.writer.clone();
let database = self.database.clone();
let connection = self.storage.writer().clone();
let database = self.storage.database().to_owned();
crate::execution::run_async(
py,
async move {
@ -83,8 +77,8 @@ impl NativeTraceStorage {
) -> PyResult<Bound<'py, PyAny>> {
let table = InsertTable::parse(table).map_err(map_error)?;
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
let connection = self.writer.clone();
let database = self.database.clone();
let connection = self.storage.writer().clone();
let database = self.storage.database().to_owned();
crate::execution::run_async(
py,
async move {
@ -104,7 +98,7 @@ impl NativeTraceStorage {
>,
) -> PyResult<Bound<'py, PyAny>> {
let query = litellm_traces::LensQuery::parse(name).map_err(map_error)?;
let connection = self.reader.clone().ok_or_else(|| {
let connection = self.storage.reader().cloned().ok_or_else(|| {
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
})?;
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
@ -127,7 +121,7 @@ impl NativeTraceStorage {
>,
) -> PyResult<Bound<'py, PyAny>> {
let query = ReadQuery::parse(query).map_err(map_error)?;
let connection = self.reader.clone().ok_or_else(|| {
let connection = self.storage.reader().cloned().ok_or_else(|| {
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
})?;
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;

View file

@ -0,0 +1,20 @@
[package]
name = "litellm-storage-clickhouse"
version = "0.1.0"
description = "Shared ClickHouse connection and HTTP storage for LiteLLM features"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
flate2.workspace = true
litellm-http.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }
rstest.workspace = true
tokio.workspace = true

View file

@ -0,0 +1,5 @@
# ClickHouse storage
`litellm-storage-clickhouse` exports `Storage`, a shared writer connection and optional reader connection for one ClickHouse database. It also exports bounded HTTP read and insert execution
The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces` supplies those rules and uses this storage for both trace rows and spend rows

View file

@ -0,0 +1,29 @@
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("invalid ClickHouse insert row")]
InvalidRow,
#[error("invalid ClickHouse insert table")]
InvalidTable,
#[error("invalid ClickHouse HTTP URL")]
InvalidUrl,
#[error("database must be a nonempty SQL identifier and retention must be positive")]
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}")]
InsertFailed(u16),
#[error("ClickHouse insert exceeds the encoded size limit")]
InsertTooLarge,
#[error("ClickHouse schema setup failed with HTTP status {0}")]
SchemaFailed(u16),
#[error("ClickHouse query exceeded the response size limit")]
ResponseTooLarge,
#[error("ClickHouse returned an invalid or failed JSON query response")]
InvalidResponse,
#[error("ClickHouse query transport failed")]
Transport,
}

View file

@ -0,0 +1,70 @@
use std::{io::Write, time::Duration};
use flate2::{Compression, write::GzEncoder};
use litellm_http::Client;
use crate::{Connection, Error, valid_identifier};
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
pub async fn insert_encoded_rows(
client: &Client,
connection: &Connection,
database: &str,
table: &str,
token: &str,
encoded: &str,
) -> Result<(), Error> {
if !valid_identifier(database) {
return Err(Error::InvalidSchema);
}
if !valid_identifier(table) {
return Err(Error::InvalidTable);
}
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder
.write_all(encoded.as_bytes())
.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!("INSERT INTO `{database}`.{} FORMAT JSONEachRow", table),
)
.append_pair("insert_deduplication_token", token)
.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)
.timeout(INSERT_TIMEOUT)
.header("Content-Encoding", "gzip")
.body(body)
.send()
.await
.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::InsertFailed(response.status().as_u16()));
}
Ok(())
}

View file

@ -0,0 +1,127 @@
mod error;
mod insert;
mod read;
pub use error::Error;
pub use insert::insert_encoded_rows;
pub use read::{Parameter, execute_read};
use url::Url;
#[derive(Clone)]
pub struct Connection {
url: Url,
}
impl Connection {
pub fn parse(value: &str) -> Result<Self, Error> {
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
return Err(Error::InvalidUrl);
}
Ok(Self { url })
}
pub fn configured(
url: &str,
database: &str,
user: &str,
password: &str,
) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
connection
.url
.set_username(user)
.map_err(|_| Error::InvalidUrl)?;
connection
.url
.set_password(Some(password))
.map_err(|_| Error::InvalidUrl)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection
.url
.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("database", database);
Ok(connection)
}
pub fn writer(url: &str) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
Ok(connection)
}
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| key != "database")
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection
.url
.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("database", database);
Ok(connection)
}
pub fn url(&self) -> &Url {
&self.url
}
}
#[derive(Clone)]
pub struct Storage {
database: String,
writer: Connection,
reader: Option<Connection>,
}
impl Storage {
pub fn new(database: String, url: &str, reader_url: Option<&str>) -> Result<Self, Error> {
if !valid_identifier(&database) {
return Err(Error::InvalidSchema);
}
Ok(Self {
writer: Connection::writer(url)?,
reader: reader_url
.map(|value| Connection::reader(value, &database))
.transpose()?,
database,
})
}
pub fn database(&self) -> &str {
&self.database
}
pub fn writer(&self) -> &Connection {
&self.writer
}
pub fn reader(&self) -> Option<&Connection> {
self.reader.as_ref()
}
}
pub(crate) fn valid_identifier(value: &str) -> bool {
!value.is_empty()
&& value
.bytes()
.all(|c| c.is_ascii_alphanumeric() || c == b'_')
}

View file

@ -0,0 +1,113 @@
use std::{collections::BTreeMap, time::Duration};
use litellm_http::Client;
use serde::Deserialize;
use crate::{Connection, Error};
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum Parameter {
Text(String),
Integer(i64),
Strings(Vec<String>),
}
impl Parameter {
fn encoded(&self) -> String {
match self {
Self::Text(value) => escaped(value),
Self::Integer(value) => value.to_string(),
Self::Strings(values) => format!(
"[{}]",
values
.iter()
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
.collect::<Vec<_>>()
.join(",")
),
}
}
}
fn escaped(value: &str) -> String {
value
.replace('\\', "\\\\")
.replace('\t', "\\t")
.replace('\n', "\\n")
.replace('\r', "\\r")
.replace('\0', "\\0")
}
pub async fn execute_read(
client: &Client,
connection: &Connection,
sql: &str,
parameters: &BTreeMap<String, Parameter>,
) -> Result<String, Error> {
if sql.trim().is_empty() {
return Err(Error::EmptySql);
}
let mut url = connection.url().clone();
let existing_pairs: Vec<(String, String)> = url
.query_pairs()
.filter(|(key, _)| {
!key.starts_with("param_")
&& !matches!(
key.as_ref(),
"query"
| "readonly"
| "default_format"
| "max_result_rows"
| "result_overflow_mode"
| "max_execution_time"
| "wait_end_of_query"
)
})
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
url.query_pairs_mut()
.clear()
.extend_pairs(existing_pairs)
.append_pair("readonly", "1")
.append_pair("max_result_rows", "1000")
.append_pair("result_overflow_mode", "throw")
.append_pair("max_execution_time", "10")
.append_pair("wait_end_of_query", "1")
.append_pair("default_format", "JSON");
url.query_pairs_mut().extend_pairs(
parameters
.iter()
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
);
let request = client
.post(url)
.timeout(Duration::from_secs(15))
.body(sql.to_owned());
let mut response = request.send().await.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::QueryFailed(response.status().as_u16()));
}
let mut body = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
if body.len() + chunk.len() > MAX_RESPONSE_BYTES {
return Err(Error::ResponseTooLarge);
}
body.extend_from_slice(&chunk);
}
let json: serde_json::Value =
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
{
return Err(Error::InvalidResponse);
}
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
}

View file

@ -0,0 +1,34 @@
use litellm_storage_clickhouse::{Connection, Storage};
use rstest::rstest;
#[rstest]
#[case::http("http://localhost:8123", true)]
#[case::https("https://localhost:8443", true)]
#[case::tcp("tcp://localhost:9000", false)]
#[case::missing_host("http://", false)]
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
assert_eq!(Connection::parse(value).is_ok(), expected);
}
#[rstest]
#[case::writer_only(None, false)]
#[case::separate_reader(Some("http://localhost:8124"), true)]
fn storage_exports_writer_and_optional_reader(
#[case] reader_url: Option<&str>,
#[case] has_reader: bool,
) {
let storage = Storage::new("litellm".to_owned(), "http://localhost:8123", reader_url)
.expect("valid ClickHouse URLs");
assert_eq!(storage.database(), "litellm");
assert_eq!(storage.writer().url().host_str(), Some("localhost"));
assert_eq!(storage.writer().url().port(), Some(8123));
assert_eq!(storage.reader().is_some(), has_reader);
}
#[rstest]
#[case::empty("")]
#[case::injection("db; DROP DATABASE default")]
fn storage_rejects_invalid_database(#[case] database: &str) {
assert!(Storage::new(database.to_owned(), "http://localhost:8123", None).is_err());
}

View file

@ -0,0 +1,34 @@
use std::collections::BTreeMap;
use litellm_http::Client;
use litellm_storage_clickhouse::{Connection, Error, execute_read, insert_encoded_rows};
use rstest::rstest;
#[rstest]
#[case::invalid_database("db; DROP DATABASE default", "spend_logs", true)]
#[case::invalid_table("litellm", "spend_logs; DROP TABLE otel_traces", false)]
#[tokio::test]
async fn insert_rejects_invalid_identifiers(
#[case] database: &str,
#[case] table: &str,
#[case] invalid_database: bool,
) {
let client = Client::no_redirect_for_test();
let connection = Connection::writer("http://localhost:8123").expect("valid URL");
let result = insert_encoded_rows(&client, &connection, database, table, "token", "{}").await;
assert!(matches!(&result, Err(Error::InvalidSchema)) == invalid_database);
assert!(matches!(&result, Err(Error::InvalidTable)) == !invalid_database);
}
#[rstest]
#[tokio::test]
async fn read_rejects_empty_sql() {
let client = Client::no_redirect_for_test();
let connection = Connection::reader("http://localhost:8123", "litellm").expect("valid URL");
assert!(matches!(
execute_read(&client, &connection, " ", &BTreeMap::new()).await,
Err(Error::EmptySql)
));
}

View file

@ -1,4 +1,4 @@
- Rust owns OTLP wire decoding, ClickHouse schema, row encoding, named reads, connection validation and transport
- Keep OTLP decoding, trace schema, row encoding and named query selection here. Generic ClickHouse connections and HTTP execution belong in `litellm-storage-clickhouse`
- Keep this crate independent of Python; PyO3 conversion and public Python exceptions belong in `python-bridge`
- Keep the SQL migrations here as the only ClickHouse schema definition
- Use typed query parameters and a dedicated SELECT-only reader with server-side limits

View file

@ -12,11 +12,11 @@ opentelemetry-proto = { version = "0.33.0", default-features = false, features =
prost = "0.14.4"
time = { workspace = true, features = ["formatting"] }
litellm-http.workspace = true
litellm-storage-clickhouse.workspace = true
sha2.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
url.workspace = true
[dev-dependencies]
litellm-http = { workspace = true, features = ["test-support"] }

View file

@ -1,33 +1,3 @@
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error("invalid ClickHouse insert row")]
InvalidRow,
#[error("invalid ClickHouse insert table")]
InvalidTable,
#[error("invalid ClickHouse HTTP URL")]
InvalidUrl,
#[error("database must be a nonempty SQL identifier and retention must be positive")]
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}")]
InsertFailed(u16),
#[error("ClickHouse insert exceeds the encoded size limit")]
InsertTooLarge,
#[error("ClickHouse schema setup failed with HTTP status {0}")]
SchemaFailed(u16),
#[error("ClickHouse query exceeded the response size limit")]
ResponseTooLarge,
#[error("ClickHouse returned an invalid or failed JSON query response")]
InvalidResponse,
#[error("ClickHouse query transport failed")]
Transport,
}
#[derive(Debug, thiserror::Error)]
pub enum DecodeError {
#[error("invalid OTLP trace payload")]

View file

@ -1,6 +1,5 @@
use std::{collections::BTreeMap, io::Write, time::Duration};
use std::collections::BTreeMap;
use flate2::{Compression, write::GzEncoder};
use litellm_http::Client;
use serde_json::Value;
use sha2::{Digest, Sha256};
@ -9,7 +8,6 @@ use time::{OffsetDateTime, format_description::well_known::Rfc3339};
use crate::{Connection, Error};
const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024;
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
pub enum InsertTable {
OtelTraces,
@ -61,55 +59,15 @@ pub async fn insert_rows(
})
.collect();
let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?;
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
encoder
.write_all(encoded.as_bytes())
.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!(
"INSERT INTO `{database}`.{} FORMAT JSONEachRow",
table.name()
),
)
.append_pair("insert_deduplication_token", &token)
.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)
.timeout(INSERT_TIMEOUT)
.header("Content-Encoding", "gzip")
.body(body)
.send()
.await
.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::InsertFailed(response.status().as_u16()));
}
Ok(())
litellm_storage_clickhouse::insert_encoded_rows(
client,
connection,
database,
table.name(),
&token,
&encoded,
)
.await
}
pub fn encode_rows(rows: Vec<BTreeMap<String, Value>>) -> Result<String, Error> {

View file

@ -4,87 +4,9 @@ mod otlp;
mod schema;
mod sql;
pub use error::{DecodeError, Error};
pub use error::DecodeError;
pub use insert::{InsertTable, encode_rows, insert_rows};
pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read};
pub use otlp::{DecodedSpan, decode_otlp};
pub use schema::{ensure_schema, schema_statements};
pub use sql::{LensQuery, Parameter, ReadQuery, execute_named_read, execute_read};
use url::Url;
#[derive(Clone)]
pub struct Connection {
url: Url,
}
impl Connection {
pub fn parse(value: &str) -> Result<Self, Error> {
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
return Err(Error::InvalidUrl);
}
Ok(Self { url })
}
pub fn configured(
url: &str,
database: &str,
user: &str,
password: &str,
) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
connection
.url
.set_username(user)
.map_err(|_| Error::InvalidUrl)?;
connection
.url
.set_password(Some(password))
.map_err(|_| Error::InvalidUrl)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection
.url
.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("database", database);
Ok(connection)
}
pub fn writer(url: &str) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
Ok(connection)
}
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
let mut connection = Self::parse(url)?;
let pairs: Vec<_> = connection
.url
.query_pairs()
.filter(|(key, _)| key != "database")
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
connection
.url
.query_pairs_mut()
.clear()
.extend_pairs(pairs)
.append_pair("database", database);
Ok(connection)
}
pub fn url(&self) -> &Url {
&self.url
}
}
pub use sql::{LensQuery, ReadQuery, execute_named_read};

View file

@ -1,12 +1,8 @@
use std::{collections::BTreeMap, time::Duration};
use serde::Deserialize;
use std::collections::BTreeMap;
use litellm_http::Client;
use crate::{Connection, Error};
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
use crate::{Connection, Error, Parameter, execute_read};
pub enum ReadQuery {
ListTraces,
@ -36,111 +32,6 @@ impl ReadQuery {
}
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
pub enum Parameter {
Text(String),
Integer(i64),
Strings(Vec<String>),
}
impl Parameter {
fn encoded(&self) -> String {
match self {
Self::Text(value) => escaped(value),
Self::Integer(value) => value.to_string(),
Self::Strings(values) => format!(
"[{}]",
values
.iter()
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
.collect::<Vec<_>>()
.join(",")
),
}
}
}
fn escaped(value: &str) -> String {
value
.replace('\\', "\\\\")
.replace('\t', "\\t")
.replace('\n', "\\n")
.replace('\r', "\\r")
.replace('\0', "\\0")
}
pub async fn execute_read(
client: &Client,
connection: &Connection,
sql: &str,
parameters: &BTreeMap<String, Parameter>,
) -> Result<String, Error> {
if sql.trim().is_empty() {
return Err(Error::EmptySql);
}
let mut url = connection.url().clone();
let existing_pairs: Vec<(String, String)> = url
.query_pairs()
.filter(|(key, _)| {
!key.starts_with("param_")
&& !matches!(
key.as_ref(),
"query"
| "readonly"
| "default_format"
| "max_result_rows"
| "result_overflow_mode"
| "max_execution_time"
| "wait_end_of_query"
)
})
.map(|(key, value)| (key.into_owned(), value.into_owned()))
.collect();
url.query_pairs_mut()
.clear()
.extend_pairs(existing_pairs)
.append_pair("readonly", "1")
.append_pair("max_result_rows", "1000")
.append_pair("result_overflow_mode", "throw")
.append_pair("max_execution_time", "10")
.append_pair("wait_end_of_query", "1")
.append_pair("default_format", "JSON");
url.query_pairs_mut().extend_pairs(
parameters
.iter()
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
);
let request = client
.post(url)
.timeout(Duration::from_secs(15))
.body(sql.to_owned());
let mut response = request.send().await.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
return Err(Error::QueryFailed(response.status().as_u16()));
}
let mut body = Vec::new();
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
if body.len() + chunk.len() > MAX_RESPONSE_BYTES {
return Err(Error::ResponseTooLarge);
}
body.extend_from_slice(&chunk);
}
let json: serde_json::Value =
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
{
return Err(Error::InvalidResponse);
}
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
}
#[derive(Clone, Copy)]
pub enum LensQuery {
Sample,

View file

@ -1,11 +0,0 @@
use litellm_traces::Connection;
use rstest::rstest;
#[rstest]
#[case::http("http://localhost:8123", true)]
#[case::https("https://localhost:8443", true)]
#[case::tcp("tcp://localhost:9000", false)]
#[case::missing_host("http://", false)]
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
assert_eq!(Connection::parse(value).is_ok(), expected);
}

View file

@ -11,7 +11,8 @@ gzip JSONEachRow insert, either every `CLICKHOUSE_FLUSH_INTERVAL_SECONDS` or as
import asyncio
import os
from collections.abc import Mapping, Sequence
from typing import Any, ClassVar
from contextlib import suppress
from typing import Any, ClassVar, Final
from litellm._logging import verbose_logger
from litellm.constants import (
@ -21,11 +22,11 @@ from litellm.constants import (
CLICKHOUSE_MAX_RETRIES,
)
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.rust_bridge.traces import TraceStorage
from litellm.rust_bridge.traces import ClickHouseStorage
def clickhouse_storage_from_env() -> TraceStorage:
return TraceStorage(
def clickhouse_storage_from_env() -> ClickHouseStorage:
return ClickHouseStorage(
database=os.getenv("CLICKHOUSE_DATABASE", "litellm"),
url=os.getenv("CLICKHOUSE_URL", ""),
)
@ -34,7 +35,7 @@ def clickhouse_storage_from_env() -> TraceStorage:
class ClickHouseBatchLogger(CustomBatchLogger):
table: ClassVar[str]
def __init__(self, storage: TraceStorage | None = None) -> None:
def __init__(self, storage: ClickHouseStorage | None = None) -> None:
self.storage = storage or clickhouse_storage_from_env()
self.rows_written = 0
self.rows_dropped = 0
@ -45,11 +46,27 @@ class ClickHouseBatchLogger(CustomBatchLogger):
flush_interval=CLICKHOUSE_FLUSH_INTERVAL_SECONDS,
)
self._flush_task: asyncio.Task[None] | None = None
self._stop: Final = asyncio.Event()
def start(self) -> None:
if self._flush_task is None or self._flush_task.done():
self._flush_task = asyncio.get_running_loop().create_task(self.periodic_flush())
async def aclose(self) -> None:
self._stop.set()
if self._flush_task is not None:
await self._flush_task
while self.log_queue:
await self.flush_queue()
async def periodic_flush(self) -> None:
while True:
with suppress(asyncio.TimeoutError):
await asyncio.wait_for(self._stop.wait(), timeout=self.flush_interval)
if self._stop.is_set():
return
await self.flush_queue()
def is_full(self) -> bool:
"""Backpressure signal: producers should reject (429) instead of enqueueing."""
return len(self.log_queue) >= CLICKHOUSE_MAX_BUFFERED_ROWS

View file

@ -1,11 +1,11 @@
from typing import Final
from litellm.rust_bridge.traces import TraceStorage
from litellm.rust_bridge.traces import ClickHouseStorage
OTEL_TRACES_TABLE: Final = "otel_traces"
AGENT_TRACES_BY_KEY_TABLE: Final = "agent_traces_by_key"
SPEND_LOGS_TABLE: Final = "spend_logs"
async def ensure_schema(storage: TraceStorage, trace_retention_days: int, spend_log_retention_days: int) -> None:
async def ensure_schema(storage: ClickHouseStorage, trace_retention_days: int, spend_log_retention_days: int) -> None:
await storage.ensure_schema(trace_retention_days, spend_log_retention_days)

View file

@ -88,11 +88,17 @@ from .types_utils.utils import get_instance_fn, validate_custom_validate_return_
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
from litellm.tracing import TraceReceiver
Span = _Span | Any
else:
Span = Any
class ProxyLifespanState(TypedDict):
tracing_receiver: ReadOnly["TraceReceiver | None"]
class ReconcileOutcome(NamedTuple):
"""What a model reconcile observed, captured while it still held the reconcile
lock.

View file

@ -35,7 +35,7 @@ from litellm.proxy.lens.models import (
WorkerCreated,
)
from litellm.proxy.lens.repository import LensRepository, WriterDatabase
from litellm.proxy.lens.sources import SourceReader, parse_execution
from litellm.proxy.lens.sources import SourceReader, Storage, parse_execution
from litellm.proxy.lens.state import (
can_access,
claim_job,
@ -45,10 +45,12 @@ from litellm.proxy.lens.state import (
replace_job,
snapshot_finding,
)
from litellm.proxy.tracing_runtime import provide_storage
router: Final = APIRouter(prefix="/lens", tags=["Lens"])
_bearer: Final = HTTPBearer()
Auth: TypeAlias = Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)]
StorageDep: TypeAlias = Annotated[Storage | None, Depends(provide_storage)]
def repository() -> LensRepository:
@ -59,10 +61,13 @@ def repository() -> LensRepository:
return LensRepository(WriterDatabase(writer_wrapper(prisma_client.db)))
def source_reader() -> SourceReader:
from litellm.proxy.tracing_endpoints import get_receiver
return SourceReader(get_receiver().store.storage)
def source_reader(storage: Storage | None) -> SourceReader:
if storage is None:
raise HTTPException(
status_code=501,
detail="Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL.",
)
return SourceReader(storage)
def user_scope(auth: UserAPIKeyAuth, write: bool = False) -> Scope:
@ -138,14 +143,12 @@ def validate_model(settings: LensSettings, auth: UserAPIKeyAuth) -> None:
@router.get("", response_model=LensList)
async def list_lenses(auth: Auth) -> LensList:
from litellm.proxy import tracing_endpoints
async def list_lenses(auth: Auth, storage: StorageDep) -> LensList:
scope: Final = user_scope(auth)
return LensList(
lenses=tuple(e for e in await repository().lenses() if can_access(scope, e.scope)),
workers=tuple(w for w in await repository().workers() if can_access(scope, w.scope)),
tracing_enabled=tracing_endpoints.receiver is not None,
tracing_enabled=storage is not None,
)
@ -265,10 +268,10 @@ class Preview(BaseModel):
@router.post("/preview/sample", response_model=Sample)
async def preview_sample(body: Preview, auth: Auth) -> Sample:
async def preview_sample(body: Preview, auth: Auth, storage: StorageDep) -> Sample:
validate_selection(body.settings)
now: Final = min(body.as_of or datetime.now(timezone.utc), datetime.now(timezone.utc))
return await source_reader().sample(
return await source_reader(storage).sample(
user_scope(auth),
body.settings,
int((now - timedelta(hours=body.lookback_hours)).timestamp() * 1000),
@ -367,14 +370,14 @@ async def progress(lens_id: str, job_id: str, body: Progress, worker: WorkerAuth
@router.get("/worker/{lens_id}/{job_id}/sample", response_model=Sample)
async def sample(lens_id: str, job_id: str, worker: WorkerAuth) -> Sample:
async def sample(lens_id: str, job_id: str, worker: WorkerAuth, storage: StorageDep) -> Sample:
lens, job = await assigned(lens_id, job_id, worker)
if job.sample is not None:
return job.sample
pages: list[Sample] = [] # mutable-ok: freeze selection after stable cursor traversal
cursor = "" # rebind-ok: advance by immutable identity, never by shifting row positions
while True:
page = await source_reader().sample(
page = await source_reader(storage).sample(
lens.scope,
job.settings,
int(job.start.timestamp() * 1000),
@ -413,6 +416,7 @@ async def content(
job_id: str,
execution_id: str,
worker: WorkerAuth,
storage: StorageDep,
cursor: str = "",
offset: int = Query(default=0, ge=0),
) -> ExecutionContent:
@ -421,7 +425,7 @@ async def content(
execution: Final = next((e for e in selected.executions if e.id == execution_id), None)
if execution is None:
raise HTTPException(404, "Execution is outside this job's sample")
return await source_reader().content(lens.scope, execution, cursor, offset)
return await source_reader(storage).content(lens.scope, execution, cursor, offset)
@router.post("/worker/{lens_id}/{job_id}/model", response_model=ModelResult)
@ -433,7 +437,7 @@ async def model(lens_id: str, job_id: str, body: ModelRequest, worker: WorkerAut
@router.post("/worker/{lens_id}/{job_id}/result", response_model=Lens)
async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth) -> Lens:
async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth, storage: StorageDep) -> Lens:
lens: Final = await get_lens(lens_id, worker.scope)
old: Final = next((j for j in lens.jobs if j.id == job_id), None)
if old and old.status in ("completed", "failed") and old.worker_id == worker.id:
@ -455,7 +459,7 @@ async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth) ->
raise HTTPException(422, "Finding references evidence outside the job")
for finding in body.findings:
await validate_finding(lens, selected, finding)
await validate_finding(lens, selected, finding, storage)
def finish(e: Lens) -> Lens:
active: Final = current_job(e)
@ -523,12 +527,12 @@ async def claim_candidate(candidate: Lens, worker: Worker, now: datetime) -> Cla
return None
async def validate_finding(lens: Lens, selected: Sample, finding: FindingDraft) -> None:
async def validate_finding(lens: Lens, selected: Sample, finding: FindingDraft, storage: Storage | None) -> None:
previous: Final = next((f for f in lens.findings if f.id == finding.existing_finding_id), None)
if finding.existing_finding_id and (previous is None or previous.check_id != finding.check_id):
raise HTTPException(422, "Existing finding must belong to the same check")
for evidence in finding.evidence:
if not await source_reader().verify_evidence(
if not await source_reader(storage).verify_evidence(
lens.scope, next(e for e in selected.executions if e.id == evidence.execution_id), evidence
):
raise HTTPException(422, "Evidence quote does not match stored content")
@ -536,7 +540,12 @@ async def validate_finding(lens: Lens, selected: Sample, finding: FindingDraft)
@router.get("/{lens_id}/executions/{execution_id}", response_model=ExecutionContent)
async def evidence_content(
lens_id: str, execution_id: str, auth: Auth, cursor: str = "", offset: int = Query(default=0, ge=0)
lens_id: str,
execution_id: str,
auth: Auth,
storage: StorageDep,
cursor: str = "",
offset: int = Query(default=0, ge=0),
) -> ExecutionContent:
lens: Final = await get_lens(lens_id, user_scope(auth))
try:
@ -556,4 +565,4 @@ async def evidence_content(
span_count=1,
root_seen=source == "requests",
)
return await source_reader().content(lens.scope, execution, cursor, offset)
return await source_reader(storage).content(lens.scope, execution, cursor, offset)

View file

@ -119,6 +119,7 @@ from litellm.proxy._types import (
PassThroughGenericEndpoint,
ProxyErrorTypes,
ProxyException,
ProxyLifespanState,
SpecialModelNames,
SupportedDBObjectType,
TeamDefaultSettings,
@ -792,6 +793,7 @@ from litellm.proxy.spend_tracking.spend_management_endpoints import (
router as spend_management_router,
)
from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload
from litellm.proxy.tracing_runtime import manage_tracing
from litellm.proxy.types_utils.utils import get_instance_fn
from litellm.proxy.ui_crud_endpoints.latest_release_endpoints import (
router as latest_release_endpoints_router,
@ -852,7 +854,6 @@ from litellm.secret_managers.main import (
secret_manager_would_be_consulted,
str_to_bool,
)
from litellm.tracing import TraceReceiver
from litellm.types.integrations.slack_alerting import AlertType, SlackAlertingArgs
from litellm.types.llms.anthropic import (
AnthropicMessagesRequest,
@ -1222,7 +1223,7 @@ async def _connect_to_count_stored_values() -> SupportsRawQueries:
@asynccontextmanager
async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]:
global \
prisma_client, \
master_key, \
@ -1529,9 +1530,6 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
_tagged.strategy._state_loaded = True
asyncio.create_task(_adaptive_router_flusher_loop())
## [Optional] Initialize agent tracing
asyncio.create_task(ProxyStartupEvent.init_tracing(general_settings))
## [Optional] Initialize dd tracer
ProxyStartupEvent._init_dd_tracer()
@ -1565,76 +1563,81 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
register_scheduled_sync(scheduler)
# End of startup event
yield
tracing_settings: Final = general_settings.get("tracing")
tracing_enabled: Final = TypeAdapter(bool).validate_python(
isinstance(tracing_settings, dict) and tracing_settings.get("store") == "clickhouse"
)
async with manage_tracing(enabled=tracing_enabled) as receiver:
state: Final[ProxyLifespanState] = {"tracing_receiver": receiver}
yield state
if model_info_scheduler is not None and model_info_scheduler.running:
model_info_scheduler.remove_job("refresh_model_info")
if model_info_scheduler is not scheduler:
model_info_scheduler.shutdown(wait=False)
if model_info_scheduler is not None and model_info_scheduler.running:
model_info_scheduler.remove_job("refresh_model_info")
if model_info_scheduler is not scheduler:
model_info_scheduler.shutdown(wait=False)
# Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window
if scheduler is not None:
pause_scheduled_jobs(scheduler)
# Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window
if scheduler is not None:
pause_scheduled_jobs(scheduler)
# Shutdown event - drain in-flight requests before tearing down dependencies
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
GracefulShutdownManager.start_shutdown()
await GracefulShutdownManager.wait_for_drain()
# Shutdown event - drain in-flight requests before tearing down dependencies
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
GracefulShutdownManager.start_shutdown()
await GracefulShutdownManager.wait_for_drain()
# Shutdown event - close shared aiohttp session
if shared_aiohttp_session is not None:
try:
await shared_aiohttp_session.close()
verbose_proxy_logger.info("SESSION REUSE: Closed shared aiohttp session")
except Exception as e:
verbose_proxy_logger.error("Error closing shared aiohttp session: %s", e)
# Shutdown event - close shared aiohttp session
if shared_aiohttp_session is not None:
try:
await shared_aiohttp_session.close()
verbose_proxy_logger.info("SESSION REUSE: Closed shared aiohttp session")
except Exception as e:
verbose_proxy_logger.error("Error closing shared aiohttp session: %s", e)
# Shutdown event - stop RDS IAM token refresh background task
if (
prisma_client is not None
and hasattr(prisma_client, "db")
and hasattr(prisma_client.db, "stop_token_refresh_task")
):
try:
await prisma_client.db.stop_token_refresh_task()
except Exception as e:
verbose_proxy_logger.error("Error stopping token refresh task: %s", e)
# Shutdown event - stop RDS IAM token refresh background task
if (
prisma_client is not None
and hasattr(prisma_client, "db")
and hasattr(prisma_client.db, "stop_token_refresh_task")
):
try:
await prisma_client.db.stop_token_refresh_task()
except Exception as e:
verbose_proxy_logger.error("Error stopping token refresh task: %s", e)
# Shutdown event - stop Prisma DB health watchdog task
if prisma_client is not None and hasattr(prisma_client, "stop_db_health_watchdog_task"):
try:
await prisma_client.stop_db_health_watchdog_task()
except Exception as e:
verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e)
# Shutdown event - stop Prisma DB health watchdog task
if prisma_client is not None and hasattr(prisma_client, "stop_db_health_watchdog_task"):
try:
await prisma_client.stop_db_health_watchdog_task()
except Exception as e:
verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e)
if prisma_client is not None and hasattr(prisma_client, "stop_view_setup_task"):
try:
await prisma_client.stop_view_setup_task()
except Exception as e:
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
if prisma_client is not None and hasattr(prisma_client, "stop_view_setup_task"):
try:
await prisma_client.stop_view_setup_task()
except Exception as e:
verbose_proxy_logger.error("Error stopping the spend view setup task: %s", e)
await _drain_spend_event_producer_on_shutdown()
await _drain_spend_event_producer_on_shutdown()
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
if scheduler is not None and scheduler_executor is not None:
try:
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
except Exception as e:
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
if scheduler is not None and scheduler_executor is not None:
try:
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
except Exception as e:
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
await flush_spend_counters_on_shutdown()
await flush_spend_counters_on_shutdown()
await _flush_spend_logs_queue_on_shutdown()
await _flush_spend_logs_queue_on_shutdown()
await proxy_config.stop_config_sync_subscriber()
await proxy_config.stop_config_sync_subscriber()
await proxy_config.stop_auth_cache_invalidation_subscriber()
await proxy_config.stop_auth_cache_invalidation_subscriber()
await proxy_shutdown_event(worker_heartbeat=worker_heartbeat)
await proxy_shutdown_event(worker_heartbeat=worker_heartbeat)
if prometheus_multiproc_dir:
mark_worker_exit(os.getpid())
if prometheus_multiproc_dir:
mark_worker_exit(os.getpid())
def _generate_stable_operation_id(route: "APIRoute") -> str:
@ -11357,39 +11360,6 @@ class ProxyStartupEvent:
)
return connected_client
@classmethod
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
"""
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 = receiver if receiver is not None else TraceReceiver.from_env()
await tracing.start()
except (KeyError, OSError, RuntimeError, ValueError) as error:
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
def _init_dd_tracer(cls):
"""

View file

@ -8,6 +8,7 @@ GET /v1/traces/{trace_id}/spans/{span_id} SpanDetail
"""
import time
from dataclasses import dataclass
from typing import Annotated, Final
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
@ -15,6 +16,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
from litellm.constants import OTLP_MAX_BODY_BYTES, OTLP_RETRY_AFTER_SECONDS
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.tracing_runtime import provide_receiver, require_receiver
from litellm.tracing import (
Tenant,
TraceReceiver,
@ -26,37 +28,42 @@ from litellm.tracing.types import SpanDetail, Trace, TracePage, TraceScope
router = APIRouter(tags=["agent tracing"])
MS_PER_DAY: Final = 24 * 60 * 60 * 1000
_ADMIN_ROLES: Final = (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
receiver: TraceReceiver | None = None
def get_receiver() -> TraceReceiver:
if receiver is None:
raise HTTPException(
status_code=501,
detail="Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL.",
)
return receiver
@dataclass(frozen=True, slots=True)
class TraceAccessContext:
receiver: TraceReceiver | None
read_scope: TraceScope | None
write_tenant: Tenant | None
def reader(self) -> tuple[TraceReceiver, TraceScope]:
tracing: Final = require_receiver(self.receiver)
if self.read_scope is None:
raise HTTPException(status_code=403, detail="Not allowed to view agent traces")
return tracing, self.read_scope
def writer(self) -> tuple[TraceReceiver, Tenant]:
if self.write_tenant is None:
raise HTTPException(status_code=403, detail="Not allowed to ingest agent traces")
return require_receiver(self.receiver), self.write_tenant
def tenant_for(user_api_key_dict: UserAPIKeyAuth) -> Tenant:
return Tenant(
team_id=user_api_key_dict.team_id or "",
api_key_hash=user_api_key_dict.token or "",
org_id=user_api_key_dict.org_id or "",
)
def scope_for(user_api_key_dict: UserAPIKeyAuth) -> TraceScope:
"""Admins see everything; team members see their team; team-less keys see their own traces."""
if user_api_key_dict.user_role in _ADMIN_ROLES:
return TraceScope(team_ids=(), api_key_hash="")
if user_api_key_dict.team_id:
return TraceScope(team_ids=(user_api_key_dict.team_id,), api_key_hash="")
if not user_api_key_dict.token:
raise HTTPException(status_code=403, detail="Not allowed to view agent traces")
return TraceScope(team_ids=("",), api_key_hash=user_api_key_dict.token)
async def provide_trace_access(
auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
tracing: Annotated[TraceReceiver | None, Depends(provide_receiver)],
) -> TraceAccessContext:
tenant: Final = Tenant(team_id=auth.team_id or "", api_key_hash=auth.token or "", org_id=auth.org_id or "")
match auth.user_role:
case LitellmUserRoles.PROXY_ADMIN:
return TraceAccessContext(tracing, TraceScope(team_ids=(), api_key_hash=""), tenant)
case LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY:
return TraceAccessContext(tracing, TraceScope(team_ids=(), api_key_hash=""), None)
case _ if auth.team_id:
return TraceAccessContext(tracing, TraceScope(team_ids=(auth.team_id,), api_key_hash=""), tenant)
case _ if auth.token:
return TraceAccessContext(tracing, TraceScope(team_ids=("",), api_key_hash=auth.token), tenant)
case _:
return TraceAccessContext(tracing, None, tenant)
async def _read_otlp_body(request: Request) -> bytes:
@ -71,18 +78,16 @@ async def _read_otlp_body(request: Request) -> bytes:
@router.post("/v1/traces", include_in_schema=False)
async def ingest_otlp_traces(
request: Request,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
context: Annotated[TraceAccessContext, Depends(provide_trace_access)],
) -> Response:
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY:
raise HTTPException(status_code=403, detail="Not allowed to ingest agent traces")
tracing: Final = get_receiver()
tracing, tenant = context.writer()
content_type: Final = request.headers.get("content-type")
try:
await tracing.ingest(
body=await _read_otlp_body(request),
content_type=content_type,
content_encoding=request.headers.get("content-encoding"),
tenant=tenant_for(user_api_key_dict),
tenant=tenant,
)
except TracingPayloadTooLargeError as e:
raise HTTPException(status_code=413, detail=str(e))
@ -99,15 +104,16 @@ async def ingest_otlp_traces(
@router.get("/v1/traces", response_model=None)
async def list_agent_traces(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
context: Annotated[TraceAccessContext, Depends(provide_trace_access)],
start_ms: Annotated[int | None, Query(description="Window start, unix ms. Default: 24h ago")] = None,
end_ms: Annotated[int | None, Query(description="Window end, unix ms. Default: now")] = None,
cursor: Annotated[str | None, Query()] = None,
) -> TracePage:
now_ms: Final = int(time.time() * 1000)
try:
return await get_receiver().list_traces(
scope=scope_for(user_api_key_dict),
tracing, scope = context.reader()
return await tracing.list_traces(
scope=scope,
start_ms=start_ms if start_ms is not None else now_ms - MS_PER_DAY,
end_ms=end_ms if end_ms is not None else now_ms,
cursor=cursor,
@ -119,10 +125,11 @@ async def list_agent_traces(
@router.get("/v1/traces/{trace_id}", response_model=None)
async def get_agent_trace(
trace_id: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
context: Annotated[TraceAccessContext, Depends(provide_trace_access)],
trace_ref: Annotated[str, Query()] = "",
) -> Trace:
trace: Final = await get_receiver().get_trace(trace_id, scope_for(user_api_key_dict), trace_ref)
tracing, scope = context.reader()
trace: Final = await tracing.get_trace(trace_id, scope, trace_ref)
if trace is None:
raise HTTPException(status_code=404, detail=f"Trace {trace_id} not found")
return trace
@ -132,10 +139,11 @@ async def get_agent_trace(
async def get_agent_trace_span(
trace_id: str,
span_id: str,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
context: Annotated[TraceAccessContext, Depends(provide_trace_access)],
trace_ref: Annotated[str, Query()] = "",
) -> SpanDetail:
span: Final = await get_receiver().get_span(trace_id, span_id, scope_for(user_api_key_dict), trace_ref)
tracing, scope = context.reader()
span: Final = await tracing.get_span(trace_id, span_id, scope, trace_ref)
if span is None:
raise HTTPException(status_code=404, detail=f"Span {span_id} not found")
return span

View file

@ -0,0 +1,66 @@
from collections.abc import AsyncGenerator, Callable
from contextlib import asynccontextmanager
from typing import Final
from fastapi import HTTPException, Request
from pydantic import ConfigDict, TypeAdapter
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger
from litellm.rust_bridge.traces import ClickHouseStorage
from litellm.tracing import TraceReceiver
_RECEIVER_ADAPTER: Final[TypeAdapter[TraceReceiver | None]] = TypeAdapter(
TraceReceiver | None, config=ConfigDict(arbitrary_types_allowed=True)
)
_UNAVAILABLE_DETAIL: Final = "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL."
def require_receiver(tracing: TraceReceiver | None) -> TraceReceiver:
if tracing is None:
raise HTTPException(status_code=501, detail=_UNAVAILABLE_DETAIL)
return tracing
async def provide_receiver(request: Request) -> TraceReceiver | None:
return _RECEIVER_ADAPTER.validate_python(getattr(request.state, "tracing_receiver", None))
async def provide_storage(request: Request) -> ClickHouseStorage | None:
tracing: Final = await provide_receiver(request)
return tracing.store.storage if tracing is not None else None
async def _start_receiver(factory: Callable[[], TraceReceiver]) -> TraceReceiver | None:
try:
tracing: Final = factory()
await tracing.start()
return tracing
except (KeyError, OSError, RuntimeError, ValueError) as error:
verbose_proxy_logger.warning("Agent tracing unavailable: %s", error)
return None
@asynccontextmanager
async def manage_tracing(
enabled: bool, receiver_factory: Callable[[], TraceReceiver] = TraceReceiver.from_env
) -> AsyncGenerator[TraceReceiver | None, None]:
tracing: Final = await _start_receiver(receiver_factory) if enabled else None
if tracing is None:
yield tracing
return
spend_logger: Final = ClickHouseSpendLogger(storage=tracing.store.storage)
manager: Final = litellm.logging_callback_manager
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)")
try:
yield tracing
finally:
manager.remove_callback_from_all_lists(spend_logger)
await spend_logger.aclose()

View file

@ -80,7 +80,7 @@ def decode_otlp(
return _native().trace_decode_otlp(body, content_type, content_encoding, max_decompressed_bytes)
class TraceStorage:
class ClickHouseStorage:
def __init__(self, database: str, url: str, reader_url: str | None = None) -> None:
self._native: Final = _native().NativeTraceStorage(database, url, reader_url)

View file

@ -1,6 +1,6 @@
- Python owns tracing endpoints, authenticated tenant scope, framework normalization and API response shaping
- Trace ingestion awaits `TraceStorage.insert_rows` before returning success; propagate storage failures so OTLP exporters can retry
- Trace ingestion awaits `ClickHouseStorage.insert_rows` before returning success; propagate storage failures so OTLP exporters can retry
- Spend logging keeps its separate batch queue in `litellm/integrations/clickhouse`
- Use `litellm.rust_bridge.traces.TraceStorage` for ClickHouse; keep schema, SQL, encoding and transport in `litellm-traces`
- Use `litellm.rust_bridge.traces.ClickHouseStorage` for ClickHouse; keep trace schema, SQL and encoding in `litellm-traces`, and generic transport in `litellm-storage-clickhouse`
- Derive tenant fields from authentication and overwrite matching fields supplied by the exporter
- Test confirmed writes, failures, tenant isolation and read behavior through public functions

View file

@ -23,9 +23,9 @@ from litellm.constants import (
OTLP_OFFLOAD_DECODE_BYTES,
)
from litellm.integrations.clickhouse.schema import ensure_schema
from litellm.rust_bridge.traces import TraceStorage
from litellm.rust_bridge.traces import ClickHouseStorage
from litellm.tracing.decode import OTLPPayloadTooLargeError, decode_otlp
from litellm.tracing.store import ClickHouseTraceStore
from litellm.tracing.store import TraceStore
from litellm.tracing.types import (
SpanDetail,
SpanRow,
@ -60,14 +60,14 @@ class Tenant:
class TraceReceiver:
def __init__(self, store: ClickHouseTraceStore) -> None:
def __init__(self, store: TraceStore) -> None:
self.store = store
@classmethod
def from_env(cls) -> "TraceReceiver":
return cls(
store=ClickHouseTraceStore(
TraceStorage(
store=TraceStore(
ClickHouseStorage(
database=os.getenv("CLICKHOUSE_DATABASE", "litellm"),
url=os.environ["CLICKHOUSE_URL"],
reader_url=os.environ["CLICKHOUSE_READER_URL"],

View file

@ -16,7 +16,7 @@ from litellm.constants import AGENT_TRACING_LIST_PAGE_SIZE
from litellm.integrations.clickhouse.schema import (
OTEL_TRACES_TABLE,
)
from litellm.rust_bridge.traces import TraceStorage
from litellm.rust_bridge.traces import ClickHouseStorage
from litellm.tracing.types import (
AgentNode,
Span,
@ -259,10 +259,10 @@ def trace_from_rows(
)
class ClickHouseTraceStore:
class TraceStore:
"""Stores spans and runs scoped trace reads."""
def __init__(self, storage: TraceStorage) -> None:
def __init__(self, storage: ClickHouseStorage) -> None:
self.storage = storage
async def insert_spans(self, rows: Sequence[SpanRow]) -> None:

View file

@ -82,7 +82,7 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
)
assert stored_worker is not None and stored_worker.id == worker.id
assert worker.id == registration.worker.id
listing: Final = await endpoints.list_lenses(admin)
listing: Final = await endpoints.list_lenses(admin, storage=None)
assert lens.id in tuple(e.id for e in listing.lenses)
assert worker.id in tuple(w.id for w in listing.workers)
claims: Final = await asyncio.gather(
@ -180,13 +180,13 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
assert needs_billing.value.status_code == 409
assert await endpoints.heartbeat(lens.id, claimed.job.id, authenticated_legacy)
finished: Final = await endpoints.result(
lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy
lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy, storage=None
)
assert finished.jobs[0].status == "completed"
assert finished.jobs[0].coverage.screened == 2
assert finished.last_scan_at == claimed.job.end
assert finished.next_run_at > finished.jobs[0].finished_at
assert await endpoints.result(lens.id, claimed.job.id, Result(coverage=Coverage()), worker) == finished
assert await endpoints.result(lens.id, claimed.job.id, Result(coverage=Coverage()), worker, storage=None) == finished
with pytest.raises(HTTPException) as stale:
await endpoints.heartbeat(lens.id, claimed.job.id, worker)
assert stale.value.status_code == 409

View file

@ -312,7 +312,7 @@ def test_db_push_timeout_hint_names_the_per_command_budget(
) -> None:
"""``db push`` keeps the per-command budget, so its timeout hint has to name that variable."""
_, log_path = toolchain_env
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.delenv("DATABASE_URL", raising=False)
monkeypatch.setenv(PRISMA_COMMAND_TIMEOUT_ENV_VAR, "1")
monkeypatch.setenv("FAKE_PRISMA_FIRST_PUSH_SLEEP", "3")

View file

@ -3,6 +3,8 @@ Tests for the CustomBatchLogger-based ClickHouse base logger.
"""
import asyncio
from collections.abc import Mapping, Sequence
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -50,8 +52,7 @@ async def test_first_enqueued_row_flushes_after_synchronous_construction():
logger.enqueue([{"i": 1}])
await asyncio.wait_for(flushed.wait(), timeout=1)
if logger._flush_task is not None:
logger._flush_task.cancel()
await logger.aclose()
@pytest.mark.asyncio
@ -79,3 +80,63 @@ async def test_failed_insert_is_requeued_then_dropped():
assert logger.rows_dropped == 2
assert logger.rows_written == 0
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_close_waits_for_active_insert_and_stops_periodic_flush() -> None:
started: Final = asyncio.Event()
release: Final = asyncio.Event()
async def insert_rows(table: str, rows: Sequence[Mapping[str, object]]) -> None:
started.set()
await release.wait()
insert: Final = AsyncMock(side_effect=insert_rows)
logger: Final = _logger(insert)
logger.flush_interval = 0.001
logger.enqueue([{"i": 1}])
await asyncio.wait_for(started.wait(), timeout=1)
closing: Final = asyncio.create_task(logger.aclose())
await asyncio.sleep(0)
assert not closing.done()
release.set()
await asyncio.wait_for(closing, timeout=1)
assert logger.rows_written == 1
insert.assert_awaited_once_with("test_table", [{"i": 1}])
assert logger._flush_task is not None and logger._flush_task.done()
assert not logger._flush_task.cancelled()
@pytest.mark.asyncio
async def test_close_wakes_idle_worker_and_drains_queued_rows() -> None:
insert: Final = AsyncMock()
logger: Final = _logger(insert)
logger.flush_interval = 3600
logger.enqueue([{"i": 1}])
await asyncio.sleep(0)
await asyncio.wait_for(logger.aclose(), timeout=1)
insert.assert_awaited_once_with("test_table", [{"i": 1}])
assert logger.rows_written == 1
assert logger.log_queue == []
assert logger._flush_task is not None and logger._flush_task.done()
assert not logger._flush_task.cancelled()
@pytest.mark.asyncio
@pytest.mark.parametrize("recovers", [True, False])
async def test_close_retries_every_batch_and_accounts_for_exhausted_rows(recovers: bool) -> None:
failure: Final = RuntimeError("ClickHouse unavailable")
insert: Final = AsyncMock(side_effect=[failure, None, None] if recovers else failure)
logger: Final = _logger(insert)
logger.batch_size = 1
logger.log_queue.extend([{"request_id": "a"}, {"request_id": "b"}])
await logger.aclose()
assert logger.log_queue == []
assert logger.rows_written == (2 if recovers else 0)
assert logger.rows_dropped == (0 if recovers else 2)
assert insert.await_count == (3 if recovers else 2 * module.CLICKHOUSE_MAX_RETRIES)
assert {call.args[1][0]["request_id"] for call in insert.await_args_list} == {"a", "b"}

View file

@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.tracing.store import (
ClickHouseTraceStore,
TraceStore,
agent_nodes,
decode_cursor,
encode_cursor,
@ -284,7 +284,7 @@ async def test_list_traces_sets_next_cursor_on_full_page():
"models": [],
}
client.query = AsyncMock(return_value=[row, {**row, "trace_id": "t1", "trace_ref": "ref1", "start_ms": 900}])
store = ClickHouseTraceStore(client)
store = TraceStore(client)
scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""}
page = await store.list_traces(scope, 0, 2000, limit=2)
@ -303,7 +303,7 @@ async def test_list_traces_sets_next_cursor_on_full_page():
async def test_get_span_not_found_and_found():
client = MagicMock()
client.query = AsyncMock(return_value=[])
store = ClickHouseTraceStore(client)
store = TraceStore(client)
scope: TraceScope = {"team_ids": (), "api_key_hash": ""}
assert await store.get_span("t", "s", scope) is None
stored_input = '[{"role": "user", "content": "hi"}]'
@ -355,7 +355,7 @@ async def test_trace_cost_is_scoped_and_counts_repeated_request_once():
},
]
client.query = AsyncMock(side_effect=[spans, spend])
store = ClickHouseTraceStore(client)
store = TraceStore(client)
scope: TraceScope = {"team_ids": ("team-a",), "api_key_hash": ""}
trace = await store.get_trace("trace-1", scope)
@ -406,7 +406,7 @@ async def test_run_list_uses_matching_spend_and_leaves_missing_cost_unavailable(
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)
page = await TraceStore(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"]
@ -428,7 +428,7 @@ async def test_ambiguous_cache_response_id_keeps_cost_unavailable():
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)
store = TraceStore(client)
scope: TraceScope = {"team_ids": ("",), "api_key_hash": "key-a"}
trace = await store.get_trace("trace-1", scope)

View file

@ -14,6 +14,7 @@ import logging
import os
import re
from collections.abc import Mapping
from contextlib import nullcontext
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
@ -44,48 +45,51 @@ from .conftest import normalize
@pytest.mark.asyncio
async def test_tracing_config_automatically_logs_spend_without_callback_setting():
@pytest.mark.parametrize("shutdown_error", [False, True])
async def test_tracing_config_automatically_logs_spend_without_callback_setting(shutdown_error: bool) -> None:
from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger
from litellm.proxy import tracing_endpoints
from litellm.proxy.proxy_server import ProxyStartupEvent
from litellm.proxy.tracing_runtime import manage_tracing
from litellm.tracing import TraceReceiver
from litellm.tracing.store import ClickHouseTraceStore
from litellm.tracing.store import TraceStore
storage = MagicMock()
storage: Final = MagicMock()
storage.ensure_schema = AsyncMock()
storage.insert_rows = AsyncMock()
receiver = TraceReceiver(ClickHouseTraceStore(storage))
prior_receiver = tracing_endpoints.receiver
receiver: Final = TraceReceiver(TraceStore(storage))
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
outcome: Final = pytest.raises(RuntimeError, match="shutdown failure") if shutdown_error else nullcontext()
with outcome:
async with manage_tracing(enabled=True, receiver_factory=lambda: receiver):
storage.ensure_schema.assert_awaited_once()
logger: Final = next(
callback
for callback in litellm._async_success_callback
if isinstance(callback, ClickHouseSpendLogger) and callback.storage is storage
)
now: Final = 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,
)
storage.insert_rows.assert_not_awaited()
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
if shutdown_error:
raise RuntimeError("shutdown failure")
assert storage.insert_rows.await_args.args[0] == "spend_logs"
assert storage.insert_rows.await_args.args[1][0]["spend"] == 0.25
assert logger not in litellm._async_success_callback
assert logger._flush_task is not None and logger._flush_task.done()
assert not logger._flush_task.cancelled()
# ---------------------------------------------------------------------------

View file

@ -2,6 +2,9 @@
Tests for the agent tracing endpoints (litellm/proxy/tracing_endpoints.py).
"""
from collections.abc import AsyncGenerator
from contextlib import asynccontextmanager
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -9,58 +12,79 @@ from fastapi import FastAPI, HTTPException
from fastapi.testclient import TestClient
from litellm.proxy import tracing_endpoints
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy._types import LitellmUserRoles, ProxyLifespanState, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.tracing_runtime import manage_tracing, provide_storage
from litellm.rust_bridge.traces import ClickHouseStorage
from litellm.tracing import TraceReceiver, TracingPayloadTooLargeError
from litellm.tracing.store import ClickHouseTraceStore
from litellm.tracing.store import TraceStore
from litellm.tracing.types import TraceScope
TEAM_KEY = UserAPIKeyAuth(
token="hashed-key", team_id="team-research", org_id="org-1", user_role=LitellmUserRoles.INTERNAL_USER
)
# ---------------------------------------------------------------- scope / tenant
@pytest.mark.parametrize(
("auth", "scope", "can_write"),
(
pytest.param(
UserAPIKeyAuth(token="admin-key", team_id="team-a", user_role=LitellmUserRoles.PROXY_ADMIN),
TraceScope(team_ids=(), api_key_hash=""),
True,
id="admin",
),
pytest.param(
UserAPIKeyAuth(token="view-key", team_id="team-a", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY),
TraceScope(team_ids=(), api_key_hash=""),
False,
id="view-only-admin",
),
pytest.param(
TEAM_KEY,
TraceScope(team_ids=("team-research",), api_key_hash=""),
True,
id="team-key",
),
pytest.param(
UserAPIKeyAuth(token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER),
TraceScope(team_ids=("",), api_key_hash="hashed-key"),
True,
id="teamless-key",
),
),
)
def test_trace_read_and_write_permissions(
client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth, scope: TraceScope, can_write: bool
) -> None:
client.app.dependency_overrides[user_api_key_auth] = lambda: auth
read: Final = client.get("/v1/traces?start_ms=1&end_ms=2")
assert read.status_code == 200, read.text
receiver.list_traces.assert_awaited_once_with(scope=scope, start_ms=1, end_ms=2, cursor=None)
def test_scope_for_admin_sees_everything():
for role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY):
auth = UserAPIKeyAuth(token="k", team_id="team-a", user_role=role)
assert tracing_endpoints.scope_for(auth) == {"team_ids": (), "api_key_hash": ""}
def test_scope_for_team_key_sees_its_team():
assert tracing_endpoints.scope_for(TEAM_KEY) == {"team_ids": ("team-research",), "api_key_hash": ""}
def test_scope_for_teamless_key_sees_only_its_own_traces():
auth = UserAPIKeyAuth(token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER)
assert tracing_endpoints.scope_for(auth) == {"team_ids": ("",), "api_key_hash": "hashed-key"}
def test_scope_for_no_team_no_token_is_forbidden():
with pytest.raises(HTTPException) as e:
tracing_endpoints.scope_for(UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER))
assert e.value.status_code == 403
def test_tenant_for_comes_from_auth():
tenant = tracing_endpoints.tenant_for(TEAM_KEY)
assert (tenant.team_id, tenant.api_key_hash, tenant.org_id) == ("team-research", "hashed-key", "org-1")
blank = tracing_endpoints.tenant_for(UserAPIKeyAuth())
assert (blank.team_id, blank.api_key_hash, blank.org_id) == ("", "", "")
# ---------------------------------------------------------------- endpoints
write: Final = client.post("/v1/traces", json={})
assert write.status_code == (200 if can_write else 403), write.text
if not can_write:
receiver.ingest.assert_not_awaited()
return
receiver.ingest.assert_awaited_once()
tenant: Final = receiver.ingest.await_args.kwargs["tenant"]
assert (tenant.team_id, tenant.api_key_hash, tenant.org_id) == (
auth.team_id or "",
auth.token or "",
auth.org_id or "",
)
@pytest.fixture
def receiver(monkeypatch) -> MagicMock:
def receiver(client) -> MagicMock:
fake = MagicMock()
fake.ingest = AsyncMock(return_value=1)
fake.list_traces = AsyncMock(return_value={"data": [], "next_cursor": None})
fake.get_trace = AsyncMock(return_value=None)
fake.get_span = AsyncMock(return_value=None)
monkeypatch.setattr(tracing_endpoints, "receiver", fake)
client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: fake
return fake
@ -72,8 +96,7 @@ def client() -> TestClient:
return TestClient(app)
def test_501_when_tracing_not_enabled(client, monkeypatch):
monkeypatch.setattr(tracing_endpoints, "receiver", None)
def test_501_when_tracing_not_enabled(client):
assert client.post("/v1/traces", content=b"").status_code == 501
assert client.get("/v1/traces").status_code == 501
@ -149,13 +172,13 @@ def test_get_span_404_and_200(client, receiver):
receiver.get_span.assert_awaited_with("t1", "s1", {"team_ids": ("team-research",), "api_key_hash": ""}, "")
def test_get_span_serves_ui_content_from_stored_payloads(client, monkeypatch):
def test_get_span_serves_ui_content_from_stored_payloads(client):
storage = MagicMock()
stored_output = '{"role": "ai", "content": "", "tool_calls": [{"name": "lookup", "args": {"id": 7}}]}'
storage.query = AsyncMock(
return_value=[{"span_id": "s1", "input": '{"city": "Paris"}', "output": stored_output, "attributes": {}}]
)
monkeypatch.setattr(tracing_endpoints, "receiver", TraceReceiver(ClickHouseTraceStore(storage)))
client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: TraceReceiver(TraceStore(storage))
body = client.get("/v1/traces/t1/spans/s1").json()
assert body["output"] == stored_output
assert body["input_ui"] == {"kind": "fields", "fields": [{"key": "city", "value": "Paris"}]}
@ -197,3 +220,259 @@ def test_view_only_admin_cannot_ingest_traces(client, receiver):
response = client.post("/v1/traces", content=b"{}")
assert response.status_code == 403
receiver.ingest.assert_not_called()
@pytest.mark.parametrize("status_code", [401, 403])
def test_auth_failure_precedes_disabled_receiver(client: TestClient, status_code: int) -> None:
def unavailable() -> None:
return None
def authenticate() -> UserAPIKeyAuth:
if status_code == 401:
raise HTTPException(status_code=401, detail="Invalid API key")
return UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
client.app.dependency_overrides[user_api_key_auth] = authenticate
client.app.dependency_overrides[tracing_endpoints.provide_receiver] = unavailable
response: Final = client.post("/v1/traces", content=b"{}")
assert response.status_code == status_code
assert response.json() == {
"detail": "Invalid API key" if status_code == 401 else "Not allowed to ingest agent traces"
}
def test_disabled_receiver_precedes_read_scope_rejection(client: TestClient) -> None:
client.app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
user_role=LitellmUserRoles.INTERNAL_USER
)
response: Final = client.get("/v1/traces")
assert response.status_code == 501
assert response.json() == {
"detail": "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL."
}
@pytest.mark.requires_rust_extension
def test_injected_receiver_persists_authenticated_tenant(client: TestClient) -> None:
storage: Final = MagicMock(spec=ClickHouseStorage)
storage.insert_rows = AsyncMock()
tracing: Final = TraceReceiver(TraceStore(storage))
client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: tracing
response: Final = client.post(
"/v1/traces",
json={
"resourceSpans": [
{
"resource": {
"attributes": [
{"key": "litellm.team_id", "value": {"stringValue": "spoofed-team"}},
{"key": "litellm.api_key_hash", "value": {"stringValue": "spoofed-key"}},
{"key": "litellm.org_id", "value": {"stringValue": "spoofed-org"}},
]
},
"scopeSpans": [
{
"spans": [
{
"traceId": "01" * 16,
"spanId": "02" * 8,
"name": "dependency-injection",
"startTimeUnixNano": "1000000000",
"endTimeUnixNano": "1000000001",
}
]
}
],
}
],
},
)
assert response.status_code == 200, response.text
assert response.json() == {}
storage.insert_rows.assert_awaited_once()
table, rows = storage.insert_rows.await_args.args
assert table == "otel_traces"
assert len(rows) == 1
assert rows[0]["TeamId"] == TEAM_KEY.team_id
assert rows[0]["ApiKeyHash"] == TEAM_KEY.token
assert rows[0]["ResourceAttributes"] == {
"litellm.team_id": TEAM_KEY.team_id,
"litellm.api_key_hash": TEAM_KEY.token,
"litellm.org_id": TEAM_KEY.org_id,
}
def test_lifespan_receivers_are_app_local() -> None:
first_storage: Final = MagicMock(spec=ClickHouseStorage)
first_storage.query = AsyncMock(
return_value=[
{
"span_id": "first-span",
"input": "first-input",
"output": "",
"attributes": {},
}
]
)
second_storage: Final = MagicMock(spec=ClickHouseStorage)
second_storage.query = AsyncMock(
return_value=[
{
"span_id": "second-span",
"input": "second-input",
"output": "",
"attributes": {},
}
]
)
first_receiver: Final = TraceReceiver(TraceStore(first_storage))
second_receiver: Final = TraceReceiver(TraceStore(second_storage))
first_storage.ensure_schema = AsyncMock()
second_storage.ensure_schema = AsyncMock()
@asynccontextmanager
async def first_lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]:
async with manage_tracing(True, lambda: first_receiver) as receiver:
state: Final[ProxyLifespanState] = {"tracing_receiver": receiver}
yield state
@asynccontextmanager
async def second_lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]:
async with manage_tracing(True, lambda: second_receiver) as receiver:
state: Final[ProxyLifespanState] = {"tracing_receiver": receiver}
yield state
first_app: Final = FastAPI(lifespan=first_lifespan)
second_app: Final = FastAPI(lifespan=second_lifespan)
first_app.include_router(tracing_endpoints.router)
second_app.include_router(tracing_endpoints.router)
first_app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY
second_app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY
with TestClient(first_app) as first_client:
with TestClient(second_app) as second_client:
second_response: Final = second_client.get("/v1/traces/t1/spans/second-span?trace_ref=second-run")
simultaneous: Final = first_client.get("/v1/traces/t1/spans/first-span?trace_ref=first-run")
first_response: Final = first_client.get("/v1/traces/t1/spans/first-span?trace_ref=first-run")
assert simultaneous.json() == first_response.json()
first_storage.ensure_schema.assert_awaited_once()
second_storage.ensure_schema.assert_awaited_once()
assert first_response.status_code == second_response.status_code == 200
assert first_response.json() == {
"span_id": "first-span",
"input": "first-input",
"output": "",
"attributes": {},
"input_ui": {"kind": "text", "text": "first-input"},
"output_ui": {"kind": "text", "text": ""},
}
assert second_response.json() == {
"span_id": "second-span",
"input": "second-input",
"output": "",
"attributes": {},
"input_ui": {"kind": "text", "text": "second-input"},
"output_ui": {"kind": "text", "text": ""},
}
assert first_storage.query.await_count == 2
first_storage.query.assert_awaited_with(
"span_detail",
{
"team_ids": (TEAM_KEY.team_id,),
"api_key_hash": "",
"trace_id": "t1",
"span_id": "first-span",
"trace_ref": "first-run",
},
)
second_storage.query.assert_awaited_once_with(
"span_detail",
{
"team_ids": (TEAM_KEY.team_id,),
"api_key_hash": "",
"trace_id": "t1",
"span_id": "second-span",
"trace_ref": "second-run",
},
)
@pytest.mark.parametrize("auth", [TEAM_KEY, UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER)])
def test_query_validation_precedes_trace_access_checks(client: TestClient, auth: UserAPIKeyAuth) -> None:
client.app.dependency_overrides[user_api_key_auth] = lambda: auth
response: Final = client.get("/v1/traces", params={"start_ms": "invalid"})
assert response.status_code == 422
assert response.json()["detail"][0]["loc"] == ["query", "start_ms"]
@pytest.mark.parametrize("enabled", [True, False])
def test_unavailable_lifespan_receiver_returns_501(enabled: bool) -> None:
storage: Final = MagicMock(spec=ClickHouseStorage)
storage.ensure_schema = AsyncMock(side_effect=RuntimeError("storage unavailable"))
tracing: Final = TraceReceiver(TraceStore(storage))
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]:
async with manage_tracing(enabled, lambda: tracing) as receiver:
state: Final[ProxyLifespanState] = {"tracing_receiver": receiver}
yield state
app: Final = FastAPI(lifespan=lifespan)
app.include_router(tracing_endpoints.router)
app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY
with TestClient(app) as client:
response: Final = client.get("/v1/traces")
assert response.status_code == 501
assert storage.ensure_schema.await_count == int(enabled)
storage.query.assert_not_called()
def test_lens_reads_from_the_lifespan_storage() -> None:
from litellm.proxy.lens.endpoints import router as lens_router
storage: Final = MagicMock(spec=ClickHouseStorage)
storage.ensure_schema = AsyncMock()
storage.lens_sample = AsyncMock(return_value=[])
tracing: Final = TraceReceiver(TraceStore(storage))
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]:
async with manage_tracing(True, lambda: tracing) as receiver:
state: Final[ProxyLifespanState] = {"tracing_receiver": receiver}
yield state
app: Final = FastAPI(lifespan=lifespan)
app.include_router(lens_router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
with TestClient(app) as client:
response: Final = client.post(
"/lens/preview/sample",
json={"settings": {"name": "Review", "model": "analysis", "context": "Find failed executions"}},
)
assert response.status_code == 200, response.text
assert response.json()["executions"] == []
storage.lens_sample.assert_awaited_once()
assert storage.lens_sample.await_args.args[0]["all_teams"] == 1
def test_lens_reads_from_injected_storage_without_receiver() -> None:
from litellm.proxy.lens.endpoints import router as lens_router
from litellm.proxy.lens.sources import Storage
storage: Final = MagicMock(spec=Storage)
storage.lens_sample = AsyncMock(return_value=[])
app: Final = FastAPI()
app.include_router(lens_router)
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
app.dependency_overrides[provide_storage] = lambda: storage
with TestClient(app) as client:
response: Final = client.post(
"/lens/preview/sample",
json={"settings": {"name": "Review", "model": "analysis", "context": "Find failed executions"}},
)
assert response.status_code == 200, response.text
assert response.json()["executions"] == []
storage.lens_sample.assert_awaited_once()