diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index e01d52a55a6..86df5cce1eb 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4469,6 +4469,19 @@ dependencies = [ "time", ] +[[package]] +name = "litellm-traces-cache" +version = "0.1.0" +dependencies = [ + "litellm-traces", + "moka", + "rstest", + "serde_json", + "sha2 0.10.9", + "thiserror 2.0.19", + "tokio", +] + [[package]] name = "litellm-traces-clickhouse" version = "0.1.0" @@ -4484,6 +4497,7 @@ dependencies = [ "litellm-migrate", "litellm-storage-clickhouse", "litellm-traces", + "litellm-traces-cache", "macro_rules_attribute", "moka", "rstest", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 7c32e82eb4c..b0766f11e87 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -13,6 +13,7 @@ litellm-config = { path = "crates/config" } litellm-router = { path = "crates/router" } litellm-tracing = { path = "crates/tracing" } litellm-traces = { path = "crates/traces" } +litellm-traces-cache = { path = "crates/traces-cache" } litellm-traces-clickhouse = { path = "crates/traces-clickhouse" } litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } litellm-migrate = { path = "crates/migrate" } diff --git a/litellm-rust/crates/traces-cache/AGENTS.md b/litellm-rust/crates/traces-cache/AGENTS.md new file mode 100644 index 00000000000..bfaea42d901 --- /dev/null +++ b/litellm-rust/crates/traces-cache/AGENTS.md @@ -0,0 +1,5 @@ +Own resolved-trace snapshot storage, cache identity, weighting, and expiry +Depend on trace domain types, never storage, HTTP, or Python +Preserve the full source and authorization scope in every cache key +Keep snapshots immutable and expose borrowed data +Keep cursor formats and database reads in their existing owners diff --git a/litellm-rust/crates/traces-cache/Cargo.toml b/litellm-rust/crates/traces-cache/Cargo.toml new file mode 100644 index 00000000000..3e58ea4c965 --- /dev/null +++ b/litellm-rust/crates/traces-cache/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "litellm-traces-cache" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-traces.workspace = true +moka.workspace = true +serde_json.workspace = true +sha2.workspace = true +thiserror.workspace = true + +[dev-dependencies] +rstest.workspace = true +tokio.workspace = true diff --git a/litellm-rust/crates/traces-cache/src/error.rs b/litellm-rust/crates/traces-cache/src/error.rs new file mode 100644 index 00000000000..ac5f19369fe --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/error.rs @@ -0,0 +1,7 @@ +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("trace snapshot serialization failed")] + Serialization(#[from] serde_json::Error), + #[error("trace snapshot exceeds the size limit")] + ReadTooLarge, +} diff --git a/litellm-rust/crates/traces-cache/src/lib.rs b/litellm-rust/crates/traces-cache/src/lib.rs new file mode 100644 index 00000000000..7dfd3bf32a1 --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/lib.rs @@ -0,0 +1,166 @@ +use std::{sync::Arc, time::Duration}; + +use litellm_traces::{Trace, query::named::ReadAccessParams}; +use moka::future::Cache; +use sha2::{Digest, Sha256}; + +mod error; + +pub use error::Error; + +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct SnapshotKey(String); + +impl SnapshotKey { + pub fn new( + source: &str, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + snapshot_ms: u64, + ) -> Result { + let encoded = serde_json::to_vec(&(source, access, trace_id, trace_ref, snapshot_ms))?; + Ok(Self(format!("{:x}", Sha256::digest(encoded)))) + } +} + +pub struct Snapshot { + trace: Trace, + version: String, + weight: u32, +} + +impl Snapshot { + pub fn trace(&self) -> &Trace { + &self.trace + } + + pub fn version(&self) -> &str { + &self.version + } +} + +pub struct SnapshotCache { + entries: Cache>, + max_graph_bytes: usize, +} + +impl SnapshotCache { + pub fn new(max_graph_bytes: usize, ttl: Duration) -> Self { + Self { + entries: Cache::builder() + .max_capacity((max_graph_bytes as u64).saturating_mul(2)) + .weigher(|_: &SnapshotKey, snapshot: &Arc| snapshot.weight) + .time_to_live(ttl) + .build(), + max_graph_bytes, + } + } + + pub async fn get(&self, key: &SnapshotKey) -> Option> { + self.entries.get(key).await + } + + pub async fn insert(&self, key: SnapshotKey, trace: Trace) -> Result, Error> { + let encoded = serde_json::to_vec(&trace)?; + if encoded.len() > self.max_graph_bytes { + return Err(Error::ReadTooLarge); + } + + let span_ids: Vec<&str> = trace + .spans + .iter() + .map(|span| span.span_id.as_str()) + .collect(); + + let version = format!("{:x}", Sha256::digest(serde_json::to_vec(&span_ids)?)); + + let snapshot = Arc::new(Snapshot { + trace, + version, + weight: u32::try_from(encoded.len().saturating_mul(2)).unwrap_or(u32::MAX), + }); + + self.entries.insert(key, Arc::clone(&snapshot)).await; + Ok(snapshot) + } +} + +#[cfg(test)] +mod tests { + use litellm_traces::{ + SpanStatus, + query::named::{SpendByResponseIdsRow, TraceSpansRow}, + resolve_trace, + }; + + use super::*; + + fn trace(span_id: &str) -> Trace { + let rows = [TraceSpansRow { + trace_id: String::new(), + span_id: span_id.into(), + parent_span_id: String::new(), + name: "run".into(), + kind: litellm_traces::ObservationType::Agent, + wrapper_candidate: false, + agent: "agent".into(), + framework: String::new(), + status: SpanStatus::Ok, + status_message: String::new(), + error_truncated: false, + start_ns: 1_790_742_989_000_000_000, + duration_ns: 10_000_000, + service: "agent-demo".into(), + input_preview: format!("input of {span_id}"), + model: String::new(), + input_tokens: 0, + output_tokens: 0, + litellm_request_id: String::new(), + call_keys: Vec::new(), + call_evidence: None, + tool_call_id: String::new(), + team_id: String::new(), + api_key_hash: String::new(), + user_id: String::new(), + }]; + resolve_trace("trace", "ref", &rows, &[] as &[SpendByResponseIdsRow]) + .expect("fixture should resolve") + } + + fn key(suffix: &str) -> SnapshotKey { + SnapshotKey::new( + "source", + &ReadAccessParams { + all_teams: false, + user_id: String::new(), + team_ids: vec!["team".into()], + }, + suffix, + "ref", + 100, + ) + .unwrap() + } + + #[tokio::test] + async fn weighted_capacity_bounds_retained_snapshots() { + let limit = ["first", "second", "third"] + .iter() + .map(|span_id| serde_json::to_vec(&trace(span_id)).unwrap().len()) + .max() + .unwrap(); + let cache = SnapshotCache::new(limit, Duration::from_secs(120)); + + for (key, span_id) in [ + (key("a"), "first"), + (key("b"), "second"), + (key("c"), "third"), + ] { + cache.insert(key, trace(span_id)).await.unwrap(); + } + + cache.entries.run_pending_tasks().await; + assert!(cache.entries.weighted_size() <= (limit as u64) * 2); + } +} diff --git a/litellm-rust/crates/traces-cache/tests/snapshots.rs b/litellm-rust/crates/traces-cache/tests/snapshots.rs new file mode 100644 index 00000000000..de199ae0301 --- /dev/null +++ b/litellm-rust/crates/traces-cache/tests/snapshots.rs @@ -0,0 +1,191 @@ +use std::time::Duration; + +use litellm_traces::{ + SpanStatus, Trace, + query::named::{ReadAccessParams, SpendByResponseIdsRow, TraceSpansRow}, + resolve_trace, +}; +use litellm_traces_cache::{Error, SnapshotCache, SnapshotKey}; +use rstest::{fixture, rstest}; + +const T0: i64 = 1_790_742_989_000_000_000; +const MS: i64 = 1_000_000; +const TTL: Duration = Duration::from_secs(120); + +fn row(span_id: &str, parent: &str, name: &str, kind: &str, agent: &str) -> TraceSpansRow { + TraceSpansRow { + trace_id: String::new(), + span_id: span_id.into(), + parent_span_id: parent.into(), + name: name.into(), + kind: kind.parse().unwrap(), + wrapper_candidate: false, + agent: agent.into(), + framework: String::new(), + status: SpanStatus::Ok, + status_message: String::new(), + error_truncated: false, + start_ns: T0, + duration_ns: 10 * MS as u64, + service: "agent-demo".into(), + input_preview: format!("input of {name}"), + model: String::new(), + input_tokens: 0, + output_tokens: 0, + litellm_request_id: String::new(), + call_keys: Vec::new(), + call_evidence: None, + tool_call_id: String::new(), + team_id: String::new(), + api_key_hash: String::new(), + user_id: String::new(), + } +} + +fn access() -> ReadAccessParams { + ReadAccessParams { + all_teams: false, + user_id: String::new(), + team_ids: vec!["team".into()], + } +} + +fn key( + source: &str, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + ms: u64, +) -> SnapshotKey { + SnapshotKey::new(source, access, trace_id, trace_ref, ms).unwrap() +} + +#[fixture] +fn trace() -> Trace { + resolve_trace( + "trace", + "ref", + &[row("root", "", "run", "agent", "agent")], + &[] as &[SpendByResponseIdsRow], + ) + .expect("fixture should resolve") +} + +#[rstest] +#[case::different_team(false, "", "other-team")] +#[case::different_user(false, "other-user", "team")] +#[case::different_scope(true, "", "team")] +#[tokio::test] +async fn cached_trace_is_isolated_by_access_scope( + trace: Trace, + #[case] all_teams: bool, + #[case] user_id: &str, + #[case] team_id: &str, +) { + let cache = SnapshotCache::new(1024 * 1024, TTL); + let stored = key("source", &access(), "trace", "ref", 100); + + cache.insert(stored.clone(), trace.clone()).await.unwrap(); + + let other_access = ReadAccessParams { + all_teams, + user_id: user_id.into(), + team_ids: vec![team_id.into()], + }; + let other = key("source", &other_access, "trace", "ref", 100); + + assert!(cache.get(&other).await.is_none()); + let cached = cache.get(&stored).await.unwrap(); + assert_eq!(cached.trace(), &trace); +} + +#[rstest] +#[case::different_source("other-source", "trace", "ref", 100)] +#[case::different_trace_id("source", "other-trace", "ref", 100)] +#[case::different_trace_ref("source", "trace", "other-ref", 100)] +#[case::different_snapshot_ms("source", "trace", "ref", 200)] +#[tokio::test] +async fn cached_trace_is_isolated_by_key_fields( + trace: Trace, + #[case] source: &str, + #[case] trace_id: &str, + #[case] trace_ref: &str, + #[case] snapshot_ms: u64, +) { + let cache = SnapshotCache::new(1024 * 1024, TTL); + let stored = key("source", &access(), "trace", "ref", 100); + + cache.insert(stored.clone(), trace.clone()).await.unwrap(); + + let other = key(source, &access(), trace_id, trace_ref, snapshot_ms); + assert!(cache.get(&other).await.is_none()); + assert!(cache.get(&stored).await.is_some()); +} + +#[rstest] +#[tokio::test] +async fn snapshot_at_the_size_limit_is_accepted(trace: Trace) { + let size = serde_json::to_vec(&trace).unwrap().len(); + let cache = SnapshotCache::new(size, TTL); + let stored = key("source", &access(), "trace", "ref", 100); + + cache.insert(stored.clone(), trace).await.unwrap(); + assert!(cache.get(&stored).await.is_some()); +} + +#[rstest] +#[tokio::test] +async fn snapshot_one_byte_over_the_size_limit_is_rejected(trace: Trace) { + let size = serde_json::to_vec(&trace).unwrap().len(); + let cache = SnapshotCache::new(size - 1, TTL); + let stored = key("source", &access(), "trace", "ref", 100); + + assert!(matches!( + cache.insert(stored.clone(), trace).await, + Err(Error::ReadTooLarge) + )); + assert!(cache.get(&stored).await.is_none()); +} + +#[rstest] +#[case::same_ids(&["root", "child"], &["root", "child"], true)] +#[case::different_ids(&["root", "child"], &["root", "other"], false)] +#[tokio::test] +async fn snapshot_version_tracks_the_ordered_span_ids( + #[case] first_ids: &[&str], + #[case] second_ids: &[&str], + #[case] equal: bool, +) { + let build = |ids: &[&str]| -> Trace { + let rows: Vec = ids + .iter() + .map(|span_id| row(span_id, "", "run", "agent", "agent")) + .collect(); + resolve_trace("trace", "ref", &rows, &[] as &[SpendByResponseIdsRow]) + .expect("fixture should resolve") + }; + let cache = SnapshotCache::new(1024 * 1024, TTL); + + let first = cache + .insert(key("source", &access(), "a", "ref", 100), build(first_ids)) + .await + .unwrap(); + let second = cache + .insert(key("source", &access(), "b", "ref", 100), build(second_ids)) + .await + .unwrap(); + + assert_eq!(first.version() == second.version(), equal); +} + +#[rstest] +#[tokio::test] +async fn snapshots_expire_after_the_ttl(trace: Trace) { + let cache = SnapshotCache::new(1024 * 1024, Duration::from_millis(50)); + let stored = key("source", &access(), "trace", "ref", 100); + + cache.insert(stored.clone(), trace).await.unwrap(); + tokio::time::sleep(Duration::from_millis(200)).await; + + assert!(cache.get(&stored).await.is_none()); +} diff --git a/litellm-rust/crates/traces-clickhouse/Cargo.toml b/litellm-rust/crates/traces-clickhouse/Cargo.toml index 15b1ca8162e..6b95eb149d3 100644 --- a/litellm-rust/crates/traces-clickhouse/Cargo.toml +++ b/litellm-rust/crates/traces-clickhouse/Cargo.toml @@ -21,6 +21,7 @@ litellm-http.workspace = true litellm-migrate.workspace = true litellm-storage-clickhouse.workspace = true litellm-traces.workspace = true +litellm-traces-cache.workspace = true moka.workspace = true serde.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/traces-clickhouse/src/error.rs b/litellm-rust/crates/traces-clickhouse/src/error.rs index fdbc3029a5d..909ac5fc60f 100644 --- a/litellm-rust/crates/traces-clickhouse/src/error.rs +++ b/litellm-rust/crates/traces-clickhouse/src/error.rs @@ -47,3 +47,12 @@ pub enum Error { #[error(transparent)] Cached(#[from] std::sync::Arc), } + +impl From for Error { + fn from(error: litellm_traces_cache::Error) -> Self { + match error { + litellm_traces_cache::Error::Serialization(_) => Self::InvalidResponse, + litellm_traces_cache::Error::ReadTooLarge => Self::ReadTooLarge, + } + } +} diff --git a/litellm-rust/crates/traces-clickhouse/src/reads.rs b/litellm-rust/crates/traces-clickhouse/src/reads.rs index 4d338e17f07..6938aab06b3 100644 --- a/litellm-rust/crates/traces-clickhouse/src/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/src/reads.rs @@ -1,6 +1,6 @@ //! Scoped trace reads: the trace list, one trace resolved with its spend, and span payloads. -use std::sync::{Arc, LazyLock}; +use std::sync::LazyLock; use std::time::Duration; use base64::{Engine, engine::general_purpose::URL_SAFE}; @@ -12,9 +12,8 @@ use litellm_traces::{ SpanDetail, SpanErrorPage, SpendLookup, Trace, TracePage, listed_summary, query::named as contracts, resolve_trace, to_ui_content, }; -use moka::future::Cache; +use litellm_traces_cache::{SnapshotCache, SnapshotKey}; use serde::{Deserialize, Serialize}; -use sha2::{Digest, Sha256}; use crate::{ Connection, Error, @@ -38,17 +37,11 @@ impl Query for RunCandidates { } // Cursor pages share a bounded snapshot so advancing does not resolve the whole graph again. -static TRACE_SNAPSHOTS: LazyLock>> = LazyLock::new(|| { - Cache::builder() - .max_capacity((2 * crate::span_batches::MAX_GRAPH_BYTES) as u64) - .weigher(|_: &String, trace: &Arc| { - serde_json::to_vec(trace.as_ref()) - .ok() - .and_then(|bytes| u32::try_from(bytes.len().saturating_mul(2)).ok()) - .unwrap_or(u32::MAX) - }) - .time_to_live(Duration::from_secs(120)) - .build() +static TRACE_SNAPSHOTS: LazyLock = LazyLock::new(|| { + SnapshotCache::new( + crate::span_batches::MAX_GRAPH_BYTES, + Duration::from_secs(120), + ) }); const NANOS_PER_MS: i64 = 1_000_000; @@ -339,74 +332,57 @@ pub async fn get_trace_page( version: String::new(), }, }; - let key_bytes = serde_json::to_vec(&( + let key = SnapshotKey::new( connection.url().as_str(), access, trace_id, &trace_ref, position.snapshot_ms, - )) + ) .map_err(|_| Error::InvalidParameters)?; - let key = format!("{:x}", Sha256::digest(key_bytes)); - let snapshot = if let Some(trace) = TRACE_SNAPSHOTS.get(&key).await { - trace - } else { - let params = TraceSpansParams { - access: access.clone(), - trace_id: trace_id.to_owned(), - trace_ref: trace_ref.clone(), - }; - let rows = - crate::span_batches::read_spans(client, connection, params, position.snapshot_ms) - .await?; - let spend_rows = spend(client, connection, access, &rows).await; - let Some(trace) = resolve_trace(trace_id, &trace_ref, &rows, &spend_rows) else { - return Ok(None); - }; - if serde_json::to_vec(&trace) - .map_err(|_| Error::InvalidResponse)? - .len() - > crate::span_batches::MAX_GRAPH_BYTES - { - return Err(Error::ReadTooLarge); + let snapshot = match TRACE_SNAPSHOTS.get(&key).await { + Some(snapshot) => snapshot, + None => { + let params = TraceSpansParams { + access: access.clone(), + trace_id: trace_id.to_owned(), + trace_ref: trace_ref.clone(), + }; + let rows = + crate::span_batches::read_spans(client, connection, params, position.snapshot_ms) + .await?; + let spend_rows = spend(client, connection, access, &rows).await; + let Some(trace) = resolve_trace(trace_id, &trace_ref, &rows, &spend_rows) else { + return Ok(None); + }; + TRACE_SNAPSHOTS.insert(key, trace).await? } - let trace = Arc::new(trace); - TRACE_SNAPSHOTS.insert(key, Arc::clone(&trace)).await; - trace }; - let span_ids: Vec<&str> = snapshot - .spans - .iter() - .map(|span| span.span_id.as_str()) - .collect(); - let version = format!( - "{:x}", - Sha256::digest(serde_json::to_vec(&span_ids).map_err(|_| Error::InvalidResponse)?) - ); - if cursor.is_some() && position.version != version { + let spans = &snapshot.trace().spans; + if cursor.is_some() && position.version != snapshot.version() { return Err(Error::TraceChanged); } let mut trace = Trace { - summary: snapshot.summary.clone(), - agents: snapshot.agents.clone(), + summary: snapshot.trace().summary.clone(), + agents: snapshot.trace().agents.clone(), spans: Vec::new(), next_cursor: None, }; - if position.offset > snapshot.spans.len() { + if position.offset > spans.len() { return Err(Error::InvalidCursor("span")); } let end = position .offset .saturating_add(page_size as usize) - .min(snapshot.spans.len()); - trace.next_cursor = (end < snapshot.spans.len()).then(|| { + .min(spans.len()); + trace.next_cursor = (end < spans.len()).then(|| { encode_cursor(&SpanPosition { offset: end, - version: version.clone(), + version: snapshot.version().to_owned(), ..position }) }); - trace.spans = snapshot.spans[position.offset..end].to_vec(); + trace.spans = spans[position.offset..end].to_vec(); while serde_json::to_vec(&trace) .map_err(|_| Error::InvalidResponse)? .len() @@ -420,7 +396,7 @@ pub async fn get_trace_page( trace_ref: trace_ref.clone(), snapshot_ms: position.snapshot_ms, offset: position.offset + trace.spans.len(), - version: version.clone(), + version: snapshot.version().to_owned(), })); } Ok(Some(trace))