diff --git a/Makefile b/Makefile index ad84e2c4014..7c8511a44d7 100644 --- a/Makefile +++ b/Makefile @@ -58,7 +58,7 @@ help: @echo " make test-unit-helm - Run helm unit tests" @echo " make test-rust-extension - Build the Rust extension and run its public Python tests" @echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container" - @echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (ARGS=\"--seed large\", LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)" + @echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (ARGS=\"--seed large --seed-logs\", LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)" @echo "" @echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide" @echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine." diff --git a/deploy/lens/README.md b/deploy/lens/README.md index 6f8c7d3940e..60e9c9991a4 100644 --- a/deploy/lens/README.md +++ b/deploy/lens/README.md @@ -178,16 +178,15 @@ Creation queues the first batch. Posting to `/lens/{id}/runs` queues another, or The default is Next.js dev with no production build (`LENS_DEV_BUILD_UI=0`). Set `LENS_DEV_BUILD_UI=1` when you also want a fresh static dashboard at `http://localhost:4000/ui/`. Build output goes to `.lens-dev/logs/ui-build.log`; a failed build stops startup. Both modes keep the live dashboard on port 3000. Startup checks the live login route before seeding and fails with the UI log path if Next.js exits. `LENS_DEV_STARTUP_TIMEOUT_SECONDS` controls startup readiness retries (default 300; `LENS_DEV_READINESS_REQUEST_TIMEOUT_SECONDS` caps each HTTP probe, default 5) -For local fixture data, run `make lens-dev ARGS=--seed`. Use `make lens-dev ARGS="--seed large"` for 2,000 fixture copies, over one million spans and linked request logs. To seed a running stack without restarting it, use `make lens-dev ARGS="--seed-only --seed large --copies 100"`. The default profile replays one copy of every checked-in capture through authenticated `/v1/traces`, including failures, retries, streaming and multiple agent frameworks. Large seeds use the same parser and compressed ClickHouse writer in batches of four copies, and write matching request logs to PostgreSQL. The first and last batches verify linked spend totals through the proxy +For local fixture data, run `make lens-dev ARGS=--seed`. Use `make lens-dev ARGS="--seed large"` for 2,000 fixture copies spread over the last 24 hours, about 860,000 spans with linked request logs, plus three long sessions of roughly 1,150, 9,200 and 92,000 spans in a single trace for drawer paging and the oversized read path. Their trace IDs are printed at the end. To seed a running stack without restarting it, use `make lens-dev ARGS="--seed-only --seed large --copies 100"`. Every profile replays one copy of every checked-in capture through authenticated `/v1/traces`, including failures, retries, streaming and multiple agent frameworks, and verifies linked spend totals through the proxy. Large seeds then copy that first copy inside ClickHouse and PostgreSQL with `INSERT ... SELECT`, rewriting trace, span and call IDs so each copy keeps its own spend, and verify the last copy through the proxy -Seeds append fresh IDs on every invocation and spread copies over recent timestamps. Restarts without `SEED` do not add data. Lens excludes activity received in the last two minutes, so wait two minutes after seeding before checking investigation previews. `LENS_DEV_SEED_COPIES` overrides total copies, and `LENS_DEV_SEED_BATCH_COPIES` overrides copies per bulk insert (default 4, about 2,000 spans). Start with four or fewer on a constrained machine. Larger batches still respect the existing ClickHouse insert size limit; each capture is decoded separately within the OTLP safety budget. Large seeds test data volume and pagination, rather than concurrent ingestion throughput or review accuracy. They can use substantial disk space; adjust `--copies` for your machine. Seeding expects the generated local tracing configuration. The old `run_tracing_proxy_local.sh --seed` command forwards to Lens dev, using its ports and saved master key +Seeds append fresh IDs on every invocation and spread copies over recent timestamps. Restarts without `SEED` do not add data. Lens excludes activity received in the last two minutes, so wait two minutes after seeding before checking investigation previews. `LENS_DEV_SEED_COPIES` overrides total copies. Large seeds test data volume and pagination, rather than concurrent ingestion throughput or review accuracy. They can use substantial disk space; adjust `--copies` for your machine. Seeding expects the generated local tracing configuration. The old `run_tracing_proxy_local.sh --seed` command forwards to Lens dev, using its ports and saved master key Local ingestion limits are explicit and configurable. Set OTLP and ClickHouse variables before starting the proxy and seeder so both processes use the same settings. Invalid, zero and negative values fail instead of silently falling back. Changing these limits does not require rebuilding Rust | Environment variable | Default | Controls | | --- | --- | --- | | `LENS_DEV_SEED_COPIES` | 1 default, 2000 large | Total fixture copies | -| `LENS_DEV_SEED_BATCH_COPIES` | 4 | Copies per bulk insert | | `LENS_DEV_SEED_TIMEOUT_SECONDS` | 120 | Seeder HTTP timeout | | `OTLP_MAX_BODY_BYTES` | 16777216 | HTTP body and decompressed payload bytes | | `OTLP_MAX_CONCURRENT_INGESTS` | 2 | Concurrent proxy ingestion requests | @@ -202,7 +201,7 @@ Local ingestion limits are explicit and configurable. Set OTLP and ClickHouse va | `CLICKHOUSE_TRACE_MAX_INSERT_BYTES` | 67108864 | Encoded trace or spend insert bytes | | `CLICKHOUSE_INSERT_TIMEOUT_SECONDS` | 30 | ClickHouse insert HTTP timeout | -The wire parsers also enforce their library recursion limits (128 levels for JSON, 100 for protobuf). Raising the configured depth does not remove those parser limits. Bulk seeding parses each capture separately, keeping the per-export limits distinct from the bulk insert limit. Use smaller batches if an insert exceeds its byte budget. For example, `LENS_DEV_SEED_COPIES=100 LENS_DEV_SEED_BATCH_COPIES=2 make lens-dev ARGS="--seed large"` +The wire parsers also enforce their library recursion limits (128 levels for JSON, 100 for protobuf). Raising the configured depth does not remove those parser limits. ## Quality evaluation diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 348beeb3813..f7b667c8ab2 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4156,6 +4156,7 @@ dependencies = [ "litellm-storage-clickhouse", "litellm-token-counter", "litellm-traces", + "litellm-traces-cache", "litellm-traces-clickhouse", "litellm-tracing", "prost", @@ -4474,13 +4475,17 @@ dependencies = [ name = "litellm-traces-cache" version = "0.1.0" dependencies = [ + "base64 0.22.1", "litellm-traces", "moka", "rstest", + "serde", "serde_json", "sha2 0.10.9", "thiserror 2.0.19", + "time", "tokio", + "tracing", ] [[package]] @@ -4488,11 +4493,9 @@ name = "litellm-traces-clickhouse" version = "0.1.0" dependencies = [ "askama", - "base64 0.22.1", "flate2", "futures-util", "hmac 0.12.1", - "itertools 0.14.0", "jsonschema", "litellm-http", "litellm-migrate", @@ -4511,7 +4514,6 @@ dependencies = [ "thiserror 2.0.19", "time", "tokio", - "tracing", "url", "wiremock", ] diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index c3a86009111..226cecd5b55 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -22,6 +22,7 @@ tiktoken = ["litellm-token-counter/tiktoken"] fancy-regex.workspace = true litellm-tracing.workspace = true litellm-traces.workspace = true +litellm-traces-cache.workspace = true litellm-traces-clickhouse.workspace = true litellm-storage-clickhouse.workspace = true litellm-host.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index dd5c6b9860f..f442d6c31de 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -1,8 +1,11 @@ -use std::collections::BTreeMap; +use std::{collections::BTreeMap, sync::Arc}; use litellm_http::ClientVariant; use litellm_traces::{QueryScope, ReadQuery, Tenant, query::named::ReadAccessParams}; -use litellm_traces_clickhouse::{Config, Error, InsertTable, Parameter, QueryReaders}; +use litellm_traces_cache::{ReadError, TraceReader}; +use litellm_traces_clickhouse::{ + ClickHouseTraces, Config, Error, InsertTable, Parameter, QueryReaders, +}; use prost::Message; use pyo3::{ exceptions::{PyOverflowError, PyRuntimeError, PyValueError}, @@ -10,6 +13,8 @@ use pyo3::{ types::PyBytes, }; +pyo3::import_exception!(litellm.rust_bridge.trace.errors, TraceChanged); + #[derive(Message)] struct OtlpErrorStatus { #[prost(int32, tag = "1")] @@ -35,15 +40,12 @@ fn map_error_ref(error: &Error) -> PyErr { use litellm_storage_clickhouse::Error as StorageError; match error { - Error::Decode(litellm_traces::Error::TooLarge) - | Error::InsertTooLarge - | Error::ReadTooLarge => PyOverflowError::new_err(error.to_string()), + Error::Decode(litellm_traces::Error::TooLarge) | Error::InsertTooLarge => { + PyOverflowError::new_err(error.to_string()) + } Error::InvalidRow | Error::InvalidLimit(_) | Error::InvalidTable - | Error::InvalidCursor(_) - | Error::AmbiguousTrace - | Error::TraceChanged | Error::Decode(_) | Error::InvalidSchema | Error::InvalidQuery @@ -78,6 +80,18 @@ fn map_error_ref(error: &Error) -> PyErr { } } +fn map_read_error(error: ReadError) -> PyErr { + match error { + error @ (ReadError::InvalidParameters + | ReadError::InvalidCursor(_) + | ReadError::AmbiguousTrace) => PyValueError::new_err(error.to_string()), + error @ ReadError::TraceChanged => TraceChanged::new_err(error.to_string()), + error @ ReadError::TooLarge => PyOverflowError::new_err(error.to_string()), + error @ ReadError::Encode(_) => PyRuntimeError::new_err(error.to_string()), + ReadError::Store(error) => map_error_ref(&error), + } +} + fn map_sql_error(error: Error) -> PyErr { match error { Error::Storage(litellm_storage_clickhouse::Error::QueryFailed(400 | 404)) => { @@ -112,6 +126,7 @@ impl NativeTraceConfig { pub struct NativeTraceStorage { config: Config, query_readers: QueryReaders, + reader: Arc, } #[pymethods] @@ -123,6 +138,9 @@ impl NativeTraceStorage { config.inner.storage().writer().clone(), config.inner.storage().database().to_owned(), ), + reader: Arc::new(TraceReader::new( + litellm_storage_clickhouse::READ_LIMITS.response_bytes, + )), config: config.inner.clone(), }) } @@ -218,21 +236,16 @@ impl NativeTraceStorage { ) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().reader().clone(); + let reader = Arc::clone(&self.reader); crate::execution::run_async( py, async move { - litellm_traces_clickhouse::list_traces( - &client, - &connection, - &scope, - start_ms, - end_ms, - cursor.as_deref(), - limit, - ) - .await + let store = ClickHouseTraces::new(client, connection); + reader + .list_traces(&store, &scope, start_ms, end_ms, cursor.as_deref(), limit) + .await }, - map_error, + map_read_error, ) } @@ -248,34 +261,31 @@ impl NativeTraceStorage { ) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().reader().clone(); + let reader = Arc::clone(&self.reader); crate::execution::run_async( py, async move { + let store = ClickHouseTraces::new(client, connection); if let Some(page_size) = page_size { - litellm_traces_clickhouse::get_trace_page( - &client, - &connection, - &scope, - &trace_id, - &trace_ref, - cursor.as_deref(), - page_size, - ) - .await + reader + .get_trace_page( + &store, + &scope, + &trace_id, + &trace_ref, + cursor.as_deref(), + page_size, + ) + .await } else if cursor.is_some() { - Err(Error::InvalidParameters) + Err(ReadError::InvalidParameters) } else { - litellm_traces_clickhouse::get_trace( - &client, - &connection, - &scope, - &trace_id, - &trace_ref, - ) - .await + reader + .get_trace(&store, &scope, &trace_id, &trace_ref) + .await } }, - map_error, + map_read_error, ) } @@ -289,20 +299,16 @@ impl NativeTraceStorage { ) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().reader().clone(); + let reader = Arc::clone(&self.reader); crate::execution::run_async( py, async move { - litellm_traces_clickhouse::get_span( - &client, - &connection, - &scope, - &trace_id, - &span_id, - &trace_ref, - ) - .await + let store = ClickHouseTraces::new(client, connection); + reader + .get_span(&store, &scope, &trace_id, &span_id, &trace_ref) + .await }, - map_error, + map_read_error, ) } @@ -318,21 +324,23 @@ impl NativeTraceStorage { ) -> PyResult> { let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; let connection = self.config.storage().reader().clone(); + let reader = Arc::clone(&self.reader); crate::execution::run_async( py, async move { - litellm_traces_clickhouse::get_span_error( - &client, - &connection, - &scope, - &trace_id, - &span_id, - &trace_ref, - cursor.as_deref(), - ) - .await + let store = ClickHouseTraces::new(client, connection); + reader + .get_span_error( + &store, + &scope, + &trace_id, + &span_id, + &trace_ref, + cursor.as_deref(), + ) + .await }, - map_error, + map_read_error, ) } @@ -495,11 +503,7 @@ mod tests { Error::Decode(litellm_traces::Error::InvalidLimit("OTLP_MAX_SPANS")), "ValueError" )] - #[case::cursor(Error::InvalidCursor("trace"), "ValueError")] - #[case::ambiguous(Error::AmbiguousTrace, "ValueError")] - #[case::changed_snapshot(Error::TraceChanged, "ValueError")] - #[case::read_budget(Error::ReadTooLarge, "OverflowError")] - fn trace_read_and_ingest_failures_preserve_public_exception_types( + fn trace_ingest_failures_preserve_public_exception_types( #[case] error: Error, #[case] exception_name: &str, ) { @@ -511,4 +515,43 @@ mod tests { ); }); } + + #[rstest] + #[case::invalid_parameters(ReadError::InvalidParameters, "ValueError")] + #[case::invalid_cursor(ReadError::InvalidCursor("trace"), "ValueError")] + #[case::ambiguous(ReadError::AmbiguousTrace, "ValueError")] + #[case::changed_snapshot(ReadError::TraceChanged, "TraceChanged")] + #[case::read_budget(ReadError::TooLarge, "OverflowError")] + #[case::encode( + ReadError::Encode(Arc::new(serde_json::Error::io(std::io::Error::other("invalid")))), + "RuntimeError" + )] + #[case::store(ReadError::Store(Arc::new(Error::InvalidScope)), "ValueError")] + fn trace_read_failures_preserve_public_exception_types( + #[case] error: ReadError, + #[case] exception_name: &str, + ) { + Python::initialize(); + Python::attach(|py| { + let repository = std::path::Path::new(env!("CARGO_MANIFEST_DIR")) + .ancestors() + .nth(3) + .unwrap() + .to_str() + .unwrap(); + pyo3::types::PyModule::import(py, "sys") + .unwrap() + .getattr("path") + .unwrap() + .call_method1("insert", (0, repository)) + .unwrap(); + let message = error.to_string(); + let exception = map_read_error(error); + assert_eq!(exception.get_type(py).name().unwrap(), exception_name); + assert_eq!( + exception.value(py).str().unwrap().to_str().unwrap(), + message + ); + }); + } } diff --git a/litellm-rust/crates/traces-cache/AGENTS.md b/litellm-rust/crates/traces-cache/AGENTS.md index bfaea42d901..7e4caacd15e 100644 --- a/litellm-rust/crates/traces-cache/AGENTS.md +++ b/litellm-rust/crates/traces-cache/AGENTS.md @@ -1,5 +1,5 @@ -Own resolved-trace snapshot storage, cache identity, weighting, and expiry -Depend on trace domain types, never storage, HTTP, or Python +Own storage-independent trace reads over `TraceStore`: the in-process read cache (identity, single-flight, freshness expiry, weighting), cursor formats, paging, response splitting, spend windows and run batching +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 +Storage adapters implement `TraceStore`; keep SQL and row encoding there diff --git a/litellm-rust/crates/traces-cache/Cargo.toml b/litellm-rust/crates/traces-cache/Cargo.toml index 3e58ea4c965..16eb01c6c8a 100644 --- a/litellm-rust/crates/traces-cache/Cargo.toml +++ b/litellm-rust/crates/traces-cache/Cargo.toml @@ -6,11 +6,15 @@ license.workspace = true repository.workspace = true [dependencies] +base64.workspace = true litellm-traces.workspace = true moka.workspace = true +serde.workspace = true serde_json.workspace = true sha2.workspace = true +time.workspace = true thiserror.workspace = true +tracing.workspace = true [dev-dependencies] rstest.workspace = true diff --git a/litellm-rust/crates/traces-cache/src/cache.rs b/litellm-rust/crates/traces-cache/src/cache.rs new file mode 100644 index 00000000000..f308714146e --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/cache.rs @@ -0,0 +1,395 @@ +use std::{future::Future, sync::Arc, time::Duration}; + +use litellm_traces::{ + Trace, TraceSummary, + query::named::{ReadAccessParams, TraceSpansRow}, +}; +use moka::{Expiry, future::Cache}; +use serde::Serialize; +use sha2::{Digest, Sha256}; + +use crate::Error; + +pub const LIVE_TTL: Duration = Duration::from_secs(5); +pub const SETTLED_TTL: Duration = Duration::from_secs(10 * 60); +const SETTLED_AFTER_MS: u64 = 5 * 60 * 1000; +const MAX_INDEX_ENTRIES: u64 = 100_000; + +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct SnapshotKey(String); + +impl SnapshotKey { + fn digest(fields: &impl Serialize) -> Result { + Ok(Self(format!( + "{:x}", + Sha256::digest(serde_json::to_vec(fields)?) + ))) + } + + pub fn new( + source: &str, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + snapshot_ms: u64, + ) -> Result { + Self::digest(&(source, access, trace_id, trace_ref, snapshot_ms)) + } + + pub fn latest( + source: &str, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + ) -> Result { + Self::digest(&(source, access, trace_id, trace_ref)) + } + + pub(crate) fn run( + source: &str, + access: &ReadAccessParams, + run: (&str, &str, &str, &str), + ) -> Result { + Self::digest(&("run", source, access, run)) + } + + pub(crate) fn scope(source: &str, access: &ReadAccessParams) -> Result { + Self::digest(&("scope", source, access)) + } +} + +/// How long a read result stays reusable: traces still receiving spans, or read with spend +/// unavailable, are re-read after `LIVE_TTL`; traces quiet for `SETTLED_AFTER_MS` are kept for +/// `SETTLED_TTL`. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum Freshness { + Live, + Settled, +} + +impl Freshness { + pub fn of(rows: &[TraceSpansRow], spend_known: bool, snapshot_ms: u64) -> Self { + let last_end_ms = rows + .iter() + .map(|row| row.start_ns.saturating_add_unsigned(row.duration_ns) / 1_000_000) + .max() + .unwrap_or(i64::MAX); + let quiet_ms = i64::try_from(snapshot_ms) + .unwrap_or(i64::MAX) + .saturating_sub(last_end_ms); + if spend_known && quiet_ms >= SETTLED_AFTER_MS as i64 { + Self::Settled + } else { + Self::Live + } + } + + fn ttl(self) -> Duration { + match self { + Self::Live => LIVE_TTL, + Self::Settled => SETTLED_TTL, + } + } +} + +trait Fresh { + fn freshness(&self) -> Freshness; +} + +struct ByFreshness; + +impl Expiry for ByFreshness { + fn expire_after_create(&self, _: &K, value: &V, _: std::time::Instant) -> Option { + Some(value.freshness().ttl()) + } +} + +pub struct Snapshot { + trace: Trace, + version: String, + snapshot_ms: u64, + freshness: Freshness, + weight: u32, +} + +impl Snapshot { + pub fn trace(&self) -> &Trace { + &self.trace + } + + pub fn version(&self) -> &str { + &self.version + } + + pub fn snapshot_ms(&self) -> u64 { + self.snapshot_ms + } + + pub fn freshness(&self) -> Freshness { + self.freshness + } +} + +#[derive(Clone, Copy)] +struct Latest { + snapshot_ms: u64, + freshness: Freshness, +} + +impl Fresh for Latest { + fn freshness(&self) -> Freshness { + self.freshness + } +} + +/// Resolved trace snapshots pinned by `snapshot_ms` for paging, plus which snapshot each trace +/// currently serves so repeated opens reuse one read until its freshness expires. +pub struct SnapshotCache { + pinned: Cache>, + latest: Cache, + max_graph_bytes: usize, +} + +impl SnapshotCache { + pub fn new(max_graph_bytes: usize, idle: Duration) -> Self { + Self { + pinned: Cache::builder() + .max_capacity((max_graph_bytes as u64).saturating_mul(2)) + .weigher(|_: &SnapshotKey, snapshot: &Arc| snapshot.weight) + .time_to_idle(idle) + .build(), + latest: Cache::builder() + .max_capacity(MAX_INDEX_ENTRIES) + .expire_after(ByFreshness) + .build(), + max_graph_bytes, + } + } + + pub async fn get(&self, key: &SnapshotKey) -> Option> { + self.pinned.get(key).await + } + + /// Returns the snapshot pinned at `key`, running `load` once for all concurrent callers on a + /// miss. A failed load is not cached. + pub async fn pinned_or_load( + &self, + key: SnapshotKey, + snapshot_ms: u64, + load: F, + ) -> Result, Arc> + where + E: From + Send + Sync + 'static, + F: Future>, + { + self.pinned + .try_get_with(key, async { + let (trace, freshness) = load.await?; + Ok(Arc::new(self.snapshot(trace, snapshot_ms, freshness)?)) + }) + .await + } + + /// Returns the snapshot `latest` currently serves. On a miss, `load_at(now_ms)` runs once for + /// all concurrent callers and its snapshot is served until its freshness expires. + pub async fn latest_or_load( + &self, + latest: SnapshotKey, + now_ms: u64, + load_at: F, + ) -> Result, Arc> + where + E: Send + Sync + 'static, + F: Fn(u64) -> Fut, + Fut: Future, Arc>>, + { + let entry = self + .latest + .try_get_with(latest, async { + let snapshot = load_at(now_ms).await?; + Ok::<_, Arc>(Latest { + snapshot_ms: snapshot.snapshot_ms, + freshness: snapshot.freshness, + }) + }) + .await + .map_err(|error| Arc::clone(&*error))?; + load_at(entry.snapshot_ms).await + } + + fn snapshot( + &self, + trace: Trace, + snapshot_ms: u64, + freshness: Freshness, + ) -> Result { + 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)?)); + Ok(Snapshot { + trace, + version, + snapshot_ms, + freshness, + weight: u32::try_from(encoded.len().saturating_mul(2)).unwrap_or(u32::MAX), + }) + } + + #[cfg(test)] + async fn weighted_size(&self) -> u64 { + self.pinned.run_pending_tasks().await; + self.pinned.weighted_size() + } +} + +#[derive(Clone)] +pub(crate) enum ListedRun { + Resolved(Box, Freshness), + Limited, +} + +impl Fresh for ListedRun { + fn freshness(&self) -> Freshness { + match self { + Self::Resolved(_, freshness) => *freshness, + Self::Limited => Freshness::Settled, + } + } +} + +pub(crate) struct ListCache { + pub(crate) runs: Cache, + pub(crate) limits: Cache, +} + +impl ListCache { + pub(crate) fn new() -> Self { + Self { + runs: Cache::builder() + .max_capacity(MAX_INDEX_ENTRIES) + .expire_after(ByFreshness) + .build(), + limits: Cache::builder() + .max_capacity(MAX_INDEX_ENTRIES) + .time_to_live(SETTLED_TTL) + .build(), + } + } +} + +#[cfg(test)] +mod tests { + use litellm_traces::{ + SpanStatus, + query::named::{SpendByResponseIdsRow, TraceSpansRow}, + resolve_trace, + }; + use rstest::rstest; + + use super::*; + + fn row(span_id: &str) -> TraceSpansRow { + 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(), + } + } + + fn trace(span_id: &str) -> Trace { + resolve_trace( + "trace", + "ref", + &[row(span_id)], + &[] 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 + .pinned_or_load(key, 100, async { + Ok::<_, Error>((trace(span_id), Freshness::Settled)) + }) + .await + .unwrap(); + } + + assert!(cache.weighted_size().await <= (limit as u64) * 2); + } + + const LAST_END_MS: u64 = 1_790_742_989_010; + + #[rstest] + #[case::just_ended(LAST_END_MS, true, Freshness::Live)] + #[case::quiet_just_under(LAST_END_MS + SETTLED_AFTER_MS - 1, true, Freshness::Live)] + #[case::quiet_long_enough(LAST_END_MS + SETTLED_AFTER_MS, true, Freshness::Settled)] + #[case::spend_unknown(LAST_END_MS + SETTLED_AFTER_MS, false, Freshness::Live)] + fn freshness_settles_once_spans_stop_and_spend_is_known( + #[case] snapshot_ms: u64, + #[case] spend_known: bool, + #[case] expected: Freshness, + ) { + assert_eq!( + Freshness::of(&[row("root")], spend_known, snapshot_ms), + expected + ); + } +} diff --git a/litellm-rust/crates/traces-cache/src/cursor.rs b/litellm-rust/crates/traces-cache/src/cursor.rs new file mode 100644 index 00000000000..8858f243f09 --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/cursor.rs @@ -0,0 +1,124 @@ +use base64::{Engine, engine::general_purpose::URL_SAFE}; +use serde::{Deserialize, Serialize}; + +use crate::ReadError; + +pub(super) fn encode_cursor(position: &T) -> String { + URL_SAFE.encode(serde_json::to_vec(position).unwrap_or_default()) +} + +pub(super) fn decode_cursor Deserialize<'de>, E>( + cursor: &str, + kind: &'static str, +) -> Result> { + URL_SAFE + .decode(cursor) + .ok() + .and_then(|json| serde_json::from_slice(&json).ok()) + .ok_or(ReadError::InvalidCursor(kind)) +} + +pub(super) fn trace_position(cursor: Option<&str>) -> Result<(i64, String), ReadError> { + let Some(cursor) = cursor.filter(|cursor| !cursor.is_empty()) else { + return Ok((0, String::new())); + }; + match decode_cursor::<(i64, String), E>(cursor, "trace")? { + (start_ms, trace_ref) if start_ms > 0 && !trace_ref.is_empty() => Ok((start_ms, trace_ref)), + _ => Err(ReadError::InvalidCursor("trace")), + } +} + +#[derive(Deserialize, Serialize)] +pub(super) struct ErrorPosition { + pub(super) offset: u64, + pub(super) version: String, +} + +pub(super) fn error_position( + cursor: Option<&str>, +) -> Result, ReadError> { + let Some(cursor) = cursor else { + return Ok(None); + }; + let position: ErrorPosition = decode_cursor(cursor, "diagnostic")?; + let valid_version = position.version.len() == 64 + && position + .version + .bytes() + .all(|byte| byte.is_ascii_digit() || (b'A'..=b'F').contains(&byte)); + if i64::try_from(position.offset).is_err() || !valid_version { + return Err(ReadError::InvalidCursor("diagnostic")); + } + Ok(Some(position)) +} + +#[derive(Deserialize, Serialize)] +pub(super) struct SpanPosition { + pub(super) trace_ref: String, + pub(super) snapshot_ms: u64, + pub(super) offset: usize, + pub(super) version: String, +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + #[rstest] + fn trace_cursor_round_trips_the_last_listed_run() { + let cursor = encode_cursor(&(1_790_742_989_377_i64, "4bad42b84e9de3ba46fc870185f8f023")); + assert_eq!( + trace_position::(Some(&cursor)).unwrap(), + ( + 1_790_742_989_377, + "4bad42b84e9de3ba46fc870185f8f023".to_owned() + ) + ); + assert_eq!( + trace_position::(None).unwrap(), + (0, String::new()) + ); + assert_eq!( + trace_position::(Some("")).unwrap(), + (0, String::new()) + ); + } + + #[rstest] + #[case::not_base64("abc")] + #[case::not_json("bm90LWpzb24=")] + #[case::numeric_reference("WzEsIDJd")] + #[case::zero_start("WzAsICJ0Il0=")] + fn malformed_trace_cursors_are_rejected(#[case] cursor: &str) { + let result: Result<(i64, String), ReadError> = trace_position(Some(cursor)); + assert!(matches!(result, Err(ReadError::InvalidCursor("trace")))); + } + + #[rstest] + #[case::not_base64("garbage")] + #[case::missing_fields("e30=")] + #[case::not_an_object("WzEsMl0=")] + fn malformed_diagnostic_cursors_are_rejected(#[case] cursor: &str) { + let result: Result, ReadError> = + error_position(Some(cursor)); + assert!(matches!( + result, + Err(ReadError::InvalidCursor("diagnostic")) + )); + } + + #[rstest] + #[case::lowercase_version("a".repeat(64))] + #[case::short_version("A".repeat(63))] + fn diagnostic_cursor_requires_a_content_version(#[case] version: String) { + let cursor = encode_cursor(&ErrorPosition { offset: 1, version }); + let result: Result, ReadError> = + error_position(Some(&cursor)); + assert!(matches!( + result, + Err(ReadError::InvalidCursor("diagnostic")) + )); + } +} diff --git a/litellm-rust/crates/traces-cache/src/error.rs b/litellm-rust/crates/traces-cache/src/error.rs index ac5f19369fe..c7aaa3fd332 100644 --- a/litellm-rust/crates/traces-cache/src/error.rs +++ b/litellm-rust/crates/traces-cache/src/error.rs @@ -1,3 +1,5 @@ +use std::sync::Arc; + #[derive(Debug, thiserror::Error)] pub enum Error { #[error("trace snapshot serialization failed")] @@ -5,3 +7,45 @@ pub enum Error { #[error("trace snapshot exceeds the size limit")] ReadTooLarge, } + +/// Cheap to clone so one failed single-flight read can be returned to every waiting caller. +#[derive(Debug, thiserror::Error)] +pub enum ReadError { + #[error("invalid trace read parameters")] + InvalidParameters, + #[error("Invalid {0} cursor")] + InvalidCursor(&'static str), + #[error("Multiple traces have this ID; provide trace_ref")] + AmbiguousTrace, + #[error("Trace changed while paging; refresh the trace to continue")] + TraceChanged, + #[error("Trace exceeds the interactive read budget; use a filtered trace query")] + TooLarge, + #[error("trace could not be encoded")] + Encode(#[source] Arc), + #[error(transparent)] + Store(Arc), +} + +impl Clone for ReadError { + fn clone(&self) -> Self { + match self { + Self::InvalidParameters => Self::InvalidParameters, + Self::InvalidCursor(kind) => Self::InvalidCursor(kind), + Self::AmbiguousTrace => Self::AmbiguousTrace, + Self::TraceChanged => Self::TraceChanged, + Self::TooLarge => Self::TooLarge, + Self::Encode(error) => Self::Encode(Arc::clone(error)), + Self::Store(error) => Self::Store(Arc::clone(error)), + } + } +} + +impl From for ReadError { + fn from(error: Error) -> Self { + match error { + Error::ReadTooLarge => Self::TooLarge, + Error::Serialization(error) => Self::Encode(Arc::new(error)), + } + } +} diff --git a/litellm-rust/crates/traces-cache/src/lib.rs b/litellm-rust/crates/traces-cache/src/lib.rs index 7dfd3bf32a1..8ed79c1a547 100644 --- a/litellm-rust/crates/traces-cache/src/lib.rs +++ b/litellm-rust/crates/traces-cache/src/lib.rs @@ -1,166 +1,12 @@ -use std::{sync::Arc, time::Duration}; - -use litellm_traces::{Trace, query::named::ReadAccessParams}; -use moka::future::Cache; -use sha2::{Digest, Sha256}; - +mod cache; +mod cursor; mod error; +mod list; +mod reader; +mod spend; +mod store; -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); - } -} +pub use cache::{Freshness, LIVE_TTL, SETTLED_TTL, Snapshot, SnapshotCache, SnapshotKey}; +pub use error::{Error, ReadError}; +pub use reader::{MAX_GRAPH_BYTES, MAX_GRAPH_SPANS, TraceReader}; +pub use store::{StoreError, TraceStore}; diff --git a/litellm-rust/crates/traces-cache/src/list.rs b/litellm-rust/crates/traces-cache/src/list.rs new file mode 100644 index 00000000000..1b1dcbbefc3 --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/list.rs @@ -0,0 +1,197 @@ +use std::collections::HashMap; + +use crate::{ + ReadError, SnapshotKey, TraceReader, TraceStore, + cache::{Freshness, ListedRun}, + reader::{map_store_error, now_ms}, + spend::{spend, spend_window, spend_within}, + store::StoreError, +}; +use litellm_traces::{ + TraceSummary, listed_summary, + query::named::{ListTracesRow, ReadAccessParams, TracePageSpansParams, TraceSpansRow}, + resolve_trace, +}; + +const RUNS_PER_SPAN_READ: usize = 16; + +fn run_key(team_id: &str, api_key_hash: &str, trace_id: &str) -> (String, String, String) { + ( + team_id.to_owned(), + api_key_hash.to_owned(), + trace_id.to_owned(), + ) +} + +fn cache_key( + source: &str, + access: &ReadAccessParams, + row: &ListTracesRow, +) -> Result> { + Ok(SnapshotKey::run( + source, + access, + ( + &row.team_id, + &row.api_key_hash, + &row.trace_id, + &row.trace_ref, + ), + )?) +} + +fn summary(row: &ListTracesRow, listed: Option<&ListedRun>) -> TraceSummary { + match listed { + Some(ListedRun::Resolved(summary, _)) => (**summary).clone(), + Some(ListedRun::Limited) | None => listed_summary(row), + } +} + +/// Summaries for one batch of listed runs. Runs resolved within their freshness window come +/// from the cache; only the rest are read from storage, with one span and one spend read. +pub(super) async fn list_summaries( + reader: &TraceReader, + store: &S, + access: &ReadAccessParams, + runs: &[ListTracesRow], +) -> Result, ReadError> { + let mut keys = Vec::with_capacity(runs.len()); + let mut listed = Vec::with_capacity(runs.len()); + for row in runs { + let key = cache_key(store.source(), access, row)?; + listed.push(reader.lists.runs.get(&key).await); + keys.push(key); + } + let misses: Vec<&ListTracesRow> = runs + .iter() + .zip(&listed) + .filter_map(|(row, listed)| listed.is_none().then_some(row)) + .collect(); + let mut resolved = resolve_runs(reader, store, access, &misses) + .await? + .into_iter(); + let mut summaries = Vec::with_capacity(runs.len()); + for ((row, key), cached) in runs.iter().zip(keys).zip(listed) { + let listed = match cached { + Some(listed) => Some(listed), + None => { + let listed = resolved.next().flatten(); + if let Some(listed) = &listed { + reader.lists.runs.insert(key, listed.clone()).await; + } + listed + } + }; + summaries.push(summary(row, listed.as_ref())); + } + Ok(summaries) +} + +/// One entry per run; `None` means the run could not be resolved and keeps its listed summary +/// without being cached. +async fn resolve_runs( + reader: &TraceReader, + store: &S, + access: &ReadAccessParams, + runs: &[&ListTracesRow], +) -> Result>, ReadError> { + let (Some(start_ms), Some(end_ms)) = ( + runs.iter().map(|row| row.start_ms).min(), + runs.iter() + .map(|row| row.start_ms.saturating_add(row.duration_ms)) + .max(), + ) else { + return Ok(Vec::new()); + }; + let params = TracePageSpansParams { + access: access.clone(), + trace_refs: runs.iter().map(|row| row.trace_ref.clone()).collect(), + start_ms, + end_ms: end_ms.saturating_add(1), + }; + let snapshot_ms = now_ms(); + let spans = match store.run_spans(¶ms, snapshot_ms).await { + Ok(spans) => spans, + Err(StoreError::TooLarge) => { + let mut resolved = Vec::with_capacity(runs.len()); + for row in runs { + resolved.push(resolve_run(reader, store, access, row).await?); + } + return Ok(resolved); + } + Err(error) => return Err(map_store_error(error)), + }; + let Some(spend_rows) = spend(store, access, &spans).await else { + // The batch's combined spend read failed; a run's own narrower window may still + // resolve, so fall back per run instead of leaving every run in the batch costless. + let mut resolved = Vec::with_capacity(runs.len()); + for row in runs { + resolved.push(resolve_run(reader, store, access, row).await?); + } + return Ok(resolved); + }; + let mut spans = spans; + spans.sort_by(|left, right| { + run_key(&left.team_id, &left.api_key_hash, &left.trace_id) + .cmp(&run_key( + &right.team_id, + &right.api_key_hash, + &right.trace_id, + )) + .then(left.start_ns.cmp(&right.start_ns)) + }); + let by_run: HashMap<_, &[TraceSpansRow]> = spans + .chunk_by(|left, right| { + (&left.team_id, &left.api_key_hash, &left.trace_id) + == (&right.team_id, &right.api_key_hash, &right.trace_id) + }) + .map(|run| { + ( + run_key(&run[0].team_id, &run[0].api_key_hash, &run[0].trace_id), + run, + ) + }) + .collect(); + Ok(runs + .iter() + .map(|row| { + let spans = by_run + .get(&run_key(&row.team_id, &row.api_key_hash, &row.trace_id)) + .copied() + .unwrap_or_default(); + let spend = + spend_window(spans).map_or(&[][..], |window| spend_within(&spend_rows, window)); + resolve_trace(&row.trace_id, &row.trace_ref, spans, spend).map(|trace| { + ListedRun::Resolved( + Box::new(trace.summary), + Freshness::of(spans, true, snapshot_ms), + ) + }) + }) + .collect()) +} + +async fn resolve_run( + reader: &TraceReader, + store: &S, + access: &ReadAccessParams, + row: &ListTracesRow, +) -> Result, ReadError> { + match reader + .current(store, access, &row.trace_id, &row.trace_ref) + .await + { + Ok(snapshot) => Ok(snapshot.map(|snapshot| { + ListedRun::Resolved( + Box::new(snapshot.trace().summary.clone()), + snapshot.freshness(), + ) + })), + Err(ReadError::TooLarge) => Ok(Some(ListedRun::Limited)), + Err(error) => Err(error), + } +} + +pub(super) fn run_batches(runs: &[T]) -> impl Iterator + '_ { + runs.chunks(RUNS_PER_SPAN_READ) +} diff --git a/litellm-rust/crates/traces-cache/src/reader.rs b/litellm-rust/crates/traces-cache/src/reader.rs new file mode 100644 index 00000000000..5c68d6e58bd --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/reader.rs @@ -0,0 +1,365 @@ +use std::{sync::Arc, time::Duration}; + +use crate::{ + ReadError, Snapshot, SnapshotCache, SnapshotKey, StoreError, TraceStore, + cache::{Freshness, ListCache}, + cursor::{ + ErrorPosition, SpanPosition, decode_cursor, encode_cursor, error_position, trace_position, + }, + list::{list_summaries, run_batches}, + spend::spend, +}; +use litellm_traces::{ + SpanDetail, SpanErrorPage, Trace, TracePage, + query::named::{ + ListTracesParams, ReadAccessParams, SpanDetailParams, SpanErrorParams, TraceIdentityParams, + TraceSpansParams, + }, + resolve_trace, to_ui_content, +}; + +pub const MAX_GRAPH_BYTES: usize = 64 * 1024 * 1024; +pub const MAX_GRAPH_SPANS: usize = 100_000; + +const SNAPSHOT_IDLE: Duration = Duration::from_secs(120); + +/// A read that found no trace, kept apart from failures so single-flight waiters share it +/// without it being cached. +pub(super) enum Miss { + Absent, + Read(ReadError), +} + +impl From for Miss { + fn from(error: crate::Error) -> Self { + Self::Read(error.into()) + } +} + +fn settle(result: Result>>) -> Result, ReadError> { + match result { + Ok(value) => Ok(Some(value)), + Err(miss) => match &*miss { + Miss::Absent => Ok(None), + Miss::Read(error) => Err(error.clone()), + }, + } +} + +pub struct TraceReader { + snapshots: SnapshotCache, + pub(super) lists: ListCache, + response_bytes: usize, +} + +impl TraceReader { + pub fn new(response_bytes: usize) -> Self { + Self { + snapshots: SnapshotCache::new(MAX_GRAPH_BYTES, SNAPSHOT_IDLE), + lists: ListCache::new(), + response_bytes, + } + } + + pub async fn list_traces( + &self, + store: &S, + access: &ReadAccessParams, + start_ms: i64, + end_ms: i64, + cursor: Option<&str>, + limit: u32, + ) -> Result> { + if limit == 0 { + return Err(ReadError::InvalidParameters); + } + let (cursor_ms, cursor_trace_id) = trace_position(cursor)?; + let scope = SnapshotKey::scope(store.source(), access)?; + let accepted = self.lists.limits.get(&scope).await.unwrap_or(u32::MAX); + let mut params = ListTracesParams { + access: access.clone(), + start_ms, + end_ms, + cursor_ms, + cursor_trace_id, + limit: limit.min(500).min(accepted), + }; + let page = loop { + match store.list_runs(¶ms).await { + Err(StoreError::TooLarge) if params.limit > 1 => { + params.limit /= 2; + self.lists.limits.insert(scope.clone(), params.limit).await; + } + Err(StoreError::TooLarge) => return Err(ReadError::TooLarge), + result => break result.map_err(map_store_error)?, + } + }; + let next_cursor = page + .last() + .filter(|_| page.len() == params.limit as usize) + .map(|last| encode_cursor(&(last.start_ms, &last.trace_ref))); + let data = { + let mut summaries = Vec::with_capacity(page.len()); + for batch in run_batches(&page) { + summaries.extend(list_summaries(self, store, access, batch).await?); + } + summaries + }; + Ok(TracePage { data, next_cursor }) + } + + pub async fn get_trace( + &self, + store: &S, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + ) -> Result, ReadError> { + let Some(trace_ref) = reference(store, access, trace_id, trace_ref).await? else { + return Ok(None); + }; + Ok(self + .current(store, access, trace_id, &trace_ref) + .await? + .map(|snapshot| snapshot.trace().clone())) + } + + pub async fn get_trace_page( + &self, + store: &S, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + cursor: Option<&str>, + page_size: u32, + ) -> Result, ReadError> { + if !(1..=500).contains(&page_size) { + return Err(ReadError::InvalidParameters); + } + let Some(trace_ref) = reference(store, access, trace_id, trace_ref).await? else { + return Ok(None); + }; + let Some(cursor) = cursor else { + let Some(snapshot) = self.current(store, access, trace_id, &trace_ref).await? else { + return Ok(None); + }; + let position = SpanPosition { + trace_ref, + snapshot_ms: snapshot.snapshot_ms(), + offset: 0, + version: snapshot.version().to_owned(), + }; + return page(&snapshot, &position, page_size, self.response_bytes).map(Some); + }; + let position: SpanPosition = decode_cursor(cursor, "span")?; + if position.trace_ref != trace_ref || position.snapshot_ms == 0 { + return Err(ReadError::InvalidCursor("span")); + } + let Some(snapshot) = settle( + self.pinned(store, access, trace_id, &trace_ref, position.snapshot_ms) + .await, + )? + else { + return Ok(None); + }; + if position.version != snapshot.version() { + return Err(ReadError::TraceChanged); + } + if position.offset > snapshot.trace().spans.len() { + return Err(ReadError::InvalidCursor("span")); + } + page(&snapshot, &position, page_size, self.response_bytes).map(Some) + } + + pub(super) async fn current( + &self, + store: &S, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + ) -> Result>, ReadError> { + let latest = SnapshotKey::latest(store.source(), access, trace_id, trace_ref)?; + settle( + self.snapshots + .latest_or_load(latest, now_ms(), |snapshot_ms| { + self.pinned(store, access, trace_id, trace_ref, snapshot_ms) + }) + .await, + ) + } + + async fn pinned( + &self, + store: &S, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, + snapshot_ms: u64, + ) -> Result, Arc>> { + let key = SnapshotKey::new(store.source(), access, trace_id, trace_ref, snapshot_ms) + .map_err(|error| Arc::new(error.into()))?; + self.snapshots + .pinned_or_load(key, snapshot_ms, async { + let params = TraceSpansParams { + access: access.clone(), + trace_id: trace_id.to_owned(), + trace_ref: trace_ref.to_owned(), + }; + let rows = store + .trace_spans(¶ms, snapshot_ms) + .await + .map_err(|error| Miss::Read(map_store_error(error)))?; + let spend_rows = spend(store, access, &rows).await; + let freshness = Freshness::of(&rows, spend_rows.is_some(), snapshot_ms); + resolve_trace( + trace_id, + trace_ref, + &rows, + spend_rows.as_deref().unwrap_or_default(), + ) + .map(|trace| (trace, freshness)) + .ok_or(Miss::Absent) + }) + .await + } + + pub async fn get_span( + &self, + store: &S, + access: &ReadAccessParams, + trace_id: &str, + span_id: &str, + trace_ref: &str, + ) -> Result, ReadError> { + let Some(trace_ref) = reference(store, access, trace_id, trace_ref).await? else { + return Ok(None); + }; + let params = SpanDetailParams { + access: access.clone(), + trace_id: trace_id.to_owned(), + trace_ref, + span_id: span_id.to_owned(), + }; + let row = store.span_detail(¶ms).await.map_err(map_store_error)?; + Ok(row.map(|row| SpanDetail { + input_ui: to_ui_content(&row.input), + output_ui: to_ui_content(&row.output), + span_id: row.span_id, + input: row.input, + output: row.output, + attributes: row.attributes, + })) + } + + pub async fn get_span_error( + &self, + store: &S, + access: &ReadAccessParams, + trace_id: &str, + span_id: &str, + trace_ref: &str, + cursor: Option<&str>, + ) -> Result, ReadError> { + let position = error_position(cursor)?; + let Some(trace_ref) = reference(store, access, trace_id, trace_ref).await? else { + return Ok(None); + }; + let offset = position.as_ref().map_or(0, |position| position.offset); + let params = SpanErrorParams { + access: access.clone(), + trace_id: trace_id.to_owned(), + trace_ref, + span_id: span_id.to_owned(), + error_offset: offset, + error_version: position + .map(|position| position.version) + .unwrap_or_default(), + }; + let Some(row) = store.span_error(¶ms).await.map_err(map_store_error)? else { + return Ok(None); + }; + let next_offset = offset + row.message.chars().count() as u64; + let next_cursor = (next_offset < row.total_chars).then(|| { + encode_cursor(&ErrorPosition { + offset: next_offset, + version: row.version, + }) + }); + Ok(Some(SpanErrorPage { + span_id: row.span_id, + message: row.message, + total_chars: row.total_chars, + next_cursor, + })) + } +} + +async fn reference( + store: &S, + access: &ReadAccessParams, + trace_id: &str, + trace_ref: &str, +) -> Result, ReadError> { + if !trace_ref.is_empty() { + return Ok(Some(trace_ref.to_owned())); + } + let params = TraceIdentityParams { + access: access.clone(), + trace_id: trace_id.to_owned(), + }; + let identities = store.trace_refs(¶ms).await.map_err(map_store_error)?; + if identities.len() > 1 { + return Err(ReadError::AmbiguousTrace); + } + Ok(identities.into_iter().next()) +} + +fn page( + snapshot: &Snapshot, + position: &SpanPosition, + page_size: u32, + response_bytes: usize, +) -> Result> { + let spans = &snapshot.trace().spans; + let create_page = |count: usize| { + let end = position.offset.saturating_add(count).min(spans.len()); + Trace { + summary: snapshot.trace().summary.clone(), + agents: snapshot.trace().agents.clone(), + spans: spans[position.offset..end].to_vec(), + next_cursor: (end < spans.len()).then(|| { + encode_cursor(&SpanPosition { + trace_ref: position.trace_ref.clone(), + snapshot_ms: position.snapshot_ms, + offset: end, + version: snapshot.version().to_owned(), + }) + }), + } + }; + let mut trace = create_page(page_size as usize); + loop { + if serde_json::to_vec(&trace) + .map_err(|error| ReadError::Encode(Arc::new(error)))? + .len() + <= response_bytes + { + return Ok(trace); + } + if trace.spans.len() <= 1 { + return Err(ReadError::TooLarge); + } + trace = create_page(trace.spans.len() / 2); + } +} + +pub(super) fn now_ms() -> u64 { + (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64 +} + +pub(super) fn map_store_error(error: StoreError) -> ReadError { + match error { + StoreError::TooLarge => ReadError::TooLarge, + StoreError::Failed(error) => ReadError::Store(Arc::new(error)), + } +} diff --git a/litellm-rust/crates/traces-cache/src/spend.rs b/litellm-rust/crates/traces-cache/src/spend.rs new file mode 100644 index 00000000000..1cfcde841d0 --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/spend.rs @@ -0,0 +1,67 @@ +use std::ops::Range; + +use crate::TraceStore; +use litellm_traces::{ + SpendLookup, + query::named::{ + ReadAccessParams, SpendByResponseIdsParams, SpendByResponseIdsRow, TraceSpansRow, + }, +}; + +const NANOS_PER_MS: i64 = 1_000_000; +const SPEND_WINDOW_MS: i64 = 30 * 60 * 1000; + +pub(super) fn spend_window(rows: &[TraceSpansRow]) -> Option> { + let start_ns = rows.iter().map(|row| row.start_ns).min()?; + let end_ns = rows + .iter() + .map(|row| row.start_ns.saturating_add_unsigned(row.duration_ns)) + .max()?; + Some( + start_ns.div_euclid(NANOS_PER_MS) - SPEND_WINDOW_MS + ..end_ns.div_euclid(NANOS_PER_MS) + SPEND_WINDOW_MS, + ) +} + +pub(super) fn spend_within( + spend: &[SpendByResponseIdsRow], + window: Range, +) -> &[SpendByResponseIdsRow] { + let first = spend.partition_point(|row| row.start_ms < window.start); + let end = spend.partition_point(|row| row.start_ms < window.end); + &spend[first..end.max(first)] +} + +/// Spend rows sorted by `start_ms`, or `None` when the lookup failed and spend is unknown. +pub(super) async fn spend( + store: &S, + access: &ReadAccessParams, + rows: &[TraceSpansRow], +) -> Option> { + let lookup = SpendLookup::new(rows); + let Some(window) = spend_window(rows) else { + return Some(Vec::new()); + }; + if lookup.is_empty() { + return Some(Vec::new()); + } + let params = SpendByResponseIdsParams { + access: access.clone(), + response_ids: lookup.response_ids, + request_ids: lookup.request_ids, + trace_ids: lookup.trace_ids, + start_ms: window.start, + end_ms: window.end, + }; + match store.spend(¶ms).await { + Ok(rows) => { + let mut rows = rows; + rows.sort_by_key(|row| row.start_ms); + Some(rows) + } + Err(error) => { + tracing::warn!(%error, "trace spend lookup unavailable"); + None + } + } +} diff --git a/litellm-rust/crates/traces-cache/src/store.rs b/litellm-rust/crates/traces-cache/src/store.rs new file mode 100644 index 00000000000..3f2a536d4cf --- /dev/null +++ b/litellm-rust/crates/traces-cache/src/store.rs @@ -0,0 +1,62 @@ +use std::future::Future; + +use litellm_traces::query::named::{ + ListTracesParams, ListTracesRow, SpanDetailParams, SpanDetailRow, SpanErrorParams, + SpanErrorRow, SpendByResponseIdsParams, SpendByResponseIdsRow, TraceIdentityParams, + TracePageSpansParams, TraceSpansParams, TraceSpansRow, +}; + +#[derive(Debug, thiserror::Error)] +pub enum StoreError { + #[error("trace read exceeds the storage read budget")] + TooLarge, + #[error(transparent)] + Failed(E), +} + +pub trait TraceStore: Sync { + type Error: std::error::Error + Send + Sync + 'static; + + /// Identifies the backing storage for snapshot cache keys. + fn source(&self) -> &str; + + fn trace_refs( + &self, + params: &TraceIdentityParams, + ) -> impl Future, StoreError>> + Send; + + /// Returns `TooLarge` when the response exceeds the storage limit so the reader can halve `limit`. + fn list_runs( + &self, + params: &ListTracesParams, + ) -> impl Future, StoreError>> + Send; + + /// Returns spans visible at `snapshot_ms`, sorted by `start_ns`, or `TooLarge` past `MAX_GRAPH_BYTES`/`MAX_GRAPH_SPANS`. + fn trace_spans( + &self, + params: &TraceSpansParams, + snapshot_ms: u64, + ) -> impl Future, StoreError>> + Send; + + /// Returns spans visible at `snapshot_ms`, sorted by `start_ns`, or `TooLarge` past `MAX_GRAPH_BYTES`/`MAX_GRAPH_SPANS`. + fn run_spans( + &self, + params: &TracePageSpansParams, + snapshot_ms: u64, + ) -> impl Future, StoreError>> + Send; + + fn spend( + &self, + params: &SpendByResponseIdsParams, + ) -> impl Future, StoreError>> + Send; + + fn span_detail( + &self, + params: &SpanDetailParams, + ) -> impl Future, StoreError>> + Send; + + fn span_error( + &self, + params: &SpanErrorParams, + ) -> impl Future, StoreError>> + Send; +} diff --git a/litellm-rust/crates/traces-cache/tests/read.rs b/litellm-rust/crates/traces-cache/tests/read.rs new file mode 100644 index 00000000000..d714d7c9ed4 --- /dev/null +++ b/litellm-rust/crates/traces-cache/tests/read.rs @@ -0,0 +1,768 @@ +use std::{ + collections::{HashMap, HashSet}, + sync::{ + Mutex, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use litellm_traces::{ + CallEvidenceKind, CallKey, ObservationType, SpanStatus, + query::named::{ + ListTracesParams, ListTracesRow, ReadAccessParams, SpanDetailParams, SpanDetailRow, + SpanErrorParams, SpanErrorRow, SpendByResponseIdsParams, SpendByResponseIdsRow, + TraceIdentityParams, TracePageSpansParams, TraceSpansParams, TraceSpansRow, + }, +}; +use litellm_traces_cache::{LIVE_TTL, ReadError, StoreError, TraceReader, TraceStore}; +use rstest::rstest; + +const START_NS: i64 = 1_790_742_989_000_000_000; + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +enum Operation { + TraceRefs, + ListRuns, + TraceSpans, + RunSpans, + Spend, + SpanDetail, + SpanError, +} + +#[derive(Clone, Copy)] +enum Failure { + TooLarge, + Failed, +} + +#[derive(Debug, thiserror::Error)] +#[error("fake trace store failed")] +struct FakeError; + +#[derive(Default)] +struct State { + failures: HashMap, + trace_refs: Vec, + list_runs: Vec, + trace_spans: HashMap>, + run_spans: Vec, + spend: Vec, + span_detail: Option, + span_error: Option, + list_runs_too_large_above: Option, + trace_too_large_refs: HashSet, + spend_fails_above_response_ids: Option, +} + +#[derive(Default)] +struct Calls { + trace_refs: AtomicUsize, + list_runs: AtomicUsize, + trace_spans: AtomicUsize, + run_spans: AtomicUsize, + spend: AtomicUsize, + span_detail: AtomicUsize, + span_error: AtomicUsize, +} + +#[derive(Default)] +struct FakeStore { + state: Mutex, + calls: Calls, +} + +impl FakeStore { + fn with_spans(trace_ref: &str, spans: Vec) -> Self { + Self { + state: Mutex::new(State { + trace_spans: HashMap::from([(trace_ref.to_owned(), spans)]), + ..State::default() + }), + calls: Calls::default(), + } + } + + fn set_failure(&self, operation: Operation, failure: Failure) { + self.state + .lock() + .unwrap() + .failures + .insert(operation, failure); + } + + fn set_trace_refs(&self, trace_refs: Vec) { + self.state.lock().unwrap().trace_refs = trace_refs; + } + + fn set_list_runs(&self, rows: Vec) { + self.state.lock().unwrap().list_runs = rows; + } + + fn set_list_runs_too_large_above(&self, limit: u32) { + self.state.lock().unwrap().list_runs_too_large_above = Some(limit); + } + + fn set_run_spans(&self, rows: Vec) { + self.state.lock().unwrap().run_spans = rows; + } + + fn set_trace_spans_too_large(&self, trace_ref: &str) { + self.state + .lock() + .unwrap() + .trace_too_large_refs + .insert(trace_ref.to_owned()); + } + + /// Fails `spend` only when the lookup covers more than `limit` response ids, so a batch + /// covering several runs fails while each run's own narrower lookup still succeeds. + fn set_spend_fails_above_response_ids(&self, limit: usize) { + self.state.lock().unwrap().spend_fails_above_response_ids = Some(limit); + } + + fn calls(&self, operation: Operation) -> usize { + match operation { + Operation::TraceRefs => self.calls.trace_refs.load(Ordering::SeqCst), + Operation::ListRuns => self.calls.list_runs.load(Ordering::SeqCst), + Operation::TraceSpans => self.calls.trace_spans.load(Ordering::SeqCst), + Operation::RunSpans => self.calls.run_spans.load(Ordering::SeqCst), + Operation::Spend => self.calls.spend.load(Ordering::SeqCst), + Operation::SpanDetail => self.calls.span_detail.load(Ordering::SeqCst), + Operation::SpanError => self.calls.span_error.load(Ordering::SeqCst), + } + } + + fn failure(state: &State, operation: Operation) -> Result<(), StoreError> { + match state.failures.get(&operation) { + Some(Failure::TooLarge) => Err(StoreError::TooLarge), + Some(Failure::Failed) => Err(StoreError::Failed(FakeError)), + None => Ok(()), + } + } +} + +impl TraceStore for FakeStore { + type Error = FakeError; + + fn source(&self) -> &str { + "fake" + } + + async fn trace_refs( + &self, + _: &TraceIdentityParams, + ) -> Result, StoreError> { + self.calls.trace_refs.fetch_add(1, Ordering::SeqCst); + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::TraceRefs)?; + Ok(state.trace_refs.clone()) + } + + async fn list_runs( + &self, + params: &ListTracesParams, + ) -> Result, StoreError> { + self.calls.list_runs.fetch_add(1, Ordering::SeqCst); + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::ListRuns)?; + if state + .list_runs_too_large_above + .is_some_and(|limit| params.limit > limit) + { + return Err(StoreError::TooLarge); + } + Ok(state + .list_runs + .iter() + .take(params.limit as usize) + .cloned() + .collect()) + } + + async fn trace_spans( + &self, + params: &TraceSpansParams, + _: u64, + ) -> Result, StoreError> { + self.calls.trace_spans.fetch_add(1, Ordering::SeqCst); + tokio::task::yield_now().await; + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::TraceSpans)?; + if state.trace_too_large_refs.contains(¶ms.trace_ref) { + return Err(StoreError::TooLarge); + } + Ok(state + .trace_spans + .get(¶ms.trace_ref) + .cloned() + .unwrap_or_default()) + } + + async fn run_spans( + &self, + _: &TracePageSpansParams, + _: u64, + ) -> Result, StoreError> { + self.calls.run_spans.fetch_add(1, Ordering::SeqCst); + tokio::task::yield_now().await; + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::RunSpans)?; + Ok(state.run_spans.clone()) + } + + async fn spend( + &self, + params: &SpendByResponseIdsParams, + ) -> Result, StoreError> { + self.calls.spend.fetch_add(1, Ordering::SeqCst); + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::Spend)?; + if state + .spend_fails_above_response_ids + .is_some_and(|limit| params.response_ids.len() > limit) + { + return Err(StoreError::Failed(FakeError)); + } + Ok(state.spend.clone()) + } + + async fn span_detail( + &self, + _: &SpanDetailParams, + ) -> Result, StoreError> { + self.calls.span_detail.fetch_add(1, Ordering::SeqCst); + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::SpanDetail)?; + Ok(state.span_detail.clone()) + } + + async fn span_error( + &self, + _: &SpanErrorParams, + ) -> Result, StoreError> { + self.calls.span_error.fetch_add(1, Ordering::SeqCst); + let state = self.state.lock().unwrap(); + Self::failure(&state, Operation::SpanError)?; + Ok(state.span_error.clone()) + } +} + +fn access() -> ReadAccessParams { + ReadAccessParams { + all_teams: true, + user_id: String::new(), + team_ids: Vec::new(), + } +} + +fn span(index: usize) -> TraceSpansRow { + TraceSpansRow { + trace_id: "trace".into(), + span_id: format!("span-{index}"), + parent_span_id: if index == 0 { + String::new() + } else { + "span-0".into() + }, + name: "agent".into(), + kind: ObservationType::Agent, + wrapper_candidate: false, + agent: "agent".into(), + framework: String::new(), + status: SpanStatus::Ok, + status_message: String::new(), + error_truncated: false, + start_ns: START_NS + index as i64 * 1_000_000, + duration_ns: 10_000_000, + service: "test".into(), + input_preview: format!("span input {index}"), + 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: "team".into(), + api_key_hash: "key".into(), + user_id: "user".into(), + } +} + +fn run(trace_id: &str, trace_ref: &str) -> ListTracesRow { + ListTracesRow { + trace_id: trace_id.into(), + trace_ref: trace_ref.into(), + team_id: "team".into(), + api_key_hash: "key".into(), + user_id: "user".into(), + name: "listed".into(), + service: "test".into(), + input_preview: String::new(), + status: SpanStatus::Ok, + start_ms: 1_790_742_989_000, + duration_ms: 10, + span_count: 1, + agent_count: 1, + agent_invocations: 1, + agent_names: vec!["agent".into()], + frameworks: Vec::new(), + llm_calls: 0, + tool_calls: 0, + input_tokens: 0, + output_tokens: 0, + models: Vec::new(), + error_count: 0, + request_ids: Vec::new(), + } +} + +#[rstest] +#[tokio::test] +async fn pages_reuse_one_trace_snapshot_and_concatenate_in_order() { + let store = FakeStore::with_spans("ref", (0..5).map(span).collect()); + let reader = TraceReader::new(usize::MAX); + let access = access(); + let first = reader + .get_trace_page(&store, &access, "trace", "ref", None, 2) + .await + .unwrap() + .unwrap(); + let second = reader + .get_trace_page( + &store, + &access, + "trace", + "ref", + first.next_cursor.as_deref(), + 2, + ) + .await + .unwrap() + .unwrap(); + let third = reader + .get_trace_page( + &store, + &access, + "trace", + "ref", + second.next_cursor.as_deref(), + 2, + ) + .await + .unwrap() + .unwrap(); + let ids: Vec<_> = first + .spans + .iter() + .chain(&second.spans) + .chain(&third.spans) + .map(|span| span.span_id.as_str()) + .collect(); + assert_eq!(ids, ["span-0", "span-1", "span-2", "span-3", "span-4"]); + assert!(third.next_cursor.is_none()); + assert_eq!(store.calls(Operation::TraceSpans), 1); +} + +#[rstest] +#[tokio::test] +async fn snapshot_versions_are_stable_across_readers_and_detect_changes() { + let access = access(); + let original = FakeStore::with_spans("ref", vec![span(0), span(1)]); + let reader_a = TraceReader::new(usize::MAX); + let first = reader_a + .get_trace_page(&original, &access, "trace", "ref", None, 1) + .await + .unwrap() + .unwrap(); + let cursor = first.next_cursor.unwrap(); + + let changed = FakeStore::with_spans("ref", vec![span(0), span(1), span(2)]); + let reader_b = TraceReader::new(usize::MAX); + let result = reader_b + .get_trace_page(&changed, &access, "trace", "ref", Some(&cursor), 1) + .await; + assert!(matches!(result, Err(ReadError::TraceChanged))); + + let unchanged = FakeStore::with_spans("ref", vec![span(0), span(1)]); + let reader_c = TraceReader::new(usize::MAX); + let next = reader_c + .get_trace_page(&unchanged, &access, "trace", "ref", Some(&cursor), 1) + .await + .unwrap() + .unwrap(); + assert_eq!(next.spans[0].span_id, "span-1"); +} + +#[rstest] +#[tokio::test] +async fn response_size_splits_pages_and_rejects_a_single_oversized_span() { + let spans: Vec<_> = (0..4) + .map(|index| { + let mut row = span(index); + row.input_preview = "x".repeat(256); + row + }) + .collect(); + let access = access(); + let full_budget_reader = TraceReader::new(usize::MAX); + let one_span = full_budget_reader + .get_trace_page( + &FakeStore::with_spans("ref", spans.clone()), + &access, + "trace", + "ref", + None, + 1, + ) + .await + .unwrap() + .unwrap(); + let response_bytes = serde_json::to_vec(&one_span).unwrap().len() + 128; + let reader = TraceReader::new(response_bytes); + let store = FakeStore::with_spans("ref", spans.clone()); + let page = reader + .get_trace_page(&store, &access, "trace", "ref", None, 4) + .await + .unwrap() + .unwrap(); + assert!(!page.spans.is_empty()); + assert!(page.spans.len() < 4); + assert!(page.next_cursor.is_some()); + let continued = reader + .get_trace_page( + &store, + &access, + "trace", + "ref", + page.next_cursor.as_deref(), + 4, + ) + .await + .unwrap() + .unwrap(); + assert!(!continued.spans.is_empty()); + assert!(matches!( + TraceReader::new(1) + .get_trace_page(&store, &access, "trace", "ref", None, 1) + .await, + Err(ReadError::TooLarge) + )); +} + +#[rstest] +#[tokio::test] +async fn list_run_budget_halves_the_limit_and_cursor_requires_a_full_page() { + let store = FakeStore::default(); + store.set_list_runs( + (0..3) + .map(|index| run(&format!("trace-{index}"), &format!("ref-{index}"))) + .collect(), + ); + store.set_list_runs_too_large_above(2); + let reader = TraceReader::new(usize::MAX); + let access = access(); + let page = reader + .list_traces(&store, &access, 0, i64::MAX, None, 8) + .await + .unwrap(); + assert_eq!(page.data.len(), 2); + assert!(page.next_cursor.is_some()); + assert_eq!(store.calls(Operation::ListRuns), 3); + + let shorter = FakeStore::default(); + shorter.set_list_runs(vec![run("only", "ref-only")]); + shorter.set_list_runs_too_large_above(2); + let page = reader + .list_traces(&shorter, &access, 0, i64::MAX, None, 8) + .await + .unwrap(); + assert_eq!(page.data.len(), 1); + assert!(page.next_cursor.is_none()); + assert_eq!(shorter.calls(Operation::ListRuns), 1); +} + +#[rstest] +#[tokio::test] +async fn oversized_run_batch_falls_back_to_each_run_and_keeps_listed_summaries() { + let store = FakeStore::with_spans("ref-good", vec![span(0)]); + store.set_list_runs(vec![ + run("trace-large", "ref-large"), + run("trace-good", "ref-good"), + ]); + store.set_trace_spans_too_large("ref-large"); + store.set_run_spans(Vec::new()); + store.set_failure(Operation::RunSpans, Failure::TooLarge); + let reader = TraceReader::new(usize::MAX); + let page = reader + .list_traces(&store, &access(), 0, i64::MAX, None, 2) + .await + .unwrap(); + assert_eq!(page.data.len(), 2); + assert!(page.data[0].resolution_limited); + assert_eq!(page.data[0].trace_ref, "ref-large"); + assert!(!page.data[1].resolution_limited); + assert_eq!(page.data[1].trace_ref, "ref-good"); + assert_eq!(store.calls(Operation::RunSpans), 1); + assert_eq!(store.calls(Operation::TraceSpans), 2); + + let again = reader + .list_traces(&store, &access(), 0, i64::MAX, None, 2) + .await + .unwrap(); + assert_eq!(again.data, page.data); + assert_eq!(store.calls(Operation::RunSpans), 1); + assert_eq!(store.calls(Operation::TraceSpans), 2); +} + +fn spend_row(response_id: &str, cost: f64) -> SpendByResponseIdsRow { + SpendByResponseIdsRow { + request_id: response_id.into(), + litellm_call_id: String::new(), + response_id: response_id.into(), + upstream_response_id: String::new(), + trace_id: String::new(), + span_id: String::new(), + team_id: "team".into(), + api_key: "key".into(), + user: "user".into(), + spend: Some(cost), + start_ms: START_NS / 1_000_000, + } +} + +#[rstest] +#[tokio::test] +async fn failed_batch_spend_lookup_falls_back_to_each_run_instead_of_losing_every_cost() { + let mut first = span(0); + first.trace_id = "trace-a".into(); + first.kind = ObservationType::Llm; + first.litellm_request_id = "response-a".into(); + first.call_keys = vec![CallKey::ProviderResponse("response-a".into())]; + first.call_evidence = Some(CallEvidenceKind::Complete); + let mut second = span(0); + second.trace_id = "trace-b".into(); + second.kind = ObservationType::Llm; + second.litellm_request_id = "response-b".into(); + second.call_keys = vec![CallKey::ProviderResponse("response-b".into())]; + second.call_evidence = Some(CallEvidenceKind::Complete); + + let store = FakeStore::default(); + store.set_list_runs(vec![run("trace-a", "ref-a"), run("trace-b", "ref-b")]); + store.set_run_spans(vec![first.clone(), second.clone()]); + { + let mut state = store.state.lock().unwrap(); + state.trace_spans.insert("ref-a".to_owned(), vec![first]); + state.trace_spans.insert("ref-b".to_owned(), vec![second]); + state.spend = vec![spend_row("response-a", 1.5), spend_row("response-b", 2.5)]; + } + // The batch covers both runs' response ids (2); each run resolved on its own only ever + // asks for its own (1), so this fails only the combined read, not the per-run fallback. + store.set_spend_fails_above_response_ids(1); + + let page = TraceReader::new(usize::MAX) + .list_traces(&store, &access(), 0, i64::MAX, None, 8) + .await + .unwrap(); + + assert_eq!(page.data.len(), 2); + let by_ref: HashMap<&str, f64> = page + .data + .iter() + .map(|run| { + ( + run.trace_ref.as_str(), + run.spend + .expect("run's own spend read should have succeeded"), + ) + }) + .collect(); + assert_eq!(by_ref["ref-a"], 1.5); + assert_eq!(by_ref["ref-b"], 2.5); + assert_eq!(store.calls(Operation::RunSpans), 1); + assert_eq!(store.calls(Operation::TraceSpans), 2); +} + +#[rstest] +#[tokio::test] +async fn failed_spend_lookup_preserves_the_trace_with_unknown_spend() { + let mut row = span(0); + row.litellm_request_id = "response".into(); + row.call_keys = vec![CallKey::ProviderResponse("response".into())]; + row.call_evidence = Some(CallEvidenceKind::Complete); + let store = FakeStore::with_spans("ref", vec![row]); + store.set_failure(Operation::Spend, Failure::Failed); + let trace = TraceReader::new(usize::MAX) + .get_trace(&store, &access(), "trace", "ref") + .await + .unwrap() + .unwrap(); + assert_eq!(trace.summary.spend, None); + assert_eq!(trace.spans[0].spend, None); + assert_eq!(store.calls(Operation::Spend), 1); +} + +#[rstest] +#[tokio::test] +async fn ambiguous_trace_references_fail_and_a_single_reference_is_resolved() { + let reader = TraceReader::new(usize::MAX); + let access = access(); + let ambiguous = FakeStore::default(); + ambiguous.set_trace_refs(vec!["ref-a".into(), "ref-b".into()]); + assert!(matches!( + reader.get_trace(&ambiguous, &access, "trace", "").await, + Err(ReadError::AmbiguousTrace) + )); + + let unique = FakeStore::with_spans("ref-only", vec![span(0)]); + unique.set_trace_refs(vec!["ref-only".into()]); + let trace = reader + .get_trace(&unique, &access, "trace", "") + .await + .unwrap() + .unwrap(); + assert_eq!(trace.summary.trace_ref, "ref-only"); +} + +#[rstest] +#[case::zero(0)] +#[case::above_max(501)] +#[tokio::test] +async fn invalid_page_sizes_are_rejected(#[case] page_size: u32) { + let store = FakeStore::default(); + let reader = TraceReader::new(usize::MAX); + let access = access(); + assert!(matches!( + reader + .get_trace_page(&store, &access, "trace", "ref", None, page_size) + .await, + Err(ReadError::InvalidParameters) + )); +} + +#[rstest] +#[tokio::test] +async fn zero_list_limit_is_rejected() { + let store = FakeStore::default(); + let reader = TraceReader::new(usize::MAX); + let access = access(); + assert!(matches!( + reader + .list_traces(&store, &access, 0, i64::MAX, None, 0) + .await, + Err(ReadError::InvalidParameters) + )); +} + +fn now_ns() -> i64 { + time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64 +} + +#[rstest] +#[tokio::test] +async fn concurrent_and_repeated_opens_share_one_storage_read() { + let store = FakeStore::with_spans("ref", (0..3).map(span).collect()); + let reader = TraceReader::new(usize::MAX); + let access = access(); + let (first, second) = tokio::join!( + reader.get_trace(&store, &access, "trace", "ref"), + reader.get_trace_page(&store, &access, "trace", "ref", None, 2), + ); + let first = first.unwrap().unwrap(); + let second = second.unwrap().unwrap(); + let reopened = reader + .get_trace_page(&store, &access, "trace", "ref", None, 2) + .await + .unwrap() + .unwrap(); + assert_eq!(first.spans.len(), 3); + assert_eq!(second.spans, first.spans[..2]); + assert_eq!(reopened.next_cursor, second.next_cursor); + assert_eq!(store.calls(Operation::TraceSpans), 1); +} + +#[rstest] +#[tokio::test] +async fn failed_reads_are_not_cached() { + let store = FakeStore::with_spans("ref", vec![span(0)]); + store.set_failure(Operation::TraceSpans, Failure::Failed); + let reader = TraceReader::new(usize::MAX); + let access = access(); + assert!(matches!( + reader.get_trace(&store, &access, "trace", "ref").await, + Err(ReadError::Store(_)) + )); + store.state.lock().unwrap().failures.clear(); + let trace = reader + .get_trace(&store, &access, "trace", "ref") + .await + .unwrap() + .unwrap(); + assert_eq!(trace.spans.len(), 1); + assert_eq!(store.calls(Operation::TraceSpans), 2); +} + +#[rstest] +#[tokio::test] +async fn listed_runs_are_read_once_until_a_live_run_expires() { + let live = TraceSpansRow { + trace_id: "trace-live".into(), + start_ns: now_ns(), + ..span(0) + }; + let settled = TraceSpansRow { + trace_id: "trace-settled".into(), + ..span(0) + }; + let store = FakeStore::default(); + store.set_list_runs(vec![ + run("trace-live", "ref-live"), + run("trace-settled", "ref-settled"), + ]); + store.set_run_spans(vec![live, settled]); + let reader = TraceReader::new(usize::MAX); + let access = access(); + let list = || reader.list_traces(&store, &access, 0, i64::MAX, None, 2); + + let first = list().await.unwrap(); + assert!(first.data.iter().all(|summary| summary.name == "agent")); + list().await.unwrap(); + assert_eq!(store.calls(Operation::RunSpans), 1); + + store.set_run_spans(Vec::new()); + tokio::time::sleep(LIVE_TTL + Duration::from_millis(200)).await; + let after = list().await.unwrap(); + assert_eq!(store.calls(Operation::RunSpans), 2); + assert_eq!(after.data[0].name, "listed"); + assert_eq!(after.data[1], first.data[1]); +} + +#[rstest] +#[tokio::test] +async fn concurrent_pages_of_an_evicted_snapshot_share_one_storage_read() { + let access = access(); + let first = TraceReader::new(usize::MAX) + .get_trace_page( + &FakeStore::with_spans("ref", (0..3).map(span).collect()), + &access, + "trace", + "ref", + None, + 1, + ) + .await + .unwrap() + .unwrap(); + let cursor = first.next_cursor.as_deref(); + let store = FakeStore::with_spans("ref", (0..3).map(span).collect()); + let reader = TraceReader::new(usize::MAX); + let (left, right) = tokio::join!( + reader.get_trace_page(&store, &access, "trace", "ref", cursor, 1), + reader.get_trace_page(&store, &access, "trace", "ref", cursor, 1), + ); + assert_eq!(left.unwrap().unwrap().spans[0].span_id, "span-1"); + assert_eq!(right.unwrap().unwrap().spans[0].span_id, "span-1"); + assert_eq!(store.calls(Operation::TraceSpans), 1); +} diff --git a/litellm-rust/crates/traces-cache/tests/snapshots.rs b/litellm-rust/crates/traces-cache/tests/snapshots.rs index de199ae0301..42085d3e95b 100644 --- a/litellm-rust/crates/traces-cache/tests/snapshots.rs +++ b/litellm-rust/crates/traces-cache/tests/snapshots.rs @@ -5,7 +5,9 @@ use litellm_traces::{ query::named::{ReadAccessParams, SpendByResponseIdsRow, TraceSpansRow}, resolve_trace, }; -use litellm_traces_cache::{Error, SnapshotCache, SnapshotKey}; +use std::sync::Arc; + +use litellm_traces_cache::{Error, Freshness, Snapshot, SnapshotCache, SnapshotKey}; use rstest::{fixture, rstest}; const T0: i64 = 1_790_742_989_000_000_000; @@ -60,6 +62,18 @@ fn key( SnapshotKey::new(source, access, trace_id, trace_ref, ms).unwrap() } +async fn insert( + cache: &SnapshotCache, + key: SnapshotKey, + trace: Trace, +) -> Result, Arc> { + cache + .pinned_or_load(key, 100, async { + Ok::<_, Error>((trace, Freshness::Settled)) + }) + .await +} + #[fixture] fn trace() -> Trace { resolve_trace( @@ -85,7 +99,7 @@ async fn cached_trace_is_isolated_by_access_scope( let cache = SnapshotCache::new(1024 * 1024, TTL); let stored = key("source", &access(), "trace", "ref", 100); - cache.insert(stored.clone(), trace.clone()).await.unwrap(); + insert(&cache, stored.clone(), trace.clone()).await.unwrap(); let other_access = ReadAccessParams { all_teams, @@ -115,7 +129,7 @@ async fn cached_trace_is_isolated_by_key_fields( let cache = SnapshotCache::new(1024 * 1024, TTL); let stored = key("source", &access(), "trace", "ref", 100); - cache.insert(stored.clone(), trace.clone()).await.unwrap(); + insert(&cache, stored.clone(), trace.clone()).await.unwrap(); let other = key(source, &access(), trace_id, trace_ref, snapshot_ms); assert!(cache.get(&other).await.is_none()); @@ -129,7 +143,7 @@ async fn snapshot_at_the_size_limit_is_accepted(trace: Trace) { let cache = SnapshotCache::new(size, TTL); let stored = key("source", &access(), "trace", "ref", 100); - cache.insert(stored.clone(), trace).await.unwrap(); + insert(&cache, stored.clone(), trace).await.unwrap(); assert!(cache.get(&stored).await.is_some()); } @@ -141,8 +155,8 @@ async fn snapshot_one_byte_over_the_size_limit_is_rejected(trace: Trace) { let stored = key("source", &access(), "trace", "ref", 100); assert!(matches!( - cache.insert(stored.clone(), trace).await, - Err(Error::ReadTooLarge) + insert(&cache, stored.clone(), trace).await, + Err(error) if matches!(*error, Error::ReadTooLarge) )); assert!(cache.get(&stored).await.is_none()); } @@ -166,25 +180,31 @@ async fn snapshot_version_tracks_the_ordered_span_ids( }; 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(); + let first = insert( + &cache, + key("source", &access(), "a", "ref", 100), + build(first_ids), + ) + .await + .unwrap(); + let second = insert( + &cache, + 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) { +async fn snapshots_expire_when_idle(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(); + insert(&cache, 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 6b95eb149d3..b90ad8a7bf8 100644 --- a/litellm-rust/crates/traces-clickhouse/Cargo.toml +++ b/litellm-rust/crates/traces-clickhouse/Cargo.toml @@ -12,11 +12,9 @@ schema = ["dep:schemars", "litellm-traces/schema"] macro_rules_attribute.workspace = true schemars = { workspace = true, optional = true } askama.workspace = true -base64.workspace = true flate2.workspace = true futures-util.workspace = true hmac = "0.12.1" -itertools = "0.14.0" litellm-http.workspace = true litellm-migrate.workspace = true litellm-storage-clickhouse.workspace = true @@ -30,7 +28,6 @@ strum.workspace = true thiserror.workspace = true time = { workspace = true, features = ["formatting"] } tokio.workspace = true -tracing.workspace = true url.workspace = true [dev-dependencies] diff --git a/litellm-rust/crates/traces-clickhouse/src/error.rs b/litellm-rust/crates/traces-clickhouse/src/error.rs index 78df9a2121c..1d15c556316 100644 --- a/litellm-rust/crates/traces-clickhouse/src/error.rs +++ b/litellm-rust/crates/traces-clickhouse/src/error.rs @@ -16,8 +16,6 @@ pub enum Error { InvalidResponse, #[error("ClickHouse insert exceeds the encoded size limit")] InsertTooLarge, - #[error("Trace exceeds the interactive read budget; use a filtered trace query")] - ReadTooLarge, #[error("ClickHouse schema setup failed with HTTP status {0}")] SchemaFailed(u16), #[error("ClickHouse schema setup transport failed")] @@ -34,12 +32,6 @@ pub enum Error { ProvisionFailed(u16), #[error("ClickHouse reader provisioning transport failed")] ProvisionTransport, - #[error("Invalid {0} cursor")] - InvalidCursor(&'static str), - #[error("Multiple traces have this ID; provide trace_ref")] - AmbiguousTrace, - #[error("Trace changed while paging; refresh the trace to continue")] - TraceChanged, #[error(transparent)] Decode(#[from] litellm_traces::Error), #[error("trace ingestion task failed")] @@ -49,12 +41,3 @@ 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/insert.rs b/litellm-rust/crates/traces-clickhouse/src/insert.rs index 9e6452844f4..d45cc8d53b8 100644 --- a/litellm-rust/crates/traces-clickhouse/src/insert.rs +++ b/litellm-rust/crates/traces-clickhouse/src/insert.rs @@ -239,8 +239,7 @@ mod tests { use rstest::rstest; use serde_json::json; - use super::Error; - use super::{shared_rows, write_rows}; + use super::{Error, shared_rows, write_rows}; #[rstest] fn encoded_limit_counts_utf8_bytes_across_rows() { diff --git a/litellm-rust/crates/traces-clickhouse/src/lib.rs b/litellm-rust/crates/traces-clickhouse/src/lib.rs index d83708d27f1..43ff0b8bf33 100644 --- a/litellm-rust/crates/traces-clickhouse/src/lib.rs +++ b/litellm-rust/crates/traces-clickhouse/src/lib.rs @@ -31,7 +31,7 @@ pub use litellm_storage_clickhouse::{Connection, Parameter}; pub use litellm_traces::{QueryScope, ReadQuery}; pub use query::{QueryHelp, execute_read, query_help, query_sql}; pub use query_access::QueryReaders; -pub use reads::{get_span, get_span_error, get_trace, get_trace_page, list_traces}; +pub use reads::ClickHouseTraces; pub use schema::{ NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, ensure_schema, schema_statements, }; diff --git a/litellm-rust/crates/traces-clickhouse/src/reads.rs b/litellm-rust/crates/traces-clickhouse/src/reads.rs index 6938aab06b3..eaaeb4fef82 100644 --- a/litellm-rust/crates/traces-clickhouse/src/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/src/reads.rs @@ -1,26 +1,13 @@ -//! Scoped trace reads: the trace list, one trace resolved with its spend, and span payloads. - -use std::sync::LazyLock; -use std::time::Duration; - -use base64::{Engine, engine::general_purpose::URL_SAFE}; -use futures_util::{StreamExt, TryStreamExt, stream}; -use itertools::Itertools; use litellm_http::Client; -use litellm_storage_clickhouse::{Query, fetch}; -use litellm_traces::{ - SpanDetail, SpanErrorPage, SpendLookup, Trace, TracePage, listed_summary, - query::named as contracts, resolve_trace, to_ui_content, -}; -use litellm_traces_cache::{SnapshotCache, SnapshotKey}; -use serde::{Deserialize, Serialize}; +use litellm_storage_clickhouse::{Error as StorageError, Query, fetch}; +use litellm_traces::query::named as contracts; +use litellm_traces_cache::{StoreError, TraceStore}; use crate::{ Connection, Error, query::named::{ - ListTracesParams, ListTracesRow, ReadAccessParams, SpanDetail as SpanDetailQuery, - SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIdsParams, TraceIdentity, - TraceIdentityParams, TracePageSpansParams, TraceSpansParams, + ListTracesParams, ListTracesRow, SpanDetail as SpanDetailQuery, SpanError, SpanErrorParams, + SpendByResponseIdsParams, TraceIdentity, TracePageSpansParams, }, }; @@ -36,501 +23,103 @@ 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(|| { - SnapshotCache::new( - crate::span_batches::MAX_GRAPH_BYTES, - Duration::from_secs(120), - ) -}); - -const NANOS_PER_MS: i64 = 1_000_000; -const SPEND_WINDOW_MS: i64 = 30 * 60 * 1000; -const SPEND_CONCURRENCY: usize = 4; - -fn encode_cursor(position: &T) -> String { - URL_SAFE.encode(serde_json::to_vec(position).unwrap_or_default()) +pub struct ClickHouseTraces { + client: Client, + connection: Connection, } -fn decode_cursor Deserialize<'de>>( - cursor: &str, - kind: &'static str, -) -> Result { - URL_SAFE - .decode(cursor) - .ok() - .and_then(|json| serde_json::from_slice(&json).ok()) - .ok_or(Error::InvalidCursor(kind)) -} - -fn trace_position(cursor: Option<&str>) -> Result<(i64, String), Error> { - let Some(cursor) = cursor.filter(|cursor| !cursor.is_empty()) else { - return Ok((0, String::new())); - }; - match decode_cursor::<(i64, String)>(cursor, "trace")? { - (start_ms, trace_ref) if start_ms > 0 && !trace_ref.is_empty() => Ok((start_ms, trace_ref)), - _ => Err(Error::InvalidCursor("trace")), +impl ClickHouseTraces { + pub fn new(client: Client, connection: Connection) -> Self { + Self { client, connection } } } -#[derive(Deserialize, Serialize)] -struct ErrorPosition { - offset: u64, - version: String, -} +impl TraceStore for ClickHouseTraces { + type Error = Error; -fn error_position(cursor: Option<&str>) -> Result, Error> { - let Some(cursor) = cursor else { - return Ok(None); - }; - let position = decode_cursor::(cursor, "diagnostic")?; - let valid_version = position.version.len() == 64 - && position - .version - .bytes() - .all(|byte| byte.is_ascii_digit() || (b'A'..=b'F').contains(&byte)); - if i64::try_from(position.offset).is_err() || !valid_version { - return Err(Error::InvalidCursor("diagnostic")); + fn source(&self) -> &str { + self.connection.url().as_str() } - Ok(Some(position)) -} -/// The stored run a trace id names for this caller; ids can repeat across tenants and runs. -async fn reference( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - trace_id: &str, - trace_ref: &str, -) -> Result, Error> { - if !trace_ref.is_empty() { - return Ok(Some(trace_ref.to_owned())); + async fn trace_refs( + &self, + params: &contracts::TraceIdentityParams, + ) -> Result, StoreError> { + fetch::(&self.client, &self.connection, params) + .await + .map(|rows| rows.into_iter().map(|row| row.trace_ref).collect()) + .map_err(failed) } - let params = TraceIdentityParams { - access: access.clone(), - trace_id: trace_id.to_owned(), - }; - let mut identities = fetch::(client, connection, ¶ms).await?; - if identities.len() > 1 { - return Err(Error::AmbiguousTrace); - } - Ok(identities.pop().map(|identity| identity.trace_ref)) -} -/// Spend records behind the spans' calls. A failed lookup leaves cost unknown instead of failing -/// the read. -async fn spend( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - rows: &[contracts::TraceSpansRow], -) -> Vec { - let lookup = SpendLookup::new(rows); - let (Some(start_ns), Some(end_ns)) = ( - rows.iter().map(|row| row.start_ns).min(), - rows.iter() - .map(|row| row.start_ns.saturating_add_unsigned(row.duration_ns)) - .max(), - ) else { - return Vec::new(); - }; - if lookup.is_empty() { - return Vec::new(); - } - let params = SpendByResponseIdsParams::from(contracts::SpendByResponseIdsParams { - access: access.clone(), - response_ids: lookup.response_ids, - request_ids: lookup.request_ids, - trace_ids: lookup.trace_ids, - start_ms: start_ns.div_euclid(NANOS_PER_MS) - SPEND_WINDOW_MS, - end_ms: end_ns.div_euclid(NANOS_PER_MS) + SPEND_WINDOW_MS, - }); - match crate::span_batches::read_spend(client, connection, params).await { - Ok(rows) => rows, - Err(error) => { - tracing::warn!(%error, "trace spend lookup unavailable"); - Vec::new() + async fn list_runs( + &self, + params: &contracts::ListTracesParams, + ) -> Result, StoreError> { + let storage_params = ListTracesParams::from(params.clone()); + match fetch::(&self.client, &self.connection, &storage_params).await { + Ok(rows) => Ok(rows.into_iter().map(|row| row.0).collect()), + Err(StorageError::ResponseTooLarge) => Err(StoreError::TooLarge), + Err(error) => Err(failed(error)), } } -} -pub async fn list_traces( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - start_ms: i64, - end_ms: i64, - cursor: Option<&str>, - limit: u32, -) -> Result { - if limit == 0 { - return Err(Error::InvalidParameters); + async fn trace_spans( + &self, + params: &contracts::TraceSpansParams, + snapshot_ms: u64, + ) -> Result, StoreError> { + crate::span_batches::read_spans(&self.client, &self.connection, params.clone(), snapshot_ms) + .await } - let (cursor_ms, cursor_trace_id) = trace_position(cursor)?; - let mut params = ListTracesParams::from(contracts::ListTracesParams { - access: access.clone(), - start_ms, - end_ms, - cursor_ms, - cursor_trace_id, - limit: limit.min(500), - }); - let page: Vec = loop { - match fetch::(client, connection, ¶ms).await { - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if params.0.limit > 1 => { - params.0.limit /= 2; - } - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { - return Err(Error::ReadTooLarge); - } - result => break result?.into_iter().map(|row| row.0).collect(), - } - }; - let next_cursor = page - .last() - .filter(|_| page.len() == params.0.limit as usize) - .map(|last| encode_cursor(&(last.start_ms, &last.trace_ref))); - let data = stream::iter(page.chunks(16)) - .then(|batch| list_summaries(client, connection, access, batch)) - .try_collect::>() - .await? - .into_iter() - .flatten() - .collect(); - Ok(TracePage { data, next_cursor }) -} -async fn list_summaries( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - runs: &[contracts::ListTracesRow], -) -> Result, Error> { - let (Some(start_ms), Some(end_ms)) = ( - runs.iter().map(|row| row.start_ms).min(), - runs.iter() - .map(|row| row.start_ms.saturating_add(row.duration_ms)) - .max(), - ) else { - return Ok(Vec::new()); - }; - let params = TracePageSpansParams::from(contracts::TracePageSpansParams { - access: access.clone(), - trace_refs: runs.iter().map(|row| row.trace_ref.clone()).collect(), - start_ms, - end_ms: end_ms.saturating_add(1), - }); - let spans = match crate::span_batches::read_list_spans(client, connection, params).await { - Ok(spans) => spans, - Err(Error::ReadTooLarge) => { - return stream::iter(runs) - .then(|row| async move { - match get_trace(client, connection, access, &row.trace_id, &row.trace_ref).await - { - Ok(trace) => { - Ok(trace.map_or_else(|| listed_summary(row), |trace| trace.summary)) - } - Err(Error::ReadTooLarge) => Ok(listed_summary(row)), - Err(error) => Err(error), - } - }) - .try_collect() - .await; - } - Err(error) => return Err(error), - }; - let by_trace = spans.into_iter().into_group_map_by(|span| { - ( - span.team_id.clone(), - span.api_key_hash.clone(), - span.trace_id.clone(), + async fn run_spans( + &self, + params: &contracts::TracePageSpansParams, + snapshot_ms: u64, + ) -> Result, StoreError> { + crate::span_batches::read_list_spans( + &self.client, + &self.connection, + TracePageSpansParams::from(params.clone()), + snapshot_ms, ) - }); - let summaries = runs - .iter() - .map(|row| { - let spans = by_trace - .get(&( - row.team_id.clone(), - row.api_key_hash.clone(), - row.trace_id.clone(), - )) - .map(Vec::as_slice) - .unwrap_or_default(); - async move { - let spend_rows = spend(client, connection, access, spans).await; - resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows) - .map_or_else(|| listed_summary(row), |trace| trace.summary) - } - }) - .collect::>(); - Ok(stream::iter(summaries) - .buffered(SPEND_CONCURRENCY) - .collect() - .await) -} - -pub async fn get_trace( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - trace_id: &str, - trace_ref: &str, -) -> Result, Error> { - let Some(trace_ref) = reference(client, connection, access, trace_id, trace_ref).await? else { - return Ok(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, u64::MAX).await?; - if rows.is_empty() { - return Ok(None); + .await } - let spend_rows = spend(client, connection, access, &rows).await; - Ok(resolve_trace(trace_id, &trace_ref, &rows, &spend_rows)) -} -#[derive(Deserialize, Serialize)] -struct SpanPosition { - trace_ref: String, - snapshot_ms: u64, - offset: usize, - version: String, -} - -pub async fn get_trace_page( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - trace_id: &str, - trace_ref: &str, - cursor: Option<&str>, - page_size: u32, -) -> Result, Error> { - if !(1..=500).contains(&page_size) { - return Err(Error::InvalidParameters); + async fn spend( + &self, + params: &contracts::SpendByResponseIdsParams, + ) -> Result, StoreError> { + crate::span_batches::read_spend( + &self.client, + &self.connection, + SpendByResponseIdsParams::from(params.clone()), + ) + .await } - let Some(trace_ref) = reference(client, connection, access, trace_id, trace_ref).await? else { - return Ok(None); - }; - let position = match cursor { - Some(cursor) => { - let position: SpanPosition = decode_cursor(cursor, "span")?; - if position.trace_ref != trace_ref || position.snapshot_ms == 0 { - return Err(Error::InvalidCursor("span")); - } - position + + async fn span_detail( + &self, + params: &contracts::SpanDetailParams, + ) -> Result, StoreError> { + match fetch::(&self.client, &self.connection, params).await { + Ok(rows) => Ok(rows.into_iter().next()), + Err(error) => Err(failed(error)), } - None => SpanPosition { - trace_ref: trace_ref.clone(), - snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) - as u64, - offset: 0, - version: String::new(), - }, - }; - let key = SnapshotKey::new( - connection.url().as_str(), - access, - trace_id, - &trace_ref, - position.snapshot_ms, - ) - .map_err(|_| Error::InvalidParameters)?; - 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? + } + + async fn span_error( + &self, + params: &contracts::SpanErrorParams, + ) -> Result, StoreError> { + let storage_params = SpanErrorParams::from(params.clone()); + match fetch::(&self.client, &self.connection, &storage_params).await { + Ok(rows) => Ok(rows.into_iter().next().map(|row| row.0)), + Err(error) => Err(failed(error)), } - }; - let spans = &snapshot.trace().spans; - if cursor.is_some() && position.version != snapshot.version() { - return Err(Error::TraceChanged); - } - let mut trace = Trace { - summary: snapshot.trace().summary.clone(), - agents: snapshot.trace().agents.clone(), - spans: Vec::new(), - next_cursor: None, - }; - if position.offset > spans.len() { - return Err(Error::InvalidCursor("span")); - } - let end = position - .offset - .saturating_add(page_size as usize) - .min(spans.len()); - trace.next_cursor = (end < spans.len()).then(|| { - encode_cursor(&SpanPosition { - offset: end, - version: snapshot.version().to_owned(), - ..position - }) - }); - trace.spans = spans[position.offset..end].to_vec(); - while serde_json::to_vec(&trace) - .map_err(|_| Error::InvalidResponse)? - .len() - > litellm_storage_clickhouse::READ_LIMITS.response_bytes - { - if trace.spans.len() <= 1 { - return Err(Error::ReadTooLarge); - } - trace.spans.truncate(trace.spans.len() / 2); - trace.next_cursor = Some(encode_cursor(&SpanPosition { - trace_ref: trace_ref.clone(), - snapshot_ms: position.snapshot_ms, - offset: position.offset + trace.spans.len(), - version: snapshot.version().to_owned(), - })); - } - Ok(Some(trace)) -} - -pub async fn get_span( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - trace_id: &str, - span_id: &str, - trace_ref: &str, -) -> Result, Error> { - let Some(trace_ref) = reference(client, connection, access, trace_id, trace_ref).await? else { - return Ok(None); - }; - let params = SpanDetailParams { - access: access.clone(), - trace_id: trace_id.to_owned(), - trace_ref, - span_id: span_id.to_owned(), - }; - let row = fetch::(client, connection, ¶ms) - .await? - .into_iter() - .next(); - Ok(row.map(|row| SpanDetail { - input_ui: to_ui_content(&row.input), - output_ui: to_ui_content(&row.output), - span_id: row.span_id, - input: row.input, - output: row.output, - attributes: row.attributes, - })) -} - -pub async fn get_span_error( - client: &Client, - connection: &Connection, - access: &ReadAccessParams, - trace_id: &str, - span_id: &str, - trace_ref: &str, - cursor: Option<&str>, -) -> Result, Error> { - let position = error_position(cursor)?; - let Some(trace_ref) = reference(client, connection, access, trace_id, trace_ref).await? else { - return Ok(None); - }; - let offset = position.as_ref().map_or(0, |position| position.offset); - let params = SpanErrorParams::from(contracts::SpanErrorParams { - access: access.clone(), - trace_id: trace_id.to_owned(), - trace_ref, - span_id: span_id.to_owned(), - error_offset: offset, - error_version: position - .map(|position| position.version) - .unwrap_or_default(), - }); - let Some(row) = fetch::(client, connection, ¶ms) - .await? - .into_iter() - .next() - else { - return Ok(None); - }; - let row = row.0; - let next_offset = offset + row.message.chars().count() as u64; - let next_cursor = (next_offset < row.total_chars).then(|| { - encode_cursor(&ErrorPosition { - offset: next_offset, - version: row.version, - }) - }); - Ok(Some(SpanErrorPage { - span_id: row.span_id, - message: row.message, - total_chars: row.total_chars, - next_cursor, - })) -} - -#[cfg(test)] -mod tests { - use rstest::rstest; - - use super::*; - - #[rstest] - fn trace_cursor_round_trips_the_last_listed_run() { - let cursor = encode_cursor(&(1_790_742_989_377_i64, "4bad42b84e9de3ba46fc870185f8f023")); - assert_eq!( - trace_position(Some(&cursor)).unwrap(), - ( - 1_790_742_989_377, - "4bad42b84e9de3ba46fc870185f8f023".to_owned() - ) - ); - assert_eq!(trace_position(None).unwrap(), (0, String::new())); - assert_eq!(trace_position(Some("")).unwrap(), (0, String::new())); - } - - #[rstest] - #[case::not_base64("abc")] - #[case::not_json("bm90LWpzb24=")] - #[case::numeric_reference("WzEsIDJd")] - #[case::zero_start("WzAsICJ0Il0=")] - fn malformed_trace_cursors_are_rejected(#[case] cursor: &str) { - assert!(matches!( - trace_position(Some(cursor)), - Err(Error::InvalidCursor("trace")) - )); - } - - #[rstest] - #[case::not_base64("garbage")] - #[case::missing_fields("e30=")] - #[case::not_an_object("WzEsMl0=")] - fn malformed_diagnostic_cursors_are_rejected(#[case] cursor: &str) { - assert!(matches!( - error_position(Some(cursor)), - Err(Error::InvalidCursor("diagnostic")) - )); - } - - #[rstest] - #[case::lowercase_version("a".repeat(64))] - #[case::short_version("A".repeat(63))] - fn diagnostic_cursor_requires_a_content_version(#[case] version: String) { - let cursor = encode_cursor(&ErrorPosition { offset: 1, version }); - assert!(matches!( - error_position(Some(&cursor)), - Err(Error::InvalidCursor("diagnostic")) - )); } } + +fn failed(error: StorageError) -> StoreError { + StoreError::Failed(Error::Storage(error)) +} diff --git a/litellm-rust/crates/traces-clickhouse/src/schema.rs b/litellm-rust/crates/traces-clickhouse/src/schema.rs index 07590b59338..6a39bd24041 100644 --- a/litellm-rust/crates/traces-clickhouse/src/schema.rs +++ b/litellm-rust/crates/traces-clickhouse/src/schema.rs @@ -3,8 +3,7 @@ use litellm_migrate::Migration; use serde::Serialize; use std::time::Duration; -use super::Connection; -use super::Error; +use super::{Connection, Error}; const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30); diff --git a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs index d66ad9506da..5e406d67cfa 100644 --- a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs +++ b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs @@ -1,15 +1,20 @@ -use futures_util::{TryStreamExt, stream}; -use itertools::Itertools; +//! Keyset-paged reads that shrink their page when ClickHouse rejects a response as too large and +//! stop accumulating once a graph exceeds the interactive budget. + +use std::{future::Future, marker::PhantomData}; + use litellm_http::Client; use litellm_storage_clickhouse::{Query, fetch}; use litellm_traces::query::named as contracts; -use serde::Serialize; +use litellm_traces_cache::{MAX_GRAPH_BYTES, MAX_GRAPH_SPANS, StoreError}; +use serde::{Serialize, de::DeserializeOwned}; -use crate::{Connection, Error, query::named::TraceSpansRow}; +use crate::{ + Connection, Error, + query::named::{SpendByResponseIdsParams, SpendByResponseIdsRow, TraceSpansRow}, +}; const PAGE_SIZE: u32 = 256; -pub(crate) const MAX_GRAPH_BYTES: usize = 64 * 1024 * 1024; -const MAX_GRAPH_SPANS: usize = 100_000; #[derive(Default)] struct ReadBudget { @@ -18,40 +23,137 @@ struct ReadBudget { } impl ReadBudget { - fn checked_add(&self, bytes: usize, rows: usize) -> Result { - let next = Self { - bytes: self.bytes.saturating_add(bytes), - rows: self.rows.saturating_add(rows), - }; - if next.bytes > MAX_GRAPH_BYTES || next.rows > MAX_GRAPH_SPANS { - return Err(Error::ReadTooLarge); + fn reserve(&mut self, bytes: usize) -> Result<(), StoreError> { + self.bytes = self.bytes.saturating_add(bytes); + if self.bytes > MAX_GRAPH_BYTES || self.rows == MAX_GRAPH_SPANS { + return Err(StoreError::TooLarge); } - Ok(next) - } - - fn record(&mut self, row: &impl Serialize) -> Result<(), Error> { - let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?; - *self = self.checked_add(bytes.len(), 1)?; + self.rows += 1; Ok(()) } + + fn record(&mut self, row: &impl Serialize) -> Result<(), StoreError> { + let bytes = + serde_json::to_vec(row).map_err(|_| StoreError::Failed(Error::InvalidResponse))?; + self.reserve(bytes.len()) + } +} + +/// One keyset position in a paged query: the SQL reads the cursor fields of `Self` plus the +/// `page_size` that [`Batch`] adds. +trait Keyset: Serialize + Sized + Send + Sync { + type Row: Serialize + DeserializeOwned + Send; + const SQL: &'static str; + + fn after(self, last: &Self::Row) -> Self; } #[derive(Serialize)] -struct Parameters { +struct Batch { + #[serde(flatten)] + keyset: K, + page_size: u32, +} + +trait PageSource { + fn page( + &self, + batch: &Batch, + ) -> impl Future, litellm_storage_clickhouse::Error>> + Send; +} + +struct Paged(PhantomData); + +impl Query for Paged { + type Params = Batch; + type Row = K::Row; + const SQL: &'static str = K::SQL; +} + +/// Reads every row after `keyset`. A page ClickHouse rejects as too large is retried at half the +/// size, and the smaller page is kept for the rest of the read because row sizes within one graph +/// rarely shrink again. Halving a one-row page means a single row exceeds the response limit. +async fn read_all>( + source: &S, + keyset: K, +) -> Result, StoreError> { + let mut batch = Batch { + keyset, + page_size: PAGE_SIZE, + }; + let mut rows = Vec::new(); + let mut budget = ReadBudget::default(); + loop { + let page = match source.page(&batch).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if batch.page_size > 1 => { + batch.page_size /= 2; + continue; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(StoreError::TooLarge); + } + result => result.map_err(|error| StoreError::Failed(Error::Storage(error)))?, + }; + let complete = page.len() < batch.page_size as usize; + for row in &page { + budget.record(row)?; + } + if let Some(last) = page.last() { + batch.keyset = batch.keyset.after(last); + } + rows.extend(page); + if complete { + return Ok(rows); + } + } +} + +struct ClickHouse<'a> { + client: &'a Client, + connection: &'a Connection, +} + +impl PageSource for ClickHouse<'_> { + fn page( + &self, + batch: &Batch, + ) -> impl Future, litellm_storage_clickhouse::Error>> + Send { + fetch::>(self.client, self.connection, batch) + } +} + +async fn read_paged( + client: &Client, + connection: &Connection, + keyset: K, +) -> Result, StoreError> { + let source = ClickHouse { client, connection }; + read_all(&source, keyset).await +} + +fn by_start(mut rows: Vec) -> Vec { + rows.sort_by_key(|row| row.start_ns); + rows +} + +#[derive(Serialize)] +struct SpanKeyset { #[serde(flatten)] trace: contracts::TraceSpansParams, after_span_id: String, - page_size: u32, snapshot_ms: u64, } -struct SpanBatch; - -impl Query for SpanBatch { - type Params = Parameters; +impl Keyset for SpanKeyset { type Row = TraceSpansRow; - const SQL: &'static str = include_str!("../query/trace_span_batch.sql"); + + fn after(self, last: &TraceSpansRow) -> Self { + Self { + after_span_id: last.0.span_id.clone(), + ..self + } + } } pub(crate) async fn read_spans( @@ -59,199 +161,109 @@ pub(crate) async fn read_spans( connection: &Connection, trace: contracts::TraceSpansParams, snapshot_ms: u64, -) -> Result, Error> { - let mut parameters = Parameters { +) -> Result, StoreError> { + let keyset = SpanKeyset { trace, after_span_id: String::new(), - page_size: PAGE_SIZE, snapshot_ms, }; - let mut spans = Vec::new(); - let mut budget = ReadBudget::default(); - loop { - let page = match fetch::(client, connection, ¶meters).await { - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) - if parameters.page_size > 1 => - { - parameters.page_size /= 2; - continue; - } - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { - return Err(Error::ReadTooLarge); - } - result => result?, - }; - let complete = page.len() < parameters.page_size as usize; - if let Some(last) = page.last() { - parameters.after_span_id.clone_from(&last.0.span_id); - } - for row in page { - budget.record(&row)?; - spans.push(row.0); - } - if complete { - spans.sort_by_key(|row| row.start_ns); - return Ok(spans); - } - parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE); - } + let rows = read_paged(client, connection, keyset).await?; + Ok(by_start(rows.into_iter().map(|row| row.0).collect())) } #[derive(Serialize)] -struct ListParameters { +struct ListSpanKeyset { #[serde(flatten)] runs: crate::query::named::TracePageSpansParams, after_team: String, after_key: String, after_trace: String, after_span: String, - page_size: u32, snapshot_ms: u64, } -struct ListSpanBatch; - -impl Query for ListSpanBatch { - type Params = ListParameters; +impl Keyset for ListSpanKeyset { type Row = TraceSpansRow; - const SQL: &'static str = include_str!("../query/trace_list_span_batch.sql"); + + fn after(self, last: &TraceSpansRow) -> Self { + Self { + after_team: last.0.team_id.clone(), + after_key: last.0.api_key_hash.clone(), + after_trace: last.0.trace_id.clone(), + after_span: last.0.span_id.clone(), + ..self + } + } } pub(crate) async fn read_list_spans( client: &Client, connection: &Connection, runs: crate::query::named::TracePageSpansParams, -) -> Result, Error> { - let parameters = ListParameters { + snapshot_ms: u64, +) -> Result, StoreError> { + let keyset = ListSpanKeyset { runs, after_team: String::new(), after_key: String::new(), after_trace: String::new(), after_span: String::new(), - page_size: PAGE_SIZE, - snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64, + snapshot_ms, }; - let pages = stream::try_unfold( - (Some(parameters), ReadBudget::default()), - |(parameters, budget)| async move { - let Some(parameters) = parameters else { - return Ok(None); - }; - let page = match fetch::(client, connection, ¶meters).await { - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) - if parameters.page_size > 1 => - { - let retry = ListParameters { - page_size: parameters.page_size / 2, - ..parameters - }; - return Ok(Some((Vec::new(), (Some(retry), budget)))); - } - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { - return Err(Error::ReadTooLarge); - } - result => result?, - }; - let next = page - .last() - .filter(|_| page.len() == parameters.page_size as usize) - .map(|last| ListParameters { - after_team: last.0.team_id.clone(), - after_key: last.0.api_key_hash.clone(), - after_trace: last.0.trace_id.clone(), - after_span: last.0.span_id.clone(), - page_size: (parameters.page_size * 2).min(PAGE_SIZE), - ..parameters - }); - let next_budget = page.iter().try_fold(budget, |budget, row| { - let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?; - budget.checked_add(bytes.len(), 1) - })?; - Ok(Some((page, (next, next_budget)))) - }, - ) - .try_collect::>() - .await?; - Ok(pages - .into_iter() - .flatten() - .map(|row| row.0) - .sorted_by_key(|row| row.start_ns) - .collect()) + let rows = read_paged(client, connection, keyset).await?; + Ok(by_start(rows.into_iter().map(|row| row.0).collect())) } #[derive(Serialize)] -struct SpendParameters { +struct SpendKeyset { #[serde(flatten)] - lookup: crate::query::named::SpendByResponseIdsParams, + lookup: SpendByResponseIdsParams, has_cursor: u8, after_team: String, after_ms: i64, after_id: String, - page_size: u32, } -struct SpendBatch; - -impl Query for SpendBatch { - type Params = SpendParameters; - type Row = crate::query::named::SpendByResponseIdsRow; - +impl Keyset for SpendKeyset { + type Row = SpendByResponseIdsRow; const SQL: &'static str = include_str!("../query/spend_batch.sql"); + + fn after(self, last: &SpendByResponseIdsRow) -> Self { + Self { + has_cursor: 1, + after_team: last.0.team_id.clone(), + after_ms: last.0.start_ms, + after_id: last.0.request_id.clone(), + ..self + } + } } pub(crate) async fn read_spend( client: &Client, connection: &Connection, - lookup: crate::query::named::SpendByResponseIdsParams, -) -> Result, Error> { - let mut parameters = SpendParameters { + lookup: SpendByResponseIdsParams, +) -> Result, StoreError> { + let keyset = SpendKeyset { lookup, has_cursor: 0, after_team: String::new(), after_ms: 0, after_id: String::new(), - page_size: PAGE_SIZE, }; - let mut rows = Vec::new(); - let mut budget = ReadBudget::default(); - loop { - let page = match fetch::(client, connection, ¶meters).await { - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) - if parameters.page_size > 1 => - { - parameters.page_size /= 2; - continue; - } - Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { - return Err(Error::ReadTooLarge); - } - result => result?, - }; - let complete = page.len() < parameters.page_size as usize; - if let Some(last) = page.last() { - parameters.has_cursor = 1; - parameters.after_team.clone_from(&last.0.team_id); - parameters.after_ms = last.0.start_ms; - parameters.after_id.clone_from(&last.0.request_id); - } - for row in page { - budget.record(&row)?; - rows.push(row.0); - } - if complete { - return Ok(rows); - } - parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE); - } + let rows = read_paged(client, connection, keyset).await?; + Ok(rows.into_iter().map(|row| row.0).collect()) } #[cfg(test)] mod tests { - use super::*; + use std::sync::Mutex; + use rstest::rstest; + use super::*; + #[rstest] #[case::byte_boundary(MAX_GRAPH_BYTES - 1, 0, 1, false)] #[case::byte_overflow(MAX_GRAPH_BYTES - 1, 0, 2, true)] @@ -264,7 +276,77 @@ mod tests { #[case] next: usize, #[case] rejected: bool, ) { - let budget = ReadBudget { bytes, rows }; - assert_eq!(budget.checked_add(next, 1).is_err(), rejected); + let mut budget = ReadBudget { bytes, rows }; + assert_eq!(budget.reserve(next).is_err(), rejected); + } + + #[derive(Serialize)] + struct Numbers { + after: u32, + } + + impl Keyset for Numbers { + type Row = u32; + const SQL: &'static str = ""; + + fn after(self, last: &u32) -> Self { + Self { after: *last } + } + } + + /// A table of `total` rows whose transport rejects any page larger than `largest_page`. + struct Table { + total: u32, + largest_page: u32, + requests: Mutex>, + } + + impl PageSource for Table { + async fn page( + &self, + batch: &Batch, + ) -> Result, litellm_storage_clickhouse::Error> { + self.requests.lock().unwrap().push(batch.page_size); + if batch.page_size > self.largest_page { + return Err(litellm_storage_clickhouse::Error::ResponseTooLarge); + } + let end = (batch.keyset.after + batch.page_size).min(self.total); + Ok((batch.keyset.after + 1..=end).collect()) + } + } + + #[rstest] + #[case::fits(1000, PAGE_SIZE, &[256, 256, 256, 256])] + #[case::uniform_large_rows(1000, 100, &[256, 128, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64, 64])] + #[tokio::test] + async fn a_rejected_page_size_is_not_retried( + #[case] total: u32, + #[case] largest_page: u32, + #[case] requests: &[u32], + ) { + let table = Table { + total, + largest_page, + requests: Mutex::new(Vec::new()), + }; + let rows = read_all(&table, Numbers { after: 0 }).await.unwrap(); + assert_eq!(rows, (1..=total).collect::>()); + assert_eq!(table.requests.lock().unwrap().as_slice(), requests); + } + + #[rstest] + #[tokio::test] + async fn a_single_oversized_row_fails_the_read() { + let table = Table { + total: 10, + largest_page: 0, + requests: Mutex::new(Vec::new()), + }; + let result = read_all(&table, Numbers { after: 0 }).await; + assert!(matches!(result, Err(StoreError::TooLarge)), "{result:?}"); + assert_eq!( + table.requests.lock().unwrap().as_slice(), + &[256, 128, 64, 32, 16, 8, 4, 2, 1] + ); } } diff --git a/litellm-rust/crates/traces-clickhouse/src/sql.rs b/litellm-rust/crates/traces-clickhouse/src/sql.rs index 80b3ec88534..dfa0618773a 100644 --- a/litellm-rust/crates/traces-clickhouse/src/sql.rs +++ b/litellm-rust/crates/traces-clickhouse/src/sql.rs @@ -3,8 +3,10 @@ use std::collections::BTreeMap; use litellm_http::Client; use litellm_traces::ReadQuery; -use super::query::{lens::*, named::*}; -use super::{Connection, Error, Parameter}; +use super::{ + Connection, Error, Parameter, + query::{lens::*, named::*}, +}; use litellm_storage_clickhouse::{Query, fetch_json}; pub async fn execute_named_read( diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index cd23484566f..09fd3e78da8 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -1887,10 +1887,10 @@ async fn named_and_sql_readers_share_request_log_visibility( #[case] legacy_key: Option<&str>, #[case] expected: Vec<&str>, ) -> TestResult { - use litellm_traces_clickhouse::query::named::{ - ReadAccessParams, SpendByResponseIds, SpendByResponseIdsParams, + use litellm_traces_clickhouse::{ + QueryReaders, QueryScope, + query::named::{ReadAccessParams, SpendByResponseIds, SpendByResponseIdsParams}, }; - use litellm_traces_clickhouse::{QueryReaders, QueryScope}; let database = database?; let writer = Connection::writer(&database.url)?; diff --git a/litellm-rust/crates/traces-clickhouse/tests/reads.rs b/litellm-rust/crates/traces-clickhouse/tests/reads.rs index 2abb308f9c1..e1996d6af04 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/reads.rs @@ -1,8 +1,10 @@ use std::collections::BTreeMap; +use litellm_http::Client; use litellm_traces::query::named::ReadAccessParams; +use litellm_traces_cache::{ReadError, TraceReader}; use litellm_traces_clickhouse::{ - Connection, InsertTable, QueryScope, get_trace, get_trace_page, insert_rows, list_traces, + ClickHouseTraces, Connection, InsertTable, QueryScope, insert_rows, }; use rstest::rstest; use serde_json::json; @@ -14,6 +16,13 @@ mod support; use fixtures::{DATABASE, SeededDatabase, migrated_database, seeded_database}; use support::TestResult; +fn make_reader(client: &Client, connection: Connection) -> (TraceReader, ClickHouseTraces) { + ( + TraceReader::new(litellm_storage_clickhouse::READ_LIMITS.response_bytes), + ClickHouseTraces::new(client.clone(), connection), + ) +} + #[rstest] #[case::api_key("key-a", "")] #[case::user("", "user-a")] @@ -74,16 +83,19 @@ async fn list_costs_match_each_run_when_response_ids_are_reused( .collect(), ) .await?; - let reader = fixture + let connection = fixture .readers .connection(client, &QueryScope::All, "fixture-secret") .await?; + let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: false, user_id: user_id.into(), team_ids: vec!["team-a".into()], }; - let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + let page = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 50) + .await?; assert_eq!(page.data.len(), runs.len()); for (trace_id, _, cost) in runs { let summary = page @@ -91,7 +103,8 @@ async fn list_costs_match_each_run_when_response_ids_are_reused( .iter() .find(|summary| summary.trace_id == trace_id) .ok_or("missing run")?; - let detail = get_trace(client, &reader, &access, trace_id, &summary.trace_ref) + let detail = reader + .get_trace(&store, &access, trace_id, &summary.trace_ref) .await? .ok_or("missing trace")?; assert_eq!(detail.summary.spend, Some(cost)); @@ -199,16 +212,19 @@ async fn large_runs_remain_complete_under_default_reader_limits( .collect::>(); insert_rows(client, &writer, DATABASE, InsertTable::SpendLogs, costs).await?; } - let reader = fixture + let connection = fixture .readers .connection(client, &QueryScope::All, "fixture-secret") .await?; + let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: false, user_id: String::new(), team_ids: vec!["team-a".into()], }; - let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 500).await?; + let page = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 500) + .await?; assert_eq!(page.data.len(), runs); assert!( page.data @@ -222,42 +238,17 @@ async fn large_runs_remain_complete_under_default_reader_limits( .send() .await? .error_for_status()?; - let read_queries = client.post(writer.url().clone()).body(format!( - "SELECT count() FROM system.query_log WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' AND query LIKE '%FROM otel_traces AS o%' AND query NOT LIKE '%system.query_log%'" - )).send().await?.error_for_status()?.text().await?; - let read_queries = read_queries.trim().parse::()?; - assert!( - read_queries > 0 && read_queries < runs, - "{read_queries} span queries for {runs} runs" - ); - if costed { - let overlapping = client - .post(writer.url().clone()) - .body(format!( - "WITH spend_reads AS ( - SELECT query_start_time_microseconds AS started, event_time_microseconds AS finished - FROM system.query_log - WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' - AND query LIKE '%FROM spend_logs FINAL%' AND query NOT LIKE '%system.query_log%' - ), events AS ( - SELECT started AS at, 1 AS delta FROM spend_reads - UNION ALL SELECT finished AS at, -1 AS delta FROM spend_reads - ) - SELECT max(active) FROM ( - SELECT sum(delta) OVER (ORDER BY at, delta ROWS UNBOUNDED PRECEDING) AS active - FROM events - )" - )) - .send() - .await? - .error_for_status()? - .text() - .await? - .trim() - .parse::()?; + for table in ["otel_traces AS o", "spend_logs FINAL"] + .into_iter() + .take(if costed { 2 } else { 1 }) + { + let read_queries = client.post(writer.url().clone()).body(format!( + "SELECT count() FROM system.query_log WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' AND query LIKE '%FROM {table}%' AND query NOT LIKE '%system.query_log%'" + )).send().await?.error_for_status()?.text().await?; + let read_queries = read_queries.trim().parse::()?; assert!( - (2..=4).contains(&overlapping), - "{overlapping} simultaneous spend reads for {runs} runs" + read_queries > 0 && read_queries < runs, + "{read_queries} {table} queries for {runs} runs" ); } } @@ -281,7 +272,8 @@ async fn large_runs_remain_complete_under_default_reader_limits( .find(|run| run.trace_id == "trace-0000") .ok_or("missing run")? .trace_ref; - let detail = get_trace(client, &reader, &access, "trace-0000", trace_ref) + let detail = reader + .get_trace(&store, &access, "trace-0000", trace_ref) .await? .ok_or("missing trace")?; assert_eq!(detail.spans.len(), steps); @@ -303,24 +295,25 @@ async fn large_runs_remain_complete_under_default_reader_limits( ..access.clone() }; assert!( - get_trace(client, &reader, &denied, "trace-0000", trace_ref) + reader + .get_trace(&store, &denied, "trace-0000", trace_ref) .await? .is_none() ); let mut cursor = None; let mut ids = Vec::new(); loop { - let page = get_trace_page( - client, - &reader, - &access, - "trace-0000", - trace_ref, - cursor.as_deref(), - 200, - ) - .await? - .ok_or("missing page")?; + let page = reader + .get_trace_page( + &store, + &access, + "trace-0000", + trace_ref, + cursor.as_deref(), + 200, + ) + .await? + .ok_or("missing page")?; assert_eq!(page.summary, detail.summary); assert!(page.spans.len() <= 200); assert!( @@ -329,17 +322,17 @@ async fn large_runs_remain_complete_under_default_reader_limits( ); if ids.is_empty() { assert!( - get_trace_page( - client, - &reader, - &denied, - "trace-0000", - trace_ref, - page.next_cursor.as_deref(), - 200, - ) - .await? - .is_none() + reader + .get_trace_page( + &store, + &denied, + "trace-0000", + trace_ref, + page.next_cursor.as_deref(), + 200, + ) + .await? + .is_none() ); client .post(writer.url().clone()) @@ -372,45 +365,43 @@ async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( ) -> TestResult { let fixture = seeded_database?; let client = &fixture.database.client; - let reader = fixture + let connection = fixture .readers .connection(client, &QueryScope::All, "fixture-secret") .await?; + let (reader, store) = make_reader(client, connection.clone()); let access = ReadAccessParams { all_teams: true, user_id: String::new(), team_ids: Vec::new(), }; - let listed = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 10).await?; + let listed = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 10) + .await?; let summary = listed .data .iter() .find(|summary| summary.span_count == 3) .ok_or("missing fixture")?; - let first = get_trace_page( - client, - &reader, - &access, - &summary.trace_id, - &summary.trace_ref, - None, - 1, - ) - .await? - .ok_or("missing first page")?; - let original_ids = get_trace( - client, - &reader, - &access, - &summary.trace_id, - &summary.trace_ref, - ) - .await? - .ok_or("missing trace")? - .spans - .into_iter() - .map(|span| span.span_id) - .collect::>(); + let first = reader + .get_trace_page( + &store, + &access, + &summary.trace_id, + &summary.trace_ref, + None, + 1, + ) + .await? + .ok_or("missing first page")?; + let original_ids = reader + .get_trace(&store, &access, &summary.trace_id, &summary.trace_ref) + .await? + .ok_or("missing trace")? + .spans + .into_iter() + .map(|span| span.span_id) + .collect::>(); let writer = Connection::writer(&fixture.database.url)?; insert_rows( client, @@ -434,17 +425,17 @@ async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( team_ids: vec!["not-this-team".into()], }; assert!( - get_trace_page( - client, - &reader, - &denied, - &summary.trace_id, - &summary.trace_ref, - first.next_cursor.as_deref(), - 1 - ) - .await? - .is_none() + reader + .get_trace_page( + &store, + &denied, + &summary.trace_id, + &summary.trace_ref, + first.next_cursor.as_deref(), + 1 + ) + .await? + .is_none() ); let first_cursor = first.next_cursor.clone(); let mut cursor = first.next_cursor; @@ -454,44 +445,45 @@ async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( .map(|span| span.span_id) .collect::>(); while let Some(current) = cursor { - let next = get_trace_page( - client, - &reader, - &access, - &summary.trace_id, - &summary.trace_ref, - Some(¤t), - 1, - ) - .await? - .ok_or("missing next page")?; + let next = reader + .get_trace_page( + &store, + &access, + &summary.trace_id, + &summary.trace_ref, + Some(¤t), + 1, + ) + .await? + .ok_or("missing next page")?; assert_eq!(next.summary.span_count, 3); ids.extend(next.spans.into_iter().map(|span| span.span_id)); cursor = next.next_cursor; } assert_eq!(ids, original_ids); - let refreshed = get_trace( - client, - &reader, - &access, - &summary.trace_id, - &summary.trace_ref, - ) - .await? - .ok_or("missing refreshed trace")?; + let cached = reader + .get_trace(&store, &access, &summary.trace_id, &summary.trace_ref) + .await? + .ok_or("missing cached trace")?; + assert_eq!(cached.spans.len(), 3); + let (fresh_reader, fresh_store) = make_reader(client, connection); + let refreshed = fresh_reader + .get_trace(&fresh_store, &access, &summary.trace_id, &summary.trace_ref) + .await? + .ok_or("missing refreshed trace")?; assert_eq!(refreshed.spans.len(), 4); assert!(matches!( - get_trace_page( - client, - &reader, - &access, - &summary.trace_id, - &summary.trace_ref, - Some("invalid"), - 1 - ) - .await, - Err(litellm_traces_clickhouse::Error::InvalidCursor("span")) + reader + .get_trace_page( + &store, + &access, + &summary.trace_id, + &summary.trace_ref, + Some("invalid"), + 1 + ) + .await, + Err(ReadError::InvalidCursor("span")) )); let backdated = json!({ "Timestamp": "2026-09-01 00:00:00.000000000", @@ -509,20 +501,21 @@ async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( .send() .await? .error_for_status()?; - let uncached_reader = + let uncached_connection = Connection::reader(&format!("{}?max_threads=1", fixture.database.url), DATABASE)?; - let changed = get_trace_page( - client, - &uncached_reader, - &access, - &summary.trace_id, - &summary.trace_ref, - first_cursor.as_deref(), - 1, - ) - .await; + let (uncached_reader, uncached_store) = make_reader(client, uncached_connection); + let changed = uncached_reader + .get_trace_page( + &uncached_store, + &access, + &summary.trace_id, + &summary.trace_ref, + first_cursor.as_deref(), + 1, + ) + .await; assert!( - matches!(changed, Err(litellm_traces_clickhouse::Error::TraceChanged)), + matches!(changed, Err(ReadError::TraceChanged)), "{changed:?}" ); Ok(()) @@ -535,16 +528,19 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( ) -> TestResult { let fixture = seeded_database?; let client = &fixture.database.client; - let reader = fixture + let connection = fixture .readers .connection(client, &QueryScope::All, "fixture-secret") .await?; + let (reader, store) = make_reader(client, connection.clone()); let access = ReadAccessParams { all_teams: true, user_id: String::new(), team_ids: Vec::new(), }; - let before = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + let before = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 50) + .await?; let run = before .data .iter() @@ -571,7 +567,14 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( ])], ) .await?; - let after = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + let cached = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 50) + .await?; + assert_eq!(cached.data, before.data); + let (reader, store) = make_reader(client, connection); + let after = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 50) + .await?; assert_eq!(after.data.len(), before.data.len()); let limited = after .data @@ -588,17 +591,10 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals( .all(|item| !item.resolution_limited) ); assert!(matches!( - get_trace_page( - client, - &reader, - &access, - &run.trace_id, - &run.trace_ref, - None, - 200 - ) - .await, - Err(litellm_traces_clickhouse::Error::ReadTooLarge) + reader + .get_trace_page(&store, &access, &run.trace_id, &run.trace_ref, None, 200) + .await, + Err(ReadError::TooLarge) )); Ok(()) } @@ -685,16 +681,19 @@ async fn gateway_ids_resolve_through_detail_and_batch_reads_with_legacy_fallback .collect(), ) .await?; - let reader = fixture + let connection = fixture .readers .connection(client, &QueryScope::All, "fixture-secret") .await?; + let (reader, store) = make_reader(client, connection); let access = ReadAccessParams { all_teams: false, user_id: String::new(), team_ids: vec!["team-a".into()], }; - let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?; + let page = reader + .list_traces(&store, &access, 0, 2_000_000_000_000, None, 50) + .await?; assert_eq!(page.data.len(), cases.len()); for (id, _, _, _, _, expected) in cases { let summary = page @@ -702,7 +701,8 @@ async fn gateway_ids_resolve_through_detail_and_batch_reads_with_legacy_fallback .iter() .find(|summary| summary.trace_id == id) .ok_or("missing run")?; - let detail = get_trace(client, &reader, &access, id, &summary.trace_ref) + let detail = reader + .get_trace(&store, &access, id, &summary.trace_ref) .await? .ok_or("missing trace")?; assert_eq!(detail.summary.spend, expected, "{id}"); diff --git a/litellm-rust/crates/traces/src/query/named.rs b/litellm-rust/crates/traces/src/query/named.rs index 5765efaa62e..a981bc46471 100644 --- a/litellm-rust/crates/traces/src/query/named.rs +++ b/litellm-rust/crates/traces/src/query/named.rs @@ -15,7 +15,7 @@ pub struct ReadAccessParams { pub team_ids: Vec, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct ListTracesParams { #[serde(flatten)] pub access: ReadAccessParams, @@ -26,7 +26,7 @@ pub struct ListTracesParams { pub limit: u32, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct ListTracesRow { pub trace_id: String, pub trace_ref: String, @@ -56,7 +56,7 @@ pub struct ListTracesRow { pub request_ids: Vec, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct TraceSpansParams { #[serde(flatten)] pub access: ReadAccessParams, @@ -113,7 +113,7 @@ pub struct TraceSpansRow { pub user_id: String, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct TracePageSpansParams { #[serde(flatten)] pub access: ReadAccessParams, @@ -131,7 +131,7 @@ pub struct SpanDetailParams { pub span_id: String, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct SpanDetailRow { pub span_id: String, pub input: String, @@ -139,7 +139,7 @@ pub struct SpanDetailRow { pub attributes: BTreeMap, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct SpanErrorParams { #[serde(flatten)] pub access: ReadAccessParams, @@ -150,7 +150,7 @@ pub struct SpanErrorParams { pub error_version: String, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct SpanErrorRow { pub span_id: String, pub message: String, @@ -158,7 +158,7 @@ pub struct SpanErrorRow { pub version: String, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct SpendByResponseIdsParams { #[serde(flatten)] pub access: ReadAccessParams, @@ -169,7 +169,7 @@ pub struct SpendByResponseIdsParams { pub end_ms: i64, } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Clone, Debug, Deserialize, Serialize)] pub struct SpendByResponseIdsRow { pub request_id: String, pub litellm_call_id: String, diff --git a/litellm/constants.py b/litellm/constants.py index d58fc8a6318..69ff3cf7a5e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -55,6 +55,7 @@ DEFAULT_AGENT_TRACING_RETENTION_DAYS: Final = 14 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) +TRACE_READ_RETRY_AFTER_SECONDS: Final = get_env_int("TRACE_READ_RETRY_AFTER_SECONDS", 2) 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) diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index 96df9728a73..ca8e507bff5 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -1,5 +1,6 @@ """``CustomLogger`` adapter on the OpenTelemetry span engine.""" +import sys from collections import OrderedDict from collections.abc import Callable, Iterator, Mapping, Sequence from contextlib import contextmanager, nullcontext @@ -966,10 +967,12 @@ def _v2_configs(in_memory_loggers: Sequence[object], logger: "OpenTelemetryV2") def _registered_v2_logger() -> "OpenTelemetryV2 | None": - try: - from litellm.proxy import proxy_server - except Exception: - return None + """The proxy's registered V2 logger, read without importing the proxy. + + Request paths call this (the router's ``route`` phase among them), so importing + ``proxy_server`` here would load the whole proxy on an SDK caller's event loop. + """ + proxy_server: Final = sys.modules.get("litellm.proxy.proxy_server") logger: Final = getattr(proxy_server, "open_telemetry_logger", None) return logger if isinstance(logger, OpenTelemetryV2) else None diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 5a3a17f338c..b2c22880f5a 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -2259,11 +2259,11 @@ class Logging(LiteLLMLoggingBaseClass): except Exception: return True - def has_run_logging( + def mark_logging_complete( self, event_type: Literal["async_success", "sync_success", "async_failure", "sync_failure"], ) -> None: - if self.stream is not None and self.stream is True: + if self.stream is not None and self.stream is True and event_type in ["async_success", "sync_success"]: """ Ignore check on stream, as there can be multiple chunks """ @@ -2271,6 +2271,13 @@ class Logging(LiteLLMLoggingBaseClass): self.model_call_details[f"has_logged_{event_type}"] = True return + def has_run_logging( + self, + event_type: Literal["async_success", "sync_success", "async_failure", "sync_failure"], + ) -> None: + """Deprecated alias of mark_logging_complete, kept for callers of the old name""" + self.mark_logging_complete(event_type=event_type) + def should_run_callback(self, callback: litellm.CALLBACK_TYPES, litellm_params: dict, event_hook: str) -> bool: if litellm.global_disable_no_log_param: return True @@ -2862,7 +2869,7 @@ class Logging(LiteLLMLoggingBaseClass): call_type=self.call_type, ) - self.has_run_logging(event_type="sync_success") + self.mark_logging_complete(event_type="sync_success") for callback in callbacks: try: should_run = self.should_run_callback( @@ -3443,7 +3450,7 @@ class Logging(LiteLLMLoggingBaseClass): ) self._handle_callback_failure(callback=callback) - self.has_run_logging(event_type="async_success") + self.mark_logging_complete(event_type="async_success") for callback in callbacks: # check if callback can run for this request @@ -3736,7 +3743,7 @@ class Logging(LiteLLMLoggingBaseClass): model_call_details=(self.model_call_details if hasattr(self, "model_call_details") else {}), result=result, ) - self.has_run_logging(event_type="sync_failure") + self.mark_logging_complete(event_type="sync_failure") for callback in callbacks: try: should_run = self.should_run_callback( @@ -3928,7 +3935,7 @@ class Logging(LiteLLMLoggingBaseClass): result: Final = None # result sent to all loggers, init this to None incase it's not created - self.has_run_logging(event_type="async_failure") + self.mark_logging_complete(event_type="async_failure") for callback in callbacks: try: litellm_params = self.model_call_details.get("litellm_params", {}) diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 806240c9749..578a4e47dfd 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -224,6 +224,11 @@ def _write_back_message_text(message: _WritableMessage, target: MessageTextTarge _TOOL_USE_INPUT_ADAPTER: Final = TypeAdapter(dict[str, object]) +_RELEASED_TOOL_USE_STOP: Final = ( + b"event: message_delta\n" + b'data: {"type": "message_delta", "delta": {"stop_reason": "tool_use", "stop_sequence": null}, ' + b'"usage": {"output_tokens": 0}}\n\n' +) def _rewritten_tool_use_input(arguments: str) -> Mapping[str, object] | None: @@ -1573,6 +1578,12 @@ class AnthropicMessagesHandler(BaseTranslation): tool_calls_in_flight=bool(tool_use_fingerprints) and not stream_ended, ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + released_key: Final = self.get_streaming_scan_key(responses_so_far) + if released_key is None or not released_key.tool_calls_in_flight: + return tuple(responses_so_far) + return (*responses_so_far, _RELEASED_TOOL_USE_STOP) + @classmethod def _streamed_tool_use_fingerprints(cls, responses_so_far: Sequence[object]) -> tuple[str, ...]: return tuple( diff --git a/litellm/llms/base_llm/guardrail_translation/base_translation.py b/litellm/llms/base_llm/guardrail_translation/base_translation.py index 89ad67f0485..70b8291d32c 100644 --- a/litellm/llms/base_llm/guardrail_translation/base_translation.py +++ b/litellm/llms/base_llm/guardrail_translation/base_translation.py @@ -203,6 +203,11 @@ class BaseTranslation(ABC): def get_streaming_scan_key(self, responses_so_far: Sequence[object]) -> StreamingScanKey | None: return None + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + """The chunks a client left the stream with, closed the way this endpoint ends a stream, so the + end-of-stream scan also inspects tool calls the stream never finished""" + return tuple(responses_so_far) + def build_block_sse_chunks( self, exc: "ModifyResponseException", diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index aa175733582..aee439862db 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -17,7 +17,7 @@ This pattern can be replicated for other message formats (e.g., Anthropic). import json import time import uuid -from collections.abc import Mapping, Sequence +from collections.abc import Iterator, Mapping, Sequence from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Union, cast @@ -54,6 +54,7 @@ from litellm.types.utils import ( ChatCompletionDeltaToolCall, ChatCompletionMessageToolCall, Choices, + Delta, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream, @@ -837,6 +838,18 @@ class OpenAIChatCompletionsHandler(BaseTranslation): tool_calls_in_flight=bool(tool_call_fingerprints) and not stream_ended, ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + released_key: Final = self.get_streaming_scan_key(responses_so_far) + if released_key is None or not released_key.tool_calls_in_flight: + return tuple(responses_so_far) + terminator: Final = ModelResponseStream( + choices=[ + StreamingChoices(index=index, delta=Delta(), finish_reason="tool_calls") + for index in _choice_indices_with_tool_calls(responses_so_far) + ] + ) + return (*responses_so_far, terminator) + @staticmethod def _streamed_tool_call_fingerprints(responses_so_far: Sequence[object]) -> tuple[str, ...]: return tuple( @@ -1388,6 +1401,20 @@ def _streamed_delta_tool_calls(delta: object) -> tuple[object, ...]: return stream_item_items(delta, "tool_calls") + legacy +def _released_choices(responses_so_far: Sequence[object]) -> Iterator[object]: + for chunk in responses_so_far: + yield from _stream_chunk_choices(chunk) + + +def _choice_indices_with_tool_calls(responses_so_far: Sequence[object]) -> tuple[int, ...]: + indices: Final = ( + index if isinstance(index := stream_item_field(choice, "index"), int) else 0 + for choice in _released_choices(responses_so_far) + if _streamed_delta_tool_calls(stream_item_field(choice, "delta")) + ) + return tuple(dict.fromkeys(indices)) + + def _blocked_stream_identity( exc: "ModifyResponseException", responses_so_far: Sequence[object] ) -> tuple[str, int, str]: diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 90cdef87ec7..ad68924ce20 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -293,6 +293,51 @@ def _is_tool_call_output_item(item: object) -> bool: return _tool_call_output_item_mapping(item) is not None +def _released_tool_call_payload(responses_so_far: Sequence[object], item_id: object) -> str | None: + events: Final = tuple(event for event in responses_so_far if stream_item_field(event, "item_id") == item_id) + finished: Final = tuple( + payload + for event in events + if isinstance(event_type := stream_item_field(event, "type"), str) + and event_type in _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS + and isinstance(payload := stream_item_field(event, _TOOL_CALL_PAYLOAD_DONE_EVENT_FIELDS[event_type]), str) + ) + if finished: + return finished[-1] + deltas: Final = tuple( + delta + for event in events + if stream_item_field(event, "type") in _TOOL_CALL_PAYLOAD_DELTA_EVENT_TYPES + and isinstance(delta := stream_item_field(event, "delta"), str) + ) + return "".join(deltas) if deltas else None + + +def _with_released_payload(item: Mapping[str, object], responses_so_far: Sequence[object]) -> Mapping[str, object]: + field: Final = _TOOL_CALL_PAYLOAD_FIELDS[str(item.get("type"))] + payload: Final = _released_tool_call_payload(responses_so_far, item.get("id")) + return item if payload is None else {**item, field: payload} + + +def _released_message_item(text: str) -> Mapping[str, object]: + content: Final = [{"type": "output_text", "text": text}] + return {"type": "message", "role": "assistant", "content": content} + + +def _released_tool_call_items(responses_so_far: Sequence[object]) -> tuple[Mapping[str, object], ...]: + announced: Final = tuple( + item + for item in ( + _tool_call_output_item_mapping(stream_item_field(event, "item")) + for event in responses_so_far + if stream_item_field(event, "type") in _OUTPUT_ITEM_EVENT_TYPES + ) + if item is not None + ) + latest_by_id: Final = MappingProxyType({item.get("id"): item for item in announced}) + return tuple(_with_released_payload(item, responses_so_far) for item in latest_by_id.values()) + + def _last_message_role(messages: Sequence[object]) -> str | None: if not messages: return None @@ -1324,6 +1369,27 @@ class OpenAIResponsesHandler(BaseTranslation): tool_calls_in_flight=self._has_streamed_tool_call_events(responses_so_far), ) + def released_stream_as_ended(self, responses_so_far: Sequence[object]) -> tuple[object, ...]: + if self._check_streaming_has_ended(responses_so_far): + return tuple(responses_so_far) + ends_on_finished_item: Final = ( + bool(responses_so_far) + and stream_item_field(responses_so_far[-1], "type") == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE.value + ) + if not ends_on_finished_item and not self._has_streamed_tool_call_events(responses_so_far): + return tuple(responses_so_far) + text_events: Final = tuple( + event for event in responses_so_far if stream_item_field(event, "type") in _OUTPUT_TEXT_EVENT_TYPES + ) + released_text: Final = self.get_streaming_string_so_far(text_events) + message_items: Final = (_released_message_item(released_text),) if released_text else () + tool_items: Final = _released_tool_call_items(responses_so_far) + output: Final = [*message_items, *tool_items] + response: Final = {"status": "incomplete", "output": output} + incomplete: Final = ResponsesAPIStreamEvents.RESPONSE_INCOMPLETE.value + envelope: Final = {"type": incomplete, "response": response} + return (*responses_so_far, envelope) + @staticmethod def _has_streamed_tool_call_events(responses_so_far: Sequence[object]) -> bool: return any( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index cd38b5b0d21..2d0985bcb64 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -6137,7 +6137,7 @@ "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, "input_cost_per_token": 4e-06, @@ -6148,7 +6148,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 2.4e-05, - "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule", "supported_endpoints": [ "/v1/realtime" ], @@ -6172,7 +6172,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, @@ -6183,7 +6183,7 @@ "mode": "realtime", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, - "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule", "supported_endpoints": [ "/v1/realtime" ], @@ -42328,14 +42328,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 6e-09, - "input_cost_per_token": 3e-07, + "cache_read_input_token_cost": 3e-09, + "input_cost_per_token": 3e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.2e-06, + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42348,14 +42348,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "cache_read_input_token_cost": 1.4e-07, - "input_cost_per_token": 2.2e-07, + "cache_read_input_token_cost": 7e-07, + "input_cost_per_token": 8.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 4.2e-06, + "output_cost_per_token": 5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -67616,8 +67616,9 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.8-27b": { - "input_cost_per_token": 4.2e-07, - "output_cost_per_token": 3e-06, + "cache_creation_input_token_cost": 5.3125e-07, + "input_cost_per_token": 4.25e-07, + "output_cost_per_token": 2.55e-06, "cache_read_input_token_cost": 8.5e-08, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -67732,7 +67733,7 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "cache_read_input_token_cost": 1.37e-08, + "cache_read_input_token_cost": 1.52e-08, "input_cost_per_token": 1.52e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, @@ -67823,14 +67824,14 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k3": { - "cache_read_input_token_cost": 4.9e-07, - "input_cost_per_token": 4.99e-07, + "cache_read_input_token_cost": 7e-07, + "input_cost_per_token": 7.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.3e-05, + "output_cost_per_token": 1.4e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -67948,14 +67949,14 @@ "supports_web_search": true }, "openrouter/z-ai/glm-5.2": { - "cache_read_input_token_cost": 2.6e-07, - "input_cost_per_token": 3e-07, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1.04e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 3.49e-06, + "output_cost_per_token": 8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68311,14 +68312,14 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "cache_read_input_token_cost": 5.6e-09, - "input_cost_per_token": 2.8e-08, + "cache_read_input_token_cost": 2.24e-08, + "input_cost_per_token": 2.24e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 5.6e-08, + "output_cost_per_token": 1.28e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index d83a4a72e2f..1dea44a84f0 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -2286,8 +2286,8 @@ async def _run_post_mcp_call_guardrails( def suppress_completed_success_logging(logging_obj: LiteLLMLoggingObj) -> None: """An interim ``InputRequiredResult`` is not a completed call, so the ``@client`` wrapper on ``call_mcp_tool`` must not run the success handlers for it when the coroutine returns.""" - logging_obj.has_run_logging(event_type="sync_success") - logging_obj.has_run_logging(event_type="async_success") + logging_obj.mark_logging_complete(event_type="sync_success") + logging_obj.mark_logging_complete(event_type="async_success") async def _fire_mcp_tool_call_logging( @@ -2335,8 +2335,8 @@ async def _fire_mcp_tool_call_logging( await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) return result - logging_obj.has_run_logging(event_type="sync_success") - logging_obj.has_run_logging(event_type="async_success") + logging_obj.mark_logging_complete(event_type="sync_success") + logging_obj.mark_logging_complete(event_type="async_success") tool_error: Final = MCPToolResultError(error_message) logging_obj.failure_handler(tool_error, "", start_time, end_time) await logging_obj.async_failure_handler(tool_error, "", start_time, end_time) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index c85169f0ba5..687df9b0348 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -272,6 +272,18 @@ def _withheld_provider_output(response: object) -> bool: return getattr(response, "has_buffered_provider_output", False) is True +async def close_guarded_stream(stream: object) -> None: + if not isinstance(stream, AsyncGenerator): + return + with anyio.CancelScope(shield=True): + try: + await stream.aclose() + except Exception as e: # noqa: BLE001 # a failing callback cleanup must not skip the refund and finalizer + verbose_proxy_logger.warning( + "Closing the guarded stream after a client disconnect raised %s", type(e).__name__ + ) + + def resolve_litellm_call_id(client_call_id: str | None) -> str: if client_call_id is not None and 0 < len(client_call_id) <= MAX_LITELLM_CALL_ID_LENGTH: return client_call_id @@ -3909,13 +3921,14 @@ class ProxyBaseLLMRequestProcessing: client_disconnected = False delivered_chunk = False recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes + guarded_stream: Final[AsyncGenerator[object, None]] = proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) try: str_so_far = "" - async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ): + async for chunk in guarded_stream: # ``.format(chunk)`` was previously evaluated for every chunk # regardless of log level; gate it behind the level check. if debug_enabled: @@ -3971,6 +3984,7 @@ class ProxyBaseLLMRequestProcessing: # Starlette closes on disconnect, so the nested iterator hook (which # only sees GeneratorExit on GC) cannot own the refund. client_disconnected = not stream_completed + await close_guarded_stream(guarded_stream) if not delivered_chunk and not _withheld_provider_output(response): from litellm.proxy.spend_tracking.budget_reservation import ( release_budget_reservation_on_cancel, diff --git a/litellm/proxy/common_utils/fips.py b/litellm/proxy/common_utils/fips.py new file mode 100644 index 00000000000..07bd4d2462c --- /dev/null +++ b/litellm/proxy/common_utils/fips.py @@ -0,0 +1,150 @@ +import hashlib +import os +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final + +from typing_extensions import assert_never + +from litellm.secret_managers.main import str_to_bool + +FIPS_MODE_ENV_VAR: Final = "LITELLM_FIPS_MODE" +SSL_VERIFY_ENV_VAR: Final = "SSL_VERIFY" +SSL_VERIFY_SETTING: Final = "litellm_settings.ssl_verify" +REFUSAL_PREFIX: Final = "LiteLLM proxy refused to start" + +_TRUE_VALUES: Final = frozenset({"true", "1", "yes", "on"}) +_FALSE_VALUES: Final = frozenset({"false", "0", "no", "off", ""}) + + +@dataclass(frozen=True, slots=True) +class FipsModeOff: + pass + + +@dataclass(frozen=True, slots=True) +class FipsModeOn: + pass + + +@dataclass(frozen=True, slots=True) +class MalformedFipsMode: + value: str + + +FipsModeSetting = FipsModeOff | FipsModeOn | MalformedFipsMode + + +@dataclass(frozen=True, slots=True) +class ProviderDoesNotEnforceFips: + pass + + +@dataclass(frozen=True, slots=True) +class TlsVerificationDisabled: + sources: tuple[str, ...] + + +FipsBootRefusal = MalformedFipsMode | ProviderDoesNotEnforceFips | TlsVerificationDisabled +FipsBootVerdict = FipsModeOff | FipsModeOn | FipsBootRefusal + + +class FipsModeError(Exception): + pass + + +def parse_fips_mode(raw: str | None) -> FipsModeSetting: + if raw is None: + return FipsModeOff() + normalized: Final = raw.strip().lower() + if normalized in _TRUE_VALUES: + return FipsModeOn() + if normalized in _FALSE_VALUES: + return FipsModeOff() + return MalformedFipsMode(value=raw) + + +def is_fips_mode(environ: Callable[[str], str | None] = os.environ.get) -> bool: + return isinstance(parse_fips_mode(environ(FIPS_MODE_ENV_VAR)), FipsModeOn) + + +def openssl_enforces_fips() -> bool: + """MD5 is not an approved digest, so an enforcing FIPS provider refuses it even when asked for security use.""" + try: + hashlib.md5(b"", usedforsecurity=True) + except ValueError: + return True + return False + + +def fips_boot_verdict( + *, + raw_fips_mode: str | None, + provider_enforces_fips: Callable[[], bool], + ssl_verify_environment: str | None, + ssl_verify_setting: object, +) -> FipsBootVerdict: + setting: Final = parse_fips_mode(raw_fips_mode) + match setting: + case FipsModeOff() | MalformedFipsMode(): + return setting + case FipsModeOn(): + pass + case _: + assert_never(setting) + disabled: Final = tuple( + source + for source, off in ( + (SSL_VERIFY_ENV_VAR, _is_off(ssl_verify_environment)), + (SSL_VERIFY_SETTING, _is_off(ssl_verify_setting)), + ) + if off + ) + if disabled: + return TlsVerificationDisabled(sources=disabled) + if not provider_enforces_fips(): + return ProviderDoesNotEnforceFips() + return setting + + +def enforce_fips_boot_verdict(verdict: FipsBootVerdict, announce: Callable[[str], object]) -> None: + match verdict: + case FipsModeOff() | FipsModeOn(): + return + case MalformedFipsMode() | ProviderDoesNotEnforceFips() | TlsVerificationDisabled(): + message: Final = render_refusal(verdict) + announce(f"\n{message}\n\n") + raise FipsModeError(message) + case _: + assert_never(verdict) + + +def render_refusal(refusal: FipsBootRefusal) -> str: + match refusal: + case MalformedFipsMode(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR}={refusal.value} is not a boolean.\n" + f"Set {FIPS_MODE_ENV_VAR} to true or false, or unset it." + ) + case ProviderDoesNotEnforceFips(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but this Python does not enforce FIPS.\n" + "Its OpenSSL still allows non-approved algorithms (MD5 succeeded), so passwords and keys would be\n" + "protected with algorithms the FIPS 140-3 policy forbids. Run the proxy from a FIPS image whose\n" + f"OpenSSL FIPS provider is enabled, or unset {FIPS_MODE_ENV_VAR} on a non-FIPS runtime." + ) + case TlsVerificationDisabled(): + return ( + f"{REFUSAL_PREFIX}: {FIPS_MODE_ENV_VAR} is on but TLS certificate verification is disabled by " + f"{' and '.join(refusal.sources)}.\nFIPS deployments must verify upstream certificates, so remove the " + "override or point ssl_verify at a CA bundle instead." + ) + return assert_never(refusal) + + +def _is_off(value: object) -> bool: + if isinstance(value, bool): + return value is False + if isinstance(value, str): + return str_to_bool(value) is False + return False diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 620b24df95d..f900d14bdc1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -10,6 +10,7 @@ import sys sys.path.insert(0, os.path.abspath("../..")) # Adds the parent directory to the system path import asyncio +import contextlib import copy import json import re @@ -2747,14 +2748,17 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): UnifiedLLMGuardrails, ) - async for streamed_chunk in UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - guardrail_to_apply=self, - buffer_until_moderated_default=False, - ): - yield streamed_chunk + async with contextlib.aclosing( + UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + guardrail_to_apply=self, + buffer_until_moderated_default=False, + ) + ) as guarded: + async for streamed_chunk in guarded: + yield streamed_chunk return # Responses-API events are neither chat-completions chunks nor raw diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 37c1829def4..0c4ea6b29b5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -6,11 +6,14 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint 3. Implements a way to call /applyGuardrail endpoint for `/chat/completions` + `/v1/messages` requests on async_post_call_streaming_iterator_hook """ +import asyncio +import contextlib import copy import json from collections.abc import AsyncGenerator, AsyncIterable, Awaitable, Callable, Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final, Protocol, TypeAlias, cast +import anyio from fastapi import HTTPException from litellm._logging import verbose_proxy_logger @@ -19,6 +22,7 @@ from litellm.cost_calculator import _infer_call_type from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route +from litellm.litellm_core_utils.core_helpers import get_or_create_metadata_bucket from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings from litellm.proxy._types import UserAPIKeyAuth from litellm.types.guardrails import GuardrailEventHooks @@ -28,6 +32,7 @@ from litellm.types.utils import ( CallTypesLiteral, Delta, ModelResponseStream, + StandardLoggingGuardrailInformation, StreamingChoices, ) @@ -45,6 +50,8 @@ A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message) GUARDRAIL_NAME: Final = "unified_llm_guardrails" +_RequestData: TypeAlias = dict[str, object] + class _EndpointTranslation(Protocol): @property @@ -59,6 +66,9 @@ class _EndpointTranslation(Protocol): @property def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ... + @property + def released_stream_as_ended(self) -> "Callable[[Sequence[object]], tuple[object, ...]]": ... + @property def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ... @@ -109,6 +119,18 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]: return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0) +def _recorded_guardrail_information(request_data: _RequestData) -> tuple[StandardLoggingGuardrailInformation, ...]: + _metadata_key, metadata_bucket = get_or_create_metadata_bucket(request_data) + entries: Final = metadata_bucket.get("standard_logging_guardrail_information") + if not isinstance(entries, list): + return () + return tuple( + cast( # cast-ok: only the guardrail logging helpers write this metadata key + "list[StandardLoggingGuardrailInformation]", entries + ) + ) + + def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool: if scan_key is None: return False @@ -601,6 +623,7 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice: dict[int, str | None], held_chars_per_choice: dict[int, int], is_final: bool, + terminated: asyncio.Event, ) -> AsyncGenerator[object, None]: """Run one guardrail processing round and emit the resulting diff chunk. @@ -632,6 +655,7 @@ class UnifiedLLMGuardrails(CustomLogger): is_final=is_final, ) except ModifyResponseException as e: + terminated.set() if e.original_response is None: e.original_response = responses_so_far async for block_chunk in self.handle_streaming_block( @@ -643,6 +667,7 @@ class UnifiedLLMGuardrails(CustomLogger): yield block_chunk raise _StreamTerminated() except HTTPException as e: + terminated.set() async for error_item in self.emit_streaming_http_error( e, call_type, @@ -664,7 +689,7 @@ class UnifiedLLMGuardrails(CustomLogger): *, guardrail_to_apply: CustomGuardrail, response: AsyncIterable[object], - request_data: dict, + request_data: _RequestData, user_api_key_dict: UserAPIKeyAuth, call_type: str, sampling_rate: int, @@ -687,6 +712,7 @@ class UnifiedLLMGuardrails(CustomLogger): held_chars_per_choice: Final[dict[int, int]] = {} chunk_counter = 0 last_chunk: object | None = None + terminated: Final = asyncio.Event() def _round(reference_chunk: object, is_final: bool) -> AsyncGenerator[object, None]: return self._emit_transform_round( @@ -702,10 +728,13 @@ class UnifiedLLMGuardrails(CustomLogger): finish_reason_per_choice=finish_reason_per_choice, held_chars_per_choice=held_chars_per_choice, is_final=is_final, + terminated=terminated, ) saw_tool_calls = False saw_text_content = False + tool_calls_released = False # rebind-ok: set once a raw tool call reaches the client unscanned + end_of_stream_inspection_started = False # rebind-ok: set once the end-of-stream inspection owns the verdict try: async for item in response: @@ -742,6 +771,7 @@ class UnifiedLLMGuardrails(CustomLogger): held_choices=_held_choices(held_chars_per_choice), ) responses_yielded.append(tool_only) + tool_calls_released = True yield tool_only continue @@ -781,6 +811,7 @@ class UnifiedLLMGuardrails(CustomLogger): # ``stream_transform_underflow`` 400 from mismatched prefixes. A shallow # list copy wouldn't help — the mutation is on the chunk objects # themselves — so we deepcopy. + end_of_stream_inspection_started = True if saw_tool_calls: async for out in self._inspect_full_response_for_block( endpoint_translation=endpoint_translation, @@ -801,6 +832,37 @@ class UnifiedLLMGuardrails(CustomLogger): yield out except _StreamTerminated: return + except (GeneratorExit, asyncio.CancelledError): + await self._scan_uninspected_tool_calls_after_disconnect( + uninspected=tool_calls_released and not end_of_stream_inspection_started and not terminated.is_set(), + endpoint_translation=endpoint_translation, + responses_released=responses_yielded, + guardrail_to_apply=guardrail_to_apply, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + raise + + @staticmethod + async def _scan_uninspected_tool_calls_after_disconnect( + *, + uninspected: bool, + endpoint_translation: _EndpointTranslation, + responses_released: Sequence[object], + guardrail_to_apply: CustomGuardrail, + user_api_key_dict: UserAPIKeyAuth, + request_data: _RequestData, + ) -> None: + if not uninspected: + return + await UnifiedLLMGuardrails._scan_released_stream_after_disconnect( + endpoint_translation=endpoint_translation, + responses_released=responses_released, + last_scan_key=None, + guardrail_to_apply=guardrail_to_apply, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) async def _emit_stream_tail( self, @@ -841,14 +903,15 @@ class UnifiedLLMGuardrails(CustomLogger): from litellm.integrations.custom_guardrail import ModifyResponseException try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=request_data, - stream_transform_sink=None, - ) + with anyio.CancelScope(shield=bool(responses_yielded)): + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + stream_transform_sink=None, + ) except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far @@ -968,6 +1031,45 @@ class UnifiedLLMGuardrails(CustomLogger): config_value: Final = config.get(name, attribute_value) if isinstance(config, dict) else attribute_value return self.optional_params.get(name, config_value) + @staticmethod + async def _scan_released_stream_after_disconnect( + *, + endpoint_translation: _EndpointTranslation, + responses_released: Sequence[object], + last_scan_key: "StreamingScanKey | None", + guardrail_to_apply: CustomGuardrail, + user_api_key_dict: UserAPIKeyAuth, + request_data: _RequestData, + ) -> None: + scanned: Final = endpoint_translation.released_stream_as_ended(copy.deepcopy(tuple(responses_released))) + if _is_redundant_scan(endpoint_translation.get_streaming_scan_key(scanned), last_scan_key): + return + recorded_before: Final = len(_recorded_guardrail_information(request_data)) + with anyio.CancelScope(shield=True): + try: + await endpoint_translation.process_output_streaming_response( + responses_so_far=scanned, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except Exception as e: # noqa: BLE001 # the client is gone, so the verdict can only be recorded + verbose_proxy_logger.warning( + "UnifiedLLMGuardrails: %s scanned a stream the client disconnected from and raised %s", + guardrail_to_apply.guardrail_name, + type(e).__name__, + ) + recorded_during_scan: Final = _recorded_guardrail_information(request_data)[recorded_before:] + if any(entry.get("guardrail_status") != "success" for entry in recorded_during_scan): + return + guardrail_to_apply.add_standard_logging_guardrail_information_to_request_data( + guardrail_json_response=e, + request_data=request_data, + guardrail_status="guardrail_failed_to_respond", + event_type=GuardrailEventHooks.post_call, + ) + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, @@ -995,6 +1097,7 @@ class UnifiedLLMGuardrails(CustomLogger): if guardrail_to_apply is None: guardrail_to_apply = request_data.pop("guardrail_to_apply", None) + typed_request_data: Final[_RequestData] = request_data def _streaming_flag(name: str, default: object) -> Any: return self.resolve_streaming_flag(guardrail_to_apply, name, default) @@ -1061,17 +1164,20 @@ class UnifiedLLMGuardrails(CustomLogger): mappings=mappings, ) if transform_call_type is not None: - async for transformed_item in self._run_incremental_transform_stream( - guardrail_to_apply=guardrail_to_apply, - response=response, - request_data=request_data, - user_api_key_dict=user_api_key_dict, - call_type=transform_call_type, - sampling_rate=sampling_rate, - end_of_stream_only=end_of_stream_only, - mappings=mappings, - ): - yield transformed_item + async with contextlib.aclosing( + self._run_incremental_transform_stream( + guardrail_to_apply=guardrail_to_apply, + response=response, + request_data=typed_request_data, + user_api_key_dict=user_api_key_dict, + call_type=transform_call_type, + sampling_rate=sampling_rate, + end_of_stream_only=end_of_stream_only, + mappings=mappings, + ) + ) as transformed: + async for transformed_item in transformed: + yield transformed_item return verbose_proxy_logger.warning( "UnifiedLLMGuardrails: streaming_transform_mode=incremental_diff is only supported " @@ -1093,217 +1199,240 @@ class UnifiedLLMGuardrails(CustomLogger): chunks_yielded = False last_scan_key: StreamingScanKey | None = None # rebind-ok: replaced after every scan round tool_calls_in_flight = False # rebind-ok: tracks the latest scan key's unscanned tool calls + verdict_settled = False # rebind-ok: set once the end-of-stream scan or a block owns the verdict - async for item in response: - chunk_counter += 1 - responses_so_far.append(item) + try: + async for item in response: + chunk_counter += 1 + responses_so_far.append(item) - # Infer call type from first chunk if not already done - if call_type is None and user_api_key_dict.request_route is not None: - call_types = get_call_types_for_route(user_api_key_dict.request_route) - if call_types is not None: - call_type = call_types[0].value + # Infer call type from first chunk if not already done + if call_type is None and user_api_key_dict.request_route is not None: + call_types = get_call_types_for_route(user_api_key_dict.request_route) + if call_types is not None: + call_type = call_types[0].value - if call_type is None: - call_type = _infer_call_type(call_type=None, completion_response=item) + if call_type is None: + call_type = _infer_call_type(call_type=None, completion_response=item) - # If call type not supported, just pass through all chunks - if call_type is None or CallTypes(call_type) not in mappings: - yield item - async for remaining_item in response: - yield remaining_item - return + # If call type not supported, just pass through all chunks + if call_type is None or CallTypes(call_type) not in mappings: + yield item + async for remaining_item in response: + yield remaining_item + return - # If end_of_stream_only mode, yield chunks without processing. - # When buffering, withhold them instead -- they are released (or - # replaced by the block message) only after end-of-stream - # moderation runs below. - if end_of_stream_only: - if not buffer_until_moderated: - endpoint_translation = mappings[CallTypes(call_type)]() - stream_has_ended = hasattr( - endpoint_translation, "_check_streaming_has_ended" - ) and endpoint_translation._check_streaming_has_ended(responses_so_far) - if pending_end_of_stream_items or stream_has_ended: - pending_end_of_stream_items.append(item) - else: - chunks_yielded = True - responses_yielded.append(item) - yield item - else: - withheld_items.append(item) - continue - - # Process chunk based on sampling rate - if buffer_until_moderated: - withheld_items.append(item) - if chunk_counter % sampling_rate == 0: - endpoint_translation = mappings[CallTypes(call_type)]() - scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far) - if scan_key is not None: - tool_calls_in_flight = scan_key.tool_calls_in_flight - hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight) - if _is_redundant_scan(scan_key, last_scan_key): - verbose_proxy_logger.debug( - "Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round", - chunk_counter, - guardrail_to_apply.guardrail_name, - ) - if buffer_until_moderated: - if hold_window: - continue - for withheld_item in withheld_items: + # If end_of_stream_only mode, yield chunks without processing. + # When buffering, withhold them instead -- they are released (or + # replaced by the block message) only after end-of-stream + # moderation runs below. + if end_of_stream_only: + if not buffer_until_moderated: + endpoint_translation = mappings[CallTypes(call_type)]() + stream_has_ended = hasattr( + endpoint_translation, "_check_streaming_has_ended" + ) and endpoint_translation._check_streaming_has_ended(responses_so_far) + if pending_end_of_stream_items or stream_has_ended: + pending_end_of_stream_items.append(item) + else: chunks_yielded = True - responses_yielded.append(withheld_item) - yield withheld_item - withheld_items.clear() + responses_yielded.append(item) + yield item else: - chunks_yielded = True - responses_yielded.append(item) - yield item + withheld_items.append(item) continue + # Process chunk based on sampling rate + if buffer_until_moderated: + withheld_items.append(item) + if chunk_counter % sampling_rate == 0: + endpoint_translation = mappings[CallTypes(call_type)]() + scan_key = endpoint_translation.get_streaming_scan_key(responses_so_far) + if scan_key is not None: + tool_calls_in_flight = scan_key.tool_calls_in_flight + hold_window = buffer_until_moderated and (scan_key is None or tool_calls_in_flight) + if _is_redundant_scan(scan_key, last_scan_key): + verbose_proxy_logger.debug( + "Skipping streaming chunk %s for guardrail %s: nothing new to scan since the last round", + chunk_counter, + guardrail_to_apply.guardrail_name, + ) + if buffer_until_moderated: + if hold_window: + continue + for withheld_item in withheld_items: + chunks_yielded = True + responses_yielded.append(withheld_item) + yield withheld_item + withheld_items.clear() + else: + chunks_yielded = True + responses_yielded.append(item) + yield item + continue + + verbose_proxy_logger.debug( + "Processing streaming chunk %s (sampling_rate=%s) with guardrail %s", + chunk_counter, + sampling_rate, + guardrail_to_apply.guardrail_name, + ) + + original_items = ( + tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),) + ) + + try: + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + except ModifyResponseException as e: + verdict_settled = True + if e.original_response is None: + e.original_response = responses_so_far + # Guardrail blocked the response mid-stream. Emit a clean + # terminating SSE sequence delivering the block message + # instead of letting the exception propagate into a bare + # `data: {"error": ...}` blob (which truncates the stream). + # Chunks have already been forwarded here, so the block + # continues the in-progress message (stream_started=True). + # The current chunk was appended to responses_so_far but not + # yet yielded, so exclude it: the continuation must reflect + # only what the client has actually received. + async for block_chunk in self.handle_streaming_block( + e, + endpoint_translation, + stream_started=chunks_yielded, + responses_so_far=responses_yielded, + ): + yield block_chunk + return + except HTTPException as e: + verdict_settled = True + # Response already started (we already yielded chunks); cannot send 400. + async for error_item in self.emit_streaming_http_error( + e, + call_type, + responses_so_far, + request_data, + endpoint_translation=endpoint_translation, + stream_started=chunks_yielded, + responses_yielded=responses_yielded, + ): + yield error_item + return + if scan_key is not None: + last_scan_key = scan_key + if hold_window: + verbose_proxy_logger.debug( + "Holding %s buffered chunks for guardrail %s: this round could not scan the whole window", + len(withheld_items), + guardrail_to_apply.guardrail_name, + ) + withheld_items[:] = original_items + continue + for original_item in original_items: + chunks_yielded = True + responses_yielded.append(original_item) + yield original_item + withheld_items.clear() + else: + if not buffer_until_moderated: + chunks_yielded = True + responses_yielded.append(item) + yield item + + # Stream has ended - do final processing with all collected chunks + if call_type is not None and CallTypes(call_type) in mappings: verbose_proxy_logger.debug( - "Processing streaming chunk %s (sampling_rate=%s) with guardrail %s", - chunk_counter, - sampling_rate, + "Processing final streaming response with all %s chunks for guardrail %s", + len(responses_so_far), guardrail_to_apply.guardrail_name, ) - original_items = ( - tuple(copy.deepcopy(withheld_items)) if buffer_until_moderated else (copy.deepcopy(item),) + endpoint_translation = mappings[CallTypes(call_type)]() + + buffered_items: Final = ( + tuple(copy.deepcopy(withheld_items)) + if buffer_until_moderated and release_on_scan and not end_of_stream_only + else tuple(withheld_items) + if buffer_until_moderated + else None ) + end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far) + verdict_settled = True + if _is_redundant_scan(end_scan_key, last_scan_key): + verbose_proxy_logger.debug( + "Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all", + guardrail_to_apply.guardrail_name, + ) + for buffered_item in buffered_items or (): + yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item + return try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - request_data=request_data, - ) + with anyio.CancelScope(shield=chunks_yielded): + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + # Moderation passed: release the withheld original chunks. + if buffered_items is not None: + for buffered_item in buffered_items: + yield buffered_item + for pending_item in pending_end_of_stream_items: + responses_yielded.append(pending_item) + yield pending_item except ModifyResponseException as e: if e.original_response is None: e.original_response = responses_so_far - # Guardrail blocked the response mid-stream. Emit a clean - # terminating SSE sequence delivering the block message - # instead of letting the exception propagate into a bare - # `data: {"error": ...}` blob (which truncates the stream). - # Chunks have already been forwarded here, so the block - # continues the in-progress message (stream_started=True). - # The current chunk was appended to responses_so_far but not - # yet yielded, so exclude it: the continuation must reflect - # only what the client has actually received. + # Block detected during end-of-stream processing. Emit a clean + # terminating SSE sequence with the block message rather than + # propagating into a bare error blob that truncates the stream. + # The withheld original chunks are never released. async for block_chunk in self.handle_streaming_block( e, endpoint_translation, - stream_started=chunks_yielded, + stream_started=bool(responses_yielded), responses_so_far=responses_yielded, ): yield block_chunk return except HTTPException as e: - # Response already started (we already yielded chunks); cannot send 400. async for error_item in self.emit_streaming_http_error( e, call_type, responses_so_far, request_data, endpoint_translation=endpoint_translation, - stream_started=chunks_yielded, + stream_started=bool(responses_yielded), responses_yielded=responses_yielded, ): yield error_item - return - if scan_key is not None: - last_scan_key = scan_key - if hold_window: - verbose_proxy_logger.debug( - "Holding %s buffered chunks for guardrail %s: this round could not scan the whole window", - len(withheld_items), - guardrail_to_apply.guardrail_name, - ) - withheld_items[:] = original_items - continue - for original_item in original_items: - chunks_yielded = True - responses_yielded.append(original_item) - yield original_item - withheld_items.clear() - else: - if not buffer_until_moderated: - chunks_yielded = True - responses_yielded.append(item) - yield item - - # Stream has ended - do final processing with all collected chunks - if call_type is not None and CallTypes(call_type) in mappings: - verbose_proxy_logger.debug( - "Processing final streaming response with all %s chunks for guardrail %s", - len(responses_so_far), - guardrail_to_apply.guardrail_name, - ) - - endpoint_translation = mappings[CallTypes(call_type)]() - - buffered_items: Final = ( - tuple(copy.deepcopy(withheld_items)) - if buffer_until_moderated and release_on_scan and not end_of_stream_only - else tuple(withheld_items) - if buffer_until_moderated - else None - ) - end_scan_key: Final = endpoint_translation.get_streaming_scan_key(responses_so_far) - if _is_redundant_scan(end_scan_key, last_scan_key): - verbose_proxy_logger.debug( - "Skipping end-of-stream scan for guardrail %s: the last sampled round already scanned it all", - guardrail_to_apply.guardrail_name, - ) - for buffered_item in buffered_items or (): - yield buffered_item - for pending_item in pending_end_of_stream_items: - responses_yielded.append(pending_item) - yield pending_item - return - - try: - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, + except (GeneratorExit, asyncio.CancelledError): + translation_class: Final = None if call_type is None else mappings.get(CallTypes(call_type)) + if ( + chunks_yielded + and not verdict_settled + and translation_class is not None + and isinstance(guardrail_to_apply, CustomGuardrail) + ): + await self._scan_released_stream_after_disconnect( + endpoint_translation=translation_class(), + responses_released=responses_yielded, + last_scan_key=last_scan_key, guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), user_api_key_dict=user_api_key_dict, - request_data=request_data, + request_data=typed_request_data, ) - # Moderation passed: release the withheld original chunks. - if buffered_items is not None: - for buffered_item in buffered_items: - yield buffered_item - for pending_item in pending_end_of_stream_items: - responses_yielded.append(pending_item) - yield pending_item - except ModifyResponseException as e: - if e.original_response is None: - e.original_response = responses_so_far - # Block detected during end-of-stream processing. Emit a clean - # terminating SSE sequence with the block message rather than - # propagating into a bare error blob that truncates the stream. - # The withheld original chunks are never released. - async for block_chunk in self.handle_streaming_block( - e, - endpoint_translation, - stream_started=bool(responses_yielded), - responses_so_far=responses_yielded, - ): - yield block_chunk - return - except HTTPException as e: - async for error_item in self.emit_streaming_http_error( - e, - call_type, - responses_so_far, - request_data, - endpoint_translation=endpoint_translation, - stream_started=bool(responses_yielded), - responses_yielded=responses_yielded, - ): - yield error_item + raise diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 3ae9db0a600..12119f07962 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -20,6 +20,7 @@ from litellm.proxy.db.routing_prisma_wrapper import writer_wrapper from litellm.proxy.lens.billing import validate_key from litellm.proxy.lens.inference import Deployment, deployment_prices from litellm.proxy.lens.models import ( + ActivitySelection, Claim, Execution, ExecutionContent, @@ -129,7 +130,7 @@ def required(lens: Lens | None) -> Lens: return lens -def validate_selection(settings: LensSettings) -> None: +def validate_selection(settings: ActivitySelection) -> None: for identity in settings.execution_ids: try: source, _, _, _ = parse_execution(identity) @@ -391,13 +392,13 @@ async def update_finding(lens_id: str, finding_id: str, body: FindingUpdate, aut class Preview(BaseModel): as_of: AwareDatetime | None = None offset: int = Field(default=0, ge=0) - settings: LensSettings + selection: ActivitySelection lookback_hours: LookbackHours = 24 @router.post("/preview/sample", response_model=Sample) async def preview_sample(body: Preview, auth: Auth, storage: StorageDep) -> Sample: - validate_selection(body.settings) + validate_selection(body.selection) now: Final = min(body.as_of or datetime.now(timezone.utc), datetime.now(timezone.utc)) try: start: Final = int((now - timedelta(hours=body.lookback_hours)).timestamp() * 1000) @@ -406,7 +407,7 @@ async def preview_sample(body: Preview, auth: Auth, storage: StorageDep) -> Samp raise HTTPException(422, "Preview window exceeds the supported calendar range") from error return await source_reader(storage).sample( user_scope(auth), - body.settings, + body.selection, start, end, offset=body.offset, diff --git a/litellm/proxy/lens/models.py b/litellm/proxy/lens/models.py index e3a08103c8d..90c92cc7acd 100644 --- a/litellm/proxy/lens/models.py +++ b/litellm/proxy/lens/models.py @@ -45,23 +45,26 @@ class Check(Record): enabled: bool = True -class LensSettings(Record): - name: str = Field(min_length=1) - context: str = Field(default="") +class ActivitySelection(Record): source: Literal["traces", "requests", "both"] = "traces" - lookback_hours: LookbackHours = 24 service: str = Field(default="") agent_name: str = Field(default="") filters: tuple[MetadataFilter, ...] = Field(default=()) + sample_size: int | None = Field(default=None, ge=1) + sample_percent: float = Field(default=100, gt=0, le=100, allow_inf_nan=False) + team_id: str = "" + execution_ids: tuple[str, ...] = () + + +class LensSettings(ActivitySelection): + name: str = Field(min_length=1) + context: str = Field(default="") + lookback_hours: LookbackHours = 24 checks: tuple[Check, ...] = () model: str = Field(min_length=1) enabled: bool = True interval_minutes: IntervalMinutes = 15 - sample_size: int | None = Field(default=None, ge=1) - sample_percent: float = Field(default=100, gt=0, le=100, allow_inf_nan=False) concurrency: int = Field(default=8, ge=1) - team_id: str = "" - execution_ids: tuple[str, ...] = () monthly_budget: float = Field(default=100, gt=0, allow_inf_nan=False) @model_validator(mode="after") diff --git a/litellm/proxy/lens/sources.py b/litellm/proxy/lens/sources.py index 9af0f3679b6..e36653aa091 100644 --- a/litellm/proxy/lens/sources.py +++ b/litellm/proxy/lens/sources.py @@ -6,10 +6,10 @@ from typing import Final, Protocol, TypeAlias from pydantic import TypeAdapter from litellm.proxy.lens.models import ( + ActivitySelection, Evidence, Execution, ExecutionContent, - LensSettings, MetadataFilter, Sample, Scope, @@ -73,7 +73,7 @@ class SourceReader: async def sample( self, scope: Scope, - settings: LensSettings, + settings: ActivitySelection, start: int, end: int, offset: int = 0, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f40e541dd86..ed2a5b89a32 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -399,6 +399,7 @@ from litellm.proxy.common_request_processing import ( ProxyBaseLLMRequestProcessing, _is_azure_model_router_request, _should_return_raw_model_name, + close_guarded_stream, create_response, log_llm_api_exception, open_sse_before_first_byte, @@ -422,6 +423,14 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import ( encrypt_value_helper, ) from litellm.proxy.common_utils.error_body_call_id import JSON_OBJECT, error_body_call_id, with_call_id +from litellm.proxy.common_utils.fips import ( + FIPS_MODE_ENV_VAR, + SSL_VERIFY_ENV_VAR, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + openssl_enforces_fips, +) from litellm.proxy.common_utils.healthy_model_filter import ( get_hidden_unhealthy_model_names, is_healthy_only_listing_default, @@ -1329,6 +1338,16 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState if isinstance(worker_config, dict): await initialize_from_worker_config(worker_config) + enforce_fips_boot_verdict( + fips_boot_verdict( + raw_fips_mode=os.getenv(FIPS_MODE_ENV_VAR), + provider_enforces_fips=openssl_enforces_fips, + ssl_verify_environment=os.getenv(SSL_VERIFY_ENV_VAR), + ssl_verify_setting=litellm.ssl_verify, + ), + announce=announce_on_stderr_at_exit, + ) + enforce_master_key_boot_verdict( await with_stored_secrets_counted( master_key_boot_verdict( @@ -1368,10 +1387,21 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState try: result: Final = await migrate_passwords_to_scrypt_async(prisma_client) verbose_proxy_logger.info("Password migration: %s", result) + except ValueError as e: + verbose_proxy_logger.error( + "Password migration failed, so plaintext passwords stay unhashed in the database: %s. " + "This is what an OpenSSL FIPS provider reports when the hashing algorithm is not approved.", + e, + ) + if is_fips_mode(): + raise except Exception as e: verbose_proxy_logger.warning("Password migration skipped: %s", e) - asyncio.create_task(_run_pw_migration()) + if is_fips_mode(): + await _run_pw_migration() + else: + asyncio.create_task(_run_pw_migration()) async def _run_agent_grant_id_migration() -> None: from litellm.proxy.agent_endpoints.agent_registry import ( @@ -9882,6 +9912,17 @@ async def async_data_generator( stream_completed = False client_disconnected = False error_state: Final = ResponsesStreamErrorState() if responses_stream_errors else None + needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap() + stream_iterator: Final[AsyncIterator[object]] = ( + proxy_logging_obj.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=response, + request_data=request_data, + ) + if needs_iterator_wrap + else response + ) + stream_source: AsyncIterator[object] | None = None # rebind-ok: bound once the keepalive policy resolves try: error_message: str | None = None requested_model_from_client: Final = _get_client_requested_model_for_streaming(request_data=request_data) @@ -9906,21 +9947,11 @@ async def async_data_generator( # per-chunk hook. Coalescing them into a single flag forced wasted # ``get_response_string`` work per chunk on every deployment that # happened to ship a streaming-iterator override (the default). - needs_iterator_wrap: Final = proxy_logging_obj.needs_iterator_wrap() needs_per_chunk_hook: Final = proxy_logging_obj.needs_per_chunk_streaming_hook() is_raw_sse_stream: Final = bool(request_data.get("_litellm_raw_sse_stream")) strip_stream_usage: Final = bool(request_data.get("_litellm_strip_stream_usage")) raw_sse_buffer = "" - if needs_iterator_wrap: - stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook( - user_api_key_dict=user_api_key_dict, - response=response, - request_data=request_data, - ) - else: - stream_iterator = response - # A stream can start on a deployment with keepalive off and fall back # mid-stream to one that enables it: only skip wrapping altogether when # there's no router to ever fall back through AND the resolved interval @@ -9929,7 +9960,7 @@ async def async_data_generator( # happens to start with it off. resolve_keepalive_seconds: Final = _make_keepalive_resolver(request_data) initial_keepalive_seconds: Final = resolve_keepalive_seconds(response) - stream_source: Final = ( + stream_source = ( _iter_with_keepalive( stream_iterator.__aiter__(), resolve_keepalive_seconds, @@ -10070,6 +10101,9 @@ async def async_data_generator( # (a nested iterator hook would only see GeneratorExit on GC). if not stream_completed: client_disconnected = True + for guarded_layer in (stream_source, stream_iterator): + if guarded_layer is not response: + await close_guarded_stream(guarded_layer) raise except Exception as e: verbose_proxy_logger.exception("litellm.proxy.proxy_server.async_data_generator(): Exception occured - %s", e) diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index 1563e4b5b55..0d3723c45b8 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -13,19 +13,21 @@ from dataclasses import dataclass from functools import partial from http.client import responses from types import MappingProxyType -from typing import Annotated, Final +from typing import Annotated, Final, Literal from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response from pydantic import BaseModel, ConfigDict +from typing_extensions import assert_never from litellm._logging import verbose_proxy_logger -from litellm.constants import OTLP_RETRY_AFTER_SECONDS +from litellm.constants import OTLP_RETRY_AFTER_SECONDS, TRACE_READ_RETRY_AFTER_SECONDS from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.authorization import AllRows, ReadScope, resolve_trace_read_scope from litellm.proxy.auth.authorization_dependencies import LogTeamLookupDependency 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.rust_bridge.trace.errors import TraceChanged from litellm.rust_bridge.trace.generated.models import TraceQueryHelp from litellm.rust_bridge.trace.generated.types import ( AllQueryScope, @@ -135,6 +137,49 @@ async def ingest_otlp_traces( return Response(content=body, media_type=media_type) +class TraceReadFailure(BaseModel): + """The body of every failed trace read. Clients branch on `code`, never on `message`.""" + + model_config = ConfigDict(frozen=True) + + code: Literal["invalid_request", "trace_changed", "too_large", "unavailable"] + message: str + + +def read_failure(error: TraceChanged | ValueError | OverflowError | RuntimeError) -> HTTPException: + """One status per failure kind, so a client can tell a bad cursor (400, fix the request) from a + traversal it must restart (409), a result it cannot page through (413), and an outage it should + retry after `Retry-After` (503).""" + match error: + case TraceChanged(): + return HTTPException( + status_code=409, + detail=TraceReadFailure(code="trace_changed", message=str(error)).model_dump(), + ) + case ValueError(): + return HTTPException( + status_code=400, detail=TraceReadFailure(code="invalid_request", message=str(error)).model_dump() + ) + case OverflowError(): + return HTTPException( + status_code=413, + detail=TraceReadFailure( + code="too_large", message="Trace is too large for this view. Use a filtered trace query." + ).model_dump(), + ) + case RuntimeError(): + verbose_proxy_logger.warning("Trace read unavailable: %s", error) + return HTTPException( + status_code=503, + detail=TraceReadFailure( + code="unavailable", message="Traces are temporarily unavailable. Please try again." + ).model_dump(), + headers={"Retry-After": str(TRACE_READ_RETRY_AFTER_SECONDS)}, + ) + case _: + return assert_never(error) + + @router.get("/v1/traces", response_model=TracePage) async def list_agent_traces( context: Annotated[TraceAccessContext, Depends(provide_trace_access)], @@ -151,15 +196,8 @@ async def list_agent_traces( end_ms=end_ms if end_ms is not None else now_ms, cursor=cursor, ) - except ValueError as error: - raise HTTPException(status_code=400, detail=str(error)) from error - except OverflowError as error: - raise HTTPException( - status_code=413, detail="Trace is too large for this view. Use a filtered trace query." - ) from error - except RuntimeError as error: - verbose_proxy_logger.warning("Trace read unavailable: %s", error) - raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error + except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: + raise read_failure(error) from error class TraceQueryRequest(BaseModel): @@ -241,15 +279,8 @@ async def get_agent_trace( tracing, scope = context.reader() try: trace: Final = await tracing.get_trace(trace_id, scope, trace_ref, cursor, page_size) - except ValueError as error: - raise HTTPException(status_code=400, detail=str(error)) from error - except OverflowError as error: - raise HTTPException( - status_code=413, detail="Trace is too large for this view. Use a filtered trace query." - ) from error - except RuntimeError as error: - verbose_proxy_logger.warning("Trace read unavailable: %s", error) - raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error + except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: + raise read_failure(error) from error if trace is None: raise HTTPException(status_code=404, detail=f"Trace {trace_id} not found") return trace @@ -265,15 +296,8 @@ async def get_agent_trace_span( tracing, scope = context.reader() try: span: Final = await tracing.get_span(trace_id, span_id, scope, trace_ref) - except ValueError as error: - raise HTTPException(status_code=400, detail=str(error)) from error - except OverflowError as error: - raise HTTPException( - status_code=413, detail="Trace is too large for this view. Use a filtered trace query." - ) from error - except RuntimeError as error: - verbose_proxy_logger.warning("Trace read unavailable: %s", error) - raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error + except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: + raise read_failure(error) from error if span is None: raise HTTPException(status_code=404, detail=f"Span {span_id} not found") return span @@ -290,15 +314,8 @@ async def get_agent_trace_span_error( 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 - except OverflowError as error: - raise HTTPException( - status_code=413, detail="Trace is too large for this view. Use a filtered trace query." - ) from error - except RuntimeError as error: - verbose_proxy_logger.warning("Trace read unavailable: %s", error) - raise HTTPException(status_code=503, detail="Traces are temporarily unavailable. Please try again.") from error + except (TraceChanged, ValueError, OverflowError, RuntimeError) as error: + raise read_failure(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/proxy/utils.py b/litellm/proxy/utils.py index 1c738aefb97..29f2f46f001 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -44,6 +44,7 @@ from typing import ( Union, cast, overload, + runtime_checkable, ) from typing_extensions import ReadOnly, TypedDict @@ -524,8 +525,13 @@ class _UpstreamStreamBoundary(Generic[_T]): raise +@runtime_checkable +class _ClosableAsyncIterator(Protocol): + def aclose(self) -> object: ... + + class _StreamIteratorHook(Protocol[_T]): - def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ... + def __call__(self, *, response: AsyncIterator[_T]) -> AsyncIterator[_T]: ... def _is_client_error_exception(exc: Exception) -> bool: @@ -2773,8 +2779,22 @@ class ProxyLogging: ) -> AsyncGenerator[_T, None]: upstream: Final = _UpstreamStreamBoundary(response) try: - async for chunk in hook(response=upstream): - yield chunk + guarded: Final = hook(response=upstream) + try: + async for chunk in guarded: + yield chunk + finally: + if isinstance(guarded, _ClosableAsyncIterator): + try: + closing: Final = guarded.aclose() + if inspect.isawaitable(closing): + await closing + except Exception as e: # noqa: BLE001 # a finished stream must not fail on callback cleanup + verbose_proxy_logger.warning( + "Closing the streaming iterator of %s raised %s", + getattr(callback, "guardrail_name", None) or type(callback).__name__, + type(e).__name__, + ) except Exception as e: if e is not upstream.failure: enrich_http_exception_with_guardrail_context(e, callback) @@ -3959,6 +3979,7 @@ class ProxyLogging: stream_needs_translation: Final = ProxyLogging._stream_requires_guardrail_translation(user_api_key_dict) pipeline_gated_names: Final = _pipeline_step_guardrail_names(post_call_pipelines) + guarded_layers: Final[list[AsyncGenerator[object, None]]] = [] # mutable-ok: closed on disconnect for resolved_callback, kind in caps.iterator_overrides: if isinstance(resolved_callback, CustomGuardrail): if resolved_callback.guardrail_name in pipeline_gated_names: @@ -4001,6 +4022,7 @@ class ProxyLogging: hook, request_data=request_data, ) + guarded_layers.append(current_response) pipeline_translation: Final = ( resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None @@ -4013,6 +4035,7 @@ class ProxyLogging: pipelines=post_call_pipelines, translation=pipeline_translation, ) + guarded_layers.append(current_response) served_chunks: Final[list[object]] = [] # mutable-ok: accumulates while yielding to the client try: @@ -4020,6 +4043,7 @@ class ProxyLogging: served_chunks.append(chunk) yield chunk except (GeneratorExit, asyncio.CancelledError): + await ProxyLogging._close_guarded_layers(guarded_layers) ProxyLogging._record_served_stream_output(request_data, served_chunks) raise except Exception as e: @@ -4100,6 +4124,16 @@ class ProxyLogging: for buffered_item in buffered: yield buffered_item + @staticmethod + async def _close_guarded_layers(layers: Sequence[AsyncGenerator[object, None]]) -> None: + for layer in reversed(layers): + try: + await layer.aclose() + except Exception as e: # noqa: BLE001 # one failing callback cleanup must not skip the inner ones + verbose_proxy_logger.warning( + "Closing a streaming callback layer after a client disconnect raised %s", type(e).__name__ + ) + @staticmethod def _record_served_stream_output(request_data: Mapping[str, object], served_chunks: Sequence[object]) -> None: logging_obj: Final = request_data.get("litellm_logging_obj") diff --git a/litellm/rust_bridge/trace/errors.py b/litellm/rust_bridge/trace/errors.py new file mode 100644 index 00000000000..3c84d233be4 --- /dev/null +++ b/litellm/rust_bridge/trace/errors.py @@ -0,0 +1,5 @@ +class TraceChanged(Exception): + """The paging snapshot no longer matches the stored trace, so the client must start a new traversal. + + Raised by the Rust trace reader when a cursor's snapshot version differs from the graph it rebuilt. + """ diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index cd38b5b0d21..2d0985bcb64 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -6137,7 +6137,7 @@ "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, "input_cost_per_token": 4e-06, @@ -6148,7 +6148,7 @@ "mode": "realtime", "output_cost_per_audio_token": 6.4e-05, "output_cost_per_token": 2.4e-05, - "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule", "supported_endpoints": [ "/v1/realtime" ], @@ -6172,7 +6172,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-07-31", + "deprecation_date": "2027-06-25", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, @@ -6183,7 +6183,7 @@ "mode": "realtime", "output_cost_per_audio_token": 2e-05, "output_cost_per_token": 2.4e-06, - "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/model-retirement-schedule", "supported_endpoints": [ "/v1/realtime" ], @@ -42328,14 +42328,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 6e-09, - "input_cost_per_token": 3e-07, + "cache_read_input_token_cost": 3e-09, + "input_cost_per_token": 3e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.2e-06, + "output_cost_per_token": 2.4e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42348,14 +42348,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "cache_read_input_token_cost": 1.4e-07, - "input_cost_per_token": 2.2e-07, + "cache_read_input_token_cost": 7e-07, + "input_cost_per_token": 8.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 4.2e-06, + "output_cost_per_token": 5e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -67616,8 +67616,9 @@ "supports_web_search": false }, "openrouter/qwen/qwen3.8-27b": { - "input_cost_per_token": 4.2e-07, - "output_cost_per_token": 3e-06, + "cache_creation_input_token_cost": 5.3125e-07, + "input_cost_per_token": 4.25e-07, + "output_cost_per_token": 2.55e-06, "cache_read_input_token_cost": 8.5e-08, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -67732,7 +67733,7 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "cache_read_input_token_cost": 1.37e-08, + "cache_read_input_token_cost": 1.52e-08, "input_cost_per_token": 1.52e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, @@ -67823,14 +67824,14 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k3": { - "cache_read_input_token_cost": 4.9e-07, - "input_cost_per_token": 4.99e-07, + "cache_read_input_token_cost": 7e-07, + "input_cost_per_token": 7.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 1.3e-05, + "output_cost_per_token": 1.4e-05, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -67948,14 +67949,14 @@ "supports_web_search": true }, "openrouter/z-ai/glm-5.2": { - "cache_read_input_token_cost": 2.6e-07, - "input_cost_per_token": 3e-07, + "cache_read_input_token_cost": 1e-07, + "input_cost_per_token": 1.04e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 3.49e-06, + "output_cost_per_token": 8e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68311,14 +68312,14 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "cache_read_input_token_cost": 5.6e-09, - "input_cost_per_token": 2.8e-08, + "cache_read_input_token_cost": 2.24e-08, + "input_cost_per_token": 2.24e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 5.6e-08, + "output_cost_per_token": 1.28e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, diff --git a/scripts/lens_dev.sh b/scripts/lens_dev.sh index d8918f02fc5..1add480c227 100755 --- a/scripts/lens_dev.sh +++ b/scripts/lens_dev.sh @@ -7,7 +7,8 @@ # (default: random, generated once into .lens-dev/master_key) # LENS_DEV_CONFIG proxy config to use instead of the generated one # LENS_DEV_DATABASE_URL Postgres URL (default: the tracing stack's litellm DB on :15432) -# LENS_DEV_REBUILD_RUST=1 rebuild the Rust bridge even if it imports +# LENS_DEV_SEED trace seed profile (default|large), same as --seed +# LENS_DEV_SEED_LOGS request-log seed profile (default|large), same as --seed-logs # # State (master key, worker token, generated config, logs) lives in .lens-dev/ (gitignored). set -euo pipefail @@ -216,36 +217,55 @@ build_dashboard() { ) > "$log_dir/ui-build.log" 2>&1 || die "UI build failed; see $log_dir/ui-build.log" } +# Trace fixtures feed Lens (ClickHouse + spend rows); request logs feed the Logs page +# (Postgres only) with rows sized to stress the log detail drawer. seed_data() { ( proxy_env "" - "$py" -m scripts.seed_tracing_fixtures --profile "$seed_profile" ${seed_options[@]+"${seed_options[@]}"} + export LENS_DEV_UI_URL="http://localhost:$ui_port" + if [ -n "$seed_profile" ]; then + "$py" -m scripts.seed_tracing_fixtures --profile "$seed_profile" ${seed_options[@]+"${seed_options[@]}"} + fi + if [ -n "$seed_logs_profile" ]; then + "$py" -m scripts.seed_request_logs --profile "$seed_logs_profile" + fi ) } +# --seed and --seed-logs take an optional profile; a bare flag means default. +seed_profile_arg() { + if [ "${1:-}" = default ] || [ "${1:-}" = large ]; then echo "$1"; else echo default; fi +} + parse_args() { seed_profile="${LENS_DEV_SEED:-}" + seed_logs_profile="${LENS_DEV_SEED_LOGS:-}" seed_only=0 seed_options=() while [ "$#" -gt 0 ]; do case "$1" in --seed) - seed_profile=default - if [ "${2:-}" = default ] || [ "${2:-}" = large ]; then seed_profile="$2"; shift; fi + seed_profile="$(seed_profile_arg "${2:-}")" + [ "$seed_profile" = "${2:-}" ] && shift + ;; + --seed-logs) + seed_logs_profile="$(seed_profile_arg "${2:-}")" + [ "$seed_logs_profile" = "${2:-}" ] && shift ;; --copies) [ "$#" -ge 2 ] && [[ "$2" =~ ^[1-9][0-9]*$ ]] || die "--copies requires a positive integer" seed_options=(--copies "$2"); shift ;; --seed-only) seed_only=1 ;; --help) - echo "Usage: $0 [--seed [default|large]] [--copies N] [--seed-only]" + echo "Usage: $0 [--seed [default|large]] [--seed-logs [default|large]] [--copies N] [--seed-only]" exit 0 ;; *) die "unknown argument: $1 (use --help)" ;; esac shift done - if [ "$seed_only" = 1 ] && [ -z "$seed_profile" ]; then seed_profile=default; fi + if [ "$seed_only" = 1 ] && [ -z "$seed_profile" ] && [ -z "$seed_logs_profile" ]; then seed_profile=default; fi case "$seed_profile" in ""|default|large) ;; *) die "seed profile must be default or large" ;; esac + case "$seed_logs_profile" in ""|default|large) ;; *) die "seed-logs profile must be default or large" ;; esac [ "${#seed_options[@]}" = 0 ] || [ -n "$seed_profile" ] || die "--copies requires --seed" } @@ -277,11 +297,14 @@ main() { ensure_services "$py" scripts/prisma_generate_if_needed.py - if [ "${LENS_DEV_REBUILD_RUST:-0}" = "1" ] || ! "$py" -c "import litellm.rust_bridge._native" >/dev/null 2>&1; then - echo "lens-dev: building the Rust bridge (litellm.rust_bridge._native); the ClickHouse trace store uses it" - VIRTUAL_ENV="$repo_root/.venv" uvx --from maturin==1.15.0 maturin develop \ - --release --manifest-path litellm-rust/crates/python-bridge/Cargo.toml --features extension-module - fi + # cargo/maturin already fingerprint every crate's sources, so re-running this on each + # start is a no-op (a couple seconds) when nothing changed and only rebuilds the + # subset that did. An import check can't tell content-stale from content-fresh: a + # `.so` built from an older commit still imports fine, it just no longer matches + # what the current Python bindings (e.g. the trace store protocol) expect. + echo "lens-dev: checking the Rust bridge (litellm.rust_bridge._native) is current; the ClickHouse trace store uses it" + PYO3_PYTHON="$py" VIRTUAL_ENV="$repo_root/.venv" uvx --from maturin==1.15.0 maturin develop \ + --release --manifest-path litellm-rust/crates/python-bridge/Cargo.toml --features extension-module if [ ! -x ui/litellm-dashboard/node_modules/.bin/next ]; then (cd ui/litellm-dashboard && "$repo_root/scripts/with_dashboard_node.sh" npm ci) @@ -317,7 +340,7 @@ main() { wait_for_ui "$ui_pid" wait_for_proxy "$proxy_pid" ensure_worker_token - if [ -n "$seed_profile" ]; then seed_data; fi + if [ -n "$seed_profile" ] || [ -n "$seed_logs_profile" ]; then seed_data; fi LITELLM_RELEASE_TAG="$source_release_tag" \ LITELLM_MODE=PRODUCTION LITELLM_URL="$proxy_url" LENS_WORKER_TOKEN="$(cat "$token_file")" \ @@ -332,6 +355,7 @@ main() { Lens dev is up. Ctrl-C stops everything. Log in: http://localhost:$ui_port/ui/login/ (admin / $key_hint) Lens: http://localhost:$ui_port/ui/lens/ (hot-reloads) + Logs: http://localhost:$ui_port/ui/?page=logs API: $proxy_url Logs: $log_dir/proxy.log $log_dir/worker.log diff --git a/scripts/seed_request_logs.py b/scripts/seed_request_logs.py new file mode 100644 index 00000000000..1217a9fb7e4 --- /dev/null +++ b/scripts/seed_request_logs.py @@ -0,0 +1,427 @@ +from __future__ import annotations + +import argparse +import asyncio +import json +import math +import os +import random +import sys +from collections.abc import Iterator, Sequence +from datetime import datetime, timedelta, timezone +from itertools import chain +from typing import TYPE_CHECKING, Final, Literal + +import httpx +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter + +from scripts.seed_tracing_fixtures import JSON_OBJECT, spend_fixtures + +if TYPE_CHECKING: + from prisma.types import LiteLLM_SpendLogsCreateWithoutRelationsInput + +REQUEST_ID_PREFIX: Final = "seed-logs-" +SESSION_ID_PREFIX: Final = "seed-logs-session-" +WINDOW_HOURS: Final = 23 +RNG_SEED: Final = 20261004 +LARGE_COPIES: Final = 3000 +PROFILES: Final = ("default", "large") +Profile = Literal["default", "large"] +JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) + +WORDS: Final = ( + "trace", "span", "token", "request", "response", "latency", "router", "fallback", "cache", "budget", + "guardrail", "stream", "deployment", "proxy", "callback", "cursor", "schema", "payload", "retry", "quota", +) + + +class SeededLog(BaseModel): + """One synthetic spend-log row before it is shaped for Postgres.""" + + model_config = ConfigDict(frozen=True) + request_id: str + label: str + call_type: str + model: str + provider: str + status: Literal["success", "failure"] + session_id: str | None + offset_minutes: int + duration_ms: int + prompt_tokens: int + completion_tokens: int + spend: float + messages: JsonValue + response: JsonValue + proxy_server_request: JsonValue + error_information: dict[str, JsonValue] | None = None + + +def prose(rng: random.Random, chars: int) -> str: + words: Final[list[str]] = [] + length = 0 + while length < chars: + word: Final = rng.choice(WORDS) + words.append(word) + length += len(word) + 1 # rebind-ok: accumulates generated text length + return " ".join(words)[:chars] + + +def tool_definition(index: int) -> dict[str, JsonValue]: + return { + "type": "function", + "function": { + "name": f"seed_tool_{index}", + "description": f"Synthetic tool number {index} used only by the request-log seeder.", + "parameters": { + "type": "object", + "properties": { + "query": {"type": "string", "description": "What to look up"}, + "limit": {"type": "integer", "minimum": 1, "maximum": 100}, + }, + "required": ["query"], + }, + }, + } + + +def tool_call(index: int, rng: random.Random) -> dict[str, JsonValue]: + return { + "id": f"call_seed_{index}", + "type": "function", + "function": {"name": f"seed_tool_{index}", "arguments": json.dumps({"query": prose(rng, 40), "limit": index})}, + } + + +def chat_turns(rng: random.Random, turns: int, chars_per_turn: int) -> list[JsonValue]: + def turn(index: int) -> Iterator[JsonValue]: + yield {"role": "user", "content": prose(rng, chars_per_turn)} + if index % 3 == 0: + yield {"role": "assistant", "content": None, "tool_calls": [tool_call(index % 7, rng)]} + yield {"role": "tool", "tool_call_id": f"call_seed_{index % 7}", "content": prose(rng, chars_per_turn * 4)} + else: + yield {"role": "assistant", "content": prose(rng, chars_per_turn)} + + return list(chain.from_iterable(turn(index) for index in range(turns))) + + +def chat_response(content: str, tool_calls: list[JsonValue] | None, prompt_tokens: int, completion_tokens: int) -> JsonValue: + message: dict[str, JsonValue] = {"role": "assistant", "content": content} + if tool_calls: + message["tool_calls"] = tool_calls + return { + "id": "chatcmpl-seed", + "object": "chat.completion", + "model": "gpt-5.5", + "choices": [{"index": 0, "finish_reason": "tool_calls" if tool_calls else "stop", "message": message}], + "usage": { + "prompt_tokens": prompt_tokens, + "completion_tokens": completion_tokens, + "total_tokens": prompt_tokens + completion_tokens, + }, + } + + +def chat_log( + rng: random.Random, + label: str, + *, + messages: list[JsonValue], + response_chars: int, + tools: int = 0, + called_tools: int = 0, + offset_minutes: int, + session_id: str | None = None, +) -> SeededLog: + prompt_tokens: Final = len(json.dumps(messages)) // 4 + completion_tokens: Final = max(response_chars // 4, 1) + tool_calls: Final = [tool_call(index, rng) for index in range(called_tools)] or None + request: dict[str, JsonValue] = {"model": "gpt-5.5", "messages": messages, "stream": False} + if tools: + request["tools"] = [tool_definition(index) for index in range(tools)] + return SeededLog( + request_id=f"{REQUEST_ID_PREFIX}{label}", + label=label, + call_type="acompletion", + model="gpt-5.5", + provider="openai", + status="success", + session_id=session_id, + offset_minutes=offset_minutes, + duration_ms=1500 + completion_tokens // 10, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + spend=prompt_tokens * 0.000002 + completion_tokens * 0.000008, + messages=messages, + response=chat_response(prose(rng, response_chars), tool_calls, prompt_tokens, completion_tokens), + proxy_server_request=request, + ) + + +def anthropic_log(rng: random.Random, offset_minutes: int) -> SeededLog: + messages: Final[list[JsonValue]] = [ + {"role": "user", "content": [{"type": "text", "text": prose(rng, 2000)}]}, + { + "role": "assistant", + "content": [ + {"type": "text", "text": prose(rng, 500)}, + {"type": "tool_use", "id": "toolu_seed_1", "name": "seed_tool_1", "input": {"query": "spend"}}, + ], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_seed_1", "content": prose(rng, 20_000)}]}, + ] + tools: Final[list[JsonValue]] = [ + {"name": f"seed_tool_{index}", "description": "Synthetic Anthropic tool", "input_schema": {"type": "object"}} + for index in range(3) + ] + response: Final[JsonValue] = { + "id": "msg_seed", + "type": "message", + "role": "assistant", + "model": "claude-opus-5-5", + "content": [ + {"type": "text", "text": prose(rng, 50_000)}, + {"type": "tool_use", "id": "toolu_seed_2", "name": "seed_tool_2", "input": {"query": "latency", "limit": 5}}, + ], + "stop_reason": "tool_use", + "usage": {"input_tokens": 6000, "output_tokens": 12_500}, + } + return SeededLog( + request_id=f"{REQUEST_ID_PREFIX}anthropic-tool-use", + label="anthropic-tool-use", + call_type="anthropic_messages", + model="claude-opus-5-5", + provider="anthropic", + status="success", + session_id=None, + offset_minutes=offset_minutes, + duration_ms=9000, + prompt_tokens=6000, + completion_tokens=12_500, + spend=6000 * 0.000015 + 12_500 * 0.000075, + messages=messages, + response=response, + proxy_server_request={"model": "claude-opus-5-5", "max_tokens": 16_000, "messages": messages, "tools": tools}, + ) + + +def failure_log(rng: random.Random, offset_minutes: int) -> SeededLog: + messages: Final[list[JsonValue]] = [{"role": "user", "content": prose(rng, 300_000)}] + return SeededLog( + request_id=f"{REQUEST_ID_PREFIX}context-window-failure", + label="context-window-failure", + call_type="acompletion", + model="gpt-5.5", + provider="openai", + status="failure", + session_id=None, + offset_minutes=offset_minutes, + duration_ms=800, + prompt_tokens=75_000, + completion_tokens=0, + spend=0.0, + messages=messages, + response={}, + proxy_server_request={"model": "gpt-5.5", "messages": messages}, + error_information={ + "error_code": "400", + "error_class": "ContextWindowExceededError", + "llm_provider": "openai", + "error_message": "This model's maximum context length is 128000 tokens. Your messages resulted in 75000 tokens plus 300000 characters of synthetic prose.", + "traceback": "Traceback (most recent call last):\n" + "\n".join(f" File seed_{index}.py, line {index}" for index in range(40)), + }, + ) + + +def seeded_logs(rng: random.Random) -> tuple[SeededLog, ...]: + """The size ladder: one axis per thing that can make the log drawer slow.""" + session: Final = f"{SESSION_ID_PREFIX}agent-run" + single: Final = [{"role": "user", "content": "Summarise the seeded request logs in one paragraph."}] + return ( + chat_log(rng, "baseline-small", messages=single, response_chars=400, offset_minutes=5), + chat_log(rng, "response-100kb", messages=single, response_chars=100_000, offset_minutes=20), + chat_log(rng, "response-1mb", messages=single, response_chars=1_000_000, offset_minutes=35), + chat_log(rng, "response-5mb", messages=single, response_chars=5_000_000, offset_minutes=50), + chat_log(rng, "turns-200", messages=chat_turns(rng, 200, 500), response_chars=2000, offset_minutes=70), + chat_log(rng, "turns-1000", messages=chat_turns(rng, 1000, 500), response_chars=2000, offset_minutes=90), + chat_log(rng, "system-prompt-200kb", messages=[{"role": "system", "content": prose(rng, 200_000)}, *single], response_chars=1500, offset_minutes=110), + chat_log(rng, "tools-50", messages=single, response_chars=800, tools=50, called_tools=6, offset_minutes=130), + anthropic_log(rng, offset_minutes=150), + failure_log(rng, offset_minutes=170), + *( + chat_log( + rng, + f"session-call-{index:02d}", + messages=chat_turns(rng, index + 1, 400), + response_chars=3000, + tools=4, + called_tools=index % 3, + offset_minutes=200 + index, + session_id=session, + ) + for index in range(30) + ), + ) + + +def spread_offsets(logs: tuple[SeededLog, ...]) -> tuple[SeededLog, ...]: + """Fit every row into the page's default 24h window, newest first.""" + last: Final = max(log.offset_minutes for log in logs) + scale: Final = min(1.0, WINDOW_HOURS * 60 / max(last, 1)) + return tuple(log.model_copy(update={"offset_minutes": int(log.offset_minutes * scale)}) for log in logs) + + +def metadata(log: SeededLog, template: dict[str, JsonValue]) -> dict[str, JsonValue]: + usage: Final[dict[str, JsonValue]] = { + "prompt_tokens": log.prompt_tokens, + "completion_tokens": log.completion_tokens, + "total_tokens": log.prompt_tokens + log.completion_tokens, + "prompt_tokens_details": {"cached_tokens": 0, "text_tokens": log.prompt_tokens}, + } + seeded: dict[str, JsonValue] = { + **template, + "status": log.status, + "model_group": log.model, + "deployment": f"{log.provider}/{log.model}", + "deployment_model_name": f"{log.provider}/{log.model}", + "user_api_key_team_alias": "seed-logs", + "usage_object": usage, + "additional_usage_values": {"cache_read_input_tokens": 0, "cache_creation_input_tokens": 0, **usage}, + "cost_breakdown": { + "input_cost": log.prompt_tokens * 0.000002, + "output_cost": log.completion_tokens * 0.000008, + "total_cost": log.spend, + }, + "litellm_overhead_time_ms": 12.5, + "attempted_retries": 0, + "max_retries": 2, + "hidden_params": {"litellm_overhead_time_ms": 12.5, "response_cost": log.spend}, + "fixture_capture": None, + "seed_label": log.label, + } + if log.error_information is not None: + seeded["error_information"] = log.error_information + return seeded + + +def postgres_row(log: SeededLog, template: dict[str, JsonValue], now: datetime) -> LiteLLM_SpendLogsCreateWithoutRelationsInput: + from prisma import Json + from prisma.types import LiteLLM_SpendLogsCreateWithoutRelationsInput + + end: Final = now - timedelta(minutes=log.offset_minutes) + start: Final = end - timedelta(milliseconds=log.duration_ms) + return LiteLLM_SpendLogsCreateWithoutRelationsInput( + request_id=log.request_id, + litellm_call_id=log.request_id, + call_type=log.call_type, + api_key=str(template.get("user_api_key", "seed-logs-key")), + user="seed-logs-user", + team_id="seed-logs-team", + spend=log.spend, + model=log.model, + model_id=f"seed-logs-{log.model}", + model_group=log.model, + custom_llm_provider=log.provider, + api_base=f"https://api.{log.provider}.example", + prompt_tokens=log.prompt_tokens, + completion_tokens=log.completion_tokens, + total_tokens=log.prompt_tokens + log.completion_tokens, + startTime=start, + endTime=end, + completionStartTime=start + timedelta(milliseconds=min(400, log.duration_ms // 2)), + request_duration_ms=log.duration_ms, + session_id=log.session_id, + status=log.status, + cache_hit="False", + request_tags=Json(["seed-logs", log.label.split("-")[0]]), + metadata=Json(metadata(log, template)), + messages=Json(log.messages), + response=Json(log.response), + proxy_server_request=Json(log.proxy_server_request), + ) + + +COPY_SQL: Final = """INSERT INTO "LiteLLM_SpendLogs" +SELECT (jsonb_populate_record(s, jsonb_build_object( + 'request_id', s.request_id || '-copy-' || c.n, + 'litellm_call_id', s.request_id || '-copy-' || c.n, + 'session_id', NULL, + 'startTime', s."startTime" - make_interval(secs => c.n * $3::bigint / 1000.0), + 'endTime', s."endTime" - make_interval(secs => c.n * $3::bigint / 1000.0), + 'completionStartTime', s."completionStartTime" - make_interval(secs => c.n * $3::bigint / 1000.0) +))).* +FROM "LiteLLM_SpendLogs" AS s CROSS JOIN generate_series(1, $2::int) AS c(n) +WHERE s.request_id = $1""" + + +class SeedOptions(BaseModel): + model_config = ConfigDict(frozen=True) + profile: Profile + timeout_seconds: float = 120 + + +def seed_arguments(argv: Sequence[str] | None = None) -> SeedOptions: + parser: Final = argparse.ArgumentParser(description="Insert synthetic request logs of controlled sizes into a local proxy DB") + parser.add_argument("--profile", choices=PROFILES, default="default") + parser.add_argument("--timeout-seconds", type=float, default=os.environ.get("LENS_DEV_SEED_TIMEOUT_SECONDS", "120")) + arguments: Final = SeedOptions.model_validate(vars(parser.parse_args(argv))) + if not math.isfinite(arguments.timeout_seconds) or arguments.timeout_seconds <= 0: + parser.error("--timeout-seconds must be finite and positive") + return arguments + + +def metadata_template() -> dict[str, JsonValue]: + """A real captured row's metadata, so the drawer sees the keys the gateway writes.""" + _, rows = spend_fixtures()[0] + return JSON_OBJECT.validate_json(rows[0]["metadata"]) + + +async def verify(client: httpx.AsyncClient, logs: tuple[SeededLog, ...], ui_base: str) -> None: + async def fetch(log: SeededLog) -> dict[str, JsonValue]: + detail: Final = await client.get(f"/spend/logs/ui/{log.request_id}") + detail.raise_for_status() + payload: Final = JSON.validate_json(detail.content) + return { + "label": log.label, + "bytes": len(detail.content), + "found": isinstance(payload, dict) and bool(payload), + "url": f"{ui_base}/ui/?page=logs&log_id={log.request_id}" + + (f"&session_id={log.session_id}" if log.session_id else ""), + } + + results: Final = tuple(await asyncio.gather(*(fetch(log) for log in logs))) + sys.stdout.write(json.dumps(list(results), indent=2) + "\n") + if not all(result["found"] for result in results): + raise RuntimeError("Seeded request logs did not round-trip through /spend/logs/ui/{request_id}") + + +async def seed(profile: Profile = "default", timeout_seconds: float = 120) -> int: + from prisma import Prisma + + logs: Final = spread_offsets(seeded_logs(random.Random(RNG_SEED))) + template: Final = metadata_template() + now: Final = datetime.now(timezone.utc) + proxy_url: Final = os.environ.get("PROXY_BASE_URL", "http://127.0.0.1:4000") + ui_base: Final = os.environ.get("LENS_DEV_UI_URL", proxy_url) + async with ( + httpx.AsyncClient( + base_url=proxy_url, + headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, + timeout=timeout_seconds, + ) as client, + Prisma(http={"timeout": httpx.Timeout(600)}) as database, + ): + await database.litellm_spendlogs.delete_many(where={"request_id": {"startswith": REQUEST_ID_PREFIX}}) + await database.litellm_spendlogs.create_many(data=[postgres_row(log, template, now) for log in logs]) + if profile == "large": + step_ms: Final = WINDOW_HOURS * 60 * 60 * 1000 // LARGE_COPIES + await database.execute_raw(COPY_SQL, f"{REQUEST_ID_PREFIX}baseline-small", LARGE_COPIES, step_ms) + await verify(client, tuple(log for log in logs if not log.session_id or log.label.endswith("-00")), ui_base) + total: Final = len(logs) + (LARGE_COPIES if profile == "large" else 0) + sys.stdout.write(f"Request log seed complete: profile={profile}, rows={total}, prefix={REQUEST_ID_PREFIX}\n") + return 0 + + +if __name__ == "__main__": + arguments: Final = seed_arguments() + raise SystemExit(asyncio.run(seed(arguments.profile, arguments.timeout_seconds))) diff --git a/scripts/seed_tracing_fixtures.py b/scripts/seed_tracing_fixtures.py index 3febe49c451..be11f7cc90b 100644 --- a/scripts/seed_tracing_fixtures.py +++ b/scripts/seed_tracing_fixtures.py @@ -43,6 +43,9 @@ TRACE: Final = TypeAdapter(Trace) NANOSECOND_FIELDS: Final = frozenset({"startTimeUnixNano", "endTimeUnixNano", "timeUnixNano"}) TRACE_ID_FIELDS: Final = frozenset({"traceId", "trace_id", "session_id"}) SPAN_ID_FIELDS: Final = frozenset({"spanId", "parentSpanId", "span_id"}) +COPY_WINDOW_MS: Final = 24 * 60 * 60 * 1000 +LONG_SESSION_SOURCE: Final = "openai_agents_swarm" +LONG_SESSION_REPEATS: Final = (50, 400, 4000) class TenantIdentity(BaseModel): @@ -245,7 +248,6 @@ class SeedOptions(BaseModel): model_config = ConfigDict(frozen=True) profile: Literal["default", "large"] copies: int | None - batch_copies: int = 4 timeout_seconds: float = 120 @@ -258,33 +260,23 @@ def seed_arguments(argv: Sequence[str] | None = None) -> SeedOptions: default=os.environ.get("LENS_DEV_SEED_COPIES"), help="Override fixture copies (default: 1, large: 2000; env: LENS_DEV_SEED_COPIES)", ) - parser.add_argument( - "--batch-copies", - type=int, - default=os.environ.get("LENS_DEV_SEED_BATCH_COPIES", "4"), - help="Copies per bulk insert (default: 4; env: LENS_DEV_SEED_BATCH_COPIES)", - ) parser.add_argument("--timeout-seconds", type=float, default=os.environ.get("LENS_DEV_SEED_TIMEOUT_SECONDS", "120")) arguments: Final = SeedOptions.model_validate(vars(parser.parse_args(argv))) if arguments.copies is not None and arguments.copies < 1: parser.error("--copies must be positive") - if arguments.batch_copies < 1: - parser.error("--batch-copies must be positive") if not math.isfinite(arguments.timeout_seconds) or arguments.timeout_seconds <= 0: parser.error("--timeout-seconds must be finite and positive") return arguments -async def seed_batch( +async def seed_copy( client: httpx.AsyncClient, storage: ClickHouseStorage, database: Prisma, replays: tuple[FixtureReplay, ...], fixtures: tuple[tuple[str, tuple[SpendLogRecord, ...]], ...], pattern: re.Pattern[str], - tenant: TenantIdentity | None, - verify: bool, -) -> TenantIdentity: +) -> tuple[tuple[str, tuple[SpendLogRecord, ...]], ...]: by_name: Final = MappingProxyType(dict(fixtures)) paired: Final = tuple( ( @@ -294,36 +286,36 @@ async def seed_batch( for replay in replays if replay.name in by_name ) - rebased_spends: Final = tuple(chain.from_iterable(rows for _, rows in paired)) - resolved_tenant: Final = await ingest_replays( - client, storage, replays, fixture_capture(*next((name, rows[0]) for name, rows in paired)).trace_id, tenant - ) - stamped_spends: Final[tuple[SpendLogRecord, ...]] = tuple( - {**row, "team_id": resolved_tenant.team_id, "api_key": resolved_tenant.api_key, "user": resolved_tenant.user} - for row in rebased_spends + tenant: Final = await ingest_replays( + client, storage, replays, fixture_capture(*next((name, rows[0]) for name, rows in paired)).trace_id ) + stamped: Final = tuple((name, tuple(stamp(row, tenant) for row in rows)) for name, rows in paired) + stamped_spends: Final = tuple(chain.from_iterable(rows for _, rows in stamped)) await storage.insert_rows("spend_logs", stamped_spends) await database.litellm_spendlogs.create_many(data=[postgres_row(row) for row in stamped_spends]) - if verify: - verified: Final = tuple(await asyncio.gather(*(verify_capture(client, name, rows) for name, rows in paired))) - sys.stdout.write(json.dumps({"spend_rows": len(stamped_spends), "captures": verified}, indent=2) + "\n") - if not all(capture["verified"] for capture in verified): - raise RuntimeError("Seed spend verification failed") - return resolved_tenant + return stamped + + +def stamp(row: SpendLogRecord, tenant: TenantIdentity) -> SpendLogRecord: + return {**row, "team_id": tenant.team_id, "api_key": tenant.api_key, "user": tenant.user} + + +async def verify( + client: httpx.AsyncClient, captures: tuple[tuple[str, tuple[SpendLogRecord, ...]], ...], trace_salt: str +) -> None: + verified: Final = tuple( + await asyncio.gather(*(verify_capture(client, name, rows, trace_salt) for name, rows in captures)) + ) + sys.stdout.write( + json.dumps({"spend_rows": sum(len(rows) for _, rows in captures), "captures": verified}, indent=2) + "\n" + ) + if not all(capture["verified"] for capture in verified): + raise RuntimeError("Seed spend verification failed") async def ingest_replays( - client: httpx.AsyncClient, - storage: ClickHouseStorage, - replays: tuple[FixtureReplay, ...], - trace_id: str, - tenant: TenantIdentity | None, + client: httpx.AsyncClient, storage: ClickHouseStorage, replays: tuple[FixtureReplay, ...], trace_id: str ) -> TenantIdentity: - if tenant is not None: - await storage.insert_rows( - "otel_traces", bulk_span_rows(replays, Tenant(tenant.team_id, tenant.api_key, user_id=tenant.user)) - ) - return tenant for replay in replays: ( await client.post( @@ -347,24 +339,161 @@ def bulk_span_rows(replays: tuple[FixtureReplay, ...], tenant: Tenant) -> tuple[ ) -def replay_batches( - count: int, now_ms: int, namespace: str, pattern: re.Pattern[str], batch_copies: int = 4 -) -> Iterator[tuple[int, tuple[FixtureReplay, ...]]]: - for start, stop in ((start, min(start + batch_copies, count)) for start in range(1, count, batch_copies)): - yield ( - stop, - tuple( - chain.from_iterable( - fixture_replays(TRACE_FIXTURES, now_ms - index * 1000, f"{namespace}-{index}", pattern) - for index in range(start, stop) - ) - ), +@dataclass(frozen=True, slots=True) +class Copies: + """Server-side copies of seeded traces and their spend. + + Each copy `n` hashes trace ids with `session` (or `n` when empty), keeps root span ids, hashes the + other span ids with `n`, rewrites seeded call ids from `source` to `{target}{n}-` and moves `n * step_ms` + earlier. A `session` folds every copy into one trace under a single root that spans all of them. + """ + + trace_ids: tuple[str, ...] + request_ids: tuple[str, ...] + numbers: range + step_ms: int + source: str + target: str + session: str = "" + + +def copied_trace_id(trace_id: str, salt: str) -> str: + return hashlib.sha256(f"{trace_id}:{salt}".encode()).hexdigest()[:32] + + +def clickhouse_call_id(column: str) -> str: + target: Final = "concat({target:String}, toString(c.n), '-')" + return ( + f"if(startsWith({column}, 'resp_'), concat('resp_', base64Encode(replaceAll(" + f"tryBase64Decode(substring({column}, 6)), {{source:String}}, {target}))), " + f"replaceAll({column}, {{source:String}}, {target}))" + ) + + +def clickhouse_hash(column: str, salt: str, length: int) -> str: + return f"if({column} = '', '', substring(lower(hex(SHA256(concat({column}, ':', {salt})))), 1, {length}))" + + +def clickhouse_copy_sql(database: str) -> tuple[str, str]: + trace_salt: Final = "if({session:String} = '', toString(c.n), {session:String})" + roots: Final = ( + f"(SELECT SpanId FROM {database}.otel_traces " + "WHERE TraceId IN {trace_ids:Array(String)} AND ParentSpanId = '')" + ) + folded_root: Final = "{session:String} != '' AND t.ParentSpanId = ''" + shift: Final = f"toIntervalMillisecond(if({folded_root}, {{last:UInt64}}, c.n) * {{step_ms:UInt64}})" + numbers: Final = "CROSS JOIN (SELECT number AS n FROM numbers({first:UInt64}, {count:UInt64})) AS c" + spans: Final = f"""INSERT INTO {database}.otel_traces +SELECT t.* REPLACE ( + t.Timestamp - {shift} AS Timestamp, + {clickhouse_hash("t.TraceId", trace_salt, 32)} AS TraceId, + if(t.ParentSpanId = '', t.SpanId, {clickhouse_hash("t.SpanId", "toString(c.n)", 16)}) AS SpanId, + if(t.ParentSpanId IN {roots}, t.ParentSpanId, {clickhouse_hash("t.ParentSpanId", "toString(c.n)", 16)}) + AS ParentSpanId, + t.Duration + if({folded_root}, {{last:UInt64}} * {{step_ms:UInt64}} * 1000000, 0) AS Duration, + arrayMap(at -> at - {shift}, t.`Events.Timestamp`) AS `Events.Timestamp`, + mapApply((name, value) -> (name, {clickhouse_call_id("value")}), t.SpanAttributes) AS SpanAttributes, + {clickhouse_call_id("t.LiteLLMRequestId")} AS LiteLLMRequestId, + arrayMap(key -> concat(extract(key, '^[^:]*:'), {clickhouse_call_id("replaceRegexpOne(key, '^[^:]*:', '')")}), + t.CallKeys) AS CallKeys +) +FROM {database}.otel_traces AS t {numbers} +WHERE t.TraceId IN {{trace_ids:Array(String)}} + AND ({{session:String}} = '' OR t.ParentSpanId != '' OR c.n = {{first:UInt64}})""" + spend: Final = f"""INSERT INTO {database}.spend_logs +SELECT s.* REPLACE ( + {clickhouse_call_id("s.request_id")} AS request_id, + {clickhouse_call_id("s.response_id")} AS response_id, + {clickhouse_call_id("s.litellm_call_id")} AS litellm_call_id, + {clickhouse_hash("s.trace_id", trace_salt, 32)} AS trace_id, + {clickhouse_hash("s.session_id", trace_salt, 32)} AS session_id, + if(s.span_id IN {roots}, s.span_id, {clickhouse_hash("s.span_id", "toString(c.n)", 16)}) AS span_id, + s.start_time - toIntervalMillisecond(c.n * {{step_ms:UInt64}}) AS start_time, + s.end_time - toIntervalMillisecond(c.n * {{step_ms:UInt64}}) AS end_time, + s.completion_start_time - toIntervalMillisecond(c.n * {{step_ms:UInt64}}) AS completion_start_time +) +FROM {database}.spend_logs AS s {numbers} +WHERE s.request_id IN {{request_ids:Array(String)}}""" + return spans, spend + + +def clickhouse_array(values: tuple[str, ...]) -> str: + return "[" + ",".join("'" + value.replace("\\", "\\\\").replace("'", "\\'") + "'" for value in values) + "]" + + +async def copy_clickhouse(client: httpx.AsyncClient, database: str, copies: Copies) -> None: + parameters: Final = { + "param_trace_ids": clickhouse_array(copies.trace_ids), + "param_request_ids": clickhouse_array(copies.request_ids), + "param_first": str(copies.numbers.start), + "param_count": str(len(copies.numbers)), + "param_last": str(copies.numbers.stop - 1), + "param_step_ms": str(copies.step_ms), + "param_source": copies.source, + "param_target": copies.target, + "param_session": copies.session, + } + for sql in clickhouse_copy_sql(database): + (await client.post("/", params=parameters, content=sql)).raise_for_status() + + +POSTGRES_COPY_TARGET: Final = "($2 || c.n || '-')" +POSTGRES_COPY_SHIFT: Final = "make_interval(secs => c.n * $6::bigint / 1000.0)" +POSTGRES_COPY_SQL: Final = f"""INSERT INTO "LiteLLM_SpendLogs" +SELECT (jsonb_populate_record(s, jsonb_build_object( + 'request_id', CASE WHEN left(s.request_id, 5) = 'resp_' + THEN 'resp_' || translate(encode(convert_to(replace(convert_from(decode(substr(s.request_id, 6), 'base64'), + 'UTF8'), $1, {POSTGRES_COPY_TARGET}), 'UTF8'), 'base64'), E'\\n', '') + ELSE replace(s.request_id, $1, {POSTGRES_COPY_TARGET}) END, + 'session_id', CASE WHEN coalesce(s.session_id, '') = '' THEN s.session_id + ELSE substr(encode(sha256(convert_to( + s.session_id || ':' || CASE WHEN $3 = '' THEN c.n::text ELSE $3 END, 'UTF8')), 'hex'), 1, 32) END, + 'startTime', s."startTime" - {POSTGRES_COPY_SHIFT}, + 'endTime', s."endTime" - {POSTGRES_COPY_SHIFT}, + 'completionStartTime', s."completionStartTime" - {POSTGRES_COPY_SHIFT} +))).* +FROM "LiteLLM_SpendLogs" AS s CROSS JOIN generate_series($4::int, $5::int) AS c(n) +WHERE s.request_id = ANY(string_to_array($7, E'\\n'))""" + + +async def copy_postgres(database: Prisma, copies: Copies) -> None: + await database.execute_raw( + POSTGRES_COPY_SQL, + copies.source, + copies.target, + copies.session, + copies.numbers.start, + copies.numbers.stop - 1, + copies.step_ms, + "\n".join(copies.request_ids), + ) + + +def long_sessions( + replays: tuple[FixtureReplay, ...], + captures: tuple[tuple[str, tuple[SpendLogRecord, ...]], ...], + source: str, + target: str, + repeats: tuple[int, ...] = LONG_SESSION_REPEATS, +) -> tuple[Copies, ...]: + replay: Final = next(replay for replay in replays if replay.name == LONG_SESSION_SOURCE) + rows: Final = dict(captures)[LONG_SESSION_SOURCE] + span_ns: Final = tuple(timestamps(replay.export)) + return tuple( + Copies( + trace_ids=(fixture_capture(LONG_SESSION_SOURCE, rows[0]).trace_id,), + request_ids=tuple(row["request_id"] for row in rows), + numbers=range(count), + step_ms=(max(span_ns) - min(span_ns)) // 1_000_000 + 1000, + source=source, + target=f"{target}s{count}x", + session=f"session{count}", ) + for count in repeats + ) -async def seed( - profile: str = "default", copies: int | None = None, batch_copies: int = 4, timeout_seconds: float = 120 -) -> int: +async def seed(profile: str = "default", copies: int | None = None, timeout_seconds: float = 120) -> int: from prisma import Prisma fixtures: Final = spend_fixtures() @@ -372,29 +501,43 @@ async def seed( count: Final = copies if copies is not None else (2000 if profile == "large" else 1) namespace: Final = uuid4().hex now_ms: Final = time.time_ns() // 1_000_000 - storage: Final = ClickHouseStorage(trace_storage_config({})) + config: Final = trace_storage_config({}) + storage: Final = ClickHouseStorage(config) + replays: Final = fixture_replays(TRACE_FIXTURES, now_ms, namespace + "-0", pattern) + source: Final = f"seed-{namespace}-0-" + target: Final = f"seed-{namespace}-" async with ( httpx.AsyncClient( base_url=os.environ.get("PROXY_BASE_URL", "http://127.0.0.1:4002"), headers={"Authorization": f"Bearer {os.environ['LITELLM_MASTER_KEY']}"}, timeout=timeout_seconds, ) as client, - Prisma() as database, + httpx.AsyncClient(base_url=config.url, params={"database": config.database}, timeout=600) as clickhouse, + Prisma(http={"timeout": httpx.Timeout(600)}) as database, ): - tenant: Final = await seed_batch( - client, - storage, - database, - fixture_replays(TRACE_FIXTURES, now_ms, namespace + "-0", pattern), - fixtures, - pattern, - None, - True, + captures: Final = await seed_copy(client, storage, database, replays, fixtures, pattern) + await verify(client, captures, "") + repeated: Final = Copies( + trace_ids=tuple( + sorted(frozenset(str(span["TraceId"]) for span in bulk_span_rows(replays, Tenant("", "")))) + ), + request_ids=tuple(row["request_id"] for _, rows in captures for row in rows), + numbers=range(1, count), + step_ms=COPY_WINDOW_MS // count, + source=source, + target=target, ) - for stop, replays in replay_batches(count, now_ms, namespace, pattern, batch_copies): - await seed_batch(client, storage, database, replays, fixtures, pattern, tenant, stop == count) - sys.stdout.write(f"Seeded {stop}/{count} fixture copies\n") - sys.stdout.flush() + sessions: Final = long_sessions(replays, captures, source, target) if profile == "large" else () + for plan in (repeated, *sessions) if count > 1 else sessions: + await copy_clickhouse(clickhouse, config.database, plan) + await copy_postgres(database, plan) + if count > 1: + await verify(client, captures, str(count - 1)) + for plan in sessions: + sys.stdout.write( + f"Long session: {len(plan.numbers)} repeats, " + f"trace_id={copied_trace_id(plan.trace_ids[0], plan.session)}\n" + ) sys.stdout.write(f"Seed complete: profile={profile}, copies={count}, namespace={namespace}\n") return 0 @@ -410,17 +553,18 @@ def fixture_capture(name: str, row: SpendLogRecord) -> FixtureCapture: async def verify_capture( - client: httpx.AsyncClient, name: str, rows: tuple[SpendLogRecord, ...] + client: httpx.AsyncClient, name: str, rows: tuple[SpendLogRecord, ...], trace_salt: str = "" ) -> Mapping[str, JsonValue]: capture: Final = fixture_capture(name, rows[0]) - detail: Final = await client.get(f"/v1/traces/{capture.trace_id}") + trace_id: Final = copied_trace_id(capture.trace_id, trace_salt) if trace_salt else capture.trace_id + detail: Final = await client.get(f"/v1/traces/{trace_id}") detail.raise_for_status() trace: Final = TRACE.validate_json(detail.content) expected: Final = sum(row["spend"] or 0 for row in rows) actual: Final = trace["summary"]["spend"] return { "fixture": name, - "trace_id": capture.trace_id, + "trace_id": trace_id, "spend_rows": len(rows), "recorded_spend": expected, "trace_spend": actual, @@ -432,6 +576,4 @@ async def verify_capture( if __name__ == "__main__": arguments: Final = seed_arguments() - raise SystemExit( - asyncio.run(seed(arguments.profile, arguments.copies, arguments.batch_copies, arguments.timeout_seconds)) - ) + raise SystemExit(asyncio.run(seed(arguments.profile, arguments.copies, arguments.timeout_seconds))) diff --git a/tests/e2e/coverage_registry/logging.yaml b/tests/e2e/coverage_registry/logging.yaml index 5424fb2d12f..9a851587fce 100644 --- a/tests/e2e/coverage_registry/logging.yaml +++ b/tests/e2e/coverage_registry/logging.yaml @@ -6,6 +6,7 @@ - {id: logging.datadog.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses, embeddings], source: "integrations/datadog/datadog.py", rationale: "Powers dashboards/alerts; cardinality regressions common"} - {id: logging.datadog.stream.exports_metric, module: logging, tier: P0, event: stream, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses], source: "integrations/datadog/datadog.py", rationale: "Streaming aggregates usage after the last chunk; delivery and cost must survive that path"} - {id: logging.datadog.failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions], source: "integrations/datadog/datadog.py", rationale: "Failure metrics for alerting/SLO"} +- {id: logging.datadog.stream_failure.exports_metric, module: logging, tier: P0, event: failure, assertions: [exports_metric], exercised_on: [chat_completions], source: "litellm_core_utils/litellm_logging.py", rationale: "Streamed failures re-enter the failure handler once per retry on the same logging object; dedup must hold on the stream path too (#42988)"} - {id: logging.prometheus.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, embeddings], source: "integrations/prometheus.py", rationale: "Standard OSS metrics; per-key cardinality (existing e2e)"} - {id: logging.prometheus.success.records_queue_time, module: logging, tier: P1, event: success, assertions: [records_queue_time], exercised_on: [chat_completions], source: "integrations/prometheus.py / LIT-2034", fail_before_fix: proven, rationale: "Queue time feeds saturation alerting; the family stayed registered while no observation was ever recorded, so presence alone is not the contract"} - {id: logging.otel.success.exports_metric, module: logging, tier: P0, event: success, assertions: [exports_metric], exercised_on: [chat_completions, messages, responses, embeddings], source: "integrations/otel/logger.py", rationale: "OTEL spans on every call path"} diff --git a/tests/e2e/logging/test_datadog_log_e2e.py b/tests/e2e/logging/test_datadog_log_e2e.py index a4821ed058b..2d57181c4bf 100644 --- a/tests/e2e/logging/test_datadog_log_e2e.py +++ b/tests/e2e/logging/test_datadog_log_e2e.py @@ -22,13 +22,12 @@ import math import time import pytest -from pydantic import BaseModel, ConfigDict - from datadog_reader import DdLogEvent, DdLogsReader from e2e_config import CHEAP_ANTHROPIC_MODEL, CHEAP_OPENAI_MODEL, unique_marker from lifecycle import ResourceManager from logging_client import INVALID_UPSTREAM_API_KEY, LoggingClient, first_ok, readiness_details_body -from models import LiteLLMParamsBody +from models import ChatMessage, LiteLLMParamsBody, ReliabilityChatBody, RouterSettingsOverride +from pydantic import BaseModel, ConfigDict pytestmark = pytest.mark.e2e @@ -398,3 +397,67 @@ class TestDataDogFailureDelivery: assert payload.error_str is not None and "AnthropicException" in payload.error_str, ( f"the event must carry the provider error, got error_str={payload.error_str!r}" ) + + @pytest.mark.covers("logging.datadog.stream_failure.exports_metric", exercised_on=["chat_completions"]) + def test_failed_chat_completions_stream_emits_one_error_event( + self, client: LoggingClient, dd_logs: DdLogsReader, resources: ResourceManager + ) -> None: + """A STREAMED /chat/completions call that fails at the provider after + its configured retries must still reach the DataDog logs intake as + exactly one error-grade event: the retry loop invokes the failure + handler once per attempt on the same logging object, so a dedup that + only works for non-streaming calls multiplies every retried stream + failure by its attempt count (issue #42988). + + The deployment fails with a connect error, not an auth error: the + router does not retry AuthenticationError when the model group has a + single deployment, so an invalid key would never reach the retry loop + this test exercises. api_base is an unroutable address, so every + attempt fails the same retryable way.""" + _assert_datadog_configured(client) + + model_name = f"dd-err-stream-{unique_marker()}" + model_id = client.create_model( + model_name, + LiteLLMParamsBody( + model="anthropic/claude-haiku-4-5", + api_key=INVALID_UPSTREAM_API_KEY, + api_base="http://localhost:1", + ), + ) + resources.defer(lambda: client.delete_model(model_id)) + key = client.key_with_alias(f"dd-err-stream-key-{unique_marker()}", models=[model_name]) + resources.defer(lambda: client.delete_key(key)) + + deadline = time.monotonic() + client.proxy.poll_timeout + while True: + outcome = client.proxy.transport.stream( + "/chat/completions", + headers=client.proxy.transport.bearer(key), + json=ReliabilityChatBody( + model=model_name, + messages=[ChatMessage(role="user", content="trigger an upstream connect failure")], + stream=True, + max_tokens=16, + router_settings_override=RouterSettingsOverride(num_retries=2), + ), + ) + assert not outcome.ok, "the call must fail; the deployment's upstream is unreachable" + assert outcome.status_code != -1, ( + "network failure between the test and the proxy while provoking the provider " + "failure; retrying now could double-log the failure payload and falsely trip " + f"the exactly-one assertion - fix the rig connectivity first: {outcome.body[:200]}" + ) + if "AnthropicException" in outcome.body or time.monotonic() >= deadline: + break + time.sleep(client.proxy.poll_interval) + assert "AnthropicException" in outcome.body, ( + "never saw the upstream provider failure before the deadline; the deployment may still " + f"be propagating - last outcome {outcome.status_code}: {outcome.body[:200]}" + ) + + events = dd_logs.poll_events_for_query(f"@model_group:{model_name}") + payload = _assert_exactly_one_failure_event(events, model_group=model_name) + assert payload.error_str is not None and "AnthropicException" in payload.error_str, ( + f"the event must carry the provider error, got error_str={payload.error_str!r}" + ) diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 3d76c7881e1..3fa7f4b0333 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -230,6 +230,60 @@ def owned_proxy_process( _stop(process) +def _is_ready(client: httpx.Client) -> bool: + try: + return client.get("/health/readiness", timeout=2).status_code == 200 + except httpx.TransportError: + return False + + +def refused_boot_log( + gateway: Gateway, + directory: Path, + overrides: Mapping[str, str], + *, + config: Path | None = None, +) -> str: + """Start the proxy and return its log once it exits non-zero instead of becoming ready.""" + root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) + environment: Final = { + **os.environ, + **proxy_database_environment(), + "LITELLM_MASTER_KEY": gateway.key, + "LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"), + "STORE_MODEL_IN_DB": "True", + **overrides, + } + output: Final = Path(os.environ.get("INTEGRATION_RESULTS_DIR", str(directory))) + output.mkdir(parents=True, exist_ok=True) + command: Final = ( + sys.executable, + "-m", + "integration._support.proxy", + "--config", + str(config or "tests/integration/proxy_config.yaml"), + "--host", + "127.0.0.1", + "--num_workers", + "1", + *DB_PUSH, + ) + launch: Final = _launch(command, root, environment, output) + try: + with httpx.Client(base_url=f"http://127.0.0.1:{launch.port}", timeout=15, trust_env=False) as client: + deadline: Final = time.monotonic() + 70 + while launch.process.poll() is None: + assert not _is_ready(client), ( + f"Proxy became ready instead of refusing to boot:\n{launch.log.read_text()}" + ) + assert time.monotonic() < deadline, "Proxy neither exited nor became ready within the deadline" + time.sleep(0.1) + assert launch.process.returncode != 0, f"Proxy exited 0 instead of refusing to boot:\n{launch.log.read_text()}" + return launch.log.read_text() + finally: + _stop(launch.process) + + _UPSTREAM_READY_SECONDS: Final = 60 diff --git a/tests/integration/configuration/test_fips_mode_boot.py b/tests/integration/configuration/test_fips_mode_boot.py new file mode 100644 index 00000000000..80904c600e1 --- /dev/null +++ b/tests/integration/configuration/test_fips_mode_boot.py @@ -0,0 +1,62 @@ +"""LITELLM_FIPS_MODE is a boot gate: the proxy refuses to serve unless the process really enforces FIPS. + +Every leg launches the real proxy binary against the suite's Postgres and asserts on what an operator sees: +exit status and the refusal text in the log. Nothing is patched inside the proxy. +""" + +import hashlib +from pathlib import Path +from typing import Final + +import pytest +import yaml + +from tests.integration._support.client import Gateway +from tests.integration._support.process import owned_proxy, refused_boot_log + +REFUSAL: Final = "LiteLLM proxy refused to start" + + +def _this_python_enforces_fips() -> bool: + try: + hashlib.md5(b"probe", usedforsecurity=True) + except ValueError: + return True + return False + + +def test_fips_mode_refuses_to_serve_when_this_python_does_not_enforce_fips(gateway: Gateway, tmp_path: Path) -> None: + if _this_python_enforces_fips(): + pytest.skip("Runner OpenSSL enforces FIPS, so this leg cannot observe the non-enforcing refusal") + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE" in log and "does not enforce FIPS" in log, log + + +@pytest.mark.parametrize("source", ("environment", "config")) +def test_fips_mode_refuses_to_serve_with_tls_verification_disabled( + gateway: Gateway, tmp_path: Path, source: str +) -> None: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path: Final = tmp_path / "ssl_verify_off.yaml" + settings: Final = {**config.get("litellm_settings", {}), "ssl_verify": False} + path.write_text(yaml.safe_dump({**config, "litellm_settings": settings})) + log: Final = ( + refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true", "SSL_VERIFY": "false"}) + if source == "environment" + else refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "true"}, config=path) + ) + assert REFUSAL in log, log + assert "TLS certificate verification is disabled" in log, log + assert ("SSL_VERIFY" if source == "environment" else "litellm_settings.ssl_verify") in log, log + + +def test_fips_mode_refuses_to_serve_on_a_value_that_is_not_a_boolean(gateway: Gateway, tmp_path: Path) -> None: + log: Final = refused_boot_log(gateway, tmp_path, {"LITELLM_FIPS_MODE": "enforced"}) + assert REFUSAL in log, log + assert "LITELLM_FIPS_MODE=enforced" in log and "true or false" in log, log + + +def test_fips_mode_off_serves_even_with_tls_verification_disabled(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {"LITELLM_FIPS_MODE": "false", "SSL_VERIFY": "false"}) as candidate: + assert candidate.client.get("/health/readiness").status_code == 200 diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index 7f7870c9fb5..79feb4f0d92 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -1,13 +1,21 @@ +import asyncio import json import os +import re import signal import socket +import threading import uuid +from collections.abc import Callable, Iterator, Mapping from concurrent.futures import ThreadPoolExecutor from datetime import datetime, timezone +from contextlib import contextmanager +from dataclasses import dataclass from pathlib import Path +from types import MappingProxyType from typing import Final +import anthropic import httpx import psutil import pytest @@ -15,9 +23,10 @@ import yaml from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows from integration._support.mcp import mcp_peer, register_mcp, tool_names -from integration._support.process import group_members, owned_proxy, owned_proxy_process -from integration._support.wire import Reply, Request, wire_server +from integration._support.process import OwnedProxy, group_members, owned_proxy, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server from openai import AsyncOpenAI, OpenAI +from pydantic import JsonValue @pytest.mark.covers("other.observability.guardrails.rewrite_reaches_correct_anthropic_positions") @@ -1797,8 +1806,6 @@ def test_responses_pre_call_denial_stream_survives_worker_kill(gateway: Gateway, for response in responses: assert response.status_code == 200, response.text assert response.headers["content-type"].startswith("text/event-stream"), response.text - - @pytest.mark.parametrize( ("logging_only_scope", "scanned_directions"), (("input", ("request",)), ("output", ("response",)), ("both", ("request", "response"))), @@ -1954,3 +1961,737 @@ def test_logging_only_scope_literal_or_mode_mismatch_is_ignored_at_load_and_keep assert len(upstream.drain()) == 0 guardrails: Final = candidate.get("/v2/guardrails/list")["guardrails"] assert any(object_value(row)["guardrail_name"] == identity for row in guardrails), guardrails +_TOKEN: Final = re.compile(rb"token-[0-9a-f]{32}-\d+") + + +def _secret_for(request: Request) -> str: + token: Final = _TOKEN.search(request.body) + assert token is not None, request.body + return "synthetic-leaked-secret-" + token.group().decode() + + +def _chat_frame(identity: str, choices: tuple[dict[str, JsonValue], ...]) -> bytes: + payload: Final = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": list(choices), + } + return b"data: " + json.dumps(payload).encode() + b"\n\n" + + +def _chat_choice(index: int, delta: dict[str, JsonValue], finish: str | None = None) -> dict[str, JsonValue]: + return {"index": index, "delta": delta, "finish_reason": finish} + + +def _chat_stream_frames(secret: str, shape: str) -> tuple[bytes, ...]: + identity: Final = "chatcmpl-" + uuid.uuid4().hex + tool_call: Final = { + "index": 0, + "id": "call_" + identity, + "type": "function", + "function": {"name": "lookup", "arguments": json.dumps({"query": secret})}, + } + released: Final = { + "text": (_chat_choice(0, {"role": "assistant", "content": secret}),), + "empty": (_chat_choice(0, {"role": "assistant", "content": ""}),), + "tool_call": (_chat_choice(0, {"role": "assistant", "tool_calls": [tool_call]}),), + "two_choices": ( + _chat_choice(0, {"role": "assistant", "content": secret + "-first"}), + _chat_choice(1, {"role": "assistant", "content": secret + "-second"}), + ), + }[shape] + finish: Final = "tool_calls" if shape == "tool_call" else "stop" + tail: Final = tuple(_chat_choice(int(str(choice["index"])), {"content": " tail"}, finish) for choice in released) + return (_chat_frame(identity, released), _chat_frame(identity, tail), b"data: [DONE]\n\n") + + +def _chat_completion_body(secret: str) -> bytes: + return json.dumps( + { + "id": "chatcmpl-" + uuid.uuid4().hex, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": secret}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13}, + } + ).encode() + + +def _responses_tool_call_frames(secret: str, *, with_text: bool) -> tuple[bytes, ...]: + identity: Final = "resp_" + uuid.uuid4().hex + arguments: Final = json.dumps({"query": secret}) + pending: Final = {"type": "function_call", "id": "fc_" + identity, "call_id": "call_" + identity, "name": "lookup"} + finished: Final = {**pending, "arguments": arguments, "status": "completed"} + envelope: Final = {"id": identity, "object": "response", "created_at": 1, "model": "gpt-4o-mini", "output": []} + message: Final = { + "type": "message", + "id": "msg_" + identity, + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": secret, "annotations": []}], + } + text_delta: Final = { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": secret, + } + text_events: Final = (text_delta,) if with_text else () + tool_index: Final = len(text_events) + output: Final = [*((message,) if with_text else ()), finished] + events: Final = ( + {"type": "response.created", "response": {**envelope, "status": "in_progress"}}, + *text_events, + { + "type": "response.output_item.added", + "output_index": tool_index, + "item": {**pending, "arguments": "", "status": "in_progress"}, + }, + { + "type": "response.function_call_arguments.delta", + "item_id": "fc_" + identity, + "output_index": tool_index, + "delta": arguments, + }, + {"type": "response.output_item.done", "output_index": tool_index, "item": finished}, + {"type": "response.completed", "response": {**envelope, "status": "completed", "output": output}}, + ) + encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + released_through: Final = len(events) - 1 if with_text else 3 + return (b"".join(encoded[:released_through]), b"".join(encoded[released_through:])) + + +def _responses_stream_frames(secret: str) -> tuple[bytes, ...]: + identity: Final = "resp_" + uuid.uuid4().hex + completed: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": "msg_" + identity, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": secret, "annotations": []}], + } + ], + "usage": { + "input_tokens": 11, + "output_tokens": 4, + "total_tokens": 15, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + } + events: Final = ( + {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}}, + { + "type": "response.output_text.delta", + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": secret, + }, + {"type": "response.completed", "response": completed}, + ) + encoded: Final = tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + return (encoded[0] + encoded[1], encoded[2]) + + +def _gemini_stream_frames(secret: str) -> tuple[bytes, ...]: + def frame(text: str, finish: str | None) -> bytes: + candidate: Final = { + "content": {"parts": [{"text": text}], "role": "model"}, + "index": 0, + **({"finishReason": finish} if finish else {}), + } + payload: Final = { + "candidates": [candidate], + "usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15}, + "modelVersion": "gemini-2.5-flash", + } + return b"data: " + json.dumps(payload).encode() + b"\r\n\r\n" + + return (frame(secret, None), frame(" tail", "STOP")) + + +def _scripted_provider(gate: threading.Event | None, pause: float, shape: str) -> Callable[[Request], Reply]: + def provider(request: Request) -> Reply: + secret: Final = _secret_for(request) + path: Final = request.target.split("?")[0] + if path.endswith("/chat/completions") and not json.loads(request.body).get("stream"): + return Reply(body=_chat_completion_body(secret)) + frames: Final = ( + _gemini_stream_frames(secret) + if "streamGenerateContent" in path + else ( + _responses_tool_call_frames(secret, with_text=shape == "text_then_tool_call") + if shape in ("tool_call", "text_then_tool_call") + else _responses_stream_frames(secret) + ) + if path.endswith("/responses") + else _chat_stream_frames(secret, shape) + ) + return Reply(content_type="text/event-stream", chunks=frames, gate_after_first=gate, pause_between_chunks=pause) + + return provider + + +def _allowing_guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + +def _failing_response_scans(reply: Reply) -> Callable[[Request], Reply]: + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + if json.loads(request.body)["input_type"] == "response": + return reply + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + return guardrail + + +def _post_call_config( + tmp_path: Path, identity: str, policy_url: str, params: Mapping[str, JsonValue], default_on: bool +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "post_call", + "default_on": default_on, + "api_base": policy_url, + "api_key": "synthetic-guardrail-key", + **params, + }, + } + ] + path: Final = tmp_path / f"{identity}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +class _ScanLog: + def __init__(self, policy: Wire) -> None: + self.policy: Final = policy + self.seen: tuple[dict[str, JsonValue], ...] = () + + def response_scans(self, secret: str) -> tuple[dict[str, JsonValue], ...]: + self.seen = (*self.seen, *(object_value(json.loads(request.body)) for request in self.policy.drain())) + return tuple(body for body in self.seen if body["input_type"] == "response" and secret in json.dumps(body)) + + +@dataclass(frozen=True, slots=True) +class _DisconnectRig: + owned: OwnedProxy + model: str + gemini: str + scans: _ScanLog + identity: str + gate: threading.Event + upstream: Wire + + @property + def candidate(self) -> Gateway: + return self.owned.gateway + + def token(self, index: int = 0) -> str: + return f"token-{self.identity.removeprefix('guardrail')}-{index}" + + def secret(self, index: int = 0) -> str: + return "synthetic-leaked-secret-" + self.token(index) + + +_END_OF_STREAM_ONLY: Final = MappingProxyType({"streaming_end_of_stream_only": True}) + + +@contextmanager +def _disconnect_rig( + gateway: Gateway, + tmp_path: Path, + *, + params: Mapping[str, JsonValue] = _END_OF_STREAM_ONLY, + guardrail: Callable[[Request], Reply] = _allowing_guardrail, + gated: bool = True, + pause: float = 0, + shape: str = "text", + default_on: bool = True, + workers: int = 1, +) -> Iterator[_DisconnectRig]: + identity: Final = "guardrail" + uuid.uuid4().hex + gate: Final = threading.Event() + with ( + wire_server(guardrail) as policy, + wire_server(_scripted_provider(gate if gated else None, pause, shape)) as upstream, + ): + config: Final = _post_call_config(tmp_path, identity, policy.url, params, default_on) + try: + with ( + owned_proxy_process(gateway, tmp_path, {}, config=config, workers=workers) as owned, + owned.gateway.scenario() as scenario, + ): + yield _DisconnectRig( + owned, + scenario.model(api_base=upstream.url + "/v1", api_key="synthetic-openai-key"), + scenario.model( + model="gemini/gemini-2.5-flash", api_base=upstream.url, api_key="synthetic-gemini-key" + ), + _ScanLog(policy), + identity, + gate, + upstream, + ) + finally: + gate.set() + + +def _close_on(response: httpx.Response, marker: str) -> str: + assert response.status_code == 200, response.read() + for line in response.iter_lines(): + if marker in line: + return line + raise AssertionError(f"The stream ended before the client received {marker}") + + +def _stream_and_close(rig: _DisconnectRig, path: str, body: Mapping[str, JsonValue], marker: str) -> str: + with rig.candidate.client.stream( + "POST", path, json=dict(body), headers={"Authorization": f"Bearer {rig.candidate.key}"} + ) as response: + return _close_on(response, marker) + + +def _chat_body(rig: _DisconnectRig, index: int, stream: bool = True) -> dict[str, JsonValue]: + return { + "model": rig.model, + "messages": [{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + "stream": stream, + } + + +def _chat_httpx(rig: _DisconnectRig, index: int = 0) -> str: + return _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, index), rig.secret(index)) + + +def _responses_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {"model": rig.model, "input": "synthetic prompt " + rig.token(index), "stream": True} + return _stream_and_close(rig, "/v1/responses", body, rig.secret(index)) + + +def _messages_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {**_chat_body(rig, index), "max_tokens": 64} + return _stream_and_close(rig, "/v1/messages", body, rig.secret(index)) + + +def _gemini_httpx(rig: _DisconnectRig, index: int = 0) -> str: + body: Final = {"contents": [{"role": "user", "parts": [{"text": "synthetic prompt " + rig.token(index)}]}]} + path: Final = f"/v1beta/models/{rig.gemini}:streamGenerateContent?alt=sse" + return _stream_and_close(rig, path, body, rig.secret(index)) + + +def _chat_async_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str: + async def read() -> str: + client: Final = AsyncOpenAI( + base_url=str(rig.candidate.client.base_url) + "/v1", + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.AsyncClient(trust_env=False, timeout=30), + ) + async with client: + stream: Final = await client.chat.completions.create( + model=rig.model, + messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + stream=True, + ) + async for chunk in stream: + if chunk.choices and rig.secret(index) in (chunk.choices[0].delta.content or ""): + await stream.close() + return chunk.choices[0].delta.content or "" + raise AssertionError("The stream ended before the client received the streamed content") + + return asyncio.run(read()) + + +def _responses_openai_sdk(rig: _DisconnectRig, index: int = 0) -> str: + client: Final = OpenAI( + base_url=str(rig.candidate.client.base_url) + "/v1", + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) + with client: + stream: Final = client.responses.create( + model=rig.model, input="synthetic prompt " + rig.token(index), stream=True + ) + for event in stream: + if event.type == "response.output_text.delta" and rig.secret(index) in event.delta: + stream.close() + return event.delta + raise AssertionError("The stream ended before the client received the streamed content") + + +def _messages_anthropic_sdk(rig: _DisconnectRig, index: int = 0) -> str: + client: Final = anthropic.Anthropic( + base_url=str(rig.candidate.client.base_url), + api_key=rig.candidate.key, + max_retries=0, + http_client=httpx.Client(trust_env=False, timeout=30), + ) + with client: + stream: Final = client.messages.create( + model=rig.model, + max_tokens=64, + messages=[{"role": "user", "content": "synthetic prompt " + rig.token(index)}], + stream=True, + ) + for event in stream: + text: Final = ( + event.delta.text if event.type == "content_block_delta" and event.delta.type == "text_delta" else "" + ) + if rig.secret(index) in text: + stream.close() + return text + raise AssertionError("The stream ended before the client received the streamed content") + + +def _scanned_while_upstream_is_held( + rig: _DisconnectRig, disconnect: Callable[[_DisconnectRig, int], str], index: int = 0 +) -> tuple[dict[str, JsonValue], ...]: + try: + received: Final = disconnect(rig, index) + assert rig.secret(index) in received, received + return eventually( + lambda: rig.scans.response_scans(rig.secret(index)), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + + +def _post_call_statuses(rig: _DisconnectRig, model: str, rows: int = 1) -> tuple[tuple[str, ...], ...]: + found: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda values: len(values) == rows, + seconds=70, + ) + + def statuses(metadata: JsonValue) -> tuple[str, ...]: + entries: Final = object_value(metadata).get("guardrail_information") or [] + assert isinstance(entries, list), metadata + post_call: Final = tuple( + entry + for entry in (object_value(value) for value in entries) + if entry.get("guardrail_name") == rig.identity and entry.get("guardrail_mode") == "post_call" + ) + return tuple(str(entry["guardrail_status"]) for entry in post_call) + + return tuple(statuses(row["metadata"]) for row in found) + + +_ENDPOINT_CLIENTS: Final = ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"), + pytest.param(_responses_httpx, id="responses-httpx"), + pytest.param(_responses_openai_sdk, id="responses-openai-sdk"), + pytest.param(_messages_httpx, id="messages-httpx"), + pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk"), + pytest.param(_gemini_httpx, id="native-gemini-stream-generate-content"), +) + + +_NO_DISCONNECT_ROW_LIT_8603: Final = pytest.mark.skip( + reason="BUG: LIT-8603 a mid-stream disconnect writes no spend row" +) + + +@pytest.mark.parametrize("disconnect", _ENDPOINT_CLIENTS) +def test_client_disconnect_mid_stream_still_scans_the_content_it_already_received( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig(gateway, tmp_path) as rig: + scans: Final = _scanned_while_upstream_is_held(rig, disconnect) + assert len(scans) == 1, scans + + +@pytest.mark.parametrize( + "disconnect", + ( + pytest.param(_chat_httpx, id="chat-httpx"), + pytest.param(_chat_async_openai_sdk, id="chat-async-openai-sdk"), + pytest.param(_responses_httpx, id="responses-httpx", marks=_NO_DISCONNECT_ROW_LIT_8603), + pytest.param(_messages_anthropic_sdk, id="messages-anthropic-sdk", marks=_NO_DISCONNECT_ROW_LIT_8603), + pytest.param( + _gemini_httpx, + id="native-gemini-stream-generate-content", + marks=pytest.mark.skip(reason="BUG: LIT-9087 a mid-stream disconnect writes no spend row"), + ), + ), +) +def test_client_disconnect_mid_stream_records_the_post_call_verdict_on_the_spend_row( + gateway: Gateway, tmp_path: Path, disconnect: Callable[[_DisconnectRig, int], str] +) -> None: + with _disconnect_rig(gateway, tmp_path) as rig: + _scanned_while_upstream_is_held(rig, disconnect) + model: Final = rig.gemini if disconnect is _gemini_httpx else rig.model + assert _post_call_statuses(rig, model) == (("success",),) + + +@pytest.mark.parametrize( + "reply", + ( + pytest.param(Reply(status=500, body=b'{"error": "synthetic guardrail outage"}'), id="guardrail-500"), + pytest.param(Reply(body=b"synthetic non-json guardrail body"), id="guardrail-malformed-200"), + ), +) +def test_client_disconnect_mid_stream_records_a_failed_scan_when_the_guardrail_errors( + gateway: Gateway, tmp_path: Path, reply: Reply +) -> None: + with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(reply)) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + assert _post_call_statuses(rig, rig.model) == (("guardrail_failed_to_respond",),) + + +def test_client_disconnect_mid_stream_records_a_blocking_verdict_and_keeps_serving( + gateway: Gateway, tmp_path: Path +) -> None: + blocked: Final = Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic leak"}).encode()) + with _disconnect_rig(gateway, tmp_path, guardrail=_failing_response_scans(blocked)) as rig: + _scanned_while_upstream_is_held(rig, _chat_httpx) + statuses: Final = _post_call_statuses(rig, rig.model) + assert len(statuses) == 1 and len(statuses[0]) == 1 and statuses[0][0] != "success", statuses + health: Final = rig.candidate.request("GET", "/health/liveliness") + assert health.status_code == 200, health.text + + +def test_client_disconnect_while_end_of_stream_scan_is_in_flight_still_records_the_verdict( + gateway: Gateway, tmp_path: Path +) -> None: + scan_started: Final = threading.Event() + scan_released: Final = threading.Event() + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + if json.loads(request.body)["input_type"] == "response": + scan_started.set() + assert scan_released.wait(timeout=30), "The in-flight scan was never released" + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False) as rig: + with rig.candidate.client.stream( + "POST", + "/v1/chat/completions", + json=_chat_body(rig, 0), + headers={"Authorization": f"Bearer {rig.candidate.key}"}, + ) as response: + try: + assert rig.secret() in _close_on(response, rig.secret()) + assert scan_started.wait(timeout=10), "The end-of-stream scan never started" + finally: + pass + scan_released.set() + assert _post_call_statuses(rig, rig.model) == (("success",),) + assert len(rig.scans.response_scans(rig.secret())) == 1 + + +@pytest.mark.parametrize( + ("params", "shape", "expected"), + ( + pytest.param( + {"streaming_buffer_until_moderated": False}, + "text", + ("synthetic-leaked-secret-",), + id="sampled-before-the-sampling-threshold", + ), + pytest.param( + {"streaming_buffer_until_moderated": False, "streaming_transform_mode": "incremental_diff"}, + "tool_call", + ('\\"query\\": \\"synthetic-leaked-secret-',), + id="incremental-diff-tool-call-in-flight", + ), + pytest.param(dict(_END_OF_STREAM_ONLY), "two_choices", ("-first", "-second"), id="two-choices"), + ), +) +def test_client_disconnect_mid_stream_scans_what_each_streaming_mode_released( + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue], shape: str, expected: tuple[str, ...] +) -> None: + with _disconnect_rig(gateway, tmp_path, params=params, shape=shape) as rig: + marker: Final = rig.secret() + ("-first" if shape == "two_choices" else "") + try: + received: Final = _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), rig.token()) + assert rig.token() in received, received + scans: Final = eventually( + lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + payload: Final = json.dumps(scans[-1]) + assert all(fragment in payload for fragment in expected), (marker, scans) + assert _post_call_statuses(rig, rig.model)[0][-1:] == ("success",) + + +def _tool_call_request(rig: _DisconnectRig, path: str) -> dict[str, JsonValue]: + if path == "/v1/responses": + return {"model": rig.model, "input": "synthetic prompt " + rig.token(), "stream": True} + if path == "/v1/messages": + return {**_chat_body(rig, 0), "max_tokens": 64} + return _chat_body(rig, 0) + + +@pytest.mark.parametrize( + "path", + ( + pytest.param("/v1/chat/completions", id="chat"), + pytest.param("/v1/responses", id="responses"), + pytest.param("/v1/messages", id="messages"), + ), +) +def test_client_disconnect_mid_tool_call_scans_the_tool_call_it_already_received( + gateway: Gateway, tmp_path: Path, path: str +) -> None: + with _disconnect_rig(gateway, tmp_path, shape="tool_call") as rig: + try: + received: Final = _stream_and_close(rig, path, _tool_call_request(rig, path), rig.token()) + assert rig.token() in received, received + scans: Final = eventually( + lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) >= 1, seconds=4 + ) + finally: + rig.gate.set() + assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans + + +def test_client_disconnect_after_a_finished_responses_tool_call_scans_the_text_and_tool_call_it_received( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig(gateway, tmp_path, shape="text_then_tool_call") as rig: + try: + received: Final = _stream_and_close( + rig, "/v1/responses", _tool_call_request(rig, "/v1/responses"), "response.output_item.done" + ) + assert rig.token() in received, received + scans: Final = eventually( + lambda: tuple(scan for scan in rig.scans.response_scans(rig.secret()) if scan.get("texts")), + lambda values: len(values) >= 1, + seconds=4, + ) + finally: + rig.gate.set() + assert rig.secret() in json.dumps(scans[-1].get("texts")), scans + assert rig.secret() in json.dumps(scans[-1].get("tool_calls")), scans + + +def test_client_disconnect_mid_stream_scans_for_a_guardrail_the_request_opted_into( + gateway: Gateway, tmp_path: Path +) -> None: + with _disconnect_rig(gateway, tmp_path, default_on=False) as rig: + body: Final = {**_chat_body(rig, 0), "guardrails": [rig.identity]} + try: + received: Final = _stream_and_close(rig, "/v1/chat/completions", body, rig.secret()) + assert rig.secret() in received, received + eventually(lambda: rig.scans.response_scans(rig.secret()), lambda values: len(values) == 1, seconds=4) + finally: + rig.gate.set() + assert _post_call_statuses(rig, rig.model) == (("success",),) + + +def test_client_disconnect_before_any_content_sends_no_response_scan(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, shape="empty") as rig: + try: + _stream_and_close(rig, "/v1/chat/completions", _chat_body(rig, 0), "data: ") + finally: + rig.gate.set() + rows: Final = _post_call_statuses(rig, rig.model) + assert rig.scans.response_scans(rig.secret()) == (), rig.scans.seen + assert len(rows) == 1 and "success" not in rows[0], rows + + +@pytest.mark.parametrize( + "params", + ( + pytest.param(dict(_END_OF_STREAM_ONLY), id="end-of-stream-only"), + pytest.param({"streaming_buffer_until_moderated": True}, id="buffered"), + ), +) +def test_a_fully_read_stream_is_scanned_exactly_once( + gateway: Gateway, tmp_path: Path, params: dict[str, JsonValue] +) -> None: + with _disconnect_rig(gateway, tmp_path, params=params, gated=False) as rig: + response: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0)) + assert response.status_code == 200, response.text + assert rig.secret() in response.text and "[DONE]" in response.text, response.text + assert _post_call_statuses(rig, rig.model) == (("success",),) + assert len(rig.scans.response_scans(rig.secret())) == 1, rig.scans.seen + + +def _cached_twin_rows(rig: _DisconnectRig) -> tuple[tuple[str, ...], ...]: + first: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False)) + second: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig, 0, stream=False)) + assert first.status_code == second.status_code == 200, (first.text, second.text) + assert first.json()["choices"] == second.json()["choices"], (first.text, second.text) + assert len(rig.upstream.drain()) == 1, "the second request must be served from the cache" + return _post_call_statuses(rig, rig.model, rows=2) + + +def test_a_non_streaming_response_and_its_cache_hit_are_each_scanned_once(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, gated=False) as rig: + rows: Final = _cached_twin_rows(rig) + assert rows[0] == ("success",), rows + assert len(rig.scans.response_scans(rig.secret())) == 2, (rows, rig.scans.seen) + + +def test_a_cache_hit_row_records_the_post_call_verdict_of_its_scan(gateway: Gateway, tmp_path: Path) -> None: + pytest.skip("BUG: LIT-9088 the cache-hit spend row drops the post_call verdict of the scan that ran on it") + with _disconnect_rig(gateway, tmp_path, gated=False) as rig: + assert _cached_twin_rows(rig) == (("success",), ("success",)) + + +def test_concurrent_disconnects_during_a_guardrail_outage_each_record_exactly_one_verdict( + gateway: Gateway, tmp_path: Path +) -> None: + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api", request.target + body: Final = json.loads(request.body) + index: Final = int(_secret_for(request).rsplit("-", 1)[1]) + if body["input_type"] == "response" and index % 3 == 0: + return Reply(status=503, body=b'{"error": "synthetic guardrail outage"}') + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + clients: Final = (_chat_httpx, _responses_httpx, _messages_httpx) + with _disconnect_rig(gateway, tmp_path, guardrail=guardrail, gated=False, pause=3, workers=2) as rig: + with ThreadPoolExecutor(max_workers=30) as pool: + received: Final = tuple(pool.map(lambda index: clients[index % 3](rig, index), range(30))) + assert all(rig.secret(index) in line for index, line in enumerate(received)), received + scanned: Final = eventually( + lambda: tuple(len(rig.scans.response_scans(rig.secret(index) + '"')) for index in range(30)), + lambda counts: all(count >= 1 for count in counts), + seconds=20, + ) + assert scanned == (1,) * 30, scanned + chat_rows: Final = _post_call_statuses(rig, rig.model, rows=10) + assert sorted(chat_rows) == sorted( + ("guardrail_failed_to_respond",) if index % 3 == 0 else ("success",) for index in range(0, 30, 3) + ), chat_rows + + +def test_disconnect_scans_keep_recording_after_a_worker_is_killed(gateway: Gateway, tmp_path: Path) -> None: + with _disconnect_rig(gateway, tmp_path, gated=False, pause=3, workers=2) as rig: + members: Final = tuple( + member for member in group_members(rig.owned.process.pid) if member.pid != rig.owned.process.pid + ) + workers: Final = tuple(member for member in members if any("spawn_main" in part for part in member.cmdline())) + assert len(workers) >= 2, members + workers[0].send_signal(signal.SIGKILL) + psutil.wait_procs((workers[0],), timeout=10) + with ThreadPoolExecutor(max_workers=8) as pool: + received: Final = tuple(pool.map(lambda index: _chat_httpx(rig, index), range(8))) + assert all(rig.secret(index) in line for index, line in enumerate(received)), received + assert _post_call_statuses(rig, rig.model, rows=8) == (("success",),) * 8 + assert rig.owned.process.poll() is None diff --git a/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py b/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py new file mode 100644 index 00000000000..ee347592d37 --- /dev/null +++ b/tests/integration/observability/test_s3_v2_streaming_failure_dedupe.py @@ -0,0 +1,408 @@ +from __future__ import annotations + +import json +import uuid +from collections import Counter +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from typing import Final, Literal + +import httpx +import pytest +from _s3_v2_support import RecordingS3Sink, call_surface, collect_payloads, matched_ids, s3_config, surface_reply +from integration._support.client import Gateway, JsonValue, Scenario, eventually, object_value +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +Surface = Literal["chat", "messages", "responses"] +SURFACES: Final[tuple[Surface, ...]] = ("chat", "messages", "responses") +FAILURE_BODY: Final[bytes] = json.dumps( + { + "error": { + "message": "synthetic upstream failure", + "type": "server_error", + "param": None, + "code": "synthetic_failure", + } + } +).encode() + + +def _failure_provider(request: Request) -> Reply: + if request.method != "POST": + return Reply(status=404) + return Reply(status=503, body=FAILURE_BODY) + + +def _surface_target(surface: Surface) -> str: + return { + "chat": "/v1/chat/completions", + "messages": "/v1/messages", + "responses": "/v1/responses", + }[surface] + + +def _surface_request( + candidate: Gateway, + surface: Surface, + openai_model: str, + anthropic_model: str, + key: str, + marker: str, + stream: bool, +) -> httpx.Response: + if surface == "chat": + body: Final = { + "model": openai_model, + "messages": [{"role": "user", "content": marker}], + "stream": stream, + } + return candidate.request("POST", _surface_target(surface), body, key=key) + if surface == "messages": + body: Final = { + "model": anthropic_model, + "max_tokens": 16, + "messages": [{"role": "user", "content": marker}], + "stream": stream, + } + return candidate.request("POST", _surface_target(surface), body, key=key) + body: Final = {"model": openai_model, "input": marker, "stream": stream} + return candidate.request("POST", _surface_target(surface), body, key=key) + + +def _responses_stream_id(response: httpx.Response) -> str: + events: Final = tuple( + object_value(json.loads(line.removeprefix("data: "))) + for line in response.text.splitlines() + if line.startswith("data: {") + ) + completed: Final = tuple( + object_value(event["response"]) + for event in events + if event.get("type") == "response.completed" + ) + assert len(completed) == 1, events + response_id: Final = completed[0]["id"] + assert isinstance(response_id, str), response_id + return response_id + + +def _register_models(scenario: Scenario, upstream_url: str) -> tuple[str, str]: + openai_model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=upstream_url + "/v1", + api_key="synthetic-provider-key", + ) + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=upstream_url, + api_key="synthetic-provider-key", + ) + return openai_model, anthropic_model + + +def _request_key(payload: dict[str, JsonValue], response: httpx.Response, marker: str) -> str | None: + call_id: Final = response.headers.get("x-litellm-call-id") + if call_id is not None and call_id in {str(payload.get("id")), str(payload.get("litellm_call_id"))}: + return call_id + if marker in json.dumps(payload): + return marker + return None + + +def _assert_upstream_requests(requests: tuple[Request, ...], surface: Surface, expected: int) -> None: + target: Final = _surface_target(surface) + posts: Final = tuple(request for request in requests if request.method == "POST") + observed: Final = tuple(request.target for request in posts) + assert len(posts) == expected, f"expected {expected} upstream POSTs, observed {len(posts)}: {observed}" + assert all(request.target == target for request in posts), f"expected {target}, observed {observed}" + + +def _assert_one_payload( + sink: RecordingS3Sink, + response: httpx.Response, + marker: str, + status: str, + upstream_posts: int, +) -> tuple[dict[str, JsonValue], ...]: + first_payloads: Final = collect_payloads(sink, 1) + assert first_payloads[0]["status"] == status + # A full three-flush-interval window is needed to detect late duplicate uploads. + payloads: Final = eventually( + lambda: sink.payloads(), + lambda observed: len(observed) >= 2, + seconds=6, + return_last_on_timeout=True, + ) + assert len(payloads) == 1, ( + f"expected one {status} payload after {upstream_posts} upstream POSTs, " + f"observed {len(payloads)} payloads" + ) + assert ( + _request_key(payloads[0], response, marker) == response.headers.get("x-litellm-call-id") + or marker in json.dumps(payloads[0]) + ), f"payload did not match request id or marker: {payloads[0]!r}" + return payloads + + +@pytest.mark.parametrize( + ("surface", "stream"), + [ + pytest.param("chat", False, id="chat-nonstream"), + pytest.param("chat", True, id="chat-stream"), + pytest.param("messages", False, id="messages-nonstream"), + pytest.param("messages", True, id="messages-stream"), + pytest.param("responses", False, id="responses-nonstream"), + pytest.param("responses", True, id="responses-stream"), + ], +) +def test_retried_failure_uploads_one_s3_object( + gateway: Gateway, + tmp_path: Path, + surface: Surface, + stream: bool, +) -> None: + marker: Final = f"s3-a-{surface}-{uuid.uuid4().hex}" + sink: Final = RecordingS3Sink() + with wire_server(_failure_provider) as upstream, wire_server(sink.respond) as bucket: + config: Final = s3_config(tmp_path, bucket.url, {}, settings={"num_retries": 2}) + with ( + owned_proxy( + gateway, + tmp_path, + {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, + config=config, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + openai_model, anthropic_model = _register_models(scenario, upstream.url) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + response: Final = _surface_request( + candidate, + surface, + openai_model, + anthropic_model, + key, + marker, + stream, + ) + requests: Final = upstream.drain() + _assert_upstream_requests(requests, surface, 3) + assert not 200 <= response.status_code < 300, f"{response.status_code}: {response.text}" + assert "synthetic upstream failure" in response.text, response.text + _assert_one_payload(sink, response, marker, "failure", len(requests)) + + +@pytest.mark.parametrize("stream", [False, True], ids=["nonstream", "stream"]) +def test_fallback_chain_failure_uploads_one_s3_object( + gateway: Gateway, + tmp_path: Path, + stream: bool, +) -> None: + marker: Final = f"s3-b-{uuid.uuid4().hex}" + sink: Final = RecordingS3Sink() + with ( + wire_server(_failure_provider) as primary_upstream, + wire_server(_failure_provider) as secondary_upstream, + wire_server(sink.respond) as bucket, + ): + config: Final = s3_config(tmp_path, bucket.url, {}, settings={"num_retries": 0}) + with ( + owned_proxy( + gateway, + tmp_path, + {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, + config=config, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + primary_model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=primary_upstream.url + "/v1", + api_key="synthetic-primary-key", + ) + secondary_model: Final = scenario.model( + model="openai/gpt-4o-mini", + api_base=secondary_upstream.url + "/v1", + api_key="synthetic-secondary-key", + ) + key: Final = scenario.key(models=[primary_model, secondary_model]) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": primary_model, + "messages": [{"role": "user", "content": marker}], + "stream": stream, + "fallbacks": [secondary_model], + "num_retries": 0, + }, + key=key, + ) + primary_requests: Final = primary_upstream.drain() + secondary_requests: Final = secondary_upstream.drain() + _assert_upstream_requests(primary_requests, "chat", 1) + _assert_upstream_requests(secondary_requests, "chat", 1) + assert not 200 <= response.status_code < 300, f"{response.status_code}: {response.text}" + assert "synthetic upstream failure" in response.text, response.text + _assert_one_payload(sink, response, marker, "failure", len(primary_requests) + len(secondary_requests)) + + +@pytest.mark.parametrize("surface", SURFACES, ids=SURFACES) +def test_streaming_success_uploads_one_s3_object( + gateway: Gateway, + tmp_path: Path, + surface: Surface, +) -> None: + marker: Final = f"s3-c-{surface}-{uuid.uuid4().hex}" + sink: Final = RecordingS3Sink() + with wire_server(surface_reply) as upstream, wire_server(sink.respond) as bucket: + config: Final = s3_config(tmp_path, bucket.url, {}, settings={"num_retries": 0}) + with ( + owned_proxy( + gateway, + tmp_path, + {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, + config=config, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + openai_model, anthropic_model = _register_models(scenario, upstream.url) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + captured_responses: Final[list[httpx.Response]] = [] + candidate.client.event_hooks["response"].append(captured_responses.append) + surface_response_id, call_id = call_surface( + candidate, + f"{surface}_stream", + openai_model, + anthropic_model, + key, + marker, + ) + if surface == "responses": + assert len(captured_responses) == 1, captured_responses + response_id: Final = ( + _responses_stream_id(captured_responses[0]) if surface == "responses" else surface_response_id + ) + payloads: Final = collect_payloads(sink, 1) + payloads_after_window: Final = eventually( + lambda: sink.payloads(), + lambda observed: len(observed) >= 2, + seconds=6, + return_last_on_timeout=True, + ) + assert len(payloads_after_window) == 1, payloads_after_window + assert payloads[0]["status"] == "success" + if surface == "responses": + assert payloads[0]["litellm_call_id"] == call_id + else: + assert payloads[0]["id"] == response_id + assert matched_ids(payloads, ((response_id, call_id),)) == frozenset({str(payloads[0]["id"])}) + + +def test_failure_burst_through_sink_outage_lands_each_request_once( + gateway: Gateway, + tmp_path: Path, +) -> None: + requests_per_variant: Final = 4 + jobs: Final = tuple( + (surface, stream, f"s3-d-{surface}-{'stream' if stream else 'nonstream'}-{index}") + for surface in SURFACES + for stream in (False, True) + for index in range(requests_per_variant) + ) + sink: Final = RecordingS3Sink() + with wire_server(_failure_provider) as upstream, wire_server(sink.respond) as bucket: + config: Final = s3_config(tmp_path, bucket.url, {}, settings={"num_retries": 2}) + with ( + owned_proxy( + gateway, + tmp_path, + {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "2"}, + config=config, + workers=2, + ) as candidate, + candidate.scenario() as scenario, + ): + openai_model, anthropic_model = _register_models(scenario, upstream.url) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + sink.fail_until = float("inf") + + def call(job: tuple[Surface, bool, str]) -> httpx.Response: + surface, stream, marker = job + return _surface_request(candidate, surface, openai_model, anthropic_model, key, marker, stream) + + with ThreadPoolExecutor(max_workers=len(jobs)) as pool: + responses: Final = tuple(pool.map(call, jobs)) + + upstream_requests: Final = tuple( + request for request in upstream.drain() if request.method == "POST" + ) + route_counts: Final = Counter(request.target for request in upstream_requests) + expected_route_counts: Final = { + _surface_target(surface): 2 * requests_per_variant * 3 for surface in SURFACES + } + assert len(upstream_requests) == len(jobs) * 3, ( + f"expected {len(jobs) * 3} upstream POSTs, observed {len(upstream_requests)}: {route_counts}" + ) + assert route_counts == expected_route_counts, f"unexpected upstream routes: {route_counts}" + assert all(not 200 <= response.status_code < 300 for response in responses) + assert all("synthetic upstream failure" in response.text for response in responses) + eventually( + lambda: sink.attempts, + lambda attempts: attempts > len(sink.store), + seconds=30, + ) + rejected: Final = sink.attempts - len(sink.store) + assert rejected >= 1, f"expected at least one rejected sink upload, rejected={rejected}" + assert len(sink.store) == 0, ( + f"expected all sink uploads to fail before recovery, rejected={rejected}, " + f"attempts={sink.attempts}" + ) + sink.fail_until = 0.0 + payloads: Final = eventually( + lambda: sink.payloads(), + lambda observed: len(observed) >= len(jobs), + seconds=90, + return_last_on_timeout=True, + ) + payloads_after_window: Final = eventually( + lambda: sink.payloads(), + lambda observed: len(observed) > len(payloads), + seconds=6, + return_last_on_timeout=True, + ) + payload_matches: Final = tuple( + tuple( + index + for index, (response, job) in enumerate(zip(responses, jobs)) + if _request_key(payload, response, job[2]) is not None + ) + for payload in payloads_after_window + ) + assert all(len(matches) == 1 for matches in payload_matches), ( + f"payloads could not be matched to one request: rejected={rejected}, " + f"matches={payload_matches}, payloads={payloads_after_window}" + ) + landed_keys: Final = tuple( + _request_key(payload, responses[matches[0]], jobs[matches[0]][2]) + for payload, matches in zip(payloads_after_window, payload_matches) + ) + duplicate_keys: Final = tuple(key for key in set(landed_keys) if landed_keys.count(key) > 1) + expected_keys: Final = tuple( + response.headers.get("x-litellm-call-id") or job[2] for response, job in zip(responses, jobs) + ) + missing_keys: Final = tuple(key for key in expected_keys if key not in landed_keys) + assert not duplicate_keys, ( + f"duplicate request ids landed: {duplicate_keys}; " + f"missing={missing_keys}; rejected={rejected}; upstream_posts={len(upstream_requests)} " + f"sink_payloads={len(payloads_after_window)}" + ) + assert not missing_keys, ( + f"requests missing after sink recovery: {missing_keys}; " + f"rejected={rejected}; upstream_posts={len(upstream_requests)} " + f"sink_payloads={len(payloads_after_window)}" + ) diff --git a/tests/integration/security/_sweeps.py b/tests/integration/security/_sweeps.py index f97a0a7fcc6..5b9fa9a613f 100644 --- a/tests/integration/security/_sweeps.py +++ b/tests/integration/security/_sweeps.py @@ -106,6 +106,7 @@ ROUTE_DENY_LIST: Final = MappingProxyType( "/openai_passthrough/{endpoint:path}": "forwards to a provider, not a proxy read", "/get/latest_release_info": "fetches the latest release from api.github.com", "/roi-calculator/repositories": "lists repositories from the configured GitHub API, api.github.com by default", + "/roi-calculator/observed/repositories": "lists repositories from the connected GitHub or GitLab API, api.github.com by default", } ) @@ -140,6 +141,7 @@ NOT_FOUND_EXPECTED: Final = MappingProxyType( "/fallback/{model}": "answers 404 when the model has no fallbacks configured", "/team/{team_id}/members/me": "answers 404 when the caller is not a member, which the admin is not", "/guardrails/submissions/{guardrail_id}": "answers 404 for a guardrail no team submitted for review", + "/credentials/{credential_name:path}/jwks": "answers 404 for any credential that is not an anthropic internal_issuer credential, which the canary credential is not", } ) diff --git a/tests/integration/spend/test_lens_billing.py b/tests/integration/spend/test_lens_billing.py index d8eded62b39..69a1899ba6d 100644 --- a/tests/integration/spend/test_lens_billing.py +++ b/tests/integration/spend/test_lens_billing.py @@ -11,6 +11,8 @@ from tests.integration._support.database import read_rows, write_rows from tests.integration._support.process import owned_proxy from tests.integration.pricing.test_off_peak_pricing import off_peak_window +RELEASE_TAG: Final = "v0.0.0-lens-integration" + def delete_lens(lens_id: str) -> None: write_rows('DELETE FROM "LiteLLM_LensRun" WHERE lens_id=%s', (lens_id,)) @@ -19,8 +21,13 @@ def delete_lens(lens_id: str) -> None: @pytest.mark.parametrize("off_peak", (False, True)) -def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, off_peak: bool) -> None: - with gateway.scenario() as scenario: +def test_lens_bills_selected_key_and_rechecks_its_permissions( + gateway: Gateway, tmp_path: Path, off_peak: bool +) -> None: + with ( + owned_proxy(gateway, tmp_path, {"LITELLM_RELEASE_TAG": RELEASE_TAG}) as isolated, + isolated.scenario() as scenario, + ): model: Final = scenario.model( input_cost_per_token=0.000001, output_cost_per_token=0.000002, @@ -36,12 +43,13 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, ) key: Final = scenario.key(models=[model], max_budget=1) key_id: Final = sha256(key.encode()).hexdigest() - worker: Final = gateway.post( + worker: Final = isolated.post( "/lens/workers/register", {"name": "Billing regression", "analysis_key_id": key_id} ) + assert worker["image"] == "ghcr.io/berriai/litellm-lens-worker:" + RELEASE_TAG worker_id: Final = string_value(object_value(worker["worker"])["id"]) scenario.cleanups.callback(write_rows, 'DELETE FROM "LiteLLM_LensWorker" WHERE id=%s', (worker_id,)) - lens: Final = gateway.post( + lens: Final = isolated.post( "/lens", { "name": "Billing regression", @@ -54,14 +62,19 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, lens_id: Final = string_value(lens["id"]) scenario.cleanups.callback(delete_lens, lens_id) worker_key: Final = string_value(worker["token"]) - unauthorized: Final = gateway.request( + unauthorized: Final = isolated.request( "POST", "/lens/workers/register", {"name": "Denied", "analysis_key_id": key_id}, key=key ) assert unauthorized.status_code == 403, unauthorized.text with ThreadPoolExecutor(max_workers=8) as pool: claims: Final = tuple( pool.map( - lambda _: gateway.request("POST", "/lens/worker/claim?protocol_version=2", {}, key=worker_key), + lambda _: isolated.request( + "POST", + "/lens/worker/claim?protocol_version=4&worker_release=" + RELEASE_TAG, + {}, + key=worker_key, + ), range(8), ) ) @@ -72,7 +85,7 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, assert claim["lens_id"] == lens_id job_id: Final = string_value(object_value(claim["job"])["id"]) path: Final = f"/lens/worker/{lens_id}/{job_id}/model" - result: Final = gateway.post(path, {"prompt": "Inspect this run", "purpose": "extract"}, key=worker_key) + result: Final = isolated.post(path, {"prompt": "Inspect this run", "purpose": "extract"}, key=worker_key) expected: Final = (20 * 0.000001 + 20 * 0.000002) * (0.5 if off_peak else 1) assert result["cost"] == pytest.approx(expected) rows: Final = eventually( @@ -81,8 +94,8 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, seconds=70, ) assert rows[0]["spend"] == pytest.approx(expected) - assert gateway.get(f"/lens/{lens_id}")["spent"] == pytest.approx(expected) - raw_hash: Final = gateway.request( + assert isolated.get(f"/lens/{lens_id}")["spent"] == pytest.approx(expected) + raw_hash: Final = isolated.request( "POST", "/v1/chat/completions", { @@ -92,31 +105,35 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, key=key_id, ) assert raw_hash.status_code == 401, raw_hash.text - gateway.post("/key/update", {"key": key, "max_budget": expected / 2}) - exhausted: Final = gateway.request( + isolated.post("/key/update", {"key": key, "max_budget": expected / 2}) + exhausted: Final = isolated.request( "POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key ) assert exhausted.status_code == 402, exhausted.text - gateway.post("/key/update", {"key": key, "max_budget": 1, "models": ["unavailable-analysis-model"]}) - restricted: Final = gateway.request( + isolated.post("/key/update", {"key": key, "max_budget": 1, "models": ["unavailable-analysis-model"]}) + restricted: Final = isolated.request( "POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key ) assert restricted.status_code == 403, restricted.text - gateway.post("/key/block", {"key": key}) - blocked: Final = gateway.request("POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key) + isolated.post("/key/block", {"key": key}) + blocked: Final = isolated.request( + "POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key + ) assert blocked.status_code == 400, blocked.text - assert gateway.get(f"/lens/{lens_id}")["spent"] == pytest.approx(expected) + assert isolated.get(f"/lens/{lens_id}")["spent"] == pytest.approx(expected) replacement: Final = scenario.key(models=[model], rpm_limit=1) replacement_id: Final = sha256(replacement.encode()).hexdigest() - changed: Final = gateway.request( + changed: Final = isolated.request( "PUT", f"/lens/workers/{worker_id}/billing-key", {"analysis_key_id": replacement_id} ) assert changed.status_code == 200, changed.text - billed_replacement: Final = gateway.post( + billed_replacement: Final = isolated.post( path, {"prompt": "Inspect another run", "purpose": "extract"}, key=worker_key ) assert billed_replacement["cost"] == pytest.approx(expected) - limited: Final = gateway.request("POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key) + limited: Final = isolated.request( + "POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key + ) assert limited.status_code == 429, limited.text second_rows: Final = eventually( lambda: read_rows('SELECT spend FROM "LiteLLM_VerificationToken" WHERE token=%s', (replacement_id,)), @@ -124,16 +141,16 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, seconds=70, ) assert second_rows[0]["spend"] == pytest.approx(expected) - active_revoke: Final = gateway.request("DELETE", f"/lens/workers/{worker_id}") + active_revoke: Final = isolated.request("DELETE", f"/lens/workers/{worker_id}") assert active_revoke.status_code == 409, active_revoke.text - gateway.post(f"/lens/{lens_id}/cancel", {}) - revoked: Final = gateway.request("DELETE", f"/lens/workers/{worker_id}") + isolated.post(f"/lens/{lens_id}/cancel", {}) + revoked: Final = isolated.request("DELETE", f"/lens/workers/{worker_id}") assert revoked.status_code == 200, revoked.text - denied_worker: Final = gateway.request( + denied_worker: Final = isolated.request( "POST", path, {"prompt": "Must not run", "purpose": "extract"}, key=worker_key ) assert denied_worker.status_code == 401, denied_worker.text - forbidden_change: Final = gateway.request( + forbidden_change: Final = isolated.request( "PUT", f"/lens/workers/{worker_id}/billing-key", {"analysis_key_id": replacement_id} ) assert forbidden_change.status_code == 409, forbidden_change.text @@ -161,7 +178,10 @@ def test_worker_spend_logs_do_not_expose_investigation_content( } ) ) - with owned_proxy(gateway, tmp_path, {}, config=config) as isolated, isolated.scenario() as scenario: + with ( + owned_proxy(gateway, tmp_path, {"LITELLM_RELEASE_TAG": RELEASE_TAG}, config=config) as isolated, + isolated.scenario() as scenario, + ): model: Final = scenario.model(input_cost_per_token=0.000001, output_cost_per_token=0.000002) key: Final = scenario.key(models=[model]) key_id: Final = sha256(key.encode()).hexdigest() @@ -185,7 +205,9 @@ def test_worker_spend_logs_do_not_expose_investigation_content( lens_id: Final = string_value(lens["id"]) scenario.cleanups.callback(delete_lens, lens_id) worker_token: Final = string_value(worker["token"]) - claim: Final = isolated.post("/lens/worker/claim?protocol_version=2", {}, key=worker_token) + claim: Final = isolated.post( + "/lens/worker/claim?protocol_version=4&worker_release=" + RELEASE_TAG, {}, key=worker_token + ) job_id: Final = string_value(object_value(claim["job"])["id"]) result: Final = isolated.post( f"/lens/worker/{lens_id}/{job_id}/model", {"prompt": marker, "purpose": "extract"}, key=worker_token diff --git a/tests/litellm_utils_tests/test_utils.py b/tests/litellm_utils_tests/test_utils.py index 92947fbf6fe..64402b5c016 100644 --- a/tests/litellm_utils_tests/test_utils.py +++ b/tests/litellm_utils_tests/test_utils.py @@ -1343,12 +1343,16 @@ def test_is_prompt_caching_enabled_error_handling(): def test_is_prompt_caching_enabled_return_default_image_dimensions(): """ - Assert that `is_prompt_caching_valid_prompt` calls token_counter with use_default_image_token_count=True + Assert that `is_prompt_caching_valid_prompt` counts tokens with use_default_image_token_count=True when processing messages containing images IMPORTANT: Ensures Get token counter does not make a GET request to the image url """ - with patch("litellm.utils.token_counter") as mock_token_counter: + mock_token_counter = MagicMock(return_value=False) + with patch( + "litellm.utils._get_messages_reach_token_count", + return_value=mock_token_counter, + ): litellm.utils.is_prompt_caching_valid_prompt( messages=[ { diff --git a/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json b/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json index 63baadaaf31..f7ebf357d83 100644 --- a/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json +++ b/tests/logging_callback_tests/gcs_pub_sub_body/spend_logs_payload.json @@ -11,7 +11,7 @@ "user": "", "team_id": "", "organization_id": "", - "metadata": "{\"actor_agent_id\": null, \"target_agent_id\": null, \"billing_agent_id\": null, \"agent_execution_mode\": null, \"verified_human_user_id\": null, \"used_client_oauth_token\": null, \"applied_guardrails\": [], \"attempted_fallbacks\": null, \"original_model_group\": null, \"batch_models\": null, \"batch_successful_requests\": null, \"batch_failed_requests\": null, \"mcp_tool_call_metadata\": null, \"vector_store_request_metadata\": null, \"routing_decision\": null, \"internal_call_origin\": null, \"router_metadata\": null, \"autorouter_savings_estimate\": null, \"autorouter_baseline_observation\": null, \"azure_spillover\": null, \"guardrail_information\": null, \"compression_savings\": null, \"litellm_gateway_injected_cache\": null, \"usage_object\": {\"completion_tokens\": 20, \"prompt_tokens\": 10, \"total_tokens\": 30, \"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"model_map_information\": {\"model_map_key\": \"gpt-4o\", \"model_map_value\": {\"key\": \"gpt-4o\", \"max_tokens\": 16384, \"max_input_tokens\": 128000, \"max_output_tokens\": 16384, \"input_cost_per_token\": 2.5e-06, \"cache_creation_input_token_cost\": null, \"cache_read_input_token_cost\": 1.25e-06, \"input_cost_per_character\": null, \"input_cost_per_token_above_128k_tokens\": null, \"input_cost_per_token_above_200k_tokens\": null, \"input_cost_per_query\": null, \"input_cost_per_second\": null, \"input_cost_per_audio_token\": null, \"input_cost_per_token_batches\": 1.25e-06, \"output_cost_per_token_batches\": 5e-06, \"output_cost_per_token\": 1e-05, \"output_cost_per_audio_token\": null, \"output_cost_per_character\": null, \"output_cost_per_token_above_128k_tokens\": null, \"output_cost_per_character_above_128k_tokens\": null, \"output_cost_per_token_above_200k_tokens\": null, \"output_cost_per_second\": null, \"output_cost_per_image\": null, \"output_vector_size\": null, \"litellm_provider\": \"openai\", \"mode\": \"chat\", \"supports_system_messages\": true, \"supports_response_schema\": true, \"supports_vision\": true, \"supports_function_calling\": true, \"supports_tool_choice\": true, \"supports_assistant_prefill\": false, \"supports_prompt_caching\": true, \"supports_audio_input\": false, \"supports_audio_output\": false, \"supports_pdf_input\": false, \"supports_embedding_image_input\": false, \"supports_native_streaming\": null, \"supports_web_search\": true, \"supports_reasoning\": false, \"search_context_cost_per_query\": {\"search_context_size_low\": 0.03, \"search_context_size_medium\": 0.035, \"search_context_size_high\": 0.05}, \"tpm\": null, \"rpm\": null, \"supported_openai_params\": [\"frequency_penalty\", \"logit_bias\", \"logprobs\", \"top_logprobs\", \"max_tokens\", \"max_completion_tokens\", \"modalities\", \"prediction\", \"n\", \"presence_penalty\", \"seed\", \"stop\", \"stream\", \"stream_options\", \"temperature\", \"top_p\", \"tools\", \"tool_choice\", \"function_call\", \"functions\", \"max_retries\", \"extra_headers\", \"parallel_tool_calls\", \"audio\", \"response_format\", \"user\"]}}, \"additional_usage_values\": {\"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"user_api_key\": null, \"user_api_key_alias\": null, \"user_api_key_team_id\": null, \"user_api_key_project_id\": null, \"user_api_key_project_alias\": null, \"user_api_key_org_id\": null, \"user_api_key_user_id\": null, \"user_api_key_team_alias\": null, \"spend_logs_metadata\": null, \"requester_ip_address\": null, \"user_agent\": null, \"status\": null, \"proxy_server_request\": null, \"error_information\": null, \"attempted_retries\": null, \"max_retries\": null}", + "metadata": "{\"actor_agent_id\": null, \"target_agent_id\": null, \"billing_agent_id\": null, \"agent_execution_mode\": null, \"verified_human_user_id\": null, \"used_client_oauth_token\": null, \"applied_guardrails\": [], \"attempted_fallbacks\": null, \"original_model_group\": null, \"batch_models\": null, \"batch_successful_requests\": null, \"batch_failed_requests\": null, \"mcp_tool_call_metadata\": null, \"vector_store_request_metadata\": null, \"routing_decision\": null, \"internal_call_origin\": null, \"litellm_roi_estimator\": false, \"router_metadata\": null, \"autorouter_savings_estimate\": null, \"autorouter_baseline_observation\": null, \"azure_spillover\": null, \"guardrail_information\": null, \"compression_savings\": null, \"litellm_gateway_injected_cache\": null, \"usage_object\": {\"completion_tokens\": 20, \"prompt_tokens\": 10, \"total_tokens\": 30, \"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"model_map_information\": {\"model_map_key\": \"gpt-4o\", \"model_map_value\": {\"key\": \"gpt-4o\", \"max_tokens\": 16384, \"max_input_tokens\": 128000, \"max_output_tokens\": 16384, \"input_cost_per_token\": 2.5e-06, \"cache_creation_input_token_cost\": null, \"cache_read_input_token_cost\": 1.25e-06, \"input_cost_per_character\": null, \"input_cost_per_token_above_128k_tokens\": null, \"input_cost_per_token_above_200k_tokens\": null, \"input_cost_per_query\": null, \"input_cost_per_second\": null, \"input_cost_per_audio_token\": null, \"input_cost_per_token_batches\": 1.25e-06, \"output_cost_per_token_batches\": 5e-06, \"output_cost_per_token\": 1e-05, \"output_cost_per_audio_token\": null, \"output_cost_per_character\": null, \"output_cost_per_token_above_128k_tokens\": null, \"output_cost_per_character_above_128k_tokens\": null, \"output_cost_per_token_above_200k_tokens\": null, \"output_cost_per_second\": null, \"output_cost_per_image\": null, \"output_vector_size\": null, \"litellm_provider\": \"openai\", \"mode\": \"chat\", \"supports_system_messages\": true, \"supports_response_schema\": true, \"supports_vision\": true, \"supports_function_calling\": true, \"supports_tool_choice\": true, \"supports_assistant_prefill\": false, \"supports_prompt_caching\": true, \"supports_audio_input\": false, \"supports_audio_output\": false, \"supports_pdf_input\": false, \"supports_embedding_image_input\": false, \"supports_native_streaming\": null, \"supports_web_search\": true, \"supports_reasoning\": false, \"search_context_cost_per_query\": {\"search_context_size_low\": 0.03, \"search_context_size_medium\": 0.035, \"search_context_size_high\": 0.05}, \"tpm\": null, \"rpm\": null, \"supported_openai_params\": [\"frequency_penalty\", \"logit_bias\", \"logprobs\", \"top_logprobs\", \"max_tokens\", \"max_completion_tokens\", \"modalities\", \"prediction\", \"n\", \"presence_penalty\", \"seed\", \"stop\", \"stream\", \"stream_options\", \"temperature\", \"top_p\", \"tools\", \"tool_choice\", \"function_call\", \"functions\", \"max_retries\", \"extra_headers\", \"parallel_tool_calls\", \"audio\", \"response_format\", \"user\"]}}, \"additional_usage_values\": {\"completion_tokens_details\": null, \"prompt_tokens_details\": null}, \"user_api_key\": null, \"user_api_key_alias\": null, \"user_api_key_team_id\": null, \"user_api_key_project_id\": null, \"user_api_key_project_alias\": null, \"user_api_key_org_id\": null, \"user_api_key_user_id\": null, \"user_api_key_team_alias\": null, \"spend_logs_metadata\": null, \"requester_ip_address\": null, \"user_agent\": null, \"status\": null, \"proxy_server_request\": null, \"error_information\": null, \"attempted_retries\": null, \"max_retries\": null}", "cache_key": "Cache OFF", "spend": 0.00022500000000000002, "total_tokens": 30, diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index c1847a52760..b8af0f0c374 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -4,13 +4,15 @@ import json import math import re import time -from collections.abc import Iterator +from collections.abc import Generator, Iterator +from contextlib import closing from dataclasses import dataclass from itertools import chain from types import MappingProxyType from typing import Final from urllib.parse import parse_qs, urlsplit +import httpx import pytest from fastapi import FastAPI from fastapi.testclient import TestClient @@ -19,16 +21,21 @@ from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter from litellm.constants import OTLP_MAX_ATTRIBUTE_VALUE_BYTES from litellm.rust_bridge._native import NativeTraceConfig, NativeTraceStorage from litellm.rust_bridge.trace.generated.models import ActivityAvailability, LensAccessParams, TraceQueryHelp -from litellm.rust_bridge.trace.generated.types import TraceScope +from litellm.rust_bridge.trace.generated.types import Trace, TraceScope from litellm.rust_bridge.trace.storage import ClickHouseStorage, TraceStorageConfig, span_rows from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError from litellm.tracing.types import SpendLogRecord from scripts.seed_tracing_fixtures import ( TRACE, TRACE_FIXTURES, + Copies, FixtureReplay, + bulk_span_rows, + copied_trace_id, + copy_clickhouse, fixture_capture, fixture_replays, + long_sessions, rebase_spend, response_pattern, spend_fixtures, @@ -503,7 +510,7 @@ def seeded_trace_api(clickhouse_url: str) -> Iterator[SeededTraceAPI]: def _fixture_trace_api( clickhouse_url: str, replays: tuple[FixtureReplay, ...], stamped: tuple[SpendLogRecord, ...] -) -> Iterator[SeededTraceAPI]: +) -> Generator[SeededTraceAPI]: from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.tracing_endpoints import provide_receiver, provide_trace_query_secret, router @@ -595,24 +602,34 @@ def test_query_correlation_requires_key_or_user_ownership_within_a_team(seeded_t assert all(row["request_id"] != unrelated["request_id"] for row in matches) -@pytest.fixture(scope="module") -def captured_trace_api() -> Iterator[SeededTraceAPI]: +def _captured_replays( + namespace: str, +) -> tuple[tuple[FixtureReplay, ...], tuple[tuple[str, tuple[SpendLogRecord, ...]], ...]]: captures: Final = spend_fixtures() - originals: Final = tuple(chain.from_iterable(rows for _, rows in captures)) - pattern: Final = response_pattern(originals) - replays: Final = fixture_replays(TRACE_FIXTURES, time.time_ns() // 1_000_000, "captured-api", pattern) + pattern: Final = response_pattern(tuple(chain.from_iterable(rows for _, rows in captures))) + replays: Final = fixture_replays(TRACE_FIXTURES, time.time_ns() // 1_000_000, namespace, pattern) by_name: Final = MappingProxyType(dict(captures)) - paired: Final = tuple( - rebase_spend(by_name[replay.name], replay.offset_ms, replay.namespace, pattern) + return replays, tuple( + ( + replay.name, + tuple( + _stamp(row) for row in rebase_spend(by_name[replay.name], replay.offset_ms, replay.namespace, pattern) + ), + ) for replay in replays if replay.name in by_name ) - stamped: Final[tuple[SpendLogRecord, ...]] = tuple( - {**row, "team_id": "team-a", "api_key": "fixture-key", "user": "fixture-user"} - for row in chain.from_iterable(paired) - ) + + +def _stamp(row: SpendLogRecord) -> SpendLogRecord: + return {**row, "team_id": "team-a", "api_key": "fixture-key", "user": "fixture-user"} + + +@pytest.fixture(scope="module") +def captured_trace_api() -> Iterator[SeededTraceAPI]: + replays, paired = _captured_replays("captured-api") with clickhouse_service() as url: - yield from _fixture_trace_api(url, replays, stamped) + yield from _fixture_trace_api(url, replays, tuple(chain.from_iterable(rows for _, rows in paired))) @pytest.mark.parametrize("name", tuple(name for name, _ in spend_fixtures())) @@ -644,3 +661,64 @@ def test_captured_sdk_cost_survives_seeding_and_is_queryable(name: str, captured assert math.isclose(sum(row.spend for row in records), sum(row["spend"] or 0 for row in rows)) assert sum(row.prompt_tokens for row in records) == sum(row["prompt_tokens"] for row in rows) assert sum(row.completion_tokens for row in records) == sum(row["completion_tokens"] for row in rows) + + +def test_server_side_copies_keep_every_capture_linked_to_its_spend() -> None: + replays, paired = _captured_replays("copied-api") + copies: Final = Copies( + trace_ids=tuple(sorted(frozenset(str(span["TraceId"]) for span in bulk_span_rows(replays, Tenant("", ""))))), + request_ids=tuple(row["request_id"] for _, rows in paired for row in rows), + numbers=range(1, 3), + step_ms=60_000, + source="seed-copied-api-", + target="seed-copied-api-c", + ) + (session,) = long_sessions(replays, paired, "seed-copied-api-", "seed-copied-api-c", (3,)) + session_spend: Final = sum(row["spend"] or 0 for row in dict(paired)["openai_agents_swarm"]) + session_spans: Final = len( + span_rows((TRACE_FIXTURES / "openai_agents_swarm.json").read_bytes(), "application/json") + ) + with ( + clickhouse_service() as url, + closing(_fixture_trace_api(url, replays, tuple(chain.from_iterable(rows for _, rows in paired)))) as seeded, + ): + api: Final = next(seeded) + assert api.client.portal is not None + for plan in (copies, session): + api.client.portal.call(_copy_clickhouse, url, plan) + for name, rows in paired: + _assert_capture(api, name, rows, fixture_capture(name, rows[0]).trace_id) + _assert_capture(api, name, rows, copied_trace_id(fixture_capture(name, rows[0]).trace_id, "2")) + trace: Final = _trace(api, copied_trace_id(session.trace_ids[0], session.session)) + assert trace["summary"]["span_count"] == 1 + 3 * (session_spans - 1) + (root,) = (span for span in trace["spans"] if not span["parent_span_id"]) + assert {span["parent_span_id"] for span in trace["spans"] if span["parent_span_id"]} <= { + span["span_id"] for span in trace["spans"] + } + assert root["start_offset_ms"] == min(span["start_offset_ms"] for span in trace["spans"]) + assert root["start_offset_ms"] + root["duration_ms"] >= max( + span["start_offset_ms"] + span["duration_ms"] for span in trace["spans"] + ) + assert trace["summary"]["spend"] == pytest.approx(3 * session_spend) + + +def _trace(api: SeededTraceAPI, trace_id: str) -> Trace: + response: Final = api.client.get(f"/v1/traces/{trace_id}") + assert response.status_code == 200, response.text + return TRACE.validate_json(response.content) + + +def _assert_capture(api: SeededTraceAPI, name: str, rows: tuple[SpendLogRecord, ...], trace_id: str) -> None: + capture: Final = fixture_capture(name, rows[0]) + summary: Final = _trace(api, trace_id)["summary"] + assert summary["span_count"] == len(span_rows((TRACE_FIXTURES / f"{name}.json").read_bytes(), "application/json")) + assert summary["spend"] == ( + pytest.approx(sum(row["spend"] or 0 for row in rows)) + if capture.spend_linked and capture.spend_complete + else None + ) + + +async def _copy_clickhouse(url: str, copies: Copies) -> None: + async with httpx.AsyncClient(base_url=url, params={"database": "trace_test"}) as client: + await copy_clickhouse(client, "trace_test", copies) diff --git a/tests/unit/integrations/otel/test_runtime.py b/tests/unit/integrations/otel/test_runtime.py index 285759c7553..5747be27656 100644 --- a/tests/unit/integrations/otel/test_runtime.py +++ b/tests/unit/integrations/otel/test_runtime.py @@ -8,6 +8,8 @@ import lock. These tests pin the import to a single resolution. """ import builtins +import importlib.abc +import sys import litellm.integrations.otel.runtime as runtime @@ -69,3 +71,25 @@ def test_phase_event_no_ops_when_runtime_absent(monkeypatch): assert runtime.phase_event("litellm.request.body_parsed") is None assert runtime.phase_event("litellm.request.body_received", {"litellm.request.body_bytes": 3}) is None + + +def test_phase_span_does_not_import_the_proxy_in_an_sdk_process(monkeypatch): + import litellm.proxy + + monkeypatch.delitem(sys.modules, "litellm.proxy.proxy_server", raising=False) + monkeypatch.delattr(litellm.proxy, "proxy_server", raising=False) + proxy_imports: list[str] = [] + + class _RefuseProxyImport(importlib.abc.MetaPathFinder): + def find_spec(self, fullname, path, target=None): + if fullname == "litellm.proxy.proxy_server": + proxy_imports.append(fullname) + raise ImportError(fullname) + return None + + monkeypatch.setattr(sys, "meta_path", [_RefuseProxyImport(), *sys.meta_path]) + + with runtime.phase_span("route gpt-5-mini") as span: + assert span is None + + assert proxy_imports == [] diff --git a/tests/unit/integrations/test_s3_v2.py b/tests/unit/integrations/test_s3_v2.py index 963586d9532..ac905335f81 100644 --- a/tests/unit/integrations/test_s3_v2.py +++ b/tests/unit/integrations/test_s3_v2.py @@ -17,7 +17,10 @@ import httpx import pytest import respx +import litellm +from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.s3_v2 import S3BatchUploadError, S3Logger +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.types.integrations.s3_v2 import s3BatchLoggingElement from litellm.types.utils import StandardLoggingPayload @@ -5069,3 +5072,60 @@ async def test_send_batch_time_grows_linearly_with_the_batch() -> None: quadrupled: Final = await _timed_send_batch(8_000) assert quadrupled / baseline < 8, f"2k took {baseline:.3f}s, 8k took {quadrupled:.3f}s" + + +class QueueOnlyS3LoggerWithoutCredentialsOrPeriodicFlush(S3Logger): + def __init__(self) -> None: + CustomLogger.__init__(self) + self.log_queue: list[s3BatchLoggingElement] = [] + self.batch_size = 100 + self.s3_strip_base64_files = False + self.s3_use_team_prefix = False + self.s3_use_key_prefix = False + self.s3_path = "" + self.s3_log_prompts_only = None + self.s3_batch_file_upload = False + self._upload_semaphore = asyncio.Semaphore(1) + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("concurrent", [False, True]) +@pytest.mark.asyncio +async def test_repeated_failure_notifications_upload_once( + monkeypatch: pytest.MonkeyPatch, + stream: bool, + concurrent: bool, +) -> None: + sink: Final = QueueOnlyS3LoggerWithoutCredentialsOrPeriodicFlush() + upload: Final = AsyncMock(return_value=True) + monkeypatch.setattr(sink, "async_upload_data_to_s3", upload) + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "_async_failure_callback", [sink]) + logging_obj: Final = Logging( + model="openai/gpt-4o-mini", + messages=[], + stream=stream, + call_type="acompletion", + start_time=datetime.now(), + litellm_call_id="synthetic-request", + function_id="synthetic-function", + ) + error: Final = litellm.ServiceUnavailableError( + message="synthetic deployment-selection failure", + llm_provider="openai", + model="openai/gpt-4o-mini", + ) + + if concurrent: + await asyncio.gather( + *(logging_obj.async_failure_handler(error, "synthetic traceback") for _ in range(3)) + ) + else: + for _ in range(3): + await logging_obj.async_failure_handler(error, "synthetic traceback") + + assert len(sink.log_queue) == 1 + assert all(entry.payload["status"] == "failure" for entry in sink.log_queue) + await sink.async_send_batch() + await asyncio.sleep(0) + assert upload.await_count == 1 diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index 4b25da2ff79..c726127f892 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -1540,11 +1540,161 @@ def test_logging_prevent_double_logging(logging_obj): This is to avoid double logging. """ logging_obj.stream = False - logging_obj.has_run_logging(event_type="sync_success") - assert logging_obj.should_run_logging(event_type="sync_success") == False - assert logging_obj.should_run_logging(event_type="sync_failure") == True - assert logging_obj.should_run_logging(event_type="async_success") == True - assert logging_obj.should_run_logging(event_type="async_failure") == True + logging_obj.mark_logging_complete(event_type="sync_success") + assert logging_obj.should_run_logging(event_type="sync_success") is False + assert logging_obj.should_run_logging(event_type="sync_failure") is True + assert logging_obj.should_run_logging(event_type="async_success") is True + assert logging_obj.should_run_logging(event_type="async_failure") is True + + +class _FailureCountingLogger(CustomLogger): + def __init__(self) -> None: + super().__init__() + self.async_failure_events: asyncio.Queue[None] = asyncio.Queue() + self.sync_failure_events: asyncio.Queue[None] = asyncio.Queue() + + async def async_log_failure_event( + self, + kwargs: dict[str, object], + response_obj: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self.async_failure_events.put_nowait(None) + + def log_failure_event( + self, + kwargs: dict[str, object], + response_obj: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self.sync_failure_events.put_nowait(None) + + +def _register_failure_counting_logger( + monkeypatch: pytest.MonkeyPatch, + logger: _FailureCountingLogger, +) -> None: + monkeypatch.setattr(litellm, "callbacks", []) + monkeypatch.setattr(litellm, "success_callback", []) + monkeypatch.setattr(litellm, "_async_success_callback", []) + monkeypatch.setattr(litellm, "failure_callback", [logger]) + monkeypatch.setattr(litellm, "_async_failure_callback", [logger]) + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize( + "event", + ["async_success", "sync_success", "async_failure", "sync_failure"], +) +def test_mark_logging_complete_flags_by_stream_and_event( + logging_obj: LitellmLogging, + stream: bool, + event: Literal["async_success", "sync_success", "async_failure", "sync_failure"], +) -> None: + events: Final[tuple[Literal["async_success", "sync_success", "async_failure", "sync_failure"], ...]] = ( + "async_success", + "sync_success", + "async_failure", + "sync_failure", + ) + logging_obj.stream = stream + + logging_obj.mark_logging_complete(event_type=event) + + expected_to_run: Final = stream and event in ("async_success", "sync_success") + assert logging_obj.should_run_logging(event_type=event) is expected_to_run + assert all(logging_obj.should_run_logging(event_type=other_event) for other_event in events if other_event != event) + + +@pytest.mark.parametrize("stream", [False, True]) +def test_has_run_logging_alias_marks_logging_complete( + logging_obj: LitellmLogging, + stream: bool, +) -> None: + logging_obj.stream = stream + + logging_obj.has_run_logging(event_type="async_failure") + assert logging_obj.should_run_logging(event_type="async_failure") is False + + logging_obj.has_run_logging(event_type="async_success") + assert logging_obj.should_run_logging(event_type="async_success") is stream + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("concurrent", [False, True]) +async def test_async_failure_handler_dispatches_once_on_repeated_notifications( + logging_obj: LitellmLogging, + monkeypatch: pytest.MonkeyPatch, + stream: bool, + concurrent: bool, +) -> None: + failure_logger: Final = _FailureCountingLogger() + _register_failure_counting_logger(monkeypatch, failure_logger) + logging_obj.stream = stream + logging_obj.model_call_details["litellm_params"] = {} + error: Final = litellm.ServiceUnavailableError("503", "anthropic", "claude") + failure_calls: Final = tuple(logging_obj.async_failure_handler(error, "tb") for _ in range(3)) + + if concurrent: + await asyncio.gather(*failure_calls) + else: + for failure_call in failure_calls: + await failure_call + + assert failure_logger.async_failure_events.qsize() == 1 + + +def test_sync_failure_handler_dispatches_once_on_repeated_streaming_notifications( + logging_obj: LitellmLogging, + monkeypatch: pytest.MonkeyPatch, +) -> None: + failure_logger: Final = _FailureCountingLogger() + _register_failure_counting_logger(monkeypatch, failure_logger) + logging_obj.stream = True + logging_obj.call_type = "completion" + logging_obj.model_call_details["litellm_params"] = {} + error: Final = litellm.ServiceUnavailableError("503", "anthropic", "claude") + + for _ in range(3): + logging_obj.failure_handler(error, "tb") + + assert failure_logger.sync_failure_events.qsize() == 1 + + +@pytest.mark.asyncio +async def test_streaming_sync_and_async_failure_dedupe_independently( + logging_obj: LitellmLogging, + monkeypatch: pytest.MonkeyPatch, +) -> None: + failure_logger: Final = _FailureCountingLogger() + _register_failure_counting_logger(monkeypatch, failure_logger) + logging_obj.stream = True + logging_obj.call_type = "completion" + logging_obj.model_call_details["litellm_params"] = {} + error: Final = litellm.ServiceUnavailableError("503", "anthropic", "claude") + + await logging_obj.async_failure_handler(error, "tb") + await logging_obj.async_failure_handler(error, "tb") + logging_obj.failure_handler(error, "tb") + logging_obj.failure_handler(error, "tb") + + assert failure_logger.async_failure_events.qsize() == 1 + assert failure_logger.sync_failure_events.qsize() == 1 + + +def test_streaming_success_is_not_marked_complete(logging_obj: LitellmLogging) -> None: + logging_obj.stream = True + + logging_obj.mark_logging_complete(event_type="async_success") + logging_obj.mark_logging_complete(event_type="sync_success") + + assert logging_obj.should_run_logging(event_type="async_success") is True + assert logging_obj.should_run_logging(event_type="sync_success") is True + assert "has_logged_async_success" not in logging_obj.model_call_details + assert "has_logged_sync_success" not in logging_obj.model_call_details @pytest.mark.asyncio diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 09921a9a85c..85e4c055f07 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -2489,6 +2489,28 @@ class TestAnthropicMessagesHandlerStreamingScanKey: assert ended_key.tool_calls_in_flight is False assert ended_key != open_key + def test_released_stream_as_ended_keys_the_tool_use_the_client_already_received(self): + handler = AnthropicMessagesHandler() + tool_use = self._sse( + "content_block_start", + { + "type": "content_block_start", + "index": 1, + "content_block": {"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}}, + }, + ) + stopped_key = handler.get_streaming_scan_key([self._text_delta("hi"), tool_use, self._stop("tool_use")]) + released_key = handler.get_streaming_scan_key( + handler.released_stream_as_ended([self._text_delta("hi"), tool_use]) + ) + assert released_key.stream_ended is True + assert released_key == stopped_key + + def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self): + released = (self._text_delta("hi"), self._text_delta(" there")) + ended = AnthropicMessagesHandler().released_stream_as_ended(released) + assert ended == released and all(a is b for a, b in zip(ended, released, strict=True)) + class PerRowTextGuardrail(CustomGuardrail): """Answers one redacted text per chat row it was shown, the way a guardrail diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 5c85faa5e13..5d71cbe9afc 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -2336,3 +2336,32 @@ class TestStreamingScanKey: handler = OpenAIChatCompletionsHandler() key = handler.get_streaming_scan_key([self._chunk("hi"), b"data: [DONE]"]) assert key.texts == ("hi",) + + def test_released_stream_as_ended_finishes_only_the_choice_whose_tool_call_was_in_flight(self): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + released = ( + self._chunk("hi", index=0), + ModelResponseStream(choices=[StreamingChoices(index=1, delta=Delta(tool_calls=[tool_call]))]), + ) + handler = OpenAIChatCompletionsHandler() + ended = handler.released_stream_as_ended(released) + assert all(a is b for a, b in zip(ended[:-1], released, strict=True)) + assert [(choice.index, choice.finish_reason) for choice in ended[-1].choices] == [(1, "tool_calls")] + ended_key = handler.get_streaming_scan_key(ended) + assert ended_key.stream_ended is True and len(ended_key.tool_calls) == 1, ended_key + + def test_released_stream_as_ended_leaves_a_stream_with_no_tool_call_in_flight_as_released(self): + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + tool_call = {"index": 0, "id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}} + text_only = (self._chunk("hi", index=0), self._chunk(" there", index=1)) + finished_tool_call = ( + ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(tool_calls=[tool_call]))]), + self._chunk(None, finish_reason="tool_calls", index=0), + ) + handler = OpenAIChatCompletionsHandler() + for released in (text_only, finished_tool_call): + ended = handler.released_stream_as_ended(released) + assert len(ended) == len(released) and all(a is b for a, b in zip(ended, released, strict=True)) diff --git a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py index 87980a47f87..8bd242308f7 100644 --- a/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/unit/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -6,6 +6,7 @@ with guardrail transformations. """ import copy +import json from collections.abc import Callable from typing import Any, Final, List, Literal, Optional, Tuple from unittest.mock import AsyncMock, MagicMock, patch @@ -3654,3 +3655,85 @@ class TestOpenAIResponsesHandlerStreamingScanKey: ended_key = handler.get_streaming_scan_key([self._delta(0, "hi"), added, self._completed(3, [function_call])]) assert ended_key.tool_calls_in_flight is False assert len(ended_key.tool_calls) == 1 + + def test_released_stream_as_ended_keys_the_tool_call_the_client_already_received(self): + handler = OpenAIResponsesHandler() + added = { + "type": "response.output_item.added", + "sequence_number": 1, + "item": {"type": "function_call", "id": "fc_1", "call_id": "call_1", "name": "get_weather", "arguments": ""}, + } + arguments_delta = { + "type": "response.function_call_arguments.delta", + "sequence_number": 2, + "item_id": "fc_1", + "delta": '{"city": "Paris"', + } + ended_key = handler.get_streaming_scan_key( + handler.released_stream_as_ended([self._delta(0, "hi"), added, arguments_delta]) + ) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + assert len(ended_key.tool_calls) == 1 and "Paris" in ended_key.tool_calls[0], ended_key + + def test_released_stream_as_ended_leaves_a_text_only_stream_as_released(self): + released = (self._delta(0, "hi"), self._delta(1, " there")) + ended = OpenAIResponsesHandler().released_stream_as_ended(released) + assert ended == released and all(a is b for a, b in zip(ended, released, strict=True)) + + @staticmethod + def _finished_function_call(sequence_number: int, item_id: str, city: str) -> tuple[dict[str, object], ...]: + arguments = json.dumps({"city": city}) + pending = {"type": "function_call", "id": item_id, "call_id": "call_" + item_id, "name": "get_weather"} + return ( + {"type": "response.output_item.added", "sequence_number": sequence_number, "item": {**pending, "arguments": ""}}, + { + "type": "response.function_call_arguments.delta", + "sequence_number": sequence_number + 1, + "item_id": item_id, + "delta": arguments, + }, + { + "type": "response.output_item.done", + "sequence_number": sequence_number + 2, + "item": {**pending, "arguments": arguments, "status": "completed"}, + }, + ) + + def test_released_stream_as_ended_keys_every_tool_call_finished_before_the_disconnect(self): + handler = OpenAIResponsesHandler() + released = ( + self._delta(0, "hi"), + *self._finished_function_call(1, "fc_1", "Paris"), + *self._finished_function_call(4, "fc_2", "Rome"), + ) + ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended(released)) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + cities = tuple(city for fingerprint in ended_key.tool_calls for city in ("Paris", "Rome") if city in fingerprint) + assert cities == ("Paris", "Rome"), ended_key + + def test_released_stream_as_ended_keys_a_message_whose_item_already_finished(self): + handler = OpenAIResponsesHandler() + message_done = { + "type": "response.output_item.done", + "sequence_number": 1, + "item": {"type": "message", "id": "msg_1", "content": [{"type": "output_text", "text": "hi"}]}, + } + ended_key = handler.get_streaming_scan_key(handler.released_stream_as_ended((self._delta(0, "hi"), message_done))) + assert ended_key.stream_ended is True + assert ended_key.texts == ("hi",) + + @pytest.mark.asyncio + async def test_scan_of_a_stream_released_through_a_finished_tool_call_covers_its_text_too(self): + handler = OpenAIResponsesHandler() + guardrail = MockRecordingGuardrail(guardrail_name="test") + released = (self._delta(0, "hi"), *self._finished_function_call(1, "fc_1", "Paris")) + await handler.process_output_streaming_response( + responses_so_far=list(handler.released_stream_as_ended(released)), + guardrail_to_apply=guardrail, + request_data={}, + ) + assert [inputs.get("texts") for inputs in guardrail.seen_inputs] == [["hi"]], guardrail.seen_inputs + tool_calls = guardrail.seen_inputs[0].get("tool_calls") or [] + assert [call["function"]["arguments"] for call in tool_calls] == ['{"city": "Paris"}'], tool_calls diff --git a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py index 5d6d47b8c38..088b232f3b6 100644 --- a/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py +++ b/tests/unit/proxy/_experimental/mcp_server/outbound_credentials/test_sso_assertion_store.py @@ -210,18 +210,18 @@ async def test_persist_and_fetch_round_trip_encrypted_at_rest(): stored = {} prisma = _make_prisma(stored) token = _make_id_token() - assertion = assertion_from_sso_login(token, "rt_1") + assertion = assertion_from_sso_login(token, "refresh.token") with patch("litellm.proxy.proxy_server.prisma_client", prisma): await persist_sso_identity_assertion("user-a", assertion) fetched = await fetch_sso_identity_assertion("user-a") assert fetched is not None assert fetched.id_token.get_secret_value() == token assert fetched.refresh_token is not None - assert fetched.refresh_token.get_secret_value() == "rt_1" + assert fetched.refresh_token.get_secret_value() == "refresh.token" assert fetched.issuer == assertion.issuer assert fetched.expires_at == assertion.expires_at assert token not in stored["user-a"] - assert "rt_1" not in stored["user-a"] + assert "refresh.token" not in stored["user-a"] decrypted = decrypt_value_helper(stored["user-a"], "test", exception_type="debug") assert json.loads(decrypted)["id_token"] == token diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py index 545b2757ffd..556fe0b2536 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_server_tool_calls_and_headers.py @@ -9191,8 +9191,8 @@ async def test_fire_mcp_tool_call_logging_iserror_logs_failure(): tool_error = logging_obj.async_failure_handler.await_args.args[0] assert isinstance(tool_error, MCPToolResultError) assert str(tool_error) == "upstream exploded" - logging_obj.has_run_logging.assert_any_call(event_type="sync_success") - logging_obj.has_run_logging.assert_any_call(event_type="async_success") + logging_obj.mark_logging_complete.assert_any_call(event_type="sync_success") + logging_obj.mark_logging_complete.assert_any_call(event_type="async_success") proxy_logging_mock.post_call_failure_hook.assert_awaited_once() hook_kwargs = proxy_logging_mock.post_call_failure_hook.await_args.kwargs assert hook_kwargs["route"] == "/mcp/call_tool" @@ -9768,7 +9768,7 @@ async def test_call_mcp_tool_modern_interim_result_passes_through_without_comple logging_obj.async_post_mcp_tool_call_hook.assert_not_awaited() proxy_logging_mock.post_mcp_call_hook.assert_not_awaited() proxy_logging_mock.post_call_failure_hook.assert_not_awaited() - assert sorted(c.kwargs["event_type"] for c in logging_obj.has_run_logging.call_args_list) == [ + assert sorted(c.kwargs["event_type"] for c in logging_obj.mark_logging_complete.call_args_list) == [ "async_success", "sync_success", ], "the @client wrapper would otherwise log the interim result as a completed success on return" diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index 6b1d8bc3a2c..d0d52dd6566 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -2218,9 +2218,9 @@ def test_proxy_admin_viewer_can_access_audit_logs(route): # layer, even though the underlying handlers already gate on PROXY_ADMIN_VIEW_ONLY. # # Each route below corresponds to a network call made by the Logs page -# (ui/litellm-dashboard/src/components/view_logs/) — see the comment on each. +# (ui/litellm-dashboard/src/components/logs/) — see the comment on each. ADMIN_VIEWER_LOGS_PAGE_ROUTES = [ - # Main paginated log list — uiSpendLogsCall in log_filter_logic.tsx & index.tsx + # Main paginated log list — uiSpendLogsCall in request/useLogFilterLogic.ts & index.tsx "/spend/logs/ui", # Single-log detail drawer — fetched on row click in LogDetailsDrawer "/spend/logs/ui/abc-request-id", diff --git a/tests/unit/proxy/common_utils/test_fips.py b/tests/unit/proxy/common_utils/test_fips.py new file mode 100644 index 00000000000..0e8138856c0 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_fips.py @@ -0,0 +1,106 @@ +import pytest + +from litellm.proxy.common_utils.fips import ( + FipsModeError, + FipsModeOff, + FipsModeOn, + MalformedFipsMode, + ProviderDoesNotEnforceFips, + TlsVerificationDisabled, + enforce_fips_boot_verdict, + fips_boot_verdict, + is_fips_mode, + parse_fips_mode, +) + + +def _verdict( + raw: str | None, + *, + enforcing: bool = True, + ssl_env: str | None = None, + ssl_setting: object = True, +): + return fips_boot_verdict( + raw_fips_mode=raw, + provider_enforces_fips=lambda: enforcing, + ssl_verify_environment=ssl_env, + ssl_verify_setting=ssl_setting, + ) + + +@pytest.mark.parametrize("raw", [None, "false", "0", "no", "off", "", " False "]) +def test_unset_and_false_spellings_leave_fips_mode_off(raw): + assert parse_fips_mode(raw) == FipsModeOff() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is False + + +@pytest.mark.parametrize("raw", ["true", "1", "yes", "on", " TRUE "]) +def test_true_spellings_turn_fips_mode_on(raw): + assert parse_fips_mode(raw) == FipsModeOn() + assert is_fips_mode({"LITELLM_FIPS_MODE": raw}.get) is True + + +@pytest.mark.parametrize("raw", ["enforced", "2", "strict", "yes please"]) +def test_anything_else_is_malformed_and_refused_with_the_offending_value(raw): + assert parse_fips_mode(raw) == MalformedFipsMode(value=raw) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict(raw, enforcing=False), announce=lambda _: None) + assert f"LITELLM_FIPS_MODE={raw} is not a boolean" in str(refused.value) + assert "true or false" in str(refused.value) + + +def test_off_never_consults_the_provider_or_tls_settings(): + def explode() -> bool: + raise AssertionError("provider probe must not run while FIPS mode is off") + + verdict = fips_boot_verdict( + raw_fips_mode=None, provider_enforces_fips=explode, ssl_verify_environment="false", ssl_verify_setting=False + ) + assert verdict == FipsModeOff() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when off")) + + +def test_on_with_an_enforcing_provider_and_verified_tls_boots(): + verdict = _verdict("true", enforcing=True, ssl_env="true", ssl_setting="/etc/ssl/certs/ca.pem") + assert verdict == FipsModeOn() + enforce_fips_boot_verdict(verdict, announce=lambda _: pytest.fail("nothing to announce when on")) + + +def test_on_with_a_non_enforcing_provider_is_refused_and_names_the_fix(): + announced = [] + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(_verdict("true", enforcing=False), announce=announced.append) + message = str(refused.value) + assert message.startswith("LiteLLM proxy refused to start") + assert "LITELLM_FIPS_MODE is on but this Python does not enforce FIPS" in message + assert "FIPS image" in message + assert announced == [f"\n{message}\n\n"] + + +@pytest.mark.parametrize( + "ssl_env, ssl_setting, sources", + [ + ("false", True, ("SSL_VERIFY",)), + (" FALSE ", True, ("SSL_VERIFY",)), + (None, False, ("litellm_settings.ssl_verify",)), + (None, "False", ("litellm_settings.ssl_verify",)), + ("false", False, ("SSL_VERIFY", "litellm_settings.ssl_verify")), + ], +) +def test_disabled_tls_verification_is_refused_naming_every_source(ssl_env, ssl_setting, sources): + verdict = _verdict("true", enforcing=True, ssl_env=ssl_env, ssl_setting=ssl_setting) + assert verdict == TlsVerificationDisabled(sources=sources) + with pytest.raises(FipsModeError) as refused: + enforce_fips_boot_verdict(verdict, announce=lambda _: None) + assert "TLS certificate verification is disabled by " + " and ".join(sources) in str(refused.value) + + +@pytest.mark.parametrize("ssl_setting", [True, "true", "/etc/ssl/certs/ca.pem", None, "", "0", "no"]) +def test_verified_or_custom_bundle_tls_settings_are_not_treated_as_disabled(ssl_setting): + assert _verdict("true", enforcing=True, ssl_setting=ssl_setting) == FipsModeOn() + + +def test_disabled_tls_is_reported_before_the_provider_so_operators_see_config_mistakes_first(): + assert _verdict("true", enforcing=False, ssl_env="false") == TlsVerificationDisabled(sources=("SSL_VERIFY",)) + assert _verdict("true", enforcing=False) == ProviderDoesNotEnforceFips() diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py index 45ad336368c..0cae9344160 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/test_bedrock_guardrails.py @@ -5707,6 +5707,41 @@ async def test_unbuffered_end_of_stream_hook_yields_chunks_before_scan(): assert len(chunk_events) == 3 +@pytest.mark.asyncio +async def test_unbuffered_end_of_stream_hook_scans_released_chunks_when_the_client_closes_early(): + guardrail = BedrockGuardrail( + guardrail_name="bedrock-audit-mode", + guardrailIdentifier="test-id", + guardrailVersion="DRAFT", + event_hook=GuardrailEventHooks.post_call, + default_on=True, + streaming_buffer_until_moderated=False, + streaming_end_of_stream_only=True, + ) + scans = [] + + async def record_scan(*args, **kwargs): + scans.append(kwargs["source"]) + return {"action": "NONE", "assessments": [], "outputs": []} + + async def mock_stream(): + yield _chat_chunk("Hello", None) + yield _chat_chunk(" world", None) + yield _chat_chunk("", "stop") + + with patch.object(guardrail, "make_bedrock_api_request", AsyncMock(side_effect=record_scan)): + stream = guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(), + response=mock_stream(), + request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hi"}]}, + ) + first = await stream.__anext__() + await stream.aclose() + + assert first.choices[0].delta.content == "Hello" + assert scans == ["OUTPUT"] + + @pytest.mark.asyncio async def test_buffered_default_hook_scans_before_any_chunk(): guardrail = BedrockGuardrail( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index e547575ef9c..e797304a1c6 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1,16 +1,21 @@ """Tests for unified guardrail.""" +import asyncio +import contextlib import io import logging +from collections.abc import AsyncGenerator, AsyncIterable, AsyncIterator from types import SimpleNamespace from typing import TYPE_CHECKING, Final, Literal +import anyio import pytest import litellm from litellm.caching import DualCache from litellm.integrations.custom_guardrail import ( CustomGuardrail, + ModifyResponseException, log_guardrail_information, ) from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route @@ -42,7 +47,15 @@ from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrai ) from litellm.types.guardrails import GuardrailEventHooks from litellm.types.llms.openai import ResponsesAPIResponse -from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + CallTypes, + Delta, + GenericGuardrailAPIInputs, + ModelResponse, + ModelResponseStream, + StandardLoggingGuardrailInformation, + StreamingChoices, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -2164,6 +2177,554 @@ class _ScanCountingGuardrail(CustomGuardrail): return inputs +class _GatedScanGuardrail(_ScanCountingGuardrail): + """End-of-stream scan that holds until released, recording scans that finished.""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + self.scan_started = anyio.Event() + self.scan_released = anyio.Event() + self.finished_scans = 0 + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + self.scan_started.set() + await self.scan_released.wait() + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + self.finished_scans += 1 + return recorded + + +class _FinishReasonRecordingGuardrail(_ScanCountingGuardrail): + """Records the finish reasons of the stream handed to each response-side scan""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + self.finish_reasons: tuple[str | None, ...] = () + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + rebuilt = request_data.get("response") + released = request_data.get("responses") + chunks = ( + tuple(chunk for chunk in released if isinstance(chunk, ModelResponseStream)) + if isinstance(released, list) + else () + ) + choices = [ + *(rebuilt.choices if isinstance(rebuilt, ModelResponse) else ()), + *(choice for chunk in chunks for choice in chunk.choices), + ] + self.finish_reasons = (*self.finish_reasons, *(choice.finish_reason for choice in choices)) + return await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + + +class _GatedToolCallGuardrail(_StreamingTextGuardrail): + """Tool-call inspection that holds until released""" + + def __init__(self) -> None: + super().__init__() + self.inspection_started = anyio.Event() + self.inspection_released = anyio.Event() + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + if input_type == "response" and inputs.get("tool_calls"): + self.inspection_started.set() + await self.inspection_released.wait() + return await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + + +class _MarkerBlockingScanGuardrail(_ScanCountingGuardrail): + """Scan-counting guardrail that blocks any scan whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in recorded.get("texts") or []): + raise ModifyResponseException( + message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name + ) + return recorded + + +class _MarkerHttpErrorScanGuardrail(_ScanCountingGuardrail): + """Scan-counting guardrail that raises an HTTPException for any scan whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in recorded.get("texts") or []): + raise unified_module.HTTPException(status_code=400, detail={"error": "Violated guardrail policy"}) + return recorded + + +class _MarkerBlockingStreamingTextGuardrail(_StreamingTextGuardrail): + """incremental_diff guardrail that blocks any round whose text contains BLOCKME""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + transformed = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + if any("BLOCKME" in text for text in inputs.get("texts") or []): + raise ModifyResponseException( + message="blocked", model="gpt-4", request_data=request_data, guardrail_name=self.guardrail_name + ) + return transformed + + +class _DisconnectRewritingGuardrail(_ScanCountingGuardrail): + """End-of-stream guardrail that rewrites every scanned text""" + + def __init__(self) -> None: + super().__init__(end_of_stream_only=True) + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + recorded = await super().apply_guardrail(inputs, request_data, input_type, **kwargs) + return {**recorded, "texts": ["REWRITTEN" for _ in recorded.get("texts") or []]} + + +class _RecordedScanGuardrail(CustomGuardrail): + """End-of-stream scan recorded through log_guardrail_information, returning ``reply`` or raising ``error``""" + + def __init__(self, *, reply: GenericGuardrailAPIInputs | None = None, error: Exception | None = None) -> None: + super().__init__(guardrail_name="recorded-scan") + self.streaming_end_of_stream_only = True + self.streaming_buffer_until_moderated = False + self.guardrail_config = {} + self._reply = reply + self._error = error + + def should_run_guardrail(self, data: dict[str, object], event_type: GuardrailEventHooks) -> bool: + return True + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict[str, object], + input_type: Literal["request", "response"], + **kwargs: object, + ) -> GenericGuardrailAPIInputs: + if self._error is not None: + raise self._error + return inputs if self._reply is None else self._reply + + +def _recorded_guardrail_statuses(request_data: dict[str, object]) -> list[str]: + metadata = request_data["metadata"] + assert isinstance(metadata, dict), request_data + entries: list[StandardLoggingGuardrailInformation] = metadata.get("standard_logging_guardrail_information", []) + return [entry["guardrail_status"] for entry in entries] + + +class TestStreamingClientDisconnectScan: + """A client that reads streamed content and then disconnects must not skip + the end-of-stream scan of what it already received.""" + + @pytest.fixture(autouse=True) + def _use_real_mappings(self, monkeypatch: pytest.MonkeyPatch) -> None: + _patch_translation_mappings(monkeypatch, load_guardrail_translation_mappings()) + + @staticmethod + def _guarded_stream(guardrail: CustomGuardrail, upstream: AsyncIterable[object]) -> AsyncGenerator[object, None]: + return UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream, + request_data={"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}}, + ) + + @pytest.mark.asyncio + async def test_closing_after_released_content_still_scans_it(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert _delta_text(received) == "synthetic secret" + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_mid_text_stream_does_not_hand_the_scan_a_tool_calls_finish(self): + guardrail = _FinishReasonRecordingGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + await stream.__anext__() + await stream.aclose() + + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + assert "tool_calls" not in guardrail.finish_reasons, guardrail.finish_reasons + + @pytest.mark.asyncio + async def test_upstream_cancellation_after_released_content_still_scans_it(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + raise asyncio.CancelledError() + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + with pytest.raises(asyncio.CancelledError): + await stream.__anext__() + + assert _delta_text(received) == "synthetic secret" + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_cancellation_while_the_disconnect_scan_is_in_flight_lets_it_finish(self): + guardrail = _GatedScanGuardrail() + first_chunk_received = anyio.Event() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + await anyio.sleep_forever() + yield _stream_chunk(" tail", finish_reason="stop") + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async with contextlib.aclosing(self._guarded_stream(guardrail, upstream())) as stream: + async for _item in stream: + first_chunk_received.set() + + scopes = [] + with anyio.fail_after(5): + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await first_chunk_received.wait() + scopes[0].cancel() + await guardrail.scan_started.wait() + await anyio.sleep(0) + guardrail.scan_released.set() + + assert guardrail.finished_scans == 1 + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_after_a_sampled_scan_covered_everything_released_does_not_scan_again(self): + guardrail = _ScanCountingGuardrail(sampling_rate=1) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic") + yield _stream_chunk(" secret") + await anyio.sleep_forever() + + stream = self._guarded_stream(guardrail, upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + scanned_texts = [scan["texts"] for scan in guardrail.scans] + assert "".join(_delta_text(chunk) for chunk in released) == "synthetic secret" + assert scanned_texts[-1] == ["synthetic secret"], scanned_texts + assert len(scanned_texts) == len({tuple(texts) for texts in scanned_texts}), scanned_texts + + @pytest.mark.asyncio + async def test_closing_before_any_content_is_released_does_not_scan(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True, buffer_until_moderated=True) + upstream_started = anyio.Event() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("withheld") + upstream_started.set() + await anyio.sleep_forever() + yield _stream_chunk("never", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + async with anyio.create_task_group() as task_group: + task_group.start_soon(stream.__anext__) + await upstream_started.wait() + task_group.cancel_scope.cancel() + + assert guardrail.scans == () + + @pytest.mark.asyncio + async def test_cancellation_during_end_of_stream_scan_lets_the_scan_finish(self): + guardrail = _GatedScanGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async for _item in self._guarded_stream(guardrail, upstream()): + pass + + scopes = [] + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await guardrail.scan_started.wait() + scopes[0].cancel() + await anyio.sleep(0) + guardrail.scan_released.set() + + assert guardrail.finished_scans == 1 + assert [scan["texts"] for scan in guardrail.scans] == [["synthetic secret tail"]], guardrail.scans + + @staticmethod + async def _close_after_first_chunk(guardrail: CustomGuardrail) -> dict[str, object]: + request_data: dict[str, object] = {"guardrail_to_apply": guardrail, "model": "gpt-4", "metadata": {}} + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = UnifiedLLMGuardrails().async_post_call_streaming_iterator_hook( + user_api_key_dict=UserAPIKeyAuth(api_key="test-key", request_route="/v1/chat/completions"), + response=upstream(), + request_data=request_data, + ) + received = await stream.__anext__() + await stream.aclose() + assert _delta_text(received) == "synthetic secret" + return request_data + + @pytest.mark.asyncio + async def test_disconnect_scan_that_fails_after_the_verdict_records_the_failure(self): + request_data = await self._close_after_first_chunk( + _RecordedScanGuardrail(reply={"texts": ["synthetic secret", "unmatched extra text"]}) + ) + + assert _recorded_guardrail_statuses(request_data) == ["success", "guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_disconnect_scan_whose_guardrail_raises_records_one_failure(self): + request_data = await self._close_after_first_chunk(_RecordedScanGuardrail(error=RuntimeError("provider down"))) + + assert _recorded_guardrail_statuses(request_data) == ["guardrail_failed_to_respond"] + + @pytest.mark.asyncio + async def test_disconnect_scan_that_passes_records_only_the_verdict(self): + request_data = await self._close_after_first_chunk(_RecordedScanGuardrail()) + + assert _recorded_guardrail_statuses(request_data) == ["success"] + + @staticmethod + async def _tool_call_upstream() -> AsyncIterator[ModelResponseStream]: + from litellm.types.utils import ChatCompletionDeltaToolCall, Function + + tool_call = ChatCompletionDeltaToolCall( + id="call_1", index=0, type="function", function=Function(name="get_weather", arguments='{"city": "Paris"}') + ) + yield ModelResponseStream( + choices=[StreamingChoices(index=0, delta=Delta(content=None, tool_calls=[tool_call]))] + ) + yield _stream_chunk(None, finish_reason="tool_calls") + + @pytest.mark.asyncio + async def test_cancellation_during_incremental_diff_tool_call_inspection_lets_it_finish_once(self): + guardrail = _GatedToolCallGuardrail() + + async def consume(scope_ready: list[anyio.CancelScope]) -> None: + with anyio.CancelScope() as scope: + scope_ready.append(scope) + async with contextlib.aclosing(self._guarded_stream(guardrail, self._tool_call_upstream())) as stream: + async for _item in stream: + pass + + scopes = [] + async with anyio.create_task_group() as task_group: + task_group.start_soon(consume, scopes) + await guardrail.inspection_started.wait() + scopes[0].cancel() + await anyio.sleep(0) + guardrail.inspection_released.set() + + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_the_incremental_diff_tool_call_inspection_does_not_inspect_again(self): + guardrail = _StreamingTextGuardrail() + + stream = self._guarded_stream(guardrail, self._tool_call_upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + assert released[-1].choices[0].finish_reason == "tool_calls" + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_a_released_tool_call_under_incremental_diff_still_inspects_it(self): + guardrail = _StreamingTextGuardrail() + + stream = self._guarded_stream(guardrail, self._tool_call_upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"] + assert [[call["function"]["name"] for call in calls] for calls in guardrail.received_tool_calls] == [ + ["get_weather"] + ], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_while_a_mid_stream_block_is_delivered_does_not_scan_the_blocked_content_again(self): + guardrail = _MarkerBlockingScanGuardrail(sampling_rate=2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + for text in ("a", "b", "c", "BLOCKME"): + yield _stream_chunk(text) + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.__anext__() + await stream.aclose() + + assert received == ["a", "b", "c"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_while_a_mid_stream_guardrail_error_is_delivered_does_not_scan_the_content_again(self): + guardrail = _MarkerHttpErrorScanGuardrail(sampling_rate=2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + for text in ("a", "b", "c", "BLOCKME"): + yield _stream_chunk(text) + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.__anext__() + await stream.aclose() + + assert received == ["a", "b", "c"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"], ["abcBLOCKME"]], guardrail.scans + + @pytest.mark.asyncio + async def test_closing_while_the_scanned_final_chunk_is_delivered_does_not_scan_the_stream_again(self): + guardrail = _ScanCountingGuardrail(end_of_stream_only=True) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("a") + yield _stream_chunk("b") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()) for _ in range(3)] + await stream.aclose() + + assert received == ["a", "b", " tail"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab tail"]], guardrail.scans + + @pytest.mark.asyncio + async def test_cancellation_with_a_withheld_window_scans_only_the_released_chunks(self): + guardrail = _ScanCountingGuardrail(sampling_rate=2, buffer_until_moderated=True) + guardrail.streaming_buffer_release_on_scan = True + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("a") + yield _stream_chunk("b") + yield _stream_chunk("WITHHELD") + raise asyncio.CancelledError() + + stream = self._guarded_stream(guardrail, upstream()) + received = [_delta_text(await stream.__anext__()), _delta_text(await stream.__anext__())] + with pytest.raises(asyncio.CancelledError): + await stream.__anext__() + + assert received == ["a", "b"] + assert [scan["texts"] for scan in guardrail.scans] == [["ab"]], guardrail.scans + + @pytest.mark.asyncio + async def test_disconnect_scan_that_rewrites_text_leaves_the_released_chunks_untouched(self): + stream = self._guarded_stream(_DisconnectRewritingGuardrail(), self._text_then_tail_upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert _delta_text(received) == "synthetic secret" + + @staticmethod + async def _text_then_tail_upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("synthetic secret") + yield _stream_chunk(" tail", finish_reason="stop") + + @pytest.mark.asyncio + async def test_closing_while_an_incremental_diff_block_is_delivered_does_not_inspect_the_tool_call_again(self): + guardrail = _MarkerBlockingStreamingTextGuardrail() + + async def upstream() -> AsyncIterator[ModelResponseStream]: + async for tool_chunk in self._tool_call_upstream(): + if tool_chunk.choices[0].finish_reason is None: + yield tool_chunk + yield _stream_chunk("BLOCKME") + yield _stream_chunk(" tail", finish_reason="stop") + + stream = self._guarded_stream(guardrail, upstream()) + released = [await stream.__anext__(), await stream.__anext__()] + await stream.aclose() + + assert [call.function.name for call in released[0].choices[0].delta.tool_calls] == ["get_weather"] + assert guardrail.received_texts == [["BLOCKME"]], guardrail.received_texts + assert guardrail.received_tool_calls == [], guardrail.received_tool_calls + + @pytest.mark.asyncio + async def test_closing_after_a_released_tool_call_under_incremental_diff_scans_no_held_back_text(self): + guardrail = _StreamingTextGuardrail(holdback_schedule=[len("held secret")] * 2) + + async def upstream() -> AsyncIterator[ModelResponseStream]: + yield _stream_chunk("held secret") + async for tool_chunk in self._tool_call_upstream(): + yield tool_chunk + + stream = self._guarded_stream(guardrail, upstream()) + received = await stream.__anext__() + await stream.aclose() + + assert [call.function.name for call in received.choices[0].delta.tool_calls] == ["get_weather"] + assert guardrail.received_tool_calls, guardrail.received_texts + assert all("held secret" not in text for text in guardrail.received_texts[-1]), guardrail.received_texts + + def _responses_delta(sequence_number, text): return { "type": "response.output_text.delta", diff --git a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py index 9c785d59830..31f1c054c59 100644 --- a/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py +++ b/tests/unit/proxy/hooks/test_async_post_call_streaming_iterator_hook.py @@ -7,6 +7,7 @@ Verifies that the hook: 3. Actually yields chunks from async generators """ +import logging from typing import AsyncGenerator, Any from unittest.mock import MagicMock, patch @@ -31,7 +32,7 @@ class MockStreamingCallback(CustomLogger): self, user_api_key_dict: UserAPIKeyAuth, response: AsyncGenerator[Any, None], - request_data: dict, + request_data: dict[str, object], ) -> AsyncGenerator[Any, None]: """Transform chunks by tracking and optionally prefixing.""" async for chunk in response: @@ -185,3 +186,340 @@ async def test_streaming_hook_propagates_callback_errors(): with pytest.raises(RuntimeError, match="Callback failed!"): async for _ in result: pass + + +class CleanupRecordingCallback(CustomLogger): + """Iterator hook whose cleanup marks when it ran.""" + + def __init__(self): + super().__init__() + self.cleaned_up = False + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> AsyncGenerator[Any, None]: + try: + async for chunk in response: + yield chunk + finally: + self.cleaned_up = True + + +@pytest.mark.asyncio +async def test_closing_the_stream_runs_every_callback_cleanup_before_returning(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callbacks = [CleanupRecordingCallback(), CleanupRecordingCallback()] + + with patch.object(litellm, "callbacks", callbacks): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + first = await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert first == {"choices": [{"delta": {"content": "Hello"}}]} + assert [callback.cleaned_up for callback in callbacks] == [True, True] + + +class RaisingCleanupCallback(CustomLogger): + """Iterator hook whose cleanup raises.""" + + async def async_post_call_streaming_iterator_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> AsyncGenerator[Any, None]: + try: + async for chunk in response: + yield chunk + finally: + raise RuntimeError("cleanup failed") + + +@pytest.mark.asyncio +async def test_closing_the_stream_still_cleans_up_inner_callbacks_when_an_outer_cleanup_raises(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + inner = CleanupRecordingCallback() + + with patch.object(litellm, "callbacks", [inner, RaisingCleanupCallback()]): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert inner.cleaned_up is True + + +class _PlainAsyncIterator: + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + self._response = response + + def __aiter__(self) -> "_PlainAsyncIterator": + return self + + async def __anext__(self) -> Any: + return await self._response.__anext__() + + +class PlainIteratorCallback(CustomLogger): + """Iterator hook that returns an async iterator with no aclose.""" + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a plain async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _PlainAsyncIterator: + return _PlainAsyncIterator(response) + + +@pytest.mark.asyncio +async def test_a_hook_returning_a_plain_async_iterator_streams_every_chunk(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + + with patch.object(litellm, "callbacks", [PlainIteratorCallback()]): + ProxyLogging._callback_capabilities_cache.clear() + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + + +class _ClosableAsyncIterator(_PlainAsyncIterator): + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + super().__init__(response) + self.closed = False + + async def aclose(self) -> None: + self.closed = True + + +class _SyncClosableAsyncIterator(_PlainAsyncIterator): + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + super().__init__(response) + self.closed = False + + def aclose(self) -> None: + self.closed = True + + +class ClosableIteratorCallback(CustomLogger): + """Iterator hook that returns a non-generator async iterator with its own aclose.""" + + def __init__( + self, + iterator_type: type[_ClosableAsyncIterator] | type[_SyncClosableAsyncIterator] = _ClosableAsyncIterator, + ) -> None: + super().__init__() + self.iterator_type = iterator_type + self.returned: tuple[_ClosableAsyncIterator | _SyncClosableAsyncIterator, ...] = () + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _ClosableAsyncIterator | _SyncClosableAsyncIterator: + iterator = self.iterator_type(response) + self.returned = (*self.returned, iterator) + return iterator + + +@pytest.mark.asyncio +async def test_closing_the_stream_closes_a_hook_iterator_that_is_not_a_generator(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + stream = proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + await stream.__anext__() + await stream.aclose() + ProxyLogging._callback_capabilities_cache.clear() + + assert [iterator.closed for iterator in callback.returned] == [True] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_with_a_synchronous_aclose_streams_everything_and_is_closed(): + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback(iterator_type=_SyncClosableAsyncIterator) + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + + +class _RaisingAcloseIterator(_ClosableAsyncIterator): + """Non-generator async iterator whose asynchronous aclose raises.""" + + def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None: + super().__init__(response) + self.error = error + + async def aclose(self) -> None: + self.closed = True + raise self.error + + +class _SyncRaisingAcloseIterator(_SyncClosableAsyncIterator): + """Non-generator async iterator whose synchronous aclose raises.""" + + def __init__(self, response: AsyncGenerator[Any, None], error: Exception) -> None: + super().__init__(response) + self.error = error + + def aclose(self) -> None: + self.closed = True + raise self.error + + +class RaisingAcloseIteratorCallback(CustomLogger): + """Iterator hook that returns a non-generator async iterator whose aclose raises.""" + + def __init__( + self, + iterator_type: type[_RaisingAcloseIterator] | type[_SyncRaisingAcloseIterator] = _RaisingAcloseIterator, + error: Exception | None = None, + ) -> None: + super().__init__() + self.iterator_type = iterator_type + self.error = error if error is not None else RuntimeError("cleanup failed") + self.returned: tuple[_RaisingAcloseIterator | _SyncRaisingAcloseIterator, ...] = () + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _RaisingAcloseIterator | _SyncRaisingAcloseIterator: + iterator = self.iterator_type(response, self.error) + self.returned = (*self.returned, iterator) + return iterator + + +@pytest.mark.asyncio +async def test_a_hook_iterator_whose_aclose_raises_still_finishes_the_stream(caplog: pytest.LogCaptureFixture) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = RaisingAcloseIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + warnings_emitted = [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage() + ] + assert len(warnings_emitted) == 1 + assert "RuntimeError" in warnings_emitted[0] + assert "cleanup failed" not in warnings_emitted[0] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_whose_synchronous_aclose_raises_still_finishes_the_stream( + caplog: pytest.LogCaptureFixture, +) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = RaisingAcloseIteratorCallback( + iterator_type=_SyncRaisingAcloseIterator, error=ValueError("sync cleanup failed") + ) + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + warnings_emitted = [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "RaisingAcloseIteratorCallback" in record.getMessage() + ] + assert len(warnings_emitted) == 1 + assert "ValueError" in warnings_emitted[0] + + +@pytest.mark.asyncio +async def test_a_hook_iterator_with_a_clean_aclose_streams_everything_without_warning( + caplog: pytest.LogCaptureFixture, +) -> None: + proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + callback = ClosableIteratorCallback() + + with patch.object(litellm, "callbacks", [callback]): + ProxyLogging._callback_capabilities_cache.clear() + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + received = [ + chunk + async for chunk in proxy_logging.async_post_call_streaming_iterator_hook( + response=mock_streaming_response(), + user_api_key_dict=UserAPIKeyAuth(api_key="test_key"), + request_data={"model": "gpt-4", "messages": []}, + ) + ] + ProxyLogging._callback_capabilities_cache.clear() + + assert [chunk async for chunk in mock_streaming_response()] == received + assert [iterator.closed for iterator in callback.returned] == [True] + assert not [ + record.getMessage() + for record in caplog.records + if record.levelname == "WARNING" and "ClosableIteratorCallback" in record.getMessage() + ] diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index a0d01878277..54dbaac1e42 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -18,7 +18,7 @@ from litellm.proxy.lens.endpoints import ( watching, worker_supports_model, ) -from litellm.proxy.lens.models import Lens, LensSettings, RunRequest, Scope +from litellm.proxy.lens.models import ActivitySelection, Lens, LensSettings, RunRequest, Scope @pytest.fixture @@ -281,6 +281,25 @@ def test_model_errors_reach_worker_with_status_and_redacted_provider_message(pro assert error.headers == {"retry-after": "60"} +@pytest.mark.asyncio +async def test_preview_samples_a_selection_without_investigation_settings() -> None: + from litellm.proxy.lens.endpoints import Preview, preview_sample + + class SelectionStorage: + async def lens_sample(self, parameters): + assert (parameters.source, parameters.agent_name, parameters.selected_team) == ("requests", "billing", "t1") + assert parameters.preview == 1 and parameters.offset == 3 + return [] + + body: Final = Preview.model_validate( + {"selection": {"source": "requests", "agent_name": "billing", "team_id": "t1"}, "offset": 3} + ) + sample: Final = await preview_sample( + body, UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), SelectionStorage() + ) + assert sample.eligible == 0 and not sample.executions + + @pytest.mark.asyncio async def test_preview_reports_calendar_overflow_as_a_validation_error() -> None: from datetime import datetime, timezone @@ -288,7 +307,7 @@ async def test_preview_reports_calendar_overflow_as_a_validation_error() -> None from litellm.proxy.lens.endpoints import Preview, preview_sample body: Final = Preview( - settings=LensSettings(name="Calendar regression", model="analysis", context="Read recorded activity"), + selection=ActivitySelection(), as_of=datetime.min.replace(tzinfo=timezone.utc), ) with pytest.raises(HTTPException) as error: @@ -323,7 +342,9 @@ async def test_unknown_gateway_release_refuses_registration_and_claims(monkeypat monkeypatch.setenv("LITELLM_RELEASE_TAG", "") monkeypatch.setenv("LENS_WORKER_IMAGE", "registry.example/lens-worker:old") with pytest.raises(HTTPException) as registration_error: - await register_worker(WorkerName(analysis_key_id="a" * 64), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)) + await register_worker( + WorkerName(analysis_key_id="a" * 64), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + ) assert registration_error.value.status_code == 503 assert "LITELLM_RELEASE_TAG" in registration_error.value.detail with pytest.raises(HTTPException) as claim_error: diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index 8485c286a30..8dc6c0477b0 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -26,6 +26,7 @@ from litellm.constants import ( CLIENT_REQUESTED_MODEL_SCOPE_KEY, MAX_LITELLM_CALL_ID_LENGTH, RETURN_RAW_MODEL_NAME_METADATA_KEY, + STREAM_SSE_KEEPALIVE_PING_BYTES, ) from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.opentelemetry import UserAPIKeyAuth @@ -4256,6 +4257,72 @@ class TestStreamCloseOnDisconnect: assert upstream.aclosed + async def test_async_streaming_data_generator_closes_the_guardrail_chain_on_client_disconnect( + self, + ): + cleanup_ran = [] + + async def guarded_chain(**_kwargs): + try: + yield {"type": "chunk"} + yield {"type": "chunk"} + finally: + cleanup_ran.append(True) + + proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) + proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain + gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=MagicMock(), + user_api_key_dict=MagicMock(spec=UserAPIKeyAuth), + request_data={"model": "mock-model"}, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=lambda c: "data: x\n\n", + serialize_error=lambda e: "data: error\n\n", + ) + + await gen.__anext__() + await gen.aclose() + + assert cleanup_ran == [True] + + async def test_async_streaming_data_generator_refunds_the_budget_when_closing_the_guardrail_chain_raises( + self, + ): + async def guarded_chain(**_kwargs): + try: + yield STREAM_SSE_KEEPALIVE_PING_BYTES + yield {"type": "chunk"} + finally: + raise RuntimeError("cleanup failed") + + proxy_logging_obj = ProxyLogging(user_api_key_cache=MagicMock()) + proxy_logging_obj.async_post_call_streaming_iterator_hook = guarded_chain + reservation = object() + user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + user_api_key_dict.budget_reservation = reservation + gen = ProxyBaseLLMRequestProcessing.async_streaming_data_generator( + response=MagicMock(), + user_api_key_dict=user_api_key_dict, + request_data={"model": "mock-model"}, + proxy_logging_obj=proxy_logging_obj, + serialize_chunk=lambda c: "data: x\n\n", + serialize_error=lambda e: "data: error\n\n", + ) + + released: list[object] = [] + + async def record_release(budget_reservation: object) -> None: + released.append(budget_reservation) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation_on_cancel", + new=record_release, + ): + await gen.__anext__() + await gen.aclose() + + assert released == [reservation] + async def test_async_streaming_data_generator_redacts_internal_details_on_error( self, ): diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index 42a9dc441dd..e2fbd5dcd98 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -1771,6 +1771,108 @@ async def test_proxy_startup_refuses_an_unsafe_master_key_even_when_the_database assert ("could not be checked" in announced[0]) == key_can_have_encrypted_the_database +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_this_python_does_not_enforce_fips(monkeypatch, tmp_path): + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + _, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: False) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert len(announced) == 1 + assert "does not enforce FIPS" in announced[0] + + +@pytest.mark.asyncio +async def test_proxy_startup_refuses_fips_mode_when_the_config_disables_tls_verification(monkeypatch, tmp_path): + import yaml + from fastapi import FastAPI + + from litellm.proxy.common_utils.fips import FipsModeError + from litellm.proxy.proxy_server import proxy_startup_event + + config_path, announced = _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + config_path.write_text( + yaml.dump({"general_settings": {"master_key": "sk-a-safe-master-key"}, "litellm_settings": {"ssl_verify": False}}) + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setattr(litellm, "ssl_verify", True) + monkeypatch.setenv("LITELLM_FIPS_MODE", "true") + + with pytest.raises(FipsModeError): + async with proxy_startup_event(FastAPI()): + pass + + assert "TLS certificate verification is disabled by litellm_settings.ssl_verify" in announced[0] + + +class _PrismaClientWhoseUserTableCannotHash: + class _Table: + async def find_many(self, where): + raise ValueError("[digital envelope routines] unsupported") + + class _Db: + litellm_usertable = None + + def __init__(self, database_url, proxy_logging_obj): + self.db = self._Db() + self.db.litellm_usertable = self._Table() + self.writer_db = self.db + + async def connect(self): + pass + + async def disconnect(self): + pass + + def start_view_setup_task(self): + pass + + async def check_view_exists(self): + pass + + async def health_check(self): + pass + + +@pytest.mark.asyncio +@pytest.mark.parametrize("fips_mode", ["true", "false"]) +async def test_proxy_startup_surfaces_a_password_migration_crypto_failure(monkeypatch, tmp_path, caplog, fips_mode): + from fastapi import FastAPI + + from litellm.proxy.proxy_server import proxy_startup_event + + _boot_with_general_settings(monkeypatch, tmp_path, {"master_key": "sk-a-safe-master-key"}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setenv("DATABASE_URL", "postgresql://nobody:nothing@127.0.0.1:1/unreachable") + monkeypatch.setattr("litellm.proxy.proxy_server.PrismaClient", _PrismaClientWhoseUserTableCannotHash) + monkeypatch.setattr("litellm.proxy.proxy_server.openssl_enforces_fips", lambda: True) + monkeypatch.setenv("LITELLM_FIPS_MODE", fips_mode) + + with caplog.at_level(logging.ERROR, logger="LiteLLM Proxy"): + if fips_mode == "true": + with pytest.raises(ValueError, match="digital envelope routines"): + async with proxy_startup_event(FastAPI()): + pass + else: + async with proxy_startup_event(FastAPI()): + await asyncio.sleep(0) + + failures = [r.getMessage() for r in caplog.records if "Password migration failed" in r.getMessage()] + assert len(failures) == 1 + assert "plaintext passwords stay unhashed" in failures[0] + assert "digital envelope routines" in failures[0] + + class _DatabaseWithOneStoredCredential: def __init__(self, ciphertext): self._ciphertext = ciphertext @@ -7323,6 +7425,69 @@ async def test_async_data_generator_cleanup_on_early_exit(): mock_response.aclose.assert_awaited_once() +def _guarded_chain_logging(chain): + from litellm.proxy.utils import ProxyLogging + + proxy_logging = MagicMock(spec=ProxyLogging) + proxy_logging.async_post_call_streaming_iterator_hook = chain + proxy_logging.async_post_call_streaming_hook = AsyncMock(side_effect=lambda **kwargs: kwargs.get("response")) + proxy_logging.post_call_failure_hook = AsyncMock() + return proxy_logging + + +@pytest.mark.asyncio +async def test_async_data_generator_closes_the_guardrail_chain_before_returning_on_client_disconnect(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + + cleanup_ran = [] + + async def guarded_chain(**_kwargs): + try: + yield {"choices": [{"delta": {"content": "Hello"}}]} + yield {"choices": [{"delta": {"content": " world"}}]} + finally: + cleanup_ran.append(True) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)): + gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"}) + first_chunk = await gen.__anext__() + await gen.aclose() + + assert first_chunk.startswith("data: ") + assert cleanup_ran == [True] + + +@pytest.mark.asyncio +async def test_async_data_generator_closes_the_guardrail_chain_while_a_keepalive_read_is_pending(): + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.proxy_server import async_data_generator + + cleanup_ran = [] + never_arrives = asyncio.Event() + + async def guarded_chain(**_kwargs): + try: + yield {"choices": [{"delta": {"content": "Hello"}}]} + await never_arrives.wait() + yield {"choices": [{"delta": {"content": " world"}}]} + finally: + cleanup_ran.append(True) + + with ( + patch.object(litellm, "sse_keepalive_ping_interval_seconds", 1.0), + patch("litellm.proxy.proxy_server.proxy_logging_obj", _guarded_chain_logging(guarded_chain)), + ): + gen = async_data_generator(MagicMock(), MagicMock(spec=UserAPIKeyAuth), {"model": "gpt-4o-mini"}) + first_chunk = await gen.__anext__() + heartbeat = await gen.__anext__() + await gen.aclose() + + assert first_chunk.startswith("data: ") + assert heartbeat == ": ping\n\n" + assert cleanup_ran == [True] + + @pytest.mark.asyncio async def test_async_data_generator_uses_direct_stream_fast_path_without_callbacks(): """ diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index cff67d1f57d..516a0415545 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -12,6 +12,7 @@ import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient +from litellm.constants import TRACE_READ_RETRY_AFTER_SECONDS from litellm.proxy import tracing_endpoints from litellm.proxy._types import LitellmUserRoles, ProxyLifespanState, UserAPIKeyAuth from litellm.proxy.auth.authorization import OwnedRows, ReadScope @@ -19,6 +20,7 @@ from litellm.proxy.auth.authorization_dependencies import get_log_team_lookup from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.tracing_runtime import manage_tracing, provide_storage from litellm.rust_bridge import loader +from litellm.rust_bridge.trace.errors import TraceChanged from litellm.rust_bridge.trace.generated.models import TraceQueryHelp from litellm.rust_bridge.trace.generated.types import AllQueryScope, TraceScope from litellm.rust_bridge.trace.queries import TraceSQLResponse @@ -297,19 +299,45 @@ def test_trace_detail_passes_scoped_reference(client, receiver, suffix, cursor, ), ) @pytest.mark.parametrize( - "error,status,message", + "error,status,code,message", ( - (RuntimeError("private database details"), 503, "Traces are temporarily unavailable. Please try again."), - (OverflowError("private query details"), 413, "Trace is too large for this view. Use a filtered trace query."), + ( + RuntimeError("private database details"), + 503, + "unavailable", + "Traces are temporarily unavailable. Please try again.", + ), + ( + OverflowError("private query details"), + 413, + "too_large", + "Trace is too large for this view. Use a filtered trace query.", + ), + ( + TraceChanged("Trace changed while paging; refresh the trace to continue"), + 409, + "trace_changed", + "Trace changed while paging; refresh the trace to continue", + ), + (ValueError("Invalid span cursor"), 400, "invalid_request", "Invalid span cursor"), ), ) -def test_read_failures_are_actionable_without_exposing_database_details( - client: TestClient, receiver: MagicMock, path: str, method: str, error: Exception, status: int, message: str +def test_read_failures_carry_a_code_per_kind_without_exposing_database_details( + client: TestClient, + receiver: MagicMock, + path: str, + method: str, + error: Exception, + status: int, + code: str, + message: str, ) -> None: getattr(receiver, method).side_effect = error response: Final = client.get(path) assert response.status_code == status - assert response.json() == {"detail": message} + assert response.json() == {"detail": {"code": code, "message": message}} + retry_after: Final = response.headers.get("Retry-After") + assert (retry_after == str(TRACE_READ_RETRY_AFTER_SECONDS)) == (status == 503), retry_after @pytest.mark.parametrize("query", ("page_size=0", "page_size=501", "cursor=" + "x" * 513)) @@ -519,12 +547,13 @@ def test_lens_reads_from_the_lifespan_storage() -> None: with TestClient(app) as client: response: Final = client.post( "/lens/preview/sample", - json={"settings": {"name": "Review", "model": "analysis", "context": "Find failed executions"}}, + json={"selection": {"source": "requests", "service": "checkout"}}, ) assert response.status_code == 200, response.text assert response.json()["executions"] == [] storage.lens_sample.assert_awaited_once() - assert storage.lens_sample.await_args.args[0].all_teams == 1 + params: Final = storage.lens_sample.await_args.args[0] + assert (params.all_teams, params.source, params.service, params.preview) == (1, "requests", "checkout", 1) def test_lens_reads_from_injected_storage_without_receiver() -> None: @@ -541,12 +570,14 @@ def test_lens_reads_from_injected_storage_without_receiver() -> None: with TestClient(app) as client: response: Final = client.post( "/lens/preview/sample", - json={"settings": {"name": "Review", "model": "analysis", "context": "Find failed executions"}}, + json={"selection": {"source": "requests", "service": "checkout"}}, ) assert response.status_code == 200, response.text assert response.json()["executions"] == [] storage.lens_sample.assert_awaited_once() + params: Final = storage.lens_sample.await_args.args[0] + assert (params.source, params.service, params.preview) == ("requests", "checkout", 1) @pytest.mark.parametrize( diff --git a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py index 8bd9dc0df8a..dbb354526d8 100644 --- a/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py +++ b/tests/unit/proxy/utils/proxy_logging/test_guardrail_pipeline.py @@ -11,9 +11,9 @@ from __future__ import annotations import asyncio import json +from collections.abc import AsyncGenerator, Iterator from copy import deepcopy import logging -from collections.abc import Iterator from typing import Any, Callable, Dict, List from unittest.mock import AsyncMock, MagicMock, patch @@ -26,6 +26,7 @@ from litellm.integrations.custom_guardrail import ( CustomGuardrail, ModifyResponseException, ) +from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.prometheus import PrometheusLogger from litellm.llms.base_llm.guardrail_translation.utils import stream_item_field from litellm.proxy._types import UserAPIKeyAuth @@ -2933,3 +2934,62 @@ async def test_streaming_iterator_hook_pipeline_discards_dropped_tool_call_on_re assert delivered == _responses_function_call_events() assert any("'gr-post'" in message and "discarded" in message for message in _warnings(caplog)) + + +class _RaisingAcloseIterator: + """Non-generator async iterator whose aclose raises after the stream ended.""" + + def __init__(self, response: AsyncGenerator[Any, None]) -> None: + self._response = response + + def __aiter__(self) -> "_RaisingAcloseIterator": + return self + + async def __anext__(self) -> Any: + return await self._response.__anext__() + + async def aclose(self) -> None: + raise RuntimeError("cleanup failed") + + +class RaisingAcloseCallback(CustomLogger): + """Iterator hook returning a non-generator async iterator whose aclose raises.""" + + def async_post_call_streaming_iterator_hook( # pyright: ignore[reportIncompatibleMethodOverride] # a custom async iterator worked before aclose handling + self, + user_api_key_dict: UserAPIKeyAuth, + response: AsyncGenerator[Any, None], + request_data: dict[str, object], + ) -> _RaisingAcloseIterator: + return _RaisingAcloseIterator(response) + + +@pytest.mark.asyncio +async def test_streaming_iterator_hook_pipeline_releases_buffered_content_when_a_callback_aclose_raises( + proxy_logging: ProxyLogging, + make_user_api_key_auth: Callable[..., UserAPIKeyAuth], + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + monkeypatch.setattr( + litellm, "callbacks", [_rewriting_stream_guardrail(lambda inputs: {}), RaisingAcloseCallback()] + ) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False) + data = _post_call_pipeline_data(stream=True) + chunks = _tool_call_stream_chunks() + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + delivered = [ + item + async for item in proxy_logging.async_post_call_streaming_iterator_hook( + user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"), + response=_async_chunk_iter(chunks), + request_data=data, + ) + ] + + assert [chunk.model_dump() for chunk in delivered] == [chunk.model_dump() for chunk in chunks] + assert any( + "RaisingAcloseCallback" in message and "RuntimeError" in message and "cleanup failed" not in message + for message in _warnings(caplog) + ) diff --git a/tests/unit/test_lens_dev.py b/tests/unit/test_lens_dev.py index af61734731e..bc73695f898 100644 --- a/tests/unit/test_lens_dev.py +++ b/tests/unit/test_lens_dev.py @@ -194,12 +194,11 @@ def test_seed_only_with_no_cli_count_preserves_env_controls(tmp_path: Path) -> N proc = _run( tmp_path, "parse_args --seed-only; master_key=sk-local; py() { " - 'printf "%s %s %s\\n" "$LENS_DEV_SEED_COPIES" "$LENS_DEV_SEED_BATCH_COPIES" "$@"; }; py=py; seed_data', + 'printf "%s %s\\n" "$LENS_DEV_SEED_COPIES" "$@"; }; py=py; seed_data', LENS_DEV_SEED_COPIES="3", - LENS_DEV_SEED_BATCH_COPIES="1", ) assert proc.returncode == 0, proc.stderr - assert proc.stdout.startswith("3 1 -m") + assert proc.stdout.startswith("3 -m") def test_proxy_uses_this_checkouts_ui_build(tmp_path: Path) -> None: diff --git a/tests/unit/test_seed_tracing_fixtures.py b/tests/unit/test_seed_tracing_fixtures.py index 30cfde9feec..8dac769b3c1 100644 --- a/tests/unit/test_seed_tracing_fixtures.py +++ b/tests/unit/test_seed_tracing_fixtures.py @@ -11,23 +11,22 @@ import pytest from prisma import Json, Prisma from pydantic import InstanceOf, TypeAdapter +from litellm.rust_bridge.trace.queries import TraceSQLResponse from litellm.rust_bridge.trace.storage import Tenant, span_rows from litellm.tracing.types import SpendLogRecord from scripts.seed_tracing_fixtures import ( JSON, TRACE_FIXTURES, - TenantIdentity, bulk_span_rows, fixture_capture, fixture_replays, postgres_row, rebase, rebase_spend, - replay_batches, response_ids, response_pattern, seed_arguments, - seed_batch, + seed_copy, seed_id, spend_fixtures, timestamps, @@ -196,29 +195,39 @@ def test_bulk_export_preserves_all_spans_and_disjoint_copy_ids() -> None: @pytest.mark.requires_rust_extension @pytest.mark.asyncio -async def test_bulk_seed_stamps_authenticated_tenant_and_writes_both_stores() -> None: - from litellm.rust_bridge.trace.storage import ClickHouseStorage, Tenant +async def test_first_copy_stamps_the_authenticated_tenant_and_writes_both_stores( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.rust_bridge.trace.storage import ClickHouseStorage + monkeypatch.setenv("LITELLM_MASTER_KEY", "sk-local") fixtures: Final = spend_fixtures() pattern: Final = response_pattern(tuple(chain.from_iterable(rows for _, rows in fixtures))) - replays: Final = fixture_replays(TRACE_FIXTURES, 1_800_000_000_000, "bulk", pattern) + replays: Final = fixture_replays(TRACE_FIXTURES, 1_800_000_000_000, "first", pattern) storage: Final = AsyncMock(spec=ClickHouseStorage) + storage.query_sql.return_value = TraceSQLResponse.model_validate( + { + "meta": (), + "data": [{"team_id": "local-team", "api_key": "local-hash", "user": "admin"}], + "rows": 1, + "statistics": {"elapsed": 0, "rows_read": 1, "bytes_read": 1}, + } + ) database: Final = AsyncMock(spec=Prisma, litellm_spendlogs=AsyncMock()) - tenant: Final = TenantIdentity(team_id="local-team", api_key="local-hash", user="admin") - async with httpx.AsyncClient() as client: - result: Final = await seed_batch(client, storage, database, replays, fixtures, pattern, tenant, False) - assert result == tenant - trace_table, trace_rows = storage.insert_rows.call_args_list[0].args - assert trace_table == "otel_traces" - assert trace_rows == bulk_span_rows(replays, Tenant("local-team", "local-hash", user_id="admin")) - table, rows = storage.insert_rows.call_args_list[1].args - assert table == "spend_logs" + client: Final = AsyncMock(spec=httpx.AsyncClient) + client.post.return_value = httpx.Response(200, request=httpx.Request("POST", "http://proxy/v1/traces")) + captures: Final = await seed_copy(client, storage, database, replays, fixtures, pattern) + assert tuple(JSON.validate_json(call.kwargs["content"]) for call in client.post.call_args_list) == tuple( + replay.export for replay in replays + ) + rows: Final = tuple(chain.from_iterable(rows for _, rows in captures)) + assert {name for name, _ in captures} == {name for name, _ in fixtures} + assert storage.insert_rows.call_args.args == ("spend_logs", rows) assert len(rows) == sum(len(original) for _, original in fixtures) assert all((row["team_id"], row["api_key"], row["user"]) == ("local-team", "local-hash", "admin") for row in rows) saved: Final = database.litellm_spendlogs.create_many.call_args.kwargs["data"] assert tuple(row["request_id"] for row in saved) == tuple(row["request_id"] for row in rows) assert tuple(row["spend"] for row in saved) == tuple(row["spend"] for row in rows) - storage.query_sql.assert_not_called() def test_seed_cli_rejects_nonpositive_copies() -> None: @@ -228,23 +237,6 @@ def test_seed_cli_rejects_nonpositive_copies() -> None: assert seed_arguments(["--profile", "large", "--copies", "5"]).copies == 5 -def test_bulk_batches_cover_every_copy_including_partial_tail() -> None: - batches: Final = tuple(replay_batches(8, 1_800_000_000_000, "batch", re.compile(r"(?!)"), 3)) - assert tuple(stop for stop, _ in batches) == (4, 7, 8) - expected: Final = tuple( - tuple(fixture_replays(TRACE_FIXTURES, 1_800_000_000_000 - index * 1000, f"batch-{index}", re.compile(r"(?!)"))) - for index in range(1, 8) - ) - assert tuple(chain.from_iterable(replays for _, replays in batches)) == tuple(chain.from_iterable(expected)) - - -def test_seed_cli_rejects_nonpositive_batch_size() -> None: - with pytest.raises(SystemExit) as error: - seed_arguments(["--batch-copies", "0"]) - assert error.value.code == 2 - assert seed_arguments(["--batch-copies", "2"]).batch_copies == 2 - - @pytest.mark.parametrize("timeout", ("0", "-1", "inf", "nan")) def test_seed_cli_rejects_invalid_http_timeouts(timeout: str) -> None: with pytest.raises(SystemExit) as error: diff --git a/ui/litellm-dashboard/.agents/skills/url-state/SKILL.md b/ui/litellm-dashboard/.agents/skills/url-state/SKILL.md new file mode 100644 index 00000000000..986118d81a7 --- /dev/null +++ b/ui/litellm-dashboard/.agents/skills/url-state/SKILL.md @@ -0,0 +1,20 @@ +--- +name: url-state +description: Read or write dashboard URL query state (tabs, filters, pagination, deep links, demo flags) with nuqs instead of raw history or URLSearchParams +--- + +# URL state + +Query params the dashboard reacts to go through nuqs: `useQueryState` or `useQueryStates` with parsers from `"nuqs"`. Do not call `window.history.pushState` or `replaceState`, hand-edit a `URLSearchParams`, or parse `useSearchParams().get(...)` for that state. Raw writes bypass nuqs, so other hooks on the same key go stale and queued nuqs updates can overwrite them + +Validate in the parser, not after reading. Use `parseAsStringLiteral` for enums, `parseAsInteger`, `parseAsBoolean`, `parseAsArrayOf` for lists, `.withDefault` for defaults, and `createParser` when the wire format is custom (for example `demo=1`). Keep the parser map in a module const and share it between every reader and writer of the same key + +Clear a param with `setter(null)`; nuqs keeps unrelated params and the hash. Updates replace history by default, pass `{ history: "push" }` when the change should be a back-button step. A param consumed once on arrival (a deep link or OAuth return flag) is captured in a `useState` initializer from the hook value and then cleared with the setter in a mount effect + +Raw reads are fine for one-time reads that never re-render on URL changes: OAuth callback pages, login, the legacy `?page=` redirect, and `networking.tsx` + +Reuse the existing helpers before adding a hook: `useUrlTab` in `src/hooks/useUrlTab.ts`, `useUrlTableState` in `src/components/shared/DataTable`, and the route modules `src/components/lens/route.ts` and `src/components/logs/request/logDetailRouting.ts` + +In tests, render with `renderWithProviders` from `tests/test-utils.tsx`, passing `searchParams` for the initial URL and `onUrlUpdate` to assert writes. When a test asserts `window.location` directly, wrap the render in `NuqsAdapter` from `nuqs/adapters/react`; writes flush asynchronously, so assert them inside `waitFor`. For inputs bound to URL state, set values with one `fireEvent.change` rather than `user.type`, which flakes when a re-render between keystrokes moves focus + +See the [nuqs docs](https://nuqs.dev/docs) for parser and option details diff --git a/ui/litellm-dashboard/AGENTS.md b/ui/litellm-dashboard/AGENTS.md index e5d876fad84..c890c33ce9b 100644 --- a/ui/litellm-dashboard/AGENTS.md +++ b/ui/litellm-dashboard/AGENTS.md @@ -1,3 +1,5 @@ +For URL query state (tabs, filters, pagination, deep links), follow [.agents/skills/url-state/SKILL.md](.agents/skills/url-state/SKILL.md) + Never put LiteLLM tokens or API keys in `localStorage`. `localStorage` survives browser close. Prefer `httpOnly` cookies, or `sessionStorage` at most, understanding that any web storage is readable by injected scripts (XSS), and only httpOnly cookies are not When you fix lint violations that are grandfathered in `eslint-suppressions.json`, run `eslint . --prune-suppressions` and commit the updated baseline so the gate ratchets down instead of leaving a stale suppression diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 61ab85d613a..74672cf950e 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -2266,12 +2266,12 @@ "count": 1 } }, - "src/components/view_logs/EvalViewer/EvalViewer.tsx": { + "src/components/logs/detail/sections/EvalViewer/EvalViewer.tsx": { "no-nested-ternary": { "count": 1 } }, - "src/components/view_logs/GuardrailViewer/CompliancePanel.tsx": { + "src/components/logs/detail/sections/GuardrailViewer/CompliancePanel.tsx": { "no-nested-ternary": { "count": 2 }, @@ -2279,17 +2279,17 @@ "count": 1 } }, - "src/components/view_logs/GuardrailViewer/ContentFilterDetails.tsx": { + "src/components/logs/detail/sections/GuardrailViewer/ContentFilterDetails.tsx": { "no-nested-ternary": { "count": 1 } }, - "src/components/view_logs/GuardrailViewer/GuardrailViewer.tsx": { + "src/components/logs/detail/sections/GuardrailViewer/GuardrailViewer.tsx": { "no-nested-ternary": { "count": 3 } }, - "src/components/view_logs/LogDetailsDrawer/LogDetailsDrawer.tsx": { + "src/components/logs/detail/LogDetailsDrawer.tsx": { "no-nested-ternary": { "count": 2 }, @@ -2297,22 +2297,22 @@ "count": 2 } }, - "src/components/view_logs/LogDetailsDrawer/useKeyboardNavigation.ts": { + "src/components/logs/detail/useKeyboardNavigation.ts": { "react-hooks/immutability": { "count": 2 } }, - "src/components/view_logs/columns.tsx": { + "src/components/logs/types.ts": { "local/filename-pascal-case": { "count": 1 } }, - "src/components/view_logs/log_filter_logic.tsx": { + "src/components/logs/request/useLogFilterLogic.ts": { "local/filename-pascal-case": { "count": 1 } }, - "src/components/view_logs/logs_utils.tsx": { + "src/components/logs/request/timeRange.ts": { "local/filename-pascal-case": { "count": 1 } diff --git a/ui/litellm-dashboard/eslint.config.mjs b/ui/litellm-dashboard/eslint.config.mjs index ad2923bf2a0..ce2d933292d 100644 --- a/ui/litellm-dashboard/eslint.config.mjs +++ b/ui/litellm-dashboard/eslint.config.mjs @@ -100,7 +100,7 @@ const eslintConfig = [ rules: { "local/no-ad-hoc-z-index": ["error", { allowPopupLayer: true }] }, }, { - files: ["src/components/view_logs/TraceView/**/*.tsx", "src/components/lens/**/*.tsx"], + files: ["src/components/lens/**/*.tsx"], ignores: ["src/**/*.test.tsx"], rules: { "local/no-arbitrary-design-value": "error" }, }, diff --git a/ui/litellm-dashboard/next.config.mjs b/ui/litellm-dashboard/next.config.mjs index d4fdc7af36f..6b67fd70695 100644 --- a/ui/litellm-dashboard/next.config.mjs +++ b/ui/litellm-dashboard/next.config.mjs @@ -13,9 +13,11 @@ const nextConfig = { async rewrites() { return { beforeFiles: [ + // Every dashboard HTTP client sends Accept: application/json; page loads and RSC + // fetches do not. That is what keeps GET /lens (API) apart from /lens (page) in dev. { source: "/:path*", - has: [{ type: "header", key: "content-type", value: "application/json.*" }], + has: [{ type: "header", key: "accept", value: "application/json.*" }], destination: `${devProxyUrl}/:path*`, }, { source: "/ui/:path*", destination: "/:path*" }, diff --git a/ui/litellm-dashboard/package-lock.json b/ui/litellm-dashboard/package-lock.json index 37027f215fe..9aa5764548c 100644 --- a/ui/litellm-dashboard/package-lock.json +++ b/ui/litellm-dashboard/package-lock.json @@ -47,6 +47,7 @@ "react-json-view-lite": "2.5.0", "react-markdown": "9.1.0", "react-reconciler": "0.33.0", + "react-resizable-panels": "4.14.1", "react-syntax-highlighter": "15.6.6", "recharts": "3.9.2", "remark-gfm": "4.0.1", @@ -10649,6 +10650,16 @@ } } }, + "node_modules/react-resizable-panels": { + "version": "4.14.1", + "resolved": "https://registry.npmjs.org/react-resizable-panels/-/react-resizable-panels-4.14.1.tgz", + "integrity": "sha512-OB1bXDNTLcGgTTbaX6Dn5efZhlMboSnkfr1w4xpsJTHPEbRJpACBhXpgRCXozgPO1ztgFeZRUe5J3mFNCVKm5g==", + "license": "MIT", + "peerDependencies": { + "react": "^18.0.0 || ^19.0.0", + "react-dom": "^18.0.0 || ^19.0.0" + } + }, "node_modules/react-syntax-highlighter": { "version": "15.6.6", "resolved": "https://registry.npmjs.org/react-syntax-highlighter/-/react-syntax-highlighter-15.6.6.tgz", diff --git a/ui/litellm-dashboard/package.json b/ui/litellm-dashboard/package.json index 388b6c7eab9..81b89337325 100644 --- a/ui/litellm-dashboard/package.json +++ b/ui/litellm-dashboard/package.json @@ -64,6 +64,7 @@ "react-json-view-lite": "2.5.0", "react-markdown": "9.1.0", "react-reconciler": "0.33.0", + "react-resizable-panels": "4.14.1", "react-syntax-highlighter": "15.6.6", "recharts": "3.9.2", "remark-gfm": "4.0.1", diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx index 2e3b000dfec..c14c2dae55b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.test.tsx @@ -1,9 +1,12 @@ -import { render, screen } from "@testing-library/react"; +import { screen } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; +import type { CreateKeyPrefillData } from "@/components/organisms/create_key_button"; +import { renderWithProviders } from "@/../tests/test-utils"; -const { teamListCall, authorizedSession } = vi.hoisted(() => ({ +const { teamListCall, authorizedSession, createKeyProps } = vi.hoisted(() => ({ teamListCall: vi.fn(() => new Promise(() => {})), authorizedSession: vi.fn(), + createKeyProps: vi.fn(), })); const session = (overrides: { userRole?: string; isViewOnly?: boolean } = {}) => ({ @@ -29,10 +32,6 @@ vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ teamListCall, })); -vi.mock("next/navigation", () => ({ - useSearchParams: () => new URLSearchParams(""), -})); - vi.mock("@/components/VirtualKeysPage/VirtualKeysTable", () => ({ VirtualKeysTable: ({ headerActions }: { headerActions?: React.ReactNode }) => (
@@ -43,7 +42,10 @@ vi.mock("@/components/VirtualKeysPage/VirtualKeysTable", () => ({ })); vi.mock("@/components/organisms/create_key_button", () => ({ - default: () => , + default: (props: { autoOpenCreate?: boolean; prefillData?: CreateKeyPrefillData }) => { + createKeyProps(props); + return ; + }, })); import ApiKeysDashboard from "./ApiKeysDashboard"; @@ -51,27 +53,50 @@ import ApiKeysDashboard from "./ApiKeysDashboard"; describe("ApiKeysDashboard", () => { beforeEach(() => { teamListCall.mockClear(); + createKeyProps.mockClear(); authorizedSession.mockReturnValue(session()); sessionStorage.clear(); }); it("scopes the team list to the signed-in user for non-admin roles", () => { authorizedSession.mockReturnValue(session({ userRole: "Internal User" })); - render(); + renderWithProviders(); expect(teamListCall).toHaveBeenCalledWith("sk-access", 1, 100, { userID: "u-123" }); }); it("renders the keys table with a Create Key action for roles that can write", () => { - render(); + renderWithProviders(); expect(screen.getByRole("table", { name: "Virtual Keys" })).toBeInTheDocument(); expect(screen.getByRole("button", { name: "Create Key" })).toBeInTheDocument(); + expect(createKeyProps).toHaveBeenLastCalledWith( + expect.objectContaining({ autoOpenCreate: false, prefillData: undefined }), + ); + }); + + it("passes validated URL prefill data to the create key action", () => { + renderWithProviders(, { + searchParams: "?create=true&owned_by=bogus&key_type=management&models=a,%20,b&team_id=%20t1%20", + }); + + expect(createKeyProps).toHaveBeenLastCalledWith( + expect.objectContaining({ + autoOpenCreate: true, + prefillData: { + owned_by: undefined, + team_id: "t1", + key_alias: undefined, + models: ["a", "b"], + key_type: "management", + }, + }), + ); }); it("hides Create Key for view-only roles", () => { authorizedSession.mockReturnValue(session({ isViewOnly: true })); - render(); + renderWithProviders(); expect(screen.getByRole("table", { name: "Virtual Keys" })).toBeInTheDocument(); expect(screen.queryByRole("button", { name: "Create Key" })).not.toBeInTheDocument(); @@ -80,7 +105,7 @@ describe("ApiKeysDashboard", () => { it("leaves other pages' session state intact when the tab reloads", () => { sessionStorage.setItem("chatHistory", '[{"role":"user","content":"hi"}]'); sessionStorage.setItem("selectedModel", "gpt-5.5"); - render(); + renderWithProviders(); window.dispatchEvent(new Event("beforeunload")); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx index 915df2f5ded..854a65c842f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/api-keys/ApiKeysDashboard.tsx @@ -4,58 +4,50 @@ import { Page } from "@/components/shared/Page"; import { teamListCall as v2TeamListCall } from "@/app/(dashboard)/hooks/teams/useTeams"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { KeyResponse, Team } from "@/components/key_team_helpers/key_list"; -import CreateKey, { CreateKeyPrefillData } from "@/components/organisms/create_key_button"; +import CreateKey, { type CreateKeyPrefillData } from "@/components/organisms/create_key_button"; import { VirtualKeysTable } from "@/components/VirtualKeysPage/VirtualKeysTable"; -import { useSearchParams } from "next/navigation"; +import { parseAsArrayOf, parseAsBoolean, parseAsString, parseAsStringLiteral, useQueryStates } from "nuqs"; import { useEffect, useMemo, useState } from "react"; +const CREATE_KEY_URL_PARAMS = { + create: parseAsBoolean.withDefault(false), + owned_by: parseAsStringLiteral(["you", "service_account", "another_user"] as const), + team_id: parseAsString, + key_alias: parseAsString, + models: parseAsArrayOf(parseAsString), + key_type: parseAsStringLiteral(["default", "llm_api", "management"] as const), +}; + export default function ApiKeysDashboard() { const { userId: userID, userRole, accessToken, isViewOnly } = useAuthorized(); - const searchParams = useSearchParams()!; + const [{ create, owned_by, team_id, key_alias, models, key_type }] = useQueryStates(CREATE_KEY_URL_PARAMS); const [teams, setTeams] = useState(null); const [keys, setKeys] = useState([]); - const autoOpenCreate = searchParams.get("create") === "true"; + const autoOpenCreate = create; const prefillData: CreateKeyPrefillData | undefined = useMemo(() => { if (!autoOpenCreate) return undefined; - const ownedBy = searchParams.get("owned_by"); - const teamId = searchParams.get("team_id"); - const keyAlias = searchParams.get("key_alias"); - const modelsParam = searchParams.get("models"); - const keyType = searchParams.get("key_type"); - - if (!ownedBy && !teamId && !keyAlias && !modelsParam && !keyType) { + if ([owned_by, team_id, key_alias, models, key_type].every((value) => value === null)) { return undefined; } - const validOwnedByValues = ["you", "service_account", "another_user"]; - const validatedOwnedBy = - ownedBy && validOwnedByValues.includes(ownedBy) ? (ownedBy as CreateKeyPrefillData["owned_by"]) : undefined; - - const validKeyTypes = ["default", "llm_api", "management"]; - const validatedKeyType = - keyType && validKeyTypes.includes(keyType) ? (keyType as CreateKeyPrefillData["key_type"]) : undefined; - - const sanitizedKeyAlias = keyAlias ? keyAlias.trim().slice(0, 256) : undefined; - - const sanitizedModels = modelsParam - ? modelsParam - .split(",") + const sanitizedModels = models + ? models .slice(0, 100) .map((m) => m.trim().slice(0, 256)) .filter((m) => m.length > 0) : undefined; return { - owned_by: validatedOwnedBy, - team_id: teamId?.trim() || undefined, - key_alias: sanitizedKeyAlias, + owned_by: owned_by ?? undefined, + team_id: team_id?.trim() || undefined, + key_alias: key_alias === null ? undefined : key_alias.trim().slice(0, 256), models: sanitizedModels && sanitizedModels.length > 0 ? sanitizedModels : undefined, - key_type: validatedKeyType, + key_type: key_type ?? undefined, }; - }, [searchParams, autoOpenCreate]); + }, [autoOpenCreate, key_alias, key_type, models, owned_by, team_id]); const addKey = (data: KeyResponse) => { setKeys((prevData) => (prevData ? [...prevData, data] : [data])); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx index ba9cfd8ca22..c6900a3b2ab 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingRequestsTable.tsx @@ -11,7 +11,7 @@ import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; -import { LOG_ID_QUERY_PARAM } from "@/components/view_logs/logDetailRouting"; +import { LOG_ID_QUERY_PARAM } from "@/components/logs/request/logDetailRouting"; import type { paths } from "@/lib/http/schema"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { uiHref } from "@/utils/uiHref"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx index dc8a57e50ba..c463a04bf05 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.integration.test.tsx @@ -17,7 +17,7 @@ vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({ useOrganizations: useOrganizationsMock, })); -vi.mock("@/components/view_logs/RequestLogsPanel", () => ({ +vi.mock("@/components/logs/request/RequestLogsPanel", () => ({ default: function RequestLogsPanelMock() { return
; }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx index c7e0a82a363..0e86e8bb9d2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.test.tsx @@ -17,13 +17,13 @@ vi.mock("@/app/(dashboard)/hooks/organizations/useOrganizations", () => ({ useOrganizations: useOrganizationsMock, })); -vi.mock("@/components/view_logs/RequestLogsPanel", () => ({ +vi.mock("@/components/logs/request/RequestLogsPanel", () => ({ default: function RequestLogsPanelMock({ isActive }: { isActive: boolean }) { return
{isActive ? "active" : "inactive"}
; }, })); -vi.mock("@/components/view_logs/AuditLogsPanel", () => ({ +vi.mock("@/components/logs/audit/AuditLogsPanel", () => ({ default: function AuditLogsPanelMock({ isActive }: { isActive: boolean }) { return
{isActive ? "active" : "inactive"}
; }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx index 51603dec29e..a9672b088ba 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/logs/page.tsx @@ -5,8 +5,8 @@ import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import useCan from "@/app/(dashboard)/hooks/useCan"; import DeletedKeysPage from "@/components/DeletedKeysPage/DeletedKeysPage"; import DeletedTeamsPage from "@/components/DeletedTeamsPage/DeletedTeamsPage"; -import AuditLogsPanel from "@/components/view_logs/AuditLogsPanel"; -import RequestLogsPanel from "@/components/view_logs/RequestLogsPanel"; +import AuditLogsPanel from "@/components/logs/audit/AuditLogsPanel"; +import RequestLogsPanel from "@/components/logs/request/RequestLogsPanel"; import { Page, PageTabs, PageTabsList, PageTabsTrigger, PageTabsContent } from "@/components/shared/Page"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx index 2f4fcfd5f69..acf1c24a197 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.test.tsx @@ -1,8 +1,10 @@ import React from "react"; -import { render, waitFor, screen, act, within } from "@testing-library/react"; +import { render as renderWithoutNuqs, waitFor, screen, act, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; -import { describe, it, expect, vi, beforeEach } from "vitest"; +import { afterEach, describe, it, expect, vi, beforeEach } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { NuqsAdapter } from "nuqs/adapters/react"; +import { renderWithProviders as render } from "@/../tests/test-utils"; import MCPServers, { compareServers, type SortKey } from "./mcp_servers"; import type { MCPServer } from "@/components/mcp_tools/types"; import * as networking from "@/components/networking"; @@ -17,6 +19,16 @@ vi.mock("@/components/networking", () => ({ getGeneralSettingsCall: vi.fn().mockResolvedValue([]), updateConfigFieldSetting: vi.fn().mockResolvedValue(undefined), deleteConfigFieldSetting: vi.fn().mockResolvedValue(undefined), + modelHubCall: vi.fn().mockResolvedValue({ data: [] }), + getMCPUserEnvVars: vi.fn((_accessToken: string, serverId: string) => + Promise.resolve({ + server_id: serverId, + required: [{ name: "API_KEY", description: "API key", is_set: false }], + missing_count: 1, + }), + ), + storeMCPUserEnvVars: vi.fn(), + clearMCPUserEnvVars: vi.fn(), listMCPUserEnvVarStatus: vi.fn().mockResolvedValue([]), fetchMCPGatewaySessions: vi.fn(), terminateMCPGatewaySessions: vi.fn(), @@ -140,6 +152,11 @@ describe("MCPServers", () => { beforeEach(() => { vi.clearAllMocks(); stubUiConfig({}); + window.history.replaceState(null, "", "/"); + }); + + afterEach(() => { + window.history.replaceState(null, "", "/"); }); it("should render the MCPServers component with title", async () => { @@ -162,6 +179,41 @@ describe("MCPServers", () => { expect(screen.getByText("MCP Servers")).toBeInTheDocument(); }); + it("opens the user env vars modal from a deep link and removes only that query parameter", async () => { + const server: MCPServer = { + server_id: "deep-link-server", + server_name: "Deep Link Server", + alias: "deep-link-server", + url: "https://example.com/mcp", + created_at: "", + updated_at: "", + created_by: "user", + updated_by: "user", + }; + vi.mocked(networking.fetchMCPServers).mockResolvedValue([server]); + vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue([{ server_id: server.server_id, status: "healthy" }]); + vi.mocked(networking.listMCPUserEnvVarStatus).mockResolvedValue([ + { + server_id: server.server_id, + server_name: server.server_name, + required: [{ name: "API_KEY", description: "API key", is_set: false }], + missing_count: 1, + }, + ]); + window.history.replaceState(null, "", "/?fill_env_vars=deep-link-server&other=1"); + renderWithoutNuqs( + + + + + , + ); + + expect(await screen.findByRole("heading", { name: "Set your credentials" })).toBeVisible(); + expect(within(screen.getByRole("dialog")).getByText("Deep Link Server")).toBeVisible(); + await waitFor(() => expect(window.location.search).toBe("?other=1")); + }); + it.each(["Admin", "Internal User"])("links a %s to their MCP connections page", async (userRole) => { vi.mocked(networking.fetchMCPServers).mockResolvedValue([]); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx index 13510cadf39..682cd5ad3cb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/mcp-servers/_components/mcp_servers.tsx @@ -18,6 +18,7 @@ import { } from "@/components/ui/alert-dialog"; import React, { useEffect, useState, useMemo, useCallback } from "react"; import { useQuery } from "@tanstack/react-query"; +import { useQueryState } from "nuqs"; import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers"; import { useMCPServerHealth } from "@/app/(dashboard)/hooks/mcpServers/useMCPServerHealth"; import { toast } from "@/lib/toast"; @@ -217,11 +218,10 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i const [prefillData, setPrefillData] = useState(null); const [isDeletingServer, setIsDeletingServer] = useState(false); const [byokModalServer, setByokModalServer] = useState(null); + const [fillEnvVarsParam, setFillEnvVarsParam] = useQueryState("fill_env_vars"); // Per-user env-var fill modal target + deep-link source captured once from the URL. const [envVarsModalServer, setEnvVarsModalServer] = useState(null); - const [deepLinkServerId, setDeepLinkServerId] = useState(() => - typeof window === "undefined" ? null : new URLSearchParams(window.location.search).get("fill_env_vars"), - ); + const [deepLinkServerId, setDeepLinkServerId] = useState(() => fillEnvVarsParam); const [searchQuery, setSearchQuery] = useState(""); const [sortKey, setSortKey] = useState("created_desc"); const isInternalUser = userRole === "Internal User"; @@ -251,19 +251,9 @@ const MCPServers: React.FC = ({ accessToken, userRole, userID, i [envVarStatuses], ); - // Deep-link via ?fill_env_vars= — the link users follow from the - // friendly error the proxy returns when a per-user var is missing. The id is - // captured into state above and resolved to a server below; here we only strip - // the param so a refresh doesn't reopen the modal. useEffect(() => { - if (!deepLinkServerId || typeof window === "undefined") return; - const params = new URLSearchParams(window.location.search); - if (!params.has("fill_env_vars")) return; - params.delete("fill_env_vars"); - const newSearch = params.toString(); - const newUrl = window.location.pathname + (newSearch ? `?${newSearch}` : "") + window.location.hash; - window.history.replaceState({}, "", newUrl); - }, [deepLinkServerId]); + if (fillEnvVarsParam !== null) setFillEnvVarsParam(null); + }, [fillEnvVarsParam, setFillEnvVarsParam]); const deepLinkServer = useMemo( () => (deepLinkServerId ? serversWithHealth.find((s) => s.server_id === deepLinkServerId) ?? null : null), diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/system_one.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/system_one.test.ts index 60156d8ef25..c039b018b8f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/system_one.test.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/system_one.test.ts @@ -62,6 +62,7 @@ describe("makeSystemOneRequest", () => { const expectedRequest: Partial = { method: "POST", headers: { + Accept: "application/json", "Content-Type": "application/json", Authorization: "Bearer session-key", }, diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx index 58307bcb869..b16fe5aa7d3 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx @@ -1,10 +1,13 @@ -import { fireEvent, render, screen, waitFor, within } from "@testing-library/react"; +import { fireEvent, render as renderWithoutNuqs, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { afterEach, describe, expect, it, vi } from "vitest"; +import { NuqsAdapter } from "nuqs/adapters/react"; import ObservedROIView from "./ObservedROIView"; import { createObservedDemo } from "./observedDemo"; import type { ObservedSettings, ObservedSnapshot, ObservedStatus } from "./observedData"; +const render = (ui: Parameters[0]) => renderWithoutNuqs(ui, { wrapper: NuqsAdapter }); + const settings: ObservedSettings = { source_provider: "gitlab", api_url: "https://gitlab.com/api/v4", @@ -151,7 +154,7 @@ describe("observed ROI dashboard", () => { expect(screen.getByRole("tab", { name: "Engineers 3", selected: true })).toBeInTheDocument(); expect(screen.queryByRole("button", { name: "Connections" })).not.toBeInTheDocument(); expect(screen.queryByRole("button", { name: "Link accounts" })).not.toBeInTheDocument(); - expect(window.location.search).toBe("?demo=1"); + await waitFor(() => expect(window.location.search).toBe("?demo=1")); await user.click(screen.getByRole("button", { name: "View Alex Rivera's merged changes" })); expect(await screen.findByRole("dialog", { name: "Alex Rivera" })).toHaveTextContent("alex-demo@example.com"); expect(screen.getByRole("heading", { name: "Merged changes" })).toBeInTheDocument(); @@ -172,7 +175,7 @@ describe("observed ROI dashboard", () => { await user.click(screen.getByRole("button", { name: "Exit demo" })); expect(screen.getByRole("heading", { name: "Connect your repositories" })).toBeInTheDocument(); expect(screen.getByRole("button", { name: "Connect GitHub or GitLab" })).toBeEnabled(); - expect(window.location.search).toBe(""); + await waitFor(() => expect(window.location.search).toBe("")); expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); }); @@ -192,12 +195,12 @@ describe("observed ROI dashboard", () => { await user.click(screen.getByRole("button", { name: "Preview sample report" })); expect(screen.queryByRole("button", { name: "Cancel sync" })).not.toBeInTheDocument(); expect(screen.getByRole("button", { name: "2 repositories" })).toBeInTheDocument(); - expect(window.location.search).toBe("?from=review&demo=1"); + await waitFor(() => expect(window.location.search).toBe("?from=review&demo=1")); await user.click(screen.getByRole("button", { name: "Exit demo" })); expect(screen.getByRole("button", { name: "Cancel sync" })).toBeEnabled(); expect(screen.getByRole("button", { name: "1 repository" })).toBeInTheDocument(); expect(screen.getByRole("tab", { name: "Merge requests", selected: true })).toBeInTheDocument(); - expect(window.location.search).toBe("?from=review"); + await waitFor(() => expect(window.location.search).toBe("?from=review")); expect(window.location.hash).toBe("#report"); expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); }); @@ -217,7 +220,7 @@ describe("observed ROI dashboard", () => { await user.click(screen.getByRole("button", { name: "Exit demo" })); expect(screen.queryByText("Alex Rivera")).not.toBeInTheDocument(); expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); - expect(window.location.search).toBe(""); + await waitFor(() => expect(window.location.search).toBe("")); if (state === "failed") expect(await screen.findByRole("alert")).toHaveTextContent("Live data unavailable"); }); @@ -367,7 +370,7 @@ describe("observed ROI dashboard", () => { .queryAllByRole("alert") .map((alert) => alert.textContent), ).toEqual(alerts); - expect(window.location.search).toBe(""); + await waitFor(() => expect(window.location.search).toBe("")); fireEvent.change(within(dialog).getByLabelText("Repositories"), { target: { value: "org/changed" } }); await user.click(within(dialog).getByRole("button", { name: "Save and sync" })); expect(await within(dialog).findByRole("alert")).toHaveTextContent("Provider unavailable"); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx index 29199b67f22..4a9d20fb870 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx @@ -1,6 +1,7 @@ "use client"; import { useEffect, useState } from "react"; +import { parseAsString, useQueryStates } from "nuqs"; import { Link2, RefreshCw } from "lucide-react"; import { apiClient } from "@/components/networking"; import { extractProxyErrorMessage } from "@/lib/http/client"; @@ -14,6 +15,14 @@ import ObservedReport from "./ObservedReport"; import { useObservedReport, type ObservedViewData } from "./useObservedReport"; import { syncMessage, type ObservedSnapshot } from "./observedData"; import { createObservedDemo } from "./observedDemo"; +import { parseAsDemoFlag } from "./demoUrlState"; + +const OBSERVED_ROI_QUERY_PARSERS = { + demo: parseAsDemoFlag, + connected: parseAsString, + connection_cancelled: parseAsString, + connection_failed: parseAsString, +}; function SyncActions({ data, @@ -107,35 +116,25 @@ export default function ObservedROIView({ isViewOnly?: boolean; }) { const { data, error, refresh } = useObservedReport(accessToken); - const [returned] = useState(() => new URLSearchParams(typeof window === "undefined" ? "" : window.location.search)); - const [sample, setSample] = useState(() => - returned.get("demo") === "1" ? createObservedDemo(28) : null, - ); + const [{ demo, connected, connection_cancelled, connection_failed }, setQueryParams] = + useQueryStates(OBSERVED_ROI_QUERY_PARSERS); + const [sample, setSample] = useState(() => (demo === true ? createObservedDemo(28) : null)); const [connections, setConnections] = useState( - ["github", "gitlab"].includes(returned.get("connected") ?? "") || - returned.has("connection_cancelled") || - returned.has("connection_failed"), + ["github", "gitlab"].includes(connected ?? "") || connection_cancelled !== null || connection_failed !== null, ); const [connectionError, setConnectionError] = useState(() => { - if (returned.has("connection_failed")) return "Connection failed or expired. Try again or use a token"; - if (returned.has("connection_cancelled")) return "Connection cancelled. Choose an app or token to try again"; + if (connection_failed !== null) return "Connection failed or expired. Try again or use a token"; + if (connection_cancelled !== null) return "Connection cancelled. Choose an app or token to try again"; return ""; }); - const [afterAuthorization, setAfterAuthorization] = useState(Boolean(returned.get("connected"))); + const [afterAuthorization, setAfterAuthorization] = useState(Boolean(connected)); const [busy, setBusy] = useState(false); const [actionError, setActionError] = useState(""); useEffect(() => { - const url = new URL(window.location.href); - url.searchParams.delete("connected"); - url.searchParams.delete("connection_cancelled"); - url.searchParams.delete("connection_failed"); - window.history.replaceState(window.history.state, "", url); - }, []); + setQueryParams({ connected: null, connection_cancelled: null, connection_failed: null }); + }, [setQueryParams]); function previewSample(enabled: boolean) { - const url = new URL(window.location.href); - if (enabled) url.searchParams.set("demo", "1"); - else url.searchParams.delete("demo"); - window.history.replaceState(window.history.state, "", url); + setQueryParams({ demo: enabled ? true : null }); setSample(enabled ? createObservedDemo(28) : null); } function closeConnections() { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx index ef618dbb02f..30cf5e09eee 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx @@ -1,10 +1,13 @@ import userEvent from "@testing-library/user-event"; -import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { act, fireEvent, render as renderWithoutNuqs, screen, waitFor } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; +import { NuqsAdapter } from "nuqs/adapters/react"; import { apiClient } from "@/components/networking"; import ROICalculatorView from "./ROICalculatorView"; +const render = (ui: Parameters[0]) => renderWithoutNuqs(ui, { wrapper: NuqsAdapter }); + vi.mock("@/components/networking", () => ({ apiClient: { delete: vi.fn(), @@ -520,7 +523,7 @@ describe("ROICalculatorView", () => { fireEvent.click(screen.getByRole("button", { name: "Preview sample report" })); expect(await screen.findByText("You’re viewing demo data")).toBeVisible(); - expect(window.location.search).toBe("?demo=1"); + await waitFor(() => expect(window.location.search).toBe("?demo=1")); expect(screen.getByRole("tab", { name: "Branches" })).toHaveAttribute("aria-selected", "true"); expect(screen.getByRole("searchbox")).toHaveValue(""); expect(screen.getByRole("cell", { name: "$9.10" })).toBeVisible(); @@ -543,7 +546,7 @@ describe("ROICalculatorView", () => { expect(screen.getByText("Improve request routing")).toBeVisible(); expect(screen.queryByText("Sample usage breakdown")).not.toBeInTheDocument(); - expect(window.location.search).toBe(""); + await waitFor(() => expect(window.location.search).toBe("")); expect(screen.getByRole("button", { name: "Syncing…" })).toBeDisabled(); expect(apiClient.post).not.toHaveBeenCalled(); expect(apiClient.put).not.toHaveBeenCalled(); @@ -569,7 +572,7 @@ describe("ROICalculatorView", () => { fireEvent.click(screen.getByRole("button", { name: "Exit demo" })); expect(screen.getByRole("progressbar")).toBeVisible(); expect(screen.getByText("$20.00")).toBeVisible(); - expect(window.location.search).toBe(""); + await waitFor(() => expect(window.location.search).toBe("")); }); it.each(["report", "sync"])("loads a demo link when the live %s request fails", async (failedRequest) => { @@ -615,7 +618,7 @@ describe("ROICalculatorView", () => { expect(screen.getByRole("button", { name: "Settings" })).toBeEnabled(); expect(screen.queryByText("You’re viewing demo data")).not.toBeInTheDocument(); expect(apiClient.post).not.toHaveBeenCalled(); - expect(window.location.search).toBe("?from=review"); + await waitFor(() => expect(window.location.search).toBe("?from=review")); expect(window.location.hash).toBe("#overview"); }); @@ -634,7 +637,7 @@ describe("ROICalculatorView", () => { expect(screen.getByText("Loading ROI Calculator…")).toBeVisible(); expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); expect(screen.queryByRole("button", { name: "Run analysis" })).not.toBeInTheDocument(); - expect(window.location.search).toBe(""); + await waitFor(() => expect(window.location.search).toBe("")); pending.resolve(pendingRequest === "report" ? { report: summary } : idleStatus); expect(await screen.findByText("Gateway AI cost")).toBeVisible(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx index 912bf35d4db..36cb3d8d736 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx @@ -15,6 +15,7 @@ import { Tabs } from "@/components/ui/tabs"; import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } from "@/components/ui/dialog"; import { extractErrorMessage } from "@/utils/errorUtils"; import { isProxyAdminTierRole } from "@/utils/roles"; +import { useQueryState } from "nuqs"; import ROISettingsPanel from "./ROISettingsPanel"; import { MatchedPeopleToggle } from "./MatchedPeopleToggle"; import { IdentityMatchDialog, type PersonMatchSelection, PullReasoningDialog } from "./ROICalculatorDialogs"; @@ -29,6 +30,7 @@ import type { ROISummary, ROISyncStatus, } from "./roiCalculatorData"; +import { parseAsDemoFlag } from "./demoUrlState"; type View = "overview" | "people" | "branches"; @@ -45,13 +47,6 @@ const IDLE_STATUS: ROISyncStatus = { error: null, }; -function updateDemoUrl(enabled: boolean) { - const url = new URL(window.location.href); - if (enabled) url.searchParams.set("demo", "1"); - else url.searchParams.delete("demo"); - window.history.replaceState(null, "", url); -} - export default function ROICalculatorView({ accessToken, userRole = null, @@ -61,6 +56,8 @@ export default function ROICalculatorView({ userRole?: string | null; isViewOnly?: boolean; }) { + const [demo, setDemo] = useQueryState("demo", parseAsDemoFlag); + const [demoRequestedOnLoad] = React.useState(demo === true); const [sampleSummary, setSampleSummary] = React.useState(null); const adminReadOnly = isViewOnly && isProxyAdminTierRole(userRole ?? ""); const readOnly = adminReadOnly || sampleSummary !== null; @@ -95,7 +92,6 @@ export default function ROICalculatorView({ React.useEffect(() => { if (!accessToken) return; let cancelled = false; - const demoRequested = new URLSearchParams(window.location.search).get("demo") === "1"; const settingsRequest = apiClient.get("/roi-calculator/settings", { accessToken }); const reportRequest = apiClient .get("/roi-calculator/report", { accessToken }) @@ -128,13 +124,13 @@ export default function ROICalculatorView({ }); Promise.all([ settingsRequest, - demoRequested + demoRequestedOnLoad ? apiClient .get("/roi-calculator/report", { accessToken, query: { mode: "demo" } }) .catch((reason: unknown) => { if (!cancelled) { setDemoError(`Could not load demo data: ${extractErrorMessage(reason)}`); - updateDemoUrl(false); + setDemo(null); } return liveData; }) @@ -156,7 +152,7 @@ export default function ROICalculatorView({ return () => { cancelled = true; }; - }, [accessToken]); + }, [accessToken, demoRequestedOnLoad, setDemo]); React.useEffect(() => { if (!accessToken || !settingsLoaded) return; @@ -280,7 +276,7 @@ export default function ROICalculatorView({ }); setSampleSummary(response.report); setDemoError(null); - updateDemoUrl(true); + setDemo(true); setView("branches"); setQuery(""); } catch (reason) { @@ -355,7 +351,7 @@ export default function ROICalculatorView({ {sampleSummary && ( { - updateDemoUrl(false); + setDemo(null); setSampleSummary(null); }} /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/demoUrlState.ts b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/demoUrlState.ts new file mode 100644 index 00000000000..ba5388c4c76 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/demoUrlState.ts @@ -0,0 +1,6 @@ +import { createParser } from "nuqs"; + +export const parseAsDemoFlag = createParser({ + parse: (value) => (value === "1" ? true : null), + serialize: () => "1", +}); diff --git a/ui/litellm-dashboard/src/app/chat/layout.test.tsx b/ui/litellm-dashboard/src/app/chat/layout.test.tsx index 6c78bca0b5a..ac4fe2068fb 100644 --- a/ui/litellm-dashboard/src/app/chat/layout.test.tsx +++ b/ui/litellm-dashboard/src/app/chat/layout.test.tsx @@ -1,5 +1,6 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { render, screen } from "@testing-library/react"; +import { screen } from "@testing-library/react"; +import { renderWithProviders as render } from "@/../tests/test-utils"; import ChatLayout from "./layout"; const { mockUseAuthorized, mockUseUISettings, mockReplace, mockUiHref, state } = vi.hoisted(() => { diff --git a/ui/litellm-dashboard/src/app/chat/layout.tsx b/ui/litellm-dashboard/src/app/chat/layout.tsx index 2e0db2c6bdc..0bed5bda4a8 100644 --- a/ui/litellm-dashboard/src/app/chat/layout.tsx +++ b/ui/litellm-dashboard/src/app/chat/layout.tsx @@ -10,7 +10,7 @@ import { ChatShellProvider } from "@/contexts/ChatShellContext"; import ChatShell from "@/components/chat/ChatShell"; import { uiHref } from "@/utils/uiHref"; -// ChatShellProvider uses useSearchParams(), which requires a Suspense boundary for static export. +// The nuqs Next adapter uses useSearchParams, so keep the chat tree behind Suspense. function ChatLayoutContent({ children }: { children: React.ReactNode }) { const { accessToken, userRole, userId, userEmail, premiumUser } = useAuthorized(); const { data: uiSettings, isLoading: isUISettingsLoading } = useUISettings(); diff --git a/ui/litellm-dashboard/src/app/chat/page.integration.test.tsx b/ui/litellm-dashboard/src/app/chat/page.integration.test.tsx index ec884a90d7b..2020a66bd1c 100644 --- a/ui/litellm-dashboard/src/app/chat/page.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/chat/page.integration.test.tsx @@ -1,8 +1,9 @@ -import React from "react"; -import { fireEvent, render, screen, waitFor, within } from "@testing-library/react"; +import { fireEvent, screen, waitFor, within } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { useChatHistory } from "@/components/chat/useChatHistory"; import ChatConversationPage from "./page"; +import { renderWithProviders } from "@/../tests/test-utils"; +import type { OnUrlUpdateFunction } from "nuqs/adapters/testing"; const { mockMakeOpenAIResponsesRequest, shellState } = vi.hoisted(() => ({ mockMakeOpenAIResponsesRequest: vi.fn(), @@ -68,8 +69,8 @@ const ON_TIMING_DATA_INDEX = 7; const ON_USAGE_DATA_INDEX = 8; const ON_TOTAL_LATENCY_INDEX = 24; -async function sendOneMessage(): Promise { - render(); +async function sendOneMessage(onUrlUpdate?: OnUrlUpdateFunction): Promise { + renderWithProviders(, { onUrlUpdate }); expect(await screen.findByRole("button", { name: /gpt-5\.4-mini/ })).toBeInTheDocument(); fireEvent.change(screen.getByPlaceholderText("How can I help you today?"), { target: { value: "How much did this cost?" }, @@ -120,6 +121,19 @@ describe("/ui/chat request metrics", () => { expect(typeof call[ON_TOTAL_LATENCY_INDEX]).toBe("function"); }); + it("puts the new conversation ID in the URL after the first message is sent", async () => { + mockMakeOpenAIResponsesRequest.mockResolvedValue(undefined); + const onUrlUpdate = vi.fn(); + + await sendOneMessage(onUrlUpdate); + + await waitFor(() => expect(onUrlUpdate).toHaveBeenCalled()); + const newConversationId = onUrlUpdate.mock.lastCall?.[0].searchParams.get("id"); + expect(newConversationId).toBeTruthy(); + expect(onUrlUpdate.mock.lastCall?.[0].options.history).toBe("push"); + expect(localStorage.getItem("litellm_chat_history_v1:metrics-test-user")).toContain(newConversationId); + }); + it("shows no metrics bar for a turn the provider reported no usage for", async () => { mockMakeOpenAIResponsesRequest.mockImplementation(async (...args: unknown[]) => { const updateTextUI = args[1] as (role: string, delta: string) => void; @@ -141,7 +155,7 @@ describe("/ui/chat storage banner", () => { }); it("keeps the dismiss control amber on hover instead of the ghost variant's foreground", async () => { - render(); + renderWithProviders(); const banner = await screen.findByText("Chat history won't be saved in this browser session"); const dismiss = within(banner.parentElement!).getByRole("button"); diff --git a/ui/litellm-dashboard/src/app/chat/page.tsx b/ui/litellm-dashboard/src/app/chat/page.tsx index 3625b3e489a..fdf6a1e02a3 100644 --- a/ui/litellm-dashboard/src/app/chat/page.tsx +++ b/ui/litellm-dashboard/src/app/chat/page.tsx @@ -18,6 +18,7 @@ import { makeOpenAIResponsesRequest } from "@/components/llm_calls/responses_api import type { TokenUsage } from "@/components/chat_ui/ResponseMetrics"; import type { MCPEvent } from "@/components/chat/types"; import { getProviderLogoAndName } from "@/components/provider_info_helpers"; +import { useQueryState } from "nuqs"; const SUGGESTIONS = ["Write", "Learn", "Code", "Brainstorm"]; const LOCALSTORAGE_MODEL_KEY = "litellm_chat_selected_model"; @@ -50,6 +51,7 @@ function getProviderFromModelName(modelName: string): string { } export default function ChatConversationPage() { + const [, setConversationIdInUrl] = useQueryState("id", { history: "push" }); const router = useRouter(); const { accessToken, @@ -140,7 +142,7 @@ export default function ChatConversationPage() { if (!convId) { convId = createConversation(model); setResponsesSessionId(null); // new conversation starts a fresh session - window.history.pushState(null, "", `${window.location.pathname}?id=${convId}`); + setConversationIdInUrl(convId); } appendMessage(convId, { role: "user", content: trimmed }); @@ -258,6 +260,7 @@ export default function ChatConversationPage() { updateLastAssistantMessage, isStreaming, responsesSessionId, + setConversationIdInUrl, ], ); diff --git a/ui/litellm-dashboard/src/app/chat/page.url-state.integration.test.tsx b/ui/litellm-dashboard/src/app/chat/page.url-state.integration.test.tsx new file mode 100644 index 00000000000..b8ab0013a05 --- /dev/null +++ b/ui/litellm-dashboard/src/app/chat/page.url-state.integration.test.tsx @@ -0,0 +1,83 @@ +import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { NuqsAdapter } from "nuqs/adapters/react"; +import { ChatShellProvider } from "@/contexts/ChatShellContext"; +import ChatConversationPage from "./page"; + +const { mockMakeOpenAIResponsesRequest } = vi.hoisted(() => ({ + mockMakeOpenAIResponsesRequest: vi.fn(), +})); + +vi.mock("next/navigation", () => ({ + useRouter: () => ({ push: vi.fn(), replace: vi.fn() }), +})); + +vi.mock("@/components/llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn(async () => [{ model_group: "gpt-5.4-mini" }]), +})); + +vi.mock("@/components/llm_calls/responses_api", () => ({ + makeOpenAIResponsesRequest: mockMakeOpenAIResponsesRequest, +})); + +vi.mock("@/components/chat/MCPConnectPicker", () => ({ + default: () =>
, +})); + +vi.mock("react-markdown", () => ({ + default: ({ children }: { children: string }) =>
{children}
, +})); + +vi.mock("remark-gfm", () => ({ default: () => undefined })); + +vi.mock("react-syntax-highlighter", () => ({ + Prism: ({ children }: { children: string }) =>
{children}
, +})); + +vi.mock("react-syntax-highlighter/dist/esm/styles/prism", () => ({ coy: {}, oneDark: {}, oneLight: {}, prism: {} })); + +describe("chat page URL state with ChatShellProvider", () => { + beforeEach(() => { + localStorage.clear(); + window.history.replaceState(null, "", "/chat"); + mockMakeOpenAIResponsesRequest.mockReset(); + }); + + it("keeps the active conversation in sync with the URL across the first send and back navigation", async () => { + mockMakeOpenAIResponsesRequest.mockResolvedValue(undefined); + + render( + + + + + , + ); + + expect(await screen.findByRole("button", { name: /gpt-5\.4-mini/ })).toBeInTheDocument(); + fireEvent.change(screen.getByPlaceholderText("How can I help you today?"), { + target: { value: "How much did this cost?" }, + }); + fireEvent.click(screen.getByRole("button", { name: "Send" })); + + await waitFor(() => expect(window.location.search).toMatch(/^\?id=/)); + const conversationId = new URLSearchParams(window.location.search).get("id"); + expect(conversationId).toBeTruthy(); + expect(await screen.findByText("How much did this cost?")).toBeInTheDocument(); + await waitFor(() => + expect(localStorage.getItem("litellm_chat_history_v1:url-test-user")).toContain(conversationId), + ); + + act(() => window.history.back()); + await waitFor(() => expect(window.location.search).toBe("")); + + expect(await screen.findByPlaceholderText("How can I help you today?")).toBeInTheDocument(); + expect(screen.queryByText("How much did this cost?")).not.toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/globals.css b/ui/litellm-dashboard/src/app/globals.css index cad79f64ea0..cfa85817ea6 100644 --- a/ui/litellm-dashboard/src/app/globals.css +++ b/ui/litellm-dashboard/src/app/globals.css @@ -208,9 +208,12 @@ --trace-tab-active: oklch(0.95 0.025 205); --trace-tab-hover: oklch(0.93 0.02 215); --trace-tag: oklch(0.95 0.02 205); - --trace-chain: oklch(0.56 0.17 255); - --trace-llm: oklch(0.6 0.13 215); - --trace-tool: oklch(0.64 0.14 165); + --trace-chain: #2a78d6; + --trace-llm: #eb6834; + --trace-tool: #1baf7a; + --trace-chain-soft: color-mix(in oklab, #2a78d6 12%, transparent); + --trace-llm-soft: color-mix(in oklab, #eb6834 12%, transparent); + --trace-tool-soft: color-mix(in oklab, #1baf7a 14%, transparent); --trace-glyph: oklch(0.99 0 0); --trace-human: oklch(0.5 0.15 260); --trace-human-glyph: oklch(0.95 0.04 210); @@ -282,9 +285,12 @@ --trace-tab-active: oklch(0.28 0.03 215); --trace-tab-hover: oklch(0.32 0.03 220); --trace-tag: oklch(0.28 0.03 215); - --trace-chain: oklch(0.6 0.17 255); - --trace-llm: oklch(0.64 0.13 215); - --trace-tool: oklch(0.68 0.14 165); + --trace-chain: #3987e5; + --trace-llm: #d95926; + --trace-tool: #199e70; + --trace-chain-soft: color-mix(in oklab, #3987e5 24%, transparent); + --trace-llm-soft: color-mix(in oklab, #d95926 24%, transparent); + --trace-tool-soft: color-mix(in oklab, #199e70 24%, transparent); --trace-glyph: oklch(0.99 0 0); --trace-human: oklch(0.56 0.15 260); --trace-human-glyph: oklch(0.95 0.04 210); @@ -320,6 +326,9 @@ --color-trace-chain: var(--trace-chain); --color-trace-llm: var(--trace-llm); --color-trace-tool: var(--trace-tool); + --color-trace-chain-soft: var(--trace-chain-soft); + --color-trace-llm-soft: var(--trace-llm-soft); + --color-trace-tool-soft: var(--trace-tool-soft); --color-trace-glyph: var(--trace-glyph); --color-trace-human: var(--trace-human); --color-trace-human-glyph: var(--trace-human-glyph); diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx index 083b1e5f3e2..c1fb0431e8e 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.test.tsx @@ -3,7 +3,7 @@ import React from "react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, screen, testQueryClient, waitFor, within } from "../../../tests/test-utils"; -import type { LogEntry as SpendLogEntry } from "@/components/view_logs/columns"; +import type { LogEntry as SpendLogEntry } from "@/components/logs/types"; import { LogViewer } from "./LogViewer"; vi.mock("@/components/networking", async (importOriginal) => { @@ -11,7 +11,7 @@ vi.mock("@/components/networking", async (importOriginal) => { return { ...actual, uiSpendLogsCall: vi.fn() }; }); -vi.mock("@/components/view_logs/LogDetailsDrawer", () => ({ +vi.mock("@/components/logs/detail", () => ({ LogDetailsDrawer: function LogDetailsDrawerMock({ open, logEntry, diff --git a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx index 2abd699ba86..1e2538100ae 100644 --- a/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx +++ b/ui/litellm-dashboard/src/components/GuardrailsMonitor/LogViewer.tsx @@ -5,8 +5,8 @@ import React, { useState } from "react"; import { Button } from "@/components/ui/button"; import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { uiSpendLogsCall } from "@/components/networking"; -import { LogDetailsDrawer } from "@/components/view_logs/LogDetailsDrawer"; -import type { LogEntry as ViewLogsLogEntry } from "@/components/view_logs/columns"; +import { LogDetailsDrawer } from "@/components/logs/detail"; +import type { LogEntry as ViewLogsLogEntry } from "@/components/logs/types"; import type { LogEntry } from "./mockData"; const actionConfig: Record< diff --git a/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx b/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx index 2b6c06e9a96..44ef0e61340 100644 --- a/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AutoRouterRoutingTest.tsx @@ -3,7 +3,7 @@ import { TriangleAlert } from "lucide-react"; import { Badge } from "@/components/ui/badge"; import { Button } from "@/components/ui/button"; import { Textarea } from "@/components/ui/textarea"; -import RoutingDecisionCard from "@/components/view_logs/LogDetailsDrawer/RoutingDecisionCard"; +import RoutingDecisionCard from "@/components/logs/detail/RoutingDecisionCard"; import { AutoRouterRoutingTestResult, testAutoRouterRouting } from "../networking"; import { ComplexityRouterConfigPayload, getHeuristicV2SuccessThresholdError } from "./build_complexity_router_config"; import { buildAutoRouterRoutingTestRequest } from "./build_auto_router_routing_test_request"; diff --git a/ui/litellm-dashboard/src/components/chat/ChatShell.test.tsx b/ui/litellm-dashboard/src/components/chat/ChatShell.test.tsx index bff6d30a5c8..ac92bb15071 100644 --- a/ui/litellm-dashboard/src/components/chat/ChatShell.test.tsx +++ b/ui/litellm-dashboard/src/components/chat/ChatShell.test.tsx @@ -1,5 +1,6 @@ import { afterEach, describe, expect, it, vi } from "vitest"; -import { fireEvent, render, screen } from "@testing-library/react"; +import { fireEvent, screen } from "@testing-library/react"; +import { renderWithProviders as render } from "@/../tests/test-utils"; import ChatShell from "./ChatShell"; const { mockPush, mockUsePathname, mockUseChatShell } = vi.hoisted(() => ({ diff --git a/ui/litellm-dashboard/src/components/common_components/simple_table.tsx b/ui/litellm-dashboard/src/components/common_components/simple_table.tsx index a4b84d28801..8d174e49cf2 100644 --- a/ui/litellm-dashboard/src/components/common_components/simple_table.tsx +++ b/ui/litellm-dashboard/src/components/common_components/simple_table.tsx @@ -28,7 +28,7 @@ interface SimpleTableProps { /** * Simple table component for forms and settings pages - * For complex tables with sorting/filtering, use DataTable from view_logs + * For complex tables with sorting/filtering, use DataTable from shared/DataTable */ export function SimpleTable({ data, diff --git a/ui/litellm-dashboard/src/components/lens/LensModeSwitch.tsx b/ui/litellm-dashboard/src/components/lens/LensModeSwitch.tsx index 11d02567799..0cfa9e47bce 100644 --- a/ui/litellm-dashboard/src/components/lens/LensModeSwitch.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensModeSwitch.tsx @@ -87,7 +87,7 @@ export function LensModeSwitch({ {label} {view === "investigations" && setup && ( - + {setup} )} diff --git a/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx index 320dbef4142..b842e6a0bd3 100644 --- a/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensPage.integration.test.tsx @@ -2,12 +2,12 @@ import { screen } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, testQueryClient } from "@/../tests/test-utils"; -import { dismissLensIntro } from "@/../tests/lens-test-utils"; +import { dismissLensIntro, requestPath } from "@/../tests/lens-test-utils"; import LensPage from "@/app/(dashboard)/lens/page"; const { auth } = vi.hoisted(() => ({ auth: vi.fn() })); vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: auth })); -vi.mock("@/components/view_logs/TraceView/AgentTracesPage", () => ({ +vi.mock("@/components/lens/traces/list/AgentTracesPage", () => ({ default: ({ isActive }: { isActive: boolean }) =>
Trace polling {isActive ? "active" : "paused"}
, })); vi.mock("./investigations/InvestigationsView", () => ({ @@ -26,7 +26,7 @@ describe("Lens navigation", () => { vi.stubGlobal( "fetch", vi.fn(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/v1/traces") return Response.json({ data: [{}] }); if (path === "/lens") return Response.json({ lenses: [], workers: [], tracing_enabled: true }); return Response.json({ traces: true, requests: false, data: [] }); diff --git a/ui/litellm-dashboard/src/components/lens/LensSetup.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/LensSetup.integration.test.tsx index cfc9a7183d4..c74b266714d 100644 --- a/ui/litellm-dashboard/src/components/lens/LensSetup.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensSetup.integration.test.tsx @@ -2,7 +2,7 @@ import { act, fireEvent, screen, within, waitFor } from "@testing-library/react" import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, testQueryClient } from "@/../tests/test-utils"; -import { dismissLensIntro } from "@/../tests/lens-test-utils"; +import { dismissLensIntro, readRequest, requestPath } from "@/../tests/lens-test-utils"; import { readStorage } from "@/lib/storage"; import { LENS_INTRO_DISMISSED, LENS_INTRO_SEEN } from "./storage"; import { LensWorkspace } from "./LensWorkspace"; @@ -23,14 +23,14 @@ const worker = () => ({ function serve({ enabled = false, traces = false, requests = false, connected = false } = {}) { list.mockResolvedValue({ lenses: [], workers: connected ? [worker()] : [], tracing_enabled: enabled }); network.mockImplementation(async (input, init) => { - const path = new URL(String(input), "http://localhost").pathname; + const { path, method, body } = await readRequest(input, init); if (path === "/v1/traces") return enabled ? Response.json({ data: traces ? [data.runs[0].trace.summary] : [] }) : Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); if (path === "/lens/activity/available") return Response.json({ traces, requests }); - if (path === "/lens" && init?.method === "POST") { - const saved = { ...data.lenses[0], settings: { ...data.lenses[0].settings, ...JSON.parse(String(init.body)) } }; + if (path === "/lens" && method === "POST") { + const saved = { ...data.lenses[0], settings: { ...data.lenses[0].settings, ...(body as object) } }; list.mockResolvedValue({ lenses: [saved], workers: [worker()], tracing_enabled: true }); return Response.json(saved); } @@ -154,9 +154,7 @@ describe("Lens setup journey", () => { serve({ enabled: true, traces: true }); const normal = network.getMockImplementation()!; network.mockImplementation((input, init) => - new URL(String(input), "http://localhost").pathname === pendingPath - ? new Promise(() => {}) - : normal(input, init), + requestPath(input) === pendingPath ? new Promise(() => {}) : normal(input, init), ); renderWorkspace(); expect(await screen.findByRole("table", { name: "Agent runs" })).toBeVisible(); @@ -171,9 +169,7 @@ describe("Lens setup journey", () => { list.mockResolvedValue({ lenses: data.lenses, workers: [worker()], tracing_enabled: false }); const normal = network.getMockImplementation()!; network.mockImplementation((input, init) => - new URL(String(input), "http://localhost").pathname === pendingPath - ? new Promise(() => {}) - : normal(input, init), + requestPath(input) === pendingPath ? new Promise(() => {}) : normal(input, init), ); renderWorkspace({ searchParams: `?lens=${data.lenses[0].id}` }); expect(await screen.findByRole("heading", { name: data.lenses[0].settings.name })).toBeVisible(); @@ -247,7 +243,7 @@ describe("Lens setup journey", () => { await intro.findByRole("button", { name: "Check for traces" }); const normal = network.getMockImplementation()!; network.mockImplementation(async (input, init) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/v1/traces") return Response.json({ detail: "Trace storage unavailable" }, { status: 503 }); return normal(input, init); }); @@ -268,7 +264,7 @@ describe("Lens setup journey", () => { expect(await screen.findByRole("button", { name: "New investigation" })).toBeEnabled(); const normal = network.getMockImplementation()!; network.mockImplementation((input, init) => - new URL(String(input), "http://localhost").pathname === "/lens/activity/available" + requestPath(input) === "/lens/activity/available" ? Promise.resolve(Response.json({ detail: "Activity unavailable" }, { status: 503 })) : normal(input, init), ); @@ -288,9 +284,7 @@ describe("Lens setup journey", () => { const intro = within(await screen.findByRole("dialog")); expect(await intro.findByText(/A gateway administrator can connect a worker/)).toBeVisible(); expect(intro.getByRole("button", { name: "Connect worker" })).toBeDisabled(); - expect(network.mock.calls.some(([input]) => new URL(String(input), "http://localhost").pathname === "/lens")).toBe( - false, - ); + expect(network.mock.calls.some(([input]) => requestPath(input) === "/lens")).toBe(false); }); it.each(["traces", "requests with trace errors", "requests with pending traces", "traces with activity errors"])( @@ -302,7 +296,7 @@ describe("Lens setup journey", () => { const normal = network.getMockImplementation()!; const failingPath = scenario === "requests with trace errors" ? "/v1/traces" : "/lens/activity/available"; network.mockImplementation((input, init) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/v1/traces" && scenario === "requests with pending traces") return new Promise(() => {}); if (scenario.endsWith("errors") && path === failingPath) @@ -328,13 +322,10 @@ describe("Lens setup journey", () => { within(screen.getByRole("tablist", { name: "Lens" })).getByRole("tab", { name: "Investigations" }), ).toHaveAttribute("aria-selected", "true"); await waitFor(() => expect(setupParam(onUrlUpdate)).toBeNull()); - const create = network.mock.calls.find( - ([input, init]) => new URL(String(input), "http://localhost").pathname === "/lens" && init?.method === "POST", - ); + const requests = await Promise.all(network.mock.calls.map(([input, init]) => readRequest(input, init))); + const create = requests.find((request) => request.path === "/lens" && request.method === "POST"); expect(create).toBeDefined(); - expect(JSON.parse(String(create?.[1]?.body))).toEqual( - expect.objectContaining({ name: "My first review", source }), - ); + expect(create?.body).toEqual(expect.objectContaining({ name: "My first review", source })); }, ); }); diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx index 9e7ece7cf7c..ce99bbae283 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx @@ -2,7 +2,7 @@ import { screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { renderWithProviders, testQueryClient } from "@/../tests/test-utils"; -import { dismissLensIntro } from "@/../tests/lens-test-utils"; +import { dismissLensIntro, readRequest, requestPath } from "@/../tests/lens-test-utils"; import { LensWorkspace } from "./LensWorkspace"; import { lensKeys } from "./data/queries"; import { createLensDemoData } from "./data/demo/fixtures"; @@ -21,7 +21,7 @@ beforeEach(() => { vi.stubGlobal("fetch", network); network.mockReset(); network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/v1/traces") return Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); if (path === "/lens") return Response.json({ lenses: [], workers: [], tracing_enabled: false }); return Response.json({ data: [], traces: false, requests: false }); @@ -173,7 +173,7 @@ describe("Lens interactive demo", () => { const data = createLensDemoData(); const saved = data.lenses[0]; network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [saved], workers: [], tracing_enabled: true }); if (path.endsWith("/runs")) return Response.json(saved.jobs); if (path === "/v1/traces") return Response.json({ data: data.runs.map((run) => run.trace.summary) }); @@ -194,7 +194,7 @@ describe("Lens interactive demo", () => { const user = userEvent.setup(); const saved = createLensDemoData().lenses[0]; network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [saved], workers: [], tracing_enabled: false }); if (path.endsWith("/runs")) return Response.json(saved.jobs); if (path === "/v1/traces") return Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); @@ -220,7 +220,7 @@ describe("Lens interactive demo", () => { }); const lenses = vi.fn(() => [withJob("running")]); network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: lenses(), workers: [], tracing_enabled: false }); if (path === "/v1/traces") return Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); return Response.json({ data: [], traces: false, requests: false }); @@ -237,7 +237,7 @@ describe("Lens interactive demo", () => { const user = userEvent.setup(); const saved = createLensDemoData().lenses[0]; network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [saved], workers: [], tracing_enabled: true }); if (path.endsWith("/runs")) return Response.json(saved.jobs); if (path === "/lens/agents") return Response.json([]); @@ -272,7 +272,7 @@ describe("Lens interactive demo", () => { }; const workers = vi.fn(() => [worker]); network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [saved], workers: workers(), tracing_enabled: true }); if (path.endsWith("/runs")) return Response.json(saved.jobs); if (path === "/v1/traces") return Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); @@ -308,7 +308,7 @@ describe("Lens interactive demo", () => { const user = userEvent.setup(); const onUrlUpdate = vi.fn(); network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [], workers: [], tracing_enabled: true }); if (path === "/v1/traces") return Response.json({ data: [{}] }); return Response.json({ data: [], traces: true, requests: false }); @@ -339,9 +339,9 @@ describe("Lens interactive demo", () => { }; const workers = vi.fn((): (typeof worker)[] => []); network.mockImplementation(async (input, init) => { - const path = new URL(String(input), "http://localhost").pathname; + const { path, method } = await readRequest(input, init); if (path === "/lens") return Response.json({ lenses: [], workers: workers(), tracing_enabled: true }); - if (path === "/lens/workers/register" && init?.method === "POST") { + if (path === "/lens/workers/register" && method === "POST") { workers.mockReturnValue([worker]); return Response.json({ token: "lens-test-token", image: "lens-worker:v1", worker }); } @@ -405,7 +405,7 @@ describe("Lens interactive demo", () => { last_seen: new Date().toISOString(), }; network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [saved], workers: [worker], tracing_enabled: true }); if (path === "/v1/traces") return Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); return Response.json({ data: [], traces: true, requests: false }); @@ -434,10 +434,9 @@ describe("Lens interactive demo", () => { last_seen: new Date(Date.now() - 600_000).toISOString(), }; const workers = vi.fn(() => [worker]); - const listCalls = () => - network.mock.calls.filter(([input]) => new URL(String(input), "http://localhost").pathname === "/lens").length; + const listCalls = () => network.mock.calls.filter(([input]) => requestPath(input) === "/lens").length; network.mockImplementation(async (input) => { - const path = new URL(String(input), "http://localhost").pathname; + const path = requestPath(input); if (path === "/lens") return Response.json({ lenses: [saved], workers: workers(), tracing_enabled: true }); if (path === "/v1/traces") return Response.json({ detail: "Tracing is not enabled" }, { status: 501 }); return Response.json({ data: [], traces: true, requests: false }); diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx index caa8576bdd7..ccd2f665564 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx @@ -3,13 +3,13 @@ import { useId, useState } from "react"; import { useQuery } from "@tanstack/react-query"; import { Aperture, ArrowUpRight } from "lucide-react"; -import AgentTracesPage from "@/components/view_logs/TraceView/AgentTracesPage"; +import AgentTracesPage from "@/components/lens/traces/list/AgentTracesPage"; import { Button } from "@/components/ui/button"; -import type { TraceSummary } from "@/components/view_logs/TraceView/traceTypes"; +import type { TraceSummary } from "@/components/lens/traces/types"; import { Switch } from "@/components/ui/switch"; import { Tabs, TabsContent } from "@/components/ui/tabs"; import { LensServicesProvider, useLensAccessToken, useLensApi, useLiveLensServices } from "./data/LensServices"; -import { LensPreviewContext } from "@/components/view_logs/TraceView/LensPreviewButton"; +import { LensPreviewContext } from "@/components/lens/ui/LensPreviewButton"; import { isProxyAdminRole, isProxyAdminTierRole } from "@/utils/roles"; import { InvestigationsView } from "./investigations/InvestigationsView"; import { LensSettings } from "./settings/LensSettings"; @@ -22,7 +22,7 @@ import { cn } from "@/lib/cva.config"; import { useDialogRoute, useLensRoute, type LensDialog, type LensTab } from "./route"; import { LensIntroDialog, useLensIntro } from "./onboarding/LensIntroDialog"; import { OnboardingProvider, type Onboarding } from "./onboarding/OnboardingContext"; -import { traceRefOf, useOpenTraceRouting } from "@/components/view_logs/TraceView/traceRouting"; +import { traceRefOf, useOpenTraceRouting } from "@/components/lens/traces/routing"; type WorkspaceProps = { accessToken: string; userRole: string; readOnly: boolean }; diff --git a/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx b/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx index 8abeab16a78..a4804212012 100644 --- a/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx +++ b/ui/litellm-dashboard/src/components/lens/data/LensServices.tsx @@ -2,7 +2,8 @@ import { createContext, useContext, useMemo, type ReactNode } from "react"; import { apiClient } from "@/components/networking"; -import { liveTracesApi, TracesApiContext, type TracesApi } from "@/components/view_logs/TraceView/tracesApi"; +import { fetchClient } from "@/lib/http/api"; +import { liveTracesApi, TracesApiContext, type TracesApi } from "@/components/lens/traces/api"; import { liveLensApi, type LensApi } from "./service"; export interface LensServices { @@ -14,7 +15,7 @@ export interface LensServices { const LensServicesContext = createContext(null); export function liveLensServices(accessToken: string): LensServices { - return { accessToken, lens: liveLensApi(apiClient, accessToken), traces: liveTracesApi(accessToken) }; + return { accessToken, lens: liveLensApi(fetchClient, apiClient, accessToken), traces: liveTracesApi(accessToken) }; } function useLensServices(): LensServices { diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts index 97ab6c84e91..c0b3aa50f41 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts @@ -1,5 +1,5 @@ import { ApiError } from "@/lib/http/client"; -import type { TracesApi } from "@/components/view_logs/TraceView/tracesApi"; +import type { TracesApi } from "@/components/lens/traces/api"; import type { LensServices } from "../LensServices"; import type { LensApi } from "../service"; import { createLensDemoData, type LensDemoData } from "./fixtures"; diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts b/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts index 51dc89e6fa6..becdff35d40 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/fixtures.ts @@ -1,4 +1,4 @@ -import type { Trace, Span, SpanDetail } from "@/components/view_logs/TraceView/traceTypes"; +import type { Trace, Span, SpanDetail } from "@/components/lens/traces/types"; import type { Lens, Finding, Job, Settings } from "../../model/types"; import { withReleaseCases } from "./lensDemoLongTrace"; import { scenarios, type Scenario } from "./scenarios"; diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts b/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts index e30cff6ee1d..866f1042b73 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/lensDemoLongTrace.ts @@ -1,4 +1,4 @@ -import type { Span, SpanDetail, Trace } from "@/components/view_logs/TraceView/traceTypes"; +import type { Span, SpanDetail, Trace } from "@/components/lens/traces/types"; export function withReleaseCases(run: { trace: Trace; details: SpanDetail[] }) { const { trace } = run; diff --git a/ui/litellm-dashboard/src/components/lens/data/service.ts b/ui/litellm-dashboard/src/components/lens/data/service.ts index 6f23aa0dda7..326ee398c0e 100644 --- a/ui/litellm-dashboard/src/components/lens/data/service.ts +++ b/ui/litellm-dashboard/src/components/lens/data/service.ts @@ -1,6 +1,8 @@ import { z } from "zod"; import type { ApiClient } from "@/lib/http/client"; -import type { components } from "@/lib/http/schema"; +import { getAuthHeaderName } from "@/lib/http/runtime"; +import type { Client } from "openapi-fetch"; +import type { components, paths } from "@/lib/http/schema"; import type { ActivitySelection, AnalysisModelInfo, @@ -42,7 +44,7 @@ export interface LensApi { /** Partitions query caches between backends (one token, or the demo). */ readonly scope: string; lenses(): Promise; - activity(): Promise<{ traces: boolean; requests: boolean }>; + activity(): Promise; runs(lensId: string, offset: number): Promise; run(lensId: string, jobId: string): Promise; execution(lensId: string, executionId: string, offset: number): Promise; @@ -57,45 +59,73 @@ export interface LensApi { watchAll(): Promise; cancelRun(lensId: string): Promise; reviewFinding(lensId: string, findingId: string, status: FindingStatus, reason: string): Promise; - registerWorker(analysisKeyId: string | null): Promise; - setWorkerBillingKey(workerId: string, analysisKeyId: string | null): Promise; + registerWorker(analysisKeyId: string): Promise; + setWorkerBillingKey(workerId: string, analysisKeyId: string): Promise; revokeWorker(workerId: string): Promise; generateAnalysisKey(request: AnalysisKeyRequest): Promise<{ token_id?: string }>; deleteKeys(keys: readonly string[]): Promise; } -export function liveLensApi(apiClient: ApiClient, accessToken: string): LensApi { - const encode = encodeURIComponent; +type LensClient = Client; + +async function required(request: Promise<{ data?: T }>): Promise { + const { data } = await request; + if (data === undefined) throw new Error("The proxy returned an empty response"); + return data; +} + +async function sent(request: Promise): Promise { + await request; +} + +export function liveLensApi(client: LensClient, apiClient: ApiClient, accessToken: string): LensApi { + const headers = { [getAuthHeaderName()]: `Bearer ${accessToken}` }; + const lens = (lens_id: string) => ({ headers, params: { path: { lens_id } } }); + const worker = (worker_id: string) => ({ headers, params: { path: { worker_id } } }); return { scope: accessToken, - lenses: () => apiClient.get("/lens", { accessToken }), - activity: () => apiClient.get("/lens/activity/available", { accessToken }), - runs: (lensId, offset) => apiClient.get(`/lens/${lensId}/runs`, { accessToken, query: { offset } }), - run: (lensId, jobId) => apiClient.get(`/lens/${lensId}/runs/${jobId}`, { accessToken }), + lenses: () => required(client.GET("/lens", { headers })), + activity: () => required(client.GET("/lens/activity/available", { headers })), + runs: (lensId, offset) => + required( + client.GET("/lens/{lens_id}/runs", { headers, params: { path: { lens_id: lensId }, query: { offset } } }), + ), + run: (lensId, jobId) => + required( + client.GET("/lens/{lens_id}/runs/{job_id}", { + headers, + params: { path: { lens_id: lensId, job_id: jobId } }, + }), + ), execution: (lensId, executionId, offset) => - apiClient.get(`/lens/${lensId}/executions/${encode(executionId)}`, { - accessToken, - query: { offset }, - }), - sample: (selection, offset, asOf) => { - const { lookback_hours, ...selectionSettings } = selection; - return apiClient.post("/lens/preview/sample", { - accessToken, - body: { - offset, - as_of: asOf, - settings: { - ...selectionSettings, - execution_ids: [], - name: "Preview", - model: "preview", - checks: [{ id: "preview", instruction: "Preview recorded activity" }], + required( + client.GET("/lens/{lens_id}/executions/{execution_id}", { + headers, + params: { path: { lens_id: lensId, execution_id: executionId }, query: { offset } }, + }), + ), + sample: (selection, offset, asOf) => + required( + client.POST("/lens/preview/sample", { + headers, + body: { + offset, + as_of: asOf, + selection: { + source: selection.source, + service: selection.service ?? "", + agent_name: selection.agent_name ?? "", + filters: selection.filters ?? [], + sample_size: selection.sample_size, + sample_percent: selection.sample_percent ?? 100, + team_id: selection.team_id ?? "", + execution_ids: [], + }, + lookback_hours: selection.lookback_hours ?? 24, }, - lookback_hours: lookback_hours ?? 24, - }, - }); - }, - agents: () => apiClient.get("/lens/agents", { accessToken }), + }), + ), + agents: () => required(client.GET("/lens/agents", { headers })), models: () => apiClient.get("/models", { accessToken }), modelDetails: () => apiClient.get("/model_group/info", { accessToken }), keys: async (alias, page, signal) => @@ -118,21 +148,37 @@ export function liveLensApi(apiClient: ApiClient, accessToken: string): LensApi keyInfo: async (keyId) => keyInfoSchema.parse(await apiClient.get("/key/info", { accessToken, query: { key: keyId } })).info, saveLens: (id, settings) => - apiClient.request(id ? "PUT" : "POST", id ? `/lens/${id}` : "/lens", { accessToken, body: settings }), - startRun: (lensId, request = {}) => apiClient.post(`/lens/${lensId}/runs`, { accessToken, body: request }), - watchAll: () => - apiClient.post("/lens/watch-all", { accessToken, body: {} }), - cancelRun: (lensId) => apiClient.post(`/lens/${lensId}/cancel`, { accessToken, body: {} }), + required( + id + ? client.PUT("/lens/{lens_id}", { ...lens(id), body: settings }) + : client.POST("/lens", { headers, body: settings }), + ), + startRun: (lensId, request = {}) => sent(client.POST("/lens/{lens_id}/runs", { ...lens(lensId), body: request })), + watchAll: () => required(client.POST("/lens/watch-all", { headers })), + cancelRun: (lensId) => sent(client.POST("/lens/{lens_id}/cancel", lens(lensId))), reviewFinding: (lensId, findingId, status, reason) => - apiClient.patch(`/lens/${lensId}/findings/${findingId}`, { accessToken, body: { status, reason } }), + sent( + client.PATCH("/lens/{lens_id}/findings/{finding_id}", { + headers, + params: { path: { lens_id: lensId, finding_id: findingId } }, + body: { status, reason }, + }), + ), registerWorker: (analysisKeyId) => - apiClient.post("/lens/workers/register", { - accessToken, - body: { name: "Lens worker", analysis_key_id: analysisKeyId }, - }), + required( + client.POST("/lens/workers/register", { + headers, + body: { name: "Lens worker", analysis_key_id: analysisKeyId }, + }), + ), setWorkerBillingKey: (workerId, analysisKeyId) => - apiClient.put(`/lens/workers/${workerId}/billing-key`, { accessToken, body: { analysis_key_id: analysisKeyId } }), - revokeWorker: (workerId) => apiClient.delete(`/lens/workers/${workerId}`, { accessToken }), + sent( + client.PUT("/lens/workers/{worker_id}/billing-key", { + ...worker(workerId), + body: { analysis_key_id: analysisKeyId }, + }), + ), + revokeWorker: (workerId) => sent(client.DELETE("/lens/workers/{worker_id}", worker(workerId))), generateAnalysisKey: (request) => apiClient.post<{ token_id?: string }>("/key/generate", { accessToken, diff --git a/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts b/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts index a41c9b909c3..36dc613d5ab 100644 --- a/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts +++ b/ui/litellm-dashboard/src/components/lens/hooks/useLensReadiness.ts @@ -1,7 +1,7 @@ "use client"; import { useQuery } from "@tanstack/react-query"; -import { isTracingNotEnabled, useTraceAvailability } from "@/components/view_logs/TraceView/useAgentTraces"; +import { isTracingNotEnabled, useTraceAvailability } from "@/components/lens/traces/list/useAgentTraces"; import { useLensAccessToken, useLensApi } from "../data/LensServices"; import { lensQueries } from "../data/queries"; import { readiness, type Readiness, type ReadinessInput } from "../model/readiness"; diff --git a/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx b/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx index 1023e636537..fa4b716691c 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/Evidence.tsx @@ -5,8 +5,8 @@ import { useQuery } from "@tanstack/react-query"; import { Inspector } from "@/components/shared/Inspector"; import { Button } from "@/components/ui/button"; -import { RunView } from "@/components/view_logs/TraceView/TraceDrawer"; -import { useLocalRunSelection } from "@/components/view_logs/TraceView/traceRouting"; +import { RunView } from "@/components/lens/traces/detail/run/RunView"; +import { useLocalRunSelection } from "@/components/lens/traces/routing"; import { lensQueries } from "../data/queries"; import { useLensAccessToken, useLensApi } from "../data/LensServices"; diff --git a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx index a8abcf9bc25..544a7272e14 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx @@ -2,7 +2,7 @@ import { fireEvent, screen, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { afterEach, expect, it, vi } from "vitest"; -import type { RunSelection } from "@/components/view_logs/TraceView/traceRouting"; +import type { RunSelection } from "@/components/lens/traces/routing"; import { renderWithLens } from "@/../tests/lens-test-utils"; import { Inspector } from "@/components/shared/Inspector"; @@ -11,7 +11,7 @@ import type { OwnedFinding } from "../model/inbox"; import type { Finding, Lens } from "../model/types"; import { FindingPanel, ownedFindingKey } from "./FindingDetails"; -vi.mock("@/components/view_logs/TraceView/TraceDrawer", () => ({ +vi.mock("@/components/lens/traces/detail/run/RunView", () => ({ RunView: ({ traceId, selection }: { traceId: string; selection: RunSelection }) => (
{traceId} at {selection.spanId} diff --git a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx index 5d90281a080..039e884c793 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.integration.test.tsx @@ -2,12 +2,11 @@ import { act, fireEvent, screen, within, waitFor } from "@testing-library/react" import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import { testQueryClient } from "@/../tests/test-utils"; -import { renderWithLens } from "@/../tests/lens-test-utils"; +import { renderWithLens, stubGateway } from "@/../tests/lens-test-utils"; import { ApiError } from "@/lib/http/client"; -import { apiClient } from "@/components/networking"; import { lensKeys } from "../data/queries"; import { InvestigationsView } from "./InvestigationsView"; -import { LensPreviewContext } from "@/components/view_logs/TraceView/LensPreviewButton"; +import { LensPreviewContext } from "@/components/lens/ui/LensPreviewButton"; import { briefMarkdown } from "../model/findings"; import { findingKey } from "../model/inbox"; import { runTime } from "../model/format"; @@ -21,16 +20,20 @@ function renderWithProviders(ui: React.ReactElement, options?: Parameters ({ - apiClient: { get: vi.fn(), post: vi.fn(), patch: vi.fn(), request: vi.fn() }, +vi.mock("@/components/networking", async (importOriginal) => ({ + ...(await importOriginal()), proxyBaseUrl: "", getProxyBaseUrl: () => "", })); +let proxy = stubGateway(); +const sentBody = (handler: typeof proxy.post, path: string) => + handler.mock.calls.filter(([called]) => called === path).map(([, request]) => request.body); + beforeEach(() => { window.history.replaceState({}, "", "/lens/?lens=lens"); - vi.mocked(apiClient.post).mockReset(); - vi.mocked(apiClient.post).mockResolvedValue({ eligible: 0, selected: 0, executions: [] }); + proxy = stubGateway(); + proxy.post.mockResolvedValue({ eligible: 0, selected: 0, executions: [] }); }); const executionId = btoa(JSON.stringify(["traces", "", "trace-42"])); @@ -161,8 +164,8 @@ const lens: Lens = { describe("Lens findings and runs", () => { beforeEach(() => { - vi.mocked(apiClient.get).mockReset(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockReset(); + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return lens.jobs; return { data: [] }; @@ -197,7 +200,7 @@ describe("Lens findings and runs", () => { async function openIssue(finding: Finding) { testQueryClient.clear(); const jobs = lens.jobs.map((job) => ({ ...job, findings: [finding] })); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [{ ...lens, findings: [finding], jobs }], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return jobs; @@ -245,7 +248,7 @@ describe("Lens findings and runs", () => { it("closes the open run when the keyboard switches to another investigation run", async () => { testQueryClient.clear(); const older = { ...lens.jobs[0], id: "older", created_at: "2026-09-29T10:00:00Z" }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return [...lens.jobs, older]; return { data: [] }; @@ -269,7 +272,7 @@ describe("Lens findings and runs", () => { it("runs with saved settings from Run now without opening setup, then accepts an agent and window", async () => { testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], @@ -289,32 +292,29 @@ it("runs with saved settings from Run now without opening setup, then accepts an if (path === "/lens/activity/available") return { traces: true, requests: false }; return { data: [] }; }); - vi.mocked(apiClient.post).mockResolvedValue(lens); + proxy.post.mockResolvedValue(lens); const user = userEvent.setup(); renderWithProviders(); await user.click(await screen.findByRole("button", { name: "Run now" })); const choices = await screen.findByRole("dialog", { name: "Run now" }); expect(within(choices).getByRole("button", { name: "Since last run" })).toHaveAttribute("aria-pressed", "true"); await user.click(within(choices).getByRole("button", { name: "Run now" })); - expect(apiClient.post).toHaveBeenCalledWith("/lens/lens/runs", { accessToken: "test", body: {} }); + expect(sentBody(proxy.post, "/lens/lens/runs")).toEqual([{}]); await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); - vi.mocked(apiClient.post).mockClear(); + proxy.post.mockClear(); await user.click(screen.getByRole("button", { name: "Run now" })); const custom = await screen.findByRole("dialog", { name: "Run now" }); fireEvent.change(within(custom).getByRole("combobox", { name: "Agent" }), { target: { value: "billing" } }); await user.click(within(custom).getByRole("button", { name: "Last 24h" })); await user.click(within(custom).getByRole("button", { name: "Run now" })); - expect(apiClient.post).toHaveBeenCalledWith("/lens/lens/runs", { - accessToken: "test", - body: { agent_name: "billing", lookback_hours: 24 }, - }); + expect(sentBody(proxy.post, "/lens/lens/runs")).toEqual([{ agent_name: "billing", lookback_hours: 24 }]); }); it("offers the interactive demo without starting an investigation", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [], workers: [], tracing_enabled: false }; return { traces: false, requests: false }; }); @@ -323,13 +323,13 @@ it("offers the interactive demo without starting an investigation", async () => renderWithProviders(withPreview(, onPreview)); await user.click(await screen.findByRole("button", { name: "Preview sample" })); expect(onPreview).toHaveBeenCalledOnce(); - expect(apiClient.post).not.toHaveBeenCalled(); + expect(proxy.post).not.toHaveBeenCalled(); }); it("guides a first-time administrator into worker connection and lens setup", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [], workers: [], tracing_enabled: true }; if (path === "/lens/agents") return []; return { traces: true, requests: false, data: [] }; @@ -341,7 +341,10 @@ it("guides a first-time administrator into worker connection and lens setup", as onboarding: { connect, create }, }); const guide = within(await screen.findByRole("region", { name: "Get Lens running" })); - expect(apiClient.get).toHaveBeenCalledWith("/lens/activity/available", { accessToken: "test" }); + expect(proxy.get).toHaveBeenCalledWith( + "/lens/activity/available", + expect.objectContaining({ authorization: "Bearer test" }), + ); expect(guide.getByRole("button", { name: /Send your first trace/ })).toContainElement( guide.getByLabelText("Step 2 complete"), ); @@ -382,7 +385,7 @@ it("opens the saved results of an older batch", async () => { finished_at: "2026-09-29T10:02:13Z", findings: [{ ...issue, title: "Earlier batch finding" }], }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return [lens.jobs[0], older]; if (path === "/lens/lens/runs/older") return older; @@ -413,11 +416,11 @@ it("reads request content from the beginning after its abbreviated preview", asy executions: [{ ...lens.jobs[0].sample!.executions[0], id: requestId, source: "requests" as const }], }, }; - vi.mocked(apiClient.get).mockImplementation(async (path, options) => { + proxy.get.mockImplementation(async (path, options) => { if (path === "/lens") return { lenses: [{ ...lens, jobs: [job] }], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return [job]; if (!path.includes("/executions/")) return { data: [] }; - const offset = options?.query?.offset ?? 0; + const offset = Number(options.query.offset ?? 0); return { parts: [ { @@ -453,7 +456,7 @@ it.each([false, true])( async (enabled) => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => + proxy.get.mockImplementation(async (path) => path === "/lens" ? { lenses: [], workers: [], tracing_enabled: enabled } : { data: [] }, ); const user = userEvent.setup(); @@ -472,7 +475,7 @@ it("enables first-lens setup when a trace arrives without leaving Investigations window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); const traceCheck = vi.fn().mockResolvedValue({ traces: false, requests: false }); - vi.mocked(apiClient.get).mockImplementation(async (path) => + proxy.get.mockImplementation(async (path) => path === "/lens" ? { lenses: [], workers: [], tracing_enabled: true } : traceCheck(), ); vi.useFakeTimers(); @@ -502,7 +505,7 @@ it("allows retrying a failed trace readiness check without treating it as an emp .fn() .mockRejectedValueOnce(new ApiError("Trace storage unavailable", 503, {})) .mockResolvedValue({ data: [] }); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [], workers: [], tracing_enabled: true }; if (path === "/lens/activity/available") return traceCheck(); return { data: [] }; @@ -523,7 +526,7 @@ it("shows a centered failure with a retry when investigations cannot load, then .fn() .mockRejectedValueOnce(new ApiError("Proxy timed out", 504, {})) .mockResolvedValue({ lenses: [], workers: [], tracing_enabled: true }); - vi.mocked(apiClient.get).mockImplementation(async (path) => (path === "/lens" ? list() : { data: [] })); + proxy.get.mockImplementation(async (path) => (path === "/lens" ? list() : { data: [] })); const user = userEvent.setup(); renderWithProviders(); const alert = await screen.findByRole("alert"); @@ -536,7 +539,7 @@ it("shows a centered failure with a retry when investigations cannot load, then it("keeps saved investigations accessible when tracing is disabled", async () => { testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: false }; if (path === "/lens/lens/runs") return lens.jobs; return { data: [] }; @@ -549,7 +552,7 @@ it("keeps saved investigations accessible when tracing is disabled", async () => it("allows request-only accounts to connect a worker without requiring agent traces", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [], workers: [], tracing_enabled: true }; if (path === "/lens/activity/available") return { traces: false, requests: true }; return { data: [] }; @@ -565,7 +568,7 @@ it("allows request-only accounts to connect a worker without requiring agent tra it("reopens the inline editor from a shared link and drops it from the URL on cancel", async () => { testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return lens.jobs; if (path === "/lens/agents") return []; @@ -587,12 +590,13 @@ it("reopens the inline editor from a shared link and drops it from the URL on ca const url = new URLSearchParams(String(onUrlUpdate.mock.lastCall?.[0].queryString ?? "")); expect(url.has("dialog")).toBe(false); expect(url.get("lens")).toBe(lens.id); - expect(apiClient.request).not.toHaveBeenCalled(); + expect(proxy.put).not.toHaveBeenCalled(); + expect(sentBody(proxy.post, "/lens")).toEqual([]); }); it("reopens a finding and a results section from shared links", async () => { testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return lens.jobs; return { data: [] }; @@ -628,7 +632,7 @@ it("steps across findings and investigations with J and K, skipping hidden findi findings: [twinIssue], jobs: lens.jobs.map((job) => ({ ...job, findings: [twinIssue] })), }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens, twin], tracing_enabled: true, workers: [] }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path.endsWith("/runs")) return []; @@ -665,7 +669,7 @@ it("steps across findings and investigations with J and K, skipping hidden findi it("opens an investigation beside the list and walks from it into its findings with J and K", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], tracing_enabled: true, workers: [] }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path.endsWith("/runs")) return []; @@ -698,13 +702,13 @@ it("lists each finding under the investigation that owns it and resolves only th window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); const twin: Lens = { ...lens, id: "twin", settings: { ...lens.settings, name: "Twin reviews" } }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens, twin], tracing_enabled: true, workers: [] }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path.endsWith("/runs")) return []; return { data: [] }; }); - vi.mocked(apiClient.patch).mockResolvedValue(undefined); + proxy.patch.mockResolvedValue(undefined); const user = userEvent.setup(); renderWithProviders(); const rows = await screen.findAllByRole("row", { name: issue.title }); @@ -714,14 +718,14 @@ it("lists each finding under the investigation that owns it and resolves only th expect(remaining.previousElementSibling).toBe(screen.getByRole("row", { name: twin.settings.name })); await user.click(remaining); await user.click(await screen.findByRole("button", { name: "Mark resolved" })); - await waitFor(() => expect(apiClient.patch).toHaveBeenCalledTimes(1)); - expect(vi.mocked(apiClient.patch).mock.calls[0][0]).toBe("/lens/twin/findings/issue"); + await waitFor(() => expect(proxy.patch).toHaveBeenCalledTimes(1)); + expect(proxy.patch.mock.calls[0][0]).toBe("/lens/twin/findings/issue"); }); it("lists investigations without edit or run controls for read-only viewers", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], tracing_enabled: true, workers: [] }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path === "/lens/lens/runs") return []; @@ -740,7 +744,7 @@ it("lists investigations without edit or run controls for read-only viewers", as it("opens investigations from the keyboard without treating nested edit keys as row activation", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], tracing_enabled: true, workers: [] }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path === "/lens/lens/runs") return []; @@ -779,7 +783,7 @@ it("opens a failed investigation's details from its row and edits only from the error: "boom", findings: [], }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [{ ...lens, jobs: [job] }], workers: [], tracing_enabled: true }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path === "/lens/lens/runs") return [job]; @@ -804,7 +808,7 @@ it("shows the actual saved failure and run context without opening backend logs" "Grouping observations failed: Clusters response invalid after 2 attempts.\n" + "candidates.0.check_id: Field required [missing]"; const job = { ...lens.jobs[0], id: "failed-run", status: "failed" as const, stage: "Failed", error, findings: [] }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [{ ...lens, jobs: [job] }], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return [job]; if (path === "/lens/lens/runs/failed-run") return job; @@ -822,14 +826,14 @@ it("keeps a finding open to retry when its update fails", async () => { window.history.replaceState({}, "", "/lens/"); testQueryClient.clear(); const twin: Lens = { ...lens, id: "twin", settings: { ...lens.settings, name: "Twin reviews" } }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens, twin], tracing_enabled: true, workers: [] }; if (path === "/lens/activity/available") return { traces: true, requests: false }; if (path.endsWith("/runs")) return []; return { data: [] }; }); - vi.mocked(apiClient.patch).mockReset(); - vi.mocked(apiClient.patch).mockImplementation(async (path) => { + proxy.patch.mockReset(); + proxy.patch.mockImplementation(async (path) => { if (String(path).startsWith("/lens/twin/")) throw new Error("Twin reviews could not be updated"); }); const user = userEvent.setup(); @@ -844,28 +848,24 @@ it("keeps a finding open to retry when its update fails", async () => { it("pauses monitoring from the detail menu by saving the investigation with monitoring off", async () => { testQueryClient.clear(); const watching: Lens = { ...lens, settings: { ...lens.settings, enabled: true } }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [watching], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return watching.jobs; return { data: [] }; }); - vi.mocked(apiClient.request).mockReset(); - vi.mocked(apiClient.request).mockResolvedValue({ ...watching, settings: lens.settings }); + proxy.put.mockResolvedValue({ ...watching, settings: lens.settings }); const user = userEvent.setup(); renderWithProviders(); await user.click(await screen.findByRole("button", { name: "Investigation actions" })); await user.click(await screen.findByRole("menuitem", { name: "Pause monitoring" })); - await waitFor(() => expect(apiClient.request).toHaveBeenCalledTimes(1)); - expect(apiClient.request).toHaveBeenCalledWith("PUT", "/lens/lens", { - accessToken: "test", - body: { ...watching.settings, enabled: false }, - }); + await waitFor(() => expect(proxy.put).toHaveBeenCalledTimes(1)); + expect(sentBody(proxy.put, "/lens/lens")).toEqual([{ ...watching.settings, enabled: false }]); }); it("cancels the running job from the progress banner", async () => { testQueryClient.clear(); const running = { ...lens.jobs[0], id: "live", status: "running" as const, stage: "Reading executions" }; - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [{ ...lens, jobs: [running, lens.jobs[0]] }], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return [running, lens.jobs[0]]; @@ -874,15 +874,13 @@ it("cancels the running job from the progress banner", async () => { const user = userEvent.setup(); renderWithProviders(); await user.click(await screen.findByRole("button", { name: "Cancel" })); - await waitFor(() => - expect(apiClient.post).toHaveBeenCalledWith("/lens/lens/cancel", { accessToken: "test", body: {} }), - ); + await waitFor(() => expect(proxy.post).toHaveBeenCalledWith("/lens/lens/cancel", expect.anything())); }); it("refreshes run history as soon as the list reports a job the scheduler started", async () => { testQueryClient.clear(); const runs = vi.fn().mockResolvedValue(lens.jobs); - vi.mocked(apiClient.get).mockImplementation(async (path) => { + proxy.get.mockImplementation(async (path) => { if (path === "/lens") return { lenses: [lens], workers: [], tracing_enabled: true }; if (path === "/lens/lens/runs") return runs(); return { data: [] }; diff --git a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx index c668bba3799..a99133d6ca2 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx @@ -5,7 +5,7 @@ import { useQuery } from "@tanstack/react-query"; import { Plus } from "lucide-react"; import { Button } from "@/components/ui/button"; import { cn } from "@/lib/cva.config"; -import { LensPreviewButton } from "@/components/view_logs/TraceView/LensPreviewButton"; +import { LensPreviewButton } from "@/components/lens/ui/LensPreviewButton"; import { useInvalidateLenses } from "../data/mutations"; import { lensQueries } from "../data/queries"; diff --git a/ui/litellm-dashboard/src/components/lens/onboarding/LensGettingStarted.tsx b/ui/litellm-dashboard/src/components/lens/onboarding/LensGettingStarted.tsx index 7f27678c344..7b64391c85a 100644 --- a/ui/litellm-dashboard/src/components/lens/onboarding/LensGettingStarted.tsx +++ b/ui/litellm-dashboard/src/components/lens/onboarding/LensGettingStarted.tsx @@ -30,9 +30,9 @@ export function LensGettingStarted({ state, onStart, onExit, onDemo }: LensGetti setupRef.current?.focus({ preventScroll: true }); }; return ( -
+
-
+
!next && close()}>