diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 7bae178791d..e01d52a55a6 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4478,6 +4478,7 @@ dependencies = [ "flate2", "futures-util", "hmac 0.12.1", + "itertools 0.14.0", "jsonschema", "litellm-http", "litellm-migrate", diff --git a/litellm-rust/crates/traces-clickhouse/Cargo.toml b/litellm-rust/crates/traces-clickhouse/Cargo.toml index 3ab0bfce9fa..15b1ca8162e 100644 --- a/litellm-rust/crates/traces-clickhouse/Cargo.toml +++ b/litellm-rust/crates/traces-clickhouse/Cargo.toml @@ -16,6 +16,7 @@ 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 diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql b/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql new file mode 100644 index 00000000000..679edfdec2e --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql @@ -0,0 +1,31 @@ +SELECT * FROM ( +SELECT o.TraceId AS trace_id, o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name, + o.ObservationType AS type, toUInt8(o.WrapperCandidate) AS wrapper_candidate, o.AgentName AS agent, + o.Framework AS framework, o.StatusCode AS status, + substringUTF8(o.StatusMessage, 1, 128) AS status_message, + lengthUTF8(o.StatusMessage) > 128 AS error_truncated, + toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns, + o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model, + o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens, + o.LiteLLMRequestId AS litellm_request_id, + o.CallKeys AS call_keys, o.CallEvidence AS call_evidence, + -- Rows written before ToolCallId keep the call id only in their attributes. + if(o.ToolCallId != '' OR o.ObservationType != 'tool', o.ToolCallId, + coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) + AS tool_call_id, + o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash +FROM otel_traces AS o +WHERE o.Timestamp >= fromUnixTimestamp64Milli({start_ms:Int64}) + AND o.Timestamp < fromUnixTimestamp64Milli({end_ms:Int64}) + AND ({all_teams:UInt8} = 1 + OR ({user_id:String} != '' AND o.UserId = {user_id:String}) + OR has({team_ids:Array(String)}, o.TeamId)) + AND hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) IN {trace_refs:Array(String)} + AND o.EngineReceivedMs <= {snapshot_ms:UInt64} +ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage +LIMIT 1 BY o.TeamId, o.ApiKeyHash, o.TraceId, o.SpanId + +) +WHERE (team_id, api_key_hash, trace_id, span_id) > ({after_team:String}, {after_key:String}, {after_trace:String}, {after_span:String}) +ORDER BY team_id, api_key_hash, trace_id, span_id +LIMIT {page_size:UInt32} diff --git a/litellm-rust/crates/traces-clickhouse/src/reads.rs b/litellm-rust/crates/traces-clickhouse/src/reads.rs index 5949b546335..4d338e17f07 100644 --- a/litellm-rust/crates/traces-clickhouse/src/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/src/reads.rs @@ -4,6 +4,8 @@ use std::sync::{Arc, 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::{ @@ -19,7 +21,7 @@ use crate::{ query::named::{ ListTracesParams, ListTracesRow, ReadAccessParams, SpanDetail as SpanDetailQuery, SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIdsParams, TraceIdentity, - TraceIdentityParams, TraceSpansParams, + TraceIdentityParams, TracePageSpansParams, TraceSpansParams, }, }; @@ -38,7 +40,7 @@ impl Query for RunCandidates { // Cursor pages share a bounded snapshot so advancing does not resolve the whole graph again. static TRACE_SNAPSHOTS: LazyLock>> = LazyLock::new(|| { Cache::builder() - .max_capacity(64 * 1024 * 1024) + .max_capacity((2 * crate::span_batches::MAX_GRAPH_BYTES) as u64) .weigher(|_: &String, trace: &Arc| { serde_json::to_vec(trace.as_ref()) .ok() @@ -51,6 +53,7 @@ static TRACE_SNAPSHOTS: LazyLock>> = LazyLock::new(|| { const NANOS_PER_MS: i64 = 1_000_000; const SPEND_WINDOW_MS: i64 = 30 * 60 * 1000; +const SPEND_CONCURRENCY: usize = 4; fn encode_cursor(position: &T) -> String { URL_SAFE.encode(serde_json::to_vec(position).unwrap_or_default()) @@ -194,19 +197,86 @@ pub async fn list_traces( .last() .filter(|_| page.len() == params.0.limit as usize) .map(|last| encode_cursor(&(last.start_ms, &last.trace_ref))); - let mut data = Vec::with_capacity(page.len()); - for row in &page { - let summary = - match get_trace(client, connection, access, &row.trace_id, &row.trace_ref).await { - Ok(trace) => trace.map_or_else(|| listed_summary(row), |trace| trace.summary), - Err(Error::ReadTooLarge) => listed_summary(row), - Err(error) => return Err(error), - }; - data.push(summary); - } + let data = stream::iter(page.chunks(16)) + .then(|batch| list_summaries(client, connection, access, batch)) + .try_collect::>() + .await? + .into_iter() + .flatten() + .collect(); Ok(TracePage { data, next_cursor }) } +async fn list_summaries( + client: &Client, + connection: &Connection, + access: &ReadAccessParams, + runs: &[contracts::ListTracesRow], +) -> Result, Error> { + let (Some(start_ms), Some(end_ms)) = ( + runs.iter().map(|row| row.start_ms).min(), + runs.iter() + .map(|row| row.start_ms.saturating_add(row.duration_ms)) + .max(), + ) else { + return Ok(Vec::new()); + }; + let params = TracePageSpansParams::from(contracts::TracePageSpansParams { + access: access.clone(), + trace_refs: runs.iter().map(|row| row.trace_ref.clone()).collect(), + start_ms, + end_ms: end_ms.saturating_add(1), + }); + let spans = match crate::span_batches::read_list_spans(client, connection, params).await { + Ok(spans) => spans, + Err(Error::ReadTooLarge) => { + return stream::iter(runs) + .then(|row| async move { + match get_trace(client, connection, access, &row.trace_id, &row.trace_ref).await + { + Ok(trace) => { + Ok(trace.map_or_else(|| listed_summary(row), |trace| trace.summary)) + } + Err(Error::ReadTooLarge) => Ok(listed_summary(row)), + Err(error) => Err(error), + } + }) + .try_collect() + .await; + } + Err(error) => return Err(error), + }; + let by_trace = spans.into_iter().into_group_map_by(|span| { + ( + span.team_id.clone(), + span.api_key_hash.clone(), + span.trace_id.clone(), + ) + }); + let summaries = runs + .iter() + .map(|row| { + let spans = by_trace + .get(&( + row.team_id.clone(), + row.api_key_hash.clone(), + row.trace_id.clone(), + )) + .map(Vec::as_slice) + .unwrap_or_default(); + async move { + let spend_rows = spend(client, connection, access, spans).await; + resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows) + .map_or_else(|| listed_summary(row), |trace| trace.summary) + } + }) + .collect::>(); + Ok(stream::iter(summaries) + .buffered(SPEND_CONCURRENCY) + .collect() + .await) +} + pub async fn get_trace( client: &Client, connection: &Connection, @@ -293,6 +363,13 @@ pub async fn get_trace_page( let Some(trace) = resolve_trace(trace_id, &trace_ref, &rows, &spend_rows) else { return Ok(None); }; + if serde_json::to_vec(&trace) + .map_err(|_| Error::InvalidResponse)? + .len() + > crate::span_batches::MAX_GRAPH_BYTES + { + return Err(Error::ReadTooLarge); + } let trace = Arc::new(trace); TRACE_SNAPSHOTS.insert(key, Arc::clone(&trace)).await; trace diff --git a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs index 1d903ff0f54..d66ad9506da 100644 --- a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs +++ b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs @@ -1,3 +1,5 @@ +use futures_util::{TryStreamExt, stream}; +use itertools::Itertools; use litellm_http::Client; use litellm_storage_clickhouse::{Query, fetch}; use litellm_traces::query::named as contracts; @@ -6,7 +8,7 @@ use serde::Serialize; use crate::{Connection, Error, query::named::TraceSpansRow}; const PAGE_SIZE: u32 = 256; -const MAX_GRAPH_BYTES: usize = 64 * 1024 * 1024; +pub(crate) const MAX_GRAPH_BYTES: usize = 64 * 1024 * 1024; const MAX_GRAPH_SPANS: usize = 100_000; #[derive(Default)] @@ -16,18 +18,21 @@ struct ReadBudget { } impl ReadBudget { - fn reserve(&mut self, bytes: usize) -> Result<(), Error> { - self.bytes = self.bytes.saturating_add(bytes); - if self.bytes > MAX_GRAPH_BYTES || self.rows == MAX_GRAPH_SPANS { + fn checked_add(&self, bytes: usize, rows: usize) -> Result { + let next = Self { + bytes: self.bytes.saturating_add(bytes), + rows: self.rows.saturating_add(rows), + }; + if next.bytes > MAX_GRAPH_BYTES || next.rows > MAX_GRAPH_SPANS { return Err(Error::ReadTooLarge); } - self.rows += 1; - Ok(()) + Ok(next) } fn record(&mut self, row: &impl Serialize) -> Result<(), Error> { let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?; - self.reserve(bytes.len()) + *self = self.checked_add(bytes.len(), 1)?; + Ok(()) } } @@ -92,6 +97,90 @@ pub(crate) async fn read_spans( } } +#[derive(Serialize)] +struct ListParameters { + #[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; + type Row = TraceSpansRow; + + const SQL: &'static str = include_str!("../query/trace_list_span_batch.sql"); +} + +pub(crate) async fn read_list_spans( + client: &Client, + connection: &Connection, + runs: crate::query::named::TracePageSpansParams, +) -> Result, Error> { + let parameters = ListParameters { + 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, + }; + let pages = stream::try_unfold( + (Some(parameters), ReadBudget::default()), + |(parameters, budget)| async move { + let Some(parameters) = parameters else { + return Ok(None); + }; + let page = match fetch::(client, connection, ¶meters).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) + if parameters.page_size > 1 => + { + let retry = ListParameters { + page_size: parameters.page_size / 2, + ..parameters + }; + return Ok(Some((Vec::new(), (Some(retry), budget)))); + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => result?, + }; + let next = page + .last() + .filter(|_| page.len() == parameters.page_size as usize) + .map(|last| ListParameters { + after_team: last.0.team_id.clone(), + after_key: last.0.api_key_hash.clone(), + after_trace: last.0.trace_id.clone(), + after_span: last.0.span_id.clone(), + page_size: (parameters.page_size * 2).min(PAGE_SIZE), + ..parameters + }); + let next_budget = page.iter().try_fold(budget, |budget, row| { + let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?; + budget.checked_add(bytes.len(), 1) + })?; + Ok(Some((page, (next, next_budget)))) + }, + ) + .try_collect::>() + .await?; + Ok(pages + .into_iter() + .flatten() + .map(|row| row.0) + .sorted_by_key(|row| row.start_ns) + .collect()) +} + #[derive(Serialize)] struct SpendParameters { #[serde(flatten)] @@ -175,7 +264,7 @@ mod tests { #[case] next: usize, #[case] rejected: bool, ) { - let mut budget = ReadBudget { bytes, rows }; - assert_eq!(budget.reserve(next).is_err(), rejected); + let budget = ReadBudget { bytes, rows }; + assert_eq!(budget.checked_add(next, 1).is_err(), rejected); } } diff --git a/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs b/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs index 49499ebfa39..73b5125929a 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs @@ -174,7 +174,7 @@ async fn admin_sql_enforces_result_row_limit( matches!( result, Err(Error::Storage( - litellm_storage_clickhouse::Error::QueryFailed(_) + litellm_storage_clickhouse::Error::ResponseTooLarge )) ), "{result:?}" diff --git a/litellm-rust/crates/traces-clickhouse/tests/reads.rs b/litellm-rust/crates/traces-clickhouse/tests/reads.rs index ff4461e2157..41820065d84 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/reads.rs @@ -14,11 +14,98 @@ mod support; use fixtures::{DATABASE, SeededDatabase, migrated_database, seeded_database}; use support::TestResult; +#[rstest] +#[case::api_key("key-a", "")] +#[case::user("", "user-a")] +#[tokio::test] +async fn list_costs_match_each_run_when_response_ids_are_reused( + #[future(awt)] migrated_database: TestResult, + #[case] api_key: &str, + #[case] user_id: &str, +) -> TestResult { + let fixture = migrated_database?; + let client = &fixture.database.client; + let writer = Connection::writer(&fixture.database.url)?; + let runs = [ + ("earlier-run", 1_790_000_000_000_i64, 0.25), + ("later-run", 1_790_007_200_000_i64, 0.75), + ]; + insert_rows( + client, + &writer, + DATABASE, + InsertTable::OtelTraces, + runs.iter() + .map(|(trace_id, start_ms, _)| { + BTreeMap::from([ + ("Timestamp".into(), json!(start_ms * 1_000_000)), + ("TraceId".into(), json!(trace_id)), + ("SpanId".into(), json!("llm-span")), + ("ObservationType".into(), json!("llm")), + ("TeamId".into(), json!("team-a")), + ("ApiKeyHash".into(), json!(api_key)), + ("UserId".into(), json!(user_id)), + ("Duration".into(), json!(1_000_000)), + ("LiteLLMRequestId".into(), json!("reused-response")), + ]) + }) + .collect(), + ) + .await?; + insert_rows( + client, + &writer, + DATABASE, + InsertTable::SpendLogs, + runs.iter() + .map(|(trace_id, start_ms, cost)| { + BTreeMap::from([ + ("request_id".into(), json!(format!("request-{trace_id}"))), + ("response_id".into(), json!("reused-response")), + ("team_id".into(), json!("team-a")), + ("api_key".into(), json!(api_key)), + ("user".into(), json!(user_id)), + ("start_time".into(), json!(start_ms)), + ("end_time".into(), json!(start_ms + 1)), + ("spend".into(), json!(cost)), + ]) + }) + .collect(), + ) + .await?; + let reader = fixture + .readers + .connection(client, &QueryScope::All, "fixture-secret") + .await?; + 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?; + assert_eq!(page.data.len(), runs.len()); + for (trace_id, _, cost) in runs { + let summary = page + .data + .iter() + .find(|summary| summary.trace_id == trace_id) + .ok_or("missing run")?; + let detail = get_trace(client, &reader, &access, trace_id, &summary.trace_ref) + .await? + .ok_or("missing trace")?; + assert_eq!(detail.summary.spend, Some(cost)); + assert_eq!(summary.spend, detail.summary.spend, "{trace_id}"); + } + Ok(()) +} + #[rstest] #[case::many_runs(50, 21, 0, false)] #[case::one_large_run(1, 1100, 0, false)] #[case::large_rows(1, 280, 20_000, false)] +#[case::large_cached_snapshot(1, 280, 140_000, false)] #[case::many_costs(1, 1101, 0, true)] +#[case::many_costed_runs(500, 2, 0, true)] #[tokio::test] async fn large_runs_remain_complete_under_default_reader_limits( #[future(awt)] migrated_database: TestResult, @@ -30,9 +117,9 @@ async fn large_runs_remain_complete_under_default_reader_limits( let fixture = migrated_database?; let client = &fixture.database.client; let writer = Connection::writer(&fixture.database.url)?; - for run in 0..runs { - let rows = (0..steps) - .map(|step| { + let rows = (0..runs) + .flat_map(|run| { + (0..steps).map(move |step| { BTreeMap::from([ ( "Timestamp".into(), @@ -75,17 +162,17 @@ async fn large_runs_remain_complete_under_default_reader_limits( ), ]) }) - .collect::>(); - for chunk in rows.chunks(100) { - insert_rows( - client, - &writer, - DATABASE, - InsertTable::OtelTraces, - chunk.to_vec(), - ) - .await?; - } + }) + .collect::>(); + for chunk in rows.chunks(100) { + insert_rows( + client, + &writer, + DATABASE, + InsertTable::OtelTraces, + chunk.to_vec(), + ) + .await?; } if costed { let costs = (1..steps) @@ -112,8 +199,59 @@ async fn large_runs_remain_complete_under_default_reader_limits( 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 = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 500).await?; assert_eq!(page.data.len(), runs); + assert!( + page.data + .windows(2) + .all(|runs| runs[0].trace_ref > runs[1].trace_ref) + ); + if runs > 1 { + client + .post(writer.url().clone()) + .body("SYSTEM FLUSH LOGS") + .send() + .await? + .error_for_status()?; + let read_queries = client.post(writer.url().clone()).body(format!( + "SELECT count() FROM system.query_log WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' AND query LIKE '%FROM otel_traces AS o%' AND query NOT LIKE '%system.query_log%'" + )).send().await?.error_for_status()?.text().await?; + let read_queries = read_queries.trim().parse::()?; + assert!( + read_queries > 0 && read_queries < runs, + "{read_queries} span queries for {runs} runs" + ); + if costed { + let overlapping = client + .post(writer.url().clone()) + .body(format!( + "WITH spend_reads AS ( + SELECT query_start_time_microseconds AS started, event_time_microseconds AS finished + FROM system.query_log + WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' + AND query LIKE '%FROM spend_logs FINAL%' AND query NOT LIKE '%system.query_log%' + ), events AS ( + SELECT started AS at, 1 AS delta FROM spend_reads + UNION ALL SELECT finished AS at, -1 AS delta FROM spend_reads + ) + SELECT max(active) FROM ( + SELECT sum(delta) OVER (ORDER BY at, delta ROWS UNBOUNDED PRECEDING) AS active + FROM events + )" + )) + .send() + .await? + .error_for_status()? + .text() + .await? + .trim() + .parse::()?; + assert!( + (2..=4).contains(&overlapping), + "{overlapping} simultaneous spend reads for {runs} runs" + ); + } + } for summary in &page.data { assert_eq!(summary.span_count, steps as u64); assert_eq!( @@ -151,6 +289,15 @@ async fn large_runs_remain_complete_under_default_reader_limits( }, (steps - 1) as u64 ); + let denied = ReadAccessParams { + team_ids: vec!["other-team".into()], + ..access.clone() + }; + assert!( + get_trace(client, &reader, &denied, "trace-0000", trace_ref) + .await? + .is_none() + ); let mut cursor = None; let mut ids = Vec::new(); loop { @@ -171,6 +318,27 @@ async fn large_runs_remain_complete_under_default_reader_limits( serde_json::to_vec(&page)?.len() <= litellm_storage_clickhouse::READ_LIMITS.response_bytes ); + if ids.is_empty() { + assert!( + get_trace_page( + client, + &reader, + &denied, + "trace-0000", + trace_ref, + page.next_cursor.as_deref(), + 200, + ) + .await? + .is_none() + ); + client + .post(writer.url().clone()) + .body(format!("TRUNCATE TABLE {DATABASE}.otel_traces")) + .send() + .await? + .error_for_status()?; + } ids.extend(page.spans.into_iter().map(|span| span.span_id)); cursor = page.next_cursor; if cursor.is_none() { @@ -185,15 +353,6 @@ async fn large_runs_remain_complete_under_default_reader_limits( .map(|span| span.span_id.clone()) .collect::>() ); - let denied = ReadAccessParams { - team_ids: vec!["other-team".into()], - ..access - }; - assert!( - get_trace(client, &reader, &denied, "trace-0000", trace_ref) - .await? - .is_none() - ); Ok(()) } diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py index 676934ba3e0..5da08e6af75 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py @@ -19,20 +19,23 @@ defaults to ``True`` so a config dict (raw, not Pydantic) without an ``auth`` key still requires authentication. """ -from unittest.mock import AsyncMock, MagicMock +from typing import Final +from unittest.mock import MagicMock import pytest from fastapi import FastAPI from fastapi.routing import APIRoute +from fastapi.testclient import TestClient -from litellm.proxy._types import PassThroughGenericEndpoint +from litellm.proxy._types import PassThroughGenericEndpoint, ProxyException from litellm.proxy.auth.user_api_key_auth import ( check_api_key_for_custom_headers_or_pass_through_endpoints, ) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( _register_pass_through_endpoint, ) +from litellm.proxy.proxy_server import openai_exception_handler def test_passthrough_auth_defaults_to_true(): @@ -58,20 +61,17 @@ def test_passthrough_auth_can_still_be_explicitly_disabled(): @pytest.mark.asyncio -async def test_register_passthrough_with_auth_true_works_for_oss(monkeypatch): - # Regression: setting ``auth: true`` used to raise at startup - # unless ``premium_user`` was True, leaving OSS with no safe - # configuration. - app = FastAPI() - visited: set = set() +async def test_register_passthrough_with_auth_true_works_for_oss(monkeypatch: pytest.MonkeyPatch) -> None: + app: Final = FastAPI(exception_handlers={ProxyException: openai_exception_handler}) + visited: Final[set[str]] = set() + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-passthrough-test") - endpoint = PassThroughGenericEndpoint( + endpoint: Final = PassThroughGenericEndpoint( path="/forwarder", target="https://example.com", auth=True, ) - # Should not raise; OSS premium_user=False is allowed to use auth=True. await _register_pass_through_endpoint( endpoint=endpoint, app=app, @@ -79,6 +79,10 @@ async def test_register_passthrough_with_auth_true_works_for_oss(monkeypatch): visited_endpoints=visited, ) assert [route.path for route in app.routes if isinstance(route, APIRoute)] == ["/forwarder"] + with TestClient(app) as client: + response: Final = client.get(endpoint.path) + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "auth_error" @pytest.mark.asyncio diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx index d055426257c..bd881730925 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.test.tsx @@ -1,4 +1,5 @@ -import { screen, within } from "@testing-library/react"; +import { screen, within, waitFor, act } from "@testing-library/react"; +import { focusManager, onlineManager } from "@tanstack/react-query"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; @@ -46,6 +47,7 @@ const rootSpanId = (trace: Trace): string => trace.spans.find((s) => s.parent_sp describe("RunView", () => { beforeEach(() => { testQueryClient.clear(); + testQueryClient.setQueryDefaults(["agentTrace"], {}); vi.mocked(copyToClipboard).mockClear(); }); @@ -188,6 +190,68 @@ describe("RunView", () => { expect(vi.mocked(agentTraceCall).mock.calls.map((call) => call[3])).toEqual([null, "next-page", "next-page"]); }); + it("refreshes a failed later page from one new snapshot", async () => { + const user = userEvent.setup(); + const summary = { ...research.summary, span_count: 3 }; + const first: Trace = { ...research, summary, spans: research.spans.slice(0, 1), next_cursor: "old-second" }; + const second: Trace = { + ...research, + summary, + spans: [ + { ...research.spans[1], type: "tool", name: "old-snapshot-tool", parent_span_id: research.spans[0].span_id }, + ], + next_cursor: "old-third", + }; + const fresh: Trace = { + ...first, + summary: { ...summary, span_count: 2 }, + spans: [{ ...research.spans[0], name: "fresh-root" }], + next_cursor: "fresh-second", + }; + const freshSecond: Trace = { ...fresh, spans: [{ ...second.spans[0], name: "fresh-tool" }], next_cursor: null }; + vi.mocked(agentTraceCall).mockReset(); + vi.mocked(agentTraceCall) + .mockResolvedValueOnce(first) + .mockResolvedValueOnce(second) + .mockRejectedValueOnce(new Error("Trace changed while paging; refresh the trace")) + .mockResolvedValueOnce(fresh) + .mockResolvedValueOnce(freshSecond); + renderWithProviders(); + await user.click(await screen.findByRole("button", { name: "Load more steps" })); + expect(await screen.findByText("old-snapshot-tool")).toBeVisible(); + await user.click(screen.getByRole("button", { name: "Load more steps" })); + await user.click(await screen.findByRole("button", { name: "Refresh trace" })); + expect(await screen.findByText("Showing 1 of 2 steps")).toBeVisible(); + expect(screen.queryByText("old-snapshot-tool")).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Load more steps" })); + expect(await screen.findByText("fresh-tool")).toBeVisible(); + expect(screen.getAllByRole("treeitem")).toHaveLength(2); + expect(vi.mocked(agentTraceCall).mock.calls.map((call) => call[3])).toEqual([ + null, + "old-second", + "old-third", + null, + "fresh-second", + ]); + }); + + it("keeps a loaded snapshot on focus and reconnect", async () => { + testQueryClient.setQueryDefaults(["agentTrace"], { refetchOnWindowFocus: true, refetchOnReconnect: true }); + vi.mocked(agentTraceCall).mockReset(); + renderRun(research); + await screen.findByTestId("detail-pane"); + await testQueryClient.invalidateQueries({ queryKey: ["agentTrace"], refetchType: "none" }); + await act(async () => { + focusManager.setFocused(false); + onlineManager.setOnline(false); + focusManager.setFocused(true); + onlineManager.setOnline(true); + }); + await waitFor(() => expect(testQueryClient.isFetching()).toBe(0)); + expect(vi.mocked(agentTraceCall)).toHaveBeenCalledTimes(1); + expect(screen.getByRole("tree", { name: "Spans in time order" })).toHaveTextContent(research.spans[0].name); + }); + it("keeps a way back to the runs table when a run fails to load", async () => { const user = userEvent.setup(); const onBack = vi.fn(); diff --git a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx index 6e8464b797e..3ff8bd0ef4d 100644 --- a/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/TraceView/TraceDrawer.tsx @@ -2,7 +2,7 @@ import { useLensDemo } from "@/components/lens/LensDemoContext"; import { useTracesApi } from "@/components/lens/services"; -import { useInfiniteQuery } from "@tanstack/react-query"; +import { useInfiniteQuery, useQueryClient } from "@tanstack/react-query"; import { ArrowLeft, Check, Copy } from "lucide-react"; import { useCallback, useEffect, useMemo, useState } from "react"; @@ -374,6 +374,7 @@ function initialSpanMissing(trace: Trace | undefined, spanId?: string): boolean export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, embedded = false }: RunViewProps) { const traces = useTracesApi(accessToken); + const queryClient = useQueryClient(); const [view, setView] = useState("steps"); const traceQueryOptions = { queryKey: ["agentTrace", traceId, traceRef, accessToken], @@ -381,9 +382,13 @@ export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, initialPageParam: null as string | null, getNextPageParam: (lastPage: Trace) => lastPage.next_cursor ?? undefined, staleTime: 30_000, + refetchOnWindowFocus: false, + refetchOnReconnect: false, + refetchOnMount: false, retry: false, }; const traceQuery = useInfiniteQuery(traceQueryOptions); + const refreshTrace = () => queryClient.resetQueries({ queryKey: traceQueryOptions.queryKey, exact: true }); const trace = useMemo(() => { const pages = traceQuery.data?.pages; if (!pages?.length) return undefined; @@ -426,7 +431,7 @@ export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack,

Could not load trace

{traceQuery.error?.message ?? "Unknown error"} - @@ -451,12 +456,7 @@ export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, : `Showing ${trace.spans.length.toLocaleString()} of ${trace.summary.span_count.toLocaleString()} steps`} {traceQuery.isError && ( - )} @@ -464,7 +464,7 @@ export function RunView({ traceId, traceRef, initialSpanId, accessToken, onBack, size="xs" variant="outline" disabled={traceQuery.isFetching} - onClick={() => void (traceQuery.hasNextPage ? traceQuery.fetchNextPage() : traceQuery.refetch())} + onClick={() => void (traceQuery.hasNextPage ? traceQuery.fetchNextPage() : refreshTrace())} > {traceQuery.isFetching ? "Loading…" : pageAction}