Merge remote-tracking branch 'origin/main' into litellm_otel_v2_team_capture_message_content
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

This commit is contained in:
mrinal 2026-10-03 19:19:13 +00:00
commit c3810ea755
149 changed files with 9770 additions and 492 deletions

View file

@ -61,7 +61,7 @@ jobs:
- shard: core-utils
artifact-name: core-utils
test-path: ""
test-path: tests/unit/decisions
unit-flag: core-utils
workers: 2
reruns: 1
@ -141,6 +141,7 @@ jobs:
artifact-name: proxy-endpoints
test-path: >-
tests/unit/proxy/analytics_endpoints
tests/unit/proxy/decisions_endpoints
tests/unit/proxy/management_endpoints
tests/unit/proxy/list_api
tests/unit/proxy/memory

3
.gitignore vendored
View file

@ -151,3 +151,6 @@ litellm.log
.coverage-rust
coverage-rust.xml
# make lens-dev worker token, generated config and logs
.lens-dev/

View file

@ -4,7 +4,7 @@
.PHONY: help test test-unit test-unit-llms test-unit-proxy-guardrails test-unit-proxy-core test-unit-proxy-misc test-unit-proxy-root \
test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \
test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \
test-rust-extension rust-sqlx-prepare \
test-rust-extension rust-sqlx-prepare lens-dev \
info lint lint-inner lint-dev lint-checks format \
lint-basedpyright lint-e2e-basedpyright lint-basedpyright-budget-update lint-type-discipline lint-type-discipline-budget-update \
lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-gate \
@ -58,6 +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 (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."
@ -311,6 +312,9 @@ test-rust-extension:
rust-sqlx-prepare:
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
lens-dev:
./scripts/lens_dev.sh
test: install-test-deps
$(UV_RUN) pytest tests/

View file

@ -390,11 +390,13 @@ Set `LITELLM_PROXY_API_BASE` and `LITELLM_PROXY_API_KEY` and every model call th
| [Sail (`sail`)](https://docs.litellm.ai/docs/providers/sail) | ✅ | ✅ | ✅ | | | | | | | |
| [Sambanova (`sambanova`)](https://docs.litellm.ai/docs/providers/sambanova) | ✅ | ✅ | ✅ | | | | | | | |
| [Snowflake (`snowflake`)](https://docs.litellm.ai/docs/providers/snowflake) | ✅ | ✅ | ✅ | | | | | | | |
| [Strands Decider (`strands_decider`)](https://docs.litellm.ai/docs/providers) | | | | | | | | | | |
| [Text Completion Codestral (`text-completion-codestral`)](https://docs.litellm.ai/docs/providers/codestral) | ✅ | ✅ | ✅ | | | | | | | |
| [Text Completion OpenAI (`text-completion-openai`)](https://docs.litellm.ai/docs/providers/text_completion_openai) | ✅ | ✅ | ✅ | | | ✅ | ✅ | ✅ | ✅ | |
| [Together AI (`together_ai`)](https://docs.litellm.ai/docs/providers/togetherai) | ✅ | ✅ | ✅ | | | | | | | |
| [Topaz (`topaz`)](https://docs.litellm.ai/docs/providers/topaz) | ✅ | ✅ | ✅ | | | | | | | |
| [Triton (`triton`)](https://docs.litellm.ai/docs/providers/triton-inference-server) | ✅ | ✅ | ✅ | | | | | | | |
| [Typesafe Decisions API (`typesafe`)](https://docs.litellm.ai/docs/providers) | | | | | | | | | | |
| [V0 (`v0`)](https://docs.litellm.ai/docs/providers/v0) | ✅ | ✅ | ✅ | | | | | | | |
| [Vercel AI Gateway (`vercel_ai_gateway`)](https://docs.litellm.ai/docs/providers/vercel_ai_gateway) | ✅ | ✅ | ✅ | | | | | | | |
| [VLLM (`vllm`)](https://docs.litellm.ai/docs/providers/vllm) | ✅ | ✅ | ✅ | | | | | | | |

View file

@ -33,7 +33,7 @@ For deployments managed with Compose, download `compose.yaml` and provide `LITEL
docker compose --env-file /path/to/lens.env -f compose.yaml up -d
```
Developers can build locally with `LENS_WORKER_IMAGE=litellm-lens-worker:local docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build`
Developers can build locally with `LENS_WORKER_IMAGE=litellm-lens-worker:local docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build`. To work on Lens itself, `make lens-dev` runs the proxy, a worker from source and the hot-reload dashboard together; set `LENS_DEV_PROXY_PORT` / `LENS_DEV_UI_PORT` to move them off 4000/3000
The generated command gives the worker 1 GiB of temporary memory-backed storage, shared across parallel reviews. Change `size=1g` in the Docker command or set `LENS_WORKER_TMP_SIZE` with Compose to fit your server and workload. A storage failure marks the scan as failed, cleans up temporary traces, and leaves the worker available for other scans; it does not silently truncate the review. Existing workers must be recreated with the new image and mount options

View file

@ -1,7 +1,7 @@
"""Path allowlist for the gateway component.
The gateway exposes the LLM data-plane surface: chat/completions, embeddings,
audio, batches, files, fine-tuning, rerank, ocr, rag, video, search, image,
audio, batches, files, fine-tuning, rerank, decisions, ocr, rag, video, search, image,
responses, vector stores, passthrough providers, realtime websockets, MCP
tool-call endpoints, and operational endpoints (/health, /metrics, and the
/debug/memory/summary read of the serving worker's RSS).
@ -60,6 +60,8 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/v1/rerank",
"/v2/rerank",
"/rerank",
"/v1/decisions",
"/decisions",
"/v1/ocr",
"/ocr",
"/v1/rag/",

View file

@ -4478,6 +4478,7 @@ dependencies = [
"flate2",
"futures-util",
"hmac 0.12.1",
"itertools 0.14.0",
"jsonschema",
"litellm-http",
"litellm-migrate",

View file

@ -35,13 +35,14 @@ fn map_error_ref(error: &Error) -> PyErr {
use litellm_storage_clickhouse::Error as StorageError;
match error {
Error::Decode(litellm_traces::Error::TooLarge) | Error::InsertTooLarge => {
PyOverflowError::new_err(error.to_string())
}
Error::Decode(litellm_traces::Error::TooLarge)
| Error::InsertTooLarge
| Error::ReadTooLarge => PyOverflowError::new_err(error.to_string()),
Error::InvalidRow
| Error::InvalidTable
| Error::InvalidCursor(_)
| Error::AmbiguousTrace
| Error::TraceChanged
| Error::Decode(_)
| Error::InvalidSchema
| Error::InvalidQuery
@ -233,26 +234,44 @@ impl NativeTraceStorage {
)
}
#[pyo3(signature = (trace_id, scope, trace_ref, cursor=None, page_size=None))]
fn get_trace<'py>(
&self,
py: Python<'py>,
trace_id: String,
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] scope: ReadAccessParams,
trace_ref: String,
cursor: Option<String>,
page_size: Option<u32>,
) -> PyResult<Bound<'py, PyAny>> {
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
let connection = self.config.storage().reader().clone();
crate::execution::run_async(
py,
async move {
litellm_traces_clickhouse::get_trace(
&client,
&connection,
&scope,
&trace_id,
&trace_ref,
)
.await
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
} else if cursor.is_some() {
Err(Error::InvalidParameters)
} else {
litellm_traces_clickhouse::get_trace(
&client,
&connection,
&scope,
&trace_id,
&trace_ref,
)
.await
}
},
map_error,
)
@ -465,6 +484,8 @@ mod tests {
#[case::invalid_export(Error::Decode(litellm_traces::Error::InvalidPayload), "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(
#[case] error: Error,
#[case] exception_name: &str,

View file

@ -110,6 +110,13 @@ pub async fn execute_read(
.body(sql.to_owned());
let mut response = request.send().await.map_err(|_| Error::Transport)?;
if !response.status().is_success() {
if response
.headers()
.get("x-clickhouse-exception-code")
.is_some_and(|code| code == "396")
{
return Err(Error::ResponseTooLarge);
}
return Err(Error::QueryFailed(response.status().as_u16()));
}

View file

@ -111,3 +111,36 @@ async fn typed_fetch_encodes_parameters_and_validates_rows(
assert!(matches!(envelope, Err(Error::InvalidResponse)));
}
}
#[rstest]
#[case::result_limit("396", true)]
#[case::memory_limit("241", false)]
#[case::timeout("159", false)]
#[case::unknown("", false)]
#[tokio::test]
async fn server_result_limits_allow_smaller_pages_without_retrying_other_failures(
#[case] code: &str,
#[case] result_limit: bool,
) {
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(500).insert_header("X-ClickHouse-Exception-Code", code))
.expect(1)
.mount(&server)
.await;
let connection = Connection::parse(&server.uri()).unwrap();
let error = execute_read(
&Client::no_redirect_for_test(),
&connection,
"SELECT 1",
&BTreeMap::new(),
)
.await
.unwrap_err();
if result_limit {
assert!(matches!(error, Error::ResponseTooLarge));
} else {
assert!(matches!(error, Error::QueryFailed(500)));
}
}

View file

@ -16,6 +16,7 @@ base64.workspace = true
flate2.workspace = true
futures-util.workspace = true
hmac = "0.12.1"
itertools = "0.14.0"
litellm-http.workspace = true
litellm-migrate.workspace = true
litellm-storage-clickhouse.workspace = true

View file

@ -0,0 +1,27 @@
SELECT * FROM (
SELECT request_id, response_id, upstream_response_id, trace_id, span_id, team_id, api_key, user, spend,
toUnixTimestamp64Milli(start_time) AS start_ms
FROM (
SELECT *,
-- A chat request served through the Responses API returns the upstream `resp_` id to the
-- client but logs LiteLLM's managed `resp_<base64>` id, which embeds it.
if(startsWith(response_id, 'resp_'),
extract(tryBase64Decode(substring(response_id, 6)), 'response_id:([^;]+)'),
'') AS upstream_response_id
FROM spend_logs FINAL
WHERE start_time >= fromUnixTimestamp64Milli({start_ms:Int64})
AND start_time < fromUnixTimestamp64Milli({end_ms:Int64})
AND ({all_teams:UInt8} = 1
OR ({user_id:String} != '' AND user = {user_id:String})
OR has({team_ids:Array(String)}, team_id))
)
WHERE response_id IN {response_ids:Array(String)}
OR upstream_response_id IN {response_ids:Array(String)}
OR request_id IN {request_ids:Array(String)}
OR (trace_id != '' AND trace_id IN {trace_ids:Array(String)})
ORDER BY start_time DESC
)
WHERE {has_cursor:UInt8} = 0
OR (team_id, start_ms, request_id) > ({after_team:String}, {after_ms:Int64}, {after_id:String})
ORDER BY team_id, start_ms, request_id
LIMIT {page_size:UInt32}

View file

@ -0,0 +1,31 @@
SELECT * FROM (
SELECT o.TraceId AS trace_id, o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name,
o.ObservationType AS type, toUInt8(o.WrapperCandidate) AS wrapper_candidate, o.AgentName AS agent,
o.Framework AS framework, o.StatusCode AS status,
substringUTF8(o.StatusMessage, 1, 128) AS status_message,
lengthUTF8(o.StatusMessage) > 128 AS error_truncated,
toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns,
o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model,
o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens,
o.LiteLLMRequestId AS litellm_request_id,
o.CallKeys AS call_keys, o.CallEvidence AS call_evidence,
-- Rows written before ToolCallId keep the call id only in their attributes.
if(o.ToolCallId != '' OR o.ObservationType != 'tool', o.ToolCallId,
coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), ''))
AS tool_call_id,
o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash
FROM otel_traces AS o
WHERE o.Timestamp >= fromUnixTimestamp64Milli({start_ms:Int64})
AND o.Timestamp < fromUnixTimestamp64Milli({end_ms:Int64})
AND ({all_teams:UInt8} = 1
OR ({user_id:String} != '' AND o.UserId = {user_id:String})
OR has({team_ids:Array(String)}, o.TeamId))
AND hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) IN {trace_refs:Array(String)}
AND o.EngineReceivedMs <= {snapshot_ms:UInt64}
ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage
LIMIT 1 BY o.TeamId, o.ApiKeyHash, o.TraceId, o.SpanId
)
WHERE (team_id, api_key_hash, trace_id, span_id) > ({after_team:String}, {after_key:String}, {after_trace:String}, {after_span:String})
ORDER BY team_id, api_key_hash, trace_id, span_id
LIMIT {page_size:UInt32}

View file

@ -0,0 +1,30 @@
SELECT * FROM (
SELECT o.TraceId AS trace_id, o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name,
o.ObservationType AS type, toUInt8(o.WrapperCandidate) AS wrapper_candidate, o.AgentName AS agent,
o.Framework AS framework, o.StatusCode AS status,
substringUTF8(o.StatusMessage, 1, 128) AS status_message,
lengthUTF8(o.StatusMessage) > 128 AS error_truncated,
toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns,
o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model,
o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens,
o.LiteLLMRequestId AS litellm_request_id,
o.CallKeys AS call_keys, o.CallEvidence AS call_evidence,
-- Rows written before ToolCallId keep the call id only in their attributes.
if(o.ToolCallId != '' OR o.ObservationType != 'tool', o.ToolCallId,
coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), ''))
AS tool_call_id,
o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash
FROM otel_traces AS o
WHERE o.TraceId = {trace_id:String}
AND ({all_teams:UInt8} = 1
OR ({user_id:String} != '' AND o.UserId = {user_id:String})
OR has({team_ids:Array(String)}, o.TeamId))
AND ({trace_ref:String} = '' OR
hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String})
AND o.EngineReceivedMs <= {snapshot_ms:UInt64}
ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage
LIMIT 1 BY o.SpanId
)
WHERE span_id > {after_span_id:String}
ORDER BY span_id
LIMIT {page_size:UInt32}

View file

@ -14,6 +14,8 @@ 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,6 +36,8 @@ pub enum Error {
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")]

View file

@ -17,6 +17,7 @@ pub mod query;
mod query_access;
mod reads;
mod schema;
mod span_batches;
mod span_row;
mod sql;
mod table;
@ -30,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, list_traces};
pub use reads::{get_span, get_span_error, get_trace, get_trace_page, list_traces};
pub use schema::{
NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, ensure_schema, schema_statements,
};

View file

@ -1,28 +1,59 @@
//! Scoped trace reads: the trace list, one trace resolved with its spend, and span payloads.
use std::collections::HashMap;
use std::sync::{Arc, LazyLock};
use std::time::Duration;
use base64::{Engine, engine::general_purpose::URL_SAFE};
use futures_util::{StreamExt, TryStreamExt, stream};
use itertools::Itertools;
use litellm_http::Client;
use litellm_storage_clickhouse::fetch;
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 moka::future::Cache;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use crate::{
Connection, Error,
query::named::{
ListTraces, ListTracesParams, ReadAccessParams, SpanDetail as SpanDetailQuery,
SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIds, SpendByResponseIdsParams,
TraceIdentity, TraceIdentityParams, TracePageSpans, TracePageSpansParams, TraceSpans,
TraceSpansParams,
ListTracesParams, ListTracesRow, ReadAccessParams, SpanDetail as SpanDetailQuery,
SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIdsParams, TraceIdentity,
TraceIdentityParams, TracePageSpansParams, TraceSpansParams,
},
};
struct RunCandidates;
impl Query for RunCandidates {
type Params = ListTracesParams;
type Row = ListTracesRow;
const SQL: &'static str = concat!(
"SELECT * EXCEPT (request_ids), [] AS request_ids FROM (",
include_str!("../query/list_traces.sql"),
") ORDER BY start_ms DESC, trace_ref DESC"
);
}
// Cursor pages share a bounded snapshot so advancing does not resolve the whole graph again.
static TRACE_SNAPSHOTS: LazyLock<Cache<String, Arc<Trace>>> = LazyLock::new(|| {
Cache::builder()
.max_capacity((2 * crate::span_batches::MAX_GRAPH_BYTES) as u64)
.weigher(|_: &String, trace: &Arc<Trace>| {
serde_json::to_vec(trace.as_ref())
.ok()
.and_then(|bytes| u32::try_from(bytes.len().saturating_mul(2)).ok())
.unwrap_or(u32::MAX)
})
.time_to_live(Duration::from_secs(120))
.build()
});
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())
@ -121,8 +152,8 @@ async fn spend(
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 fetch::<SpendByResponseIds>(client, connection, &params).await {
Ok(rows) => rows.into_iter().map(|row| row.0).collect(),
match crate::span_batches::read_spend(client, connection, params).await {
Ok(rows) => rows,
Err(error) => {
tracing::warn!(%error, "trace spend lookup unavailable");
Vec::new()
@ -139,57 +170,90 @@ pub async fn list_traces(
cursor: Option<&str>,
limit: u32,
) -> Result<TracePage, Error> {
if limit == 0 {
return Err(Error::InvalidParameters);
}
let (cursor_ms, cursor_trace_id) = trace_position(cursor)?;
let params = ListTracesParams::from(contracts::ListTracesParams {
let mut params = ListTracesParams::from(contracts::ListTracesParams {
access: access.clone(),
start_ms,
end_ms,
cursor_ms,
cursor_trace_id,
limit,
limit: limit.min(500),
});
let page: Vec<contracts::ListTracesRow> = fetch::<ListTraces>(client, connection, &params)
.await?
.into_iter()
.map(|row| row.0)
.collect();
let page: Vec<contracts::ListTracesRow> = loop {
match fetch::<RunCandidates>(client, connection, &params).await {
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if params.0.limit > 1 => {
params.0.limit /= 2;
}
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
return Err(Error::ReadTooLarge);
}
result => break result?.into_iter().map(|row| row.0).collect(),
}
};
let next_cursor = page
.last()
.filter(|_| page.len() == limit as usize)
.filter(|_| page.len() == params.0.limit as usize)
.map(|last| encode_cursor(&(last.start_ms, &last.trace_ref)));
let (Some(page_start), Some(page_end)) = (
page.iter().map(|row| row.start_ms).min(),
page.iter().map(|row| row.start_ms + row.duration_ms).max(),
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(TracePage {
data: Vec::new(),
next_cursor,
});
return Ok(Vec::new());
};
let span_params = TracePageSpansParams::from(contracts::TracePageSpansParams {
let params = TracePageSpansParams::from(contracts::TracePageSpansParams {
access: access.clone(),
trace_refs: page.iter().map(|row| row.trace_ref.clone()).collect(),
start_ms: page_start,
end_ms: page_end + 1,
trace_refs: runs.iter().map(|row| row.trace_ref.clone()).collect(),
start_ms,
end_ms: end_ms.saturating_add(1),
});
let span_rows: Vec<contracts::TraceSpansRow> =
fetch::<TracePageSpans>(client, connection, &span_params)
.await?
.into_iter()
.map(|row| row.0)
.collect();
let spend_rows = spend(client, connection, access, &span_rows).await;
let mut by_trace: HashMap<(String, String, String), Vec<contracts::TraceSpansRow>> =
HashMap::new();
for span in span_rows {
let key = (
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(),
);
by_trace.entry(key).or_default().push(span);
}
let data = page
)
});
let summaries = runs
.iter()
.map(|row| {
let spans = by_trace
@ -200,11 +264,17 @@ pub async fn list_traces(
))
.map(Vec::as_slice)
.unwrap_or_default();
resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows)
.map_or_else(|| listed_summary(row), |trace| trace.summary)
async move {
let spend_rows = spend(client, connection, access, spans).await;
resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows)
.map_or_else(|| listed_summary(row), |trace| trace.summary)
}
})
.collect();
Ok(TracePage { data, next_cursor })
.collect::<Vec<_>>();
Ok(stream::iter(summaries)
.buffered(SPEND_CONCURRENCY)
.collect()
.await)
}
pub async fn get_trace(
@ -222,11 +292,7 @@ pub async fn get_trace(
trace_id: trace_id.to_owned(),
trace_ref: trace_ref.clone(),
};
let rows: Vec<contracts::TraceSpansRow> = fetch::<TraceSpans>(client, connection, &params)
.await?
.into_iter()
.map(|row| row.0)
.collect();
let rows = crate::span_batches::read_spans(client, connection, params, u64::MAX).await?;
if rows.is_empty() {
return Ok(None);
}
@ -234,6 +300,132 @@ pub async fn get_trace(
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);
}
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
}
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_bytes = serde_json::to_vec(&(
connection.url().as_str(),
access,
trace_id,
&trace_ref,
position.snapshot_ms,
))
.map_err(|_| Error::InvalidParameters)?;
let key = format!("{:x}", Sha256::digest(key_bytes));
let snapshot = if let Some(trace) = TRACE_SNAPSHOTS.get(&key).await {
trace
} else {
let params = TraceSpansParams {
access: access.clone(),
trace_id: trace_id.to_owned(),
trace_ref: trace_ref.clone(),
};
let rows =
crate::span_batches::read_spans(client, connection, params, position.snapshot_ms)
.await?;
let spend_rows = spend(client, connection, access, &rows).await;
let Some(trace) = resolve_trace(trace_id, &trace_ref, &rows, &spend_rows) else {
return Ok(None);
};
if serde_json::to_vec(&trace)
.map_err(|_| Error::InvalidResponse)?
.len()
> crate::span_batches::MAX_GRAPH_BYTES
{
return Err(Error::ReadTooLarge);
}
let trace = Arc::new(trace);
TRACE_SNAPSHOTS.insert(key, Arc::clone(&trace)).await;
trace
};
let span_ids: Vec<&str> = snapshot
.spans
.iter()
.map(|span| span.span_id.as_str())
.collect();
let version = format!(
"{:x}",
Sha256::digest(serde_json::to_vec(&span_ids).map_err(|_| Error::InvalidResponse)?)
);
if cursor.is_some() && position.version != version {
return Err(Error::TraceChanged);
}
let mut trace = Trace {
summary: snapshot.summary.clone(),
agents: snapshot.agents.clone(),
spans: Vec::new(),
next_cursor: None,
};
if position.offset > snapshot.spans.len() {
return Err(Error::InvalidCursor("span"));
}
let end = position
.offset
.saturating_add(page_size as usize)
.min(snapshot.spans.len());
trace.next_cursor = (end < snapshot.spans.len()).then(|| {
encode_cursor(&SpanPosition {
offset: end,
version: version.clone(),
..position
})
});
trace.spans = snapshot.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: version.clone(),
}));
}
Ok(Some(trace))
}
pub async fn get_span(
client: &Client,
connection: &Connection,

View file

@ -0,0 +1,270 @@
use futures_util::{TryStreamExt, stream};
use itertools::Itertools;
use litellm_http::Client;
use litellm_storage_clickhouse::{Query, fetch};
use litellm_traces::query::named as contracts;
use serde::Serialize;
use crate::{Connection, Error, query::named::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 {
bytes: usize,
rows: usize,
}
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);
}
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)?;
Ok(())
}
}
#[derive(Serialize)]
struct Parameters {
#[serde(flatten)]
trace: contracts::TraceSpansParams,
after_span_id: String,
page_size: u32,
snapshot_ms: u64,
}
struct SpanBatch;
impl Query for SpanBatch {
type Params = Parameters;
type Row = TraceSpansRow;
const SQL: &'static str = include_str!("../query/trace_span_batch.sql");
}
pub(crate) async fn read_spans(
client: &Client,
connection: &Connection,
trace: contracts::TraceSpansParams,
snapshot_ms: u64,
) -> Result<Vec<contracts::TraceSpansRow>, Error> {
let mut parameters = Parameters {
trace,
after_span_id: String::new(),
page_size: PAGE_SIZE,
snapshot_ms,
};
let mut spans = Vec::new();
let mut budget = ReadBudget::default();
loop {
let page = match fetch::<SpanBatch>(client, connection, &parameters).await {
Err(litellm_storage_clickhouse::Error::ResponseTooLarge)
if parameters.page_size > 1 =>
{
parameters.page_size /= 2;
continue;
}
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
return Err(Error::ReadTooLarge);
}
result => result?,
};
let complete = page.len() < parameters.page_size as usize;
if let Some(last) = page.last() {
parameters.after_span_id.clone_from(&last.0.span_id);
}
for row in page {
budget.record(&row)?;
spans.push(row.0);
}
if complete {
spans.sort_by_key(|row| row.start_ns);
return Ok(spans);
}
parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE);
}
}
#[derive(Serialize)]
struct ListParameters {
#[serde(flatten)]
runs: crate::query::named::TracePageSpansParams,
after_team: String,
after_key: String,
after_trace: String,
after_span: String,
page_size: u32,
snapshot_ms: u64,
}
struct ListSpanBatch;
impl Query for ListSpanBatch {
type Params = ListParameters;
type Row = TraceSpansRow;
const SQL: &'static str = include_str!("../query/trace_list_span_batch.sql");
}
pub(crate) async fn read_list_spans(
client: &Client,
connection: &Connection,
runs: crate::query::named::TracePageSpansParams,
) -> Result<Vec<contracts::TraceSpansRow>, Error> {
let parameters = ListParameters {
runs,
after_team: String::new(),
after_key: String::new(),
after_trace: String::new(),
after_span: String::new(),
page_size: PAGE_SIZE,
snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64,
};
let pages = stream::try_unfold(
(Some(parameters), ReadBudget::default()),
|(parameters, budget)| async move {
let Some(parameters) = parameters else {
return Ok(None);
};
let page = match fetch::<ListSpanBatch>(client, connection, &parameters).await {
Err(litellm_storage_clickhouse::Error::ResponseTooLarge)
if parameters.page_size > 1 =>
{
let retry = ListParameters {
page_size: parameters.page_size / 2,
..parameters
};
return Ok(Some((Vec::new(), (Some(retry), budget))));
}
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
return Err(Error::ReadTooLarge);
}
result => result?,
};
let next = page
.last()
.filter(|_| page.len() == parameters.page_size as usize)
.map(|last| ListParameters {
after_team: last.0.team_id.clone(),
after_key: last.0.api_key_hash.clone(),
after_trace: last.0.trace_id.clone(),
after_span: last.0.span_id.clone(),
page_size: (parameters.page_size * 2).min(PAGE_SIZE),
..parameters
});
let next_budget = page.iter().try_fold(budget, |budget, row| {
let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?;
budget.checked_add(bytes.len(), 1)
})?;
Ok(Some((page, (next, next_budget))))
},
)
.try_collect::<Vec<_>>()
.await?;
Ok(pages
.into_iter()
.flatten()
.map(|row| row.0)
.sorted_by_key(|row| row.start_ns)
.collect())
}
#[derive(Serialize)]
struct SpendParameters {
#[serde(flatten)]
lookup: crate::query::named::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;
const SQL: &'static str = include_str!("../query/spend_batch.sql");
}
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,
has_cursor: 0,
after_team: String::new(),
after_ms: 0,
after_id: String::new(),
page_size: PAGE_SIZE,
};
let mut rows = Vec::new();
let mut budget = ReadBudget::default();
loop {
let page = match fetch::<SpendBatch>(client, connection, &parameters).await {
Err(litellm_storage_clickhouse::Error::ResponseTooLarge)
if parameters.page_size > 1 =>
{
parameters.page_size /= 2;
continue;
}
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
return Err(Error::ReadTooLarge);
}
result => result?,
};
let complete = page.len() < parameters.page_size as usize;
if let Some(last) = page.last() {
parameters.has_cursor = 1;
parameters.after_team.clone_from(&last.0.team_id);
parameters.after_ms = last.0.start_ms;
parameters.after_id.clone_from(&last.0.request_id);
}
for row in page {
budget.record(&row)?;
rows.push(row.0);
}
if complete {
return Ok(rows);
}
parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE);
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[rstest]
#[case::byte_boundary(MAX_GRAPH_BYTES - 1, 0, 1, false)]
#[case::byte_overflow(MAX_GRAPH_BYTES - 1, 0, 2, true)]
#[case::integer_overflow(MAX_GRAPH_BYTES, 0, usize::MAX, true)]
#[case::row_boundary(0, MAX_GRAPH_SPANS - 1, 1, false)]
#[case::row_overflow(0, MAX_GRAPH_SPANS, 1, true)]
fn accumulation_stops_at_the_graph_budget(
#[case] bytes: usize,
#[case] rows: usize,
#[case] next: usize,
#[case] rejected: bool,
) {
let budget = ReadBudget { bytes, rows };
assert_eq!(budget.checked_add(next, 1).is_err(), rejected);
}
}

View file

@ -174,7 +174,7 @@ async fn admin_sql_enforces_result_row_limit(
matches!(
result,
Err(Error::Storage(
litellm_storage_clickhouse::Error::QueryFailed(_)
litellm_storage_clickhouse::Error::ResponseTooLarge
))
),
"{result:?}"

View file

@ -0,0 +1,595 @@
use std::collections::BTreeMap;
use litellm_traces::query::named::ReadAccessParams;
use litellm_traces_clickhouse::{
Connection, InsertTable, QueryScope, get_trace, get_trace_page, insert_rows, list_traces,
};
use rstest::rstest;
use serde_json::json;
#[path = "queries/support.rs"]
mod fixtures;
mod support;
use fixtures::{DATABASE, SeededDatabase, migrated_database, seeded_database};
use support::TestResult;
#[rstest]
#[case::api_key("key-a", "")]
#[case::user("", "user-a")]
#[tokio::test]
async fn list_costs_match_each_run_when_response_ids_are_reused(
#[future(awt)] migrated_database: TestResult<SeededDatabase>,
#[case] api_key: &str,
#[case] user_id: &str,
) -> TestResult {
let fixture = migrated_database?;
let client = &fixture.database.client;
let writer = Connection::writer(&fixture.database.url)?;
let runs = [
("earlier-run", 1_790_000_000_000_i64, 0.25),
("later-run", 1_790_007_200_000_i64, 0.75),
];
insert_rows(
client,
&writer,
DATABASE,
InsertTable::OtelTraces,
runs.iter()
.map(|(trace_id, start_ms, _)| {
BTreeMap::from([
("Timestamp".into(), json!(start_ms * 1_000_000)),
("TraceId".into(), json!(trace_id)),
("SpanId".into(), json!("llm-span")),
("ObservationType".into(), json!("llm")),
("TeamId".into(), json!("team-a")),
("ApiKeyHash".into(), json!(api_key)),
("UserId".into(), json!(user_id)),
("Duration".into(), json!(1_000_000)),
("LiteLLMRequestId".into(), json!("reused-response")),
])
})
.collect(),
)
.await?;
insert_rows(
client,
&writer,
DATABASE,
InsertTable::SpendLogs,
runs.iter()
.map(|(trace_id, start_ms, cost)| {
BTreeMap::from([
("request_id".into(), json!(format!("request-{trace_id}"))),
("response_id".into(), json!("reused-response")),
("team_id".into(), json!("team-a")),
("api_key".into(), json!(api_key)),
("user".into(), json!(user_id)),
("start_time".into(), json!(start_ms)),
("end_time".into(), json!(start_ms + 1)),
("spend".into(), json!(cost)),
])
})
.collect(),
)
.await?;
let reader = fixture
.readers
.connection(client, &QueryScope::All, "fixture-secret")
.await?;
let access = ReadAccessParams {
all_teams: false,
user_id: user_id.into(),
team_ids: vec!["team-a".into()],
};
let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
assert_eq!(page.data.len(), runs.len());
for (trace_id, _, cost) in runs {
let summary = page
.data
.iter()
.find(|summary| summary.trace_id == trace_id)
.ok_or("missing run")?;
let detail = get_trace(client, &reader, &access, trace_id, &summary.trace_ref)
.await?
.ok_or("missing trace")?;
assert_eq!(detail.summary.spend, Some(cost));
assert_eq!(summary.spend, detail.summary.spend, "{trace_id}");
}
Ok(())
}
#[rstest]
#[case::many_runs(50, 21, 0, false)]
#[case::one_large_run(1, 1100, 0, false)]
#[case::large_rows(1, 280, 20_000, false)]
#[case::large_cached_snapshot(1, 280, 140_000, false)]
#[case::many_costs(1, 1101, 0, true)]
#[case::many_costed_runs(500, 2, 0, true)]
#[tokio::test]
async fn large_runs_remain_complete_under_default_reader_limits(
#[future(awt)] migrated_database: TestResult<SeededDatabase>,
#[case] runs: usize,
#[case] steps: usize,
#[case] name_bytes: usize,
#[case] costed: bool,
) -> TestResult {
let fixture = migrated_database?;
let client = &fixture.database.client;
let writer = Connection::writer(&fixture.database.url)?;
let rows = (0..runs)
.flat_map(|run| {
(0..steps).map(move |step| {
BTreeMap::from([
(
"Timestamp".into(),
json!(1_790_000_000_000_000_000_i64 + step as i64),
),
("TraceId".into(), json!(format!("trace-{run:04}"))),
("SpanId".into(), json!(format!("span-{step:04}"))),
(
"ParentSpanId".into(),
json!(if step == 0 { "" } else { "span-0000" }),
),
(
"SpanName".into(),
json!(if name_bytes == 0 {
format!("step-{step}")
} else {
"x".repeat(name_bytes)
}),
),
(
"ObservationType".into(),
json!(if step == 0 {
"agent"
} else if costed {
"llm"
} else {
"tool"
}),
),
("TeamId".into(), json!("team-a")),
("ApiKeyHash".into(), json!("key-a")),
("Duration".into(), json!(1000)),
(
"LiteLLMRequestId".into(),
json!(if costed && step > 0 {
format!("response-{step}")
} else {
String::new()
}),
),
])
})
})
.collect::<Vec<_>>();
for chunk in rows.chunks(100) {
insert_rows(
client,
&writer,
DATABASE,
InsertTable::OtelTraces,
chunk.to_vec(),
)
.await?;
}
if costed {
let costs = (1..steps)
.map(|step| {
BTreeMap::from([
("request_id".into(), json!(format!("request-{step}"))),
("response_id".into(), json!(format!("response-{step}"))),
("team_id".into(), json!("team-a")),
("api_key".into(), json!("key-a")),
("start_time".into(), json!(1_790_000_000_000_i64)),
("end_time".into(), json!(1_790_000_000_001_i64)),
("spend".into(), json!(0.25)),
])
})
.collect::<Vec<_>>();
insert_rows(client, &writer, DATABASE, InsertTable::SpendLogs, costs).await?;
}
let reader = fixture
.readers
.connection(client, &QueryScope::All, "fixture-secret")
.await?;
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?;
assert_eq!(page.data.len(), runs);
assert!(
page.data
.windows(2)
.all(|runs| runs[0].trace_ref > runs[1].trace_ref)
);
if runs > 1 {
client
.post(writer.url().clone())
.body("SYSTEM FLUSH LOGS")
.send()
.await?
.error_for_status()?;
let read_queries = client.post(writer.url().clone()).body(format!(
"SELECT count() FROM system.query_log WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' AND query LIKE '%FROM otel_traces AS o%' AND query NOT LIKE '%system.query_log%'"
)).send().await?.error_for_status()?.text().await?;
let read_queries = read_queries.trim().parse::<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>()?;
assert!(
(2..=4).contains(&overlapping),
"{overlapping} simultaneous spend reads for {runs} runs"
);
}
}
for summary in &page.data {
assert_eq!(summary.span_count, steps as u64);
assert_eq!(
if costed {
summary.llm_calls
} else {
summary.tool_calls
},
(steps - 1) as u64
);
if costed {
assert_eq!(summary.spend, Some((steps - 1) as f64 * 0.25));
}
}
let trace_ref = &page
.data
.iter()
.find(|run| run.trace_id == "trace-0000")
.ok_or("missing run")?
.trace_ref;
let detail = get_trace(client, &reader, &access, "trace-0000", trace_ref)
.await?
.ok_or("missing trace")?;
assert_eq!(detail.spans.len(), steps);
assert_eq!(detail.spans[0].span_id, "span-0000");
assert_eq!(
detail.spans[steps - 1].span_id,
format!("span-{:04}", steps - 1)
);
assert_eq!(
if costed {
detail.summary.llm_calls
} else {
detail.summary.tool_calls
},
(steps - 1) as u64
);
let denied = ReadAccessParams {
team_ids: vec!["other-team".into()],
..access.clone()
};
assert!(
get_trace(client, &reader, &denied, "trace-0000", trace_ref)
.await?
.is_none()
);
let mut cursor = None;
let mut ids = Vec::new();
loop {
let page = get_trace_page(
client,
&reader,
&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!(
serde_json::to_vec(&page)?.len()
<= litellm_storage_clickhouse::READ_LIMITS.response_bytes
);
if ids.is_empty() {
assert!(
get_trace_page(
client,
&reader,
&denied,
"trace-0000",
trace_ref,
page.next_cursor.as_deref(),
200,
)
.await?
.is_none()
);
client
.post(writer.url().clone())
.body(format!("TRUNCATE TABLE {DATABASE}.otel_traces"))
.send()
.await?
.error_for_status()?;
}
ids.extend(page.spans.into_iter().map(|span| span.span_id));
cursor = page.next_cursor;
if cursor.is_none() {
break;
}
}
assert_eq!(
ids,
detail
.spans
.iter()
.map(|span| span.span_id.clone())
.collect::<Vec<_>>()
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive(
#[future(awt)] seeded_database: TestResult<SeededDatabase>,
) -> TestResult {
let fixture = seeded_database?;
let client = &fixture.database.client;
let reader = fixture
.readers
.connection(client, &QueryScope::All, "fixture-secret")
.await?;
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 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 writer = Connection::writer(&fixture.database.url)?;
insert_rows(
client,
&writer,
DATABASE,
InsertTable::OtelTraces,
vec![BTreeMap::from([
("Timestamp".into(), json!(1_790_000_000_000_000_000_i64)),
("TraceId".into(), json!(summary.trace_id)),
("SpanId".into(), json!("late-span")),
("ParentSpanId".into(), json!(first.spans[0].span_id)),
("TeamId".into(), json!("team-a")),
("ApiKeyHash".into(), json!("key-a")),
("EngineReceivedMs".into(), json!(u64::MAX / 2)),
])],
)
.await?;
let denied = ReadAccessParams {
all_teams: false,
user_id: String::new(),
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()
);
let first_cursor = first.next_cursor.clone();
let mut cursor = first.next_cursor;
let mut ids = first
.spans
.into_iter()
.map(|span| span.span_id)
.collect::<Vec<_>>();
while let Some(current) = cursor {
let next = get_trace_page(
client,
&reader,
&access,
&summary.trace_id,
&summary.trace_ref,
Some(&current),
1,
)
.await?
.ok_or("missing next page")?;
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")?;
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"))
));
let backdated = json!({
"Timestamp": "2026-09-01 00:00:00.000000000",
"TraceId": summary.trace_id,
"SpanId": "backdated-span",
"EngineReceivedMs": 1,
"TeamId": "team-a",
"ApiKeyHash": "key-a"
});
client
.post(writer.url().clone())
.body(format!(
"INSERT INTO {DATABASE}.otel_traces FORMAT JSONEachRow\n{backdated}"
))
.send()
.await?
.error_for_status()?;
let uncached_reader =
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;
assert!(
matches!(changed, Err(litellm_traces_clickhouse::Error::TraceChanged)),
"{changed:?}"
);
Ok(())
}
#[rstest]
#[tokio::test]
async fn an_oversized_span_keeps_the_run_list_available_with_partial_totals(
#[future(awt)] seeded_database: TestResult<SeededDatabase>,
) -> TestResult {
let fixture = seeded_database?;
let client = &fixture.database.client;
let reader = fixture
.readers
.connection(client, &QueryScope::All, "fixture-secret")
.await?;
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 run = before
.data
.iter()
.find(|run| run.span_count == 3)
.ok_or("missing fixture")?;
let writer = Connection::writer(&fixture.database.url)?;
insert_rows(
client,
&writer,
DATABASE,
InsertTable::OtelTraces,
vec![BTreeMap::from([
("Timestamp".into(), json!(1_790_000_000_000_000_000_i64)),
("TraceId".into(), json!(run.trace_id)),
("SpanId".into(), json!("oversized-child")),
("ParentSpanId".into(), json!("0101010101010101")),
(
"SpanName".into(),
json!("x".repeat(litellm_storage_clickhouse::READ_LIMITS.response_bytes + 1)),
),
("ObservationType".into(), json!("tool")),
("TeamId".into(), json!("team-a")),
("ApiKeyHash".into(), json!("key-a")),
])],
)
.await?;
let after = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
assert_eq!(after.data.len(), before.data.len());
let limited = after
.data
.iter()
.find(|item| item.trace_ref == run.trace_ref)
.ok_or("missing run")?;
assert!(limited.resolution_limited);
assert_eq!(limited.span_count, 4);
assert!(
after
.data
.iter()
.filter(|item| item.trace_ref != run.trace_ref)
.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)
));
Ok(())
}

View file

@ -171,6 +171,7 @@ pub fn resolve_trace(
.map(|(_, (span, _))| span.input_preview.clone())
.unwrap_or_default();
let summary = TraceSummary {
resolution_limited: false,
trace_id: trace_id.to_owned(),
trace_ref: trace_ref.to_owned(),
name: spans[root].name.clone(),
@ -209,11 +210,13 @@ pub fn resolve_trace(
summary,
agents,
spans,
next_cursor: None,
})
}
pub fn listed_summary(row: &ListTracesRow) -> TraceSummary {
TraceSummary {
resolution_limited: true,
trace_id: row.trace_id.clone(),
trace_ref: row.trace_ref.clone(),
name: row.name.clone(),

View file

@ -17,7 +17,7 @@ pub enum SpanStatus {
}
#[macro_rules_attribute::apply(response_type)]
#[derive(Debug, PartialEq)]
#[derive(Clone, Debug, PartialEq)]
pub struct Span {
pub span_id: String,
pub parent_span_id: Option<String>,
@ -41,7 +41,7 @@ pub struct Span {
/// One distinct agent in a trace: 200 invocations of `researcher` are one node.
#[macro_rules_attribute::apply(response_type)]
#[derive(Debug, PartialEq)]
#[derive(Clone, Debug, PartialEq)]
pub struct AgentNode {
pub name: String,
pub parent_agent: Option<String>,
@ -53,8 +53,10 @@ pub struct AgentNode {
}
#[macro_rules_attribute::apply(response_type)]
#[derive(Debug, PartialEq)]
#[derive(Clone, Debug, PartialEq)]
pub struct TraceSummary {
#[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))]
pub resolution_limited: bool,
pub trace_id: String,
#[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))]
pub trace_ref: String,
@ -81,11 +83,13 @@ pub struct TraceSummary {
}
#[macro_rules_attribute::apply(response_type)]
#[derive(Debug, PartialEq)]
#[derive(Clone, Debug, PartialEq)]
pub struct Trace {
pub summary: TraceSummary,
pub agents: Vec<AgentNode>,
pub spans: Vec<Span>,
#[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))]
pub next_cursor: Option<String>,
}
#[macro_rules_attribute::apply(response_type)]

View file

@ -1473,6 +1473,7 @@ from .embeddings.dispatch import *
from .rust_bridge import rust
from .rag.main import *
from .sandbox.main import *
from .decisions.main import *
from .search.main import *
from .realtime_api.main import (
_arealtime,
@ -2302,6 +2303,7 @@ _AGENT_EXPORTS: Final = frozenset(
"CodexOptions",
"OpenCodeOptions",
"DeepAgentsOptions",
"ToolLoopOptions",
}
)

View file

@ -99,6 +99,7 @@ from litellm.llms.vertex_ai.cost_calculator import cost_router as google_cost_ro
from litellm.llms.xai.cost_calculator import cost_per_token as xai_cost_per_token
from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.agents import LiteLLMSendMessageResponse
from litellm.types.decisions import DecisionsResponse, DecisionsUsage
from litellm.types.llms.base import CachedTokensDetails
from litellm.types.llms.openai import (
HttpxBinaryResponseContent,
@ -1058,6 +1059,7 @@ def _is_known_usage_objects(usage_obj):
return (
isinstance(usage_obj, litellm.Usage)
or isinstance(usage_obj, ResponseAPIUsage)
or isinstance(usage_obj, DecisionsUsage)
or TranscriptionUsageObjectTransformation.is_transcription_usage_object(usage_obj)
)
@ -1466,7 +1468,12 @@ def completion_cost(
"usage",
litellm.Usage(**_usage_for_dump.model_dump()),
)
if usage_obj is None:
if isinstance(usage_obj, DecisionsUsage):
_usage = {
"prompt_tokens": usage_obj.input_tokens,
"completion_tokens": usage_obj.output_tokens,
}
elif usage_obj is None:
_usage = {}
elif isinstance(usage_obj, BaseModel):
_usage = cast(BaseModel, usage_obj).model_dump()
@ -1957,7 +1964,8 @@ def response_cost_calculator(
| LiteLLMRealtimeStreamLoggingObject
| OpenAIModerationResponse
| Response
| SearchResponse,
| SearchResponse
| DecisionsResponse,
model: str,
custom_llm_provider: str | None,
call_type: Literal[
@ -1979,6 +1987,8 @@ def response_cost_calculator(
"arerank",
"search",
"asearch",
"decisions",
"adecisions",
],
optional_params: dict,
cache_hit: bool | None = None,

View file

@ -0,0 +1,3 @@
from litellm.decisions.main import adecisions, decisions
__all__ = ["adecisions", "decisions"]

299
litellm/decisions/main.py Normal file
View file

@ -0,0 +1,299 @@
from collections.abc import Mapping
from dataclasses import dataclass, field
from types import MappingProxyType
from typing import Final
import httpx
from pydantic import TypeAdapter, ValidationError
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.decisions.transformation import DecisionsProviderConfig
from litellm.llms.cloudflare.decisions.transformation import CLOUDFLARE_DECISIONS_ENDPOINT
from litellm.llms.custom_httpx.http_handler import _get_httpx_client, get_async_httpx_client
from litellm.llms.openrouter.decisions.transformation import OPENROUTER_DECISIONS_ENDPOINT
from litellm.llms.perplexity.decisions.transformation import PERPLEXITY_DECISIONS_ENDPOINT
from litellm.llms.strands_decider.decisions.transformation import STRANDS_DECIDER_DECISIONS_ENDPOINT
from litellm.llms.typesafe.decisions.transformation import TYPESAFE_DECISIONS_ENDPOINT
from litellm.secret_managers.main import get_secret_str
from litellm.types.decisions import (
DecisionQuestion,
DecisionsJSON,
DecisionsRequest,
DecisionsResponse,
)
from litellm.utils import client
DECISIONS_ENDPOINTS: Final[Mapping[str, DecisionsProviderConfig]] = MappingProxyType(
{
"perplexity": PERPLEXITY_DECISIONS_ENDPOINT,
"typesafe": TYPESAFE_DECISIONS_ENDPOINT,
"openrouter": OPENROUTER_DECISIONS_ENDPOINT,
"cloudflare": CLOUDFLARE_DECISIONS_ENDPOINT,
"strands_decider": STRANDS_DECIDER_DECISIONS_ENDPOINT,
}
)
_DECISIONS_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequest]] = TypeAdapter(DecisionsRequest)
_DECISIONS_PAYLOAD_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object)
_DECISIONS_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse)
@dataclass(frozen=True, slots=True, repr=False)
class _PreparedDecisionsRequest:
config: DecisionsProviderConfig
provider: str
upstream_model: str
url: str
api_key: str | None = field(repr=False)
headers: Mapping[str, str] = field(repr=False)
body: Mapping[str, object] = field(repr=False)
def _resolve_provider_model(model: str, custom_llm_provider: str | None) -> tuple[str, str]:
provider: Final = model.partition("/")[0] if custom_llm_provider is None else custom_llm_provider
if provider not in DECISIONS_ENDPOINTS:
supported: Final = ", ".join(DECISIONS_ENDPOINTS)
raise litellm.BadRequestError(
message=f"Unknown Decisions provider '{provider}'. Supported providers: {supported}",
model=model,
llm_provider=provider,
)
upstream_model: Final = model.removeprefix(f"{provider}/") if model.startswith(f"{provider}/") else model
if not upstream_model:
raise litellm.BadRequestError(
message="A model name is required for the Decisions API",
model=model,
llm_provider=provider,
)
return provider, upstream_model
def _resolve_api_key(
*,
provider: str,
model: str,
endpoint: DecisionsProviderConfig,
api_key: str | None,
) -> str | None:
if api_key is not None:
return api_key
server_api_key: Final = next(
(key for key in (get_secret_str(name) for name in endpoint.api_key_env) if key),
None,
)
if server_api_key is None:
if not endpoint.api_key_required:
return None
raise litellm.AuthenticationError(
message=f"Missing API key for Decisions provider '{provider}'",
model=model,
llm_provider=provider,
)
return server_api_key
def _prepare_request(
*,
model: str,
state: DecisionsJSON,
questions: Mapping[str, DecisionQuestion | Mapping[str, object]],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, str] | None,
) -> _PreparedDecisionsRequest:
provider, upstream_model = _resolve_provider_model(model, custom_llm_provider)
try:
validated_request: Final = _DECISIONS_REQUEST_ADAPTER.validate_python(
{"model": model, "state": state, "questions": questions}
)
except ValidationError as error:
raise litellm.BadRequestError(
message=f"Invalid Decisions request: {error}",
model=model,
llm_provider=provider,
) from error
endpoint: Final = DECISIONS_ENDPOINTS[provider]
env_api_base: Final = get_secret_str(endpoint.api_base_env)
default_api_base: Final = endpoint.default_api_base()
resolved_api_base: Final = api_base or env_api_base or default_api_base
if resolved_api_base is None:
raise litellm.BadRequestError(
message=endpoint.missing_api_base_message(provider),
model=model,
llm_provider=provider,
)
resolved_api_key: Final = _resolve_api_key(
provider=provider,
model=model,
endpoint=endpoint,
api_key=api_key,
)
canonical_model: Final = endpoint.canonical_model(upstream_model)
outbound_headers: Final = MappingProxyType(
{
**{
name: value
for name, value in (extra_headers or {}).items()
if name.lower() not in {"authorization", "content-type"}
},
**({"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key is not None else {}),
"Content-Type": "application/json",
}
)
body: Final = MappingProxyType(
{
"model": endpoint.request_model(canonical_model),
"state": validated_request.state,
"questions": {
name: question.model_dump(mode="json", exclude_none=True)
for name, question in validated_request.questions.items()
},
}
)
return _PreparedDecisionsRequest(
config=endpoint,
provider=provider,
upstream_model=canonical_model,
url=endpoint.endpoint_url(resolved_api_base, canonical_model),
api_key=resolved_api_key,
headers=outbound_headers,
body=body,
)
def _log_request(
prepared: _PreparedDecisionsRequest,
kwargs: Mapping[str, object],
) -> LiteLLMLoggingObj | None:
logging_obj: Final = kwargs.get("litellm_logging_obj")
if not isinstance(logging_obj, LiteLLMLoggingObj):
return None
logging_obj.update_from_kwargs(
kwargs=dict(kwargs),
model=prepared.upstream_model,
litellm_params={
"litellm_call_id": kwargs.get("litellm_call_id"),
"api_base": prepared.url,
},
custom_llm_provider=prepared.provider,
)
request_body: Final = dict(prepared.body)
request_headers: Final = dict(prepared.headers)
logging_obj.pre_call(
input=request_body,
api_key=prepared.api_key,
model=prepared.upstream_model,
additional_args={
"api_base": prepared.url,
"complete_input_dict": request_body,
"headers": request_headers,
},
)
return logging_obj
def _parse_response(
response: httpx.Response,
prepared: _PreparedDecisionsRequest,
) -> DecisionsResponse:
response.raise_for_status()
payload: Final[object] = _DECISIONS_PAYLOAD_ADAPTER.validate_json(response.content)
result: Final = _DECISIONS_RESPONSE_ADAPTER.validate_python(prepared.config.unwrap_response(payload))
result._hidden_params.update(
{
"model": f"{prepared.provider}/{prepared.upstream_model}",
"custom_llm_provider": prepared.provider,
"provider_response_model": f"{prepared.provider}/{prepared.upstream_model}",
}
)
return result
def _map_upstream_exception(error: Exception, prepared: _PreparedDecisionsRequest) -> Exception:
return litellm.exception_type(
model=f"{prepared.provider}/{prepared.upstream_model}",
custom_llm_provider=prepared.provider,
original_exception=error,
)
@client
async def adecisions(
model: str,
state: DecisionsJSON,
questions: Mapping[str, DecisionQuestion | Mapping[str, object]],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, str] | None = None,
**kwargs: object,
) -> DecisionsResponse:
prepared: Final = _prepare_request(
model=model,
state=state,
questions=questions,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
)
logging_obj: Final = _log_request(prepared, kwargs)
try:
handler: Final = get_async_httpx_client(llm_provider=prepared.provider)
response: Final = await handler.post(
prepared.url,
json=dict(prepared.body),
headers=dict(prepared.headers),
timeout=timeout,
logging_obj=logging_obj,
)
return _parse_response(response=response, prepared=prepared)
except Exception as error:
raise _map_upstream_exception(error, prepared) from error
@client
def decisions(
model: str,
state: DecisionsJSON,
questions: Mapping[str, DecisionQuestion | Mapping[str, object]],
api_key: str | None = None,
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
custom_llm_provider: str | None = None,
extra_headers: Mapping[str, str] | None = None,
**kwargs: object,
) -> DecisionsResponse:
prepared: Final = _prepare_request(
model=model,
state=state,
questions=questions,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
)
logging_obj: Final = _log_request(prepared, kwargs)
try:
handler: Final = _get_httpx_client()
response: Final = handler.post(
prepared.url,
json=dict(prepared.body),
headers=dict(prepared.headers),
timeout=timeout,
logging_obj=logging_obj,
)
return _parse_response(response=response, prepared=prepared)
except Exception as error:
raise _map_upstream_exception(error, prepared) from error
__all__ = ["DECISIONS_ENDPOINTS", "adecisions", "decisions"]

View file

@ -1,4 +1,4 @@
"""Agent harnesses: run Claude Code, Codex, OpenCode or Deep Agents on any LiteLLM model.
"""Agent harnesses: run Claude Code, Codex, OpenCode, Deep Agents or Tool Loop on any LiteLLM model.
The entrypoints live on the top-level package:
@ -30,6 +30,7 @@ from litellm.harness.options import (
CodexOptions,
DeepAgentsOptions,
OpenCodeOptions,
ToolLoopOptions,
)
from litellm.harness.runtime import (
AsyncEventStream,
@ -86,6 +87,7 @@ __all__ = (
"StateIncompatible",
"Text",
"ToolCall",
"ToolLoopOptions",
"ToolResult",
"Usage",
"aagent",

View file

@ -1,4 +1,4 @@
"""Handlers run a harness config: CLI runtimes as subprocesses, Deep Agents in-process."""
"""Handlers run a harness config: CLI runtimes and in-process harnesses."""
from __future__ import annotations
@ -29,6 +29,10 @@ def get_harness_handler(config: BaseHarnessConfig) -> BaseHarnessHandler:
from litellm.harness.handlers.deepagents_handler import DeepAgentsHandler
return DeepAgentsHandler(config)
if config.harness is Harness.TOOL_LOOP:
from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler
return ToolLoopHandler(config)
raise HarnessError(f"No handler for Harness.{config.harness.name}")

View file

@ -0,0 +1,345 @@
"""In-process tool-calling loop handler."""
from __future__ import annotations
import asyncio
import copy
import inspect
import json
import math
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Protocol, TypeAlias
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
import litellm
from litellm.harness.context import SessionContext
from litellm.harness.errors import CapabilityUnsupported
from litellm.harness.handlers.base import BaseHarnessHandler
from litellm.harness.types import Approval, Event, Reasoning, Text, ToolCall, ToolResult
from litellm.llms.base_llm.harness.transformation import HarnessTurnError
from litellm.llms.tool_loop.harness.transformation import (
TOOL_LOOP_MAX_MODEL_CALLS,
FunctionTool,
ToolLoopHarnessConfig,
completion_kwargs,
function_tool,
)
from litellm.types.completion import ChatCompletionMessageParam
from litellm.types.utils import (
ChatCompletionMessageCustomToolCall,
ChatCompletionMessageToolCall,
ChatCompletionToolParam,
ModelResponse,
)
@dataclass(frozen=True, slots=True)
class _FunctionToolCall:
id: str
name: str
arguments: str
AsyncCompletion: TypeAlias = Callable[..., Awaitable[ModelResponse]]
_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object])
_MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_MODEL_RESPONSE_ADAPTER: Final = TypeAdapter(ModelResponse)
_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int])
_JSON_DECODER: Final = json.JSONDecoder()
_HISTORY_ADAPTER: Final = TypeAdapter(list[dict[str, object]])
class _Usage(BaseModel):
model_config = ConfigDict(from_attributes=True)
prompt_tokens: int | None = None
completion_tokens: int | None = None
_USAGE_ADAPTER: Final = TypeAdapter(_Usage)
class _AwaitableObject(Protocol):
def __await__(self) -> Generator[object, None, object]: ...
async def _await_tool_result(result: _AwaitableObject) -> object:
return await result
@dataclass(frozen=True, slots=True)
class _ToolOutcome:
output: str
is_error: bool
def _normalize_tool_call(
call: ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall,
) -> _FunctionToolCall:
if isinstance(call, ChatCompletionMessageToolCall):
return _FunctionToolCall(
id=call.id,
name=call.function.name or "",
arguments=call.function.arguments,
)
return _FunctionToolCall(id=call.id, name=call.custom.name, arguments=call.custom.input)
def _parse_arguments(raw: str) -> tuple[dict[str, object], str | None]: # mutable-ok: event input uses a dict
text: Final = raw.strip()
try:
raw_decoded: object = _JSON_DECODER.raw_decode(text)
except json.JSONDecodeError as error:
return {}, f"{type(error).__name__}: {error}"
parsed, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded)
if text[end:].strip():
return {}, "JSONDecodeError: Extra data after tool arguments"
try:
arguments: Final = _ARGUMENTS_ADAPTER.validate_python(parsed)
except ValidationError as error:
if not isinstance(parsed, dict):
return {}, "ValueError: tool arguments must be a JSON object"
return {}, f"{type(error).__name__}: {error}"
return arguments, None
def _cost_value(value: object) -> float | None:
if isinstance(value, bool):
return None
if not isinstance(value, str | int | float):
return None
try:
cost: Final = float(value)
except (OverflowError, TypeError, ValueError):
return None
return cost if math.isfinite(cost) else None
def _as_mapping(value: object) -> Mapping[str, object] | None:
try:
return _MAPPING_ADAPTER.validate_python(value)
except ValidationError:
return None
def _response_cost(response: ModelResponse) -> float:
hidden_params_value: Final[object] = getattr(response, "_hidden_params", {})
hidden_params: Final = _as_mapping(hidden_params_value)
if hidden_params is not None:
additional_headers: Final = _as_mapping(hidden_params.get("additional_headers"))
if additional_headers is not None:
header_cost: Final = _cost_value(additional_headers.get("llm_provider-x-litellm-response-cost"))
if header_cost is not None:
return header_cost
hidden_cost: Final = _cost_value(hidden_params.get("response_cost"))
if hidden_cost is not None:
return hidden_cost
try:
calculated_cost: Final = litellm.completion_cost(completion_response=response)
except Exception:
return 0.0
return _cost_value(calculated_cost) or 0.0
async def _approval_error(approval: Approval | None) -> str | None:
if approval is None:
return None
allowed, reason = await approval.wait()
return None if allowed else f"denied: {reason}"
async def _tool_outcome(
tool: FunctionTool | None,
tool_name: str,
arguments: dict[str, object],
parse_error: str | None,
approval_error: str | None,
) -> _ToolOutcome:
if parse_error is not None:
return _ToolOutcome(output=parse_error, is_error=True)
if approval_error is not None:
return _ToolOutcome(output=approval_error, is_error=True)
if tool is None:
return _ToolOutcome(output=f"ValueError: unknown tool {tool_name!r}", is_error=True)
try:
result: Final = await _execute_tool(tool, arguments)
output: Final = result if isinstance(result, str) else json.dumps(result, default=str)
return _ToolOutcome(output=output, is_error=False)
except Exception as error:
return _ToolOutcome(output=f"{type(error).__name__}: {error}", is_error=True)
def _record_usage(ctx: SessionContext, response: ModelResponse) -> None:
try:
usage_value: Final[object] = getattr(response, "usage", None)
usage: Final = _USAGE_ADAPTER.validate_python(usage_value)
input_tokens: Final = usage.prompt_tokens or 0
output_tokens: Final = usage.completion_tokens or 0
ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink
ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink
except Exception:
return
async def _execute_tool(tool: FunctionTool, arguments: Mapping[str, object]) -> object:
validated_model: Final = tool.args_model.model_validate(arguments)
values_object: Final[object] = validated_model.model_dump()
validated: Final = _MAPPING_ADAPTER.validate_python(values_object)
parameters: Final = tuple(inspect.signature(tool.fn).parameters.values())
positional_args: Final = tuple(
validated[parameter.name] for parameter in parameters if parameter.kind is inspect.Parameter.POSITIONAL_ONLY
)
keyword_args: Final[dict[str, object]] = { # mutable-ok: tool calls need keyword arguments
parameter.name: validated[parameter.name]
for parameter in parameters
if parameter.kind is not inspect.Parameter.POSITIONAL_ONLY
}
if inspect.iscoroutinefunction(tool.fn):
async_result: Final[object] = tool.fn(*positional_args, **keyword_args)
if inspect.isawaitable(async_result):
return await _await_tool_result(async_result)
return async_result
sync_result: Final[object] = await asyncio.to_thread(tool.fn, *positional_args, **keyword_args)
if inspect.isawaitable(sync_result):
return await _await_tool_result(sync_result)
return sync_result
async def _default_acompletion(**kwargs: object) -> ModelResponse: # kwargs-ok: provider-specific completion options
response: Final[object] = await litellm.acompletion(**kwargs)
return _MODEL_RESPONSE_ADAPTER.validate_python(response)
class ToolLoopHandler(BaseHarnessHandler):
def __init__(
self,
config: ToolLoopHarnessConfig,
acompletion: AsyncCompletion | None = None,
) -> None:
super().__init__(config) # pyright: ignore[reportUnknownMemberType] # base handler config is unparameterized
self._config = config
self._acompletion = acompletion if acompletion is not None else _default_acompletion
self._messages: tuple[ChatCompletionMessageParam, ...] = ()
self._tools: Mapping[str, FunctionTool] = MappingProxyType({})
self._tool_specs: tuple[ChatCompletionToolParam, ...] = ()
self._completion_kwargs: Mapping[str, object] = MappingProxyType({})
async def start(self, ctx: SessionContext) -> None:
self._config.validate_environment(ctx)
tools: Final = tuple(function_tool(fn) for fn in ctx.tools)
if len({tool.name for tool in tools}) != len(tools):
raise ValueError("Harness.TOOL_LOOP tool names must be unique")
self._tools = MappingProxyType({tool.name: tool for tool in tools})
self._tool_specs = tuple(tool.spec for tool in tools)
self._completion_kwargs = MappingProxyType(completion_kwargs(ctx))
if not self._messages and ctx.instructions:
self._messages = ({"role": "system", "content": ctx.instructions},)
def native_session_id(self) -> str | None:
return None
async def resume(self, ctx: SessionContext, native_session_id: str) -> None:
raise CapabilityUnsupported("Harness.TOOL_LOOP does not support resume")
async def history(
self, ctx: SessionContext
) -> list[dict[str, object]]: # mutable-ok: public API returns copied message dictionaries
history_object: Final[object] = copy.deepcopy(list(self._messages))
return _HISTORY_ADAPTER.validate_python(history_object) # pyright: ignore[reportIncompatibleMethodOverride] # base history uses Any
async def stop(self, ctx: SessionContext) -> None:
return None
async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
ctx.final_text = "" # rebind-ok: SessionContext is the runtime's per-turn result sink
ctx.output_json = None # rebind-ok: SessionContext is the runtime's per-turn result sink
user_message: Final[ChatCompletionMessageParam] = {"role": "user", "content": prompt}
self._messages = (*self._messages, user_message)
for _ in range(TOOL_LOOP_MAX_MODEL_CALLS):
messages: list[ChatCompletionMessageParam] = copy.deepcopy( # mutable-ok: acompletion takes list messages
list(self._messages)
)
tool_specs: list[ChatCompletionToolParam] = copy.deepcopy( # mutable-ok: acompletion takes tool list
list(self._tool_specs)
)
request_kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}
}
kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
**request_kwargs,
"messages": messages,
**({"tools": tool_specs} if tool_specs else {}),
}
response = await self._acompletion(**kwargs)
_record_usage(ctx, response)
message = response.choices[0].message
reasoning_value: object = getattr(message, "reasoning_content", None)
reasoning = reasoning_value if isinstance(reasoning_value, str) else None
content = message.content
if reasoning:
yield Reasoning(reasoning)
if content:
yield Text(content)
tool_calls = message.tool_calls or ()
if not tool_calls:
final_text = content or ""
ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink
ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink
final_message: ChatCompletionMessageParam = {
"role": "assistant",
"content": content,
}
self._messages = (*self._messages, final_message)
return
normalized_calls = tuple(_normalize_tool_call(call) for call in tool_calls)
assistant_message: ChatCompletionMessageParam = {
"role": "assistant",
"content": content,
"tool_calls": [
{
"id": call.id,
"type": "function",
"function": {"name": call.name, "arguments": call.arguments},
}
for call in normalized_calls
],
}
self._messages = (*self._messages, assistant_message)
for call in normalized_calls:
arguments, parse_error = _parse_arguments(call.arguments)
yield ToolCall(
id=call.id,
name=call.name,
native_name=call.name,
input=arguments,
builtin=False,
)
approval = (
Approval(tool=call.name, input=arguments)
if parse_error is None and ctx.permissions == "ask"
else None
)
if approval is not None:
yield approval
approval_error = await _approval_error(approval)
tool = self._tools.get(call.name)
outcome = await _tool_outcome(
tool,
call.name,
arguments,
parse_error,
approval_error,
)
yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error)
tool_message: ChatCompletionMessageParam = {
"role": "tool",
"tool_call_id": call.id,
"content": outcome.output,
}
self._messages = (*self._messages, tool_message)
raise HarnessTurnError(f"Harness.TOOL_LOOP exceeded {TOOL_LOOP_MAX_MODEL_CALLS} model calls in one turn")

View file

@ -34,4 +34,9 @@ class DeepAgentsOptions:
recursion_limit: int | None = None
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions
@dataclass(frozen=True)
class ToolLoopOptions:
completion_kwargs: Mapping[str, Any] = field(default_factory=dict)
HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions | ToolLoopOptions

View file

@ -411,7 +411,7 @@ def agent(
options: HarnessOptions | None = None,
install: bool = False,
) -> Result | EventStream:
"""Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents) on one prompt.
"""Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents, Tool Loop) on one prompt.
Returns a Result. With stream=True it returns an iterator of events instead.
Prefix the model with `litellm_proxy/` to route every model call through your

View file

@ -27,6 +27,7 @@ class Harness(Enum):
CODEX = "codex"
OPENCODE = "opencode"
DEEPAGENTS = "deepagents"
TOOL_LOOP = "tool_loop"
def require_harness(harness: object) -> Harness:

View file

@ -290,6 +290,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset(
"LiteLLM_Config",
"LiteLLM_SpendLogs",
"LiteLLM_BudgetWindowSpend",
"LiteLLM_BackgroundInteractionSettlement",
"LiteLLM_ErrorLogs",
"LiteLLM_UserNotifications",
"LiteLLM_TeamMembership",

View file

@ -168,6 +168,7 @@ class HealthCheckHelpers:
"batch",
"responses",
"ocr",
"evaluation",
],
Callable,
]:
@ -190,7 +191,7 @@ class HealthCheckHelpers:
from litellm.litellm_core_utils.audio_utils.utils import (
get_audio_file_for_health_check,
)
from litellm.litellm_core_utils.health_check_utils import _filter_model_params
from litellm.litellm_core_utils.health_check_utils import DECISIONS_CALL_PARAMS, _filter_model_params
from litellm.realtime_api.main import _realtime_health_check
return {
@ -257,4 +258,13 @@ class HealthCheckHelpers:
**_filter_model_params(model_params=model_params),
document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider),
),
"evaluation": lambda: litellm.adecisions(
**DECISIONS_CALL_PARAMS.validate_python(
{
"state": prompt or "health check",
"questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}},
**_filter_model_params(model_params=model_params),
}
)
),
}

View file

@ -4,6 +4,12 @@ Utils used for litellm.ahealth_check()
from typing import Final
from pydantic import TypeAdapter
from litellm.types.decisions import DecisionsCallParams
DECISIONS_CALL_PARAMS: Final[TypeAdapter[DecisionsCallParams]] = TypeAdapter(DecisionsCallParams)
def _filter_model_params(model_params: dict) -> dict:
"""Remove 'messages' param from model params."""

View file

@ -118,6 +118,7 @@ from litellm.llms.base_llm.search.transformation import SearchResponse
from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.agents import LiteLLMSendMessageResponse
from litellm.types.containers.main import ContainerObject
from litellm.types.decisions import DecisionsResponse
from litellm.types.integrations.s3_v2 import S3PartitionGranularity
from litellm.types.interactions import (
InteractionsAPIResponse,
@ -1815,6 +1816,7 @@ class Logging(LiteLLMLoggingBaseClass):
LiteLLMRealtimeStreamLoggingObject,
OpenAIModerationResponse,
"SearchResponse",
DecisionsResponse,
dict,
list,
],
@ -2600,6 +2602,7 @@ class Logging(LiteLLMLoggingBaseClass):
or isinstance(logging_result, OpenAIModerationResponse)
or isinstance(logging_result, OCRResponse) # OCR
or isinstance(logging_result, SearchResponse) # Search API
or isinstance(logging_result, DecisionsResponse)
or (
isinstance(logging_result, InteractionsAPIResponse)
and logging_result.usage is not None

View file

@ -0,0 +1,3 @@
from .transformation import DecisionsProviderConfig, JevCompatibleDecisionsEndpoint
__all__ = ["DecisionsProviderConfig", "JevCompatibleDecisionsEndpoint"]

View file

@ -0,0 +1,52 @@
from dataclasses import dataclass
from typing import Protocol
@dataclass(frozen=True, slots=True)
class JevCompatibleDecisionsEndpoint:
default_api_base_value: str | None
path: str
api_key_env: tuple[str, ...]
api_base_env: str
api_key_required: bool = True
def default_api_base(self) -> str | None:
return self.default_api_base_value
def missing_api_base_message(self, provider: str) -> str:
return f"api_base is required for Decisions provider '{provider}'"
def canonical_model(self, model: str) -> str:
return model
def request_model(self, model: str) -> str:
return model
def endpoint_url(self, api_base: str, model: str) -> str:
return f"{api_base.rstrip('/').removesuffix('/v1')}{self.path}"
def unwrap_response(self, payload: object) -> object:
return payload
class DecisionsProviderConfig(Protocol):
@property
def api_key_env(self) -> tuple[str, ...]: ...
@property
def api_base_env(self) -> str: ...
@property
def api_key_required(self) -> bool: ...
def default_api_base(self) -> str | None: ...
def missing_api_base_message(self, provider: str) -> str: ...
def canonical_model(self, model: str) -> str: ...
def request_model(self, model: str) -> str: ...
def endpoint_url(self, api_base: str, model: str) -> str: ...
def unwrap_response(self, payload: object) -> object: ...

View file

@ -5,17 +5,33 @@ from __future__ import annotations
import itertools
import json
import os
import typing
from collections.abc import Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import Any, Final, TypeAlias
from typing import Final, TypeAlias
from pydantic import TypeAdapter, ValidationError
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
if typing.TYPE_CHECKING:
from litellm.harness.context import SessionContext
# A decoded JSON document: what json.loads / model_json_schema() produce.
JSONValue: TypeAlias = "dict[str, JSONValue] | list[JSONValue] | str | int | float | bool | None"
SKILL_MANIFEST: Final = "SKILL.md"
_JSON_DECODER: Final = json.JSONDecoder()
_JSON_OBJECT_ADAPTER: Final = TypeAdapter(dict[str, object])
_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int])
def gateway_headers(ctx: SessionContext) -> dict[str, str]: # mutable-ok: acompletion(extra_headers=) requires dict
metadata_json: Final = json.dumps(dict(ctx.metadata), default=str) if ctx.metadata else None
return {
"x-litellm-tags": f"harness,{ctx.harness.value}",
**({"x-litellm-spend-logs-metadata": metadata_json} if metadata_json is not None else {}),
}
def normalize_tool_name(native_name: str, mapping: Mapping[str, str]) -> str:
@ -35,17 +51,18 @@ def last_json_object(text: str) -> str | None:
index = text.find("{")
while index != -1:
try:
obj, end = _JSON_DECODER.raw_decode(text, index)
raw_decoded: object = _JSON_DECODER.raw_decode(text, index)
except json.JSONDecodeError:
index = text.find("{", index + 1)
continue
obj, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded)
if isinstance(obj, dict):
last = json.dumps(obj)
index = text.find("{", end)
return last
def structured_output_instruction(schema: Mapping[str, Any]) -> str:
def structured_output_instruction(schema: Mapping[str, object]) -> str:
return (
"When you have finished, your final message must be a single JSON object that "
"matches this JSON schema, with no other text before or after it:\n"
@ -82,16 +99,15 @@ def strict_json_schema(schema: JSONValue, depth: int = 0) -> JSONValue:
return result
def decode_json_line(line: bytes | str) -> Mapping[str, Any] | None:
def decode_json_line(line: bytes | str) -> Mapping[str, object] | None:
"""One JSONL line as a dict, or None for blank / non-JSON / non-object lines."""
text = line.strip()
if not text:
return None
try:
obj = json.loads(text)
except json.JSONDecodeError:
return _JSON_OBJECT_ADAPTER.validate_json(text)
except ValidationError:
return None
return obj if isinstance(obj, dict) else None
def stderr_tail_text(stderr_tail: Sequence[str]) -> str:

View file

@ -1125,23 +1125,42 @@ def bedrock_model_accepts_cache_points(model: str | None) -> bool:
``cachePoint`` blocks. Bedrock rejects requests carrying cachePoint blocks for
models without prompt caching support ("You invoked an unsupported model or your
request did not allow prompt caching"), so a model whose cost-map entry does not declare
``supports_prompt_caching`` must not receive them. A model absent from the map
(an application inference profile ARN, a model newer than the map) keeps emitting
so existing caching setups never silently degrade. ``litellm.utils.supports_prompt_caching``
is not reusable here: it returns False for unmapped models, the opposite polarity.
``supports_prompt_caching`` must not receive them. An explicit
``supports_prompt_cache_breakpoint`` on the entry wins over that flag: a model can price
cached tokens through implicit caching yet reject the marker on Converse ("This model
doesn't support the cachePoint field", Kimi K3). The router registers a deployment's
``model_info`` under ``bedrock/<model>`` as configured, route prefix included, while the
Converse transformation sees the model with ``converse/`` or ``converse_like/`` already
stripped, so every registration form is read. That flag set there covers an application
inference profile ARN or a model newer than the map, while only the map decides whether
a model is known: absent a map entry the model keeps emitting so existing caching setups
never silently degrade. ``litellm.utils.supports_prompt_caching`` is not reusable here:
it returns False for unmapped models, the opposite polarity.
"""
if model is None:
return True
if _OPENAI_FAMILY_MODEL_RE.search(model):
return False
entries: Final = tuple(
entry
for candidate in (model, get_bedrock_base_model(model))
if (entry := litellm.model_cost.get(candidate)) is not None
map_keys: Final = (model, get_bedrock_base_model(model))
registered_keys: Final = tuple(f"bedrock/{route}{model}" for route in ("", "converse/", "converse_like/"))
explicit_marker_support: Final = next(
(
entry.get("supports_prompt_cache_breakpoint") is True
for key in (*registered_keys, *map_keys)
if (entry := litellm.model_cost.get(key)) is not None
and entry.get("supports_prompt_cache_breakpoint") is not None
),
None,
)
if not entries:
if explicit_marker_support is not None:
return explicit_marker_support
if not any(key in litellm.model_cost for key in map_keys):
return True
return any(entry.get("supports_prompt_caching") is True for entry in entries)
return any(
entry.get("supports_prompt_caching") is True
for key in map_keys
if (entry := litellm.model_cost.get(key)) is not None
)
def bedrock_supports_tool_search(model: str) -> bool:

View file

@ -0,0 +1,58 @@
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Final
from pydantic import TypeAdapter
from litellm.secret_managers.main import (
get_secret_str,
normalize_nonempty_secret_str,
)
_RESPONSE_MAPPING_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
@dataclass(frozen=True, slots=True)
class CloudflareDecisionsEndpoint:
api_key_env: tuple[str, ...] = ("CLOUDFLARE_API_KEY",)
api_base_env: str = "CLOUDFLARE_API_BASE"
api_key_required: bool = True
def default_api_base(self) -> str | None:
account_id: Final = normalize_nonempty_secret_str(get_secret_str("CLOUDFLARE_ACCOUNT_ID"))
if account_id is None:
return None
return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run"
def missing_api_base_message(self, provider: str) -> str:
return "Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID or pass api_base explicitly"
def canonical_model(self, model: str) -> str:
if model.startswith("@cf/"):
return model
return f"@cf/cloudflare/{model}"
def request_model(self, model: str) -> str:
return model.rsplit("/", maxsplit=1)[-1]
def endpoint_url(self, api_base: str, model: str) -> str:
normalized_api_base: Final = api_base.rstrip("/")
if normalized_api_base.endswith("/ai/v1"):
return f"{normalized_api_base.removesuffix('/ai/v1')}/ai/run/{model}"
if normalized_api_base.endswith("/ai/run"):
return f"{normalized_api_base}/{model}"
return f"{normalized_api_base}/ai/run/{model}"
def unwrap_response(self, payload: object) -> object:
if not isinstance(payload, Mapping):
return payload
response_mapping: Final = _RESPONSE_MAPPING_ADAPTER.validate_python(payload)
if "answers" in response_mapping:
return payload
result: Final = response_mapping.get("result")
if isinstance(result, Mapping):
return result
return payload
CLOUDFLARE_DECISIONS_ENDPOINT: Final[CloudflareDecisionsEndpoint] = CloudflareDecisionsEndpoint()

View file

@ -27,6 +27,7 @@ from litellm.harness.types import (
ToolResult,
)
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
from litellm.llms.base_llm.harness.utils import gateway_headers
if TYPE_CHECKING:
from litellm.harness.context import SessionContext
@ -70,18 +71,6 @@ APPROVAL_TOOLS: Final = WRITE_TOOLS | EXECUTE_TOOLS
_APPROVAL_DECISIONS: Final = ("approve", "reject")
def gateway_headers(
ctx: SessionContext,
) -> dict[str, str]: # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field
"""Same attribution headers the session endpoint adds for CLI harnesses."""
metadata = ctx.metadata
metadata_json = json.dumps(dict(metadata), default=str) if metadata else None # mutable-ok: for json.dumps
metadata_header = (("x-litellm-spend-logs-metadata", metadata_json),) if metadata_json is not None else ()
return dict( # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field
(("x-litellm-tags", f"harness,{ctx.harness.value}"), *metadata_header)
)
def chat_model_kwargs(
ctx: SessionContext,
) -> dict[str, Any]: # mutable-ok: ChatLiteLLM constructor kwargs, splatted as **kwargs

View file

@ -0,0 +1,10 @@
from typing import Final
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
OPENROUTER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
default_api_base_value="https://openrouter.ai/api",
path="/alpha/decisions",
api_key_env=("OPENROUTER_API_KEY",),
api_base_env="OPENROUTER_API_BASE",
)

View file

@ -0,0 +1,10 @@
from typing import Final
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
PERPLEXITY_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
default_api_base_value="https://api.perplexity.ai",
path="/v1/decisions",
api_key_env=("PERPLEXITYAI_API_KEY", "PERPLEXITY_API_KEY"),
api_base_env="PERPLEXITY_API_BASE",
)

View file

@ -0,0 +1,11 @@
from typing import Final
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
STRANDS_DECIDER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
default_api_base_value=None,
path="/v1/systemone",
api_key_env=("STRANDS_DECIDER_API_KEY",),
api_base_env="STRANDS_DECIDER_API_BASE",
api_key_required=False,
)

View file

@ -0,0 +1 @@
"""In-process tool loop harness."""

View file

@ -0,0 +1 @@
"""Tool Loop harness configuration."""

View file

@ -0,0 +1,119 @@
"""Configuration and tool-schema helpers for the in-process Tool Loop harness."""
from __future__ import annotations
import inspect
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict, TypeAdapter, create_model
from litellm.harness.context import SessionContext
from litellm.harness.options import ToolLoopOptions
from litellm.harness.types import Capabilities, Harness
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
from litellm.llms.base_llm.harness.utils import gateway_headers
from litellm.types.utils import ChatCompletionToolParam
TOOL_LOOP_MAX_MODEL_CALLS: Final = 100
_ANNOTATIONS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_OBJECT_ADAPTER: Final = TypeAdapter(object)
_MODEL_FACTORY: Final[Callable[..., type[BaseModel]]] = create_model
@dataclass(frozen=True, slots=True)
class FunctionTool:
name: str
fn: Callable[..., object]
args_model: type[BaseModel]
spec: ChatCompletionToolParam
def _field_definition(
parameter: inspect.Parameter,
annotations: Mapping[str, object],
) -> tuple[object, object]:
annotation: Final = annotations.get(parameter.name, object)
raw_default: Final[object] = parameter.default # pyright: ignore[reportAny] # inspect exposes defaults as Any
if raw_default is inspect.Parameter.empty:
return annotation, ...
default: Final = _OBJECT_ADAPTER.validate_python(raw_default)
return annotation, default
def function_tool(fn: Callable[..., object]) -> FunctionTool:
signature: Final = inspect.signature(fn)
parameters: Final = tuple(signature.parameters.values())
raw_annotations: Final[object] = inspect.get_annotations(fn, eval_str=True)
annotations: Final = _ANNOTATIONS_ADAPTER.validate_python(raw_annotations)
if any(
parameter.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD) for parameter in parameters
):
raise ValueError(f"Tool {fn.__name__} cannot use variadic parameters")
fields: Final = MappingProxyType(
{parameter.name: _field_definition(parameter, annotations) for parameter in parameters}
)
args_model: Final[type[BaseModel]] = _MODEL_FACTORY(
f"{fn.__name__}_args",
__config__=ConfigDict(extra="forbid"),
**fields, # pyright: ignore[reportCallIssue, reportArgumentType] # Pydantic creates fields dynamically
)
spec: Final[ChatCompletionToolParam] = {
"type": "function",
"function": {
"name": fn.__name__,
"description": inspect.getdoc(fn) or "",
"parameters": args_model.model_json_schema(),
},
}
return FunctionTool(name=fn.__name__, fn=fn, args_model=args_model, spec=spec)
def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]:
if not ctx.model:
raise ValueError("Harness.TOOL_LOOP needs model=")
if ctx.gateway is not None:
return {
"model": f"litellm_proxy/{ctx.model}",
"api_base": ctx.gateway.api_base,
"api_key": ctx.gateway.api_key,
"extra_headers": gateway_headers(ctx),
}
return {
"model": ctx.model,
**({"api_key": ctx.api_key} if ctx.api_key is not None else {}),
**({"api_base": ctx.api_base} if ctx.api_base is not None else {}),
}
def completion_kwargs(ctx: SessionContext) -> Mapping[str, object]:
options: Final = ToolLoopHarnessConfig().get_options(ctx)
routing: Final = _routing_kwargs(ctx)
kwargs: Final[Mapping[str, object]] = MappingProxyType({**options.completion_kwargs, **routing})
if ctx.output is None:
return kwargs
return {**kwargs, "response_format": ctx.output}
class ToolLoopHarnessConfig(BaseHarnessConfig[ToolLoopOptions]):
harness = Harness.TOOL_LOOP
options_type = ToolLoopOptions
uses_model_endpoint = False
capabilities = Capabilities(
structured_output=True,
tool_approval=True,
tool_filtering=False,
history=True,
custom_tools=True,
skills=False,
resume=False,
permission_modes=frozenset({"ask", "full"}),
)
def validate_environment(self, ctx: SessionContext) -> None:
self.get_options(ctx)
if not ctx.model:
raise ValueError("Harness.TOOL_LOOP needs model=")

View file

@ -0,0 +1,10 @@
from typing import Final
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
TYPESAFE_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
default_api_base_value="https://api.typesafe.ai",
path="/v1/systemone",
api_key_env=("TYPESAFE_API_KEY",),
api_base_env="TYPESAFE_API_BASE",
)

View file

@ -11542,9 +11542,11 @@
},
"azure_ai/flux.2-pro": {
"litellm_provider": "azure_ai",
"max_input_tokens": 32000,
"max_tokens": 32000,
"mode": "image_generation",
"output_cost_per_image": 0.04,
"source": "https://ai.azure.com/explore/models/flux.2-pro/version/1/registry/azureml-blackforestlabs",
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure",
"supported_endpoints": [
"/v1/images/generations"
]
@ -15612,6 +15614,38 @@
"prompt_cache_min_tokens": 1024,
"source": "https://platform.claude.com/docs/en/about-claude/pricing"
},
"cloudflare/clef": {
"input_cost_per_token": 2.4e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/clef-flash": {
"input_cost_per_token": 9e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/@cf/cloudflare/clef": {
"input_cost_per_token": 2.4e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/@cf/cloudflare/clef-flash": {
"input_cost_per_token": 9e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/@cf/meta/llama-2-7b-chat-fp16": {
"input_cost_per_token": 1.923e-06,
"litellm_provider": "cloudflare",
@ -44421,6 +44455,14 @@
"mode": "chat",
"output_cost_per_token": 2.8e-07
},
"perplexity/pplx-decider-v1-27b": {
"input_cost_per_token": 4e-08,
"litellm_provider": "perplexity",
"max_input_tokens": 262144,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://docs.perplexity.ai/api-reference/decisions-post"
},
"perplexity/sonar": {
"input_cost_per_token": 1e-06,
"litellm_provider": "perplexity",
@ -72800,6 +72842,16 @@
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"strands_decider/strands-decider-2B-hobson-v19": {
"input_cost_per_token": 0.0,
"litellm_provider": "strands_decider",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://huggingface.co/StrandsAgents/strands-decider-2B-hobson-v19",
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"typesafe/jev-1.13.0": {
"input_cost_per_token": 4.2e-08,
"litellm_provider": "typesafe",
@ -76881,6 +76933,7 @@
"source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_prompt_cache_breakpoint": false,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
@ -76902,6 +76955,7 @@
"source": "https://aws.amazon.com/bedrock/pricing/",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_prompt_cache_breakpoint": false,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
@ -76923,6 +76977,7 @@
"source": "https://aws.amazon.com/bedrock/pricing/",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_prompt_cache_breakpoint": false,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,

View file

@ -16,6 +16,7 @@ from collections.abc import Set as AbstractSet
from contextlib import asynccontextmanager
from dataclasses import dataclass, field
from functools import partial
from itertools import chain
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
@ -268,6 +269,11 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
module_path="litellm.proxy.openai_evals_endpoints.endpoints",
path_prefixes=("/v1/evals", "/evals"),
),
LazyFeature(
name="decisions",
module_path="litellm.proxy.decisions_endpoints.endpoints",
path_prefixes=("/v1/decisions", "/decisions"),
),
LazyFeature(
name="claude_code_marketplace",
module_path="litellm.proxy.anthropic_endpoints.claude_code_endpoints",
@ -378,6 +384,10 @@ def _lazy_slots(app: "FastAPI") -> Mapping[str, BaseRoute | None]:
return app.state.lazy_slots if hasattr(app.state, "lazy_slots") else MappingProxyType({})
def _lazy_routes(app: "FastAPI") -> Mapping[str, tuple[BaseRoute, ...]]:
return app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({})
def reserve_lazy_slot(app: "FastAPI", name: str, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> None:
"""Record the route the feature's router used to be included after, so its routes
are spliced back in there once it loads and keep the same precedence. Anchoring on
@ -474,11 +484,8 @@ def _lazy_lock(app: "FastAPI", module_path: str) -> asyncio.Lock:
def _register_feature(app: "FastAPI", feat: LazyFeature, module: object, features: tuple[LazyFeature, ...]) -> None:
before: Final = len(app.router.routes)
feat.register_fn(app, module)
previous: Final[Mapping[str, tuple[BaseRoute, ...]]] = (
app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({})
)
lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]]] = MappingProxyType(
{**previous, feat.module_path: tuple(app.router.routes[before:])}
{**_lazy_routes(app), feat.module_path: tuple(app.router.routes[before:])}
)
app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added
app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table
@ -543,11 +550,8 @@ def _register_all_on_startup(inner: "Lifespan[FastAPI]", features: tuple[LazyFea
def _restore_registry_order(app: "FastAPI", features: tuple[LazyFeature, ...]) -> None:
present: Final = frozenset(id(route) for route in app.router.routes)
registered: Final[Mapping[str, tuple[BaseRoute, ...]]] = (
app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({})
)
still_routed: Final = MappingProxyType(
{module_path: tuple(r for r in routes if id(r) in present) for module_path, routes in registered.items()}
{module_path: tuple(r for r in routes if id(r) in present) for module_path, routes in _lazy_routes(app).items()}
)
app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table
_in_registry_order(app.router.routes, still_routed, features, _lazy_slots(app))
@ -599,6 +603,13 @@ def _make_warmup_router(app: "FastAPI", features: tuple[LazyFeature, ...] = LAZY
return router
def lazy_owned_routes(app: "FastAPI") -> frozenset[int]:
"""ids of the routes lazy features have registered on this app. A route added later at
one of their paths (a config pass-through at /v1/decisions) goes ahead of them, the
precedence lazy mode gives it when the feature has not loaded by the time the config is read."""
return frozenset(id(route) for route in chain.from_iterable(_lazy_routes(app).values()))
def loaded_lazy_modules(app: "FastAPI") -> frozenset[str]:
"""The set of lazy feature modules whose routers are actually registered
on this app (tracked by _install), empty until a feature loads or eager startup runs.

View file

@ -9492,6 +9492,61 @@
}
}
},
"decisions": {
"components": {
"schemas": {}
},
"paths": {
"/decisions": {
"post": {
"operationId": "decisions_decisions_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Decisions",
"tags": [
"decisions"
]
}
},
"/v1/decisions": {
"post": {
"operationId": "decisions_v1_decisions_post",
"responses": {
"200": {
"content": {
"application/json": {
"schema": {}
}
},
"description": "Successful Response"
}
},
"security": [
{
"APIKeyHeader": []
}
],
"summary": "Decisions",
"tags": [
"decisions"
]
}
}
}
},
"evals": {
"components": {
"schemas": {
@ -34854,7 +34909,6 @@
},
"name": {
"default": "Lens worker",
"maxLength": 100,
"minLength": 1,
"title": "Name",
"type": "string"

View file

@ -479,6 +479,8 @@ class LiteLLMRoutes(enum.Enum):
"/v1/search",
"/search/{search_tool_name}",
"/v1/search/{search_tool_name}",
"/decisions",
"/v1/decisions",
# OCR
"/ocr",
"/v1/ocr",

View file

@ -29,6 +29,7 @@ _MANAGED_MODEL_ROUTES: Final = frozenset(
"audio/speech",
"moderations",
"rerank",
"decisions",
"ocr",
),
)
@ -72,6 +73,7 @@ _MODEL_ROUTE_KINDS: Final[
"/audio/transcriptions": "moderation",
"/audio/speech": "speech",
"/rerank": "body",
"/decisions": "body",
"/messages/count_tokens": "body",
":countTokens": "path",
}

View file

@ -176,6 +176,7 @@ ProxyRouteType: TypeAlias = Literal[
"avector_store_file_delete",
"aocr",
"asearch",
"adecisions",
"avideo_generation",
"avideo_list",
"avideo_status",
@ -1961,6 +1962,7 @@ class ProxyBaseLLMRequestProcessing:
"avector_store_file_delete",
"aocr",
"asearch",
"adecisions",
"avideo_generation",
"avideo_list",
"avideo_status",

View file

@ -0,0 +1 @@
__all__ = ()

View file

@ -0,0 +1,102 @@
from typing import Annotated, Final
from fastapi import APIRouter, Depends, Request, Response
from fastapi.responses import ORJSONResponse # pyright: ignore[reportDeprecated] # required endpoint contract
from pydantic import TypeAdapter, ValidationError
from litellm.exceptions import BadRequestError
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.types.decisions import DecisionsRequestBody
router: Final = APIRouter()
_REQUEST_DATA_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
_DECISIONS_REQUEST_BODY_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody)
_GENERAL_SETTINGS_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
_OPTIONAL_STRING_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
_OPTIONAL_FLOAT_ADAPTER: Final[TypeAdapter[float | None]] = TypeAdapter(float | None)
@router.post(
"/v1/decisions",
dependencies=[Depends(user_api_key_auth)],
response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract
tags=["decisions"],
)
@router.post(
"/decisions",
dependencies=[Depends(user_api_key_auth)],
response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract
tags=["decisions"],
)
async def decisions(
request: Request,
fastapi_response: Response,
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
):
from litellm.proxy.proxy_server import (
general_settings as proxy_general_settings,
)
from litellm.proxy.proxy_server import (
llm_router,
proxy_config,
proxy_logging_obj,
user_max_tokens,
user_request_timeout,
version,
)
from litellm.proxy.proxy_server import (
user_api_base as proxy_user_api_base,
)
from litellm.proxy.proxy_server import (
user_model as proxy_user_model,
)
from litellm.proxy.proxy_server import (
user_temperature as proxy_user_temperature,
)
data: Final = _REQUEST_DATA_ADAPTER.validate_json(await request.body())
general_settings: Final = _GENERAL_SETTINGS_ADAPTER.validate_python(proxy_general_settings)
user_api_base: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_api_base)
user_model: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_model)
user_temperature: Final = _OPTIONAL_FLOAT_ADAPTER.validate_python(proxy_user_temperature)
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
try:
_DECISIONS_REQUEST_BODY_ADAPTER.validate_python(data)
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="adecisions",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=None,
model=None,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except ValidationError as error:
bad_request_error: Final = BadRequestError(
message=f"Invalid Decisions request: {error}",
model=str(data.get("model", "")),
llm_provider="",
)
raise await processor._handle_llm_api_exception(
e=bad_request_error,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
except Exception as error:
raise await processor._handle_llm_api_exception(
e=error,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)

View file

@ -7,7 +7,7 @@ from itertools import chain, islice
from types import MappingProxyType
from typing import Final, Literal, TypeAlias, TypeVar
from pydantic import Field, ValidationError
from pydantic import Field, TypeAdapter, ValidationError
from .models import (
Claim,
@ -31,8 +31,8 @@ from .trace_store import TraceStore, overview_content, trace_store
class Observation(Record):
check_id: str
kind: Literal["issue", "pattern"] = "issue"
summary: str = Field(max_length=2000)
evidence: tuple[Evidence, ...] = Field(default=(), max_length=6)
summary: str
evidence: tuple[Evidence, ...] = Field(default=())
class Extraction(Record):
@ -47,14 +47,14 @@ class SpanRead(Record):
class TraceReview(Extraction):
feedback_page: int | None = Field(default=None, ge=0)
reads: tuple[SpanRead, ...] = Field(default=(), max_length=2)
reads: tuple[SpanRead, ...] = Field(default=())
class Candidate(Record):
check_id: str
kind: Literal["issue", "pattern"] = "issue"
title: str = Field(max_length=160)
hypothesis: str = Field(max_length=2000)
title: str
hypothesis: str
execution_ids: tuple[str, ...]
existing_finding_id: str | None = None
@ -64,7 +64,7 @@ class Clusters(Record):
class Decision(Record):
action: Literal["read", "observations", "catalog", "feedback", "submit", "inconclusive"]
action: Literal["read", "evidence", "observations", "catalog", "feedback", "submit", "inconclusive"]
page: int = Field(default=0, ge=0)
execution_id: str | None = None
cursor: str = ""
@ -83,11 +83,13 @@ class Examined(Record):
parts: tuple[TracePart, ...]
partial: bool
cannot_assess: bool
error: str = ""
class Investigation(Record):
finding: FindingDraft | None
parts: tuple[TracePart, ...]
error: str = ""
ModelCall: TypeAlias = Callable[[ModelRequest], Awaitable[ModelResult]]
@ -98,6 +100,28 @@ ReportProgress: TypeAlias = Callable[[str, Coverage], Awaitable[None]]
ResponseT = TypeVar("ResponseT", bound=Record)
class ValidationIssue(Record):
type: str
loc: tuple[str | int, ...]
msg: str
def validation_details(error: ValidationError) -> str:
issues: Final = TypeAdapter(tuple[ValidationIssue, ...]).validate_json(
error.json(include_input=False, include_context=False, include_url=False)
)
return "\n".join(
f"{'.'.join(str(part) for part in issue.loc) or '$'}: {issue.msg} [{issue.type}]"
if issue.type != "extra_forbidden"
else "Unexpected field: Extra inputs are not permitted [extra_forbidden]"
for issue in issues
)
class AnalysisResponseError(ValueError):
pass
async def structured_response(
request: ModelRequest,
schema: type[ResponseT],
@ -107,6 +131,8 @@ async def structured_response(
response: Final = await model(request)
try:
parsed: Final = schema.model_validate_json(response.content)
if response.finish_reason:
raise ValueError(f"Model did not finish its response (finish_reason={response.finish_reason})")
invalid: Final = validate(parsed)
if invalid:
raise ValueError(invalid)
@ -124,11 +150,34 @@ async def structured_response(
}
)
)
corrected: Final = schema.model_validate_json((await model(repair)).content)
remaining: Final = validate(corrected)
if remaining:
raise ValueError(remaining)
return corrected
repaired: Final = await model(repair)
try:
corrected: Final = schema.model_validate_json(repaired.content)
if repaired.finish_reason:
raise ValueError(f"Model did not finish its response (finish_reason={repaired.finish_reason})")
remaining: Final = validate(corrected)
if remaining:
raise ValueError(remaining)
return corrected
except ValueError as error:
stage: Final = MappingProxyType(
{
"extract": "Reading executions",
"cluster": "Grouping observations",
"investigate": "Checking original evidence",
}
)[request.purpose]
detail: Final = validation_details(error) if isinstance(error, ValidationError) else str(error)
stopped: Final = (
" Model output was truncated (finish_reason=length)."
if repaired.finish_reason == "length"
else " Model output was blocked (finish_reason=content_filter)."
if repaired.finish_reason == "content_filter"
else ""
)
raise AnalysisResponseError(
f"{stage} failed: {schema.__name__} response invalid after 2 attempts.{stopped}\n{detail}"
) from error
def evidence_valid(evidence: Evidence, parts: tuple[TracePart, ...]) -> bool:
@ -203,8 +252,15 @@ async def extract(claim: Claim, execution: Execution, read: ReadContent, model:
with trace_store() as store:
try:
return await extract_stored(claim, execution, read, model, store)
except ValidationError:
return Examined(execution=execution, observations=(), parts=(), partial=True, cannot_assess=True)
except (ValidationError, AnalysisResponseError) as error:
return Examined(
execution=execution,
observations=(),
parts=(),
partial=True,
cannot_assess=True,
error=validation_details(error) if isinstance(error, ValidationError) else str(error),
)
async def extract_stored(
@ -254,7 +310,7 @@ async def extract_stored(
for p in (first_root,)
if p is not None
),
"read_evidence": tuple(p.model_dump() for p in additional[-2:]),
"read_evidence": tuple(p.model_dump() for p in additional),
"previous_observations": tuple(o.model_dump() for o in previous.observations),
"completed_read_count": len(reads),
"last_completed_read": reads[-1].model_dump() if reads else None,
@ -294,7 +350,7 @@ async def extract_stored(
previous = response
continue
fetched = tuple([parts async for parts in concurrent_results(requested, fetch)])
if not any(p.content and p not in additional for p in chain.from_iterable(fetched)):
if not any(p.content for p in chain.from_iterable(fetched)):
must_decide = True
previous = response
continue
@ -355,8 +411,12 @@ async def investigate(
with trace_store() as store:
try:
return await investigate_stored(claim, candidate, examined, read, model, store)
except ValidationError:
return Investigation(finding=None, parts=())
except (ValidationError, AnalysisResponseError) as error:
return Investigation(
finding=None,
parts=(),
error=validation_details(error) if isinstance(error, ValidationError) else str(error),
)
async def investigate_stored(
@ -371,6 +431,8 @@ async def investigate_stored(
navigation: ExecutionContent | None = None # rebind-ok: last fetched page
reads: tuple[Decision, ...] = () # rebind-ok: track completed tool requests to detect loops
observation_page = 0 # rebind-ok: model controls navigation through observations
evidence_page = 0 # rebind-ok: navigate all content in the fetched evidence batch
evidence_seen = frozenset((0,)) # rebind-ok: reset navigation history when evidence changes
catalog_page = 0 # rebind-ok: model controls navigation through the run catalog
feedback_page = 0 # rebind-ok: navigate bounded prior finding pages
feedback: Final = feedback_pages(claim, candidate.check_id)
@ -381,6 +443,7 @@ async def investigate_stored(
navigation: ExecutionContent | None,
reads: tuple[Decision, ...],
observation_page: int,
evidence_page: int,
catalog_page: int,
feedback_page: int,
stalled: bool,
@ -411,7 +474,7 @@ async def investigate_stored(
)
)
bounded: Final = partition_content(prioritized, 30000)
evidence: Final = bounded[0] if bounded else ()
evidence: Final = bounded[evidence_page] if evidence_page < len(bounded) else ()
catalog_batches: Final = partition_items(
(*relevant, *(item for item in examined if item not in relevant)),
lambda item: len(item.execution.model_dump_json()),
@ -452,13 +515,13 @@ async def investigate_stored(
"feedback_page": feedback_page,
"feedback_pages": len(feedback),
"evidence": tuple(p.model_dump() for p in evidence),
"evidence_page": evidence_page,
"evidence_pages": len(bounded),
"must_decide": stalled,
"last_read": navigation.model_dump(exclude=MappingProxyType({"parts": True})) if navigation else None,
},
ensure_ascii=False,
)
if len(prompt) > 100000:
return Investigation(finding=None, parts=evidence)
request: Final = ModelRequest(purpose="investigate", prompt=prompt)
decision: Final = await investigation_decision(request, model, 1 if stalled else 2)
if decision.action == "submit" and decision.finding:
@ -479,11 +542,12 @@ async def investigate_stored(
)
):
return Investigation(finding=finding, parts=evidence)
if stalled or decision.action not in ("read", "observations", "catalog", "feedback"):
if stalled or decision.action not in ("read", "evidence", "observations", "catalog", "feedback"):
return Investigation(finding=None, parts=evidence)
page_count: Final = MappingProxyType(
{
"observations": len(supporting_batches),
"evidence": len(bounded),
"catalog": len(catalog_batches),
"feedback": len(feedback),
}
@ -497,13 +561,20 @@ async def investigate_stored(
)
while True:
step_result = await decide(
additional, navigation, reads, observation_page, catalog_page, feedback_page, stalled
additional, navigation, reads, observation_page, evidence_page, catalog_page, feedback_page, stalled
)
if isinstance(step_result, Decision) and step_result.action == "inconclusive":
stalled = True
continue
if isinstance(step_result, Investigation):
return step_result
if step_result.action == "evidence":
if step_result.page in evidence_seen:
stalled = True
else:
evidence_page = step_result.page
evidence_seen = evidence_seen | frozenset((evidence_page,))
continue
if any(
(r.action, r.execution_id, r.cursor, r.offset, r.page)
== (step_result.action, step_result.execution_id, step_result.cursor, step_result.offset, step_result.page)
@ -514,16 +585,20 @@ async def investigate_stored(
reads = (*reads, step_result)
if step_result.action == "observations":
observation_page = step_result.page
evidence_page = 0
evidence_seen = frozenset((0,))
elif step_result.action == "catalog":
catalog_page = step_result.page
elif step_result.action == "feedback":
feedback_page = step_result.page
elif any(e.execution.id == step_result.execution_id for e in examined):
navigation = await read(step_result.execution_id or "", step_result.cursor, step_result.offset)
if not any(p.content and p not in additional for p in navigation.parts):
if not any(p.content for p in navigation.parts):
stalled = True
store.add_reads(navigation.parts)
additional = navigation.parts
evidence_page = 0
evidence_seen = frozenset((0,))
else:
return Investigation(finding=None, parts=additional)
@ -619,7 +694,11 @@ async def _analyze_sample(
await progress("Grouping observations", coverage)
observations: Final = tuple(chain.from_iterable(item.observations for item in examined))
if not observations:
return Result(coverage=coverage, assessments=assessments)
return Result(
coverage=coverage,
assessments=assessments,
error="\n\n".join(dict.fromkeys(item.error for item in examined if item.error)),
)
batches: Final = observation_batches(observations)
grouping: Final = coverage.model_copy(update=MappingProxyType({"grouping_batches": len(batches)}))
clusters: Final = await cluster_batches(batches, limited_model, progress, grouping)
@ -638,6 +717,7 @@ async def _analyze_sample(
return Result(
findings=tuple(item.finding for item in investigated if item.finding is not None),
assessments=assessments,
error="\n\n".join(dict.fromkeys(item.error for item in (*examined, *investigated) if item.error)),
coverage=investigating.model_copy(
update=MappingProxyType(
{"investigated": len(candidates), "inconclusive": sum(item.finding is None for item in investigated)}
@ -657,7 +737,7 @@ async def cluster_batches(
Candidate(
check_id=o.check_id,
kind=o.kind,
title=o.summary[:160],
title=o.summary,
hypothesis=f"{o.kind}: {o.summary}",
execution_ids=tuple(sorted(frozenset(e.execution_id for e in o.evidence))),
)

View file

@ -1,4 +1,5 @@
from collections.abc import Awaitable, Callable, Mapping
from collections.abc import Callable, Mapping
from contextlib import AbstractAsyncContextManager
from typing import Final
import orjson
@ -32,7 +33,7 @@ async def validate_key(key_id: str | None) -> UserAPIKeyAuth | None:
async def complete(
key_id: str, data: dict[str, object], reserve: Callable[[], Awaitable[None]], incoming: Request
key_id: str, data: dict[str, object], reserve: Callable[[], AbstractAsyncContextManager[None]], incoming: Request
) -> tuple[ModelResponse, float | None]:
from litellm.proxy import proxy_server
from litellm.proxy.proxy_server import llm_router, proxy_config, proxy_logging_obj, version
@ -67,23 +68,23 @@ async def complete(
)
try:
auth: Final = await authorize_internal_virtual_key(key_id, request, data)
await reserve()
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
fastapi_response: Final = Response()
try:
response: Final = TypeAdapter(ModelResponse).validate_python(
await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=auth,
route_type="acompletion",
proxy_logging_obj=proxy_logging_obj,
general_settings=TypeAdapter(dict[str, object]).validate_python(proxy_server.general_settings), # pyright: ignore[reportUnknownMemberType] # Validate the legacy untyped config at the request boundary
proxy_config=proxy_config,
llm_router=llm_router,
version=version,
async with reserve():
response: Final = TypeAdapter(ModelResponse).validate_python(
await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=auth,
route_type="acompletion",
proxy_logging_obj=proxy_logging_obj,
general_settings=TypeAdapter(dict[str, object]).validate_python(proxy_server.general_settings), # pyright: ignore[reportUnknownMemberType] # Validate the legacy untyped config at the request boundary
proxy_config=proxy_config,
llm_router=llm_router,
version=version,
)
)
)
billed: Final = fastapi_response.headers.get("x-litellm-response-cost")
return response, float(billed) if billed not in (None, "", "None") else None
except Exception as exc:

View file

@ -7,11 +7,12 @@ from types import MappingProxyType
from typing import Annotated, Final, TypeAlias
from uuid import uuid4
from fastapi import APIRouter, Depends, HTTPException, Query, Request
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import AwareDatetime, BaseModel, Field
from litellm.proxy._types import LitellmUserRoles, ModelAccessDeniedProxyException, UserAPIKeyAuth
from litellm.litellm_core_utils.secret_redaction import redact_internal_details
from litellm.proxy._types import LitellmUserRoles, ModelAccessDeniedProxyException, ProxyException, UserAPIKeyAuth
from litellm.proxy.auth.auth_checks import can_key_call_model
from litellm.proxy.auth.resolvers.exceptions import KeyNotFoundError
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -28,6 +29,7 @@ from litellm.proxy.lens.models import (
Lens,
LensList,
LensSettings,
LookbackHours,
ModelRequest,
ModelResult,
Progress,
@ -324,18 +326,23 @@ class Preview(BaseModel):
as_of: AwareDatetime | None = None
offset: int = Field(default=0, ge=0)
settings: LensSettings
lookback_hours: int = Field(default=24, ge=1, le=8760)
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)
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)
end: Final = int((now - timedelta(minutes=2)).timestamp() * 1000)
except (OverflowError, ValueError) as error:
raise HTTPException(422, "Preview window exceeds the supported calendar range") from error
return await source_reader(storage).sample(
user_scope(auth),
body.settings,
int((now - timedelta(hours=body.lookback_hours)).timestamp() * 1000),
int((now - timedelta(minutes=2)).timestamp() * 1000),
start,
end,
offset=body.offset,
preview=True,
)
@ -346,7 +353,7 @@ class WorkerBilling(BaseModel):
class WorkerName(WorkerBilling):
name: str = Field(default="Lens worker", min_length=1, max_length=100)
name: str = Field(default="Lens worker", min_length=1)
@router.post("/workers/register", response_model=WorkerCreated)
@ -395,7 +402,7 @@ async def revoke_worker(worker_id: str, auth: Auth) -> bool:
@router.post("/worker/claim", response_model=Claim | None)
async def claim(worker: WorkerAuth, protocol_version: int = 1) -> Claim | None:
if protocol_version != 2:
if protocol_version not in (2, 3):
raise HTTPException(409, "Upgrade the Lens worker using the current Connect worker command")
if worker.analysis_key_id is None:
raise HTTPException(409, "Assign an analysis key to this worker in Lens setup")
@ -491,12 +498,31 @@ async def content(
return await source_reader(storage).content(lens.scope, execution, cursor, offset)
def model_failure(error: HTTPException | ProxyException) -> HTTPException:
if isinstance(error, ProxyException):
status: Final = int(error.code) if error.code.isdigit() else 500
return HTTPException(status, {"lens_error": redact_internal_details(error.message)}, headers=error.headers)
if isinstance(error.detail, str):
return HTTPException(
error.status_code, {"lens_error": redact_internal_details(error.detail)}, headers=error.headers
)
return error
@router.post("/worker/{lens_id}/{job_id}/model", response_model=ModelResult)
async def model(lens_id: str, job_id: str, body: ModelRequest, worker: WorkerAuth, request: Request) -> ModelResult:
async def model(
lens_id: str, job_id: str, body: ModelRequest, worker: WorkerAuth, request: Request, response: Response
) -> ModelResult:
from litellm.proxy.lens.inference import analyze
lens, job = await assigned(lens_id, job_id, worker)
return await analyze(repository(), lens, job, worker, body, request)
try:
completion: Final = await analyze(repository(), lens, job, worker, body, request)
except (ProxyException, HTTPException) as error:
raise model_failure(error) from error
if completion.finish_reason:
response.headers["x-litellm-lens-finish-reason"] = completion.finish_reason
return completion
@router.post("/worker/{lens_id}/{job_id}/result", response_model=Lens)

View file

@ -1,3 +1,5 @@
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Final
@ -9,6 +11,7 @@ import litellm
from litellm.exceptions import ModelNotMappedError
from litellm.integrations.clickhouse.context import lens_analysis
from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy
from litellm.litellm_core_utils.token_counter import get_modified_max_tokens
from litellm.proxy.lens.billing import complete, validate_key
from litellm.proxy.lens.models import Job, Lens, ModelRequest, ModelResult, Worker
from litellm.proxy.lens.repository import LensRepository
@ -21,11 +24,19 @@ class DeploymentParams(BaseModel):
model: str
input_cost_per_token: float | None = None
output_cost_per_token: float | None = None
max_tokens: int | None = Field(default=None, gt=0)
max_completion_tokens: int | None = Field(default=None, gt=0)
class ModelCapacity(BaseModel):
model_config = ConfigDict(extra="ignore")
max_output_tokens: int | None = Field(default=None, gt=0)
class Deployment(BaseModel):
model_config = ConfigDict(extra="ignore")
litellm_params: DeploymentParams
model_info: ModelCapacity = ModelCapacity()
class Message(BaseModel):
@ -36,6 +47,7 @@ class Message(BaseModel):
class Choice(BaseModel):
model_config = ConfigDict(extra="ignore")
message: Message
finish_reason: str | None = None
class Completion(BaseModel):
@ -88,6 +100,36 @@ def deployment_prices(deployment: Deployment) -> Prices:
) from exc
def catalog_capacity(model: str) -> ModelCapacity:
try:
return ModelCapacity.model_validate(litellm.get_model_info(model=model))
except (ModelNotMappedError, ValueError):
return ModelCapacity()
def output_tokens(deployment: Deployment, prompt: str | None = None) -> int:
params: Final = deployment.litellm_params
configured: Final = params.max_completion_tokens or params.max_tokens or deployment.model_info.max_output_tokens
capacity: Final = configured or catalog_capacity(params.model).max_output_tokens
if capacity is None:
raise HTTPException(
400,
f"Output capacity is unknown for {params.model}. Set model_info.max_output_tokens to the model's "
"supported output capacity or configure max_tokens on its deployment.",
)
if prompt is None:
return capacity
adjusted: Final = get_modified_max_tokens(
model=params.model,
base_model=params.model,
messages=[{"role": "system", "content": _SYSTEM}, {"role": "user", "content": prompt}],
user_max_tokens=capacity,
buffer_perc=0,
buffer_num=0,
)
return adjusted if adjusted is not None else capacity
def quote(deployments: tuple[Deployment, ...], prompt: str) -> float:
prices: Final = tuple(deployment_prices(d) for d in deployments)
input_rate: Final = max(
@ -102,7 +144,15 @@ def quote(deployments: tuple[Deployment, ...], prompt: str) -> float:
)
for p in prices
)
return ((len((prompt + _SYSTEM).encode()) + 1024) * input_rate + 4096 * output_rate) * 2
output: Final = min(output_tokens(d, prompt) for d in deployments)
input_tokens: Final = max(
litellm.token_counter(
model=d.litellm_params.model,
messages=[{"role": "system", "content": _SYSTEM}, {"role": "user", "content": prompt}],
)
for d in deployments
)
return input_tokens * input_rate + output * output_rate
async def analyze(
@ -142,37 +192,7 @@ async def analyze(
current, active.model_copy(update=MappingProxyType({"cost": active.cost + estimate}))
).model_copy(update=MappingProxyType({"spent": current.spent + estimate}))
async def reserve_budget() -> None:
if await repo.update(lens.id, reserve) is None:
raise HTTPException(409, "Could not reserve analysis budget")
data: Final[dict[str, object]] = { # mutable-ok: proxy processing enriches request data
"model": job.settings.model,
"messages": [
{"role": "system", "content": _SYSTEM},
{"role": "user", "content": body.prompt},
],
"max_tokens": 4096,
"stream": False,
"timeout": 120,
"num_retries": 0,
"disable_fallbacks": True,
"response_format": {"type": "json_object"},
"metadata": {
"tags": ["litellm-lens"],
"lens_id": lens.id,
"lens_run_id": job.id,
"lens_worker_id": worker.id,
"user_api_key_team_id": team_id,
},
}
with lens_analysis(), inherit_message_logging_privacy(True):
response, billed_cost = await complete(worker.analysis_key_id, data, reserve_budget, request)
parsed: Final = Completion.model_validate_json(response.model_dump_json())
cost: Final = billed_cost if billed_cost is not None else completion_charge(deployments, response, estimate)
def settle(e: Lens) -> Lens:
def settle(e: Lens, cost: float) -> Lens:
charged: Final = next((j for j in e.jobs if j.id == job.id), None)
adjusted: Final = (
e.model_copy(update=MappingProxyType({"spent": max(0, e.spent - estimate + cost)}))
@ -187,8 +207,50 @@ async def analyze(
else adjusted
)
await repo.update(lens.id, settle)
return ModelResult(content=parsed.choices[0].message.content or "{}", cost=cost)
@asynccontextmanager
async def reserve_budget() -> AsyncIterator[None]:
if await repo.update(lens.id, reserve) is None:
raise HTTPException(409, "Could not reserve analysis budget")
try:
yield
except BaseException:
await repo.update(lens.id, lambda e: settle(e, 0))
raise
data: Final[dict[str, object]] = { # mutable-ok: proxy processing enriches request data
"model": job.settings.model,
"messages": [
{"role": "system", "content": _SYSTEM},
{"role": "user", "content": body.prompt},
],
"max_tokens": min(output_tokens(d, body.prompt) for d in deployments),
"stream": False,
"num_retries": 0,
"disable_fallbacks": True,
"response_format": {"type": "json_object"},
"metadata": {
"tags": ["litellm-lens"],
"lens_id": lens.id,
"lens_run_id": job.id,
"lens_worker_id": worker.id,
"user_api_key_team_id": team_id,
},
}
with lens_analysis(), inherit_message_logging_privacy(True):
response, billed_cost = await complete(worker.analysis_key_id, data, reserve_budget, request)
cost: Final = billed_cost if billed_cost is not None else completion_charge(deployments, response, estimate)
await repo.update(lens.id, lambda e: settle(e, cost))
parsed: Final = Completion.model_validate_json(response.model_dump_json())
choice: Final = parsed.choices[0]
return ModelResult(
content=choice.message.content or "",
cost=cost,
finish_reason="length"
if choice.finish_reason == "length"
else ("content_filter" if choice.finish_reason == "content_filter" else None),
)
def completion_charge(deployments: tuple[Deployment, ...], response: ModelResponse, estimate: float) -> float:

View file

@ -1,7 +1,27 @@
from datetime import datetime
from typing import Final, Literal
from datetime import datetime, timedelta, timezone
from typing import Annotated, Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, Field, model_validator
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, model_validator
def calendar_lookback(hours: int) -> int:
try:
datetime.now(timezone.utc) - timedelta(hours=hours)
except OverflowError as error:
raise ValueError("Lookback exceeds the supported calendar range") from error
return hours
def calendar_interval(minutes: int) -> int:
try:
datetime.now(timezone.utc) + timedelta(minutes=minutes)
except OverflowError as error:
raise ValueError("Interval exceeds the supported calendar range") from error
return minutes
LookbackHours: TypeAlias = Annotated[int, Field(ge=1), AfterValidator(calendar_lookback)]
IntervalMinutes: TypeAlias = Annotated[int, Field(ge=1), AfterValidator(calendar_interval)]
class Record(BaseModel):
@ -15,34 +35,34 @@ class Scope(Record):
class MetadataFilter(Record):
key: str = Field(min_length=1, max_length=200)
value: str = Field(min_length=1, max_length=500)
key: str = Field(min_length=1)
value: str = Field(min_length=1)
class Check(Record):
id: str = Field(min_length=1, max_length=80)
instruction: str = Field(min_length=3, max_length=3000)
id: str = Field(min_length=1)
instruction: str = Field(min_length=3)
enabled: bool = True
class LensSettings(Record):
name: str = Field(min_length=1, max_length=100)
context: str = Field(default="", max_length=6000)
name: str = Field(min_length=1)
context: str = Field(default="")
source: Literal["traces", "requests", "both"] = "traces"
lookback_hours: int = Field(default=24, ge=1, le=8760)
service: str = Field(default="", max_length=200)
agent_name: str = Field(default="", max_length=200)
filters: tuple[MetadataFilter, ...] = Field(default=(), max_length=8)
lookback_hours: LookbackHours = 24
service: str = Field(default="")
agent_name: str = Field(default="")
filters: tuple[MetadataFilter, ...] = Field(default=())
checks: tuple[Check, ...] = ()
model: str = Field(min_length=1, max_length=200)
model: str = Field(min_length=1)
enabled: bool = True
interval_minutes: int = Field(default=15, ge=1, le=10080)
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, le=100000, allow_inf_nan=False)
monthly_budget: float = Field(default=100, gt=0, allow_inf_nan=False)
@model_validator(mode="after")
def unique_checks(self) -> "LensSettings":
@ -72,32 +92,32 @@ class LensSettings(Record):
class Evidence(Record):
execution_id: str
span_id: str
quote: str = Field(min_length=1, max_length=1000)
quote: str = Field(min_length=1)
role: Literal["support", "counterexample"] = "support"
class AgentTestCase(Record):
input: str = Field(min_length=1, max_length=1000)
expected: str = Field(min_length=1, max_length=1000)
input: str = Field(min_length=1)
expected: str = Field(min_length=1)
class IssueBrief(Record):
problem: str = Field(min_length=10, max_length=400)
user_goal: str = Field(min_length=3, max_length=400)
what_happened: str = Field(min_length=3, max_length=1500)
test_cases: tuple[AgentTestCase, ...] = Field(min_length=1, max_length=5)
problem: str = Field(min_length=10)
user_goal: str = Field(min_length=3)
what_happened: str = Field(min_length=3)
test_cases: tuple[AgentTestCase, ...] = Field(min_length=1)
class FindingDraft(Record):
title: str = Field(min_length=3, max_length=160)
description: str = Field(min_length=10, max_length=4000)
title: str = Field(min_length=3)
description: str = Field(min_length=10)
check_id: str
kind: Literal["issue", "pattern"] = "issue"
priority: Literal["high", "medium", "low"] = "medium"
suggestion: str = Field(default="", max_length=2000)
limitation: str = Field(default="", max_length=600)
suggestion: str = Field(default="")
limitation: str = Field(default="")
brief: IssueBrief | None = None
evidence: tuple[Evidence, ...] = Field(min_length=1, max_length=20)
evidence: tuple[Evidence, ...] = Field(min_length=1)
existing_finding_id: str | None = None
@ -228,12 +248,12 @@ class LensList(Record):
class RunRequest(Record):
settings: LensSettings | None = None
lookback_hours: int | None = Field(default=None, ge=1, le=8760)
lookback_hours: LookbackHours | None = None
class FindingUpdate(Record):
status: Literal["open", "resolved", "dismissed"]
reason: str = Field(default="", max_length=2000)
reason: str = Field(default="")
class Claim(Record):
@ -243,7 +263,7 @@ class Claim(Record):
class Progress(Record):
stage: str = Field(max_length=100)
stage: str = Field()
coverage: Coverage = Coverage()
@ -251,14 +271,15 @@ class Result(Record):
assessments: tuple[RunAssessment, ...] = ()
findings: tuple[FindingDraft, ...] = ()
coverage: Coverage
error: str = Field(default="", max_length=1000)
error: str = Field(default="")
class ModelRequest(Record):
prompt: str = Field(min_length=1, max_length=100000)
prompt: str = Field(min_length=1)
purpose: Literal["extract", "cluster", "investigate"]
class ModelResult(Record):
content: str
cost: float
finish_reason: Literal["length", "content_filter"] | None = Field(default=None, exclude=True)

View file

@ -10,6 +10,8 @@ Return action='read' with execution_id, cursor (span ID; default empty), offset
Reads return up to 40 spans; advance cursor from next_cursor for more spans or offset by 8000 for longer content; offset=1 reads original beginning after an abbreviated excerpt.
Read any execution in the supplied catalog.
Use action='catalog' or 'observations' with page to fetch another page of runs or supporting observations.
Use action='evidence' with page to read the remaining content in a fetched batch; evidence_pages includes every supplied span.
Read needed evidence pages before advancing the span cursor. Evidence pages reset to zero after a read or observation-page change.
Use action=feedback to read prior findings and dismissal reasons only when feedback_pages>1.
The current page is already supplied; feedback_pages=0 means no prior findings or feedback exist, so do not request feedback.
Request only page numbers below the corresponding page count.
@ -19,16 +21,16 @@ Mark quotes from runs that demonstrate the opposite behavior as counterexample,
Include at least one supporting quote.
Never put internal run aliases in prose; the evidence links identify the runs.
Write for a busy person, in plain English.
Title: a short, concrete outcome in at most 12 words.
Description: one or two short sentences saying what happened and why it matters, at most 60 words.
Put uncertainty or counterexamples in limitation, not in the main description; use at most 40 words.
Suggestion: one specific action, at most 25 words, or empty if no action is needed.
Title: a short, concrete outcome.
Description: one or two short sentences saying what happened and why it matters.
Put uncertainty or counterexamples in limitation, not in the main description.
Suggestion: one specific action, or empty if no action is needed.
For issues, also return brief, which describes the failure so anyone can reproduce and verify it without access to the agent's code.
Scope what went wrong from the evidence: compare each failed or empty tool result with the tools, permissions, working directory, and configuration visible in the recorded requests, and name the most specific cause the evidence supports.
brief.problem: the root cause in one or two sentences.
brief.user_goal: what the end user was trying to achieve.
brief.what_happened: what the agent actually output or did, quoting the recorded output where possible.
brief.test_cases: one to five user inputs drawn from the evidence, each with the behavior a correct agent should show.
brief.test_cases: user inputs drawn from the evidence, each with the behavior a correct agent should show.
Do not prescribe code or configuration changes in brief.
Omit brief for patterns.
Avoid jargon such as document-borne, visible noncompliance, instruction-bearing, or evaluator-directed.

View file

@ -20,7 +20,6 @@ Respect prior feedback about accepted behavior, but do not suppress different pr
Request reads with span_id and offset=0 for initial evidence.
If an excerpt omits content, offset=1 reads the original beginning; later offsets advance by 8000 characters through the original stored span.
Do not repeat a completed read.
At most two reads per turn.
Return observations using an enabled check ID, exact quotes, and the correct execution_id/span_id.
Never quote an omission marker or join text from either side of one.
If you need more evidence, return reads; otherwise return reads=[] and your final observations.

View file

@ -126,7 +126,7 @@ class SourceReader:
metadata=tuple(
MetadataFilter(key=k, value=v)
for k, v in row.attributes
if k != "litellm.api_key_hash" and 0 < len(k) <= 200 and 0 < len(v) <= 500
if k != "litellm.api_key_hash" and k and v
),
)
for row in rows

View file

@ -3,14 +3,13 @@ import logging
import os
import sqlite3
from collections.abc import Awaitable, Callable
from contextlib import suppress
from types import MappingProxyType
from typing import Final
import httpx
from pydantic import BaseModel, ConfigDict, ValidationError
from .analysis import analyze_sample
from .analysis import AnalysisResponseError, analyze_sample, validation_details
from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult, Progress, Result, Sample
logger: Final = logging.getLogger("litellm.lens.worker")
@ -27,7 +26,21 @@ class ClaimIdentity(BaseModel):
job: ClaimedJobIdentity
class PublicModelError(BaseModel):
model_config = ConfigDict(extra="ignore")
lens_error: str
class ModelErrorEnvelope(BaseModel):
model_config = ConfigDict(extra="ignore")
detail: PublicModelError
def failure_message(error: Exception) -> str:
if isinstance(error, AnalysisResponseError):
return str(error)
if isinstance(error, ValidationError):
return f"Invalid {error.title} response (ValidationError):\n{validation_details(error)}"
if isinstance(error, (OSError, sqlite3.Error)):
return "Worker temporary storage failed. Increase its capacity or reduce analysis parallelism."
if isinstance(error, httpx.TimeoutException):
@ -46,6 +59,12 @@ def failure_message(error: Exception) -> str:
else "Worker request"
)
status: Final = error.response.status_code
if path.endswith("/model"):
try:
diagnostic: Final = ModelErrorEnvelope.model_validate_json(error.response.content)
return f"Model request failed (HTTP {status}):\n{diagnostic.detail.lens_error}"
except ValueError:
pass
guidance: Final = MappingProxyType(
{
400: "Check the configured model and whether the worker's billing key is enabled.",
@ -62,15 +81,33 @@ def failure_message(error: Exception) -> str:
class LensWorker:
def __init__(self, client: httpx.AsyncClient, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep) -> None:
def __init__(
self,
client: httpx.AsyncClient,
sleep: Callable[[float], Awaitable[None]] = asyncio.sleep,
heartbeat_wait: Callable[[float], Awaitable[None]] = asyncio.sleep,
) -> None:
self.client: Final = client
self.sleep: Final = sleep
self.heartbeat_wait: Final = heartbeat_wait
async def model_request(self, path: str, body: ModelRequest, attempt: int = 0) -> ModelResult:
try:
result: Final = await self.client.post(path, json=body.model_dump())
timeout: Final = httpx.Timeout(
None,
connect=self.client.timeout.connect,
write=self.client.timeout.write,
pool=self.client.timeout.pool,
)
result: Final = await self.client.post(path, json=body.model_dump(), timeout=timeout)
result.raise_for_status()
return ModelResult.model_validate(result.json())
parsed: Final = ModelResult.model_validate(result.json())
reason: Final = result.headers.get("x-litellm-lens-finish-reason")
return (
parsed.model_copy(update=MappingProxyType({"finish_reason": reason}))
if reason in ("length", "content_filter")
else parsed
)
except (httpx.TransportError, httpx.HTTPStatusError) as exc:
retryable: Final = not isinstance(exc, httpx.HTTPStatusError) or exc.response.status_code in (
429,
@ -84,7 +121,7 @@ class LensWorker:
return await self.model_request(path, body, attempt + 1)
async def run_once(self) -> bool:
response: Final = await self.client.post("/lens/worker/claim", params=MappingProxyType({"protocol_version": 2}))
response: Final = await self.client.post("/lens/worker/claim", params=MappingProxyType({"protocol_version": 3}))
response.raise_for_status()
payload: Final = response.json()
if payload is None:
@ -131,17 +168,30 @@ class LensWorker:
async def heartbeat() -> None:
while True:
await asyncio.sleep(30)
(await self.client.post(prefix + "/heartbeat")).raise_for_status()
await self.heartbeat_wait(30)
try:
(await self.client.post(prefix + "/heartbeat")).raise_for_status()
except (httpx.TransportError, httpx.HTTPStatusError) as exc:
if isinstance(exc, httpx.HTTPStatusError) and (
exc.response.status_code < 500 and exc.response.status_code != 429
):
raise
logger.warning("Analysis %s heartbeat will retry (%s)", claim.job.id, type(exc).__name__)
pulse_task: Final = asyncio.create_task(heartbeat())
try:
async def investigate() -> None:
data: Final = await self.client.get(prefix + "/sample")
data.raise_for_status()
sample: Final = Sample.model_validate(data.json())
result: Final = await analyze_sample(claim, sample, read, model, progress)
saved: Final = await self.client.post(prefix + "/result", json=result.model_dump(mode="json"))
saved.raise_for_status()
pulse_task: Final = asyncio.create_task(heartbeat())
work_task: Final = asyncio.create_task(investigate())
try:
finished, _ = await asyncio.wait((pulse_task, work_task), return_when=asyncio.FIRST_COMPLETED)
for task in finished:
await task
except (httpx.HTTPError, ValueError, OSError, sqlite3.Error) as exc:
message: Final = failure_message(exc)
logger.warning("Analysis %s interrupted (%s)", claim.job.id, type(exc).__name__)
@ -152,8 +202,8 @@ class LensWorker:
failed.raise_for_status()
finally:
pulse_task.cancel()
with suppress(asyncio.CancelledError, httpx.HTTPError):
await pulse_task
work_task.cancel()
await asyncio.gather(pulse_task, work_task, return_exceptions=True)
return True

View file

@ -28,6 +28,7 @@ from fastapi import (
from fastapi.responses import StreamingResponse
from pydantic import TypeAdapter
from starlette.datastructures import UploadFile as StarletteUploadFile
from starlette.routing import BaseRoute, Route
from starlette.websockets import WebSocketState
from websockets.asyncio.client import connect
from websockets.exceptions import (
@ -67,6 +68,7 @@ from litellm.llms.base_llm.managed_resources.utils import (
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.oss_decision import validate_oss_request
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._lazy_features import lazy_owned_routes
from litellm.proxy._types import (
ConfigFieldInfo,
ConfigFieldUpdate,
@ -2942,35 +2944,34 @@ def _extract_model_from_vertex_ai_setup(setup_response: Mapping[str, object]) ->
return None
def _placed_ahead(routes: Sequence[BaseRoute], moving: BaseRoute, before: BaseRoute) -> tuple[BaseRoute, ...]:
kept: Final = tuple(route for route in routes if route is not moving)
at: Final = next(index for index, route in enumerate(kept) if route is before)
return (*kept[:at], moving, *kept[at:])
class SafeRouteAdder:
"""
Wrapper class for adding routes to FastAPI app.
Only adds routes if they don't already exist on the app.
Only adds routes if they don't already exist on the app. A route a lazy feature registered
does not count: a route added at its path goes ahead of it, the precedence a config
pass-through at /v1/decisions gets in lazy mode, where the feature has not loaded yet.
"""
@staticmethod
def _colliding_routes(app: FastAPI, path: str, methods: Sequence[str]) -> tuple[Route, ...]:
wanted: Final = frozenset(methods)
return tuple(
route
for route in app.routes
if isinstance(route, Route) and route.path == path and not wanted.isdisjoint(route.methods or ())
)
@staticmethod
def _is_path_registered(app: FastAPI, path: str, methods: list[str]) -> bool:
"""
Check if a path with any of the specified methods is already registered on the app.
Args:
app: The FastAPI application instance
path: The path to check (e.g., "/v1/chat/completions")
methods: List of HTTP methods to check (e.g., ["GET", "POST"])
Returns:
True if the path is already registered with any of the methods, False otherwise
"""
for route in app.routes:
# Use getattr to safely access route attributes
route_path = getattr(route, "path", None)
route_methods = getattr(route, "methods", None)
if route_path == path and route_methods is not None:
# Check if any of the methods overlap
if any(method in route_methods for method in methods):
return True
return False
"""True when a route the app itself defines already serves the path with one of the methods."""
lazy_owned: Final = lazy_owned_routes(app)
return any(id(route) not in lazy_owned for route in SafeRouteAdder._colliding_routes(app, path, methods))
@staticmethod
def add_api_route_if_not_exists(
@ -3001,12 +3002,17 @@ class SafeRouteAdder:
)
return False
shadowed: Final = SafeRouteAdder._colliding_routes(app, path, methods)
app.add_api_route(
path=path,
endpoint=endpoint,
methods=methods,
dependencies=dependencies,
)
if shadowed:
app.router.routes[:] = _placed_ahead( # rebind-ok: the app owns its route table
app.router.routes, app.router.routes[-1], shadowed[0]
)
verbose_proxy_logger.debug(
"Successfully added route: %s with methods %s",
path,

View file

@ -93,6 +93,7 @@ ROUTE_ENDPOINT_MAPPING: Final = {
"acompact_responses": "/responses/compact",
"aocr": "/ocr",
"asearch": "/search",
"adecisions": "/decisions",
"avideo_generation": "/videos",
"avideo_list": "/videos",
"avideo_status": "/videos/{video_id}",
@ -487,6 +488,7 @@ RouteType = Literal[
"avector_store_file_delete",
"aocr",
"asearch",
"adecisions",
"avideo_generation",
"avideo_list",
"avideo_status",

View file

@ -140,7 +140,7 @@ async def list_agent_traces(
context: Annotated[TraceAccessContext, Depends(provide_trace_access)],
start_ms: Annotated[int | None, Query(description="Window start, unix ms. Default: 24h ago")] = None,
end_ms: Annotated[int | None, Query(description="Window end, unix ms. Default: now")] = None,
cursor: Annotated[str | None, Query()] = None,
cursor: Annotated[str | None, Query(max_length=512)] = None,
) -> TracePage:
now_ms: Final = int(time.time() * 1000)
try:
@ -153,6 +153,13 @@ async def list_agent_traces(
)
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
class TraceQueryRequest(BaseModel):
@ -228,12 +235,21 @@ async def get_agent_trace(
trace_id: str,
context: Annotated[TraceAccessContext, Depends(provide_trace_access)],
trace_ref: Annotated[str, Query()] = "",
cursor: Annotated[str | None, Query(max_length=512)] = None,
page_size: Annotated[int | None, Query(ge=1, le=500)] = None,
) -> Trace:
tracing, scope = context.reader()
try:
trace: Final = await tracing.get_trace(trace_id, scope, trace_ref)
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
if trace is None:
raise HTTPException(status_code=404, detail=f"Trace {trace_id} not found")
return trace
@ -251,6 +267,13 @@ async def get_agent_trace_span(
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
if span is None:
raise HTTPException(status_code=404, detail=f"Span {span_id} not found")
return span
@ -269,6 +292,13 @@ async def get_agent_trace_span_error(
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
if page is None:
raise HTTPException(status_code=404, detail="Span diagnostic not found or no longer available")
return page

View file

@ -2110,6 +2110,11 @@ class Router:
self.asearch = self.factory_function(asearch, call_type="asearch")
self.search = self.factory_function(search, call_type="search")
from litellm.decisions import adecisions, decisions
self.adecisions = self.factory_function(adecisions, call_type="adecisions")
self.decisions = self.factory_function(decisions, call_type="decisions")
def _initialize_video_endpoints(self):
"""Initialize video endpoints."""
from litellm.videos import (
@ -6663,6 +6668,8 @@ class Router:
"ocr",
"asearch",
"search",
"adecisions",
"decisions",
"aadapter_generate_content",
"avideo_generation",
"video_generation",
@ -6736,6 +6743,7 @@ class Router:
"generate_content_stream",
"ocr",
"search",
"decisions",
"video_generation",
"video_list",
"video_status",
@ -6903,6 +6911,7 @@ class Router:
"agenerate_content_stream",
"aocr",
"ocr",
"adecisions",
"avideo_generation",
"avideo_list",
"avideo_status",

View file

@ -45,7 +45,9 @@ class NativeTraceStorage:
def list_traces(
self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None, limit: int
) -> Future[JsonValue]: ...
def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str) -> Future[JsonValue]: ...
def get_trace(
self, trace_id: str, scope: TraceScope, trace_ref: str, cursor: str | None = None, page_size: int | None = None
) -> Future[JsonValue]: ...
def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str) -> Future[JsonValue]: ...
def get_span_error(
self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str, cursor: str | None

View file

@ -99,6 +99,7 @@ class UIMessage(typing_extensions.TypedDict):
class TraceSummary(typing_extensions.TypedDict):
resolution_limited: ReadOnly[NotRequired[bool]]
trace_id: ReadOnly[str]
trace_ref: ReadOnly[NotRequired[str]]
name: ReadOnly[str]
@ -145,6 +146,7 @@ class Trace(typing_extensions.TypedDict):
summary: ReadOnly[TraceSummary]
agents: ReadOnly[tuple[AgentNode, ...]]
spans: ReadOnly[tuple[Span, ...]]
next_cursor: ReadOnly[NotRequired[str | None]]
class TracePage(typing_extensions.TypedDict):

View file

@ -67,7 +67,9 @@ class NativeStore(Protocol):
self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None, limit: int
) -> Awaitable[JsonValue]: ...
def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str) -> Awaitable[JsonValue]: ...
def get_trace(
self, trace_id: str, scope: TraceScope, trace_ref: str, cursor: str | None = None, page_size: int | None = None
) -> Awaitable[JsonValue]: ...
def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str) -> Awaitable[JsonValue]: ...
@ -189,8 +191,15 @@ class ClickHouseStorage:
result: Final = await self._native.list_traces(scope, start_ms, end_ms, cursor, limit)
return _validate_query_response(_TRACE_PAGE, result)
async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None:
result: Final = await self._native.get_trace(trace_id, scope, trace_ref)
async def get_trace(
self,
trace_id: str,
scope: TraceScope,
trace_ref: str = "",
cursor: str | None = None,
page_size: int | None = None,
) -> Trace | None:
result: Final = await self._native.get_trace(trace_id, scope, trace_ref, cursor, page_size)
return _validate_query_response(_TRACE, result)
async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None:

View file

@ -99,8 +99,15 @@ class TraceReceiver:
async def list_traces(self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None = None) -> TracePage:
return await self.storage.list_traces(scope, start_ms, end_ms, cursor, AGENT_TRACING_LIST_PAGE_SIZE)
async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None:
return await self.storage.get_trace(trace_id, scope, trace_ref)
async def get_trace(
self,
trace_id: str,
scope: TraceScope,
trace_ref: str = "",
cursor: str | None = None,
page_size: int | None = None,
) -> Trace | None:
return await self.storage.get_trace(trace_id, scope, trace_ref, cursor, page_size)
async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None:
return await self.storage.get_span(trace_id, span_id, scope, trace_ref)

123
litellm/types/decisions.py Normal file
View file

@ -0,0 +1,123 @@
from collections.abc import Mapping, Sequence
from typing import Annotated, Literal, TypeAlias
from pydantic import ConfigDict, Field, PrivateAttr, model_validator, with_config
from typing_extensions import ReadOnly, Required, TypedDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
DecisionsJSON: TypeAlias = str | Mapping[str, object] | Sequence[object]
NoulCriteria: TypeAlias = Mapping[Literal["true", "false"], DecisionsJSON | None]
class NoulQuestion(LiteLLMPydanticObjectBase):
type: Literal["noul"]
instructions: DecisionsJSON | None = None
criteria: NoulCriteria | None = None
model_config = ConfigDict(extra="allow", frozen=True)
@model_validator(mode="after")
def require_instructions_or_criteria(self) -> "NoulQuestion":
if self.instructions is None and self.criteria is None:
raise ValueError("A noul question requires instructions or criteria")
return self
class ChoiceQuestion(LiteLLMPydanticObjectBase):
type: Literal["choice"]
instructions: DecisionsJSON | None = None
criteria: Annotated[Mapping[str, DecisionsJSON | None], Field(min_length=1, max_length=255)]
model_config = ConfigDict(extra="allow", frozen=True)
class ScoreQuestion(LiteLLMPydanticObjectBase):
type: Literal["score"]
instructions: DecisionsJSON | None = None
criteria: Annotated[Sequence[DecisionsJSON], Field(min_length=1, max_length=10)]
model_config = ConfigDict(extra="allow", frozen=True)
DecisionQuestion: TypeAlias = Annotated[
NoulQuestion | ChoiceQuestion | ScoreQuestion,
Field(discriminator="type"),
]
DecisionQuestionMap: TypeAlias = Annotated[
Mapping[Annotated[str, Field(min_length=1)], DecisionQuestion],
Field(min_length=1, max_length=128),
]
class DecisionsRequestBody(LiteLLMPydanticObjectBase):
state: DecisionsJSON
questions: DecisionQuestionMap
model_config = ConfigDict(extra="allow", frozen=True)
class DecisionsRequest(DecisionsRequestBody):
model: str
@with_config(ConfigDict(extra="allow"))
class DecisionsCallParams(TypedDict, total=False):
model: Required[ReadOnly[str]]
state: Required[ReadOnly[DecisionsJSON]]
questions: Required[ReadOnly[DecisionQuestionMap]]
api_key: ReadOnly[str | None]
api_base: ReadOnly[str | None]
timeout: ReadOnly[float | None]
custom_llm_provider: ReadOnly[str | None]
extra_headers: ReadOnly[Mapping[str, str] | None]
class NoulAnswer(LiteLLMPydanticObjectBase):
type: Literal["noul"]
noul: float
model_config = ConfigDict(extra="allow", frozen=True)
class ChoiceAnswer(LiteLLMPydanticObjectBase):
type: Literal["choice"]
choice: str
confidence: float
probabilities: Mapping[str, float]
model_config = ConfigDict(extra="allow", frozen=True)
class ScoreAnswer(LiteLLMPydanticObjectBase):
type: Literal["score"]
score: float
confidence: float
legend: Mapping[str, DecisionsJSON]
probabilities: Mapping[str, float]
model_config = ConfigDict(extra="allow", frozen=True)
DecisionAnswer: TypeAlias = Annotated[
NoulAnswer | ChoiceAnswer | ScoreAnswer,
Field(discriminator="type"),
]
class DecisionsUsage(LiteLLMPydanticObjectBase):
input_tokens: int = 0
output_tokens: int = 0
model_config = ConfigDict(extra="allow", frozen=True)
class DecisionsResponse(LiteLLMPydanticObjectBase):
model: str | None = None
answers: Mapping[str, DecisionAnswer]
usage: DecisionsUsage | None = None
model_config = ConfigDict(extra="allow", frozen=True)
_hidden_params: dict[str, object] = PrivateAttr(default_factory=dict)

View file

@ -468,6 +468,8 @@ class CallTypes(str, Enum):
arerank = "arerank"
search = "search"
asearch = "asearch"
decisions = "decisions"
adecisions = "adecisions"
arealtime = "_arealtime"
aresponses_websocket = "_aresponses_websocket"
create_batch = "create_batch"
@ -654,6 +656,8 @@ CallTypesLiteral = Literal[
"arerank",
"search",
"asearch",
"decisions",
"adecisions",
"_arealtime",
"_aresponses_websocket",
"create_batch",
@ -763,6 +767,8 @@ API_ROUTE_TO_CALL_TYPES: Final[Mapping[str, Sequence[CallTypes]]] = {
# Search
"/search": [CallTypes.asearch, CallTypes.search],
"/v1/search": [CallTypes.asearch, CallTypes.search],
"/decisions": [CallTypes.adecisions, CallTypes.decisions],
"/v1/decisions": [CallTypes.adecisions, CallTypes.decisions],
# Batches
"/batches": [CallTypes.acreate_batch, CallTypes.create_batch],
"/v1/batches": [CallTypes.acreate_batch, CallTypes.create_batch],
@ -4052,6 +4058,8 @@ class LlmProviders(str, Enum):
OLLAMA_CHAT = "ollama_chat"
DEEPINFRA = "deepinfra"
PERPLEXITY = "perplexity"
TYPESAFE = "typesafe"
STRANDS_DECIDER = "strands_decider"
MISTRAL = "mistral"
MILVUS = "milvus"
GROQ = "groq"

View file

@ -1218,6 +1218,9 @@ def function_setup(
if isinstance(search_query, list)
else search_query
)
elif call_type in (CallTypes.decisions.value, CallTypes.adecisions.value):
decisions_state: Final = args[1] if len(args) > 1 else kwargs.get("state", "")
messages = decisions_state if isinstance(decisions_state, str) else json.dumps(decisions_state)
elif call_type in (CallTypes.image_edit.value, CallTypes.aimage_edit.value):
messages = args[1] if len(args) > 1 else kwargs.get("prompt")
elif call_type in (CallTypes.ocr.value, CallTypes.aocr.value):
@ -9909,7 +9912,7 @@ class ProviderConfigManager:
@staticmethod
def get_provider_harness_config(harness: Harness) -> BaseHarnessConfig | None:
"""
Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents).
Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents, Tool Loop).
"""
from litellm.harness.types import Harness as _Harness
@ -9935,6 +9938,10 @@ class ProviderConfigManager:
)
return DeepAgentsHarnessConfig()
if harness == _Harness.TOOL_LOOP:
from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig
return ToolLoopHarnessConfig()
return None
@staticmethod

View file

@ -11542,9 +11542,11 @@
},
"azure_ai/flux.2-pro": {
"litellm_provider": "azure_ai",
"max_input_tokens": 32000,
"max_tokens": 32000,
"mode": "image_generation",
"output_cost_per_image": 0.04,
"source": "https://ai.azure.com/explore/models/flux.2-pro/version/1/registry/azureml-blackforestlabs",
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure",
"supported_endpoints": [
"/v1/images/generations"
]
@ -15612,6 +15614,38 @@
"prompt_cache_min_tokens": 1024,
"source": "https://platform.claude.com/docs/en/about-claude/pricing"
},
"cloudflare/clef": {
"input_cost_per_token": 2.4e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/clef-flash": {
"input_cost_per_token": 9e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/@cf/cloudflare/clef": {
"input_cost_per_token": 2.4e-07,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/@cf/cloudflare/clef-flash": {
"input_cost_per_token": 9e-08,
"litellm_provider": "cloudflare",
"max_input_tokens": 65536,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://developers.cloudflare.com/workers-ai/platform/pricing/"
},
"cloudflare/@cf/meta/llama-2-7b-chat-fp16": {
"input_cost_per_token": 1.923e-06,
"litellm_provider": "cloudflare",
@ -44421,6 +44455,14 @@
"mode": "chat",
"output_cost_per_token": 2.8e-07
},
"perplexity/pplx-decider-v1-27b": {
"input_cost_per_token": 4e-08,
"litellm_provider": "perplexity",
"max_input_tokens": 262144,
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://docs.perplexity.ai/api-reference/decisions-post"
},
"perplexity/sonar": {
"input_cost_per_token": 1e-06,
"litellm_provider": "perplexity",
@ -72800,6 +72842,16 @@
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"strands_decider/strands-decider-2B-hobson-v19": {
"input_cost_per_token": 0.0,
"litellm_provider": "strands_decider",
"mode": "evaluation",
"output_cost_per_token": 0.0,
"source": "https://huggingface.co/StrandsAgents/strands-decider-2B-hobson-v19",
"metadata": {
"notes": "Self-hosted decision model; infrastructure costs are paid separately"
}
},
"typesafe/jev-1.13.0": {
"input_cost_per_token": 4.2e-08,
"litellm_provider": "typesafe",
@ -76881,6 +76933,7 @@
"source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrock/current/us-east-1/index.json",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_prompt_cache_breakpoint": false,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
@ -76902,6 +76955,7 @@
"source": "https://aws.amazon.com/bedrock/pricing/",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_prompt_cache_breakpoint": false,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
@ -76923,6 +76977,7 @@
"source": "https://aws.amazon.com/bedrock/pricing/",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_prompt_cache_breakpoint": false,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,

View file

@ -2566,6 +2566,20 @@
"image_variations": true
}
},
"typesafe": {
"display_name": "TypeSafe (`typesafe`)",
"url": "https://docs.typesafe.ai/models",
"endpoints": {
"systemone": true
}
},
"strands_decider": {
"display_name": "Strands Decider (`strands_decider`)",
"url": "https://docs.litellm.ai/docs/providers",
"endpoints": {
"systemone": true
}
},
"tavily": {
"display_name": "Tavily (`tavily`)",
"url": "https://docs.litellm.ai/docs/search/tavily",

268
scripts/lens_dev.sh Executable file
View file

@ -0,0 +1,268 @@
#!/usr/bin/env bash
# One-command Lens local dev loop: proxy + Lens worker + hot-reload dashboard.
#
# LENS_DEV_PROXY_PORT proxy port (default 4000)
# LENS_DEV_UI_PORT next dev port (default 3000)
# LENS_DEV_MASTER_KEY master key, also the admin UI password
# (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
#
# State (master key, worker token, generated config, logs) lives in .lens-dev/ (gitignored).
set -euo pipefail
repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
proxy_port="${LENS_DEV_PROXY_PORT:-4000}"
ui_port="${LENS_DEV_UI_PORT:-3000}"
state_dir="${LENS_DEV_STATE_DIR:-$repo_root/.lens-dev}"
log_dir="$state_dir/logs"
token_file="$state_dir/worker_token"
key_file="$state_dir/master_key"
proxy_url="http://localhost:$proxy_port"
py="${LENS_DEV_PYTHON:-$repo_root/.venv/bin/python}"
database_url="${LENS_DEV_DATABASE_URL:-postgresql://litellm:litellm@127.0.0.1:15432/litellm}"
clickhouse_url=http://default:local-tracing@127.0.0.1:18123
master_key=""
pids=()
die() { echo "lens-dev: $*" >&2; exit 1; }
listening() { lsof -nP -iTCP:"$1" -sTCP:LISTEN >/dev/null 2>&1; }
# A fixed key would let anyone who can reach the proxy sign in as admin, so default to a
# random key generated once per checkout and kept next to the worker token.
load_master_key() {
if [ -n "${LENS_DEV_MASTER_KEY:-}" ]; then
master_key="$LENS_DEV_MASTER_KEY"
return
fi
if [ ! -s "$key_file" ]; then
(umask 077 && printf 'sk-%s\n' "$(openssl rand -hex 24)" > "$key_file")
fi
master_key="$(cat "$key_file")"
}
# Only reuse a listener on 15432/18123 if it accepts the tracing stack's credentials;
# start the compose service when nothing is listening; fail if something else is.
# A LENS_DEV_DATABASE_URL is left to the proxy, which may use Prisma-only URL params.
ensure_services() {
local services=()
if [ -z "${LENS_DEV_DATABASE_URL:-}" ] && listening 15432; then
"$py" -c 'import sys, psycopg; psycopg.connect(sys.argv[1], connect_timeout=5).close()' "$database_url" 2>/dev/null \
|| die "port 15432 is taken by something that isn't the tracing Postgres (litellm/litellm)"
elif [ -z "${LENS_DEV_DATABASE_URL:-}" ]; then
services+=(db)
fi
if listening 18123; then
[ "$(curl -fsS --max-time 5 "$clickhouse_url/?query=SELECT%201" 2>/dev/null)" = 1 ] \
|| die "port 18123 is taken by something that isn't the tracing ClickHouse (default/local-tracing)"
else
services+=(clickhouse)
fi
if [ "${#services[@]}" -gt 0 ]; then
docker compose -f docker/docker-compose.tracing.yml up -d --wait "${services[@]}"
else
echo "lens-dev: reusing running Postgres and ClickHouse"
fi
}
write_default_config() {
cat > "$1" <<'EOF'
model_list:
- model_name: gpt-4.1-mini
litellm_params:
model: openai/gpt-4.1-mini
api_key: os.environ/OPENAI_API_KEY
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
store_prompts_in_spend_logs: true
tracing:
store:
type: clickhouse
url: os.environ/CLICKHOUSE_URL
retention_days: 14
EOF
}
# litellm's implicit load_dotenv() walks up from a worktree into the parent checkout's
# .env and picks up REDIS_* / UI_* from there. LITELLM_MODE=PRODUCTION turns that off;
# this prints export lines for the same .env minus those vars, so provider keys still load.
dotenv_exports() {
"$py" - <<'PY'
import os, re, shlex
from dotenv import dotenv_values, find_dotenv
path = find_dotenv(usecwd=True)
skip = re.compile(r"REDIS_.*|UI_USERNAME|UI_PASSWORD|LITELLM_MODE|ANTHROPIC_BASE_URL|ANTHROPIC_AUTH_TOKEN|ANTHROPIC_CUSTOM_HEADERS|OPENAI_BASE_URL|OPENAI_API_BASE")
for key, value in (dotenv_values(path) if path else {}).items():
if value is not None and key not in os.environ and re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", key) and not skip.fullmatch(key):
print(f"export {key}={shlex.quote(value)}")
PY
}
# Run in the proxy's subshell: drop inherited settings that would point it at someone
# else's services, then set the local stack's.
proxy_env() {
local var
# Claude Code and similar tools export these; provider calls would go to them.
unset ANTHROPIC_BASE_URL ANTHROPIC_AUTH_TOKEN ANTHROPIC_CUSTOM_HEADERS OPENAI_BASE_URL OPENAI_API_BASE
for var in $(compgen -e | grep '^REDIS_' || true); do unset "$var"; done
eval "$1"
export LITELLM_MODE=PRODUCTION
export LITELLM_MASTER_KEY="$master_key"
if [ "$master_key" = sk-1234 ]; then export LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true; fi
export LITELLM_SALT_KEY=sk-local-tracing-salt-key
export DATABASE_URL="$database_url"
export STORE_MODEL_IN_DB=True
export CLICKHOUSE_URL="$clickhouse_url"
export CLICKHOUSE_DATABASE=litellm
export LITELLM_LOCAL_MODEL_COST_MAP=True
export PROXY_BASE_URL="$proxy_url"
export UI_USERNAME=admin
export UI_PASSWORD="$master_key"
}
# POST JSON as the admin and print one field of the response; dies with the body on failure.
admin_post() {
local body
body="$(curl -sS --fail-with-body "$proxy_url$1" -H "Authorization: Bearer $master_key" \
-H "Content-Type: application/json" -d "$2")" || die "POST $1 failed: $body"
"$py" -c 'import json, sys; print(json.loads(sys.argv[1])[sys.argv[2]])' "$body" "$3"
}
register_worker() {
local key_hash worker_token
key_hash="$(admin_post /key/generate "{\"key_alias\": \"lens-dev-$(date +%s)\"}" token)"
worker_token="$(admin_post /lens/workers/register "{\"name\": \"lens-dev\", \"analysis_key_id\": \"$key_hash\"}" token)"
(umask 077 && printf '%s\n' "$worker_token" > "$token_file")
echo "lens-dev: registered a new Lens worker (token in $token_file)"
}
# Auth runs before the handler, so protocol_version=1 answers 401 for a bad token and
# 409 for a good one without claiming a job.
ensure_worker_token() {
local status
if [ ! -s "$token_file" ]; then
register_worker
return
fi
status="$(curl -sS -o /dev/null -w '%{http_code}' -X POST "$proxy_url/lens/worker/claim?protocol_version=1" \
-H "Authorization: Bearer $(cat "$token_file")")"
case "$status" in
409) echo "lens-dev: reusing worker token from $token_file" ;;
401) echo "lens-dev: stored worker token was rejected"; register_worker ;;
*) die "unexpected HTTP $status checking the worker token" ;;
esac
}
wait_for_proxy() {
local proxy_pid="$1"
echo "lens-dev: waiting for the proxy (log: $log_dir/proxy.log)"
for _ in $(seq 1 300); do
kill -0 "$proxy_pid" 2>/dev/null || die "proxy exited; see $log_dir/proxy.log"
curl -fsS "$proxy_url/health/readiness" -H "Authorization: Bearer $master_key" >/dev/null 2>&1 && return
sleep 1
done
die "proxy not ready after 300s; see $log_dir/proxy.log"
}
# Children run in their own process groups (set -m), so killing -pid takes their trees too.
cleanup() {
local alive pid
trap - EXIT INT TERM
[ "${#pids[@]}" -gt 0 ] || return 0
echo "lens-dev: stopping"
for pid in "${pids[@]}"; do kill -TERM -- "-$pid" 2>/dev/null || true; done
for _ in $(seq 1 20); do
alive=0
for pid in "${pids[@]}"; do kill -0 "$pid" 2>/dev/null && alive=1; done
[ "$alive" = 0 ] && break
sleep 0.5
done
for pid in "${pids[@]}"; do kill -KILL -- "-$pid" 2>/dev/null || true; done
}
main() {
local config_file exports proxy_pid pid key_hint
if [ -n "${LENS_DEV_CONFIG:-}" ]; then
[ -f "$LENS_DEV_CONFIG" ] || die "LENS_DEV_CONFIG not found: $LENS_DEV_CONFIG"
config_file="$(cd "$(dirname "$LENS_DEV_CONFIG")" && pwd)/$(basename "$LENS_DEV_CONFIG")"
fi
cd "$repo_root"
listening "$proxy_port" && die "port $proxy_port is in use; set LENS_DEV_PROXY_PORT"
listening "$ui_port" && die "port $ui_port is in use; set LENS_DEV_UI_PORT"
[ "$proxy_port" != "$ui_port" ] || die "proxy and UI ports must differ"
mkdir -p "$log_dir"
load_master_key
uv sync --inexact --frozen --extra proxy --group proxy-dev --no-install-project
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
if [ ! -x ui/litellm-dashboard/node_modules/.bin/next ]; then
(cd ui/litellm-dashboard && "$repo_root/scripts/with_dashboard_node.sh" npm ci)
fi
if [ -z "${config_file:-}" ]; then
config_file="$state_dir/config.yaml"
write_default_config "$config_file"
fi
exports="$(dotenv_exports)"
trap cleanup EXIT
trap 'exit 130' INT TERM
set -m
(
proxy_env "$exports"
exec "$py" litellm/proxy/proxy_cli.py --config "$config_file" --host 127.0.0.1 --port "$proxy_port"
) < /dev/null > "$log_dir/proxy.log" 2>&1 &
proxy_pid=$!
pids+=("$proxy_pid")
(
cd ui/litellm-dashboard
NEXT_PUBLIC_BASE_URL="$proxy_url" exec "$repo_root/scripts/with_dashboard_node.sh" npx next dev -p "$ui_port"
) < /dev/null > "$log_dir/ui.log" 2>&1 &
pids+=("$!")
wait_for_proxy "$proxy_pid"
ensure_worker_token
LITELLM_MODE=PRODUCTION LITELLM_URL="$proxy_url" LENS_WORKER_TOKEN="$(cat "$token_file")" \
"$py" -c "import asyncio, logging; from litellm.proxy.lens.worker import main; logging.basicConfig(level=logging.INFO); asyncio.run(main())" \
< /dev/null > "$log_dir/worker.log" 2>&1 &
pids+=("$!")
key_hint="password in $key_file"
[ -z "${LENS_DEV_MASTER_KEY:-}" ] || key_hint="password from LENS_DEV_MASTER_KEY"
cat <<EOF
Lens dev is up. Ctrl-C stops everything.
Log in: $proxy_url/ui/login (admin / $key_hint)
Lens: http://localhost:$ui_port/lens
Logs: $log_dir/proxy.log
$log_dir/worker.log
$log_dir/ui.log
Restart (Ctrl-C, make lens-dev) to pick up backend or worker edits; the UI hot-reloads.
EOF
while :; do
for pid in "${pids[@]}"; do
kill -0 "$pid" 2>/dev/null || die "a child process (pid $pid) exited; check the logs above"
done
sleep 2
done
}
# Sourcing (tests) only defines the functions.
if [ "${BASH_SOURCE[0]}" = "$0" ]; then
main "$@"
fi

View file

@ -245,6 +245,10 @@
"minimum": 0,
"type": "integer"
},
"resolution_limited": {
"type": "boolean",
"x-python-optional": true
},
"service": {
"type": "string"
},
@ -282,6 +286,7 @@
}
},
"required": [
"resolution_limited",
"trace_id",
"trace_ref",
"name",
@ -314,6 +319,13 @@
},
"type": "array"
},
"next_cursor": {
"type": [
"string",
"null"
],
"x-python-optional": true
},
"spans": {
"items": {
"$ref": "#/$defs/Span"
@ -327,7 +339,8 @@
"required": [
"summary",
"agents",
"spans"
"spans",
"next_cursor"
],
"title": "Trace",
"type": "object"

View file

@ -76,6 +76,10 @@
"minimum": 0,
"type": "integer"
},
"resolution_limited": {
"type": "boolean",
"x-python-optional": true
},
"service": {
"type": "string"
},
@ -113,6 +117,7 @@
}
},
"required": [
"resolution_limited",
"trace_id",
"trace_ref",
"name",

View file

@ -33,6 +33,7 @@ EXCLUDED_FOLDERS = {
"codex",
"opencode",
"deepagents",
"tool_loop",
}

View file

@ -50,12 +50,16 @@ def signal_group(group: int, action: int) -> None:
pass
def graceful_stop_seconds() -> float:
return max(30.0, float(os.environ.get("INTEGRATION_PROXY_READY_SECONDS", "70")))
def stop_root_process(process: subprocess.Popen[bytes]) -> bool:
if process.poll() is not None:
return True
process.terminate()
try:
process.wait(timeout=30)
process.wait(timeout=graceful_stop_seconds())
except subprocess.TimeoutExpired:
return False
return True
@ -183,6 +187,7 @@ def owned_proxy_process(
remove_environment: tuple[str, ...] = (),
workers: int = 1,
database_setup: tuple[str, ...] = DB_PUSH,
extra_arguments: tuple[str, ...] = (),
) -> Iterator[OwnedProxy]:
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
environment: Final = {
@ -209,6 +214,7 @@ def owned_proxy_process(
"--num_workers",
str(workers),
*database_setup,
*extra_arguments,
)
launch: Final = _launch_until_bound(command, root, environment, output, _PORT_ATTEMPTS)
process: Final = launch.process

View file

@ -246,6 +246,7 @@ class CostTrackingTestCase(BaseModel):
"/v1/audio/speech",
"/v1/images/generations",
"/v1/images/edits",
"/v1/decisions",
]
| Annotated[str, Field(pattern=r"^/(gemini|anthropic|bedrock)/")]
) = "/v1/chat/completions"

View file

@ -57,6 +57,13 @@
"search_context_size_high": 0.012
}
},
"perplexity/pplx-decider-v1-27b": {
"input_cost_per_token": 4e-08,
"litellm_provider": "perplexity",
"max_input_tokens": 262144,
"mode": "evaluation",
"output_cost_per_token": 0.0
},
"deepseek/deepseek-v4-chat": {
"litellm_provider": "deepseek",
"mode": "chat",
@ -27848,6 +27855,257 @@
"tool_usage_cost": 0.0125
}
},
{
"name": "gpt-5.6-responses_web_search_reported_count_mixed_actions",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"endpoint": "/v1/responses",
"request": {
"model": "$MODEL",
"input": "search this text",
"tools": [
{
"type": "web_search_preview",
"search_context_size": "medium"
}
]
},
"response": {
"content_type": "application/json",
"body": {
"id": "resp_$REQUEST_ID",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"type": "web_search_call",
"id": "ws_0_$REQUEST_ID",
"status": "completed",
"action": {
"type": "search",
"query": "scripted query"
}
},
{
"type": "web_search_call",
"id": "ws_1_$REQUEST_ID",
"status": "completed",
"action": {
"type": "open_page",
"url": "https://scripted.example/a"
}
},
{
"type": "web_search_call",
"id": "ws_2_$REQUEST_ID",
"status": "completed",
"action": {
"type": "open_page",
"url": "https://scripted.example/b"
}
},
{
"type": "web_search_call",
"id": "ws_3_$REQUEST_ID",
"status": "completed",
"action": {
"type": "find_in_page",
"pattern": "release",
"url": "https://scripted.example/a"
}
},
{
"type": "message",
"id": "msg_$REQUEST_ID",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "scripted response",
"annotations": []
}
]
}
],
"usage": {
"input_tokens": 1840,
"output_tokens": 412,
"total_tokens": 2252,
"input_tokens_details": {
"cached_tokens": 0
},
"output_tokens_details": {
"reasoning_tokens": 0
}
},
"tool_usage": {
"image_gen": {
"input_tokens": 0,
"input_tokens_details": {
"image_tokens": 0,
"text_tokens": 0
},
"output_tokens": 0,
"output_tokens_details": {
"image_tokens": 0,
"text_tokens": 0
},
"total_tokens": 0
},
"web_search": {
"num_requests": 1
}
}
}
},
"expected": {
"spend": 0.021488,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412,
"tool_usage_cost": 0.0125
}
},
{
"name": "gpt-5.6-responses_web_search_reported_count_open_page_only",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"endpoint": "/v1/responses",
"request": {
"model": "$MODEL",
"input": "search this text",
"tools": [
{
"type": "web_search_preview",
"search_context_size": "medium"
}
]
},
"response": {
"content_type": "application/json",
"body": {
"id": "resp_$REQUEST_ID",
"object": "response",
"created_at": 1700000000,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"type": "web_search_call",
"id": "ws_0_$REQUEST_ID",
"status": "completed",
"action": {
"type": "open_page",
"url": "https://scripted.example/a"
}
},
{
"type": "web_search_call",
"id": "ws_1_$REQUEST_ID",
"status": "completed",
"action": {
"type": "find_in_page",
"pattern": "release",
"url": "https://scripted.example/a"
}
},
{
"type": "message",
"id": "msg_$REQUEST_ID",
"status": "completed",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "scripted response",
"annotations": []
}
]
}
],
"usage": {
"input_tokens": 1840,
"output_tokens": 412,
"total_tokens": 2252,
"input_tokens_details": {
"cached_tokens": 0
},
"output_tokens_details": {
"reasoning_tokens": 0
}
},
"tool_usage": {
"image_gen": {
"input_tokens": 0,
"input_tokens_details": {
"image_tokens": 0,
"text_tokens": 0
},
"output_tokens": 0,
"output_tokens_details": {
"image_tokens": 0,
"text_tokens": 0
},
"total_tokens": 0
},
"web_search": {
"num_requests": 0
}
}
}
},
"expected": {
"spend": 0.008988,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412,
"tool_usage_cost": 0.0
}
},
{
"name": "gpt-5.6-responses_web_search_reported_count_mixed_actions_stream",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
"model": "gpt-5.6",
"endpoint": "/v1/responses",
"request": {
"model": "$MODEL",
"input": "search this text",
"tools": [
{
"type": "web_search_preview",
"search_context_size": "medium"
}
],
"stream": true
},
"response": {
"content_type": "text/event-stream",
"frames": [
"event: response.created\ndata: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_$REQUEST_ID\",\"object\":\"response\",\"created_at\":1700000000,\"status\":\"in_progress\",\"model\":\"gpt-5.6\",\"output\":[{\"type\":\"message\",\"id\":\"msg_$REQUEST_ID\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"scripted response\",\"annotations\":[]}]}],\"usage\":null}}",
"event: response.output_item.added\ndata: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"type\":\"message\",\"id\":\"msg_$REQUEST_ID\",\"status\":\"in_progress\",\"role\":\"assistant\",\"content\":[]}}",
"event: response.content_part.added\ndata: {\"type\":\"response.content_part.added\",\"item_id\":\"msg_$REQUEST_ID\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"content_part\",\"text\":\"\"}}",
"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_$REQUEST_ID\",\"output_index\":0,\"content_index\":0,\"delta\":\"scripted \"}",
"event: response.output_text.delta\ndata: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg_$REQUEST_ID\",\"output_index\":0,\"content_index\":0,\"delta\":\"response\"}",
"event: response.output_text.done\ndata: {\"type\":\"response.output_text.done\",\"item_id\":\"msg_$REQUEST_ID\",\"output_index\":0,\"content_index\":0,\"text\":\"scripted response\"}",
"event: response.content_part.done\ndata: {\"type\":\"response.content_part.done\",\"item_id\":\"msg_$REQUEST_ID\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"scripted response\",\"annotations\":[]}}",
"event: response.output_item.done\ndata: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"message\",\"id\":\"msg_$REQUEST_ID\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"scripted response\",\"annotations\":[]}]}}",
"event: response.completed\ndata: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_$REQUEST_ID\",\"object\":\"response\",\"created_at\":1700000000,\"status\":\"completed\",\"model\":\"gpt-5.6\",\"output\":[{\"type\":\"web_search_call\",\"id\":\"ws_0_$REQUEST_ID\",\"status\":\"completed\",\"action\":{\"type\":\"search\",\"query\":\"scripted query\"}},{\"type\":\"web_search_call\",\"id\":\"ws_1_$REQUEST_ID\",\"status\":\"completed\",\"action\":{\"type\":\"open_page\",\"url\":\"https://scripted.example/a\"}},{\"type\":\"web_search_call\",\"id\":\"ws_2_$REQUEST_ID\",\"status\":\"completed\",\"action\":{\"type\":\"open_page\",\"url\":\"https://scripted.example/b\"}},{\"type\":\"web_search_call\",\"id\":\"ws_3_$REQUEST_ID\",\"status\":\"completed\",\"action\":{\"type\":\"find_in_page\",\"pattern\":\"release\",\"url\":\"https://scripted.example/a\"}},{\"type\":\"message\",\"id\":\"msg_$REQUEST_ID\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"scripted response\",\"annotations\":[]}] }],\"usage\":{\"input_tokens\":1840,\"output_tokens\":412,\"total_tokens\":2252,\"input_tokens_details\":{\"cached_tokens\":0},\"output_tokens_details\":{\"reasoning_tokens\":0}},\"tool_usage\":{\"image_gen\":{\"input_tokens\":0,\"input_tokens_details\":{\"image_tokens\":0,\"text_tokens\":0},\"output_tokens\":0,\"output_tokens_details\":{\"image_tokens\":0,\"text_tokens\":0},\"total_tokens\":0},\"web_search\":{\"num_requests\":1}}}}"
]
},
"expected": {
"spend": 0.021488,
"input_cost": 0.00322,
"output_cost": 0.005768,
"prompt_tokens": 1840,
"completion_tokens": 412,
"tool_usage_cost": 0.0125
}
},
{
"name": "gpt-5.3-codex-responses_file_search",
"covers": "quota_management.spend_tracking.cost_matrix.logs_cost",
@ -30925,6 +31183,47 @@
"completion_tokens": 412
}
},
{
"name": "perplexity/pplx-decider-v1-27b-decisions",
"covers": "quota_management.spend_tracking.decisions_costs",
"model": "perplexity/pplx-decider-v1-27b",
"endpoint": "/v1/decisions",
"request": {
"model": "$MODEL",
"state": {
"source": "cost-tracking"
},
"questions": {
"is_defect": {
"type": "noul",
"instructions": "Is this a defect?"
}
}
},
"response": {
"content_type": "application/json",
"body": {
"model": "pplx-decider-v1-27b",
"answers": {
"is_defect": {
"type": "noul",
"noul": 0.9
}
},
"usage": {
"input_tokens": 367,
"output_tokens": 3
}
}
},
"expected": {
"spend": 1.468e-05,
"input_cost": 1.468e-05,
"output_cost": 0.0,
"prompt_tokens": 367,
"completion_tokens": 3
}
},
{
"name": "gpt-5.6-client_disconnect_mid_stream",
"covers": "quota_management.spend_tracking.scripted_wire.client_disconnect",

View file

@ -268,6 +268,31 @@ def test_case_bills_expected_cost(gateway: Gateway, case: CostTrackingTestCase)
assert row.spend == 0, f"{case.name}: failure spend was {row.spend}"
return
assert response.is_success, f"{case.name}: proxy returned {response.status_code}: {response.text[:400]}"
if case.endpoint == "/v1/decisions":
observed: Final = JSON_OBJECT.validate_json(
httpx.get(f"{gateway.upstream_url}/__observations", timeout=5, trust_env=False).content
)
decision_observations: Final = tuple(
value
for value in observed["requests"]
if isinstance(value, dict) and value.get("path") == f"/{scenario_id}/v1/decisions"
)
assert decision_observations == (
{
"path": f"/{scenario_id}/v1/decisions",
"authorization": "Bearer sk-scripted-provider",
"body": {
"model": "pplx-decider-v1-27b",
"state": {"source": "cost-tracking"},
"questions": {
"is_defect": {
"type": "noul",
"instructions": "Is this a defect?",
}
},
},
},
)
if case.response.content_type == "text/event-stream":
_assert_stream_has_no_error(response.text)
rows: Final = poll_rows(key, len(responses) + (prior_response_id is not None))

View file

@ -1,8 +1,41 @@
import os
import uuid
from typing import Final
import httpx
from integration._support.client import Gateway, object_value
from integration._support.client import Gateway, object_value, string_value
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse
from pydantic import JsonValue
_DECISIONS_PROBE_REPLY: Final = JsonResponse(
content_type="application/json",
body={
"model": "pplx-decider-v1-27b",
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
"usage": {"input_tokens": 10, "output_tokens": 1},
},
)
_CONFIGURED_PROBE_REPLY: Final = JsonResponse(
content_type="application/json",
body={
"model": "jev-custom",
"answers": {"alive": {"type": "choice", "choice": "yes", "confidence": 0.9, "probabilities": {"yes": 0.9}}},
"usage": {"input_tokens": 10, "output_tokens": 1},
},
)
_STRANDS_PROBE_REPLY: Final = JsonResponse(
content_type="application/json",
body={
"model": "strands-decider-2B-hobson-v19",
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
"usage": {"input_tokens": 10, "output_tokens": 1},
},
)
_CONFIGURED_STATE: Final[dict[str, JsonValue]] = {"ticket": "health probe"}
_CONFIGURED_QUESTIONS: Final[dict[str, JsonValue]] = {
"alive": {"type": "choice", "criteria": {"yes": "the service answers", "no": "the service is down"}}
}
def test_health_check_of_a_model_added_through_the_api_calls_its_upstream_and_reports_it_healthy(
@ -32,3 +65,85 @@ def test_health_check_of_a_model_added_through_the_api_calls_its_upstream_and_re
assert [request["body"]["model"] for request in upstream.get("/__observations").json()["requests"]] == [
provider_model
]
def _health_report(gateway: Gateway, model: str) -> dict[str, JsonValue]:
health: Final = gateway.request("GET", "/health", params={"model": model})
assert health.status_code == 200, health.text
return health.json()
def _probes_sent_to(gateway: Gateway, handle: ScenarioHandle) -> list[tuple[str, JsonValue]]:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
requests: Final = upstream.get("/__observations").json()["requests"]
return [
(string_value(request["path"]), request["body"])
for request in map(object_value, requests)
if string_value(request["path"]).startswith(f"/{handle.scenario_id}/")
]
def test_evaluation_mode_health_check_resolves_the_mode_from_the_cost_map_and_sends_the_default_probe(
gateway: Gateway,
) -> None:
with gateway.scenario() as scenario:
handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _DECISIONS_PROBE_REPLY)
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(model="perplexity/pplx-decider-v1-27b", api_base=handle.api_base())
report: Final = _health_report(gateway, model)
assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report
assert _probes_sent_to(gateway, handle) == [
(
f"/{handle.scenario_id}/v1/decisions",
{
"model": "pplx-decider-v1-27b",
"state": os.environ.get("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"),
"questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}},
},
)
]
def test_evaluation_mode_health_check_sends_the_configured_state_and_questions(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _CONFIGURED_PROBE_REPLY)
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(
model="typesafe/jev-custom",
api_base=handle.api_base(),
model_info={
"mode": "evaluation",
"health_check_params": {"state": _CONFIGURED_STATE, "questions": _CONFIGURED_QUESTIONS},
},
)
report: Final = _health_report(gateway, model)
assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report
assert _probes_sent_to(gateway, handle) == [
(
f"/{handle.scenario_id}/v1/systemone",
{"model": "jev-custom", "state": _CONFIGURED_STATE, "questions": _CONFIGURED_QUESTIONS},
)
]
def test_evaluation_mode_health_check_of_the_self_hosted_strands_model_resolves_the_mode_from_the_cost_map(
gateway: Gateway,
) -> None:
with gateway.scenario() as scenario:
handle: Final = register_scenario(f"health-decisions-{uuid.uuid4().hex[:12]}", _STRANDS_PROBE_REPLY)
scenario.cleanups.callback(delete_scenario, handle)
model: Final = scenario.model(
model="strands_decider/strands-decider-2B-hobson-v19", api_base=handle.api_base(), api_key=None
)
report: Final = _health_report(gateway, model)
assert (report["healthy_count"], report["unhealthy_count"]) == (1, 0), report
assert _probes_sent_to(gateway, handle) == [
(
f"/{handle.scenario_id}/v1/systemone",
{
"model": "strands-decider-2B-hobson-v19",
"state": os.environ.get("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm"),
"questions": {"reachable": {"type": "noul", "instructions": "Is the service reachable?"}},
},
)
]

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,262 @@
import asyncio
import json
import re
import signal
import threading
import uuid
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
from urllib.parse import urlsplit
import httpx
import psutil
import pytest
import yaml
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "pplx-decider-v1-27b"
_CONFIG_MODEL: Final = "decisions-chaos"
_API_KEY: Final = "synthetic-decisions-key"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
_QUESTIONS: Final[dict[str, JsonValue]] = {"fine": {"type": "noul", "instructions": "Is the state fine?"}}
_ROUTES: Final = ("/v1/decisions", "/decisions")
@dataclass(frozen=True, slots=True)
class _Call:
route: str
marker: str
fail: bool
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
call_id: str
model_group: str
text: str
def _calls(count: int, *, fail: bool) -> tuple[_Call, ...]:
return tuple(
_Call(
route=_ROUTES[index % len(_ROUTES)],
marker=f"{'fail' if fail and index % 2 else 'ok'}-{uuid.uuid4().hex}",
fail=fail and index % 2 == 1,
)
for index in range(count)
)
def _marker_of(request: Request) -> str:
state: Final = _JSON_OBJECT.validate_json(request.body)["state"]
assert isinstance(state, str), request.body
return state
def _reply(request: Request) -> Reply:
marker: Final = _marker_of(request)
if marker.startswith("fail-"):
return Reply(status=500, body=json.dumps({"error": {"message": f"scripted outage {marker}"}}).encode())
answer: Final = {
"model": f"model-{marker}",
"answers": {"fine": {"type": "noul", "noul": 0.5}},
"usage": {"input_tokens": 12, "output_tokens": 1},
}
return Reply(body=json.dumps(answer).encode())
async def _send(client: httpx.AsyncClient, key: str, model: str | None, call: _Call) -> _Served:
body: Final[dict[str, JsonValue]] = {
**({"model": model} if model is not None else {}),
"state": call.marker,
"questions": _QUESTIONS,
"num_retries": 0,
}
response: Final = await client.post(call.route, json=body, headers={"Authorization": f"Bearer {key}"})
return _Served(
call=call,
status=response.status_code,
call_id=response.headers.get("x-litellm-call-id", ""),
model_group=response.headers.get("x-litellm-model-group", ""),
text=response.text,
)
async def _burst(
base_url: str, key: str, model: str | None, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client:
results: Final = await asyncio.gather(
*(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Served))
def _assert_served_its_own(served: _Served) -> None:
assert served.call_id, served.text
if served.call.fail:
assert served.status == 500, (served.status, served.text)
assert served.call.marker in served.text, served.text
return
assert served.status == 200, (served.status, served.text)
assert _JSON_OBJECT.validate_json(served.text)["model"] == f"model-{served.call.marker}", served.text
def _statuses_by_call_id(call_ids: tuple[str, ...]) -> dict[str, JsonValue]:
placeholders: Final = ", ".join("%s" for _ in call_ids)
rows: Final = eventually(
lambda: read_rows(
f'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id IN ({placeholders})', call_ids
),
lambda found: len(found) >= len(call_ids),
seconds=70,
)
assert len(rows) == len(call_ids), rows
return {str(row["request_id"]): row["status"] for row in rows}
def _expected_statuses(served: tuple[_Served, ...]) -> dict[str, JsonValue]:
return {item.call_id: "failure" if item.call.fail else "success" for item in served}
def _health(gateway: Gateway, model: str) -> tuple[int, int]:
health: Final = gateway.request("GET", "/health", params={"model": model})
assert health.status_code in (200, 503), health.text
report: Final = health.json()
return (report["healthy_count"], report["unhealthy_count"])
def _chaos_config(wire: Wire, tmp_path: Path) -> Path:
config: Final = {
**_JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())),
"model_list": [
{
"model_name": _CONFIG_MODEL,
"litellm_params": {"model": f"perplexity/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY},
}
],
}
path: Final = tmp_path / "decisions-chaos.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]:
text: Final = log.read_text()
return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.")
def _open_upstream_connections(pid: int, upstream: str) -> int:
port: Final = urlsplit(upstream).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
@pytest.mark.timeout(180)
async def test_burst_over_both_routes_bills_each_call_once_with_its_own_status(gateway: Gateway) -> None:
calls: Final = _calls(30, fail=True)
with wire_server(_reply) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"perplexity/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
assert len(served) == 30
for item in served:
_assert_served_its_own(item)
assert len({item.call_id for item in served}) == 30
assert _statuses_by_call_id(tuple(item.call_id for item in served)) == _expected_statuses(served)
received: Final = wire.drain()
assert sorted(_marker_of(request) for request in received) == sorted(call.marker for call in calls)
assert {request.target for request in received} == {"/v1/decisions"}, received
@pytest.mark.timeout(420)
async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_the_default_model(
gateway: Gateway, tmp_path: Path
) -> None:
calls: Final = _calls(20, fail=False)
release: Final = threading.Event()
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
def held(request: Request) -> Reply:
held_markers.put(_marker_of(request))
assert release.wait(timeout=60), "The burst was never released"
return _reply(request)
with wire_server(held) as wire:
path: Final = _chaos_config(wire, tmp_path)
with owned_proxy_process(
gateway, tmp_path, {}, config=path, workers=2, extra_arguments=("--model", _CONFIG_MODEL)
) as owned:
candidate: Final = owned.gateway
base_url: Final = str(candidate.client.base_url)
workers, _ = eventually(
lambda: _worker_startups(owned.log),
lambda found: len(found[0]) == 2 and found[1] == 2,
seconds=120,
)
burst: Final = asyncio.create_task(
_burst(base_url, candidate.key, None, calls, tolerate_transport_errors=True)
)
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60)
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers})
assert sum(held_by.values()) == 20, held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
victim: Final = psutil.Process(victim_pid)
victim.suspend()
victim.send_signal(signal.SIGKILL)
release.set()
served: Final = await burst
assert len(served) == held_by[survivor_pid], (held_by, len(served))
follow_up: Final = _Call(route="/decisions", marker=f"ok-{uuid.uuid4().hex}", fail=False)
(answered,) = await _burst(base_url, candidate.key, None, (follow_up,))
await asyncio.to_thread(
eventually,
lambda: _worker_startups(owned.log),
lambda found: len(found[0]) == 3 and found[1] == 3,
180,
)
for item in (*served, answered):
_assert_served_its_own(item)
assert item.model_group == _CONFIG_MODEL, item.model_group
call_ids: Final = tuple(item.call_id for item in (*served, answered))
assert set(_statuses_by_call_id(call_ids).values()) == {"success"}
assert len({_marker_of(request) for request in wire.drain()}) == 21
@pytest.mark.timeout(180)
async def test_upstream_outage_fails_its_calls_and_recovery_on_the_same_port_restores_them(gateway: Gateway) -> None:
base_url: Final = str(gateway.client.base_url)
with gateway.scenario() as scenario:
with wire_server(_reply) as wire:
port: Final = urlsplit(wire.url).port
assert port is not None
model: Final = scenario.model(model=f"perplexity/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
before: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False))
assert _health(gateway, model) == (1, 0)
during: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False))
assert [item.status for item in during] == [500] * 5, [item.text for item in during]
assert _health(gateway, model) == (0, 1)
with wire_server(_reply, port=port):
after: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False))
assert _health(gateway, model) == (1, 0)
for item in (*before, *after):
_assert_served_its_own(item)
statuses: Final = _statuses_by_call_id(tuple(item.call_id for item in (*before, *during, *after)))
assert statuses == {
**{item.call_id: "success" for item in (*before, *after)},
**{item.call_id: "failure" for item in during},
}

View file

@ -0,0 +1,456 @@
import json
import math
import socket
import uuid
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
from integration.cost_calculation.cost_tracking_case import JsonResponse
from pydantic import JsonValue, TypeAdapter
import litellm
_API_KEY: Final = "synthetic-decisions-key"
_ENV_KEY: Final = "synthetic-decisions-env-key"
_PASS_THROUGH_MODEL: Final = "gpt-6-luna"
_PASS_THROUGH_AUTHORIZATION: Final = "Bearer customer-held-upstream-key"
_PASS_THROUGH_NEIGHBOUR: Final = "decisions-beside-a-pass-through"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_USAGE: Final[dict[str, JsonValue]] = {"input_tokens": 367, "output_tokens": 3}
_STATE: Final[dict[str, JsonValue]] = {"ticket": "The export job hangs at 99%", "component": "billing"}
_QUESTIONS: Final[dict[str, JsonValue]] = {
"defect": {"type": "noul", "instructions": "Is this a defect?"},
"severity": {"type": "choice", "criteria": {"low": "cosmetic", "high": "blocks users"}, "weight": 2},
"confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]},
}
_ANSWERS: Final[dict[str, JsonValue]] = {
"defect": {"type": "noul", "noul": 0.93},
"severity": {"type": "choice", "choice": "high", "confidence": 0.8, "probabilities": {"low": 0.2, "high": 0.8}},
"confidence": {
"type": "score",
"score": 1.0,
"confidence": 0.7,
"legend": {"0": "unsure", "1": "sure"},
"probabilities": {"0": 0.3, "1": 0.7},
},
}
_CHAT_BODY: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": "hi"}]}
_CHAT_REPLY: Final[dict[str, JsonValue]] = {
"id": "chatcmpl-decisions-parity",
"object": "chat.completion",
"created": 1700000000,
"model": "pplx-decider-v1-27b",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
}
_SPEND_QUERY: Final = (
"SELECT spend, status, call_type, model_group, custom_llm_provider, api_base, prompt_tokens, completion_tokens, "
'request_tags FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
)
@dataclass(frozen=True, slots=True)
class _Provider:
name: str
model: str
path: str
body_model: str
api_key: str | None
wraps_result: bool
cost_map_key: str | None
_PROVIDERS: Final = (
_Provider(
"perplexity",
"perplexity/pplx-decider-v1-27b",
"/v1/decisions",
"pplx-decider-v1-27b",
_API_KEY,
False,
"perplexity/pplx-decider-v1-27b",
),
_Provider("typesafe", "typesafe/jev-1.13.0", "/v1/systemone", "jev-1.13.0", _API_KEY, False, "typesafe/jev-1.13.0"),
_Provider(
"openrouter",
"openrouter/typesafe/jev-1.13",
"/alpha/decisions",
"typesafe/jev-1.13",
_API_KEY,
False,
"openrouter/typesafe/jev-1.13",
),
_Provider(
"strands_decider", "strands_decider/systemone-decider", "/v1/systemone", "systemone-decider", None, False, None
),
_Provider(
"cloudflare",
"cloudflare/clef",
"/ai/run/@cf/cloudflare/clef",
"clef",
_API_KEY,
True,
"cloudflare/@cf/cloudflare/clef",
),
)
_PERPLEXITY: Final = _PROVIDERS[0]
_INVALID_BODIES: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = (
("missing questions", {"state": _STATE}),
("missing state", {"questions": _QUESTIONS}),
("numeric state", {"state": 5, "questions": _QUESTIONS}),
("empty questions", {"state": _STATE, "questions": {}}),
("noul without instructions or criteria", {"state": _STATE, "questions": {"q": {"type": "noul"}}}),
("choice without criteria", {"state": _STATE, "questions": {"q": {"type": "choice", "criteria": {}}}}),
(
"score with eleven criteria",
{"state": _STATE, "questions": {"q": {"type": "score", "criteria": [f"level-{index}" for index in range(11)]}}},
),
("unknown question type", {"state": _STATE, "questions": {"q": {"type": "ranking", "criteria": ["a"]}}}),
)
def _number(value: JsonValue) -> float:
assert isinstance(value, (int, float)) and not isinstance(value, bool), value
return float(value)
def _expected_spend(cost_map_key: str | None) -> float:
if cost_map_key is None:
return 0.0
prices: Final = object_value(json.loads(Path("model_prices_and_context_window.json").read_text())[cost_map_key])
return _number(_USAGE["input_tokens"]) * _number(prices["input_cost_per_token"]) + _number(
_USAGE["output_tokens"]
) * _number(prices["output_cost_per_token"])
def _answer_body(provider: _Provider) -> dict[str, JsonValue]:
answer: Final[dict[str, JsonValue]] = {"model": provider.body_model, "answers": _ANSWERS, "usage": _USAGE}
return {"result": answer, "success": True} if provider.wraps_result else answer
def _register(scenario: Scenario, body: dict[str, JsonValue], *, status: int = 200) -> ScenarioHandle:
handle: Final = register_scenario(
f"decisions-{uuid.uuid4().hex[:12]}", JsonResponse(content_type="application/json", body=body, status=status)
)
scenario.cleanups.callback(delete_scenario, handle)
return handle
def _deployment(scenario: Scenario, handle: ScenarioHandle, provider: _Provider) -> str:
return scenario.model(model=provider.model, api_base=handle.api_base(), api_key=provider.api_key)
def _decide(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response:
return gateway.request(
"POST", "/v1/decisions", {"model": model, "state": _STATE, "questions": _QUESTIONS, **extra}, key=key
)
def _chat(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response:
return gateway.request("POST", "/v1/chat/completions", {"model": model, **_CHAT_BODY, **extra}, key=key)
def _observed_requests(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]:
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
return tuple(map(object_value, upstream.get("/__observations").json()["requests"]))
def _calls_to(requests: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> list[dict[str, JsonValue]]:
return [request for request in requests if string_value(request["path"]).startswith(f"/{handle.scenario_id}/")]
def _upstream_calls(gateway: Gateway, handle: ScenarioHandle) -> list[dict[str, JsonValue]]:
return _calls_to(_observed_requests(gateway), handle)
def _spend_row(call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(lambda: read_rows(_SPEND_QUERY, (call_id,)), lambda found: len(found) == 1, seconds=70)
return rows[0]
def _free_closed_port() -> int:
with socket.socket() as probe:
probe.bind(("127.0.0.1", 0))
return probe.getsockname()[1]
def _pass_through_config(directory: Path, pass_through_target: str, native_api_base: str) -> Path:
base: Final = _JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
config: Final = {
**base,
"general_settings": {
**object_value(base["general_settings"]),
"pass_through_endpoints": [
{
"path": "/v1/decisions",
"target": pass_through_target,
"headers": {"Authorization": _PASS_THROUGH_AUTHORIZATION},
}
],
},
"model_list": [
{
"model_name": _PASS_THROUGH_NEIGHBOUR,
"litellm_params": {"model": _PERPLEXITY.model, "api_base": native_api_base, "api_key": _API_KEY},
}
],
}
path: Final = directory / "decisions-pass-through.yaml"
path.write_text(yaml.safe_dump(config))
return path
@pytest.mark.parametrize("provider", _PROVIDERS, ids=lambda provider: provider.name)
def test_each_provider_gets_its_own_path_key_and_body_and_is_billed_from_the_cost_map(
gateway: Gateway, provider: _Provider
) -> None:
expected_spend: Final = _expected_spend(provider.cost_map_key)
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(provider))
model: Final = _deployment(scenario, handle, provider)
response: Final = _decide(gateway, model)
assert response.status_code == 200, response.text
assert response.json() == {"model": provider.body_model, "answers": _ANSWERS, "usage": _USAGE}
assert response.headers["x-litellm-model-group"] == model
assert math.isclose(float(response.headers.get("x-litellm-response-cost", "0")), expected_spend, rel_tol=1e-9)
(call,) = _upstream_calls(gateway, handle)
assert call["path"] == f"/{handle.scenario_id}{provider.path}"
assert call["authorization"] == (f"Bearer {provider.api_key}" if provider.api_key else "")
assert call["body"] == {"model": provider.body_model, "state": _STATE, "questions": _QUESTIONS}
row: Final = _spend_row(response.headers["x-litellm-call-id"])
assert (
row["status"],
row["call_type"],
row["custom_llm_provider"],
row["model_group"],
row["api_base"],
row["prompt_tokens"],
row["completion_tokens"],
) == ("success", "adecisions", provider.name, model, f"{handle.api_base()}{provider.path}", 367, 3)
assert math.isclose(_number(row["spend"]), expected_spend, rel_tol=1e-9), row
def test_repeated_identical_requests_each_reach_the_upstream_and_are_each_billed(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
model: Final = _deployment(scenario, handle, _PERPLEXITY)
responses: Final = tuple(_decide(gateway, model) for _ in range(2))
assert [response.status_code for response in responses] == [200, 200], [r.text for r in responses]
call_ids: Final = tuple(response.headers["x-litellm-call-id"] for response in responses)
assert len(set(call_ids)) == 2, call_ids
assert len(_upstream_calls(gateway, handle)) == 2
for call_id in call_ids:
assert _spend_row(call_id)["status"] == "success"
async def test_sdk_sync_and_async_clients_send_the_same_request(gateway: Gateway) -> None:
provider: Final = _PROVIDERS[1]
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(provider))
synchronous: Final = litellm.decisions(
model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY
)
asynchronous: Final = await litellm.adecisions(
model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY
)
for response in (synchronous, asynchronous):
assert response.model_dump(mode="json") == {
"model": provider.body_model,
"answers": _ANSWERS,
"usage": _USAGE,
}
calls: Final = _upstream_calls(gateway, handle)
assert len(calls) == 2, calls
for call in calls:
assert call["path"] == f"/{handle.scenario_id}{provider.path}"
assert call["authorization"] == f"Bearer {_API_KEY}"
assert call["body"] == {"model": provider.body_model, "state": _STATE, "questions": _QUESTIONS}
def test_gateway_only_fields_stay_at_the_gateway_and_tags_reach_the_spend_log(gateway: Gateway) -> None:
tag: Final = f"decisions-audit-{uuid.uuid4().hex[:8]}"
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
model: Final = _deployment(scenario, handle, _PERPLEXITY)
response: Final = _decide(
gateway, model, user="auditor", num_retries=0, temperature=0.2, metadata={"tags": [tag]}
)
assert response.status_code == 200, response.text
(call,) = _upstream_calls(gateway, handle)
assert call["body"] == {"model": _PERPLEXITY.body_model, "state": _STATE, "questions": _QUESTIONS}
row: Final = _spend_row(response.headers["x-litellm-call-id"])
tags: Final = row["request_tags"]
assert isinstance(tags, list) and tag in tags, row
def test_invalid_bodies_are_refused_at_the_gateway_without_an_upstream_call(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
model: Final = _deployment(scenario, handle, _PERPLEXITY)
for label, body in _INVALID_BODIES:
response: Final = gateway.request("POST", "/v1/decisions", {"model": model, **body})
assert response.status_code == 400, (label, response.text)
assert "Invalid Decisions request" in response.text, (label, response.text)
assert _upstream_calls(gateway, handle) == []
def test_unknown_model_is_refused_like_chat(gateway: Gateway) -> None:
model: Final = f"missing-{uuid.uuid4().hex}"
decisions: Final = _decide(gateway, model)
chat: Final = _chat(gateway, model)
assert 400 <= decisions.status_code < 500, decisions.text
assert decisions.status_code == chat.status_code, (decisions.text, chat.text)
def test_key_checks_match_chat(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
model: Final = _deployment(scenario, handle, _PERPLEXITY)
anonymous: Final = gateway.client.post(
"/v1/decisions", json={"model": model, "state": _STATE, "questions": _QUESTIONS}
)
assert anonymous.status_code == 401, anonymous.text
restricted: Final = scenario.key(models=[f"other-{uuid.uuid4().hex}"])
refused: Final = _decide(gateway, model, key=restricted)
assert 400 <= refused.status_code < 500, refused.text
assert refused.status_code == _chat(gateway, model, key=restricted).status_code, refused.text
assert _upstream_calls(gateway, handle) == []
spender: Final = scenario.key(max_budget=1e-06)
first: Final = _decide(gateway, model, key=spender)
assert first.status_code == 200, first.text
blocked: Final = eventually(
lambda: _decide(gateway, model, key=spender), lambda response: response.status_code != 200, seconds=70
)
assert 400 <= blocked.status_code < 500, blocked.text
assert blocked.status_code == _chat(gateway, model, key=spender).status_code, blocked.text
def test_request_body_api_base_is_refused_like_chat_without_an_upstream_call(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
model: Final = _deployment(scenario, handle, _PERPLEXITY)
decisions: Final = _decide(gateway, model, api_base=f"http://127.0.0.1:{_free_closed_port()}")
chat: Final = _chat(gateway, model, api_base=f"http://127.0.0.1:{_free_closed_port()}")
assert 400 <= decisions.status_code < 500, decisions.text
assert decisions.status_code == chat.status_code, (decisions.text, chat.text)
assert _upstream_calls(gateway, handle) == []
def test_a_deployment_without_a_key_sends_the_provider_env_key_to_its_configured_api_base(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
model: Final = scenario.model(model=_PERPLEXITY.model, api_base=handle.api_base(), api_key=None)
response: Final = _decide(gateway, model)
assert response.status_code == 200, response.text
(call,) = _upstream_calls(gateway, handle)
assert (call["path"], call["authorization"]) == (f"/{handle.scenario_id}/v1/decisions", f"Bearer {_ENV_KEY}")
def test_a_deployment_opted_into_client_api_base_sends_decisions_and_chat_to_the_body_api_base(
gateway: Gateway,
) -> None:
with gateway.scenario() as scenario:
configured: Final = _register(scenario, _answer_body(_PERPLEXITY))
decisions_target: Final = _register(scenario, _answer_body(_PERPLEXITY))
chat_target: Final = _register(scenario, _CHAT_REPLY)
model: Final = scenario.model(
model=_PERPLEXITY.model,
api_base=configured.api_base(),
api_key=_API_KEY,
configurable_clientside_auth_params=["api_base"],
)
decisions: Final = _decide(gateway, model, api_base=decisions_target.api_base())
chat: Final = _chat(gateway, model, api_base=chat_target.api_base())
assert decisions.status_code == 200, decisions.text
assert chat.status_code == 200, chat.text
observed: Final = _observed_requests(gateway)
assert [call["path"] for call in _calls_to(observed, decisions_target)] == [
f"/{decisions_target.scenario_id}/v1/decisions"
]
assert [call["path"] for call in _calls_to(observed, chat_target)] == [
f"/{chat_target.scenario_id}/chat/completions"
]
assert _calls_to(observed, configured) == []
def test_a_config_pass_through_at_v1_decisions_keeps_answering_and_the_native_api_serves_decisions(
gateway: Gateway, tmp_path: Path
) -> None:
with gateway.scenario() as scenario:
pass_through_target: Final = _register(
scenario, {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE}
)
native_target: Final = _register(scenario, _answer_body(_PERPLEXITY))
config: Final = _pass_through_config(
tmp_path, f"{pass_through_target.api_base()}/v1/decisions", native_target.api_base()
)
with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned:
through: Final = _decide(owned.gateway, _PASS_THROUGH_MODEL)
native: Final = owned.gateway.request(
"POST", "/decisions", {"model": _PASS_THROUGH_NEIGHBOUR, "state": _STATE, "questions": _QUESTIONS}
)
assert through.status_code == 200, through.text
assert through.json() == {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE}
observed: Final = _observed_requests(gateway)
(forwarded,) = _calls_to(observed, pass_through_target)
assert (forwarded["path"], forwarded["authorization"], object_value(forwarded["body"])["model"]) == (
f"/{pass_through_target.scenario_id}/v1/decisions",
_PASS_THROUGH_AUTHORIZATION,
_PASS_THROUGH_MODEL,
)
assert native.status_code == 200, native.text
assert [call["path"] for call in _calls_to(observed, native_target)] == [
f"/{native_target.scenario_id}/v1/decisions"
]
@pytest.mark.parametrize("status", (401, 429, 500))
def test_upstream_errors_keep_their_status_and_log_an_unbilled_failure(gateway: Gateway, status: int) -> None:
marker: Final = f"scripted-{status}-{uuid.uuid4().hex[:8]}"
with gateway.scenario() as scenario:
handle: Final = _register(scenario, {"error": {"message": marker}}, status=status)
model: Final = _deployment(scenario, handle, _PERPLEXITY)
response: Final = _decide(gateway, model, num_retries=0)
assert response.status_code == status, response.text
assert marker in response.text
row: Final = _spend_row(response.headers["x-litellm-call-id"])
assert (row["status"], row["call_type"], row["model_group"], _number(row["spend"])) == (
"failure",
"adecisions",
model,
0.0,
)
def test_upstream_success_without_answers_is_a_gateway_side_server_error(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, {"model": _PERPLEXITY.body_model, "usage": _USAGE})
model: Final = _deployment(scenario, handle, _PERPLEXITY)
response: Final = _decide(gateway, model, num_retries=0)
assert 500 <= response.status_code < 600, response.text
assert "answers" in response.text
assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure"
def test_unreachable_upstream_fails_only_its_own_deployment(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
healthy: Final = _deployment(scenario, handle, _PERPLEXITY)
dead: Final = scenario.model(
model=_PERPLEXITY.model, api_base=f"http://127.0.0.1:{_free_closed_port()}", api_key=_API_KEY
)
failed: Final = _decide(gateway, dead, num_retries=0)
assert 500 <= failed.status_code < 600, failed.text
assert _spend_row(failed.headers["x-litellm-call-id"])["status"] == "failure"
served: Final = _decide(gateway, healthy)
assert served.status_code == 200, served.text
assert len(_upstream_calls(gateway, handle)) == 1

View file

@ -23,3 +23,5 @@ vector_store_registry:
api_base: os.environ/INTEGRATION_UPSTREAM_URL
api_key: integration-provider-key
vector_store_description: declared in tests/integration/proxy_config.yaml
environment_variables:
PERPLEXITYAI_API_KEY: synthetic-decisions-env-key

View file

@ -8,7 +8,7 @@ from uuid import uuid4
import pytest
import pytest_asyncio
from fastapi import HTTPException, Request
from fastapi import HTTPException, Request, Response
from fastapi.security import HTTPAuthorizationCredentials
from pydantic import TypeAdapter
@ -57,6 +57,18 @@ async def lens_database() -> AsyncIterator[PrismaClient]:
"model": "openai/lens-test-analysis",
"api_key": "test-only",
"mock_response": '{"observations":[]}',
"max_tokens": 16384,
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
},
},
{
"model_name": "lens-failing-analysis",
"litellm_params": {
"model": "openai/lens-failing-analysis",
"api_key": "test-only",
"mock_response": "litellm.RateLimitError",
"max_tokens": 16384,
"input_cost_per_token": 0.000001,
"output_cost_per_token": 0.000002,
},
@ -273,6 +285,7 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
"client": ("127.0.0.1", 1234),
}
),
response=Response(),
)
assert '"observations"' in response.content
with pytest.raises(HTTPException) as denied_ip:
@ -290,6 +303,7 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
"client": ("192.0.2.1", 1234),
}
),
response=Response(),
)
assert denied_ip.value.status_code == 403
forwarded: Final = await endpoints.model(
@ -306,6 +320,7 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
"client": ("192.0.2.100", 1234),
}
),
response=Response(),
)
assert '"observations"' in forwarded.content
with pytest.raises(HTTPException) as spoofed_chain:
@ -323,6 +338,7 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
"client": ("192.0.2.100", 1234),
}
),
response=Response(),
)
assert spoofed_chain.value.status_code == 403
charged: Final = await endpoints.get_lens(lens.id, worker.scope)
@ -387,3 +403,46 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database:
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', lens.id)
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_LensWorker" WHERE id=$1', worker.id)
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_VerificationToken" WHERE token=$1', key_id)
@pytest.mark.asyncio
async def test_failed_model_requests_release_lens_budget_reservations(lens_database: PrismaClient) -> None:
admin: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
settings: Final = LensSettings(
name="Failed billing regression", model="lens-failing-analysis", context="Verify outcomes", enabled=False
)
lens: Final = await endpoints.create_lens(settings, admin)
key_id: Final = hashlib.sha256(uuid4().bytes).hexdigest()
await lens_database.db.litellm_verificationtoken.create(data={"token": key_id, "models": [settings.model]})
registration: Final = await endpoints.register_worker(endpoints.WorkerName(analysis_key_id=key_id), admin)
worker: Final = registration.worker
try:
claimed: Final = await endpoints.claim_candidate(lens, worker, datetime.now(timezone.utc))
assert claimed is not None
for _ in range(3):
with pytest.raises(HTTPException) as failed:
await endpoints.model(
lens.id,
claimed.job.id,
ModelRequest(prompt="Return JSON", purpose="extract"),
worker,
Request(
{
"type": "http",
"scheme": "http",
"path": "/lens/worker/model",
"headers": [],
"client": ("127.0.0.1", 1234),
}
),
response=Response(),
)
assert failed.value.status_code == 429
stored: Final = await endpoints.get_lens(lens.id, worker.scope)
assert stored.spent == 0
assert stored.jobs[0].cost == 0
finally:
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_LensRun" WHERE lens_id=$1', lens.id)
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', lens.id)
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_LensWorker" WHERE id=$1', worker.id)
await lens_database.db.execute_raw('DELETE FROM "LiteLLM_VerificationToken" WHERE token=$1', key_id)

View file

@ -61,11 +61,12 @@ async def test_ingest_rejects_oversized_body_before_storage() -> None:
@pytest.mark.asyncio
async def test_reads_delegate_to_storage() -> None:
@pytest.mark.parametrize("cursor,page_size", ((None, None), ("next", 200)))
async def test_reads_delegate_to_storage(cursor: str | None, page_size: int | None) -> None:
storage: Final = _fake_storage()
scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-research",)}
assert await TraceReceiver(storage).get_trace("t1", scope) is None
storage.get_trace.assert_awaited_once_with("t1", scope, "")
assert await TraceReceiver(storage).get_trace("t1", scope, "", cursor, page_size) is None
storage.get_trace.assert_awaited_once_with("t1", scope, "", cursor, page_size)
@pytest.mark.asyncio

View file

View file

@ -0,0 +1,592 @@
from __future__ import annotations
import asyncio
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
import pytest
import respx
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.types.decisions import (
ChoiceAnswer,
DecisionsResponse,
DecisionsUsage,
NoulAnswer,
ScoreAnswer,
)
_QUESTIONS: Final[Mapping[str, object]] = MappingProxyType(
{
"is_defect": {"type": "noul", "instructions": "Is this a defect?", "provider_field": "kept"},
"sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}},
"severity": {"type": "score", "criteria": ["none", "low", "high"]},
}
)
_INPUT_TOKENS: Final[int] = 367
_OUTPUT_TOKENS: Final[int] = 3
_RESPONSE: Final[Mapping[str, object]] = {
"model": "jev-1.13",
"answers": {
"is_defect": {"type": "noul", "noul": 0.9},
"sentiment": {
"type": "choice",
"choice": "positive",
"confidence": 0.8,
"probabilities": {"positive": 0.8, "negative": 0.2},
},
"severity": {
"type": "score",
"score": 1,
"confidence": 0.7,
"legend": {"0": "none", "1": "low", "2": "high"},
"probabilities": {"0": 0.1, "1": 0.8, "2": 0.1},
},
},
"usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS},
}
_STRANDS_RESPONSE: Final[Mapping[str, object]] = {
"model": "strands-decider-2B-hobson-v19",
"answers": {
"severity": {
"type": "score",
"score": 1,
"confidence": 0.7,
"legend": {"0": "none", "1": "low", "2": "high"},
"probabilities": {"0": 0.1, "1": 0.8, "2": 0.1},
}
},
"usage": {"input_tokens": 216, "output_tokens": 3},
"latency_ms": 3722.17,
}
_PROVIDERS: Final[tuple[tuple[str, str, str, str], ...]] = (
(
"perplexity",
"perplexity/pplx-decider-v1-27b",
"https://api.perplexity.ai/v1/decisions",
"pplx-decider-v1-27b",
),
("typesafe", "typesafe/jev-1.13", "https://api.typesafe.ai/v1/systemone", "jev-1.13"),
(
"openrouter",
"openrouter/typesafe/jev-1.13",
"https://openrouter.ai/api/alpha/decisions",
"typesafe/jev-1.13",
),
)
class _RecordingLogger(CustomLogger):
def __init__(self) -> None:
super().__init__()
self.standard_logging_object: Mapping[str, object] | None = None
async def async_log_success_event(
self,
kwargs: Mapping[str, object],
response_obj: object,
start_time: object,
end_time: object,
) -> None:
standard_logging_object: Final = kwargs.get("standard_logging_object")
if isinstance(standard_logging_object, dict):
self.standard_logging_object = standard_logging_object
async def _drain_logging_worker() -> None:
await asyncio.sleep(0)
GLOBAL_LOGGING_WORKER.start()
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
@pytest.fixture(autouse=True)
def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
@pytest.mark.parametrize(("provider", "model", "url", "upstream_model"), _PROVIDERS)
async def test_adecisions_sends_the_provider_wire_contract(
provider: str,
model: str,
url: str,
upstream_model: str,
respx_mock: respx.MockRouter,
) -> None:
route: Final = respx_mock.post(url).respond(json=_RESPONSE)
response: Final = await litellm.adecisions(
model=model,
state={"source": "unit-test"},
questions=_QUESTIONS,
api_key="caller-key",
extra_headers={
"x-request-tag": "decisions-test",
"AUTHORIZATION": "attacker-key",
"Content-Type": "text/plain",
},
internal_kwarg="must-not-leak",
)
assert route.called
assert len(respx_mock.calls) == 1
request: Final = respx_mock.calls[0].request
assert request.headers["authorization"] == "Bearer caller-key"
assert request.headers["content-type"] == "application/json"
assert request.headers["x-request-tag"] == "decisions-test"
assert json.loads(request.content) == {
"model": upstream_model,
"state": {"source": "unit-test"},
"questions": {
"is_defect": {
"type": "noul",
"instructions": "Is this a defect?",
"provider_field": "kept",
},
"sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}},
"severity": {"type": "score", "criteria": ["none", "low", "high"]},
},
}
assert isinstance(response.answers["is_defect"], NoulAnswer)
assert isinstance(response.answers["sentiment"], ChoiceAnswer)
assert isinstance(response.answers["severity"], ScoreAnswer)
assert response._hidden_params["custom_llm_provider"] == provider
@pytest.mark.asyncio
async def test_router_dispatches_typesafe_decisions_without_api_base(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.delenv("TYPESAFE_API_KEY", raising=False)
monkeypatch.delenv("TYPESAFE_API_BASE", raising=False)
provider_resolution: Final = litellm.get_llm_provider("typesafe/jev-latest")
assert provider_resolution[:2] == ("jev-latest", "typesafe")
router: Final = litellm.Router(
model_list=[
{
"model_name": "jev",
"litellm_params": {
"model": "typesafe/jev-latest",
"api_key": "k",
},
}
]
)
upstream: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE)
response: Final = await router.adecisions(
model="jev",
state="router-test",
questions={
"sentiment": {
"type": "choice",
"criteria": {"positive": None, "negative": "unhappy"},
}
},
)
assert upstream.called
assert len(respx_mock.calls) == 1
assert json.loads(respx_mock.calls[0].request.content) == {
"model": "jev-latest",
"state": "router-test",
"questions": {
"sentiment": {
"type": "choice",
"criteria": {"positive": None, "negative": "unhappy"},
}
},
}
assert respx_mock.calls[0].request.headers["authorization"] == "Bearer k"
assert isinstance(response.answers["sentiment"], ChoiceAnswer)
assert response.answers["sentiment"].choice == "positive"
def test_decisions_uses_the_same_wire_contract_for_sync_calls(respx_mock: respx.MockRouter) -> None:
route: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE)
response: Final = litellm.decisions(
model="perplexity/pplx-decider-v1-27b",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
api_key="caller-key",
)
assert route.called
assert response.model == "jev-1.13"
def test_openrouter_response_keeps_provider_fields(respx_mock: respx.MockRouter) -> None:
payload: Final = {
**_RESPONSE,
"id": "decision-1",
"provider": "typesafe",
"usage": {**_RESPONSE["usage"], "cost": 0.25},
}
respx_mock.post("https://openrouter.ai/api/alpha/decisions").respond(json=payload)
response: Final = litellm.decisions(
model="openrouter/typesafe/jev-1.13",
state="review",
questions=_QUESTIONS,
api_key="caller-key",
)
assert response.model_extra["id"] == "decision-1"
assert response.model_extra["provider"] == "typesafe"
assert response.usage is not None
assert response.usage.model_extra["cost"] == 0.25
def test_decisions_cost_uses_litellm_token_pricing() -> None:
response: Final = DecisionsResponse(
model="pplx-decider-v1-27b",
answers={},
usage=DecisionsUsage(input_tokens=_INPUT_TOKENS, output_tokens=_OUTPUT_TOKENS),
)
response._hidden_params = {
"model": "perplexity/pplx-decider-v1-27b",
"custom_llm_provider": "perplexity",
}
cost: Final = litellm.completion_cost(completion_response=response)
perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"]
expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float(
perplexity_cost["output_cost_per_token"]
)
assert expected_cost > 0
assert cost == pytest.approx(expected_cost)
@pytest.mark.asyncio
async def test_decisions_cost_is_in_standard_logging_object(respx_mock: respx.MockRouter) -> None:
respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE)
recording_logger: Final = _RecordingLogger()
original_callbacks: Final = litellm.callbacks
litellm.callbacks = [recording_logger]
try:
await litellm.adecisions(
model="perplexity/pplx-decider-v1-27b",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
api_key="caller-key",
)
await _drain_logging_worker()
finally:
litellm.callbacks = original_callbacks
assert recording_logger.standard_logging_object is not None
perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"]
expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float(
perplexity_cost["output_cost_per_token"]
)
assert expected_cost > 0
assert recording_logger.standard_logging_object["response_cost"] == pytest.approx(expected_cost)
assert recording_logger.standard_logging_object["prompt_tokens"] == _INPUT_TOKENS
assert recording_logger.standard_logging_object["completion_tokens"] == _OUTPUT_TOKENS
@pytest.mark.asyncio
async def test_unknown_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None:
with pytest.raises(litellm.BadRequestError, match="Supported providers"):
await litellm.adecisions(
model="unknown/jev-1.13",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
api_key="caller-key",
)
assert len(respx_mock.calls) == 0
@pytest.mark.asyncio
async def test_empty_custom_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None:
with pytest.raises(litellm.BadRequestError, match="Supported providers"):
await litellm.adecisions(
model="perplexity/pplx-decider-v1-27b",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
api_key="caller-key",
custom_llm_provider="",
)
assert len(respx_mock.calls) == 0
@pytest.mark.asyncio
async def test_invalid_question_is_rejected_before_http(respx_mock: respx.MockRouter) -> None:
with pytest.raises(litellm.BadRequestError, match="Invalid Decisions request"):
await litellm.adecisions(
model="perplexity/pplx-decider-v1-27b",
state="review",
questions={"sentiment": {"type": "choice"}},
api_key="caller-key",
)
assert len(respx_mock.calls) == 0
def test_upstream_bad_request_maps_to_litellm_error(respx_mock: respx.MockRouter) -> None:
respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(
status_code=400,
json={"error": {"message": "invalid decision"}},
)
with pytest.raises(litellm.BadRequestError):
litellm.decisions(
model="perplexity/pplx-decider-v1-27b",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
api_key="caller-key",
)
def test_server_key_is_sent_to_an_explicit_api_base(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setenv("PERPLEXITYAI_API_KEY", "server-key")
monkeypatch.delenv("PERPLEXITY_API_KEY", raising=False)
route: Final = respx_mock.post("https://egress.example/perplexity/v1/decisions").respond(json=_RESPONSE)
litellm.decisions(
model="perplexity/pplx-decider-v1-27b",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
api_base="https://egress.example/perplexity",
)
assert route.call_count == 1
assert route.calls[0].request.headers["authorization"] == "Bearer server-key"
@pytest.mark.asyncio
@pytest.mark.parametrize("model", ("cloudflare/clef", "cloudflare/@cf/cloudflare/clef"))
@pytest.mark.parametrize("wrapped", (False, True))
async def test_cloudflare_clef_resolves_model_and_response_envelope(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
model: str,
wrapped: bool,
) -> None:
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
response_body: Final[Mapping[str, object]] = (
{"result": _RESPONSE, "success": True, "errors": [], "messages": []} if wrapped else _RESPONSE
)
route: Final = respx_mock.post(
"https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef"
).respond(json=response_body)
response: Final = await litellm.adecisions(
model=model,
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
)
assert route.called
request: Final = respx_mock.calls[0].request
assert request.headers["authorization"] == "Bearer cloudflare-key"
assert json.loads(request.content) == {
"model": "clef",
"state": "review",
"questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
}
assert response.answers == DecisionsResponse.model_validate(_RESPONSE).answers
assert response._hidden_params["model"] == "cloudflare/@cf/cloudflare/clef"
@pytest.mark.asyncio
async def test_cloudflare_clef_flash_uses_flash_endpoint_and_request_model(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
route: Final = respx_mock.post(
"https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef-flash"
).respond(json=_RESPONSE)
await litellm.adecisions(
model="cloudflare/clef-flash",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
)
assert route.called
assert json.loads(respx_mock.calls[0].request.content)["model"] == "clef-flash"
@pytest.mark.asyncio
async def test_cloudflare_api_base_from_env_uses_workers_ai_run_path(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setenv("CLOUDFLARE_API_BASE", "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1")
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
route: Final = respx_mock.post(
"https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef"
).respond(json=_RESPONSE)
await litellm.adecisions(
model="cloudflare/clef",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
)
assert route.called
@pytest.mark.asyncio
async def test_cloudflare_requires_account_id_or_api_base_before_http(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
with pytest.raises(litellm.BadRequestError, match="Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID"):
await litellm.adecisions(
model="cloudflare/clef",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
)
assert len(respx_mock.calls) == 0
@pytest.mark.asyncio
async def test_cloudflare_clef_cost_uses_the_model_cost_map(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
respx_mock.post("https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef").respond(
json=_RESPONSE
)
response: Final = await litellm.adecisions(
model="cloudflare/clef",
state="review",
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
)
cost: Final = litellm.completion_cost(completion_response=response)
clef_cost: Final = litellm.model_cost["cloudflare/@cf/cloudflare/clef"]
expected_cost: Final = _INPUT_TOKENS * float(clef_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float(
clef_cost["output_cost_per_token"]
)
assert expected_cost > 0
assert cost == pytest.approx(expected_cost)
@pytest.mark.asyncio
async def test_strands_decider_requires_api_base_before_http(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
with pytest.raises(litellm.BadRequestError, match="api_base is required"):
await litellm.adecisions(
model="strands_decider/strands-decider-2B-hobson-v19",
state="review",
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
)
assert len(respx_mock.calls) == 0
@pytest.mark.asyncio
async def test_strands_decider_without_key_preserves_response_extras(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE)
response: Final = await litellm.adecisions(
model="strands_decider/strands-decider-2B-hobson-v19",
state="review",
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
api_base="https://strands.example",
)
assert route.called
assert "authorization" not in respx_mock.calls[0].request.headers
assert response.model_extra["latency_ms"] == _STRANDS_RESPONSE["latency_ms"]
severity: Final = response.answers["severity"]
assert isinstance(severity, ScoreAnswer)
assert severity.legend == {"0": "none", "1": "low", "2": "high"}
@pytest.mark.asyncio
async def test_strands_decider_uses_key_from_matching_environment_base(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setenv("STRANDS_DECIDER_API_BASE", "https://strands.example")
monkeypatch.setenv("STRANDS_DECIDER_API_KEY", "strands-key")
route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE)
await litellm.adecisions(
model="strands_decider/strands-decider-2B-hobson-v19",
state="review",
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
api_base="https://strands.example",
)
assert route.called
assert respx_mock.calls[0].request.headers["authorization"] == "Bearer strands-key"
@pytest.mark.asyncio
async def test_strands_decider_provider_resolution_and_router_dispatch(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
provider_resolution: Final = litellm.get_llm_provider("strands_decider/strands-decider-2B-hobson-v19")
router: Final = litellm.Router(
model_list=[
{
"model_name": "strands",
"litellm_params": {
"model": "strands_decider/strands-decider-2B-hobson-v19",
"api_base": "https://strands.example",
},
}
]
)
route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE)
response: Final = await router.adecisions(
model="strands",
state="review",
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
)
assert provider_resolution[:2] == ("strands-decider-2B-hobson-v19", "strands_decider")
assert route.called
assert response.model == _STRANDS_RESPONSE["model"]

View file

@ -0,0 +1,476 @@
"""Tests for the in-process Tool Loop handler."""
from __future__ import annotations
from collections.abc import Callable, Iterator, Mapping
from pathlib import Path
from typing import Final, Literal
import litellm
import pytest
from pydantic import BaseModel
from litellm import sandbox
from litellm.harness.context import GatewayTarget, SessionContext
from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler
from litellm.harness.options import ToolLoopOptions
from litellm.harness.types import Approval, Event, Harness, Text, ToolCall, ToolResult
from litellm.llms.base_llm.harness.transformation import HarnessTurnError
from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig
from litellm.types.utils import ModelResponse
class ScriptedCompletion:
def __init__(self, responses: tuple[ModelResponse, ...]) -> None:
self.responses: Iterator[ModelResponse] = iter(responses)
self.calls: list[dict[str, object]] = [] # mutable-ok: captures injected completion requests
async def __call__(self, **kwargs: object) -> ModelResponse:
self.calls.append(dict(kwargs))
return next(self.responses)
def model_response(
*,
content: str | None = None,
tool_calls: tuple[dict[str, object], ...] = (),
prompt_tokens: int = 0,
completion_tokens: int = 0,
hidden_params: dict[str, object] | None = None,
) -> ModelResponse:
response: Final = ModelResponse(
model="gpt-test",
choices=[
{
"message": {
"role": "assistant",
"content": content,
"tool_calls": list(tool_calls) if tool_calls else None,
}
}
],
usage={
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
)
if hidden_params is not None:
response._hidden_params = hidden_params
return response
def function_call(name: str, arguments: str, call_id: str = "call-1") -> dict[str, object]:
return {
"id": call_id,
"type": "function",
"function": {"name": name, "arguments": arguments},
}
def make_context(
tmp_path: Path,
*,
model: str | None = "gpt-4o-mini",
gateway: GatewayTarget | None = None,
api_key: str | None = None,
api_base: str | None = None,
instructions: str | None = None,
tools: tuple[Callable[..., object], ...] = (),
permissions: Literal["ask", "full"] = "full",
output: type[BaseModel] | None = None,
metadata: Mapping[str, object] | None = None,
options: ToolLoopOptions | None = None,
) -> SessionContext:
return SessionContext(
harness=Harness.TOOL_LOOP,
sandbox=sandbox.local(tmp_path),
session_id="tool-loop-session",
model=model,
gateway=gateway,
api_key=api_key,
api_base=api_base,
instructions=instructions,
tools=tools,
permissions=permissions,
output=output,
metadata={} if metadata is None else metadata,
options=options,
)
def make_handler(completion: ScriptedCompletion) -> ToolLoopHandler:
return ToolLoopHandler(ToolLoopHarnessConfig(), acompletion=completion)
async def run_turn(
handler: ToolLoopHandler,
ctx: SessionContext,
prompt: str,
allow: bool = True,
) -> tuple[Event, ...]:
events: list[Event] = [] # mutable-ok: gathers this async turn for assertions
async for event in handler.turn(ctx, prompt):
events.append(event)
if isinstance(event, Approval):
event.allow() if allow else event.deny("not approved")
return tuple(events)
def add(a: int, b: int) -> int:
return a + b
def duplicate_add() -> Callable[..., object]:
def add(a: int, b: int) -> int:
return a + b
return add
async def test_tool_round_trip_appends_assistant_and_tool_messages(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
model_response(content="The sum is 5"),
)
)
ctx: Final = make_context(tmp_path, instructions="Use tools when needed", tools=(add,))
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Add two and three")
assert events == (
ToolCall(
id="call-1",
name="add",
native_name="add",
input={"a": 2, "b": 3},
builtin=False,
),
ToolResult(id="call-1", output="5", is_error=False),
Text(delta="The sum is 5"),
)
assert ctx.final_text == "The sum is 5"
assert completion.calls[1]["messages"] == [
{"role": "system", "content": "Use tools when needed"},
{"role": "user", "content": "Add two and three"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call-1",
"type": "function",
"function": {"name": "add", "arguments": '{"a": 2, "b": 3}'},
}
],
},
{"role": "tool", "tool_call_id": "call-1", "content": "5"},
]
async def test_multiple_tool_calls_run_in_order(tmp_path: Path) -> None:
values: list[int] = [] # mutable-ok: records the order of calls from the injected model response
def record(value: int) -> int:
values.append(value)
return value
completion: Final = ScriptedCompletion(
(
model_response(
tool_calls=(
function_call("record", '{"value": 1}', "call-1"),
function_call("record", '{"value": 2}', "call-2"),
)
),
model_response(content="Recorded"),
)
)
ctx: Final = make_context(tmp_path, tools=(record,))
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Record both")
assert values == [1, 2]
assert tuple(event for event in events if isinstance(event, ToolResult)) == (
ToolResult(id="call-1", output="1", is_error=False),
ToolResult(id="call-2", output="2", is_error=False),
)
async def test_tool_exception_is_returned_to_model_and_loop_continues(tmp_path: Path) -> None:
def fail() -> str:
raise RuntimeError("tool failed")
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("fail", "{}"),)),
model_response(content="Recovered"),
)
)
ctx: Final = make_context(tmp_path, tools=(fail,))
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Run fail")
assert ToolResult(id="call-1", output="RuntimeError: tool failed", is_error=True) in events
assert completion.calls[1]["messages"][-1] == {
"role": "tool",
"tool_call_id": "call-1",
"content": "RuntimeError: tool failed",
}
assert ctx.final_text == "Recovered"
async def test_unknown_tool_is_returned_and_tools_are_omitted_when_empty(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("missing", "{}"),)),
model_response(content="Unknown tool handled"),
)
)
ctx: Final = make_context(tmp_path)
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Call a missing tool")
assert "tools" not in completion.calls[0]
assert (
ToolResult(
id="call-1",
output="ValueError: unknown tool 'missing'",
is_error=True,
)
in events
)
assert ctx.final_text == "Unknown tool handled"
@pytest.mark.parametrize(
"arguments",
['{"a": 2}', '{"a": 2, "b": 3, "extra": 4}'],
)
async def test_invalid_tool_arguments_return_validation_error(
tmp_path: Path,
arguments: str,
) -> None:
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("add", arguments),)),
model_response(content="Arguments were invalid"),
)
)
ctx: Final = make_context(tmp_path, tools=(add,))
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Add")
result: Final = next(event for event in events if isinstance(event, ToolResult))
assert result.is_error
assert result.output.startswith("ValidationError:")
assert completion.calls[1]["messages"][-1]["content"] == result.output
async def test_malformed_tool_arguments_return_error_without_calling_tool(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("add", "{"),)),
model_response(content="Malformed arguments"),
)
)
ctx: Final = make_context(tmp_path, tools=(add,))
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Add")
result: Final = next(event for event in events if isinstance(event, ToolResult))
assert result.is_error
assert result.output.startswith("JSONDecodeError:")
assert completion.calls[1]["messages"][-1]["content"] == result.output
async def test_ask_permission_denial_skips_tool_and_returns_reason(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
model_response(content="Denied"),
)
)
ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask")
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Add", allow=False)
assert any(isinstance(event, Approval) for event in events)
assert ToolResult(id="call-1", output="denied: not approved", is_error=True) in events
assert completion.calls[1]["messages"][-1]["content"] == "denied: not approved"
async def test_ask_permission_allow_runs_tool(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
model_response(content="Allowed"),
)
)
ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask")
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Add")
assert any(isinstance(event, Approval) for event in events)
assert ToolResult(id="call-1", output="5", is_error=False) in events
async def test_async_tool_is_awaited(tmp_path: Path) -> None:
async def multiply(a: int, b: int) -> int:
return a * b
completion: Final = ScriptedCompletion(
(
model_response(tool_calls=(function_call("multiply", '{"a": 3, "b": 4}'),)),
model_response(content="12"),
)
)
ctx: Final = make_context(tmp_path, tools=(multiply,))
handler: Final = make_handler(completion)
await handler.start(ctx)
events: Final = await run_turn(handler, ctx, "Multiply")
assert ToolResult(id="call-1", output="12", is_error=False) in events
class Answer(BaseModel):
value: int
async def test_structured_output_is_forwarded_and_retained(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion((model_response(content='{"value": 7}'),))
ctx: Final = make_context(tmp_path, output=Answer)
handler: Final = make_handler(completion)
await handler.start(ctx)
await run_turn(handler, ctx, "Return a value")
assert completion.calls[0]["response_format"] is Answer
assert ctx.output_json == '{"value": 7}'
assert ctx.final_text == '{"value": 7}'
async def test_gateway_routing_uses_proxy_model_and_tool_loop_tag(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion((model_response(content="done"),))
gateway: Final = GatewayTarget(api_base="http://gateway", api_key="sk-virtual")
ctx: Final = make_context(tmp_path, gateway=gateway, metadata={"team": "test"})
handler: Final = make_handler(completion)
await handler.start(ctx)
await run_turn(handler, ctx, "Hi")
assert completion.calls[0]["model"] == "litellm_proxy/gpt-4o-mini"
assert completion.calls[0]["api_base"] == "http://gateway"
assert completion.calls[0]["api_key"] == "sk-virtual"
headers: Final = completion.calls[0]["extra_headers"]
assert isinstance(headers, dict)
assert headers["x-litellm-tags"] == "harness,tool_loop"
assert '"team": "test"' in headers["x-litellm-spend-logs-metadata"]
async def test_usage_and_cost_accumulate_across_model_calls(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion(
(
model_response(
tool_calls=(function_call("add", '{"a": 2, "b": 3}'),),
prompt_tokens=10,
completion_tokens=4,
hidden_params={
"additional_headers": {"llm_provider-x-litellm-response-cost": "0.4"},
"response_cost": 0.1,
},
),
model_response(
content="done",
prompt_tokens=20,
completion_tokens=5,
hidden_params={"response_cost": 0.2},
),
)
)
ctx: Final = make_context(tmp_path, tools=(add,))
handler: Final = make_handler(completion)
await handler.start(ctx)
await run_turn(handler, ctx, "Add")
assert ctx.calls == 2
assert ctx.input_tokens == 30
assert ctx.output_tokens == 9
assert ctx.cost == pytest.approx(0.6)
async def test_history_survives_stop_and_start(tmp_path: Path) -> None:
completion: Final = ScriptedCompletion((model_response(content="first"), model_response(content="second")))
ctx: Final = make_context(tmp_path, instructions="Keep answers concise")
handler: Final = make_handler(completion)
await handler.start(ctx)
await run_turn(handler, ctx, "first prompt")
await handler.stop(ctx)
await handler.start(ctx)
await run_turn(handler, ctx, "second prompt")
assert completion.calls[1]["messages"] == [
{"role": "system", "content": "Keep answers concise"},
{"role": "user", "content": "first prompt"},
{"role": "assistant", "content": "first"},
{"role": "user", "content": "second prompt"},
]
history: Final = await handler.history(ctx)
history[0]["content"] = "changed"
assert (await handler.history(ctx))[0]["content"] == "Keep answers concise"
async def test_duplicate_tool_names_are_rejected(tmp_path: Path) -> None:
ctx: Final = make_context(tmp_path, tools=(add, duplicate_add()))
handler: Final = make_handler(ScriptedCompletion(()))
with pytest.raises(ValueError, match="tool names must be unique"):
await handler.start(ctx)
async def test_model_call_limit_raises_harness_turn_error(tmp_path: Path) -> None:
repeating_response: Final = model_response(tool_calls=(function_call("ping", "{}"),))
completion: Final = ScriptedCompletion((repeating_response,) * 100)
def ping() -> str:
return "pong"
ctx: Final = make_context(tmp_path, tools=(ping,))
handler: Final = make_handler(completion)
await handler.start(ctx)
with pytest.raises(HarnessTurnError, match="exceeded 100 model calls"):
await run_turn(handler, ctx, "Ping repeatedly")
async def test_public_aagent_uses_tool_loop_with_mock_response(tmp_path: Path) -> None:
result: Final = await litellm.aagent(
Harness.TOOL_LOOP,
"Say done",
sandbox=sandbox.local(tmp_path),
model="gpt-4o-mini",
options=ToolLoopOptions(completion_kwargs={"mock_response": "done"}),
)
assert result.text == "done"
assert result.stop_reason == "done"

View file

@ -35,6 +35,7 @@ PUBLIC_NAMES = [
"CodexOptions",
"OpenCodeOptions",
"DeepAgentsOptions",
"ToolLoopOptions",
"HarnessError",
"CapabilityUnsupported",
"OptionsMismatch",
@ -89,7 +90,9 @@ def test_adapter_registry_paths_cover_every_harness():
def test_litellm_agent_is_top_level_and_lazy():
code = (
"import sys, litellm; assert 'litellm.harness' not in sys.modules; "
"assert litellm.agent is litellm.harness.agent; assert litellm.Harness.CODEX.value == 'codex'"
"assert litellm.agent is litellm.harness.agent; "
"assert litellm.Harness.CODEX.value == 'codex'; "
"assert litellm.ToolLoopOptions is litellm.harness.ToolLoopOptions"
)
out = run_child_interpreter(code, timeout=120)
assert out.returncode == 0, out.stderr

View file

@ -20,6 +20,7 @@ from litellm.harness.types import (
def test_harness_is_plain_enum():
assert Harness.CODEX.value == "codex"
assert Harness.TOOL_LOOP.value == "tool_loop"
assert not isinstance(Harness.CODEX, str)

View file

@ -1,5 +1,6 @@
"""Test health check helper functions"""
import json
import socket
import struct
import zlib
@ -8,6 +9,7 @@ from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import respx
import litellm
from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
@ -548,3 +550,100 @@ def test_ocr_health_check_document_raises_without_the_extension():
_ocr_health_check_document(model="mistral/mistral-ocr-latest", custom_llm_provider="mistral")
finally:
NATIVE_OCR_HEALTH_CHECK_DOCUMENT.reset()
@pytest.mark.parametrize(
("model", "upstream_url"),
(
("perplexity/pplx-decider-v1-27b", "https://api.perplexity.ai/v1/decisions"),
("cloudflare/clef", "https://api.cloudflare.com/client/v4/accounts/acct-1/ai/run/@cf/cloudflare/clef"),
),
)
async def test_ahealth_check_probes_evaluation_models_through_the_decisions_api(
model: str,
upstream_url: str,
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct-1")
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post(upstream_url).respond(
json={
"model": model,
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
"usage": {"input_tokens": 12, "output_tokens": 1},
}
)
result: Final = await ahealth_check({"model": model, "api_key": "sk-test"}, mode=None)
assert "error" not in result, result
assert upstream.called
sent: Final = json.loads(upstream.calls[0].request.content)
assert sent["state"] == "health check"
assert sent["questions"]["reachable"]["type"] == "noul"
@pytest.mark.asyncio
async def test_ahealth_check_evaluation_uses_configured_probe_state_and_questions(
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(
json={
"model": "perplexity/pplx-decider-v1-27b",
"answers": {"ok": {"type": "noul", "noul": 1.0}},
"usage": {"input_tokens": 12, "output_tokens": 1},
}
)
result: Final = await ahealth_check(
{
"model": "perplexity/pplx-decider-v1-27b",
"api_key": "sk-test",
"state": "custom probe",
"questions": {"ok": {"type": "noul", "instructions": "Is it ok?"}},
},
mode=None,
)
assert "error" not in result, result
assert upstream.called
sent: Final = json.loads(upstream.calls[0].request.content)
assert sent["state"] == "custom probe"
assert set(sent["questions"]) == {"ok"}
@pytest.mark.asyncio
async def test_ahealth_check_probes_strands_through_decisions_without_mode(
local_model_cost_map: None,
monkeypatch: pytest.MonkeyPatch,
respx_mock: respx.MockRouter,
) -> None:
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
upstream: Final = respx_mock.post("http://strands.local:8080/v1/systemone").respond(
json={
"model": "strands-decider-2B-hobson-v19",
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
"usage": {"input_tokens": 12, "output_tokens": 1},
}
)
result: Final = await ahealth_check(
{
"model": "strands_decider/strands-decider-2B-hobson-v19",
"api_base": "http://strands.local:8080",
},
mode=None,
)
assert "error" not in result, result
assert upstream.called
assert "authorization" not in upstream.calls[0].request.headers

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