diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 57e9e803e4a..ff0eafee47e 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4090,6 +4090,7 @@ dependencies = [ "litellm-token-counter", "litellm-traces", "litellm-tracing", + "prost", "pyo3", "pyo3-async-runtimes", "qdrant-client", @@ -4384,6 +4385,7 @@ name = "litellm-traces" version = "0.1.0" dependencies = [ "base64 0.22.1", + "criterion", "flate2", "litellm-http", "litellm-storage-clickhouse", @@ -4393,10 +4395,12 @@ dependencies = [ "serde", "serde_json", "sha2 0.10.9", + "strum", "testcontainers-modules", "thiserror 2.0.19", "time", "tokio", + "wiremock", ] [[package]] @@ -4819,6 +4823,7 @@ dependencies = [ "js-sys", "pin-project-lite", "thiserror 2.0.19", + "tracing", ] [[package]] @@ -4833,6 +4838,8 @@ dependencies = [ "opentelemetry_sdk 0.33.0", "prost", "serde", + "tonic", + "tonic-prost", ] [[package]] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 450253ea768..8d837c2d31b 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -82,6 +82,7 @@ reqwest = { version = "0.12", default-features = false, features = ["json", "mul qdrant-client = { version = "1.19.0", default-features = false } uuid = { version = "1", features = ["v4"] } rstest = "0.26.1" +wiremock = "0.6.5" rstest_reuse = "0.7.0" rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12"] } rustify = "=0.7.0" @@ -116,6 +117,8 @@ time = { version = "0.3.53", features = ["parsing"] } criterion = "0.8.2" fancy-regex = "0.19.2" veil = "0.3.0" +prost = "0.14.4" +opentelemetry-proto = "0.33" [profile.release] opt-level = 3 diff --git a/litellm-rust/crates/cache-azure-blob/Cargo.toml b/litellm-rust/crates/cache-azure-blob/Cargo.toml index 5bdfa16ef53..c28cb90d84a 100644 --- a/litellm-rust/crates/cache-azure-blob/Cargo.toml +++ b/litellm-rust/crates/cache-azure-blob/Cargo.toml @@ -26,4 +26,4 @@ litellm-cache-testing.workspace = true rstest.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/cache-gcs/Cargo.toml b/litellm-rust/crates/cache-gcs/Cargo.toml index 1a06683e615..91630879cbe 100644 --- a/litellm-rust/crates/cache-gcs/Cargo.toml +++ b/litellm-rust/crates/cache-gcs/Cargo.toml @@ -21,4 +21,4 @@ litellm-cache-testing.workspace = true rstest.workspace = true serde_json.workspace = true tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/cache-response/Cargo.toml b/litellm-rust/crates/cache-response/Cargo.toml index 1379573e505..869c40a12ab 100644 --- a/litellm-rust/crates/cache-response/Cargo.toml +++ b/litellm-rust/crates/cache-response/Cargo.toml @@ -21,4 +21,4 @@ redis = "1.7.0" redis-test = "1.0.4" rstest.workspace = true tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/cache-s3/Cargo.toml b/litellm-rust/crates/cache-s3/Cargo.toml index 680f2da8215..eb3a2fff1ac 100644 --- a/litellm-rust/crates/cache-s3/Cargo.toml +++ b/litellm-rust/crates/cache-s3/Cargo.toml @@ -23,6 +23,6 @@ tokio.workspace = true litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-testing.workspace = true rstest.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true serde_json.workspace = true tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 8410aff1d6a..85362fd90d2 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -47,4 +47,4 @@ litellm-host-native.workspace = true litellm-llms = { workspace = true, features = ["test-support"] } rstest.workspace = true rstest_reuse.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/gateway-inference/Cargo.toml b/litellm-rust/crates/gateway-inference/Cargo.toml index c854f0ea1ad..4544152d059 100644 --- a/litellm-rust/crates/gateway-inference/Cargo.toml +++ b/litellm-rust/crates/gateway-inference/Cargo.toml @@ -30,4 +30,4 @@ futures-util.workspace = true tokio = { workspace = true, features = ["io-util"] } rstest.workspace = true tower = { version = "0.5.3", features = ["util"] } -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/host-python/src/conversion_cache.rs b/litellm-rust/crates/host-python/src/conversion_cache.rs new file mode 100644 index 00000000000..c78ed42bcee --- /dev/null +++ b/litellm-rust/crates/host-python/src/conversion_cache.rs @@ -0,0 +1,57 @@ +use std::collections::{HashMap, hash_map::Entry}; + +use pyo3::prelude::*; + +pub struct ToPythonCache<'a, 'py, T> { + entries: HashMap)>, +} + +impl Default for ToPythonCache<'_, '_, T> { + fn default() -> Self { + Self { + entries: HashMap::new(), + } + } +} + +impl<'a, 'py, T> ToPythonCache<'a, 'py, T> { + pub fn get_or_try_insert_with( + &mut self, + value: &'a T, + convert: impl FnOnce(&'a T) -> PyResult>, + ) -> PyResult<&Bound<'py, PyAny>> { + let identity = std::ptr::from_ref(value) as usize; + let entry = match self.entries.entry(identity) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => entry.insert((value, convert(value)?)), + }; + Ok(&entry.1) + } +} + +pub struct FromPythonCache<'py, T> { + entries: HashMap, T)>, +} + +impl Default for FromPythonCache<'_, T> { + fn default() -> Self { + Self { + entries: HashMap::new(), + } + } +} + +impl<'py, T> FromPythonCache<'py, T> { + pub fn get_or_try_insert_with( + &mut self, + value: &Bound<'py, PyAny>, + convert: impl FnOnce(&Bound<'py, PyAny>) -> PyResult, + ) -> PyResult<&T> { + let identity = value.as_ptr() as usize; + let entry = match self.entries.entry(identity) { + Entry::Occupied(entry) => entry.into_mut(), + Entry::Vacant(entry) => entry.insert((value.clone(), convert(value)?)), + }; + Ok(&entry.1) + } +} diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 4de404e3624..00543f64085 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -5,6 +5,7 @@ mod argument; mod binding; +mod conversion_cache; mod driver; mod error; mod file_reader; @@ -20,6 +21,7 @@ mod services; pub use argument::lookup; pub use binding::PythonBinding; +pub use conversion_cache::{FromPythonCache, ToPythonCache}; pub use driver::{CallOptions, run_call}; pub use error::{InvokeError, missing_state}; pub use file_reader::{FileContent, PythonFileReader, py_bytes}; diff --git a/litellm-rust/crates/host-python/tests/conversion_cache.rs b/litellm-rust/crates/host-python/tests/conversion_cache.rs new file mode 100644 index 00000000000..70ad838e001 --- /dev/null +++ b/litellm-rust/crates/host-python/tests/conversion_cache.rs @@ -0,0 +1,121 @@ +use std::{cell::Cell, rc::Rc}; + +use litellm_host_python::{FromPythonCache, Pythonized, ToPythonCache}; +use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; +use rstest::{fixture, rstest}; + +#[fixture] +fn python() { + Python::initialize(); +} + +#[rstest] +fn rust_identity_reuses_python_objects_without_merging_equal_values(#[from(python)] _python: ()) { + Python::attach(|py| { + let original = Rc::new(vec![1, 2]); + let cloned = original.clone(); + let equal = Rc::new(vec![1, 2]); + let mut cache = ToPythonCache::default(); + let first = cache + .get_or_try_insert_with(original.as_ref(), |value| { + Pythonized(value).into_pyobject(py) + }) + .unwrap() + .clone(); + let second = cache + .get_or_try_insert_with(cloned.as_ref(), |_| panic!("must reuse conversion")) + .unwrap() + .clone(); + let third = cache + .get_or_try_insert_with(equal.as_ref(), |value| Pythonized(value).into_pyobject(py)) + .unwrap(); + assert!(first.is(&second)); + assert!(!first.is(third)); + assert!(first.eq(third).unwrap()); + }); +} + +#[rstest] +fn python_identity_reuses_rust_values_without_merging_equal_objects(#[from(python)] _python: ()) { + Python::attach(|py| { + let original = PyDict::new(py); + original.set_item("value", 1).unwrap(); + let equal = original.copy().unwrap(); + let calls = Cell::new(0); + let mut cache = FromPythonCache::default(); + let convert = |value: &Bound<'_, PyAny>| { + calls.set(calls.get() + 1); + value.get_item("value")?.extract::().map(Rc::new) + }; + let first = cache + .get_or_try_insert_with(original.as_any(), convert) + .unwrap() + .clone(); + let second = cache + .get_or_try_insert_with(original.as_any(), convert) + .unwrap() + .clone(); + let third = cache + .get_or_try_insert_with(equal.as_any(), convert) + .unwrap(); + assert!(Rc::ptr_eq(&first, &second)); + assert!(!Rc::ptr_eq(&first, third)); + assert_eq!(&first, third); + assert_eq!(calls.get(), 2); + }); +} + +#[rstest] +fn python_sources_stay_alive_until_the_cache_is_dropped(#[from(python)] _python: ()) { + Python::attach(|py| { + let value = py + .eval(pyo3::ffi::c_str!("type('Tracked', (), {})()"), None, None) + .unwrap(); + let weak = py + .import("weakref") + .unwrap() + .call_method1("ref", (&value,)) + .unwrap(); + let mut cache = FromPythonCache::default(); + cache.get_or_try_insert_with(&value, |_| Ok(42)).unwrap(); + drop(value); + assert!(!weak.call0().unwrap().is_none()); + drop(cache); + assert!(weak.call0().unwrap().is_none()); + }); +} + +#[rstest] +#[case::to_python(true)] +#[case::from_python(false)] +fn failed_conversions_preserve_exceptions_and_can_be_retried( + #[from(python)] _python: (), + #[case] to_python: bool, +) { + Python::attach(|py| { + let failure = PyValueError::new_err("conversion failed"); + if to_python { + let source = vec![1, 2]; + let mut cache = ToPythonCache::default(); + let error = cache + .get_or_try_insert_with(&source, |_| Err(failure.clone_ref(py))) + .unwrap_err(); + assert!(error.value(py).is(failure.value(py))); + let result = cache + .get_or_try_insert_with(&source, |value| Pythonized(value).into_pyobject(py)) + .unwrap(); + assert_eq!(result.extract::>().unwrap(), source); + } else { + let source = PyDict::new(py).into_any(); + let mut cache = FromPythonCache::default(); + let error = cache + .get_or_try_insert_with(&source, |_| Err(failure.clone_ref(py))) + .unwrap_err(); + assert!(error.value(py).is(failure.value(py))); + assert_eq!( + *cache.get_or_try_insert_with(&source, |_| Ok(42)).unwrap(), + 42 + ); + } + }); +} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index a1d1f63d6f3..2b505f08eca 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -52,6 +52,7 @@ litellm-llms-types.workspace = true litellm-host-python.workspace = true litellm-token-counter = { path = "../token-counter", default-features = false } pyo3.workspace = true +prost.workspace = true pyo3-async-runtimes.workspace = true reqwest.workspace = true redis = { version = "1.7.0", features = ["tls-rustls"] } @@ -73,7 +74,7 @@ futures-util.workspace = true rstest.workspace = true sha2.workspace = true tokio-tungstenite.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true aws-sdk-secretsmanager = "1.117.0" [[bench]] diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 0d4df996552..d269fa4015f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -44,7 +44,7 @@ mod _native { #[pymodule_export] use crate::routes::token_counter::TokenCounter; #[pymodule_export] - use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp}; + use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp, trace_encode_error}; #[cfg(feature = "huggingface")] #[pymodule_export] use crate::tokenizer::HuggingFaceEncoding; @@ -111,6 +111,7 @@ mod tests { "NativeDiagnosticProcessor", "NativeTraceStorage", "trace_decode_otlp", + "trace_encode_error", "TokenCounter", "Tokenizer", "gil_stats", diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 47c924f4842..ca66e2e46be 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -1,13 +1,33 @@ use std::collections::BTreeMap; +use litellm_host_python::{FromPythonCache, ToPythonCache}; use litellm_http::ClientVariant; use litellm_storage_clickhouse::Storage; -use litellm_traces::{Error, InsertTable, Parameter, ReadQuery}; +use litellm_traces::{Error, InsertTable, Parameter, ReadQuery, Shared}; +use prost::Message; use pyo3::{ exceptions::{PyOverflowError, PyRuntimeError, PyValueError}, prelude::*, + types::{PyBytes, PyDict, PyList, PyMapping, PyString}, }; +#[derive(Message)] +struct OtlpErrorStatus { + #[prost(int32, tag = "1")] + code: i32, + #[prost(string, tag = "2")] + message: String, +} + +#[pyfunction] +pub fn trace_encode_error<'py>(py: Python<'py>, message: &str) -> Bound<'py, PyBytes> { + let status = OtlpErrorStatus { + code: 0, + message: message.to_owned(), + }; + PyBytes::new(py, &status.encode_to_vec()) +} + fn map_error(error: Error) -> PyErr { match error { Error::InvalidRow @@ -71,9 +91,7 @@ impl NativeTraceStorage { &self, py: Python<'py>, table: &str, - #[pyo3(from_py_with = litellm_host_python::from_py_argument)] rows: Vec< - BTreeMap, - >, + #[pyo3(from_py_with = insert_rows_from_py)] rows: Vec, ) -> PyResult> { let table = InsertTable::parse(table).map_err(map_error)?; let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; @@ -82,7 +100,8 @@ impl NativeTraceStorage { crate::execution::run_async( py, async move { - litellm_traces::insert_rows(&client, &connection, &database, table, rows).await + litellm_traces::insert_shared_rows(&client, &connection, &database, table, rows) + .await }, map_error, ) @@ -140,21 +159,132 @@ pub fn trace_decode_otlp<'py>( py: Python<'py>, body: &[u8], content_type: Option<&str>, - content_encoding: Option<&str>, - max_decompressed_bytes: usize, ) -> PyResult> { let spans = py - .detach(|| { - litellm_traces::decode_otlp( - body, - content_type, - content_encoding, - max_decompressed_bytes, - ) - }) + .detach(|| litellm_traces::decode_otlp(body, content_type)) .map_err(|error| match error { litellm_traces::DecodeError::TooLarge => PyOverflowError::new_err(error.to_string()), _ => PyValueError::new_err(error.to_string()), })?; - litellm_host_python::Pythonized(spans).into_pyobject(py) + spans_to_py(py, &spans).map(Bound::into_any) +} + +fn insert_rows_from_py(value: &Bound<'_, PyAny>) -> PyResult> { + let mut resources = FromPythonCache::default(); + value + .try_iter()? + .map(|row| { + let row = row?; + let mut fields = BTreeMap::new(); + for item in row.cast::()?.items()?.iter() { + let (key, value): (String, Bound<'_, PyAny>) = item.extract()?; + let converted = if matches!( + key.as_str(), + "ResourceAttributes" | "ScopeName" | "ScopeVersion" + ) { + resources + .get_or_try_insert_with(&value, |value| { + litellm_host_python::from_py_argument::(value) + .map(Shared::new) + })? + .clone() + } else { + Shared::new(litellm_host_python::from_py_argument(&value)?) + }; + fields.insert(key, converted); + } + Ok(fields) + }) + .collect() +} + +fn spans_to_py<'py>( + py: Python<'py>, + spans: &[litellm_traces::DecodedSpan], +) -> PyResult> { + let mut resources = ToPythonCache::default(); + let mut scopes = ToPythonCache::default(); + let result = PyList::empty(py); + for span in spans { + let resource = resources + .get_or_try_insert_with(span.resource_attributes.as_ref(), |value| { + litellm_host_python::Pythonized(value).into_pyobject(py) + })?; + let row = PyDict::new(py); + row.set_item("trace_id", &span.trace_id)?; + row.set_item("span_id", &span.span_id)?; + row.set_item("parent_span_id", &span.parent_span_id)?; + row.set_item("trace_state", &span.trace_state)?; + row.set_item("name", &span.name)?; + row.set_item("kind", &span.kind)?; + row.set_item("resource_attributes", resource)?; + for (key, value) in [ + ("scope_name", &span.scope_name), + ("scope_version", &span.scope_version), + ] { + let value = scopes.get_or_try_insert_with(value.as_ref(), |value| { + Ok(PyString::new(py, value).into_any()) + })?; + row.set_item(key, value)?; + } + row.set_item("attributes", &span.attributes)?; + row.set_item("start_ns", span.start_ns)?; + row.set_item("end_ns", span.end_ns)?; + row.set_item("status_code", &span.status_code)?; + row.set_item("status_message", &span.status_message)?; + row.set_item( + "events", + litellm_host_python::Pythonized(&span.events).into_pyobject(py)?, + )?; + result.append(row)?; + } + Ok(result) +} + +#[cfg(test)] +mod tests { + use super::*; + use rstest::rstest; + + #[rstest] + fn insert_projection_preserves_identity_without_merging_equal_resources() { + Python::initialize(); + Python::attach(|py| { + let resource = PyDict::new(py); + resource.set_item("service.name", "shared").unwrap(); + let equal_resource = resource.copy().unwrap(); + let rows = PyList::empty(py); + for value in [&resource, &resource, &equal_resource] { + let row = PyDict::new(py); + row.set_item("ResourceAttributes", value).unwrap(); + rows.append(row).unwrap(); + } + let projected = insert_rows_from_py(rows.as_any()).unwrap(); + assert!(Shared::shares_storage_with( + &projected[0]["ResourceAttributes"], + &projected[1]["ResourceAttributes"] + )); + assert!(!Shared::shares_storage_with( + &projected[0]["ResourceAttributes"], + &projected[2]["ResourceAttributes"] + )); + assert_eq!(projected[0], projected[2]); + }); + } + + #[rstest] + fn shared_conversion_preserves_every_decoded_field() { + Python::initialize(); + Python::attach(|py| { + let spans = litellm_traces::decode_otlp( + include_bytes!("../../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json"), + Some("application/json"), + ).unwrap(); + let expected = litellm_host_python::Pythonized(&spans) + .into_pyobject(py) + .unwrap(); + let actual = spans_to_py(py, &spans).unwrap(); + assert!(actual.eq(expected).unwrap()); + }); + } } diff --git a/litellm-rust/crates/secrets-aws/Cargo.toml b/litellm-rust/crates/secrets-aws/Cargo.toml index e7a394bd247..5d3bd413484 100644 --- a/litellm-rust/crates/secrets-aws/Cargo.toml +++ b/litellm-rust/crates/secrets-aws/Cargo.toml @@ -21,5 +21,5 @@ aws-credential-types = "1.3.0" base64.workspace = true rstest.workspace = true tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true tempfile = "3" diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml index efdf681e2bc..7ec03fb98da 100644 --- a/litellm-rust/crates/secrets-azure/Cargo.toml +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -20,7 +20,7 @@ percent-encoding = "2.3" [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } -wiremock = "0.6.5" +wiremock.workspace = true rstest.workspace = true serde_json.workspace = true sha2.workspace = true diff --git a/litellm-rust/crates/secrets-cyberark/Cargo.toml b/litellm-rust/crates/secrets-cyberark/Cargo.toml index 0a91c61ade9..f630d5857d8 100644 --- a/litellm-rust/crates/secrets-cyberark/Cargo.toml +++ b/litellm-rust/crates/secrets-cyberark/Cargo.toml @@ -25,6 +25,6 @@ rcgen = "0.14.10" rstest.workspace = true tempfile = "3.27.0" tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true serde.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/secrets-google/Cargo.toml b/litellm-rust/crates/secrets-google/Cargo.toml index 208b5ddd03f..3ce14fe7a12 100644 --- a/litellm-rust/crates/secrets-google/Cargo.toml +++ b/litellm-rust/crates/secrets-google/Cargo.toml @@ -28,4 +28,4 @@ reqwest.workspace = true litellm-http = { workspace = true, features = ["test-support"] } google-cloud-auth.workspace = true rstest.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/secrets-hashicorp/Cargo.toml b/litellm-rust/crates/secrets-hashicorp/Cargo.toml index c049ba127e5..7dd3d3c674f 100644 --- a/litellm-rust/crates/secrets-hashicorp/Cargo.toml +++ b/litellm-rust/crates/secrets-hashicorp/Cargo.toml @@ -21,4 +21,4 @@ veil.workspace = true rstest.workspace = true tempfile = "3" tokio.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index f855a8a64a6..3655ce8bbc2 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -36,7 +36,7 @@ tokio = { workspace = true, features = ["fs"] } [dev-dependencies] litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true -wiremock = "0.6.5" +wiremock.workspace = true tempfile = "3" aws-sdk-kms = "1.120.0" google-cloud-kms-v1 = "1.14.0" diff --git a/litellm-rust/crates/storage-clickhouse/src/insert.rs b/litellm-rust/crates/storage-clickhouse/src/insert.rs index 3be73b086a0..81528ded907 100644 --- a/litellm-rust/crates/storage-clickhouse/src/insert.rs +++ b/litellm-rust/crates/storage-clickhouse/src/insert.rs @@ -26,6 +26,23 @@ pub async fn insert_encoded_rows( .write_all(encoded.as_bytes()) .map_err(|_| Error::InvalidRow)?; let body = encoder.finish().map_err(|_| Error::InvalidRow)?; + insert_compressed_rows(client, connection, database, table, token, body).await +} + +pub async fn insert_compressed_rows( + client: &Client, + connection: &Connection, + database: &str, + table: &str, + token: &str, + body: Vec, +) -> Result<(), Error> { + if !valid_identifier(database) { + return Err(Error::InvalidSchema); + } + if !valid_identifier(table) { + return Err(Error::InvalidTable); + } let mut url = connection.url().clone(); let existing_pairs: Vec<(String, String)> = url .query_pairs() diff --git a/litellm-rust/crates/storage-clickhouse/src/lib.rs b/litellm-rust/crates/storage-clickhouse/src/lib.rs index f2b34eddbf8..d11ee9d5cde 100644 --- a/litellm-rust/crates/storage-clickhouse/src/lib.rs +++ b/litellm-rust/crates/storage-clickhouse/src/lib.rs @@ -3,7 +3,7 @@ mod insert; mod read; pub use error::Error; -pub use insert::insert_encoded_rows; +pub use insert::{insert_compressed_rows, insert_encoded_rows}; pub use read::{Parameter, execute_read}; use url::Url; diff --git a/litellm-rust/crates/traces/Cargo.toml b/litellm-rust/crates/traces/Cargo.toml index b7f6e6ae52e..74de400764c 100644 --- a/litellm-rust/crates/traces/Cargo.toml +++ b/litellm-rust/crates/traces/Cargo.toml @@ -8,18 +8,25 @@ repository.workspace = true [dependencies] base64.workspace = true flate2.workspace = true -opentelemetry-proto = { version = "0.33.0", default-features = false, features = ["gen-tonic-messages", "trace", "with-serde"] } -prost = "0.14.4" +opentelemetry-proto = { workspace = true, features = ["gen-tonic-messages", "trace", "with-serde"] } +prost.workspace = true time = { workspace = true, features = ["formatting"] } litellm-http.workspace = true litellm-storage-clickhouse.workspace = true sha2.workspace = true -serde.workspace = true +serde = { workspace = true, features = ["rc"] } serde_json.workspace = true +strum.workspace = true thiserror.workspace = true [dev-dependencies] +criterion.workspace = true litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true testcontainers-modules = { version = "0.15.0", features = ["clickhouse"] } tokio.workspace = true +wiremock.workspace = true + +[[bench]] +name = "resource-fanout" +harness = false diff --git a/litellm-rust/crates/traces/benches/resource-fanout.rs b/litellm-rust/crates/traces/benches/resource-fanout.rs new file mode 100644 index 00000000000..edf5d2eb055 --- /dev/null +++ b/litellm-rust/crates/traces/benches/resource-fanout.rs @@ -0,0 +1,39 @@ +use std::{collections::BTreeMap, hint::black_box, time::Duration}; + +use criterion::{BenchmarkId, Criterion, Throughput, criterion_group, criterion_main}; +use litellm_traces::Shared; + +fn fanout(resource: &T, spans: usize) -> Vec { + (0..spans).map(|_| resource.clone()).collect() +} + +fn resource_fanout(c: &mut Criterion) { + let mut group = c.benchmark_group("resource_fanout"); + for (attribute_bytes, spans) in [(256, 1), (256, 64), (8192, 1024), (16384, 1024)] { + let attributes = BTreeMap::from([ + ("service.name".to_owned(), "benchmark".to_owned()), + ("payload".to_owned(), "x".repeat(attribute_bytes)), + ]); + let owned = Box::new(attributes.clone()); + let shared = Shared::new(attributes); + let case = format!("{attribute_bytes}B_{spans}_spans"); + group.throughput(Throughput::Elements(spans as u64)); + group.bench_with_input(BenchmarkId::new("owned", &case), &owned, |b, resource| { + b.iter(|| black_box(fanout(black_box(resource), spans))); + }); + group.bench_with_input(BenchmarkId::new("shared", &case), &shared, |b, resource| { + b.iter(|| black_box(fanout(black_box(resource), spans))); + }); + } + group.finish(); +} + +criterion_group! { + name = benches; + config = Criterion::default() + .sample_size(20) + .warm_up_time(Duration::from_secs(1)) + .measurement_time(Duration::from_secs(2)); + targets = resource_fanout +} +criterion_main!(benches); diff --git a/litellm-rust/crates/traces/query/span_error.sql b/litellm-rust/crates/traces/query/span_error.sql new file mode 100644 index 00000000000..b4710006389 --- /dev/null +++ b/litellm-rust/crates/traces/query/span_error.sql @@ -0,0 +1,13 @@ +SELECT SpanId AS span_id, + substringUTF8(StatusMessage, {error_offset:UInt64} + 1, 16384) AS message, + lengthUTF8(StatusMessage) AS total_chars, + hex(SHA256(StatusMessage)) AS version +FROM otel_traces +WHERE TraceId = {trace_id:String} AND SpanId = {span_id:String} + AND (empty({team_ids:Array(String)}) OR TeamId IN {team_ids:Array(String)}) + AND ({api_key_hash:String} = '' OR ApiKeyHash = {api_key_hash:String}) + AND ({trace_ref:String} = '' OR + hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String}) + AND ({error_version:String} = '' OR hex(SHA256(StatusMessage)) = {error_version:String}) +ORDER BY Timestamp, EngineReceivedMs, StatusMessage +LIMIT 1 diff --git a/litellm-rust/crates/traces/query/trace_spans.sql b/litellm-rust/crates/traces/query/trace_spans.sql index 409e6328198..dab3ac2e877 100644 --- a/litellm-rust/crates/traces/query/trace_spans.sql +++ b/litellm-rust/crates/traces/query/trace_spans.sql @@ -1,6 +1,7 @@ SELECT o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name, o.ObservationType AS type, o.AgentName AS agent, o.StatusCode AS status, - o.StatusMessage AS status_message, + substringUTF8(o.StatusMessage, 1, 128) AS status_message, + lengthUTF8(o.StatusMessage) > 128 AS error_truncated, toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns, o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model, o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens, @@ -12,5 +13,5 @@ WHERE o.TraceId = {trace_id:String} AND ({api_key_hash:String} = '' OR o.ApiKeyHash = {api_key_hash:String}) AND ({trace_ref:String} = '' OR hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String}) -ORDER BY o.Timestamp +ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage LIMIT 1 BY o.SpanId diff --git a/litellm-rust/crates/traces/src/error.rs b/litellm-rust/crates/traces/src/error.rs index 2ccfe0ea8d9..18fa4af9b53 100644 --- a/litellm-rust/crates/traces/src/error.rs +++ b/litellm-rust/crates/traces/src/error.rs @@ -2,6 +2,6 @@ pub enum DecodeError { #[error("invalid OTLP trace payload")] InvalidPayload, - #[error("OTLP trace payload exceeds the decompressed size limit")] + #[error("OTLP trace payload exceeds the decoding budget")] TooLarge, } diff --git a/litellm-rust/crates/traces/src/insert.rs b/litellm-rust/crates/traces/src/insert.rs index 6d9a2cab813..01a1eecfd7b 100644 --- a/litellm-rust/crates/traces/src/insert.rs +++ b/litellm-rust/crates/traces/src/insert.rs @@ -1,14 +1,23 @@ -use std::collections::BTreeMap; +use std::{ + borrow::Cow, + collections::BTreeMap, + io::{BufWriter, Write}, +}; +use serde::{Serialize, Serializer, ser::SerializeMap}; + +use flate2::{Compression, write::GzEncoder}; use litellm_http::Client; use serde_json::Value; use sha2::{Digest, Sha256}; use time::{OffsetDateTime, format_description::well_known::Rfc3339}; -use crate::{Connection, Error}; +use crate::{Connection, Error, Shared}; const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024; +pub type InsertRow = BTreeMap>; + pub enum InsertTable { OtelTraces, SpendLogs, @@ -37,85 +46,176 @@ pub async fn insert_rows( database: &str, table: InsertTable, rows: Vec>, +) -> Result<(), Error> { + insert_shared_rows(client, connection, database, table, shared_rows(rows)).await +} + +pub async fn insert_shared_rows( + client: &Client, + connection: &Connection, + database: &str, + table: InsertTable, + rows: Vec, ) -> Result<(), Error> { if rows.is_empty() { return Ok(()); } - let token = format!( - "{:x}", - Sha256::digest(encode_rows_with_limit(rows.clone(), MAX_INSERT_BYTES)?.as_bytes()) - ); - let received_ms = OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000; - let rows = rows - .into_iter() - .map(|row| { - row.into_iter() - .filter(|(key, _)| key != "EngineReceivedMs") - .chain(std::iter::once(( - "EngineReceivedMs".to_owned(), - Value::from(received_ms as u64), - ))) - .collect() - }) - .collect(); - let encoded = encode_rows_with_limit(rows, MAX_INSERT_BYTES)?; - litellm_storage_clickhouse::insert_encoded_rows( + let received_ms = (OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64; + let (token, body) = prepare_insert(&rows, received_ms, MAX_INSERT_BYTES)?; + litellm_storage_clickhouse::insert_compressed_rows( client, connection, database, table.name(), &token, - &encoded, + body, ) .await } -pub fn encode_rows(rows: Vec>) -> Result { - encode_rows_with_limit(rows, usize::MAX) +fn shared_rows(rows: Vec>) -> Vec { + rows.into_iter() + .map(|row| { + row.into_iter() + .map(|(key, value)| (key, Shared::new(value))) + .collect() + }) + .collect() } -fn encode_rows_with_limit( - rows: Vec>, - limit: usize, -) -> Result { - let mut body = Vec::new(); - for row in rows { - let encoded = row - .into_iter() - .map(|(name, value)| insert_value(&name, value).map(|value| (name, value))) - .collect::, _>>()?; - let record = serde_json::to_vec(&encoded).map_err(|_| Error::InvalidRow)?; - let size = body - .len() - .checked_add(record.len()) - .and_then(|size| size.checked_add(usize::from(!body.is_empty()))) - .ok_or(Error::InsertTooLarge)?; - if size > limit { - return Err(Error::InsertTooLarge); - } - if !body.is_empty() { - body.push(b'\n'); - } - body.extend_from_slice(&record); - } +pub fn encode_rows(rows: Vec>) -> Result { + let body = write_rows(&shared_rows(rows), None, Vec::new(), usize::MAX)?; String::from_utf8(body).map_err(|_| Error::InvalidRow) } -fn insert_value(name: &str, value: Value) -> Result { +fn prepare_insert( + rows: &[InsertRow], + received_ms: u64, + limit: usize, +) -> Result<(String, Vec), Error> { + let hash = write_rows(rows, None, HashWriter(Sha256::new()), limit)?; + let token = format!("{:x}", hash.0.finalize()); + let encoder = write_rows( + rows, + Some(received_ms), + BufWriter::new(GzEncoder::new(Vec::new(), Compression::default())), + limit, + )?; + let body = encoder + .into_inner() + .map_err(|_| Error::InvalidRow)? + .finish() + .map_err(|_| Error::InvalidRow)?; + Ok((token, body)) +} + +struct HashWriter(Sha256); + +impl Write for HashWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0.update(bytes); + Ok(bytes.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +struct LimitedWriter { + inner: W, + remaining: usize, + exceeded: bool, +} + +impl Write for LimitedWriter { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + if bytes.len() > self.remaining { + self.exceeded = true; + return Err(std::io::Error::other(Error::InsertTooLarge)); + } + let written = self.inner.write(bytes)?; + self.remaining -= written; + Ok(written) + } + + fn flush(&mut self) -> std::io::Result<()> { + self.inner.flush() + } +} + +fn write_rows( + rows: &[InsertRow], + received_ms: Option, + writer: W, + limit: usize, +) -> Result { + let mut writer = LimitedWriter { + inner: writer, + remaining: limit, + exceeded: false, + }; + for (index, row) in rows.iter().enumerate() { + let result = (|| { + if index != 0 { + writer.write_all(b"\n").map_err(serde_json::Error::io)?; + } + serde_json::to_writer(&mut writer, &EncodedRow { row, received_ms }) + })(); + if result.is_err() { + return Err(if writer.exceeded { + Error::InsertTooLarge + } else { + Error::InvalidRow + }); + } + } + Ok(writer.inner) +} + +struct EncodedRow<'a> { + row: &'a InsertRow, + received_ms: Option, +} + +impl Serialize for EncodedRow<'_> { + fn serialize(&self, serializer: S) -> Result { + let mut map = serializer.serialize_map(None)?; + let mut received_ms = self.received_ms; + for (name, value) in self.row { + if name.as_str() >= "EngineReceivedMs" + && let Some(timestamp) = received_ms.take() + { + map.serialize_entry("EngineReceivedMs", ×tamp)?; + } + if name == "EngineReceivedMs" && self.received_ms.is_some() { + continue; + } + let value = insert_value(name, value).map_err(serde::ser::Error::custom)?; + map.serialize_entry(name, &value)?; + } + if let Some(timestamp) = received_ms { + map.serialize_entry("EngineReceivedMs", ×tamp)?; + } + map.end() + } +} + +fn insert_value<'a>(name: &str, value: &'a Value) -> Result, Error> { let multiplier = match name { "Timestamp" => 1, "start_time" | "end_time" | "completion_start_time" => 1_000_000, - _ => return Ok(value), + _ => return Ok(Cow::Borrowed(value)), }; if name == "completion_start_time" && value.is_null() { - return Ok(value); + return Ok(Cow::Borrowed(value)); } let timestamp = value.as_i64().ok_or(Error::InvalidRow)?; let datetime = OffsetDateTime::from_unix_timestamp_nanos(i128::from(timestamp) * multiplier) .map_err(|_| Error::InvalidRow)?; datetime .format(&Rfc3339) - .map(Value::String) + .map(|value| Cow::Owned(Value::String(value))) .map_err(|_| Error::InvalidRow) } @@ -126,20 +226,83 @@ mod tests { use rstest::rstest; use serde_json::json; - use super::encode_rows_with_limit; + use super::{shared_rows, write_rows}; use crate::Error; #[rstest] fn encoded_limit_counts_utf8_bytes_across_rows() { - let rows = vec![ + let rows = shared_rows(vec![ BTreeMap::from([("Input".to_owned(), json!("雪"))]), BTreeMap::from([("Input".to_owned(), json!("雪"))]), - ]; - let encoded = encode_rows_with_limit(rows.clone(), usize::MAX).expect("valid rows"); + ]); + let encoded = write_rows(&rows, None, Vec::new(), usize::MAX).expect("valid rows"); - assert!(encode_rows_with_limit(rows.clone(), encoded.len()).is_ok()); + assert!(write_rows(&rows, None, Vec::new(), encoded.len()).is_ok()); assert!(matches!( - encode_rows_with_limit(rows, encoded.len() - 1), + write_rows(&rows, None, Vec::new(), encoded.len() - 1), + Err(Error::InsertTooLarge) + )); + } + + #[rstest] + #[case::absent(None)] + #[case::submitted(Some(123))] + fn streamed_insert_preserves_token_and_stamps_receive_time(#[case] submitted: Option) { + use flate2::read::GzDecoder; + use sha2::{Digest, Sha256}; + use std::io::Read; + let mut row = BTreeMap::from([ + ("ApiKeyHash".into(), json!("key")), + ("ResourceAttributes".into(), json!({"message": "雪\n\""})), + ("Timestamp".into(), json!(1_234_567_890)), + ]); + if let Some(value) = submitted { + row.insert("EngineReceivedMs".into(), json!(value)); + } + let legacy = match submitted { + Some(_) => { + "{\"ApiKeyHash\":\"key\",\"EngineReceivedMs\":123,\"ResourceAttributes\":{\"message\":\"雪\\n\\\"\"},\"Timestamp\":\"1970-01-01T00:00:01.23456789Z\"}" + } + None => { + "{\"ApiKeyHash\":\"key\",\"ResourceAttributes\":{\"message\":\"雪\\n\\\"\"},\"Timestamp\":\"1970-01-01T00:00:01.23456789Z\"}" + } + }; + let rows = shared_rows(vec![row.clone(), row]); + let (token, body) = super::prepare_insert(&rows, 456, 4096).unwrap(); + assert_eq!( + token, + format!("{:x}", Sha256::digest(format!("{legacy}\n{legacy}"))) + ); + let mut decoded = String::new(); + GzDecoder::new(body.as_slice()) + .read_to_string(&mut decoded) + .unwrap(); + let expected = json!({ + "ApiKeyHash": "key", "EngineReceivedMs": 456, + "ResourceAttributes": {"message": "雪\n\""}, + "Timestamp": "1970-01-01T00:00:01.23456789Z", + }); + assert_eq!( + decoded + .lines() + .map(|line| serde_json::from_str::(line).unwrap()) + .collect::>(), + vec![expected.clone(), expected] + ); + assert_eq!( + rows[0] + .get("EngineReceivedMs") + .map(|value| value.as_u64().unwrap()), + submitted + ); + } + + #[rstest] + fn stamped_insert_enforces_the_encoded_limit() { + let rows = shared_rows(vec![BTreeMap::new()]); + assert!(super::prepare_insert(&rows, 1, 22).is_ok()); + assert!(matches!( + super::prepare_insert(&rows, 1, 21), Err(Error::InsertTooLarge) )); } diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index f5defb36cc2..1489b44c118 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -2,11 +2,13 @@ mod error; mod insert; mod otlp; mod schema; +mod shared; mod sql; pub use error::DecodeError; -pub use insert::{InsertTable, encode_rows, insert_rows}; +pub use insert::{InsertRow, InsertTable, encode_rows, insert_rows, insert_shared_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 shared::{Shared, SharedIdentity}; pub use sql::{LensQuery, ReadQuery, execute_named_read}; diff --git a/litellm-rust/crates/traces/src/otlp.rs b/litellm-rust/crates/traces/src/otlp.rs deleted file mode 100644 index f162256ef1f..00000000000 --- a/litellm-rust/crates/traces/src/otlp.rs +++ /dev/null @@ -1,221 +0,0 @@ -use std::{collections::BTreeMap, io::Read}; - -use base64::Engine; -use flate2::read::GzDecoder; -use opentelemetry_proto::tonic::{ - collector::trace::v1::ExportTraceServiceRequest, - common::v1::{AnyValue, KeyValue, any_value::Value as AttributeValue}, - trace::v1::{Span, span::SpanKind, status::StatusCode}, -}; -use prost::Message; -use serde::Serialize; -use serde_json::Value; - -use crate::DecodeError; - -#[derive(Serialize)] -pub struct DecodedEvent { - pub name: String, - pub attributes: BTreeMap, -} - -#[derive(Serialize)] -pub struct DecodedSpan { - pub trace_id: String, - pub span_id: String, - pub parent_span_id: String, - pub trace_state: String, - pub name: String, - pub kind: String, - pub resource_attributes: BTreeMap, - pub scope_name: String, - pub scope_version: String, - pub attributes: BTreeMap, - pub start_ns: u64, - pub end_ns: u64, - pub status_code: String, - pub status_message: String, - pub events: Vec, -} - -pub fn decode_otlp( - body: &[u8], - content_type: Option<&str>, - content_encoding: Option<&str>, - max_decompressed_bytes: usize, -) -> Result, DecodeError> { - let payload = if content_encoding == Some("gzip") || body.starts_with(&[0x1f, 0x8b]) { - let limit = u64::try_from(max_decompressed_bytes).map_err(|_| DecodeError::TooLarge)?; - let mut decoded = Vec::new(); - GzDecoder::new(body) - .take(limit + 1) - .read_to_end(&mut decoded) - .map_err(|_| DecodeError::InvalidPayload)?; - decoded - } else { - body.to_vec() - }; - if payload.len() > max_decompressed_bytes { - return Err(DecodeError::TooLarge); - } - let request = if content_type.is_some_and(|value| value.contains("json")) { - let value: Value = - serde_json::from_slice(&payload).map_err(|_| DecodeError::InvalidPayload)?; - serde_json::from_value(normalize_json_ids(value)?) - .map_err(|_| DecodeError::InvalidPayload)? - } else { - ExportTraceServiceRequest::decode(payload.as_slice()) - .map_err(|_| DecodeError::InvalidPayload)? - }; - Ok(request - .resource_spans - .into_iter() - .flat_map(|resource_spans| { - let resource_attributes = attributes( - resource_spans - .resource - .map(|resource| resource.attributes) - .unwrap_or_default(), - ); - resource_spans - .scope_spans - .into_iter() - .flat_map(move |scope_spans| { - let scope = scope_spans.scope.unwrap_or_default(); - let resource_attributes = resource_attributes.clone(); - scope_spans.spans.into_iter().map(move |span| { - decoded_span(span, &resource_attributes, &scope.name, &scope.version) - }) - }) - }) - .collect()) -} - -fn normalize_json_ids(value: Value) -> Result { - match value { - Value::Object(fields) => fields - .into_iter() - .map(|(name, value)| { - let normalized = if matches!(name.as_str(), "traceId" | "spanId" | "parentSpanId") { - let encoded = value.as_str().ok_or(DecodeError::InvalidPayload)?; - let bytes = base64::engine::general_purpose::STANDARD - .decode(encoded) - .map_err(|_| DecodeError::InvalidPayload)?; - Value::String(hex_bytes(&bytes)) - } else if name == "kind" && value.is_string() { - let kind = SpanKind::from_str_name(value.as_str().unwrap_or_default()) - .ok_or(DecodeError::InvalidPayload)?; - Value::from(kind as i32) - } else if name == "code" && value.is_string() { - let code = StatusCode::from_str_name(value.as_str().unwrap_or_default()) - .ok_or(DecodeError::InvalidPayload)?; - Value::from(code as i32) - } else { - normalize_json_ids(value)? - }; - Ok((name, normalized)) - }) - .collect::, _>>() - .map(Value::Object), - Value::Array(values) => values - .into_iter() - .map(normalize_json_ids) - .collect::, _>>() - .map(Value::Array), - value => Ok(value), - } -} - -fn hex_bytes(bytes: &[u8]) -> String { - bytes.iter().map(|byte| format!("{byte:02x}")).collect() -} - -fn decoded_span( - span: Span, - resource_attributes: &BTreeMap, - scope_name: &str, - scope_version: &str, -) -> DecodedSpan { - let status = span.status.unwrap_or_default(); - DecodedSpan { - trace_id: hex_bytes(&span.trace_id), - span_id: hex_bytes(&span.span_id), - parent_span_id: hex_bytes(&span.parent_span_id), - trace_state: span.trace_state, - name: span.name, - kind: SpanKind::try_from(span.kind) - .unwrap_or(SpanKind::Unspecified) - .as_str_name() - .to_owned(), - resource_attributes: resource_attributes.clone(), - scope_name: scope_name.to_owned(), - scope_version: scope_version.to_owned(), - attributes: attributes(span.attributes), - start_ns: span.start_time_unix_nano, - end_ns: span.end_time_unix_nano, - status_code: StatusCode::try_from(status.code) - .unwrap_or(StatusCode::Unset) - .as_str_name() - .to_owned(), - status_message: status.message, - events: span - .events - .into_iter() - .map(|event| DecodedEvent { - name: event.name, - attributes: attributes(event.attributes), - }) - .collect(), - } -} - -fn attributes(values: Vec) -> BTreeMap { - values - .into_iter() - .map(|entry| { - ( - entry.key, - entry.value.as_ref().map(attribute_text).unwrap_or_default(), - ) - }) - .collect() -} - -fn attribute_text(value: &AnyValue) -> String { - match value.value.as_ref() { - Some(AttributeValue::StringValue(value)) => value.clone(), - Some(AttributeValue::BoolValue(value)) => value.to_string(), - Some(AttributeValue::IntValue(value)) => value.to_string(), - Some(AttributeValue::DoubleValue(value)) => { - serde_json::to_string(value).unwrap_or_default() - } - Some(AttributeValue::BytesValue(value)) => String::from_utf8_lossy(value).into_owned(), - Some(AttributeValue::ArrayValue(value)) => format!( - "[{}]", - value - .values - .iter() - .map(|value| serde_json::to_string(&attribute_text(value)).unwrap_or_default()) - .collect::>() - .join(", ") - ), - Some(AttributeValue::KvlistValue(value)) => format!( - "{{{}}}", - value - .values - .iter() - .map(|entry| format!( - "{}: {}", - serde_json::to_string(&entry.key).unwrap_or_default(), - serde_json::to_string( - &entry.value.as_ref().map(attribute_text).unwrap_or_default() - ) - .unwrap_or_default() - )) - .collect::>() - .join(", ") - ), - Some(AttributeValue::StringValueStrindex(value)) => value.to_string(), - None => String::new(), - } -} diff --git a/litellm-rust/crates/traces/src/otlp/attributes.rs b/litellm-rust/crates/traces/src/otlp/attributes.rs new file mode 100644 index 00000000000..50e063e4582 --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/attributes.rs @@ -0,0 +1,101 @@ +use std::{collections::BTreeMap, io::Write}; + +use opentelemetry_proto::tonic::common::v1::{ + AnyValue, KeyValue, any_value::Value as AttributeValue, +}; +use serde::{ + Serialize, Serializer, + ser::{SerializeMap, SerializeSeq}, +}; + +use super::limits::{Budget, MAX_ATTRIBUTES}; +use crate::DecodeError; + +struct AttributeWriter<'a> { + body: Vec, + budget: &'a mut Budget, +} + +impl Write for AttributeWriter<'_> { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.budget + .consume(bytes.len()) + .map_err(std::io::Error::other)?; + self.body.extend_from_slice(bytes); + Ok(bytes.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +pub(super) fn attributes( + values: Vec, + budget: &mut Budget, +) -> Result, DecodeError> { + if values.len() > MAX_ATTRIBUTES { + return Err(DecodeError::TooLarge); + } + values + .into_iter() + .map(|entry| { + budget.consume(entry.key.len() + 96)?; + let text = match entry.value { + Some(AnyValue { + value: Some(AttributeValue::StringValue(value)), + }) => { + budget.consume(value.len())?; + value + } + Some(AnyValue { + value: Some(AttributeValue::BytesValue(value)), + }) => { + budget.consume(value.len().saturating_mul(3))?; + String::from_utf8_lossy(&value).into_owned() + } + value => { + let mut writer = AttributeWriter { + body: Vec::new(), + budget, + }; + serde_json::to_writer(&mut writer, &AttributeJson(value.as_ref())) + .map_err(|_| DecodeError::TooLarge)?; + String::from_utf8(writer.body).map_err(|_| DecodeError::InvalidPayload)? + } + }; + Ok((entry.key, text)) + }) + .collect() +} + +struct AttributeJson<'a>(Option<&'a AnyValue>); + +impl Serialize for AttributeJson<'_> { + fn serialize(&self, serializer: S) -> Result { + match self.0.and_then(|value| value.value.as_ref()) { + Some(AttributeValue::StringValue(value)) => serializer.serialize_str(value), + Some(AttributeValue::BoolValue(value)) => serializer.serialize_bool(*value), + Some(AttributeValue::IntValue(value)) => serializer.serialize_i64(*value), + Some(AttributeValue::DoubleValue(value)) => serializer.serialize_f64(*value), + Some(AttributeValue::BytesValue(value)) => { + serializer.serialize_str(&String::from_utf8_lossy(value)) + } + Some(AttributeValue::ArrayValue(value)) => { + let mut sequence = serializer.serialize_seq(Some(value.values.len()))?; + for entry in &value.values { + sequence.serialize_element(&AttributeJson(Some(entry)))?; + } + sequence.end() + } + Some(AttributeValue::KvlistValue(value)) => { + let mut map = serializer.serialize_map(Some(value.values.len()))?; + for entry in &value.values { + map.serialize_entry(&entry.key, &AttributeJson(entry.value.as_ref()))?; + } + map.end() + } + Some(AttributeValue::StringValueStrindex(value)) => serializer.serialize_i32(*value), + None => serializer.serialize_unit(), + } + } +} diff --git a/litellm-rust/crates/traces/src/otlp/limits.rs b/litellm-rust/crates/traces/src/otlp/limits.rs new file mode 100644 index 00000000000..f6b56ccf12d --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/limits.rs @@ -0,0 +1,212 @@ +use std::fmt; + +use prost::encoding::{DecodeContext, WireType, decode_key, decode_varint, skip_field}; +use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor}; + +use crate::{DecodeError, Shared}; + +pub(super) const MAX_DEPTH: usize = 32; +pub(super) const MAX_NODES: usize = 65_536; +pub(super) const MAX_SPANS: usize = 4_096; +pub(super) const MAX_ATTRIBUTES: usize = 256; +pub(super) const MAX_EVENTS: usize = 256; +pub(super) const MAX_DECODED_SPAN_BYTES: usize = 16 * 1024 * 1024; + +pub(super) fn json_preflight(payload: &[u8]) -> Result<(), DecodeError> { + let mut nodes = 0; + let mut exceeded = false; + let mut decoder = serde_json::Deserializer::from_slice(payload); + let result = JsonBudget { + nodes: &mut nodes, + exceeded: &mut exceeded, + depth: 0, + } + .deserialize(&mut decoder) + .and_then(|()| decoder.end()); + if exceeded { + return Err(DecodeError::TooLarge); + } + result.map_err(|_| DecodeError::InvalidPayload) +} + +struct JsonBudget<'a> { + nodes: &'a mut usize, + exceeded: &'a mut bool, + depth: usize, +} + +impl<'de> DeserializeSeed<'de> for JsonBudget<'_> { + type Value = (); + + fn deserialize>(self, decoder: D) -> Result<(), D::Error> { + *self.nodes += 1; + if *self.nodes > MAX_NODES || self.depth > MAX_DEPTH { + *self.exceeded = true; + return Err(serde::de::Error::custom("OTLP structure exceeds budget")); + } + decoder.deserialize_any(self) + } +} + +impl<'de> Visitor<'de> for JsonBudget<'_> { + type Value = (); + + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("OTLP JSON") + } + fn visit_bool(self, _: bool) -> Result<(), E> { + Ok(()) + } + fn visit_i64(self, _: i64) -> Result<(), E> { + Ok(()) + } + fn visit_u64(self, _: u64) -> Result<(), E> { + Ok(()) + } + fn visit_f64(self, _: f64) -> Result<(), E> { + Ok(()) + } + fn visit_str(self, _: &str) -> Result<(), E> { + Ok(()) + } + fn visit_unit(self) -> Result<(), E> { + Ok(()) + } + + fn visit_seq>(self, mut sequence: A) -> Result<(), A::Error> { + while sequence + .next_element_seed(JsonBudget { + nodes: self.nodes, + exceeded: self.exceeded, + depth: self.depth + 1, + })? + .is_some() + {} + Ok(()) + } + + fn visit_map>(self, mut map: A) -> Result<(), A::Error> { + while map + .next_key_seed(JsonBudget { + nodes: self.nodes, + exceeded: self.exceeded, + depth: self.depth + 1, + })? + .is_some() + { + map.next_value_seed(JsonBudget { + nodes: self.nodes, + exceeded: self.exceeded, + depth: self.depth + 1, + })?; + } + Ok(()) + } +} + +#[derive(Clone, Copy)] +enum MessageKind { + Export, + ResourceSpans, + Resource, + ScopeSpans, + Scope, + Span, + Event, + Link, + Status, + KeyValue, + AnyValue, + Array, + KvList, +} + +impl MessageKind { + fn child(self, tag: u32) -> Option { + match (self, tag) { + (Self::Export, 1) => Some(Self::ResourceSpans), + (Self::ResourceSpans, 1) => Some(Self::Resource), + (Self::ResourceSpans, 2) => Some(Self::ScopeSpans), + (Self::Resource, 1) + | (Self::Scope, 3) + | (Self::Span, 9) + | (Self::Event, 3) + | (Self::Link, 4) + | (Self::KvList, 1) => Some(Self::KeyValue), + (Self::ScopeSpans, 1) => Some(Self::Scope), + (Self::ScopeSpans, 2) => Some(Self::Span), + (Self::Span, 11) => Some(Self::Event), + (Self::Span, 13) => Some(Self::Link), + (Self::Span, 15) => Some(Self::Status), + (Self::KeyValue, 2) | (Self::Array, 1) => Some(Self::AnyValue), + (Self::AnyValue, 5) => Some(Self::Array), + (Self::AnyValue, 6) => Some(Self::KvList), + _ => None, + } + } +} + +pub(super) fn protobuf_preflight(payload: &[u8]) -> Result<(), DecodeError> { + scan_message(payload, MessageKind::Export, 0, &mut 0) +} + +fn scan_message( + mut payload: &[u8], + kind: MessageKind, + depth: usize, + nodes: &mut usize, +) -> Result<(), DecodeError> { + if depth > MAX_DEPTH { + return Err(DecodeError::TooLarge); + } + while !payload.is_empty() { + *nodes += 1; + if *nodes > MAX_NODES { + return Err(DecodeError::TooLarge); + } + let (tag, wire) = decode_key(&mut payload).map_err(|_| DecodeError::InvalidPayload)?; + if let (WireType::LengthDelimited, Some(child)) = (wire, kind.child(tag)) { + let length = decode_varint(&mut payload).map_err(|_| DecodeError::InvalidPayload)?; + let length = usize::try_from(length).map_err(|_| DecodeError::InvalidPayload)?; + let (message, rest) = payload + .split_at_checked(length) + .ok_or(DecodeError::InvalidPayload)?; + scan_message(message, child, depth + 1, nodes)?; + payload = rest; + } else { + skip_field(wire, tag, &mut payload, DecodeContext::default()) + .map_err(|_| DecodeError::InvalidPayload)?; + } + } + Ok(()) +} + +pub(super) struct Budget { + remaining: usize, +} + +impl Budget { + pub(super) fn new(remaining: usize) -> Self { + Self { remaining } + } + + pub(super) fn clone_shared( + &mut self, + value: &Shared, + allocated_bytes: impl FnOnce(&T) -> usize, + ) -> Result, DecodeError> { + let cloned = value.clone(); + if !value.shares_storage_with(&cloned) { + self.consume(allocated_bytes(value))?; + } + Ok(cloned) + } + + pub(super) fn consume(&mut self, bytes: usize) -> Result<(), DecodeError> { + self.remaining = self + .remaining + .checked_sub(bytes) + .ok_or(DecodeError::TooLarge)?; + Ok(()) + } +} diff --git a/litellm-rust/crates/traces/src/otlp/mod.rs b/litellm-rust/crates/traces/src/otlp/mod.rs new file mode 100644 index 00000000000..fcc42082151 --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/mod.rs @@ -0,0 +1,42 @@ +mod attributes; +mod limits; +mod span; +mod wire; + +use serde::Serialize; +use std::collections::BTreeMap; + +use crate::{DecodeError, Shared}; + +#[derive(Serialize)] +pub struct DecodedEvent { + pub name: String, + pub attributes: BTreeMap, +} + +#[derive(Serialize)] +pub struct DecodedSpan { + pub trace_id: String, + pub span_id: String, + pub parent_span_id: String, + pub trace_state: String, + pub name: String, + pub kind: String, + pub resource_attributes: Shared>, + pub scope_name: Shared, + pub scope_version: Shared, + pub attributes: BTreeMap, + pub start_ns: u64, + pub end_ns: u64, + pub status_code: String, + pub status_message: String, + pub events: Vec, +} + +pub fn decode_otlp( + body: &[u8], + content_type: Option<&str>, +) -> Result, DecodeError> { + let request = wire::decode(body, content_type)?; + span::flatten(request) +} diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs new file mode 100644 index 00000000000..fa993f71e3c --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/span.rs @@ -0,0 +1,166 @@ +use std::collections::BTreeMap; + +use opentelemetry_proto::tonic::{ + collector::trace::v1::ExportTraceServiceRequest, + trace::v1::{ResourceSpans, ScopeSpans, Span, span::SpanKind, status::StatusCode}, +}; + +use super::{ + DecodedEvent, DecodedSpan, + attributes::attributes, + limits::{Budget, MAX_ATTRIBUTES, MAX_DECODED_SPAN_BYTES, MAX_EVENTS, MAX_SPANS}, +}; +use crate::{DecodeError, Shared}; + +pub(super) fn flatten(request: ExportTraceServiceRequest) -> Result, DecodeError> { + let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES); + let mut spans = Vec::new(); + for resource in request.resource_spans { + append_resource(resource, &mut budget, &mut spans)?; + } + Ok(spans) +} + +fn append_resource( + resource: ResourceSpans, + budget: &mut Budget, + spans: &mut Vec, +) -> Result<(), DecodeError> { + let attributes = Shared::new(attributes( + resource + .resource + .map(|resource| resource.attributes) + .unwrap_or_default(), + budget, + )?); + for scope in resource.scope_spans { + append_scope(scope, &attributes, budget, spans)?; + } + Ok(()) +} + +fn append_scope( + scope_spans: ScopeSpans, + resource: &Shared>, + budget: &mut Budget, + spans: &mut Vec, +) -> Result<(), DecodeError> { + let scope = scope_spans.scope.unwrap_or_default(); + if scope.attributes.len() > MAX_ATTRIBUTES { + return Err(DecodeError::TooLarge); + } + budget.consume(scope.name.len() + scope.version.len())?; + let scope_name: Shared = scope.name.into(); + let scope_version: Shared = scope.version.into(); + for span in scope_spans.spans { + if spans.len() >= MAX_SPANS { + return Err(DecodeError::TooLarge); + } + validate_span(&span)?; + budget.consume( + span.name.len() + + span.trace_state.len() + + span + .status + .as_ref() + .map_or(0, |status| status.message.len()) + + size_of::() + + 128, + )?; + spans.push(decoded_span( + span, + resource, + &scope_name, + &scope_version, + budget, + )?); + } + Ok(()) +} + +fn valid_id(value: &[u8], length: usize) -> bool { + value.len() == length && value.iter().any(|byte| *byte != 0) +} + +fn validate_span(span: &Span) -> Result<(), DecodeError> { + if !valid_id(&span.trace_id, 16) + || !valid_id(&span.span_id, 8) + || (!span.parent_span_id.is_empty() && !valid_id(&span.parent_span_id, 8)) + || span.start_time_unix_nano > i64::MAX as u64 + || span.end_time_unix_nano > i64::MAX as u64 + || span.end_time_unix_nano < span.start_time_unix_nano + || span + .links + .iter() + .any(|link| !valid_id(&link.trace_id, 16) || !valid_id(&link.span_id, 8)) + { + return Err(DecodeError::InvalidPayload); + } + if span.events.len() > MAX_EVENTS + || span.links.len() > MAX_EVENTS + || span.attributes.len() > MAX_ATTRIBUTES + || span + .links + .iter() + .any(|link| link.attributes.len() > MAX_ATTRIBUTES) + || span + .events + .iter() + .any(|event| event.attributes.len() > MAX_ATTRIBUTES) + { + return Err(DecodeError::TooLarge); + } + Ok(()) +} + +fn hex_bytes(bytes: &[u8]) -> String { + bytes.iter().map(|byte| format!("{byte:02x}")).collect() +} + +fn decoded_span( + span: Span, + resource_attributes: &Shared>, + scope_name: &Shared, + scope_version: &Shared, + budget: &mut Budget, +) -> Result { + let status = span.status.unwrap_or_default(); + Ok(DecodedSpan { + trace_id: hex_bytes(&span.trace_id), + span_id: hex_bytes(&span.span_id), + parent_span_id: hex_bytes(&span.parent_span_id), + trace_state: span.trace_state, + name: span.name, + kind: SpanKind::try_from(span.kind) + .unwrap_or(SpanKind::Unspecified) + .as_str_name() + .to_owned(), + resource_attributes: budget.clone_shared(resource_attributes, |attributes| { + attributes + .iter() + .map(|(key, value)| key.len() + value.len() + 96) + .sum() + })?, + scope_name: budget.clone_shared(scope_name, String::len)?, + scope_version: budget.clone_shared(scope_version, String::len)?, + attributes: attributes(span.attributes, budget)?, + start_ns: span.start_time_unix_nano, + end_ns: span.end_time_unix_nano, + status_code: StatusCode::try_from(status.code) + .unwrap_or(StatusCode::Unset) + .as_str_name() + .to_owned(), + status_message: status.message, + events: span + .events + .into_iter() + .map(|event| { + budget.consume(event.name.len() + 96)?; + Ok(DecodedEvent { + name: event.name, + attributes: attributes(event.attributes, budget)?, + }) + }) + .collect::, DecodeError>>()?, + }) +} diff --git a/litellm-rust/crates/traces/src/otlp/wire.rs b/litellm-rust/crates/traces/src/otlp/wire.rs new file mode 100644 index 00000000000..bac29ba49e4 --- /dev/null +++ b/litellm-rust/crates/traces/src/otlp/wire.rs @@ -0,0 +1,43 @@ +use opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest; +use prost::Message; + +use super::limits::{json_preflight, protobuf_preflight}; +use crate::DecodeError; + +#[derive(strum::EnumString)] +#[strum(ascii_case_insensitive)] +enum OtlpMediaType { + #[strum(serialize = "application/json")] + Json, + #[strum( + serialize = "application/x-protobuf", + serialize = "application/protobuf" + )] + Protobuf, +} + +pub(super) fn decode( + body: &[u8], + content_type: Option<&str>, +) -> Result { + let media_type = content_type + .unwrap_or("application/x-protobuf") + .split(';') + .next() + .unwrap_or_default() + .trim() + .parse::() + .map_err(|_| DecodeError::InvalidPayload)?; + + let request = match media_type { + OtlpMediaType::Json => { + json_preflight(body)?; + serde_json::from_slice(body).map_err(|_| DecodeError::InvalidPayload)? + } + OtlpMediaType::Protobuf => { + protobuf_preflight(body)?; + ExportTraceServiceRequest::decode(body).map_err(|_| DecodeError::InvalidPayload)? + } + }; + Ok(request) +} diff --git a/litellm-rust/crates/traces/src/shared.rs b/litellm-rust/crates/traces/src/shared.rs new file mode 100644 index 00000000000..dafd08b72dc --- /dev/null +++ b/litellm-rust/crates/traces/src/shared.rs @@ -0,0 +1,46 @@ +use std::ops::Deref; + +use serde::Serialize; + +type Storage = std::sync::Arc; + +#[derive(Clone, Debug, PartialEq, Serialize)] +#[serde(transparent)] +pub struct Shared(Storage); + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct SharedIdentity(usize); + +impl Shared { + pub fn new(value: T) -> Self { + Self(Storage::new(value)) + } + + pub fn identity(&self) -> SharedIdentity { + SharedIdentity(std::ptr::from_ref(self.as_ref()) as usize) + } + + pub fn shares_storage_with(&self, other: &Self) -> bool { + self.identity() == other.identity() + } +} + +impl From for Shared { + fn from(value: T) -> Self { + Self::new(value) + } +} + +impl AsRef for Shared { + fn as_ref(&self) -> &T { + self.0.as_ref() + } +} + +impl Deref for Shared { + type Target = T; + + fn deref(&self) -> &T { + self.as_ref() + } +} diff --git a/litellm-rust/crates/traces/src/sql.rs b/litellm-rust/crates/traces/src/sql.rs index 9acb8de0a7a..36d6e3b4521 100644 --- a/litellm-rust/crates/traces/src/sql.rs +++ b/litellm-rust/crates/traces/src/sql.rs @@ -8,6 +8,7 @@ pub enum ReadQuery { ListTraces, TraceSpans, SpanDetail, + SpanError, SpendByResponseIds, } @@ -17,6 +18,7 @@ impl ReadQuery { "list_traces" => Ok(Self::ListTraces), "trace_spans" => Ok(Self::TraceSpans), "span_detail" => Ok(Self::SpanDetail), + "span_error" => Ok(Self::SpanError), "spend_by_response_ids" => Ok(Self::SpendByResponseIds), _ => Err(Error::InvalidQuery), } @@ -27,6 +29,7 @@ impl ReadQuery { Self::ListTraces => include_str!("../query/list_traces.sql"), Self::TraceSpans => include_str!("../query/trace_spans.sql"), Self::SpanDetail => include_str!("../query/span_detail.sql"), + Self::SpanError => include_str!("../query/span_error.sql"), Self::SpendByResponseIds => include_str!("../query/spend_by_response_ids.sql"), } } diff --git a/litellm-rust/crates/traces/tests/insert.rs b/litellm-rust/crates/traces/tests/insert.rs index cba678152b9..9dcb9cddf1f 100644 --- a/litellm-rust/crates/traces/tests/insert.rs +++ b/litellm-rust/crates/traces/tests/insert.rs @@ -1,8 +1,107 @@ -use std::collections::BTreeMap; +use std::{ + collections::BTreeMap, + io::{BufRead, BufReader}, +}; -use litellm_traces::encode_rows; -use rstest::rstest; +use flate2::read::GzDecoder; +use litellm_http::Client; +use litellm_traces::{ + Connection, Error, InsertRow, InsertTable, Shared, encode_rows, insert_shared_rows, +}; +use rstest::{fixture, rstest}; use serde_json::{Value, json}; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{header, method}, +}; + +#[fixture] +fn shared_rows(#[default(16 * 1024)] attribute_bytes: usize) -> Vec { + let resource = Shared::new(json!({"shared": "x".repeat(attribute_bytes)})); + (0..1024) + .map(|index| { + BTreeMap::from([ + ("ResourceAttributes".into(), resource.clone()), + ("SpanId".into(), Shared::new(json!(format!("{index:016x}")))), + ("Timestamp".into(), Shared::new(json!(1))), + ]) + }) + .collect() +} + +#[rstest] +#[case::one_request(1)] +#[case::concurrent_requests(2)] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn shared_fanout_survives_gzip_insert_over_http( + shared_rows: Vec, + #[case] concurrency: usize, +) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(header("Content-Encoding", "gzip")) + .respond_with(ResponseTemplate::new(200)) + .expect(concurrency as u64) + .mount(&server) + .await; + let client = Client::no_redirect_for_test(); + let connection = Connection::parse(&server.uri()).unwrap(); + let expected_resource = shared_rows[0]["ResourceAttributes"].clone(); + let expected_count = shared_rows.len(); + let mut requests = tokio::task::JoinSet::new(); + for _ in 0..concurrency { + let client = client.clone(); + let connection = connection.clone(); + let rows = shared_rows.clone(); + requests.spawn(async move { + insert_shared_rows( + &client, + &connection, + "traces", + InsertTable::OtelTraces, + rows, + ) + .await + }); + } + while let Some(result) = requests.join_next().await { + result.unwrap().unwrap(); + } + let received = server.received_requests().await.unwrap(); + assert_eq!(received.len(), concurrency); + for request in received { + let decoder = GzDecoder::new(request.body.as_slice()); + let mut count = 0; + for (index, line) in BufReader::new(decoder).lines().enumerate() { + let row: Value = serde_json::from_str(&line.unwrap()).unwrap(); + assert_eq!(&row["ResourceAttributes"], expected_resource.as_ref()); + assert_eq!(row["SpanId"], format!("{index:016x}")); + assert_eq!(row["Timestamp"], "1970-01-01T00:00:00.000000001Z"); + assert!(row["EngineReceivedMs"].as_u64().unwrap() > 0); + count += 1; + } + assert_eq!(count, expected_count); + } +} + +#[rstest] +#[tokio::test] +async fn shared_fanout_over_insert_limit_never_reaches_http( + #[with(64 * 1024)] shared_rows: Vec, +) { + let server = MockServer::start().await; + let connection = Connection::parse(&server.uri()).unwrap(); + let result = insert_shared_rows( + &Client::no_redirect_for_test(), + &connection, + "traces", + InsertTable::OtelTraces, + shared_rows, + ) + .await; + assert!(matches!(result, Err(Error::InsertTooLarge))); + assert!(server.received_requests().await.unwrap().is_empty()); +} #[rstest] #[case::span("Timestamp", json!(1_234_567_890), json!("1970-01-01T00:00:01.23456789Z"))] diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index cc8fe51a469..01e8982423f 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -835,3 +835,148 @@ async fn lens_content_keeps_output_visible_after_long_input( assert_eq!(recovered, original); Ok(()) } + +#[rstest] +#[case::ascii(10, format!("ParentCommand: {}", "x".repeat(460_000)))] +#[case::multibyte(1_000, "\u{1f9ea}".repeat(1_024))] +#[case::escaped(1_000, "\0\n\"\\".repeat(1_024))] +#[tokio::test] +async fn trace_error_previews_preserve_paginated_diagnostics( + #[future(awt)] database: TestResult, + #[case] span_count: usize, + #[case] message: String, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let rows = (0..span_count) + .map(|index| { + serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp + index as i64, "TraceId": "diagnostic-trace", + "SpanId": format!("span-{index}"), "SpanName": "tool", + "StatusCode": "STATUS_CODE_ERROR", "StatusMessage": message, + })) + }) + .collect::>, _>>()?; + insert_rows(&database, "otel_traces", rows).await?; + let reader = Connection::reader(&database.url, "trace_test")?; + let mut parameters = BTreeMap::from([ + ( + "trace_id".into(), + Parameter::Text("diagnostic-trace".into()), + ), + ("team_ids".into(), Parameter::Strings(vec![])), + ("api_key_hash".into(), Parameter::Text(String::new())), + ("trace_ref".into(), Parameter::Text(String::new())), + ]); + let body = execute_named_read( + &database.client, + &reader, + ReadQuery::TraceSpans, + ¶meters, + ) + .await?; + let response: serde_json::Value = serde_json::from_str(&body)?; + let spans = response["data"].as_array().expect("trace spans"); + assert_eq!(spans.len(), span_count); + let prefix: String = message.chars().take(128).collect(); + assert!(!prefix.is_empty()); + assert!( + spans + .iter() + .all(|span| span["status_message"] == prefix && span["error_truncated"] == 1) + ); + parameters.insert("span_id".into(), Parameter::Text("span-0".into())); + parameters.insert("error_version".into(), Parameter::Text(String::new())); + let mut recovered = String::new(); + loop { + parameters.insert( + "error_offset".into(), + Parameter::Integer(recovered.chars().count() as i64), + ); + let body = execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters) + .await?; + assert!(body.len() < 128 * 1024); + let response: serde_json::Value = serde_json::from_str(&body)?; + let chunk = response["data"][0]["message"] + .as_str() + .expect("diagnostic chunk"); + assert!(!chunk.is_empty()); + recovered.push_str(chunk); + let version = response["data"][0]["version"] + .as_str() + .expect("diagnostic version"); + parameters.insert("error_version".into(), Parameter::Text(version.into())); + if recovered.chars().count() >= message.chars().count() { + break; + } + } + assert_eq!(recovered, message); + parameters.insert( + "api_key_hash".into(), + Parameter::Text("unrelated-key".into()), + ); + let denied = + execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?; + assert_eq!( + serde_json::from_str::(&denied)?["data"], + serde_json::json!([]) + ); + Ok(()) +} + +#[rstest] +#[case::different_start(1, 0)] +#[case::different_receive(0, 1)] +#[case::tied_timestamps(0, 0)] +#[tokio::test] +async fn duplicate_span_preview_matches_diagnostic( + #[future(awt)] database: TestResult, + #[case] start_delta: i64, + #[case] receive_delta: i64, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let message = "a".repeat(200); + let rows = [ + (start_delta, receive_delta, "z".repeat(200)), + (0, 0, message.clone()), + ] + .into_iter() + .map(|(start_delta, receive_delta, message)| { + serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp + start_delta, "EngineReceivedMs": 100 + receive_delta, + "TraceId": "duplicate-trace", "SpanId": "duplicate-span", "StatusMessage": message, + })) + }) + .collect::>, _>>()?; + insert_rows(&database, "otel_traces", rows).await?; + let reader = Connection::reader(&database.url, "trace_test")?; + let parameters = BTreeMap::from([ + ("trace_id".into(), Parameter::Text("duplicate-trace".into())), + ("span_id".into(), Parameter::Text("duplicate-span".into())), + ("team_ids".into(), Parameter::Strings(vec![])), + ("api_key_hash".into(), Parameter::Text(String::new())), + ("trace_ref".into(), Parameter::Text(String::new())), + ("error_version".into(), Parameter::Text(String::new())), + ("error_offset".into(), Parameter::Integer(0)), + ]); + let preview = execute_named_read( + &database.client, + &reader, + ReadQuery::TraceSpans, + ¶meters, + ) + .await?; + let diagnostic = + execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?; + let preview: serde_json::Value = serde_json::from_str(&preview)?; + let diagnostic: serde_json::Value = serde_json::from_str(&diagnostic)?; + assert_eq!(preview["data"].as_array().unwrap().len(), 1); + assert_eq!(preview["data"][0]["status_message"], message[..128]); + assert_eq!(diagnostic["data"][0]["message"], message); + Ok(()) +} diff --git a/litellm-rust/crates/traces/tests/otlp.rs b/litellm-rust/crates/traces/tests/otlp.rs index 002ba159ef9..8aa2cbedeb3 100644 --- a/litellm-rust/crates/traces/tests/otlp.rs +++ b/litellm-rust/crates/traces/tests/otlp.rs @@ -1,33 +1,19 @@ -use flate2::{Compression, write::GzEncoder}; +use litellm_traces::Shared; use litellm_traces::decode_otlp; use rstest::rstest; -use std::io::Write; const FIXTURE: &[u8] = include_bytes!( "../../../../tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json" ); #[rstest] -#[case::json(FIXTURE, Some("application/json"), None)] -#[case::gzip_json(FIXTURE, Some("application/json"), Some("gzip"))] -fn decodes_neutral_spans( - #[case] body: &[u8], - #[case] content_type: Option<&str>, - #[case] content_encoding: Option<&str>, -) { - let payload = if content_encoding == Some("gzip") { - let mut encoder = GzEncoder::new(Vec::new(), Compression::default()); - encoder.write_all(body).expect("gzip input"); - encoder.finish().expect("gzip payload") - } else { - body.to_vec() - }; - let spans = decode_otlp(&payload, content_type, content_encoding, 8 * 1024 * 1024) - .expect("valid OTLP export"); +#[case::json(FIXTURE, Some("application/json"))] +fn decodes_neutral_spans(#[case] body: &[u8], #[case] content_type: Option<&str>) { + let spans = decode_otlp(body, content_type).expect("valid OTLP export"); assert_eq!(spans.len(), 6); assert_eq!(spans[0].trace_id, "4bad42b84e9de3ba46fc870185f8f023"); assert_eq!(spans[0].resource_attributes["service.name"], "agent-demo"); - assert_eq!(spans[0].scope_name, "langsmith"); + assert_eq!(spans[0].scope_name.as_ref(), "langsmith"); assert!( spans .iter() @@ -36,12 +22,322 @@ fn decodes_neutral_spans( } #[rstest] -#[case::invalid(b"not protobuf", None, 8 * 1024 * 1024)] -#[case::too_large(FIXTURE, Some("application/json"), 1)] -fn rejects_invalid_or_oversized_payload( - #[case] body: &[u8], - #[case] content_type: Option<&str>, - #[case] limit: usize, -) { - assert!(decode_otlp(body, content_type, None, limit).is_err()); +fn accepts_trace_larger_than_eight_mib(mut span: opentelemetry_proto::tonic::trace::v1::Span) { + use prost::Message; + + span.name = "x".repeat(9 * 1024 * 1024); + let body = request_with(span).encode_to_vec(); + let decoded = decode_otlp(&body, None).expect("16 MiB default accepts a 9 MiB trace"); + assert_eq!(decoded[0].name.len(), 9 * 1024 * 1024); +} + +#[rstest] +fn rejects_invalid_payload() { + assert!(decode_otlp(b"not protobuf", None).is_err()); +} + +#[rstest] +fn decoder_does_not_enforce_the_http_body_limit() { + let body = format!("{{\"ignored\":\"{}\"}}", "x".repeat(16 * 1024 * 1024 + 1)); + assert!( + decode_otlp(body.as_bytes(), Some("application/json")) + .unwrap() + .is_empty() + ); +} + +fn request_with( + span: opentelemetry_proto::tonic::trace::v1::Span, +) -> opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest { + use opentelemetry_proto::tonic::{ + collector::trace::v1::ExportTraceServiceRequest, + trace::v1::{ResourceSpans, ScopeSpans}, + }; + ExportTraceServiceRequest { + resource_spans: vec![ResourceSpans { + scope_spans: vec![ScopeSpans { + spans: vec![span], + ..Default::default() + }], + ..Default::default() + }], + } +} + +#[rstest::fixture] +fn span() -> opentelemetry_proto::tonic::trace::v1::Span { + opentelemetry_proto::tonic::trace::v1::Span { + trace_id: vec![1; 16], + span_id: vec![2; 8], + start_time_unix_nano: 1, + end_time_unix_nano: 2, + ..Default::default() + } +} + +#[rstest] +fn standard_json_and_protobuf_preserve_the_same_identifiers( + span: opentelemetry_proto::tonic::trace::v1::Span, +) { + use prost::Message; + let request = request_with(span); + let json = serde_json::to_vec(&request).unwrap(); + let binary = request.encode_to_vec(); + let json_spans = decode_otlp(&json, Some("application/json; charset=utf-8")).unwrap(); + let binary_spans = decode_otlp(&binary, Some("application/x-protobuf")).unwrap(); + assert_eq!( + serde_json::to_value(&json_spans).unwrap(), + serde_json::to_value(&binary_spans).unwrap() + ); + assert_eq!(json_spans[0].trace_id, "01".repeat(16)); + assert_eq!(json_spans[0].span_id, "02".repeat(8)); +} + +#[rstest] +#[case::json("APPLICATION/JSON; charset=utf-8", b"{}")] +#[case::protobuf("application/x-protobuf; charset=binary", b"")] +#[case::protobuf_alias("APPLICATION/PROTOBUF", b"")] +fn supported_content_types_select_the_decoder(#[case] content_type: &str, #[case] body: &[u8]) { + assert!(decode_otlp(body, Some(content_type)).is_ok()); +} + +#[rstest] +#[case::missing_content_type(None)] +#[case::unsupported_content_type(Some("text/plain"))] +fn content_type_defaults_to_protobuf_and_rejects_unknown_values( + #[case] content_type: Option<&str>, +) { + let result = decode_otlp(b"", content_type); + assert_eq!(result.is_ok(), content_type.is_none()); +} + +#[rstest] +#[case::short_trace(vec![1; 15], vec![2;8], 1, 2)] +#[case::zero_trace(vec![0; 16], vec![2;8], 1, 2)] +#[case::short_span(vec![1; 16], vec![2;7], 1, 2)] +#[case::timestamp_overflow(vec![1;16], vec![2;8], i64::MAX as u64 + 1, i64::MAX as u64 + 1)] +#[case::negative_duration(vec![1;16], vec![2;8], 3, 2)] +fn rejects_ids_and_timestamps_that_cannot_be_stored( + #[case] trace_id: Vec, + #[case] span_id: Vec, + #[case] start: u64, + #[case] end: u64, +) { + use prost::Message; + let span = opentelemetry_proto::tonic::trace::v1::Span { + trace_id, + span_id, + start_time_unix_nano: start, + end_time_unix_nano: end, + ..Default::default() + }; + assert!(matches!( + decode_otlp(&request_with(span).encode_to_vec(), None), + Err(litellm_traces::DecodeError::InvalidPayload) + )); +} + +#[rstest] +fn resource_fanout_shares_one_allocation(span: opentelemetry_proto::tonic::trace::v1::Span) { + use opentelemetry_proto::tonic::{ + common::v1::{AnyValue, KeyValue, any_value::Value}, + resource::v1::Resource, + }; + use prost::Message; + let mut request = request_with(span.clone()); + request.resource_spans[0].resource = Some(Resource { + attributes: vec![KeyValue { + key: "shared".into(), + value: Some(AnyValue { + value: Some(Value::StringValue("x".repeat(16 * 1024))), + }), + ..Default::default() + }], + ..Default::default() + }); + request.resource_spans[0].scope_spans[0].spans = vec![span; 1024]; + let second_scope = request.resource_spans[0].scope_spans[0].clone(); + request.resource_spans[0].scope_spans.push(second_scope); + request + .resource_spans + .push(request.resource_spans[0].clone()); + let body = request.encode_to_vec(); + let decoded = decode_otlp(&body, None).expect("shared resources do not expand with span count"); + assert_eq!(decoded.len(), 4096); + assert!(decoded[..2048].iter().all(|span| { + Shared::shares_storage_with(&span.resource_attributes, &decoded[0].resource_attributes) + })); + assert!(!Shared::shares_storage_with( + &decoded[0].resource_attributes, + &decoded[2048].resource_attributes + )); + assert_eq!( + *decoded[0].resource_attributes, + *decoded[2048].resource_attributes + ); +} + +#[rstest] +fn nested_values_are_serialized_once(span: opentelemetry_proto::tonic::trace::v1::Span) { + use opentelemetry_proto::tonic::common::v1::{ + AnyValue, ArrayValue, KeyValue, any_value::Value, + }; + use prost::Message; + let nested = (0..8).fold( + AnyValue { + value: Some(Value::StringValue("quoted \"value\"".into())), + }, + |child, _| AnyValue { + value: Some(Value::ArrayValue(ArrayValue { + values: vec![child], + })), + }, + ); + let mut request = request_with(span); + request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue { + key: "nested".into(), + value: Some(nested), + ..Default::default() + }]; + let spans = decode_otlp(&request.encode_to_vec(), None).unwrap(); + let expected = (0..8).fold(serde_json::json!("quoted \"value\""), |child, _| { + serde_json::json!([child]) + }); + assert_eq!( + serde_json::from_str::(&spans[0].attributes["nested"]).unwrap(), + expected + ); + assert!(spans[0].attributes["nested"].len() < 64); +} + +#[rstest] +#[case::nesting(format!("{}0{}", "[".repeat(40), "]".repeat(40)).into_bytes())] +#[case::nodes(format!("[{}]", vec!["0"; 65537].join(",")).into_bytes())] +fn rejects_json_structure_before_building_a_tree(#[case] body: Vec) { + assert!(matches!( + decode_otlp(&body, Some("application/json")), + Err(litellm_traces::DecodeError::TooLarge) + )); +} + +#[rstest] +#[case::depth(40, 1)] +#[case::nodes(0, 65537)] +fn protobuf_preflight_rejects_expansion_before_prost_allocates( + span: opentelemetry_proto::tonic::trace::v1::Span, + #[case] depth: usize, + #[case] count: usize, +) { + use opentelemetry_proto::tonic::common::v1::{ + AnyValue, ArrayValue, KeyValue, any_value::Value, + }; + use prost::Message; + let value = (0..depth).fold( + AnyValue { + value: Some(Value::BoolValue(true)), + }, + |child, _| AnyValue { + value: Some(Value::ArrayValue(ArrayValue { + values: vec![child], + })), + }, + ); + let mut request = request_with(span); + request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue { + key: "deep".into(), + value: Some(value), + ..Default::default() + }]; + request.resource_spans = vec![request.resource_spans[0].clone(); count]; + let body = request.encode_to_vec(); + assert!(matches!( + decode_otlp(&body, None), + Err(litellm_traces::DecodeError::TooLarge) + )); +} + +#[rstest] +fn scope_fanout_shares_name_and_version(span: opentelemetry_proto::tonic::trace::v1::Span) { + use opentelemetry_proto::tonic::common::v1::InstrumentationScope; + use prost::Message; + let mut request = request_with(span.clone()); + request.resource_spans[0].scope_spans[0].scope = Some(InstrumentationScope { + name: "n".repeat(16 * 1024), + version: "v".repeat(16 * 1024), + ..Default::default() + }); + request.resource_spans[0].scope_spans[0].spans = vec![span; 1024]; + let decoded = decode_otlp(&request.encode_to_vec(), None).unwrap(); + assert!( + decoded + .iter() + .all(|span| Shared::shares_storage_with(&span.scope_name, &decoded[0].scope_name)) + ); + assert!( + decoded.iter().all(|span| Shared::shares_storage_with( + &span.scope_version, + &decoded[0].scope_version + )) + ); + assert_eq!(decoded[0].scope_name.len(), 16 * 1024); + assert_eq!(decoded[0].scope_version.len(), 16 * 1024); +} + +#[rstest] +fn unique_attribute_expansion_still_respects_decoded_budget( + span: opentelemetry_proto::tonic::trace::v1::Span, +) { + use opentelemetry_proto::tonic::common::v1::{AnyValue, KeyValue, any_value::Value}; + use prost::Message; + let mut request = request_with(span.clone()); + request.resource_spans[0].scope_spans[0].spans = (0..1024) + .map(|index| { + let mut span = span.clone(); + span.attributes = vec![KeyValue { + key: "unique".into(), + value: Some(AnyValue { + value: Some(Value::StringValue(format!( + "{index:04}{}", + "x".repeat(16_300) + ))), + }), + ..Default::default() + }]; + span + }) + .collect(); + let body = request.encode_to_vec(); + assert!(body.len() < 16 * 1024 * 1024); + assert!(matches!( + decode_otlp(&body, None), + Err(litellm_traces::DecodeError::TooLarge) + )); +} + +#[rstest] +fn escaped_attribute_expansion_is_bounded_below_four_mib( + span: opentelemetry_proto::tonic::trace::v1::Span, +) { + use opentelemetry_proto::tonic::common::v1::{ + AnyValue, ArrayValue, KeyValue, any_value::Value, + }; + use prost::Message; + let mut request = request_with(span); + request.resource_spans[0].scope_spans[0].spans[0].attributes = vec![KeyValue { + key: "escaped".into(), + value: Some(AnyValue { + value: Some(Value::ArrayValue(ArrayValue { + values: vec![AnyValue { + value: Some(Value::StringValue("\0".repeat(3 * 1024 * 1024))), + }], + })), + }), + ..Default::default() + }]; + let body = request.encode_to_vec(); + assert!(body.len() < 4 * 1024 * 1024); + assert!(matches!( + decode_otlp(&body, None), + Err(litellm_traces::DecodeError::TooLarge) + )); } diff --git a/litellm-rust/crates/traces/tests/shared.rs b/litellm-rust/crates/traces/tests/shared.rs new file mode 100644 index 00000000000..2e76e6321db --- /dev/null +++ b/litellm-rust/crates/traces/tests/shared.rs @@ -0,0 +1,23 @@ +use litellm_traces::Shared; +use rstest::rstest; + +#[rstest] +fn clones_preserve_values_and_serialize_transparently() { + let original = Shared::new(vec!["value".to_owned()]); + let cloned = original.clone(); + assert_eq!(cloned.as_ref(), original.as_ref()); + assert_eq!( + serde_json::to_value(&cloned).unwrap(), + serde_json::json!(["value"]) + ); +} + +#[rstest] +fn clones_share_storage_without_merging_equal_values() { + let original = Shared::new("value".to_owned()); + let cloned = original.clone(); + let equal = Shared::new("value".to_owned()); + assert!(original.shares_storage_with(&cloned)); + assert!(!original.shares_storage_with(&equal)); + assert_eq!(*original, *equal); +} diff --git a/litellm/constants.py b/litellm/constants.py index 7e1e63a112b..76ab419272f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -52,10 +52,10 @@ CLICKHOUSE_MAX_BUFFERED_ROWS: Final = get_env_int("CLICKHOUSE_MAX_BUFFERED_ROWS" CLICKHOUSE_MAX_RETRIES: Final = get_env_int("CLICKHOUSE_MAX_RETRIES", 3) AGENT_TRACING_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_RETENTION_DAYS", 30) AGENT_TRACING_SPEND_LOG_RETENTION_DAYS: Final = get_env_int("AGENT_TRACING_SPEND_LOG_RETENTION_DAYS", 90) -OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 8 * 1024 * 1024) +OTLP_MAX_BODY_BYTES: Final = get_env_int("OTLP_MAX_BODY_BYTES", 16 * 1024 * 1024) OTLP_MAX_ATTRIBUTE_VALUE_BYTES: Final = get_env_int("OTLP_MAX_ATTRIBUTE_VALUE_BYTES", 64 * 1024) OTLP_RETRY_AFTER_SECONDS: Final = get_env_int("OTLP_RETRY_AFTER_SECONDS", 2) -OTLP_OFFLOAD_DECODE_BYTES: Final = get_env_int("OTLP_OFFLOAD_DECODE_BYTES", 256 * 1024) +OTLP_MAX_CONCURRENT_INGESTS: Final = get_env_int("OTLP_MAX_CONCURRENT_INGESTS", 2) AGENT_TRACING_INPUT_PREVIEW_CHARS: Final = get_env_int("AGENT_TRACING_INPUT_PREVIEW_CHARS", 240) AGENT_TRACING_LIST_PAGE_SIZE: Final = get_env_int("AGENT_TRACING_LIST_PAGE_SIZE", 50) DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10)) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 6e5114b2f87..aa4d6a39f25 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -6,6 +6,7 @@ from typing import Annotated, Any, Final, Literal, Union, get_args, get_origin import orjson from fastapi import Request, UploadFile, status +from starlette._utils import get_route_path from typing_extensions import NotRequired, ReadOnly, Required, assert_never from litellm._logging import verbose_proxy_logger @@ -167,6 +168,10 @@ def _parse_binary_body(body: bytes) -> dict: return {} +def is_otlp_trace_request(request: Request) -> bool: + return request.method == "POST" and get_route_path(request.scope) == "/v1/traces" + + async def _read_request_body(request: Request | None) -> dict: """ Safely read the request body and parse it as JSON. @@ -181,6 +186,9 @@ async def _read_request_body(request: Request | None) -> dict: if request is None: return {} + if is_otlp_trace_request(request): + return {} + # Check if we already read and parsed the body _cached_request_body: Final[dict | None] = _safe_get_request_parsed_body(request=request) if _cached_request_body is not None: @@ -189,11 +197,7 @@ async def _read_request_body(request: Request | None) -> dict: _request_headers: Final[dict] = _safe_get_request_headers(request=request) content_type: Final = _request_headers.get("content-type", "") - if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES or ( - request.scope.get("path") == "/v1/traces" - and request.scope.get("method") == "POST" - and _request_headers.get("content-encoding", "").lower() == "gzip" - ): + if _normalize_media_type(content_type) in _BINARY_CONTENT_TYPES: parsed_body = _parse_binary_body(await request.body()) elif _is_form_content_type(content_type): try: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index cf09bdbef9b..9abf949ec1e 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -720,6 +720,9 @@ try: except ImportError: build_billing_metrics_recorder = None shutdown_billing_metrics_recorder = None +from fastapi.exception_handlers import http_exception_handler +from starlette.exceptions import HTTPException as StarletteHTTPException + from litellm.proxy import tracing_endpoints from litellm.proxy.middleware.admission_control_middleware import ( AdmissionControlMiddleware, @@ -1895,6 +1898,9 @@ async def openai_exception_handler(request: Request, exc: ProxyException): ) status_code: Final = int(exc.code) if exc.code else status.HTTP_500_INTERNAL_SERVER_ERROR _close_dangling_otel_server_span(request, status_code, exc=exc) + otlp_response: Final = tracing_endpoints.otlp_error_response(request, status_code, headers) + if otlp_response is not None: + return otlp_response return JSONResponse( status_code=status_code, content={"error": error_dict}, @@ -1902,6 +1908,15 @@ async def openai_exception_handler(request: Request, exc: ProxyException): ) +@app.exception_handler(StarletteHTTPException) +async def otlp_http_exception_handler(request: Request, exc: StarletteHTTPException) -> Response: + response: Final = tracing_endpoints.otlp_error_response(request, exc.status_code, exc.headers) + if response is not None: + _close_dangling_otel_server_span(request, exc.status_code, exc=exc) + return response + return await http_exception_handler(request, exc) + + def _log_model_access_denial(exc: ProxyException) -> None: if not isinstance(exc, ModelAccessDeniedProxyException): return @@ -2023,6 +2038,9 @@ async def otel_request_validation_exception_handler(request: Request, exc: Reque _close_dangling_otel_server_span(request, problem.status, exc=public_exc) return problem_response(problem) _close_dangling_otel_server_span(request, 422, exc=public_exc) + otlp_response: Final = tracing_endpoints.otlp_error_response(request, 422) + if otlp_response is not None: + return otlp_response return JSONResponse(status_code=422, content={"detail": public_errors}) @@ -2046,6 +2064,9 @@ async def otel_unhandled_exception_handler(request: Request, exc: Exception): ) ) _close_dangling_otel_server_span(request, 500, exc=exc) + otlp_response: Final = tracing_endpoints.otlp_error_response(request, 500) + if otlp_response is not None: + return otlp_response return JSONResponse( status_code=500, content={ diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index bd885282859..46d29c50c1b 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -8,14 +8,18 @@ GET /v1/traces/{trace_id}/spans/{span_id} SpanDetail """ import time +from collections.abc import Mapping from dataclasses import dataclass +from http.client import responses +from types import MappingProxyType from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response -from litellm.constants import OTLP_MAX_BODY_BYTES, OTLP_RETRY_AFTER_SECONDS +from litellm.constants import 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.common_utils.http_parsing_utils import is_otlp_trace_request from litellm.proxy.tracing_runtime import provide_receiver, require_receiver from litellm.tracing import ( Tenant, @@ -23,7 +27,7 @@ from litellm.tracing import ( TracingPayloadTooLargeError, ) from litellm.tracing.decode import InvalidOTLPPayloadError, encode_otlp_response -from litellm.tracing.types import SpanDetail, Trace, TracePage, TraceScope +from litellm.tracing.types import SpanDetail, SpanErrorPage, Trace, TracePage, TraceScope router = APIRouter(tags=["agent tracing"]) @@ -66,13 +70,25 @@ async def provide_trace_access( return TraceAccessContext(tracing, None, tenant) -async def _read_otlp_body(request: Request) -> bytes: - body: Final = bytearray() - async for chunk in request.stream(): - if len(body) + len(chunk) > OTLP_MAX_BODY_BYTES: - raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") - body.extend(chunk) - return bytes(body) +def otlp_error_response( + request: Request, status_code: int, headers: Mapping[str, str] | None = None +) -> Response | None: + if not is_otlp_trace_request(request): + return None + body, media_type = encode_otlp_response( + request.headers.get("content-type"), responses.get(status_code, "Trace request failed") + ) + return Response(content=body, status_code=status_code, media_type=media_type, headers=headers) + + +def _otlp_error(content_type: str | None, status_code: int, message: str, retry: bool = False) -> Response: + body, media_type = encode_otlp_response(content_type, message) + return Response( + content=body, + status_code=status_code, + media_type=media_type, + headers=MappingProxyType({"Retry-After": str(OTLP_RETRY_AFTER_SECONDS)}) if retry else None, + ) @router.post("/v1/traces", include_in_schema=False) @@ -80,24 +96,23 @@ async def ingest_otlp_traces( request: Request, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], ) -> Response: - tracing, tenant = context.writer() content_type: Final = request.headers.get("content-type") try: + tracing, tenant = context.writer() await tracing.ingest( - body=await _read_otlp_body(request), + body=request.stream(), content_type=content_type, content_encoding=request.headers.get("content-encoding"), tenant=tenant, ) except TracingPayloadTooLargeError as e: - raise HTTPException(status_code=413, detail=str(e)) + return _otlp_error(content_type, 413, str(e)) except InvalidOTLPPayloadError as error: - raise HTTPException(status_code=400, detail=str(error)) from error + return _otlp_error(content_type, 400, str(error)) except RuntimeError: - raise HTTPException( - status_code=503, - headers={"Retry-After": str(OTLP_RETRY_AFTER_SECONDS)}, - ) + return _otlp_error(content_type, 503, "Trace ingestion is temporarily unavailable", retry=True) + except HTTPException as error: + return _otlp_error(content_type, error.status_code, str(error.detail)) body, media_type = encode_otlp_response(content_type) return Response(content=body, media_type=media_type) @@ -147,3 +162,21 @@ async def get_agent_trace_span( if span is None: raise HTTPException(status_code=404, detail=f"Span {span_id} not found") return span + + +@router.get("/v1/traces/{trace_id}/spans/{span_id}/error", response_model=SpanErrorPage) +async def get_agent_trace_span_error( + trace_id: str, + span_id: str, + context: Annotated[TraceAccessContext, Depends(provide_trace_access)], + trace_ref: Annotated[str, Query()] = "", + cursor: Annotated[str | None, Query(max_length=512)] = None, +) -> SpanErrorPage: + try: + tracing, scope = context.reader() + page: Final = await tracing.get_span_error(trace_id, span_id, scope, trace_ref, cursor) + except ValueError as error: + raise HTTPException(status_code=400, detail=str(error)) from error + if page is None: + raise HTTPException(status_code=404, detail="Span diagnostic not found or no longer available") + return page diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 206c0f78ed8..a3d8ba0e582 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -21,15 +21,14 @@ class RustUpstreamError(Exception): ... class ForkedAfterNativeRuntimeStarted(RuntimeError): ... class ProcessReservedForForking(RuntimeError): ... -def trace_decode_otlp( - body: bytes, content_type: str | None, content_encoding: str | None, max_decompressed_bytes: int -) -> list[DecodedSpan]: ... +def trace_decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan]: ... +def trace_encode_error(message: str) -> bytes: ... @final class NativeTraceStorage: def __new__(cls, database: str, url: str, reader_url: str | None = None) -> NativeTraceStorage: ... def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Future[None]: ... - def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Future[None]: ... + def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Future[None]: ... def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... def query(self, query: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... @@ -353,6 +352,7 @@ __all__ = [ "reserve_process_for_forking", "responses", "trace_decode_otlp", + "trace_encode_error", "transcription", ] diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index 1c20e408709..6724db41ad3 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -31,7 +31,7 @@ class DecodedSpan(TypedDict): events: ReadOnly[list[DecodedEvent]] -ReadQueryName = Literal["list_traces", "trace_spans", "span_detail", "spend_by_response_ids"] +ReadQueryName = Literal["list_traces", "trace_spans", "span_detail", "span_error", "spend_by_response_ids"] class NativeStore(Protocol): @@ -39,7 +39,7 @@ class NativeStore(Protocol): def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Awaitable[None]: ... - def insert_rows(self, table: str, rows: Sequence[Mapping[str, JsonValue]]) -> Awaitable[None]: ... + def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Awaitable[None]: ... def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Awaitable[str]: ... @@ -53,17 +53,16 @@ class NativeTraces(Protocol): self, body: bytes, content_type: str | None, - content_encoding: str | None, - max_decompressed_bytes: int, ) -> list[DecodedSpan]: ... + def trace_encode_error(self, message: str) -> bytes: ... + class QueryResponse(BaseModel): model_config = ConfigDict(frozen=True) data: list[dict[str, JsonValue]] -INSERT_ROWS: Final = TypeAdapter(list[dict[str, JsonValue]]) QUERY_PARAMETERS: Final = TypeAdapter(dict[str, str | int | list[str]]) @@ -74,10 +73,14 @@ def _native() -> NativeTraces: return cast(NativeTraces, native) # cast-ok: the native extension is validated against this protocol at call sites -def decode_otlp( - body: bytes, content_type: str | None, content_encoding: str | None, max_decompressed_bytes: int -) -> list[DecodedSpan]: - return _native().trace_decode_otlp(body, content_type, content_encoding, max_decompressed_bytes) +def decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan]: + return _native().trace_decode_otlp(body, content_type) + + +def encode_error(message: str) -> bytes: + if get_native_bridge() is None: + return b"" + return _native().trace_encode_error(message) class ClickHouseStorage: @@ -88,7 +91,7 @@ class ClickHouseStorage: await self._native.ensure_schema(trace_retention_days, spend_log_retention_days) async def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> None: - await self._native.insert_rows(table, INSERT_ROWS.validate_python(rows)) + await self._native.insert_rows(table, rows) async def query( self, name: ReadQueryName, parameters: Mapping[str, object] | None = None diff --git a/litellm/tracing/decode.py b/litellm/tracing/decode.py index ce46bed5e22..d8b5f70de68 100644 --- a/litellm/tracing/decode.py +++ b/litellm/tracing/decode.py @@ -8,37 +8,43 @@ Pure functions, no I/O. Two steps: Deep Agents), OTEL GenAI semconv, OpenInference. """ +import gzip import json +import zlib from collections.abc import Mapping +from dataclasses import dataclass +from io import BytesIO from itertools import accumulate from types import MappingProxyType from typing import Final from pydantic import JsonValue, TypeAdapter, ValidationError -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.constants import OTLP_MAX_ATTRIBUTE_VALUE_BYTES, OTLP_MAX_BODY_BYTES from litellm.rust_bridge.traces import DecodedSpan from litellm.rust_bridge.traces import decode_otlp as native_decode_otlp -from litellm.tracing.normalizers import select_normalizer -from litellm.tracing.normalizers.base import to_int -from litellm.tracing.types import SpanRow +from litellm.rust_bridge.traces import encode_error as native_encode_error +from litellm.tracing.normalizers.messages import content_text +from litellm.tracing.types import SpanRow, SpanType +_FRAMEWORK_SUFFIXES: Final = ( + ".wrap_model_call", + ".wrap_tool_call", + ".before_agent", + ".after_agent", + ".before_model", + ".after_model", +) +_LLM_OPERATIONS: Final = frozenset({"chat", "text_completion", "generate_content"}) +_LC_ROLES: Final = MappingProxyType({"human": "user", "ai": "assistant", "system": "system", "tool": "tool"}) +_OPENINFERENCE_TYPES: Final[Mapping[str, SpanType]] = MappingProxyType({"AGENT": "agent", "LLM": "llm", "TOOL": "tool"}) + + +_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) _MESSAGE_LIST: Final = TypeAdapter(tuple[dict[str, JsonValue], ...]) _MAX_JSON_ESCAPE_BYTES: Final = 6 - -# attributes whose content we lift into Input/Output and drop from SpanAttributes -_HEAVY_ATTRIBUTES: Final = frozenset( - { - "gen_ai.prompt", - "gen_ai.completion", - "gen_ai.tool.definitions", - "gen_ai.input.messages", - "gen_ai.output.messages", - "input.value", - "output.value", - } -) +_MAX_TOKENS: Final = (1 << 32) - 1 class InvalidOTLPPayloadError(ValueError): @@ -49,12 +55,39 @@ class OTLPPayloadTooLargeError(OverflowError): pass +class MessageExtras(TypedDict): + tool_calls: ReadOnly[NotRequired[JsonValue]] + name: ReadOnly[NotRequired[str]] + + +class NormalizedMessage(MessageExtras): + role: ReadOnly[str] + content: ReadOnly[str] + + +class OTLPError(TypedDict): + message: ReadOnly[str] + + +@dataclass(frozen=True, slots=True) +class NormalizedSpan: + kind: SpanType + agent: str = "" + model: str = "" + request_id: str = "" + input: str = "" + output: str = "" + input_tokens: int = 0 + output_tokens: int = 0 + consumed: frozenset[str] = frozenset() + + def _truncate(value: str) -> str: - size = len(value.encode("utf-8")) - if size <= OTLP_MAX_ATTRIBUTE_VALUE_BYTES: + encoded: Final = value.encode("utf-8") + if len(encoded) <= OTLP_MAX_ATTRIBUTE_VALUE_BYTES: return value - kept = value.encode("utf-8")[:OTLP_MAX_ATTRIBUTE_VALUE_BYTES].decode("utf-8", "ignore") - return f"{kept}…[truncated {size - OTLP_MAX_ATTRIBUTE_VALUE_BYTES} bytes]" + kept: Final = encoded[:OTLP_MAX_ATTRIBUTE_VALUE_BYTES].decode("utf-8", "ignore") + return f"{kept}…[truncated {len(encoded) - len(kept.encode('utf-8'))} bytes]" def _size(value: str) -> int: @@ -137,9 +170,9 @@ def _truncate_payload(value: str) -> str: def decode_otlp( body: bytes, content_type: str | None = None, content_encoding: str | None = None ) -> tuple[SpanRow, ...]: - """Decode an OTLP trace export and normalize every span.""" + payload: Final = _decode_content_encoding(body, content_encoding) try: - spans: Final = native_decode_otlp(body, content_type, content_encoding, OTLP_MAX_BODY_BYTES) + spans: Final = native_decode_otlp(payload, content_type) except OverflowError as error: raise OTLPPayloadTooLargeError(str(error)) from error except ValueError as error: @@ -147,19 +180,34 @@ def decode_otlp( return tuple(_span_row(span) for span in spans) +def _decode_content_encoding(body: bytes, content_encoding: str | None) -> bytes: + if len(body) > OTLP_MAX_BODY_BYTES: + raise OTLPPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") + if content_encoding is None or content_encoding.lower() == "identity": + return body + if content_encoding.lower() != "gzip": + raise InvalidOTLPPayloadError("Unsupported OTLP content encoding") + try: + with gzip.GzipFile(fileobj=BytesIO(body)) as stream: + payload: Final = stream.read(OTLP_MAX_BODY_BYTES + 1) + except (EOFError, OSError, zlib.error) as error: + raise InvalidOTLPPayloadError("Invalid OTLP gzip body") from error + if len(payload) > OTLP_MAX_BODY_BYTES: + raise OTLPPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") + return payload + + def _exception_message(span: DecodedSpan) -> str: - """`span.record_exception()` writes an `exception` event; surface it when status.message is empty.""" for event in span["events"]: if event["name"] == "exception": - attributes = event["attributes"] - return attributes.get("exception.message") or attributes.get("exception.type", "") + return event["attributes"].get("exception.message") or event["attributes"].get("exception.type", "") return "" def _span_row(span: DecodedSpan) -> SpanRow: - attributes = span["attributes"] - resource = span["resource_attributes"] - row = SpanRow( + attributes: Final = span["attributes"] + normalized: Final = normalize(span) + return SpanRow( Timestamp=span["start_ns"], TraceId=span["trace_id"], SpanId=span["span_id"], @@ -167,44 +215,201 @@ def _span_row(span: DecodedSpan) -> SpanRow: TraceState=span["trace_state"], SpanName=span["name"], SpanKind=span["kind"], - ServiceName=resource.get("service.name", ""), - ResourceAttributes=resource, + ServiceName=span["resource_attributes"].get("service.name", ""), + ResourceAttributes=span["resource_attributes"], ScopeName=span["scope_name"], ScopeVersion=span["scope_version"], - SpanAttributes=attributes, - Duration=max(span["end_ns"] - span["start_ns"], 0), + SpanAttributes=MappingProxyType( + {key: _truncate(value) for key, value in attributes.items() if key not in normalized.consumed} + ), + Duration=span["end_ns"] - span["start_ns"], StatusCode=span["status_code"], StatusMessage=span["status_message"] or _exception_message(span), TeamId="", ApiKeyHash="", - ObservationType="chain", - AgentName="", - LiteLLMRequestId="", - Model="", - InputTokens=0, - OutputTokens=0, - Input="", - Output="", + ObservationType=normalized.kind, + AgentName=normalized.agent, + Model=normalized.model, + LiteLLMRequestId=attributes.get("gen_ai.response.id") or normalized.request_id, + InputTokens=normalized.input_tokens, + OutputTokens=normalized.output_tokens, + Input=_truncate_payload(normalized.input), + Output=_truncate(normalized.output), ) - normalize(row, attributes) - row["SpanAttributes"] = {k: _truncate(v) for k, v in attributes.items() if k not in _HEAVY_ATTRIBUTES} - row["Input"], row["Output"] = _truncate_payload(row["Input"]), _truncate(row["Output"]) - return row -def _set_tokens(row: SpanRow, attributes: Mapping[str, str]) -> None: - row["InputTokens"] = to_int(attributes.get("gen_ai.usage.input_tokens")) - row["OutputTokens"] = to_int(attributes.get("gen_ai.usage.output_tokens")) +def _loads(value: str) -> JsonValue: + if len(value.encode("utf-8")) > OTLP_MAX_BODY_BYTES: + return None + try: + return _JSON.validate_json(value) + except ValidationError: + return None -def normalize(row: SpanRow, attributes: Mapping[str, str]) -> None: - select_normalizer(row["ScopeName"], attributes).normalize(row, attributes) - if not row["InputTokens"] and not row["OutputTokens"]: - _set_tokens(row, attributes) +def _text(value: JsonValue) -> str: + return value if isinstance(value, str) else "" -def encode_otlp_response(content_type: str | None) -> tuple[bytes, str]: - """Empty ExportTraceServiceResponse in the caller's encoding.""" - if content_type and "json" in content_type: - return b"{}", "application/json" - return b"", "application/x-protobuf" +def _message(value: JsonValue) -> NormalizedMessage | None: + if not isinstance(value, dict): + return None + kwargs: Final = value.get("kwargs", value) + if not isinstance(kwargs, dict): + return None + kind: Final = _text(kwargs.get("type")) or _text(kwargs.get("role")) + if not kind: + return None + calls: Final = kwargs.get("tool_calls") + if calls is not None and (not isinstance(calls, list) or not all(isinstance(call, dict) for call in calls)): + return None + role: Final = _LC_ROLES.get(kind, kind) + content: Final = kwargs.get("content", "") + name: Final = kwargs.get("name") + tool_calls: Final = MessageExtras(tool_calls=calls) if calls else MessageExtras() + tool_name: Final = MessageExtras(name=name) if role == "tool" and isinstance(name, str) else MessageExtras() + message: Final[NormalizedMessage] = { + "role": role, + "content": content_text(content), + **tool_calls, + **tool_name, + } + return message + + +def _messages(value: JsonValue, raw: str) -> str: + if not isinstance(value, list): + return raw + messages: Final = tuple(_message(item) for item in value) + return json.dumps(messages) if all(message is not None for message in messages) else raw + + +def _langsmith_type(span: DecodedSpan) -> SpanType: + attributes: Final = span["attributes"] + kind: Final = attributes.get("langsmith.span.kind", "chain") + if kind in ("llm", "tool"): + return "llm" if kind == "llm" else "tool" + if not span["parent_span_id"] or span["name"] == attributes.get("langsmith.metadata.lc_agent_name"): + return "agent" + return "framework" if span["name"].endswith(_FRAMEWORK_SUFFIXES) else "chain" + + +def _langsmith_io(kind: SpanType, attributes: Mapping[str, str]) -> tuple[str, str, str]: + raw_prompt: Final = attributes.get("gen_ai.prompt", "") + raw_completion: Final = attributes.get("gen_ai.completion", "") + prompt: Final = _loads(raw_prompt) + completion: Final = _loads(raw_completion) + messages: Final = prompt.get("messages") if isinstance(prompt, dict) else None + if kind == "llm": + batch: Final = ( + messages[0] if isinstance(messages, list) and messages and isinstance(messages[0], list) else messages + ) + generations: Final = completion.get("generations") if isinstance(completion, dict) else None + first: Final = generations[0] if isinstance(generations, list) and generations else None + item: Final = first[0] if isinstance(first, list) and first else first + message: Final = item.get("message") if isinstance(item, dict) else None + parsed: Final = _message(message) + kwargs: Final = message.get("kwargs", message) if isinstance(message, dict) else None + metadata: Final = kwargs.get("response_metadata") if isinstance(kwargs, dict) else None + request_id: Final = _text(metadata.get("id")) if isinstance(metadata, dict) else "" + return _messages(batch, raw_prompt), json.dumps(parsed) if parsed is not None else raw_completion, request_id + if kind == "tool": + output: Final = completion.get("output", completion) if isinstance(completion, dict) else completion + update: Final = output.get("update") if isinstance(output, dict) else None + updates: Final = update.get("messages") if isinstance(update, dict) else None + final: Final = updates[-1] if isinstance(updates, list) and updates else output + content: Final = final.get("content", final) if isinstance(final, dict) else final + return ( + raw_prompt, + (content if isinstance(content, str) else json.dumps(content)) if content is not None else raw_completion, + "", + ) + if kind == "agent": + outputs: Final = completion.get("messages") if isinstance(completion, dict) else None + last: Final = _message(outputs[-1]) if isinstance(outputs, list) and outputs else None + return _messages(messages, raw_prompt), json.dumps(last) if last is not None else raw_completion, "" + return raw_prompt, raw_completion, "" + + +def _to_int(value: str | None) -> int: + try: + number: Final = int(value) if value else 0 + except ValueError: + return 0 + if not 0 <= number <= _MAX_TOKENS: + raise InvalidOTLPPayloadError("OTLP token count is outside the storage range") + return number + + +def normalize(span: DecodedSpan) -> NormalizedSpan: + attributes: Final = span["attributes"] + fallback: Final[SpanType] = "agent" if not span["parent_span_id"] else "chain" + input_tokens: Final = _to_int(attributes.get("gen_ai.usage.input_tokens")) + output_tokens: Final = _to_int(attributes.get("gen_ai.usage.output_tokens")) + if span["scope_name"] == "langsmith" or "langsmith.span.kind" in attributes: + kind: Final = _langsmith_type(span) + prompt, completion, request_id = _langsmith_io(kind, attributes) + return NormalizedSpan( + kind, + attributes.get("langsmith.metadata.lc_agent_name", ""), + attributes.get("gen_ai.request.model", ""), + request_id, + prompt, + completion, + input_tokens, + output_tokens, + frozenset({"gen_ai.prompt", "gen_ai.completion"}), + ) + if "openinference.span.kind" in attributes: + return NormalizedSpan( + _OPENINFERENCE_TYPES.get(attributes["openinference.span.kind"].upper(), fallback), + attributes.get("agent.name", ""), + attributes.get("llm.model_name", ""), + "", + attributes.get("input.value", ""), + attributes.get("output.value", ""), + _to_int(attributes.get("llm.token_count.prompt")) + if "llm.token_count.prompt" in attributes + else input_tokens, + _to_int(attributes.get("llm.token_count.completion")) + if "llm.token_count.completion" in attributes + else output_tokens, + frozenset({"input.value", "output.value"}), + ) + operation: Final = attributes.get("gen_ai.operation.name", "") + genai_kind: Final[SpanType] = ( + "llm" + if operation in _LLM_OPERATIONS + else "tool" + if operation == "execute_tool" + else "agent" + if operation == "invoke_agent" + else fallback + ) + input_key: Final = ( + "gen_ai.input.messages" if attributes.get("gen_ai.input.messages") else "gen_ai.tool.call.arguments" + ) + output_key: Final = ( + "gen_ai.output.messages" if attributes.get("gen_ai.output.messages") else "gen_ai.tool.call.result" + ) + return NormalizedSpan( + genai_kind, + attributes.get("gen_ai.agent.name", ""), + attributes.get("gen_ai.request.model") or attributes.get("gen_ai.response.model", ""), + "", + attributes.get(input_key, ""), + attributes.get(output_key, ""), + input_tokens, + output_tokens, + frozenset({input_key, output_key}), + ) + + +def encode_otlp_response(content_type: str | None, error: str | None = None) -> tuple[bytes, str]: + media_type: Final = (content_type or "application/x-protobuf").split(";", 1)[0].strip().lower() + if media_type == "application/json": + response: Final[OTLPError] = {"message": error or ""} + return (json.dumps(response).encode() if error else b"{}"), "application/json" + if error is None: + return b"", "application/x-protobuf" + return native_encode_error(error), "application/x-protobuf" diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 700f33a8a0e..6cef84ec6d0 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -14,13 +14,17 @@ The proxy endpoints are thin wrappers: auth -> build tenant/scope -> call one me import asyncio import os +from collections.abc import AsyncIterable, Callable, Mapping +from io import BytesIO +from threading import BoundedSemaphore +from types import MappingProxyType from typing import Final from litellm.constants import ( AGENT_TRACING_RETENTION_DAYS, AGENT_TRACING_SPEND_LOG_RETENTION_DAYS, OTLP_MAX_BODY_BYTES, - OTLP_OFFLOAD_DECODE_BYTES, + OTLP_MAX_CONCURRENT_INGESTS, ) from litellm.integrations.clickhouse.schema import ensure_schema from litellm.rust_bridge.traces import ClickHouseStorage @@ -28,6 +32,7 @@ from litellm.tracing.decode import OTLPPayloadTooLargeError, decode_otlp from litellm.tracing.store import TraceStore from litellm.tracing.types import ( SpanDetail, + SpanErrorPage, SpanRow, Trace, TracePage, @@ -39,6 +44,10 @@ class TracingPayloadTooLargeError(Exception): pass +class TracingOverloadedError(RuntimeError): + pass + + class Tenant: """Who sent the spans. Always taken from auth, never from span attributes.""" @@ -48,20 +57,49 @@ class Tenant: self.org_id = org_id def stamp(self, row: SpanRow) -> SpanRow: - row["TeamId"] = self.team_id - row["ApiKeyHash"] = self.api_key_hash - row["ResourceAttributes"] = { - **row["ResourceAttributes"], - "litellm.team_id": self.team_id, - "litellm.api_key_hash": self.api_key_hash, - "litellm.org_id": self.org_id, + return self.stamp_rows((row,))[0] + + def stamp_rows(self, rows: tuple[SpanRow, ...]) -> tuple[SpanRow, ...]: + resources: Final = MappingProxyType({id(row["ResourceAttributes"]): row["ResourceAttributes"] for row in rows}) + stamped: Final = MappingProxyType( + { + identity: MappingProxyType( + { + **attributes, + "litellm.team_id": self.team_id, + "litellm.api_key_hash": self.api_key_hash, + "litellm.org_id": self.org_id, + } + ) + for identity, attributes in resources.items() + } + ) + return tuple(self._stamp_row(row, stamped[id(row["ResourceAttributes"])]) for row in rows) + + def _stamp_row(self, row: SpanRow, resource: Mapping[str, str]) -> SpanRow: + stamped: Final[SpanRow] = { + **row, + "TeamId": self.team_id, + "ApiKeyHash": self.api_key_hash, + "ResourceAttributes": resource, } - return row + return stamped class TraceReceiver: - def __init__(self, store: TraceStore) -> None: + def __init__( + self, + store: TraceStore, + max_concurrent_ingests: int = OTLP_MAX_CONCURRENT_INGESTS, + decoder: Callable[[bytes, str | None, str | None], tuple[SpanRow, ...]] = decode_otlp, + body_read_timeout: float = 30, + ) -> None: + if max_concurrent_ingests < 1: + raise ValueError("OTLP ingestion concurrency must be positive") self.store = store + self._decoder: Final = decoder + self._body_read_timeout: Final = body_read_timeout + self._ingest_slots: Final = BoundedSemaphore(max_concurrent_ingests) @classmethod def from_env(cls) -> "TraceReceiver": @@ -84,24 +122,45 @@ class TraceReceiver: async def ingest( self, - body: bytes, + body: bytes | AsyncIterable[bytes], content_type: str | None, content_encoding: str | None, tenant: Tenant, ) -> int: - """Decode an OTLP trace export and store its authenticated spans.""" - if len(body) > OTLP_MAX_BODY_BYTES: + if not self._ingest_slots.acquire(blocking=False): + raise TracingOverloadedError("OTLP ingestion is at capacity") + task: Final = asyncio.create_task(self._ingest(body, content_type, content_encoding, tenant)) + task.add_done_callback(self._release_ingest) + return await asyncio.shield(task) + + def _release_ingest(self, task: asyncio.Task[int]) -> None: + self._ingest_slots.release() + if not task.cancelled(): + task.exception() + + async def _ingest( + self, + body: bytes | AsyncIterable[bytes], + content_type: str | None, + content_encoding: str | None, + tenant: Tenant, + ) -> int: + try: + payload: Final = ( + body + if isinstance(body, bytes) + else await asyncio.wait_for(_read_body(body), timeout=self._body_read_timeout) + ) + except asyncio.TimeoutError as error: + raise TracingOverloadedError("OTLP body upload timed out") from error + if len(payload) > OTLP_MAX_BODY_BYTES: raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") try: - rows: Final = ( - await asyncio.to_thread(decode_otlp, body, content_type, content_encoding) - if len(body) > OTLP_OFFLOAD_DECODE_BYTES - else decode_otlp(body, content_type, content_encoding) - ) + rows: Final = await asyncio.to_thread(self._decoder, payload, content_type, content_encoding) except OTLPPayloadTooLargeError as error: raise TracingPayloadTooLargeError(str(error)) from error try: - await self.store.insert_spans(tuple(tenant.stamp(r) for r in rows)) + await self.store.insert_spans(tenant.stamp_rows(rows)) except OverflowError as error: raise TracingPayloadTooLargeError(str(error)) from error return len(rows) @@ -114,3 +173,17 @@ class TraceReceiver: async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None: return await self.store.get_span(trace_id, span_id, scope, trace_ref) + + async def get_span_error( + self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "", cursor: str | None = None + ) -> SpanErrorPage | None: + return await self.store.get_span_error(trace_id, span_id, scope, trace_ref, cursor) + + +async def _read_body(chunks: AsyncIterable[bytes]) -> bytes: + with BytesIO() as body: + async for chunk in chunks: + if body.tell() + len(chunk) > OTLP_MAX_BODY_BYTES: + raise TracingPayloadTooLargeError(f"OTLP body exceeds {OTLP_MAX_BODY_BYTES} bytes") + body.write(chunk) + return body.getvalue() diff --git a/litellm/tracing/store.py b/litellm/tracing/store.py index 9d1f64f77f0..91420ffd025 100644 --- a/litellm/tracing/store.py +++ b/litellm/tracing/store.py @@ -9,7 +9,7 @@ from itertools import chain from types import MappingProxyType from typing import Any, Final -from pydantic import BaseModel, ConfigDict, TypeAdapter +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter from litellm._logging import verbose_logger from litellm.constants import AGENT_TRACING_LIST_PAGE_SIZE @@ -21,6 +21,7 @@ from litellm.tracing.types import ( AgentNode, Span, SpanDetail, + SpanErrorPage, SpanRow, SpanStatus, Trace, @@ -35,6 +36,20 @@ SPEND_WINDOW_MS: Final = 30 * 60 * 1000 _STATUS: Final = MappingProxyType({"STATUS_CODE_OK": "ok", "STATUS_CODE_ERROR": "error"}) +class _ErrorCursor(BaseModel): + model_config = ConfigDict(frozen=True) + offset: int = Field(ge=0, le=(1 << 63) - 1) + version: str = Field(pattern=r"^[A-F0-9]{64}$") + + +class _ErrorRow(BaseModel): + model_config = ConfigDict(frozen=True) + span_id: str + message: str + total_chars: int + version: str + + class _SpendRow(BaseModel): model_config = ConfigDict(frozen=True) @@ -134,6 +149,7 @@ def span_from_row(row: dict[str, Any], trace_start_ns: int, spend_rows: Sequence duration_ms=int(row["duration_ns"]) / NANOS_PER_MS, status=_status(row["status"]), error=row.get("status_message") or None, + error_truncated=bool(row.get("error_truncated", False)), input_preview=row["input_preview"], model=row["model"] or None, input_tokens=int(row["input_tokens"]), @@ -349,3 +365,41 @@ class TraceStore: output_ui=to_ui_content(rows[0]["output"]), attributes=rows[0]["attributes"], ) + + async def get_span_error( + self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "", cursor: str | None = None + ) -> SpanErrorPage | None: + try: + position: Final = ( + _ErrorCursor.model_validate_json(base64.b64decode(cursor, altchars=b"-_", validate=True)) + if cursor + else None + ) + except (ValueError, binascii.Error) as error: + raise ValueError("Invalid diagnostic cursor") from error + rows: Final = await self.storage.query( + "span_error", + MappingProxyType( + { + **scope, + "trace_id": trace_id, + "span_id": span_id, + "trace_ref": trace_ref, + "error_offset": position.offset if position else 0, + "error_version": position.version if position else "", + } + ), + ) + if not rows: + return None + row: Final = _ErrorRow.model_validate(rows[0]) + offset: Final = (position.offset if position else 0) + len(row.message) + continuation: Final = _ErrorCursor(offset=offset, version=row.version) if offset < row.total_chars else None + return SpanErrorPage( + span_id=row.span_id, + message=row.message, + total_chars=row.total_chars, + next_cursor=base64.urlsafe_b64encode(continuation.model_dump_json().encode()).decode() + if continuation + else None, + ) diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index fcf0d83fb8f..ff965483013 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -9,7 +9,7 @@ A trace is one agent run. It's made of spans (agent / llm / tool / chain / frame """ -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from typing import Literal from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -29,7 +29,8 @@ class Span(TypedDict): start_offset_ms: ReadOnly[float] # relative to trace start duration_ms: ReadOnly[float] status: ReadOnly[SpanStatus] - error: ReadOnly[str | None] # exception message when status == "error" + error: ReadOnly[str | None] + error_truncated: ReadOnly[bool] input_preview: ReadOnly[str] model: ReadOnly[str | None] input_tokens: ReadOnly[int] @@ -91,6 +92,13 @@ class SpanDetail(TypedDict): attributes: ReadOnly[dict[str, str]] +class SpanErrorPage(TypedDict): + span_id: ReadOnly[str] + message: ReadOnly[str] + total_chars: ReadOnly[int] + next_cursor: ReadOnly[str | None] + + class TraceScope(TypedDict): """Who is asking. Empty team_ids = all teams (admins only).""" @@ -109,15 +117,15 @@ class SpanRow(TypedDict): SpanName: ReadOnly[str] SpanKind: ReadOnly[str] ServiceName: ReadOnly[str] - ResourceAttributes: dict[str, str] + ResourceAttributes: ReadOnly[Mapping[str, str]] ScopeName: ReadOnly[str] ScopeVersion: ReadOnly[str] - SpanAttributes: dict[str, str] + SpanAttributes: ReadOnly[Mapping[str, str]] Duration: ReadOnly[int] # ns StatusCode: ReadOnly[str] StatusMessage: ReadOnly[str] - TeamId: str - ApiKeyHash: str + TeamId: ReadOnly[str] + ApiKeyHash: ReadOnly[str] ObservationType: SpanType AgentName: str LiteLLMRequestId: str diff --git a/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json b/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json index 9bd8e67633b..48d8ef0f1dc 100644 --- a/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json +++ b/tests/test_litellm/tracing/fixtures/langsmith_deep_agent_export.json @@ -42,10 +42,10 @@ }, "spans": [ { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "XnnztbUEmF4=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "5e79f3b5b504985e", "name": "deep_research_agent", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742989377137920", "endTimeUnixNano": "1790743040762587136", "attributes": [ @@ -123,16 +123,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "imocMZQNB68=", - "parentSpanId": "Hfr3D90RhPI=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "8a6a1c31940d07af", + "parentSpanId": "1dfaf70fdd1184f2", "name": "ChatOpenAI", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742989383207936", "endTimeUnixNano": "1790742998893985024", "attributes": [ @@ -354,16 +354,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "zwThqgPzRPo=", - "parentSpanId": "g0UfMjWEf2w=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "cf04e1aa03f344fa", + "parentSpanId": "83451f3235847f6c", "name": "FilesystemMiddleware.wrap_model_call", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742989379030016", "endTimeUnixNano": "1790742998895730944", "attributes": [ @@ -477,16 +477,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "svs6j1ovzgE=", - "parentSpanId": "Vt73x+GSQ0o=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "b2fb3a8f5a2fce01", + "parentSpanId": "56def7c7e192434a", "name": "task", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742998900896000", "endTimeUnixNano": "1790743034076956160", "attributes": [ @@ -624,16 +624,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "gUmbSS/ZP4U=", - "parentSpanId": "svs6j1ovzgE=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "81499b492fd93f85", + "parentSpanId": "b2fb3a8f5a2fce01", "name": "researcher", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790742998901422080", "endTimeUnixNano": "1790743034076699904", "attributes": [ @@ -759,16 +759,16 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 }, { - "traceId": "S61CuE6d47pG/IcBhfjwIw==", - "spanId": "/mLyrQOgEWw=", - "parentSpanId": "SUm+6tN4+TU=", + "traceId": "4bad42b84e9de3ba46fc870185f8f023", + "spanId": "fe62f2ad03a0116c", + "parentSpanId": "4949beead378f935", "name": "search_docs", - "kind": "SPAN_KIND_INTERNAL", + "kind": 1, "startTimeUnixNano": "1790743004976721920", "endTimeUnixNano": "1790743004977214208", "attributes": [ @@ -912,7 +912,7 @@ } ], "status": { - "code": "STATUS_CODE_OK" + "code": 1 }, "flags": 256 } @@ -921,4 +921,4 @@ ] } ] -} \ No newline at end of file +} diff --git a/tests/test_litellm/tracing/test_decode.py b/tests/test_litellm/tracing/test_decode.py index 7185341d038..21a79dd6b87 100644 --- a/tests/test_litellm/tracing/test_decode.py +++ b/tests/test_litellm/tracing/test_decode.py @@ -5,13 +5,14 @@ The fixture is a trimmed real export from a Deep Agents run (LangSmith OTEL mode deep_research_agent -> task (tool) -> researcher (subagent) -> search_docs (tool). """ +import base64 import gzip import json from pathlib import Path from unittest.mock import patch import pytest -from google.protobuf.json_format import Parse +from google.protobuf.json_format import ParseDict from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest from opentelemetry.proto.common.v1.common_pb2 import AnyValue, KeyValue from opentelemetry.proto.trace.v1.trace_pb2 import ResourceSpans, ScopeSpans, Span, Status @@ -31,7 +32,14 @@ def _fixture_json() -> bytes: def _fixture_protobuf() -> bytes: request = ExportTraceServiceRequest() - Parse(_fixture_json().decode(), request) + payload = json.loads(_fixture_json()) + for resource in payload["resourceSpans"]: + for scope in resource["scopeSpans"]: + for span in scope["spans"]: + for field in ("traceId", "spanId", "parentSpanId"): + if field in span: + span[field] = base64.b64encode(bytes.fromhex(span[field])).decode() + ParseDict(payload, request) return request.SerializeToString() @@ -191,7 +199,7 @@ def test_plain_tool_input_output(rows_by_name): def test_heavy_attributes_are_lifted_out_of_span_attributes(rows_by_name): for row in rows_by_name.values(): - assert not set(row["SpanAttributes"]) & decode._HEAVY_ATTRIBUTES + assert not set(row["SpanAttributes"]) & {"gen_ai.prompt", "gen_ai.completion"} assert rows_by_name["ChatOpenAI"]["SpanAttributes"]["langsmith.span.kind"] == "llm" @@ -217,12 +225,40 @@ def test_content_type_defaults_to_protobuf(): assert len(decode_otlp(_fixture_protobuf(), None)) == 6 -@pytest.mark.parametrize("content_encoding", ["gzip", None]) -def test_gzip_body_by_header_or_magic_bytes(content_encoding): - rows = decode_otlp(gzip.compress(_fixture_protobuf()), "application/x-protobuf", content_encoding) +def test_gzip_body_by_header(): + rows = decode_otlp(gzip.compress(_fixture_protobuf()), "application/x-protobuf", "gzip") assert len(rows) == 6 +def test_gzip_requires_content_encoding_header(): + with pytest.raises(decode.InvalidOTLPPayloadError): + decode_otlp(gzip.compress(_fixture_protobuf()), "application/x-protobuf") + + +def test_invalid_gzip_body_is_rejected(): + with pytest.raises(decode.InvalidOTLPPayloadError): + decode_otlp(b"not gzip", "application/x-protobuf", "gzip") + + +def test_gzip_expansion_respects_body_limit(): + with patch.object(decode, "OTLP_MAX_BODY_BYTES", 1024): + with pytest.raises(decode.OTLPPayloadTooLargeError): + decode_otlp(gzip.compress(b" " * 16384), "application/json", "gzip") + + +def test_concatenated_gzip_members_are_decoded(): + body = _fixture_json() + midpoint = len(body) // 2 + compressed = gzip.compress(body[:midpoint]) + gzip.compress(body[midpoint:]) + assert len(decode_otlp(compressed, "application/json", "gzip")) == 6 + + +@pytest.mark.parametrize("encoding", ["br", "gzip, identity"]) +def test_unsupported_content_encoding_is_rejected(encoding): + with pytest.raises(decode.InvalidOTLPPayloadError): + decode_otlp(_fixture_protobuf(), "application/x-protobuf", encoding) + + def test_long_values_are_truncated_with_marker(): with patch.object(decode, "OTLP_MAX_ATTRIBUTE_VALUE_BYTES", 100): rows = {r["SpanName"]: r for r in decode_otlp(_fixture_json(), "application/json")} @@ -402,7 +438,7 @@ def test_non_string_attribute_values_are_stringified(): assert row["SpanAttributes"]["flag"] == "true" assert row["SpanAttributes"]["ratio"] == "0.5" assert row["SpanAttributes"]["raw"] == "abc" - assert json.loads(row["SpanAttributes"]["list"]) == ["a", "1"] + assert json.loads(row["SpanAttributes"]["list"]) == ["a", 1] # ---------------------------------------------------------------- helpers @@ -412,3 +448,53 @@ def test_encode_otlp_response_matches_request_encoding(): assert encode_otlp_response("application/json") == (b"{}", "application/json") assert encode_otlp_response("application/x-protobuf") == (b"", "application/x-protobuf") assert encode_otlp_response(None) == (b"", "application/x-protobuf") + body, media_type = encode_otlp_response("application/x-protobuf", "invalid trace") + assert media_type == "application/x-protobuf" + from google.rpc.status_pb2 import Status + + assert Status.FromString(body).message == "invalid trace" + + +@pytest.mark.parametrize( + "attributes, expected", + [ + ({"langsmith__span__kind": "llm"}, "llm"), + ({"langsmith__span__kind": "tool"}, "tool"), + ({"gen_ai__operation__name": "chat"}, "llm"), + ({"gen_ai__operation__name": "execute_tool"}, "tool"), + ({"openinference__span__kind": "LLM"}, "llm"), + ], +) +def test_explicit_root_span_semantics_and_response_id_are_preserved(attributes, expected): + exported = _span("root", b"\x01" * 8, gen_ai__response__id="response-123", **attributes) + (row,) = decode_otlp(_export(exported)) + assert (row["ObservationType"], row["LiteLLMRequestId"]) == (expected, "response-123") + + +@pytest.mark.parametrize( + "payload", + [ + '{"messages": 7}', + '{"messages": {"0": "wrong"}}', + '{"messages": [{"kwargs": []}]}', + '{"messages": [{"role": "assistant", "tool_calls": [1]}]}', + ], +) +def test_malformed_framework_messages_preserve_raw_content_without_rejecting_the_batch(payload): + exported = _span("agent", b"\x01" * 8, langsmith__span__kind="chain", gen_ai__prompt=payload) + (row,) = decode_otlp(_export(exported)) + assert row["Input"] == payload + + +def test_unrecognized_heavy_attributes_are_retained(): + exported = _span("root", b"\x01" * 8, gen_ai__prompt="unknown convention", gen_ai__tool__definitions="tools") + (row,) = decode_otlp(_export(exported)) + assert row["SpanAttributes"]["gen_ai.prompt"] == "unknown convention" + assert row["SpanAttributes"]["gen_ai.tool.definitions"] == "tools" + + +@pytest.mark.parametrize("count", [-1, 1 << 32]) +def test_token_counts_outside_storage_range_are_rejected(count): + exported = _span("root", b"\x01" * 8, gen_ai__usage__input_tokens=count) + with pytest.raises(decode.InvalidOTLPPayloadError, match="storage range"): + decode_otlp(_export(exported)) diff --git a/tests/test_litellm/tracing/test_receiver.py b/tests/test_litellm/tracing/test_receiver.py index d492844db79..0d9aa8d034d 100644 --- a/tests/test_litellm/tracing/test_receiver.py +++ b/tests/test_litellm/tracing/test_receiver.py @@ -2,7 +2,10 @@ Tests for TraceReceiver.ingest (litellm/tracing/receiver.py) with a fake store. """ +import asyncio +from collections.abc import AsyncIterator from pathlib import Path +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -89,18 +92,6 @@ async def test_ingest_rejects_oversized_body(): store.insert_spans.assert_not_awaited() -@pytest.mark.asyncio -async def test_large_body_is_decoded_off_the_event_loop(): - store = _fake_store() - with ( - patch.object(receiver_module, "OTLP_OFFLOAD_DECODE_BYTES", 0), - patch.object(receiver_module.asyncio, "to_thread", wraps=receiver_module.asyncio.to_thread) as to_thread, - ): - count = await TraceReceiver(store).ingest(FIXTURE.read_bytes(), "application/json", None, TENANT) - assert count == 6 - to_thread.assert_called_once() - - @pytest.mark.asyncio async def test_empty_export_writes_nothing(): store = _fake_store() @@ -115,3 +106,57 @@ async def test_reads_delegate_to_store(): scope: TraceScope = {"team_ids": ("team-research",), "api_key_hash": ""} assert await tracing.get_trace("t1", scope) is None store.get_trace.assert_awaited_once_with("t1", scope, "") + + +@pytest.mark.asyncio +async def test_cancelled_request_keeps_its_worker_slot_until_decode_finishes(): + import asyncio + import threading + + from litellm.tracing.receiver import TracingOverloadedError + + loop = asyncio.get_running_loop() + owner = threading.get_ident() + started = asyncio.Event() + stored = asyncio.Event() + release = threading.Event() + + def decoder(body, content_type, content_encoding): + assert threading.get_ident() != owner + loop.call_soon_threadsafe(started.set) + assert release.wait(5) + return () + + store = _fake_store() + store.insert_spans.side_effect = lambda _: stored.set() + tracing = TraceReceiver(store, max_concurrent_ingests=1, decoder=decoder) + pending = asyncio.create_task(tracing.ingest(b"small gzip", None, "gzip", TENANT)) + try: + await asyncio.wait_for(started.wait(), 5) + pending.cancel() + with pytest.raises(asyncio.CancelledError): + await pending + with pytest.raises(TracingOverloadedError): + await tracing.ingest(b"", None, None, TENANT) + finally: + release.set() + await asyncio.wait_for(stored.wait(), 5) + await asyncio.sleep(0) + assert await tracing.ingest(b"", None, None, TENANT) == 0 + + +@pytest.mark.asyncio +async def test_expired_upload_releases_ingestion_slot_without_writing() -> None: + from litellm.tracing.receiver import TracingOverloadedError + + async def unfinished_body() -> AsyncIterator[bytes]: + await asyncio.Event().wait() + yield b"" + + store: Final = _fake_store() + receiver: Final = TraceReceiver(store, max_concurrent_ingests=1, body_read_timeout=0) + with pytest.raises(TracingOverloadedError, match="upload timed out"): + await receiver.ingest(unfinished_body(), "application/json", None, TENANT) + store.insert_spans.assert_not_awaited() + assert await receiver.ingest(b"{}", "application/json", None, TENANT) == 0 + store.insert_spans.assert_awaited_once_with(()) diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 30d7a9b5b0b..3f43e42842c 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -436,3 +436,47 @@ async def test_ambiguous_cache_response_id_keeps_cost_unavailable(): assert trace is not None assert trace["summary"]["spend"] is None assert trace["spans"][0]["spend"] is None + + +@pytest.mark.asyncio +async def test_diagnostic_continuation_preserves_content_version_scope_and_unicode_offset(): + from hashlib import sha256 + + message = "first 🧪\nlast" + version = sha256(message.encode()).hexdigest().upper() + client = MagicMock() + client.query = AsyncMock( + side_effect=[ + [{"span_id": "span-1", "message": "first 🧪", "total_chars": len(message), "version": version}], + [{"span_id": "span-1", "message": "\nlast", "total_chars": len(message), "version": version}], + ] + ) + store = TraceStore(client) + scope = {"team_ids": ("team-a",), "api_key_hash": "key-a"} + first = await store.get_span_error("trace-1", "span-1", scope, "scoped-run") + assert first is not None and first["next_cursor"] is not None + last = await store.get_span_error("trace-1", "span-1", scope, "scoped-run", first["next_cursor"]) + assert last is not None + assert first["message"] + last["message"] == message + assert last["next_cursor"] is None + client.query.assert_awaited_with( + "span_error", + { + **scope, + "trace_id": "trace-1", + "span_id": "span-1", + "trace_ref": "scoped-run", + "error_offset": len(first["message"]), + "error_version": version, + }, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cursor", ["garbage", "e30=", "WzEsMl0="]) +async def test_malformed_diagnostic_cursor_never_reaches_storage(cursor): + client = MagicMock() + client.query = AsyncMock() + with pytest.raises(ValueError, match="Invalid diagnostic cursor"): + await TraceStore(client).get_span_error("trace", "span", {"team_ids": (), "api_key_hash": ""}, cursor=cursor) + client.query.assert_not_awaited() diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index fc750d88e42..e6492c9bca6 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -2,12 +2,17 @@ import base64 import gzip import json import time +from types import MappingProxyType from typing import Final from urllib.parse import parse_qs, urlsplit import pytest -from litellm.rust_bridge._native import NativeTraceStorage +from litellm.rust_bridge._native import NativeTraceStorage, trace_decode_otlp +from litellm.rust_bridge.traces import ClickHouseStorage +from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError +from litellm.tracing.decode import decode_otlp +from litellm.tracing.store import TraceStore from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec pytestmark = pytest.mark.requires_rust_extension @@ -61,7 +66,9 @@ async def test_schema_binding_rejects_non_positive_retention() -> None: @pytest.mark.asyncio -async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement(recording_server: RecordingServer) -> None: +async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement( + recording_server: RecordingServer, +) -> None: recording_server.expected_requests = 2 recording_server.enqueue(ResponseSpec(body="")) recording_server.enqueue(ResponseSpec(status=403, body="denied")) @@ -73,9 +80,10 @@ async def test_schema_setup_uses_writer_credentials_and_rejects_failed_statement assert recording_server.requests[0].raw_body.startswith(b"CREATE DATABASE IF NOT EXISTS") assert recording_server.requests[1].raw_body.startswith(b"CREATE TABLE IF NOT EXISTS") assert "readonly" not in parse_qs(urlsplit(recording_server.requests[0].path).query) - assert recording_server.requests[0].headers["authorization"] == "Basic " + base64.b64encode( - b"writer:p@ss/word%" - ).decode() + assert ( + recording_server.requests[0].headers["authorization"] + == "Basic " + base64.b64encode(b"writer:p@ss/word%").decode() + ) @pytest.mark.asyncio @@ -93,5 +101,100 @@ async def test_insert_encodes_and_sends_rows(recording_server: RecordingServer) "Timestamp": "1970-01-01T00:00:01.23456789Z", "EngineReceivedMs": row["EngineReceivedMs"], } - assert parse_qs(urlsplit(request.path).query)["query"] == ["INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow"] + assert parse_qs(urlsplit(request.path).query)["query"] == [ + "INSERT INTO `trace_test`.otel_traces FORMAT JSONEachRow" + ] assert request.headers["content-encoding"] == "gzip" + + +def _resource_export(attribute_bytes: int, span_count: int, groups: int = 1) -> bytes: + span: Final = { + "traceId": "01" * 16, + "spanId": "02" * 8, + "name": "shared-resource", + "startTimeUnixNano": "1", + "endTimeUnixNano": "2", + } + resource: Final = { + "resource": { + "attributes": [ + {"key": "shared", "value": {"stringValue": "x" * attribute_bytes}}, + {"key": "litellm.team_id", "value": {"stringValue": "spoofed"}}, + ] + }, + "scopeSpans": [ + { + "scope": {"name": "scope-" * 32, "version": "v" * 128}, + "spans": [{**span, "spanId": f"{index + 1:016x}"} for index in range(span_count)], + } + ], + } + return json.dumps({"resourceSpans": [resource] * groups}).encode() + + +def test_decode_and_tenant_stamping_share_resources_without_crossing_groups() -> None: + body: Final = _resource_export(128, 2, 2) + native: Final = trace_decode_otlp(body, "application/json") + assert native[0]["scope_name"] is native[1]["scope_name"] + assert native[0]["scope_version"] is native[1]["scope_version"] + assert native[0]["resource_attributes"] is native[1]["resource_attributes"] + assert native[2]["resource_attributes"] is native[3]["resource_attributes"] + assert native[0]["resource_attributes"] is not native[2]["resource_attributes"] + rows: Final = decode_otlp(body, "application/json") + first: Final = Tenant("team-a", "key-a", "org-a").stamp_rows(rows) + second: Final = Tenant("team-b", "key-b", "org-b").stamp_rows(rows) + assert first[0]["ResourceAttributes"] is first[1]["ResourceAttributes"] + assert first[2]["ResourceAttributes"] is first[3]["ResourceAttributes"] + assert first[0]["ResourceAttributes"] is not first[2]["ResourceAttributes"] + assert first[0]["ResourceAttributes"] is not second[0]["ResourceAttributes"] + assert first[0]["ResourceAttributes"] == { + "shared": "x" * 128, + "litellm.team_id": "team-a", + "litellm.api_key_hash": "key-a", + "litellm.org_id": "org-a", + } + assert second[0]["ResourceAttributes"]["litellm.team_id"] == "team-b" + assert rows[0]["ResourceAttributes"] == {"shared": "x" * 128, "litellm.team_id": "spoofed"} + + +@pytest.mark.asyncio +async def test_resource_fanout_reaches_insert_with_identical_values(recording_server: RecordingServer) -> None: + body: Final = _resource_export(16 * 1024, 1024) + receiver: Final = TraceReceiver(TraceStore(ClickHouseStorage("trace_test", recording_server.base_url))) + tenant: Final = Tenant("team-a", "key-a", "org-a") + assert await receiver.ingest(body, "application/json", None, tenant) == 1024 + encoded: Final = gzip.decompress(recording_server.requests[0].raw_body) + actual: Final = tuple(json.loads(line) for line in encoded.splitlines()) + expected: Final = tenant.stamp_rows(decode_otlp(body, "application/json")) + assert len(encoded) < 64 * 1024 * 1024 + assert tuple({key: value for key, value in row.items() if key != "EngineReceivedMs"} for row in actual) == tuple( + {**row, "Timestamp": "1970-01-01T00:00:00.000000001Z"} for row in expected + ) + assert len({row["EngineReceivedMs"] for row in actual}) == 1 + + +@pytest.mark.asyncio +async def test_shared_resource_still_hits_insert_limit_before_transport(recording_server: RecordingServer) -> None: + recording_server.expected_requests = 0 + body: Final = _resource_export(64 * 1024, 1024) + receiver: Final = TraceReceiver(TraceStore(ClickHouseStorage("trace_test", recording_server.base_url))) + with pytest.raises(TracingPayloadTooLargeError, match="encoded size limit"): + await receiver.ingest(body, "application/json", None, Tenant("team-a", "key-a")) + assert recording_server.requests == [] + + +@pytest.mark.asyncio +async def test_insert_validates_values_without_pydantic_copy(recording_server: RecordingServer) -> None: + storage: Final = ClickHouseStorage("trace_test", recording_server.base_url) + invalid: Final = object() + with pytest.raises(ValueError, match=type(invalid).__name__): + await storage.insert_rows("otel_traces", [{"ResourceAttributes": invalid}]) + attributes: Final = MappingProxyType({"service.name": "trace-test"}) + await storage.insert_rows( + "otel_traces", + (MappingProxyType({"Timestamp": 1, "ResourceAttributes": attributes, "SpanAttributes": attributes}),), + ) + stored: Final = json.loads(gzip.decompress(recording_server.requests[0].raw_body)) + assert stored["Timestamp"] == "1970-01-01T00:00:00.000000001Z" + assert stored["ResourceAttributes"] == attributes + assert stored["SpanAttributes"] == attributes diff --git a/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py b/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py index dc99df24c50..56d067b16cf 100644 --- a/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py +++ b/tests/unit/integrations/open_telemetry/test_otel_exception_handler.py @@ -3,10 +3,9 @@ that fail after auth but before the route handler runs (e.g. /model/new TypeError or RequestValidationError).""" import asyncio -import types import pytest -from fastapi import HTTPException +from fastapi import HTTPException, Request from fastapi.exceptions import RequestValidationError import litellm.proxy.proxy_server as proxy_server_module @@ -23,13 +22,11 @@ from litellm.integrations._types.open_inference import ErrorAttributes from ._helpers import assert_server_span_attrs, get_server_span -def _fake_request(parent_otel_span=None, path="/key/generate"): - """A real Request always carries a url; the validation handler reads its path to - decide whether the caller is on a surface with its own error contract.""" - state = types.SimpleNamespace() - if parent_otel_span is not None: - state.parent_otel_span = parent_otel_span - return types.SimpleNamespace(state=state, url=types.SimpleNamespace(path=path)) +def _fake_request(parent_otel_span: object | None = None, path: str = "/key/generate") -> Request: + return Request({ + "type": "http", "method": "POST", "path": path, "headers": [], + "state": {"parent_otel_span": parent_otel_span}, + }) @pytest.fixture diff --git a/tests/unit/proxy/auth/test_user_api_key_auth.py b/tests/unit/proxy/auth/test_user_api_key_auth.py index 08c2f02a83c..1cfef5b3a6c 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth.py @@ -124,7 +124,7 @@ async def test_check_blocked_team(): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") await user_api_key_auth(request=request, api_key="Bearer " + user_key) @@ -162,7 +162,7 @@ async def test_team_object_has_object_permission_id(): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "test-client") - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") with patch("litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock) as mock_common_checks: @@ -263,7 +263,7 @@ async def test_aaauser_personal_budgets(key_ownership): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", _NoMembershipRowPrisma()) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") test_user_cache = getattr(litellm.proxy.proxy_server, "user_api_key_cache") @@ -294,7 +294,7 @@ async def test_user_api_key_auth_fails_with_prohibited_params(prohibited_param): setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") # Create request with prohibited parameter in body - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") async def return_body(): @@ -334,7 +334,7 @@ async def test_auth_with_allowed_routes(route, should_raise_error): setattr(proxy_server, "master_key", "sk-1234") setattr(proxy_server, "general_settings", general_settings) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": route, "headers": []}) request._url = URL(url=route) if should_raise_error: @@ -411,7 +411,7 @@ def test_ui_token_route_access(route, user_role, should_be_allowed): from starlette.datastructures import URL from fastapi import Request - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": route, "headers": []}) request._url = URL(url=route) if should_be_allowed: @@ -494,7 +494,7 @@ async def test_auth_not_connected_to_db(): {"allow_requests_on_db_unavailable": True}, ) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") valid_token = await user_api_key_auth(request=request, api_key="Bearer " + user_key) @@ -676,7 +676,7 @@ async def test_soft_budget_alert(): setattr(litellm.proxy.proxy_server, "prisma_client", AsyncMock()) # Create request - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") # Track if budget_alerts was called @@ -1162,7 +1162,7 @@ async def test_x_litellm_api_key(): ignored_key = "aj12445" # Create request with headers as bytes - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") valid_token = await user_api_key_auth( @@ -1336,7 +1336,7 @@ async def test_user_model_budget_is_enforced_through_user_api_key_auth(over_budg ttl=600, ) - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") async def return_body(): diff --git a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py index ef6832ef77b..781d0a13bfd 100644 --- a/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py +++ b/tests/unit/proxy/auth/test_user_api_key_auth_request_flow.py @@ -6267,6 +6267,7 @@ async def test_user_api_key_auth_sets_end_user_id_when_builder_skips_it(): "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6321,6 +6322,7 @@ async def test_user_api_key_auth_does_not_overwrite_end_user_id_set_by_builder() "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6376,6 +6378,7 @@ async def test_user_api_key_auth_authenticates_before_raising_malformed_body_err "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6435,6 +6438,7 @@ async def _run_auth_with_malformed_body(post_call_failure_hook): "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6507,6 +6511,7 @@ async def test_user_api_key_auth_malformed_body_with_rejected_key_still_returns_ "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") @@ -6557,6 +6562,7 @@ async def test_user_api_key_auth_does_not_double_log_a_malformed_body_from_a_rej "type": "http", "headers": [(b"content-type", b"application/json")], "method": "POST", + "path": "/chat/completions", } ) request._url = URL(url="/chat/completions") diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index e3851f6c21a..fd747d5a6f2 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -2,14 +2,14 @@ import gzip import io import json from collections.abc import Mapping -from typing import Literal, get_type_hints +from typing import Final, Literal, get_type_hints from unittest.mock import AsyncMock, MagicMock, patch import orjson import pytest -from fastapi import Request from fastapi.testclient import TestClient from starlette.datastructures import FormData +from starlette.requests import Request @@ -1109,7 +1109,7 @@ class TestGetRequestBody: mock_request.method = "POST" mock_request.body = AsyncMock(return_value=orjson.dumps(payload)) mock_request.headers = {"content-type": "application/json; charset=utf-8"} - mock_request.scope = {} + mock_request.scope = {"type": "http", "method": "POST", "path": "/v1/chat/completions"} result = await get_request_body(mock_request) assert result == payload @@ -1120,7 +1120,7 @@ class TestGetRequestBody: mock_request.method = "POST" mock_request.headers = {"content-type": "multipart/form-data; boundary=x"} mock_request.form = AsyncMock(return_value=FormData({"k": "v"})) - mock_request.scope = {} + mock_request.scope = {"type": "http", "method": "POST", "path": "/v1/chat/completions"} result = await get_request_body(mock_request) assert result == {"k": "v"} @@ -1273,3 +1273,96 @@ def test_shared_inference_model_selection_preserves_handler_precedence( from litellm.proxy.common_utils.http_parsing_utils import resolve_inference_model assert resolve_inference_model(body, settings, cli, path, kind=kind) == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method,path,skip_parse", + [ + ("POST", "/v1/traces", True), + ("GET", "/v1/traces", False), + ("POST", "/v1/messages", False), + ("POST", "/v1/traces/other", False), + ], +) +@pytest.mark.parametrize("root_path", ["", "/tenant-a"]) +async def test_only_trace_ingest_skips_json_body(method: str, path: str, skip_parse: bool, root_path: str) -> None: + body: Final = b'{"key":"value"}' + receive: Final = AsyncMock(return_value={"type": "http.request", "body": body, "more_body": False}) + request: Final = Request( + { + "type": "http", "method": method, "path": root_path + path, "root_path": root_path, + "headers": [(b"content-type", b"application/json")], + }, + receive, + ) + + parsed: Final = await _read_request_body(request) + if skip_parse: + assert parsed == {} + receive.assert_not_awaited() + else: + assert parsed == {"key": "value"} + receive.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("content_type, encoding", [ + ("application/json", ""), ("application/x-protobuf", ""), ("application/json", "gzip"), +]) +async def test_otlp_auth_does_not_consume_chunked_bodies_before_the_receiver_limit(content_type, encoding): + from litellm.constants import OTLP_MAX_BODY_BYTES + from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError + + received = [] + chunk = b"x" * (OTLP_MAX_BODY_BYTES // 2 + 1) + + async def receive(): + received.append(1) + assert len(received) <= 2, "receiver must reject without consuming subsequent chunks" + return {"type": "http.request", "body": chunk, "more_body": True} + + request = Request({"type": "http", "method": "POST", "path": "/v1/traces", "headers": [ + (b"content-type", content_type.encode()), (b"content-encoding", encoding.encode()), + ]}, receive) + assert await _read_request_body(request) == {} + assert received == [] + store = MagicMock() + store.insert_spans = AsyncMock() + with pytest.raises(TracingPayloadTooLargeError): + await TraceReceiver(store).ingest(request.stream(), content_type, encoding, Tenant("team", "key")) + assert len(received) == 2 + store.insert_spans.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_auth_body_read_and_trace_handler_leave_stream_for_receiver_limit() -> None: + from litellm.constants import OTLP_MAX_BODY_BYTES + from litellm.proxy import tracing_endpoints + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import _read_request_body_deferring_parse_failure + from litellm.tracing import TraceReceiver + + chunk: Final = b"x" * (OTLP_MAX_BODY_BYTES // 2 + 1) + receive: Final = AsyncMock( + side_effect=[{"type": "http.request", "body": chunk, "more_body": True}] * 2 + ) + request: Final = Request( + {"type": "http", "method": "POST", "path": "/v1/traces", "headers": [(b"content-type", b"application/json")]}, + receive, + ) + store: Final = MagicMock() + store.insert_spans = AsyncMock() + context: Final = await tracing_endpoints.provide_trace_access( + auth=UserAPIKeyAuth(token="key", team_id="team"), tracing=TraceReceiver(store) + ) + + parsed, parse_error = await _read_request_body_deferring_parse_failure(request) + assert parsed == {} + assert parse_error is None + receive.assert_not_awaited() + + response: Final = await tracing_endpoints.ingest_otlp_traces(request, context) + assert response.status_code == 413 + assert receive.await_count == 2 + store.insert_spans.assert_not_awaited() diff --git a/tests/unit/proxy/proxy_server/test_exception_handlers.py b/tests/unit/proxy/proxy_server/test_exception_handlers.py index 16cb1146ff5..0aff43057f9 100644 --- a/tests/unit/proxy/proxy_server/test_exception_handlers.py +++ b/tests/unit/proxy/proxy_server/test_exception_handlers.py @@ -16,7 +16,7 @@ from unittest.mock import MagicMock import httpx import pytest -from fastapi import HTTPException +from fastapi import HTTPException, Request from fastapi.exceptions import RequestValidationError from litellm.proxy._types import ProxyException @@ -31,10 +31,10 @@ from .conftest import normalize def _make_request(parent_otel_span=None, path="/chat/completions"): - """A real Request always carries a url; the validation handler reads its path to - decide whether the caller is on a surface with its own error contract.""" - state = SimpleNamespace(parent_otel_span=parent_otel_span) - return SimpleNamespace(state=state, url=SimpleNamespace(path=path)) + return Request({ + "type": "http", "method": "POST", "path": path, "headers": [], + "state": {"parent_otel_span": parent_otel_span}, + }) # --------------------------------------------------------------------------- @@ -477,3 +477,42 @@ async def test_otel_unhandled_exception_handler_reraises_http_exception_invalid( request = _make_request() with pytest.raises(HTTPException): await otel_unhandled_exception_handler(request=request, exc=HTTPException(status_code=418, detail="teapot")) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("media_type", ["application/json", "application/x-protobuf"]) +@pytest.mark.parametrize("root_path", ["", "/tenant-a"]) +@pytest.mark.parametrize("native_available", [True, False]) +@pytest.mark.parametrize("error", [ + ProxyException("database credentials: secret", "auth_error", None, 401), + HTTPException(403, "database credentials: secret"), +]) +async def test_otlp_auth_errors_hide_internal_details_and_survive_missing_native( + media_type: str, root_path: str, native_available: bool, + error: ProxyException | HTTPException, monkeypatch: pytest.MonkeyPatch, +) -> None: + from google.rpc.status_pb2 import Status + + from litellm.proxy.proxy_server import otlp_http_exception_handler + from litellm.rust_bridge import loader + + if not native_available: + monkeypatch.setattr(loader, "_cached_bridge", None) + request: Final = Request({ + "type": "http", "method": "POST", "path": root_path + "/v1/traces", "root_path": root_path, + "headers": [(b"content-type", media_type.encode())], + }) + response: Final = ( + await openai_exception_handler(request, error) + if isinstance(error, ProxyException) + else await otlp_http_exception_handler(request, error) + ) + assert response.status_code == (401 if isinstance(error, ProxyException) else 403) + assert response.headers["content-type"].startswith(media_type) + message: Final = ( + json.loads(response.body)["message"] + if media_type == "application/json" + else Status.FromString(response.body).message + ) + expected: Final = "Unauthorized" if isinstance(error, ProxyException) else "Forbidden" + assert message == (expected if native_available or media_type == "application/json" else "") diff --git a/tests/unit/proxy/test_proxy_reject_logging.py b/tests/unit/proxy/test_proxy_reject_logging.py index eb5c5a52f0a..d5a3acb2cd7 100644 --- a/tests/unit/proxy/test_proxy_reject_logging.py +++ b/tests/unit/proxy/test_proxy_reject_logging.py @@ -152,6 +152,7 @@ async def test_chat_completion_request_with_redaction(route, body): scope={ "type": "http", "method": "POST", + "path": route, "headers": [(b"content-type", b"application/json")], "query_string": query_params.encode(), } diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index 8947da4d9fc..300edc8e435 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -263,7 +263,7 @@ def test_add_headers_to_request(litellm_key_header_name): "X-Stainless-Header": "Stainless-Value", "anthropic-beta": "beta-value", } - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") request._body = json.dumps({"model": "gpt-3.5-turbo"}).encode("utf-8") request_headers = clean_headers(headers, litellm_key_header_name) @@ -466,7 +466,7 @@ async def test_team_disable_guardrails(mock_acompletion, client_no_auth, monkeyp setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": "/chat/completions", "headers": []}) request._url = URL(url="/chat/completions") body = {"metadata": {"guardrails": {"hide_secrets": False}}} @@ -1347,7 +1347,7 @@ async def test_create_team_member_add_team_admin_user_api_key_auth( from starlette.datastructures import URL - request = Request(scope={"type": "http"}) + request = Request(scope={"type": "http", "method": "POST", "path": team_route, "headers": []}) request._url = URL(url=team_route) body = {} diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 2e34172acfd..aa1403b8db9 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -96,8 +96,22 @@ def client() -> TestClient: return TestClient(app) -def test_501_when_tracing_not_enabled(client): - assert client.post("/v1/traces", content=b"").status_code == 501 +@pytest.mark.parametrize("native_available", [True, False]) +def test_501_when_tracing_not_enabled( + client: TestClient, native_available: bool, monkeypatch: pytest.MonkeyPatch +) -> None: + from google.rpc.status_pb2 import Status + + from litellm.rust_bridge import loader + + if not native_available: + monkeypatch.setattr(loader, "_cached_bridge", None) + response: Final = client.post("/v1/traces", content=b"") + assert response.status_code == 501 + assert response.headers["content-type"] == "application/x-protobuf" + assert Status.FromString(response.content).message == ( + "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL." if native_available else "" + ) assert client.get("/v1/traces").status_code == 501 @@ -111,7 +125,7 @@ def test_post_protobuf_returns_empty_protobuf(client, receiver): assert response.content == b"" assert response.headers["content-type"] == "application/x-protobuf" kwargs = receiver.ingest.call_args.kwargs - assert kwargs["body"] == b"\x0a\x00" + assert kwargs["body"] is not None assert kwargs["content_type"] == "application/x-protobuf" assert kwargs["content_encoding"] == "gzip" assert kwargs["tenant"].team_id == "team-research" @@ -134,7 +148,9 @@ def test_post_too_large_is_413(client, receiver): receiver.ingest.side_effect = TracingPayloadTooLargeError("OTLP body exceeds 10 bytes") response = client.post("/v1/traces", content=b"x" * 20) assert response.status_code == 413 - assert "exceeds" in response.json()["detail"] + from google.rpc.status_pb2 import Status + + assert "exceeds" in Status.FromString(response.content).message def test_list_traces_passes_scope_window_and_cursor(client, receiver): @@ -222,8 +238,13 @@ def test_view_only_admin_cannot_ingest_traces(client, receiver): 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: +@pytest.mark.parametrize( + "status_code, field, message", + [(401, "detail", "Invalid API key"), (403, "message", "Not allowed to ingest agent traces")], +) +def test_auth_failure_precedes_disabled_receiver( + client: TestClient, status_code: int, field: str, message: str +) -> None: def unavailable() -> None: return None @@ -234,11 +255,9 @@ def test_auth_failure_precedes_disabled_receiver(client: TestClient, status_code 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"{}") + response: Final = client.post("/v1/traces", content=b"{}", headers={"content-type": "application/json"}) assert response.status_code == status_code - assert response.json() == { - "detail": "Invalid API key" if status_code == 401 else "Not allowed to ingest agent traces" - } + assert response.json() == {field: message} def test_disabled_receiver_precedes_read_scope_rejection(client: TestClient) -> None: diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index d12d0a3219c..97d1782b08c 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -115,7 +115,7 @@ import type { ComplexityRouterConfigPayload } from "./add_model/build_complexity import type { AutoRouterPresetsResponse } from "@/lib/autorouter_presets"; import type { VectorStoreIndex } from "@/app/(dashboard)/vector-stores/_components/IndexesTab"; import type { RoutingDecision } from "./view_logs/LogDetailsDrawer/RoutingDecisionCard"; -import type { SpanDetail, Trace, TracePage } from "./view_logs/TraceView/traceTypes"; +import type { SpanDetail, SpanErrorPage, Trace, TracePage } from "./view_logs/TraceView/traceTypes"; import { createApiClient, deriveErrorMessage, @@ -2139,6 +2139,17 @@ export const agentTraceSpanCall = async ( query: { trace_ref: traceRef || undefined }, }); +export const agentTraceSpanErrorCall = async ( + accessToken: string, + traceId: string, + spanId: string, + options: { traceRef?: string; cursor?: string | null }, +): Promise => + apiClient.get(`/v1/traces/${encodeURIComponent(traceId)}/spans/${encodeURIComponent(spanId)}/error`, { + accessToken, + query: { trace_ref: options.traceRef || undefined, cursor: options.cursor || undefined }, + }); + export const adminSpendLogsCall = async (accessToken: string) => { try { const data = await apiClient.get(`/global/spend/logs`, { accessToken }); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx index 9c156d88bd5..86ec3007f72 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailContent.tsx @@ -1,15 +1,17 @@ "use client"; import { useQuery, type UseQueryOptions } from "@tanstack/react-query"; +import { useState } from "react"; import { AlertTriangle } from "lucide-react"; +import { Button } from "@/components/ui/button"; import { cn } from "@/lib/cva.config"; -import { agentTraceSpanCall } from "../../networking"; +import { agentTraceSpanCall, agentTraceSpanErrorCall } from "../../networking"; import { type KeyValue, KeyValueRows, objectEntries } from "./KeyValueRows"; import { Card, MessageCard, Section, ToolResultCard } from "./MessageCard"; import type { ErrorSource } from "./traceTree"; -import type { Span, SpanDetail, TraceMessage, UIContent, UIMessage } from "./traceTypes"; +import type { Span, SpanDetail, SpanErrorPage, TraceMessage, UIContent, UIMessage } from "./traceTypes"; import { errorSource, parseJson, parseMessages, prettyPayload } from "./traceUtils"; const ERROR_SOURCE_LABEL: Record = { tool: "Tool", model: "Model", litellm: "LiteLLM" }; @@ -143,6 +145,58 @@ interface DetailContentProps { span: Span; } +function DiagnosticContent({ accessToken, traceId, traceRef, span }: DetailContentProps) { + const [opened, setOpened] = useState(false); + const [cursor, setCursor] = useState(null); + const queryOptions: UseQueryOptions = { + queryKey: ["agentTraceSpanError", traceId, traceRef, span.span_id, accessToken, cursor], + queryFn: () => agentTraceSpanErrorCall(accessToken, traceId, span.span_id, { traceRef, cursor }), + enabled: opened, + staleTime: Infinity, + gcTime: 0, + retry: false, + }; + const query = useQuery(queryOptions); + return ( +
+ {span.error_truncated &&

Error preview truncated

} + {!opened && ( + + )} + {opened && query.isPending &&

Loading diagnostic…

} + {opened && query.isError && ( +
+ Could not load diagnostic: {query.error.message} + +
+ )} + {opened && query.data && ( + <> + +

+ {cursor ? "Continuation" : "Beginning"} of stored diagnostic ({query.data.total_chars.toLocaleString()}{" "} + characters) +

+ {query.data.next_cursor && ( + + )} + {cursor && ( + + )} + + )} +
+ ); +} + /** Content tab: the error first (if any), then collapsible Input and Output rendered as chat cards. */ export function DetailContent({ accessToken, traceId, traceRef, span }: DetailContentProps) { const detailQuery = useSpanDetail(accessToken, traceId, span.span_id, traceRef); @@ -152,6 +206,15 @@ export function DetailContent({ accessToken, traceId, traceRef, span }: DetailCo return (
+ {span.error && ( + + )} {detailQuery.isLoading &&
Loading span…
} {detailQuery.isError &&
Could not load span: {detailQuery.error.message}
} {detail?.input ? ( diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx similarity index 89% rename from ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx rename to ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx index 1adcda0f6cc..be5eefa1546 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/DetailPane.integration.test.tsx @@ -6,14 +6,15 @@ import { renderWithProviders, testQueryClient } from "../../../../tests/test-uti import { DetailPane } from "./DetailPane"; import { absoluteTime, SpanHoverCard, spanFacts } from "./SpanHoverCard"; import type { GroupRowData, SpanRowData } from "./traceTree"; -import type { Span, SpanDetail, Trace } from "./traceTypes"; +import type { Span, SpanDetail, SpanErrorPage, Trace } from "./traceTypes"; vi.mock("../../networking", () => ({ agentTraceSpanCall: vi.fn(), + agentTraceSpanErrorCall: vi.fn(), getProxyBaseUrl: () => "http://proxy.test/", })); -import { agentTraceSpanCall } from "../../networking"; +import { agentTraceSpanCall, agentTraceSpanErrorCall } from "../../networking"; type SpanFields = Partial & Pick; @@ -322,3 +323,31 @@ describe("SpanHoverCard", () => { expect(within(card).getByRole("region", { name: "Tags" })).toHaveTextContent("agent:support_triage_agent"); }); }); + +it("retrieves the retained diagnostic one section at a time", async () => { + const firstPage: SpanErrorPage = { + span_id: "tool1", + message: "First diagnostic section", + total_chars: 100, + next_cursor: "next-section", + }; + const lastPage: SpanErrorPage = { + span_id: "tool1", + message: "Last diagnostic section", + total_chars: 100, + next_cursor: null, + }; + vi.mocked(agentTraceSpanErrorCall).mockResolvedValueOnce(firstPage).mockResolvedValueOnce(lastPage); + renderPane(spanRow({ ...failedTool, error_truncated: true })); + expect(screen.getByText("Error preview truncated")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "View stored diagnostic" })); + expect(await screen.findByText("First diagnostic section")).toBeInTheDocument(); + await userEvent.click(screen.getByRole("button", { name: "Next section" })); + expect(await screen.findByText("Last diagnostic section")).toBeInTheDocument(); + expect(screen.queryByText("First diagnostic section")).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Next section" })).not.toBeInTheDocument(); + expect(agentTraceSpanErrorCall).toHaveBeenLastCalledWith("sk-test", "t1", "tool1", { + traceRef: undefined, + cursor: "next-section", + }); +}); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts index 4711799bd1a..d3080f8aab4 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/traceTypes.ts @@ -20,6 +20,7 @@ export interface Span { status: SpanStatus; /** Exception message when status is "error". */ error?: string | null; + error_truncated?: boolean; input_preview: string; model: string | null; input_tokens: number; @@ -120,3 +121,10 @@ export interface TraceMessage { name?: string; tool_calls?: TraceToolCall[]; } + +export interface SpanErrorPage { + span_id: string; + message: string; + total_chars: number; + next_cursor: string | null; +} diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index c6bb9be41df..ebc6d0e70cc 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -21779,6 +21779,23 @@ export interface paths { patch?: never; trace?: never; }; + "/v1/traces/{trace_id}/spans/{span_id}/error": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Agent Trace Span Error */ + get: operations["get_agent_trace_span_error_v1_traces__trace_id__spans__span_id__error_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/v1/unified_access_group": { parameters: { query?: never; @@ -44115,6 +44132,17 @@ export interface components { /** Version */ version?: string; }; + /** SpanErrorPage */ + SpanErrorPage: { + /** Message */ + message: string; + /** Next Cursor */ + next_cursor: string | null; + /** Span Id */ + span_id: string; + /** Total Chars */ + total_chars: number; + }; /** SpendAnalyticsPaginatedResponse */ SpendAnalyticsPaginatedResponse: { metadata?: components["schemas"]["DailySpendMetadata"]; @@ -77335,6 +77363,41 @@ export interface operations { }; }; }; + get_agent_trace_span_error_v1_traces__trace_id__spans__span_id__error_get: { + parameters: { + query?: { + trace_ref?: string; + cursor?: string | null; + }; + header?: never; + path: { + trace_id: string; + span_id: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["SpanErrorPage"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; list_access_groups_v1_unified_access_group_get: { parameters: { query?: never;