From be67fce26a19669e0696082ec3d6395cbdbcc703 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Thu, 1 Oct 2026 13:45:32 -0700 Subject: [PATCH] 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 --- litellm-rust/Cargo.lock | 17 +- litellm-rust/Cargo.toml | 1 + litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../crates/python-bridge/src/routes/traces.rs | 26 +- .../crates/storage-clickhouse/Cargo.toml | 20 + .../crates/storage-clickhouse/README.md | 5 + .../crates/storage-clickhouse/src/error.rs | 29 ++ .../crates/storage-clickhouse/src/insert.rs | 70 ++++ .../crates/storage-clickhouse/src/lib.rs | 127 +++++++ .../crates/storage-clickhouse/src/read.rs | 113 ++++++ .../storage-clickhouse/tests/connection.rs | 34 ++ .../storage-clickhouse/tests/transport.rs | 34 ++ litellm-rust/crates/traces/AGENTS.md | 2 +- litellm-rust/crates/traces/Cargo.toml | 2 +- litellm-rust/crates/traces/src/error.rs | 30 -- litellm-rust/crates/traces/src/insert.rs | 62 +-- litellm-rust/crates/traces/src/lib.rs | 84 +---- litellm-rust/crates/traces/src/sql.rs | 113 +----- litellm-rust/crates/traces/tests/queries.rs | 11 - .../clickhouse/clickhouse_batch_logger.py | 27 +- litellm/integrations/clickhouse/schema.py | 4 +- litellm/proxy/_types.py | 6 + litellm/proxy/lens/endpoints.py | 49 ++- litellm/proxy/proxy_server.py | 156 ++++---- litellm/proxy/tracing_endpoints.py | 86 +++-- litellm/proxy/tracing_runtime.py | 66 ++++ litellm/rust_bridge/traces.py | 2 +- litellm/tracing/AGENTS.md | 4 +- litellm/tracing/receiver.py | 10 +- litellm/tracing/store.py | 6 +- tests/proxy_behavior/lens/test_lifecycle.py | 6 +- .../test_prisma_toolchain.py | 2 +- .../test_clickhouse_batch_logger.py | 65 +++- tests/test_litellm/tracing/test_store.py | 12 +- .../proxy/proxy_server/test_proxy_config.py | 74 ++-- tests/unit/proxy/test_tracing_endpoints.py | 355 ++++++++++++++++-- 36 files changed, 1152 insertions(+), 559 deletions(-) create mode 100644 litellm-rust/crates/storage-clickhouse/Cargo.toml create mode 100644 litellm-rust/crates/storage-clickhouse/README.md create mode 100644 litellm-rust/crates/storage-clickhouse/src/error.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/insert.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/lib.rs create mode 100644 litellm-rust/crates/storage-clickhouse/src/read.rs create mode 100644 litellm-rust/crates/storage-clickhouse/tests/connection.rs create mode 100644 litellm-rust/crates/storage-clickhouse/tests/transport.rs delete mode 100644 litellm-rust/crates/traces/tests/queries.rs create mode 100644 litellm/proxy/tracing_runtime.py diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0ad05d99e76..57e9e803e4a 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 257a47268e4..450253ea768 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -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" } diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 99c95632bb3..a1d1f63d6f3 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 2e7a6b178a8..47c924f4842 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -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, + storage: Storage, } #[pymethods] @@ -39,12 +38,7 @@ impl NativeTraceStorage { fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult { 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> { 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> { 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> { 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> { 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)?; diff --git a/litellm-rust/crates/storage-clickhouse/Cargo.toml b/litellm-rust/crates/storage-clickhouse/Cargo.toml new file mode 100644 index 00000000000..f7f85c0dd8d --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/storage-clickhouse/README.md b/litellm-rust/crates/storage-clickhouse/README.md new file mode 100644 index 00000000000..7c4f86e4589 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/README.md @@ -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 diff --git a/litellm-rust/crates/storage-clickhouse/src/error.rs b/litellm-rust/crates/storage-clickhouse/src/error.rs new file mode 100644 index 00000000000..d283fb9021e --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/error.rs @@ -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, +} diff --git a/litellm-rust/crates/storage-clickhouse/src/insert.rs b/litellm-rust/crates/storage-clickhouse/src/insert.rs new file mode 100644 index 00000000000..3be73b086a0 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/insert.rs @@ -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(()) +} diff --git a/litellm-rust/crates/storage-clickhouse/src/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs new file mode 100644 index 00000000000..f2b34eddbf8 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -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 { + 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 { + 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 { + 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 { + 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, +} + +impl Storage { + pub fn new(database: String, url: &str, reader_url: Option<&str>) -> Result { + 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'_') +} diff --git a/litellm-rust/crates/storage-clickhouse/src/read.rs b/litellm-rust/crates/storage-clickhouse/src/read.rs new file mode 100644 index 00000000000..99c6a5120f3 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/src/read.rs @@ -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), +} + +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::>() + .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, +) -> Result { + 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) +} diff --git a/litellm-rust/crates/storage-clickhouse/tests/connection.rs b/litellm-rust/crates/storage-clickhouse/tests/connection.rs new file mode 100644 index 00000000000..0874b693249 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/tests/connection.rs @@ -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()); +} diff --git a/litellm-rust/crates/storage-clickhouse/tests/transport.rs b/litellm-rust/crates/storage-clickhouse/tests/transport.rs new file mode 100644 index 00000000000..0c7ea233522 --- /dev/null +++ b/litellm-rust/crates/storage-clickhouse/tests/transport.rs @@ -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) + )); +} diff --git a/litellm-rust/crates/traces/AGENTS.md b/litellm-rust/crates/traces/AGENTS.md index a5e2d4be53a..645e88dfae1 100644 --- a/litellm-rust/crates/traces/AGENTS.md +++ b/litellm-rust/crates/traces/AGENTS.md @@ -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 diff --git a/litellm-rust/crates/traces/Cargo.toml b/litellm-rust/crates/traces/Cargo.toml index 7d5facaa71e..b7f6e6ae52e 100644 --- a/litellm-rust/crates/traces/Cargo.toml +++ b/litellm-rust/crates/traces/Cargo.toml @@ -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"] } diff --git a/litellm-rust/crates/traces/src/error.rs b/litellm-rust/crates/traces/src/error.rs index 4a4fdaa00f7..2ccfe0ea8d9 100644 --- a/litellm-rust/crates/traces/src/error.rs +++ b/litellm-rust/crates/traces/src/error.rs @@ -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")] diff --git a/litellm-rust/crates/traces/src/insert.rs b/litellm-rust/crates/traces/src/insert.rs index bbee66f6fa5..6d9a2cab813 100644 --- a/litellm-rust/crates/traces/src/insert.rs +++ b/litellm-rust/crates/traces/src/insert.rs @@ -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>) -> Result { diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index c37602cade4..f5defb36cc2 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -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 { - 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 { - 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 { - 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 { - 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}; diff --git a/litellm-rust/crates/traces/src/sql.rs b/litellm-rust/crates/traces/src/sql.rs index 8346e06cb71..9acb8de0a7a 100644 --- a/litellm-rust/crates/traces/src/sql.rs +++ b/litellm-rust/crates/traces/src/sql.rs @@ -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), -} - -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::>() - .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, -) -> Result { - 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, diff --git a/litellm-rust/crates/traces/tests/queries.rs b/litellm-rust/crates/traces/tests/queries.rs deleted file mode 100644 index 75dfe0adc19..00000000000 --- a/litellm-rust/crates/traces/tests/queries.rs +++ /dev/null @@ -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); -} diff --git a/litellm/integrations/clickhouse/clickhouse_batch_logger.py b/litellm/integrations/clickhouse/clickhouse_batch_logger.py index 81601ea2a78..ac782ffebb2 100644 --- a/litellm/integrations/clickhouse/clickhouse_batch_logger.py +++ b/litellm/integrations/clickhouse/clickhouse_batch_logger.py @@ -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 diff --git a/litellm/integrations/clickhouse/schema.py b/litellm/integrations/clickhouse/schema.py index 6bec35c5630..5bf2b21cda5 100644 --- a/litellm/integrations/clickhouse/schema.py +++ b/litellm/integrations/clickhouse/schema.py @@ -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) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2da31d9f0fd..a471fb6f6f8 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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. diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index f24715d8265..0349c594adf 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ff0dd9df16f..cf09bdbef9b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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): """ diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index b2b10ebfbc6..bd885282859 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -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 diff --git a/litellm/proxy/tracing_runtime.py b/litellm/proxy/tracing_runtime.py new file mode 100644 index 00000000000..0b706d66a40 --- /dev/null +++ b/litellm/proxy/tracing_runtime.py @@ -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() diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index 98607aa9206..1c20e408709 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -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) diff --git a/litellm/tracing/AGENTS.md b/litellm/tracing/AGENTS.md index f69866c0419..ee1c8870edd 100644 --- a/litellm/tracing/AGENTS.md +++ b/litellm/tracing/AGENTS.md @@ -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 diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 05d9dcb3307..700f33a8a0e 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -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"], diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index fdb1c7820f2..9d1f64f77f0 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -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: diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index 3196b68bd83..849b2186a62 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -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 diff --git a/tests/proxy_migration_tests/test_prisma_toolchain.py b/tests/proxy_migration_tests/test_prisma_toolchain.py index 556c680a84a..ebe2390db16 100644 --- a/tests/proxy_migration_tests/test_prisma_toolchain.py +++ b/tests/proxy_migration_tests/test_prisma_toolchain.py @@ -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") 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 bae94ba6100..5eb14e73855 100644 --- a/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py +++ b/tests/test_litellm/integrations/clickhouse/test_clickhouse_batch_logger.py @@ -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"} diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 2e80594ff9b..30d7a9b5b0b 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -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) diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index c3709ceae3f..9785bdd5e32 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -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() # --------------------------------------------------------------------------- diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 6391c1577f3..2e34172acfd 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -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()