mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
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:
commit
443f9e3db4
342 changed files with 11922 additions and 4235 deletions
2
Makefile
2
Makefile
|
|
@ -58,7 +58,7 @@ help:
|
|||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
|
||||
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
|
||||
@echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (ARGS=\"--seed large\", LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)"
|
||||
@echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (ARGS=\"--seed large --seed-logs\", LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)"
|
||||
@echo ""
|
||||
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
|
||||
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
8
litellm-rust/Cargo.lock
generated
8
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
395
litellm-rust/crates/traces-cache/src/cache.rs
Normal file
395
litellm-rust/crates/traces-cache/src/cache.rs
Normal 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
|
||||
);
|
||||
}
|
||||
}
|
||||
124
litellm-rust/crates/traces-cache/src/cursor.rs
Normal file
124
litellm-rust/crates/traces-cache/src/cursor.rs
Normal 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"))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -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)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
197
litellm-rust/crates/traces-cache/src/list.rs
Normal file
197
litellm-rust/crates/traces-cache/src/list.rs
Normal 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(¶ms, snapshot_ms).await {
|
||||
Ok(spans) => spans,
|
||||
Err(StoreError::TooLarge) => {
|
||||
let mut resolved = Vec::with_capacity(runs.len());
|
||||
for row in runs {
|
||||
resolved.push(resolve_run(reader, store, access, row).await?);
|
||||
}
|
||||
return Ok(resolved);
|
||||
}
|
||||
Err(error) => return Err(map_store_error(error)),
|
||||
};
|
||||
let Some(spend_rows) = spend(store, access, &spans).await else {
|
||||
// The batch's combined spend read failed; a run's own narrower window may still
|
||||
// resolve, so fall back per run instead of leaving every run in the batch costless.
|
||||
let mut resolved = Vec::with_capacity(runs.len());
|
||||
for row in runs {
|
||||
resolved.push(resolve_run(reader, store, access, row).await?);
|
||||
}
|
||||
return Ok(resolved);
|
||||
};
|
||||
let mut spans = spans;
|
||||
spans.sort_by(|left, right| {
|
||||
run_key(&left.team_id, &left.api_key_hash, &left.trace_id)
|
||||
.cmp(&run_key(
|
||||
&right.team_id,
|
||||
&right.api_key_hash,
|
||||
&right.trace_id,
|
||||
))
|
||||
.then(left.start_ns.cmp(&right.start_ns))
|
||||
});
|
||||
let by_run: HashMap<_, &[TraceSpansRow]> = spans
|
||||
.chunk_by(|left, right| {
|
||||
(&left.team_id, &left.api_key_hash, &left.trace_id)
|
||||
== (&right.team_id, &right.api_key_hash, &right.trace_id)
|
||||
})
|
||||
.map(|run| {
|
||||
(
|
||||
run_key(&run[0].team_id, &run[0].api_key_hash, &run[0].trace_id),
|
||||
run,
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
Ok(runs
|
||||
.iter()
|
||||
.map(|row| {
|
||||
let spans = by_run
|
||||
.get(&run_key(&row.team_id, &row.api_key_hash, &row.trace_id))
|
||||
.copied()
|
||||
.unwrap_or_default();
|
||||
let spend =
|
||||
spend_window(spans).map_or(&[][..], |window| spend_within(&spend_rows, window));
|
||||
resolve_trace(&row.trace_id, &row.trace_ref, spans, spend).map(|trace| {
|
||||
ListedRun::Resolved(
|
||||
Box::new(trace.summary),
|
||||
Freshness::of(spans, true, snapshot_ms),
|
||||
)
|
||||
})
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn resolve_run<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)
|
||||
}
|
||||
365
litellm-rust/crates/traces-cache/src/reader.rs
Normal file
365
litellm-rust/crates/traces-cache/src/reader.rs
Normal 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(¶ms).await {
|
||||
Err(StoreError::TooLarge) if params.limit > 1 => {
|
||||
params.limit /= 2;
|
||||
self.lists.limits.insert(scope.clone(), params.limit).await;
|
||||
}
|
||||
Err(StoreError::TooLarge) => return Err(ReadError::TooLarge),
|
||||
result => break result.map_err(map_store_error)?,
|
||||
}
|
||||
};
|
||||
let next_cursor = page
|
||||
.last()
|
||||
.filter(|_| page.len() == params.limit as usize)
|
||||
.map(|last| encode_cursor(&(last.start_ms, &last.trace_ref)));
|
||||
let data = {
|
||||
let mut summaries = Vec::with_capacity(page.len());
|
||||
for batch in run_batches(&page) {
|
||||
summaries.extend(list_summaries(self, store, access, batch).await?);
|
||||
}
|
||||
summaries
|
||||
};
|
||||
Ok(TracePage { data, next_cursor })
|
||||
}
|
||||
|
||||
pub async fn get_trace<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(¶ms, snapshot_ms)
|
||||
.await
|
||||
.map_err(|error| Miss::Read(map_store_error(error)))?;
|
||||
let spend_rows = spend(store, access, &rows).await;
|
||||
let freshness = Freshness::of(&rows, spend_rows.is_some(), snapshot_ms);
|
||||
resolve_trace(
|
||||
trace_id,
|
||||
trace_ref,
|
||||
&rows,
|
||||
spend_rows.as_deref().unwrap_or_default(),
|
||||
)
|
||||
.map(|trace| (trace, freshness))
|
||||
.ok_or(Miss::Absent)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn get_span<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(¶ms).await.map_err(map_store_error)?;
|
||||
Ok(row.map(|row| SpanDetail {
|
||||
input_ui: to_ui_content(&row.input),
|
||||
output_ui: to_ui_content(&row.output),
|
||||
span_id: row.span_id,
|
||||
input: row.input,
|
||||
output: row.output,
|
||||
attributes: row.attributes,
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn get_span_error<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(¶ms).await.map_err(map_store_error)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let next_offset = offset + row.message.chars().count() as u64;
|
||||
let next_cursor = (next_offset < row.total_chars).then(|| {
|
||||
encode_cursor(&ErrorPosition {
|
||||
offset: next_offset,
|
||||
version: row.version,
|
||||
})
|
||||
});
|
||||
Ok(Some(SpanErrorPage {
|
||||
span_id: row.span_id,
|
||||
message: row.message,
|
||||
total_chars: row.total_chars,
|
||||
next_cursor,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
async fn reference<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(¶ms).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)),
|
||||
}
|
||||
}
|
||||
67
litellm-rust/crates/traces-cache/src/spend.rs
Normal file
67
litellm-rust/crates/traces-cache/src/spend.rs
Normal 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(¶ms).await {
|
||||
Ok(rows) => {
|
||||
let mut rows = rows;
|
||||
rows.sort_by_key(|row| row.start_ms);
|
||||
Some(rows)
|
||||
}
|
||||
Err(error) => {
|
||||
tracing::warn!(%error, "trace spend lookup unavailable");
|
||||
None
|
||||
}
|
||||
}
|
||||
}
|
||||
62
litellm-rust/crates/traces-cache/src/store.rs
Normal file
62
litellm-rust/crates/traces-cache/src/store.rs
Normal 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;
|
||||
}
|
||||
768
litellm-rust/crates/traces-cache/tests/read.rs
Normal file
768
litellm-rust/crates/traces-cache/tests/read.rs
Normal 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(¶ms.trace_ref) {
|
||||
return Err(StoreError::TooLarge);
|
||||
}
|
||||
Ok(state
|
||||
.trace_spans
|
||||
.get(¶ms.trace_ref)
|
||||
.cloned()
|
||||
.unwrap_or_default())
|
||||
}
|
||||
|
||||
async fn run_spans(
|
||||
&self,
|
||||
_: &TracePageSpansParams,
|
||||
_: u64,
|
||||
) -> Result<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);
|
||||
}
|
||||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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, ¶ms).await?;
|
||||
if identities.len() > 1 {
|
||||
return Err(Error::AmbiguousTrace);
|
||||
}
|
||||
Ok(identities.pop().map(|identity| identity.trace_ref))
|
||||
}
|
||||
|
||||
/// Spend records behind the spans' calls. A failed lookup leaves cost unknown instead of failing
|
||||
/// the read.
|
||||
async fn spend(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
access: &ReadAccessParams,
|
||||
rows: &[contracts::TraceSpansRow],
|
||||
) -> Vec<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, ¶ms).await {
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if params.0.limit > 1 => {
|
||||
params.0.limit /= 2;
|
||||
}
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
result => break result?.into_iter().map(|row| row.0).collect(),
|
||||
}
|
||||
};
|
||||
let next_cursor = page
|
||||
.last()
|
||||
.filter(|_| page.len() == params.0.limit as usize)
|
||||
.map(|last| encode_cursor(&(last.start_ms, &last.trace_ref)));
|
||||
let data = stream::iter(page.chunks(16))
|
||||
.then(|batch| list_summaries(client, connection, access, batch))
|
||||
.try_collect::<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, ¶ms)
|
||||
.await?
|
||||
.into_iter()
|
||||
.next();
|
||||
Ok(row.map(|row| SpanDetail {
|
||||
input_ui: to_ui_content(&row.input),
|
||||
output_ui: to_ui_content(&row.output),
|
||||
span_id: row.span_id,
|
||||
input: row.input,
|
||||
output: row.output,
|
||||
attributes: row.attributes,
|
||||
}))
|
||||
}
|
||||
|
||||
pub async fn get_span_error(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
access: &ReadAccessParams,
|
||||
trace_id: &str,
|
||||
span_id: &str,
|
||||
trace_ref: &str,
|
||||
cursor: Option<&str>,
|
||||
) -> Result<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, ¶ms)
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let row = row.0;
|
||||
let next_offset = offset + row.message.chars().count() as u64;
|
||||
let next_cursor = (next_offset < row.total_chars).then(|| {
|
||||
encode_cursor(&ErrorPosition {
|
||||
offset: next_offset,
|
||||
version: row.version,
|
||||
})
|
||||
});
|
||||
Ok(Some(SpanErrorPage {
|
||||
span_id: row.span_id,
|
||||
message: row.message,
|
||||
total_chars: row.total_chars,
|
||||
next_cursor,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
fn trace_cursor_round_trips_the_last_listed_run() {
|
||||
let cursor = encode_cursor(&(1_790_742_989_377_i64, "4bad42b84e9de3ba46fc870185f8f023"));
|
||||
assert_eq!(
|
||||
trace_position(Some(&cursor)).unwrap(),
|
||||
(
|
||||
1_790_742_989_377,
|
||||
"4bad42b84e9de3ba46fc870185f8f023".to_owned()
|
||||
)
|
||||
);
|
||||
assert_eq!(trace_position(None).unwrap(), (0, String::new()));
|
||||
assert_eq!(trace_position(Some("")).unwrap(), (0, String::new()));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::not_base64("abc")]
|
||||
#[case::not_json("bm90LWpzb24=")]
|
||||
#[case::numeric_reference("WzEsIDJd")]
|
||||
#[case::zero_start("WzAsICJ0Il0=")]
|
||||
fn malformed_trace_cursors_are_rejected(#[case] cursor: &str) {
|
||||
assert!(matches!(
|
||||
trace_position(Some(cursor)),
|
||||
Err(Error::InvalidCursor("trace"))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::not_base64("garbage")]
|
||||
#[case::missing_fields("e30=")]
|
||||
#[case::not_an_object("WzEsMl0=")]
|
||||
fn malformed_diagnostic_cursors_are_rejected(#[case] cursor: &str) {
|
||||
assert!(matches!(
|
||||
error_position(Some(cursor)),
|
||||
Err(Error::InvalidCursor("diagnostic"))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::lowercase_version("a".repeat(64))]
|
||||
#[case::short_version("A".repeat(63))]
|
||||
fn diagnostic_cursor_requires_a_content_version(#[case] version: String) {
|
||||
let cursor = encode_cursor(&ErrorPosition { offset: 1, version });
|
||||
assert!(matches!(
|
||||
error_position(Some(&cursor)),
|
||||
Err(Error::InvalidCursor("diagnostic"))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
fn failed(error: StorageError) -> StoreError<Error> {
|
||||
StoreError::Failed(Error::Storage(error))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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, ¶meters).await {
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge)
|
||||
if parameters.page_size > 1 =>
|
||||
{
|
||||
parameters.page_size /= 2;
|
||||
continue;
|
||||
}
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
result => result?,
|
||||
};
|
||||
let complete = page.len() < parameters.page_size as usize;
|
||||
if let Some(last) = page.last() {
|
||||
parameters.after_span_id.clone_from(&last.0.span_id);
|
||||
}
|
||||
for row in page {
|
||||
budget.record(&row)?;
|
||||
spans.push(row.0);
|
||||
}
|
||||
if complete {
|
||||
spans.sort_by_key(|row| row.start_ns);
|
||||
return Ok(spans);
|
||||
}
|
||||
parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE);
|
||||
}
|
||||
let rows = read_paged(client, connection, keyset).await?;
|
||||
Ok(by_start(rows.into_iter().map(|row| row.0).collect()))
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct ListParameters {
|
||||
struct ListSpanKeyset {
|
||||
#[serde(flatten)]
|
||||
runs: crate::query::named::TracePageSpansParams,
|
||||
after_team: String,
|
||||
after_key: String,
|
||||
after_trace: String,
|
||||
after_span: String,
|
||||
page_size: u32,
|
||||
snapshot_ms: u64,
|
||||
}
|
||||
|
||||
struct ListSpanBatch;
|
||||
|
||||
impl Query for ListSpanBatch {
|
||||
type Params = ListParameters;
|
||||
impl Keyset for ListSpanKeyset {
|
||||
type Row = TraceSpansRow;
|
||||
|
||||
const SQL: &'static str = include_str!("../query/trace_list_span_batch.sql");
|
||||
|
||||
fn after(self, last: &TraceSpansRow) -> Self {
|
||||
Self {
|
||||
after_team: last.0.team_id.clone(),
|
||||
after_key: last.0.api_key_hash.clone(),
|
||||
after_trace: last.0.trace_id.clone(),
|
||||
after_span: last.0.span_id.clone(),
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn read_list_spans(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
runs: crate::query::named::TracePageSpansParams,
|
||||
) -> Result<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, ¶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::<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, ¶meters).await {
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge)
|
||||
if parameters.page_size > 1 =>
|
||||
{
|
||||
parameters.page_size /= 2;
|
||||
continue;
|
||||
}
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
result => result?,
|
||||
};
|
||||
let complete = page.len() < parameters.page_size as usize;
|
||||
if let Some(last) = page.last() {
|
||||
parameters.has_cursor = 1;
|
||||
parameters.after_team.clone_from(&last.0.team_id);
|
||||
parameters.after_ms = last.0.start_ms;
|
||||
parameters.after_id.clone_from(&last.0.request_id);
|
||||
}
|
||||
for row in page {
|
||||
budget.record(&row)?;
|
||||
rows.push(row.0);
|
||||
}
|
||||
if complete {
|
||||
return Ok(rows);
|
||||
}
|
||||
parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE);
|
||||
}
|
||||
let rows = read_paged(client, connection, keyset).await?;
|
||||
Ok(rows.into_iter().map(|row| row.0).collect())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::Mutex;
|
||||
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[rstest]
|
||||
#[case::byte_boundary(MAX_GRAPH_BYTES - 1, 0, 1, false)]
|
||||
#[case::byte_overflow(MAX_GRAPH_BYTES - 1, 0, 2, true)]
|
||||
|
|
@ -264,7 +276,77 @@ mod tests {
|
|||
#[case] next: usize,
|
||||
#[case] rejected: bool,
|
||||
) {
|
||||
let budget = ReadBudget { bytes, rows };
|
||||
assert_eq!(budget.checked_add(next, 1).is_err(), rejected);
|
||||
let mut budget = ReadBudget { bytes, rows };
|
||||
assert_eq!(budget.reserve(next).is_err(), rejected);
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct Numbers {
|
||||
after: u32,
|
||||
}
|
||||
|
||||
impl Keyset for Numbers {
|
||||
type Row = u32;
|
||||
const SQL: &'static str = "";
|
||||
|
||||
fn after(self, last: &u32) -> Self {
|
||||
Self { after: *last }
|
||||
}
|
||||
}
|
||||
|
||||
/// A table of `total` rows whose transport rejects any page larger than `largest_page`.
|
||||
struct Table {
|
||||
total: u32,
|
||||
largest_page: u32,
|
||||
requests: Mutex<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]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -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(¤t),
|
||||
1,
|
||||
)
|
||||
.await?
|
||||
.ok_or("missing next page")?;
|
||||
let next = reader
|
||||
.get_trace_page(
|
||||
&store,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
Some(¤t),
|
||||
1,
|
||||
)
|
||||
.await?
|
||||
.ok_or("missing next page")?;
|
||||
assert_eq!(next.summary.span_count, 3);
|
||||
ids.extend(next.spans.into_iter().map(|span| span.span_id));
|
||||
cursor = next.next_cursor;
|
||||
}
|
||||
assert_eq!(ids, original_ids);
|
||||
let refreshed = get_trace(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
)
|
||||
.await?
|
||||
.ok_or("missing refreshed trace")?;
|
||||
let cached = reader
|
||||
.get_trace(&store, &access, &summary.trace_id, &summary.trace_ref)
|
||||
.await?
|
||||
.ok_or("missing cached trace")?;
|
||||
assert_eq!(cached.spans.len(), 3);
|
||||
let (fresh_reader, fresh_store) = make_reader(client, connection);
|
||||
let refreshed = fresh_reader
|
||||
.get_trace(&fresh_store, &access, &summary.trace_id, &summary.trace_ref)
|
||||
.await?
|
||||
.ok_or("missing refreshed trace")?;
|
||||
assert_eq!(refreshed.spans.len(), 4);
|
||||
assert!(matches!(
|
||||
get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
Some("invalid"),
|
||||
1
|
||||
)
|
||||
.await,
|
||||
Err(litellm_traces_clickhouse::Error::InvalidCursor("span"))
|
||||
reader
|
||||
.get_trace_page(
|
||||
&store,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
Some("invalid"),
|
||||
1
|
||||
)
|
||||
.await,
|
||||
Err(ReadError::InvalidCursor("span"))
|
||||
));
|
||||
let backdated = json!({
|
||||
"Timestamp": "2026-09-01 00:00:00.000000000",
|
||||
|
|
@ -509,20 +501,21 @@ async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive(
|
|||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
let uncached_reader =
|
||||
let uncached_connection =
|
||||
Connection::reader(&format!("{}?max_threads=1", fixture.database.url), DATABASE)?;
|
||||
let changed = get_trace_page(
|
||||
client,
|
||||
&uncached_reader,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
first_cursor.as_deref(),
|
||||
1,
|
||||
)
|
||||
.await;
|
||||
let (uncached_reader, uncached_store) = make_reader(client, uncached_connection);
|
||||
let changed = uncached_reader
|
||||
.get_trace_page(
|
||||
&uncached_store,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
first_cursor.as_deref(),
|
||||
1,
|
||||
)
|
||||
.await;
|
||||
assert!(
|
||||
matches!(changed, Err(litellm_traces_clickhouse::Error::TraceChanged)),
|
||||
matches!(changed, Err(ReadError::TraceChanged)),
|
||||
"{changed:?}"
|
||||
);
|
||||
Ok(())
|
||||
|
|
@ -535,16 +528,19 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals(
|
|||
) -> TestResult {
|
||||
let fixture = seeded_database?;
|
||||
let client = &fixture.database.client;
|
||||
let reader = fixture
|
||||
let connection = fixture
|
||||
.readers
|
||||
.connection(client, &QueryScope::All, "fixture-secret")
|
||||
.await?;
|
||||
let (reader, store) = make_reader(client, connection.clone());
|
||||
let access = ReadAccessParams {
|
||||
all_teams: true,
|
||||
user_id: String::new(),
|
||||
team_ids: Vec::new(),
|
||||
};
|
||||
let before = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
|
||||
let before = reader
|
||||
.list_traces(&store, &access, 0, 2_000_000_000_000, None, 50)
|
||||
.await?;
|
||||
let run = before
|
||||
.data
|
||||
.iter()
|
||||
|
|
@ -571,7 +567,14 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals(
|
|||
])],
|
||||
)
|
||||
.await?;
|
||||
let after = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
|
||||
let cached = reader
|
||||
.list_traces(&store, &access, 0, 2_000_000_000_000, None, 50)
|
||||
.await?;
|
||||
assert_eq!(cached.data, before.data);
|
||||
let (reader, store) = make_reader(client, connection);
|
||||
let after = reader
|
||||
.list_traces(&store, &access, 0, 2_000_000_000_000, None, 50)
|
||||
.await?;
|
||||
assert_eq!(after.data.len(), before.data.len());
|
||||
let limited = after
|
||||
.data
|
||||
|
|
@ -588,17 +591,10 @@ async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals(
|
|||
.all(|item| !item.resolution_limited)
|
||||
);
|
||||
assert!(matches!(
|
||||
get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&run.trace_id,
|
||||
&run.trace_ref,
|
||||
None,
|
||||
200
|
||||
)
|
||||
.await,
|
||||
Err(litellm_traces_clickhouse::Error::ReadTooLarge)
|
||||
reader
|
||||
.get_trace_page(&store, &access, &run.trace_id, &run.trace_ref, None, 200)
|
||||
.await,
|
||||
Err(ReadError::TooLarge)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -685,16 +681,19 @@ async fn gateway_ids_resolve_through_detail_and_batch_reads_with_legacy_fallback
|
|||
.collect(),
|
||||
)
|
||||
.await?;
|
||||
let reader = fixture
|
||||
let connection = fixture
|
||||
.readers
|
||||
.connection(client, &QueryScope::All, "fixture-secret")
|
||||
.await?;
|
||||
let (reader, store) = make_reader(client, connection);
|
||||
let access = ReadAccessParams {
|
||||
all_teams: false,
|
||||
user_id: String::new(),
|
||||
team_ids: vec!["team-a".into()],
|
||||
};
|
||||
let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
|
||||
let page = reader
|
||||
.list_traces(&store, &access, 0, 2_000_000_000_000, None, 50)
|
||||
.await?;
|
||||
assert_eq!(page.data.len(), cases.len());
|
||||
for (id, _, _, _, _, expected) in cases {
|
||||
let summary = page
|
||||
|
|
@ -702,7 +701,8 @@ async fn gateway_ids_resolve_through_detail_and_batch_reads_with_legacy_fallback
|
|||
.iter()
|
||||
.find(|summary| summary.trace_id == id)
|
||||
.ok_or("missing run")?;
|
||||
let detail = get_trace(client, &reader, &access, id, &summary.trace_ref)
|
||||
let detail = reader
|
||||
.get_trace(&store, &access, id, &summary.trace_ref)
|
||||
.await?
|
||||
.ok_or("missing trace")?;
|
||||
assert_eq!(detail.summary.spend, expected, "{id}");
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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", {})
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
150
litellm/proxy/common_utils/fips.py
Normal file
150
litellm/proxy/common_utils/fips.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
5
litellm/rust_bridge/trace/errors.py
Normal file
5
litellm/rust_bridge/trace/errors.py
Normal 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.
|
||||
"""
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
427
scripts/seed_request_logs.py
Normal file
427
scripts/seed_request_logs.py
Normal 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)))
|
||||
|
|
@ -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)))
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
62
tests/integration/configuration/test_fips_mode_boot.py
Normal file
62
tests/integration/configuration/test_fips_mode_boot.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
)
|
||||
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
106
tests/unit/proxy/common_utils/test_fips.py
Normal file
106
tests/unit/proxy/common_utils/test_fips.py
Normal 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()
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
20
ui/litellm-dashboard/.agents/skills/url-state/SKILL.md
Normal file
20
ui/litellm-dashboard/.agents/skills/url-state/SKILL.md
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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" },
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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*" },
|
||||
|
|
|
|||
11
ui/litellm-dashboard/package-lock.json
generated
11
ui/litellm-dashboard/package-lock.json
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
|
||||
|
|
|
|||
|
|
@ -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]));
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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" />;
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
||||
|
|
|
|||
|
|
@ -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([]);
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ describe("makeSystemOneRequest", () => {
|
|||
const expectedRequest: Partial<RequestInit> = {
|
||||
method: "POST",
|
||||
headers: {
|
||||
Accept: "application/json",
|
||||
"Content-Type": "application/json",
|
||||
Authorization: "Bearer session-key",
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue