mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
refactor(proxy): inject tracing receiver and access context (#44035)
* refactor(proxy): inject tracing receiver and access context * refactor(proxy): own tracing resources through FastAPI lifespan * test(proxy): pass tracing dependency in Lens lifecycle * refactor(proxy): stop tracing logger cooperatively * refactor(proxy): derive tracing permissions in one place * refactor(proxy): compose application lifespan state * refactor(proxy): give Lens tracing storage directly * refactor(tracing): name shared ClickHouse storage explicitly * refactor(tracing): extract shared ClickHouse storage crate * test(proxy): isolate db push timeout from Lens safety check * fix(tracing): drain spend retries during shutdown
This commit is contained in:
parent
163ebccad5
commit
be67fce26a
36 changed files with 1152 additions and 559 deletions
17
litellm-rust/Cargo.lock
generated
17
litellm-rust/Cargo.lock
generated
|
|
@ -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]]
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_http::ClientVariant;
|
||||
use litellm_traces::{Connection, Error, InsertTable, Parameter, ReadQuery};
|
||||
use litellm_storage_clickhouse::Storage;
|
||||
use litellm_traces::{Error, InsertTable, Parameter, ReadQuery};
|
||||
use pyo3::{
|
||||
exceptions::{PyOverflowError, PyRuntimeError, PyValueError},
|
||||
prelude::*,
|
||||
|
|
@ -27,9 +28,7 @@ fn map_error(error: Error) -> PyErr {
|
|||
|
||||
#[pyclass]
|
||||
pub struct NativeTraceStorage {
|
||||
database: String,
|
||||
writer: Connection,
|
||||
reader: Option<Connection>,
|
||||
storage: Storage,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
|
|
@ -39,12 +38,7 @@ impl NativeTraceStorage {
|
|||
fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult<Self> {
|
||||
litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?;
|
||||
Ok(Self {
|
||||
writer: Connection::writer(url).map_err(map_error)?,
|
||||
reader: reader_url
|
||||
.map(|value| Connection::reader(value, &database))
|
||||
.transpose()
|
||||
.map_err(map_error)?,
|
||||
database,
|
||||
storage: Storage::new(database, url, reader_url).map_err(map_error)?,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -55,8 +49,8 @@ impl NativeTraceStorage {
|
|||
spend_log_retention_days: u32,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
let connection = self.writer.clone();
|
||||
let database = self.database.clone();
|
||||
let connection = self.storage.writer().clone();
|
||||
let database = self.storage.database().to_owned();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
|
|
@ -83,8 +77,8 @@ impl NativeTraceStorage {
|
|||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let table = InsertTable::parse(table).map_err(map_error)?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
let connection = self.writer.clone();
|
||||
let database = self.database.clone();
|
||||
let connection = self.storage.writer().clone();
|
||||
let database = self.storage.database().to_owned();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
|
|
@ -104,7 +98,7 @@ impl NativeTraceStorage {
|
|||
>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let query = litellm_traces::LensQuery::parse(name).map_err(map_error)?;
|
||||
let connection = self.reader.clone().ok_or_else(|| {
|
||||
let connection = self.storage.reader().cloned().ok_or_else(|| {
|
||||
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
|
||||
})?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
|
|
@ -127,7 +121,7 @@ impl NativeTraceStorage {
|
|||
>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let query = ReadQuery::parse(query).map_err(map_error)?;
|
||||
let connection = self.reader.clone().ok_or_else(|| {
|
||||
let connection = self.storage.reader().cloned().ok_or_else(|| {
|
||||
PyRuntimeError::new_err("Trace reads require a separate ClickHouse reader URL")
|
||||
})?;
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
|
|
|
|||
20
litellm-rust/crates/storage-clickhouse/Cargo.toml
Normal file
20
litellm-rust/crates/storage-clickhouse/Cargo.toml
Normal file
|
|
@ -0,0 +1,20 @@
|
|||
[package]
|
||||
name = "litellm-storage-clickhouse"
|
||||
version = "0.1.0"
|
||||
description = "Shared ClickHouse connection and HTTP storage for LiteLLM features"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
flate2.workspace = true
|
||||
litellm-http.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
url.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
5
litellm-rust/crates/storage-clickhouse/README.md
Normal file
5
litellm-rust/crates/storage-clickhouse/README.md
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
# ClickHouse storage
|
||||
|
||||
`litellm-storage-clickhouse` exports `Storage`, a shared writer connection and optional reader connection for one ClickHouse database. It also exports bounded HTTP read and insert execution
|
||||
|
||||
The crate has no trace tables, OTLP types, or named trace queries. `litellm-traces` supplies those rules and uses this storage for both trace rows and spend rows
|
||||
29
litellm-rust/crates/storage-clickhouse/src/error.rs
Normal file
29
litellm-rust/crates/storage-clickhouse/src/error.rs
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("invalid ClickHouse insert row")]
|
||||
InvalidRow,
|
||||
#[error("invalid ClickHouse insert table")]
|
||||
InvalidTable,
|
||||
#[error("invalid ClickHouse HTTP URL")]
|
||||
InvalidUrl,
|
||||
#[error("database must be a nonempty SQL identifier and retention must be positive")]
|
||||
InvalidSchema,
|
||||
#[error("SQL query must not be empty")]
|
||||
EmptySql,
|
||||
#[error("unknown ClickHouse read query")]
|
||||
InvalidQuery,
|
||||
#[error("ClickHouse query failed with HTTP status {0}")]
|
||||
QueryFailed(u16),
|
||||
#[error("ClickHouse insert failed with HTTP status {0}")]
|
||||
InsertFailed(u16),
|
||||
#[error("ClickHouse insert exceeds the encoded size limit")]
|
||||
InsertTooLarge,
|
||||
#[error("ClickHouse schema setup failed with HTTP status {0}")]
|
||||
SchemaFailed(u16),
|
||||
#[error("ClickHouse query exceeded the response size limit")]
|
||||
ResponseTooLarge,
|
||||
#[error("ClickHouse returned an invalid or failed JSON query response")]
|
||||
InvalidResponse,
|
||||
#[error("ClickHouse query transport failed")]
|
||||
Transport,
|
||||
}
|
||||
70
litellm-rust/crates/storage-clickhouse/src/insert.rs
Normal file
70
litellm-rust/crates/storage-clickhouse/src/insert.rs
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
use std::{io::Write, time::Duration};
|
||||
|
||||
use flate2::{Compression, write::GzEncoder};
|
||||
use litellm_http::Client;
|
||||
|
||||
use crate::{Connection, Error, valid_identifier};
|
||||
|
||||
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
pub async fn insert_encoded_rows(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
database: &str,
|
||||
table: &str,
|
||||
token: &str,
|
||||
encoded: &str,
|
||||
) -> Result<(), Error> {
|
||||
if !valid_identifier(database) {
|
||||
return Err(Error::InvalidSchema);
|
||||
}
|
||||
if !valid_identifier(table) {
|
||||
return Err(Error::InvalidTable);
|
||||
}
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
encoder
|
||||
.write_all(encoded.as_bytes())
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
let body = encoder.finish().map_err(|_| Error::InvalidRow)?;
|
||||
let mut url = connection.url().clone();
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "async_insert"
|
||||
| "async_insert_deduplicate"
|
||||
| "wait_for_async_insert"
|
||||
| "input_format_skip_unknown_fields"
|
||||
| "date_time_input_format"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair(
|
||||
"query",
|
||||
&format!("INSERT INTO `{database}`.{} FORMAT JSONEachRow", table),
|
||||
)
|
||||
.append_pair("insert_deduplication_token", token)
|
||||
.append_pair("async_insert", "1")
|
||||
.append_pair("async_insert_deduplicate", "1")
|
||||
.append_pair("wait_for_async_insert", "1")
|
||||
.append_pair("input_format_skip_unknown_fields", "0")
|
||||
.append_pair("date_time_input_format", "best_effort");
|
||||
let response = client
|
||||
.post(url)
|
||||
.timeout(INSERT_TIMEOUT)
|
||||
.header("Content-Encoding", "gzip")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::InsertFailed(response.status().as_u16()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
127
litellm-rust/crates/storage-clickhouse/src/lib.rs
Normal file
127
litellm-rust/crates/storage-clickhouse/src/lib.rs
Normal file
|
|
@ -0,0 +1,127 @@
|
|||
mod error;
|
||||
mod insert;
|
||||
mod read;
|
||||
|
||||
pub use error::Error;
|
||||
pub use insert::insert_encoded_rows;
|
||||
pub use read::{Parameter, execute_read};
|
||||
use url::Url;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Connection {
|
||||
url: Url,
|
||||
}
|
||||
|
||||
impl Connection {
|
||||
pub fn parse(value: &str) -> Result<Self, Error> {
|
||||
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
|
||||
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
|
||||
return Err(Error::InvalidUrl);
|
||||
}
|
||||
Ok(Self { url })
|
||||
}
|
||||
|
||||
pub fn configured(
|
||||
url: &str,
|
||||
database: &str,
|
||||
user: &str,
|
||||
password: &str,
|
||||
) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
connection
|
||||
.url
|
||||
.set_username(user)
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection
|
||||
.url
|
||||
.set_password(Some(password))
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn writer(url: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| key != "database")
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn url(&self) -> &Url {
|
||||
&self.url
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Storage {
|
||||
database: String,
|
||||
writer: Connection,
|
||||
reader: Option<Connection>,
|
||||
}
|
||||
|
||||
impl Storage {
|
||||
pub fn new(database: String, url: &str, reader_url: Option<&str>) -> Result<Self, Error> {
|
||||
if !valid_identifier(&database) {
|
||||
return Err(Error::InvalidSchema);
|
||||
}
|
||||
Ok(Self {
|
||||
writer: Connection::writer(url)?,
|
||||
reader: reader_url
|
||||
.map(|value| Connection::reader(value, &database))
|
||||
.transpose()?,
|
||||
database,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn database(&self) -> &str {
|
||||
&self.database
|
||||
}
|
||||
|
||||
pub fn writer(&self) -> &Connection {
|
||||
&self.writer
|
||||
}
|
||||
|
||||
pub fn reader(&self) -> Option<&Connection> {
|
||||
self.reader.as_ref()
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn valid_identifier(value: &str) -> bool {
|
||||
!value.is_empty()
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|c| c.is_ascii_alphanumeric() || c == b'_')
|
||||
}
|
||||
113
litellm-rust/crates/storage-clickhouse/src/read.rs
Normal file
113
litellm-rust/crates/storage-clickhouse/src/read.rs
Normal file
|
|
@ -0,0 +1,113 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use litellm_http::Client;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{Connection, Error};
|
||||
|
||||
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum Parameter {
|
||||
Text(String),
|
||||
Integer(i64),
|
||||
Strings(Vec<String>),
|
||||
}
|
||||
|
||||
impl Parameter {
|
||||
fn encoded(&self) -> String {
|
||||
match self {
|
||||
Self::Text(value) => escaped(value),
|
||||
Self::Integer(value) => value.to_string(),
|
||||
Self::Strings(values) => format!(
|
||||
"[{}]",
|
||||
values
|
||||
.iter()
|
||||
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
|
||||
.collect::<Vec<_>>()
|
||||
.join(",")
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn escaped(value: &str) -> String {
|
||||
value
|
||||
.replace('\\', "\\\\")
|
||||
.replace('\t', "\\t")
|
||||
.replace('\n', "\\n")
|
||||
.replace('\r', "\\r")
|
||||
.replace('\0', "\\0")
|
||||
}
|
||||
|
||||
pub async fn execute_read(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
sql: &str,
|
||||
parameters: &BTreeMap<String, Parameter>,
|
||||
) -> Result<String, Error> {
|
||||
if sql.trim().is_empty() {
|
||||
return Err(Error::EmptySql);
|
||||
}
|
||||
|
||||
let mut url = connection.url().clone();
|
||||
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!key.starts_with("param_")
|
||||
&& !matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "readonly"
|
||||
| "default_format"
|
||||
| "max_result_rows"
|
||||
| "result_overflow_mode"
|
||||
| "max_execution_time"
|
||||
| "wait_end_of_query"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair("readonly", "1")
|
||||
.append_pair("max_result_rows", "1000")
|
||||
.append_pair("result_overflow_mode", "throw")
|
||||
.append_pair("max_execution_time", "10")
|
||||
.append_pair("wait_end_of_query", "1")
|
||||
.append_pair("default_format", "JSON");
|
||||
|
||||
url.query_pairs_mut().extend_pairs(
|
||||
parameters
|
||||
.iter()
|
||||
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
|
||||
);
|
||||
|
||||
let request = client
|
||||
.post(url)
|
||||
.timeout(Duration::from_secs(15))
|
||||
.body(sql.to_owned());
|
||||
let mut response = request.send().await.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::QueryFailed(response.status().as_u16()));
|
||||
}
|
||||
|
||||
let mut body = Vec::new();
|
||||
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
|
||||
if body.len() + chunk.len() > MAX_RESPONSE_BYTES {
|
||||
return Err(Error::ResponseTooLarge);
|
||||
}
|
||||
body.extend_from_slice(&chunk);
|
||||
}
|
||||
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
|
||||
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
|
||||
{
|
||||
return Err(Error::InvalidResponse);
|
||||
}
|
||||
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
|
||||
}
|
||||
34
litellm-rust/crates/storage-clickhouse/tests/connection.rs
Normal file
34
litellm-rust/crates/storage-clickhouse/tests/connection.rs
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
use litellm_storage_clickhouse::{Connection, Storage};
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::http("http://localhost:8123", true)]
|
||||
#[case::https("https://localhost:8443", true)]
|
||||
#[case::tcp("tcp://localhost:9000", false)]
|
||||
#[case::missing_host("http://", false)]
|
||||
fn accepts_only_clickhouse_http_urls(#[case] value: &str, #[case] expected: bool) {
|
||||
assert_eq!(Connection::parse(value).is_ok(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::writer_only(None, false)]
|
||||
#[case::separate_reader(Some("http://localhost:8124"), true)]
|
||||
fn storage_exports_writer_and_optional_reader(
|
||||
#[case] reader_url: Option<&str>,
|
||||
#[case] has_reader: bool,
|
||||
) {
|
||||
let storage = Storage::new("litellm".to_owned(), "http://localhost:8123", reader_url)
|
||||
.expect("valid ClickHouse URLs");
|
||||
|
||||
assert_eq!(storage.database(), "litellm");
|
||||
assert_eq!(storage.writer().url().host_str(), Some("localhost"));
|
||||
assert_eq!(storage.writer().url().port(), Some(8123));
|
||||
assert_eq!(storage.reader().is_some(), has_reader);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty("")]
|
||||
#[case::injection("db; DROP DATABASE default")]
|
||||
fn storage_rejects_invalid_database(#[case] database: &str) {
|
||||
assert!(Storage::new(database.to_owned(), "http://localhost:8123", None).is_err());
|
||||
}
|
||||
34
litellm-rust/crates/storage-clickhouse/tests/transport.rs
Normal file
34
litellm-rust/crates/storage-clickhouse/tests/transport.rs
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_http::Client;
|
||||
use litellm_storage_clickhouse::{Connection, Error, execute_read, insert_encoded_rows};
|
||||
use rstest::rstest;
|
||||
|
||||
#[rstest]
|
||||
#[case::invalid_database("db; DROP DATABASE default", "spend_logs", true)]
|
||||
#[case::invalid_table("litellm", "spend_logs; DROP TABLE otel_traces", false)]
|
||||
#[tokio::test]
|
||||
async fn insert_rejects_invalid_identifiers(
|
||||
#[case] database: &str,
|
||||
#[case] table: &str,
|
||||
#[case] invalid_database: bool,
|
||||
) {
|
||||
let client = Client::no_redirect_for_test();
|
||||
let connection = Connection::writer("http://localhost:8123").expect("valid URL");
|
||||
let result = insert_encoded_rows(&client, &connection, database, table, "token", "{}").await;
|
||||
|
||||
assert!(matches!(&result, Err(Error::InvalidSchema)) == invalid_database);
|
||||
assert!(matches!(&result, Err(Error::InvalidTable)) == !invalid_database);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn read_rejects_empty_sql() {
|
||||
let client = Client::no_redirect_for_test();
|
||||
let connection = Connection::reader("http://localhost:8123", "litellm").expect("valid URL");
|
||||
|
||||
assert!(matches!(
|
||||
execute_read(&client, &connection, " ", &BTreeMap::new()).await,
|
||||
Err(Error::EmptySql)
|
||||
));
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
use std::{collections::BTreeMap, io::Write, time::Duration};
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use flate2::{Compression, write::GzEncoder};
|
||||
use litellm_http::Client;
|
||||
use serde_json::Value;
|
||||
use sha2::{Digest, Sha256};
|
||||
|
|
@ -9,7 +8,6 @@ use time::{OffsetDateTime, format_description::well_known::Rfc3339};
|
|||
use crate::{Connection, Error};
|
||||
|
||||
const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024;
|
||||
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
|
||||
|
||||
pub enum InsertTable {
|
||||
OtelTraces,
|
||||
|
|
@ -61,55 +59,15 @@ pub async fn insert_rows(
|
|||
})
|
||||
.collect();
|
||||
let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?;
|
||||
let mut encoder = GzEncoder::new(Vec::new(), Compression::default());
|
||||
encoder
|
||||
.write_all(encoded.as_bytes())
|
||||
.map_err(|_| Error::InvalidRow)?;
|
||||
let body = encoder.finish().map_err(|_| Error::InvalidRow)?;
|
||||
let mut url = connection.url().clone();
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "async_insert"
|
||||
| "async_insert_deduplicate"
|
||||
| "wait_for_async_insert"
|
||||
| "input_format_skip_unknown_fields"
|
||||
| "date_time_input_format"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair(
|
||||
"query",
|
||||
&format!(
|
||||
"INSERT INTO `{database}`.{} FORMAT JSONEachRow",
|
||||
table.name()
|
||||
),
|
||||
)
|
||||
.append_pair("insert_deduplication_token", &token)
|
||||
.append_pair("async_insert", "1")
|
||||
.append_pair("async_insert_deduplicate", "1")
|
||||
.append_pair("wait_for_async_insert", "1")
|
||||
.append_pair("input_format_skip_unknown_fields", "0")
|
||||
.append_pair("date_time_input_format", "best_effort");
|
||||
let response = client
|
||||
.post(url)
|
||||
.timeout(INSERT_TIMEOUT)
|
||||
.header("Content-Encoding", "gzip")
|
||||
.body(body)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::InsertFailed(response.status().as_u16()));
|
||||
}
|
||||
Ok(())
|
||||
litellm_storage_clickhouse::insert_encoded_rows(
|
||||
client,
|
||||
connection,
|
||||
database,
|
||||
table.name(),
|
||||
&token,
|
||||
&encoded,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub fn encode_rows(rows: Vec<BTreeMap<String, Value>>) -> Result<String, Error> {
|
||||
|
|
|
|||
|
|
@ -4,87 +4,9 @@ mod otlp;
|
|||
mod schema;
|
||||
mod sql;
|
||||
|
||||
pub use error::{DecodeError, Error};
|
||||
pub use error::DecodeError;
|
||||
pub use insert::{InsertTable, encode_rows, insert_rows};
|
||||
pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read};
|
||||
pub use otlp::{DecodedSpan, decode_otlp};
|
||||
pub use schema::{ensure_schema, schema_statements};
|
||||
pub use sql::{LensQuery, Parameter, ReadQuery, execute_named_read, execute_read};
|
||||
use url::Url;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct Connection {
|
||||
url: Url,
|
||||
}
|
||||
|
||||
impl Connection {
|
||||
pub fn parse(value: &str) -> Result<Self, Error> {
|
||||
let url = Url::parse(value).map_err(|_| Error::InvalidUrl)?;
|
||||
if !matches!(url.scheme(), "http" | "https") || url.host().is_none() {
|
||||
return Err(Error::InvalidUrl);
|
||||
}
|
||||
Ok(Self { url })
|
||||
}
|
||||
|
||||
pub fn configured(
|
||||
url: &str,
|
||||
database: &str,
|
||||
user: &str,
|
||||
password: &str,
|
||||
) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
connection
|
||||
.url
|
||||
.set_username(user)
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
connection
|
||||
.url
|
||||
.set_password(Some(password))
|
||||
.map_err(|_| Error::InvalidUrl)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "user" | "password"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn writer(url: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| !matches!(key.as_ref(), "database" | "readonly" | "query"))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection.url.query_pairs_mut().clear().extend_pairs(pairs);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn reader(url: &str, database: &str) -> Result<Self, Error> {
|
||||
let mut connection = Self::parse(url)?;
|
||||
let pairs: Vec<_> = connection
|
||||
.url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| key != "database")
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
connection
|
||||
.url
|
||||
.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(pairs)
|
||||
.append_pair("database", database);
|
||||
Ok(connection)
|
||||
}
|
||||
|
||||
pub fn url(&self) -> &Url {
|
||||
&self.url
|
||||
}
|
||||
}
|
||||
pub use sql::{LensQuery, ReadQuery, execute_named_read};
|
||||
|
|
|
|||
|
|
@ -1,12 +1,8 @@
|
|||
use std::{collections::BTreeMap, time::Duration};
|
||||
|
||||
use serde::Deserialize;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_http::Client;
|
||||
|
||||
use crate::{Connection, Error};
|
||||
|
||||
const MAX_RESPONSE_BYTES: usize = 4 * 1024 * 1024;
|
||||
use crate::{Connection, Error, Parameter, execute_read};
|
||||
|
||||
pub enum ReadQuery {
|
||||
ListTraces,
|
||||
|
|
@ -36,111 +32,6 @@ impl ReadQuery {
|
|||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub enum Parameter {
|
||||
Text(String),
|
||||
Integer(i64),
|
||||
Strings(Vec<String>),
|
||||
}
|
||||
|
||||
impl Parameter {
|
||||
fn encoded(&self) -> String {
|
||||
match self {
|
||||
Self::Text(value) => escaped(value),
|
||||
Self::Integer(value) => value.to_string(),
|
||||
Self::Strings(values) => format!(
|
||||
"[{}]",
|
||||
values
|
||||
.iter()
|
||||
.map(|value| format!("'{}'", escaped(value).replace('\'', "\\'")))
|
||||
.collect::<Vec<_>>()
|
||||
.join(",")
|
||||
),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn escaped(value: &str) -> String {
|
||||
value
|
||||
.replace('\\', "\\\\")
|
||||
.replace('\t', "\\t")
|
||||
.replace('\n', "\\n")
|
||||
.replace('\r', "\\r")
|
||||
.replace('\0', "\\0")
|
||||
}
|
||||
|
||||
pub async fn execute_read(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
sql: &str,
|
||||
parameters: &BTreeMap<String, Parameter>,
|
||||
) -> Result<String, Error> {
|
||||
if sql.trim().is_empty() {
|
||||
return Err(Error::EmptySql);
|
||||
}
|
||||
|
||||
let mut url = connection.url().clone();
|
||||
|
||||
let existing_pairs: Vec<(String, String)> = url
|
||||
.query_pairs()
|
||||
.filter(|(key, _)| {
|
||||
!key.starts_with("param_")
|
||||
&& !matches!(
|
||||
key.as_ref(),
|
||||
"query"
|
||||
| "readonly"
|
||||
| "default_format"
|
||||
| "max_result_rows"
|
||||
| "result_overflow_mode"
|
||||
| "max_execution_time"
|
||||
| "wait_end_of_query"
|
||||
)
|
||||
})
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect();
|
||||
url.query_pairs_mut()
|
||||
.clear()
|
||||
.extend_pairs(existing_pairs)
|
||||
.append_pair("readonly", "1")
|
||||
.append_pair("max_result_rows", "1000")
|
||||
.append_pair("result_overflow_mode", "throw")
|
||||
.append_pair("max_execution_time", "10")
|
||||
.append_pair("wait_end_of_query", "1")
|
||||
.append_pair("default_format", "JSON");
|
||||
|
||||
url.query_pairs_mut().extend_pairs(
|
||||
parameters
|
||||
.iter()
|
||||
.map(|(name, value)| (format!("param_{name}"), value.encoded())),
|
||||
);
|
||||
|
||||
let request = client
|
||||
.post(url)
|
||||
.timeout(Duration::from_secs(15))
|
||||
.body(sql.to_owned());
|
||||
let mut response = request.send().await.map_err(|_| Error::Transport)?;
|
||||
if !response.status().is_success() {
|
||||
return Err(Error::QueryFailed(response.status().as_u16()));
|
||||
}
|
||||
|
||||
let mut body = Vec::new();
|
||||
while let Some(chunk) = response.chunk().await.map_err(|_| Error::Transport)? {
|
||||
if body.len() + chunk.len() > MAX_RESPONSE_BYTES {
|
||||
return Err(Error::ResponseTooLarge);
|
||||
}
|
||||
body.extend_from_slice(&chunk);
|
||||
}
|
||||
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_slice(&body).map_err(|_| Error::InvalidResponse)?;
|
||||
if json.get("exception").is_some() || !json.get("data").is_some_and(serde_json::Value::is_array)
|
||||
{
|
||||
return Err(Error::InvalidResponse);
|
||||
}
|
||||
String::from_utf8(body).map_err(|_| Error::InvalidResponse)
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
pub enum LensQuery {
|
||||
Sample,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
66
litellm/proxy/tracing_runtime.py
Normal file
66
litellm/proxy/tracing_runtime.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
from collections.abc import AsyncGenerator, Callable
|
||||
from contextlib import asynccontextmanager
|
||||
from typing import Final
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
from pydantic import ConfigDict, TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.clickhouse.clickhouse_spend_logger import ClickHouseSpendLogger
|
||||
from litellm.rust_bridge.traces import ClickHouseStorage
|
||||
from litellm.tracing import TraceReceiver
|
||||
|
||||
_RECEIVER_ADAPTER: Final[TypeAdapter[TraceReceiver | None]] = TypeAdapter(
|
||||
TraceReceiver | None, config=ConfigDict(arbitrary_types_allowed=True)
|
||||
)
|
||||
_UNAVAILABLE_DETAIL: Final = "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL."
|
||||
|
||||
|
||||
def require_receiver(tracing: TraceReceiver | None) -> TraceReceiver:
|
||||
if tracing is None:
|
||||
raise HTTPException(status_code=501, detail=_UNAVAILABLE_DETAIL)
|
||||
return tracing
|
||||
|
||||
|
||||
async def provide_receiver(request: Request) -> TraceReceiver | None:
|
||||
return _RECEIVER_ADAPTER.validate_python(getattr(request.state, "tracing_receiver", None))
|
||||
|
||||
|
||||
async def provide_storage(request: Request) -> ClickHouseStorage | None:
|
||||
tracing: Final = await provide_receiver(request)
|
||||
return tracing.store.storage if tracing is not None else None
|
||||
|
||||
|
||||
async def _start_receiver(factory: Callable[[], TraceReceiver]) -> TraceReceiver | None:
|
||||
try:
|
||||
tracing: Final = factory()
|
||||
await tracing.start()
|
||||
return tracing
|
||||
except (KeyError, OSError, RuntimeError, ValueError) as error:
|
||||
verbose_proxy_logger.warning("Agent tracing unavailable: %s", error)
|
||||
return None
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def manage_tracing(
|
||||
enabled: bool, receiver_factory: Callable[[], TraceReceiver] = TraceReceiver.from_env
|
||||
) -> AsyncGenerator[TraceReceiver | None, None]:
|
||||
tracing: Final = await _start_receiver(receiver_factory) if enabled else None
|
||||
if tracing is None:
|
||||
yield tracing
|
||||
return
|
||||
|
||||
spend_logger: Final = ClickHouseSpendLogger(storage=tracing.store.storage)
|
||||
manager: Final = litellm.logging_callback_manager
|
||||
manager.add_litellm_callback(spend_logger)
|
||||
manager.add_litellm_success_callback(spend_logger)
|
||||
manager.add_litellm_failure_callback(spend_logger)
|
||||
manager.add_litellm_async_success_callback(spend_logger)
|
||||
manager.add_litellm_async_failure_callback(spend_logger)
|
||||
verbose_proxy_logger.info("Agent tracing enabled (store=clickhouse)")
|
||||
try:
|
||||
yield tracing
|
||||
finally:
|
||||
manager.remove_callback_from_all_lists(spend_logger)
|
||||
await spend_logger.aclose()
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue