Merge remote-tracking branch 'origin/main' into litellm_logging_only_scope

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

# Conflicts:
#	tests/integration/observability/test_guardrail_effects.py
This commit is contained in:
yucheng 2026-10-05 09:00:54 +00:00
commit 443f9e3db4
342 changed files with 11922 additions and 4235 deletions

View file

@ -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."

View file

@ -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

View file

@ -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",
]

View file

@ -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

View file

@ -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<Error>) -> 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<TraceReader>,
}
#[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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Error>,
#[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
);
});
}
}

View file

@ -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

View file

@ -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

View file

@ -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<Self, Error> {
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, Error> {
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, Error> {
Self::digest(&(source, access, trace_id, trace_ref))
}
pub(crate) fn run(
source: &str,
access: &ReadAccessParams,
run: (&str, &str, &str, &str),
) -> Result<Self, Error> {
Self::digest(&("run", source, access, run))
}
pub(crate) fn scope(source: &str, access: &ReadAccessParams) -> Result<Self, Error> {
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<K, V: Fresh> Expiry<K, V> for ByFreshness {
fn expire_after_create(&self, _: &K, value: &V, _: std::time::Instant) -> Option<Duration> {
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<SnapshotKey, Arc<Snapshot>>,
latest: Cache<SnapshotKey, Latest>,
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>| 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<Arc<Snapshot>> {
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<E, F>(
&self,
key: SnapshotKey,
snapshot_ms: u64,
load: F,
) -> Result<Arc<Snapshot>, Arc<E>>
where
E: From<Error> + Send + Sync + 'static,
F: Future<Output = Result<(Trace, Freshness), E>>,
{
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<E, F, Fut>(
&self,
latest: SnapshotKey,
now_ms: u64,
load_at: F,
) -> Result<Arc<Snapshot>, Arc<E>>
where
E: Send + Sync + 'static,
F: Fn(u64) -> Fut,
Fut: Future<Output = Result<Arc<Snapshot>, Arc<E>>>,
{
let entry = self
.latest
.try_get_with(latest, async {
let snapshot = load_at(now_ms).await?;
Ok::<_, Arc<E>>(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<Snapshot, 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)?));
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<TraceSummary>, 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<SnapshotKey, ListedRun>,
pub(crate) limits: Cache<SnapshotKey, u32>,
}
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
);
}
}

View file

@ -0,0 +1,124 @@
use base64::{Engine, engine::general_purpose::URL_SAFE};
use serde::{Deserialize, Serialize};
use crate::ReadError;
pub(super) fn encode_cursor<T: Serialize>(position: &T) -> String {
URL_SAFE.encode(serde_json::to_vec(position).unwrap_or_default())
}
pub(super) fn decode_cursor<T: for<'de> Deserialize<'de>, E>(
cursor: &str,
kind: &'static str,
) -> Result<T, ReadError<E>> {
URL_SAFE
.decode(cursor)
.ok()
.and_then(|json| serde_json::from_slice(&json).ok())
.ok_or(ReadError::InvalidCursor(kind))
}
pub(super) fn trace_position<E>(cursor: Option<&str>) -> Result<(i64, String), ReadError<E>> {
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<E>(
cursor: Option<&str>,
) -> Result<Option<ErrorPosition>, ReadError<E>> {
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::<std::io::Error>(Some(&cursor)).unwrap(),
(
1_790_742_989_377,
"4bad42b84e9de3ba46fc870185f8f023".to_owned()
)
);
assert_eq!(
trace_position::<std::io::Error>(None).unwrap(),
(0, String::new())
);
assert_eq!(
trace_position::<std::io::Error>(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<std::io::Error>> = 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<Option<ErrorPosition>, ReadError<std::io::Error>> =
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<Option<ErrorPosition>, ReadError<std::io::Error>> =
error_position(Some(&cursor));
assert!(matches!(
result,
Err(ReadError::InvalidCursor("diagnostic"))
));
}
}

View file

@ -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<E> {
#[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<serde_json::Error>),
#[error(transparent)]
Store(Arc<E>),
}
impl<E> Clone for ReadError<E> {
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<E> From<Error> for ReadError<E> {
fn from(error: Error) -> Self {
match error {
Error::ReadTooLarge => Self::TooLarge,
Error::Serialization(error) => Self::Encode(Arc::new(error)),
}
}
}

View file

@ -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<Self, Error> {
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<SnapshotKey, Arc<Snapshot>>,
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>| snapshot.weight)
.time_to_live(ttl)
.build(),
max_graph_bytes,
}
}
pub async fn get(&self, key: &SnapshotKey) -> Option<Arc<Snapshot>> {
self.entries.get(key).await
}
pub async fn insert(&self, key: SnapshotKey, trace: Trace) -> Result<Arc<Snapshot>, 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};

View file

@ -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<E>(
source: &str,
access: &ReadAccessParams,
row: &ListTracesRow,
) -> Result<SnapshotKey, ReadError<E>> {
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<S: TraceStore>(
reader: &TraceReader,
store: &S,
access: &ReadAccessParams,
runs: &[ListTracesRow],
) -> Result<Vec<TraceSummary>, ReadError<S::Error>> {
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<S: TraceStore>(
reader: &TraceReader,
store: &S,
access: &ReadAccessParams,
runs: &[&ListTracesRow],
) -> Result<Vec<Option<ListedRun>>, ReadError<S::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 {
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(&params, 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<S: TraceStore>(
reader: &TraceReader,
store: &S,
access: &ReadAccessParams,
row: &ListTracesRow,
) -> Result<Option<ListedRun>, ReadError<S::Error>> {
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<T>(runs: &[T]) -> impl Iterator<Item = &[T]> + '_ {
runs.chunks(RUNS_PER_SPAN_READ)
}

View file

@ -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<E> {
Absent,
Read(ReadError<E>),
}
impl<E> From<crate::Error> for Miss<E> {
fn from(error: crate::Error) -> Self {
Self::Read(error.into())
}
}
fn settle<T, E>(result: Result<T, Arc<Miss<E>>>) -> Result<Option<T>, ReadError<E>> {
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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
start_ms: i64,
end_ms: i64,
cursor: Option<&str>,
limit: u32,
) -> Result<TracePage, ReadError<S::Error>> {
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(&params).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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
trace_id: &str,
trace_ref: &str,
) -> Result<Option<Trace>, ReadError<S::Error>> {
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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
trace_id: &str,
trace_ref: &str,
cursor: Option<&str>,
page_size: u32,
) -> Result<Option<Trace>, ReadError<S::Error>> {
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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
trace_id: &str,
trace_ref: &str,
) -> Result<Option<Arc<Snapshot>>, ReadError<S::Error>> {
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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
trace_id: &str,
trace_ref: &str,
snapshot_ms: u64,
) -> Result<Arc<Snapshot>, Arc<Miss<S::Error>>> {
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(&params, 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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
trace_id: &str,
span_id: &str,
trace_ref: &str,
) -> Result<Option<SpanDetail>, ReadError<S::Error>> {
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(&params).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<S: TraceStore>(
&self,
store: &S,
access: &ReadAccessParams,
trace_id: &str,
span_id: &str,
trace_ref: &str,
cursor: Option<&str>,
) -> Result<Option<SpanErrorPage>, ReadError<S::Error>> {
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(&params).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<S: TraceStore>(
store: &S,
access: &ReadAccessParams,
trace_id: &str,
trace_ref: &str,
) -> Result<Option<String>, ReadError<S::Error>> {
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(&params).await.map_err(map_store_error)?;
if identities.len() > 1 {
return Err(ReadError::AmbiguousTrace);
}
Ok(identities.into_iter().next())
}
fn page<E>(
snapshot: &Snapshot,
position: &SpanPosition,
page_size: u32,
response_bytes: usize,
) -> Result<Trace, ReadError<E>> {
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<E>(error: StoreError<E>) -> ReadError<E> {
match error {
StoreError::TooLarge => ReadError::TooLarge,
StoreError::Failed(error) => ReadError::Store(Arc::new(error)),
}
}

View file

@ -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<Range<i64>> {
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<i64>,
) -> &[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<S: TraceStore>(
store: &S,
access: &ReadAccessParams,
rows: &[TraceSpansRow],
) -> Option<Vec<SpendByResponseIdsRow>> {
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(&params).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
}
}
}

View file

@ -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<E> {
#[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<Output = Result<Vec<String>, StoreError<Self::Error>>> + Send;
/// Returns `TooLarge` when the response exceeds the storage limit so the reader can halve `limit`.
fn list_runs(
&self,
params: &ListTracesParams,
) -> impl Future<Output = Result<Vec<ListTracesRow>, StoreError<Self::Error>>> + 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<Output = Result<Vec<TraceSpansRow>, StoreError<Self::Error>>> + 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<Output = Result<Vec<TraceSpansRow>, StoreError<Self::Error>>> + Send;
fn spend(
&self,
params: &SpendByResponseIdsParams,
) -> impl Future<Output = Result<Vec<SpendByResponseIdsRow>, StoreError<Self::Error>>> + Send;
fn span_detail(
&self,
params: &SpanDetailParams,
) -> impl Future<Output = Result<Option<SpanDetailRow>, StoreError<Self::Error>>> + Send;
fn span_error(
&self,
params: &SpanErrorParams,
) -> impl Future<Output = Result<Option<SpanErrorRow>, StoreError<Self::Error>>> + Send;
}

View file

@ -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<Operation, Failure>,
trace_refs: Vec<String>,
list_runs: Vec<ListTracesRow>,
trace_spans: HashMap<String, Vec<TraceSpansRow>>,
run_spans: Vec<TraceSpansRow>,
spend: Vec<SpendByResponseIdsRow>,
span_detail: Option<SpanDetailRow>,
span_error: Option<SpanErrorRow>,
list_runs_too_large_above: Option<u32>,
trace_too_large_refs: HashSet<String>,
spend_fails_above_response_ids: Option<usize>,
}
#[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<State>,
calls: Calls,
}
impl FakeStore {
fn with_spans(trace_ref: &str, spans: Vec<TraceSpansRow>) -> 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<String>) {
self.state.lock().unwrap().trace_refs = trace_refs;
}
fn set_list_runs(&self, rows: Vec<ListTracesRow>) {
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<TraceSpansRow>) {
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<FakeError>> {
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<Vec<String>, StoreError<Self::Error>> {
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<Vec<ListTracesRow>, StoreError<Self::Error>> {
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<Vec<TraceSpansRow>, StoreError<Self::Error>> {
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(&params.trace_ref) {
return Err(StoreError::TooLarge);
}
Ok(state
.trace_spans
.get(&params.trace_ref)
.cloned()
.unwrap_or_default())
}
async fn run_spans(
&self,
_: &TracePageSpansParams,
_: u64,
) -> Result<Vec<TraceSpansRow>, StoreError<Self::Error>> {
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<Vec<SpendByResponseIdsRow>, StoreError<Self::Error>> {
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<Option<SpanDetailRow>, StoreError<Self::Error>> {
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<Option<SpanErrorRow>, StoreError<Self::Error>> {
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);
}

View file

@ -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<Snapshot>, Arc<Error>> {
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());

View file

@ -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]

View file

@ -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<Error>),
}
impl From<litellm_traces_cache::Error> 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,
}
}
}

View file

@ -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() {

View file

@ -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,
};

View file

@ -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<SnapshotCache> = 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<T: Serialize>(position: &T) -> String {
URL_SAFE.encode(serde_json::to_vec(position).unwrap_or_default())
pub struct ClickHouseTraces {
client: Client,
connection: Connection,
}
fn decode_cursor<T: for<'de> Deserialize<'de>>(
cursor: &str,
kind: &'static str,
) -> Result<T, Error> {
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<Option<ErrorPosition>, Error> {
let Some(cursor) = cursor else {
return Ok(None);
};
let position = decode_cursor::<ErrorPosition>(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<Option<String>, Error> {
if !trace_ref.is_empty() {
return Ok(Some(trace_ref.to_owned()));
async fn trace_refs(
&self,
params: &contracts::TraceIdentityParams,
) -> Result<Vec<String>, StoreError<Self::Error>> {
fetch::<TraceIdentity>(&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::<TraceIdentity>(client, connection, &params).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<contracts::SpendByResponseIdsRow> {
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<Vec<contracts::ListTracesRow>, StoreError<Self::Error>> {
let storage_params = ListTracesParams::from(params.clone());
match fetch::<RunCandidates>(&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<TracePage, Error> {
if limit == 0 {
return Err(Error::InvalidParameters);
async fn trace_spans(
&self,
params: &contracts::TraceSpansParams,
snapshot_ms: u64,
) -> Result<Vec<contracts::TraceSpansRow>, StoreError<Self::Error>> {
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<contracts::ListTracesRow> = loop {
match fetch::<RunCandidates>(client, connection, &params).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::<Vec<_>>()
.await?
.into_iter()
.flatten()
.collect();
Ok(TracePage { data, next_cursor })
}
async fn list_summaries(
client: &Client,
connection: &Connection,
access: &ReadAccessParams,
runs: &[contracts::ListTracesRow],
) -> Result<Vec<litellm_traces::TraceSummary>, 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<Vec<contracts::TraceSpansRow>, StoreError<Self::Error>> {
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::<Vec<_>>();
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<Option<Trace>, 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<Option<Trace>, Error> {
if !(1..=500).contains(&page_size) {
return Err(Error::InvalidParameters);
async fn spend(
&self,
params: &contracts::SpendByResponseIdsParams,
) -> Result<Vec<contracts::SpendByResponseIdsRow>, StoreError<Self::Error>> {
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<Option<contracts::SpanDetailRow>, StoreError<Self::Error>> {
match fetch::<SpanDetailQuery>(&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<Option<contracts::SpanErrorRow>, StoreError<Self::Error>> {
let storage_params = SpanErrorParams::from(params.clone());
match fetch::<SpanError>(&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<Option<SpanDetail>, 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::<SpanDetailQuery>(client, connection, &params)
.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<Option<SpanErrorPage>, 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::<SpanError>(client, connection, &params)
.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<Error> {
StoreError::Failed(Error::Storage(error))
}

View file

@ -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);

View file

@ -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<Self, Error> {
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<Error>> {
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<Error>> {
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<K> {
#[serde(flatten)]
keyset: K,
page_size: u32,
}
trait PageSource<K: Keyset> {
fn page(
&self,
batch: &Batch<K>,
) -> impl Future<Output = Result<Vec<K::Row>, litellm_storage_clickhouse::Error>> + Send;
}
struct Paged<K>(PhantomData<K>);
impl<K: Keyset> Query for Paged<K> {
type Params = Batch<K>;
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<K: Keyset, S: PageSource<K>>(
source: &S,
keyset: K,
) -> Result<Vec<K::Row>, StoreError<Error>> {
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<K: Keyset> PageSource<K> for ClickHouse<'_> {
fn page(
&self,
batch: &Batch<K>,
) -> impl Future<Output = Result<Vec<K::Row>, litellm_storage_clickhouse::Error>> + Send {
fetch::<Paged<K>>(self.client, self.connection, batch)
}
}
async fn read_paged<K: Keyset>(
client: &Client,
connection: &Connection,
keyset: K,
) -> Result<Vec<K::Row>, StoreError<Error>> {
let source = ClickHouse { client, connection };
read_all(&source, keyset).await
}
fn by_start(mut rows: Vec<contracts::TraceSpansRow>) -> Vec<contracts::TraceSpansRow> {
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<Vec<contracts::TraceSpansRow>, Error> {
let mut parameters = Parameters {
) -> Result<Vec<contracts::TraceSpansRow>, StoreError<Error>> {
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::<SpanBatch>(client, connection, &parameters).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<Vec<contracts::TraceSpansRow>, Error> {
let parameters = ListParameters {
snapshot_ms: u64,
) -> Result<Vec<contracts::TraceSpansRow>, StoreError<Error>> {
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::<ListSpanBatch>(client, connection, &parameters).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::<Vec<_>>()
.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<Vec<contracts::SpendByResponseIdsRow>, Error> {
let mut parameters = SpendParameters {
lookup: SpendByResponseIdsParams,
) -> Result<Vec<contracts::SpendByResponseIdsRow>, StoreError<Error>> {
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::<SpendBatch>(client, connection, &parameters).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<Vec<u32>>,
}
impl PageSource<Numbers> for Table {
async fn page(
&self,
batch: &Batch<Numbers>,
) -> Result<Vec<u32>, 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::<Vec<_>>());
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]
);
}
}

View file

@ -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(

View file

@ -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)?;

View file

@ -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::<Vec<_>>();
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::<usize>()?;
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::<usize>()?;
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::<usize>()?;
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::<Vec<_>>();
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::<Vec<_>>();
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::<Vec<_>>();
while let Some(current) = cursor {
let next = get_trace_page(
client,
&reader,
&access,
&summary.trace_id,
&summary.trace_ref,
Some(&current),
1,
)
.await?
.ok_or("missing next page")?;
let next = reader
.get_trace_page(
&store,
&access,
&summary.trace_id,
&summary.trace_ref,
Some(&current),
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}");

View file

@ -15,7 +15,7 @@ pub struct ReadAccessParams {
pub team_ids: Vec<String>,
}
#[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<String>,
}
#[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<String, String>,
}
#[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,

View file

@ -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)

View file

@ -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

View file

@ -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", {})

View file

@ -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(

View file

@ -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",

View file

@ -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]:

View file

@ -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(

View file

@ -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,

View file

@ -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)

View file

@ -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,

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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")

View file

@ -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,

View file

@ -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)

View file

@ -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

View file

@ -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")

View file

@ -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.
"""

View file

@ -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,

View file

@ -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

View file

@ -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)))

View file

@ -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)))

View file

@ -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"}

View file

@ -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}"
)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)}"
)

View file

@ -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",
}
)

View file

@ -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

View file

@ -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=[
{

View file

@ -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,

View file

@ -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)

View file

@ -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 == []

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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))

View file

@ -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

View file

@ -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

View file

@ -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"

View file

@ -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",

View file

@ -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()

View file

@ -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(

View file

@ -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",

View file

@ -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()
]

View file

@ -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:

View file

@ -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,
):

View file

@ -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():
"""

View file

@ -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(

View file

@ -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)
)

View file

@ -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:

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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
}

View file

@ -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" },
},

View file

@ -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*" },

View file

@ -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",

View file

@ -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",

View file

@ -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 }) => (
<div>
@ -43,7 +42,10 @@ vi.mock("@/components/VirtualKeysPage/VirtualKeysTable", () => ({
}));
vi.mock("@/components/organisms/create_key_button", () => ({
default: () => <button type="button">Create Key</button>,
default: (props: { autoOpenCreate?: boolean; prefillData?: CreateKeyPrefillData }) => {
createKeyProps(props);
return <button type="button">Create Key</button>;
},
}));
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(<ApiKeysDashboard />);
renderWithProviders(<ApiKeysDashboard />);
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(<ApiKeysDashboard />);
renderWithProviders(<ApiKeysDashboard />);
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(<ApiKeysDashboard />, {
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(<ApiKeysDashboard />);
renderWithProviders(<ApiKeysDashboard />);
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(<ApiKeysDashboard />);
renderWithProviders(<ApiKeysDashboard />);
window.dispatchEvent(new Event("beforeunload"));

View file

@ -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<Team[] | null>(null);
const [keys, setKeys] = useState<KeyResponse[] | null>([]);
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]));

View file

@ -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";

View file

@ -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 <div data-testid="request-logs-panel" />;
},

View file

@ -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 <div data-testid="request-logs-panel">{isActive ? "active" : "inactive"}</div>;
},
}));
vi.mock("@/components/view_logs/AuditLogsPanel", () => ({
vi.mock("@/components/logs/audit/AuditLogsPanel", () => ({
default: function AuditLogsPanelMock({ isActive }: { isActive: boolean }) {
return <div data-testid="audit-logs-panel">{isActive ? "active" : "inactive"}</div>;
},

View file

@ -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";

View file

@ -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(
<NuqsAdapter>
<QueryClientProvider client={createQueryClient()}>
<MCPServers {...defaultProps} />
</QueryClientProvider>
</NuqsAdapter>,
);
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([]);

View file

@ -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<MCPServerProps> = ({ accessToken, userRole, userID, i
const [prefillData, setPrefillData] = useState<DiscoverableMCPServer | null>(null);
const [isDeletingServer, setIsDeletingServer] = useState(false);
const [byokModalServer, setByokModalServer] = useState<MCPServer | null>(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<MCPServer | null>(null);
const [deepLinkServerId, setDeepLinkServerId] = useState<string | null>(() =>
typeof window === "undefined" ? null : new URLSearchParams(window.location.search).get("fill_env_vars"),
);
const [deepLinkServerId, setDeepLinkServerId] = useState<string | null>(() => fillEnvVarsParam);
const [searchQuery, setSearchQuery] = useState<string>("");
const [sortKey, setSortKey] = useState<SortKey>("created_desc");
const isInternalUser = userRole === "Internal User";
@ -251,19 +251,9 @@ const MCPServers: React.FC<MCPServerProps> = ({ accessToken, userRole, userID, i
[envVarStatuses],
);
// Deep-link via ?fill_env_vars=<server_id> — 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),

View file

@ -62,6 +62,7 @@ describe("makeSystemOneRequest", () => {
const expectedRequest: Partial<RequestInit> = {
method: "POST",
headers: {
Accept: "application/json",
"Content-Type": "application/json",
Authorization: "Bearer session-key",
},

View file

@ -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<typeof renderWithoutNuqs>[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");

Some files were not shown because too many files have changed in this diff Show more