diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index 06c2990a4a4..b769c8f6a3d 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -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 diff --git a/.gitignore b/.gitignore index 399458eced5..763f8db2940 100644 --- a/.gitignore +++ b/.gitignore @@ -151,3 +151,6 @@ litellm.log .coverage-rust coverage-rust.xml + +# make lens-dev worker token, generated config and logs +.lens-dev/ diff --git a/Makefile b/Makefile index e512960949c..40573fb68a3 100644 --- a/Makefile +++ b/Makefile @@ -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/ diff --git a/README.md b/README.md index 4004e6474ee..7ffc44854bb 100644 --- a/README.md +++ b/README.md @@ -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) | ✅ | ✅ | ✅ | | | | | | | | diff --git a/deploy/lens/README.md b/deploy/lens/README.md index 4b65c88afcd..4f78b7bfb59 100644 --- a/deploy/lens/README.md +++ b/deploy/lens/README.md @@ -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 diff --git a/gateway/routes/allowlist.py b/gateway/routes/allowlist.py index 6e91f5486d0..fc11c059c85 100644 --- a/gateway/routes/allowlist.py +++ b/gateway/routes/allowlist.py @@ -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/", diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 7bae178791d..e01d52a55a6 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -4478,6 +4478,7 @@ dependencies = [ "flate2", "futures-util", "hmac 0.12.1", + "itertools 0.14.0", "jsonschema", "litellm-http", "litellm-migrate", diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index 644627b05fd..77a98b3980a 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -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, + page_size: Option, ) -> PyResult> { 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, diff --git a/litellm-rust/crates/storage-clickhouse/src/read.rs b/litellm-rust/crates/storage-clickhouse/src/read.rs index 59d5b0de558..c4bfdef393a 100644 --- a/litellm-rust/crates/storage-clickhouse/src/read.rs +++ b/litellm-rust/crates/storage-clickhouse/src/read.rs @@ -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())); } diff --git a/litellm-rust/crates/storage-clickhouse/tests/transport.rs b/litellm-rust/crates/storage-clickhouse/tests/transport.rs index f58e52b4940..1dab575e21b 100644 --- a/litellm-rust/crates/storage-clickhouse/tests/transport.rs +++ b/litellm-rust/crates/storage-clickhouse/tests/transport.rs @@ -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))); + } +} diff --git a/litellm-rust/crates/traces-clickhouse/Cargo.toml b/litellm-rust/crates/traces-clickhouse/Cargo.toml index 3ab0bfce9fa..15b1ca8162e 100644 --- a/litellm-rust/crates/traces-clickhouse/Cargo.toml +++ b/litellm-rust/crates/traces-clickhouse/Cargo.toml @@ -16,6 +16,7 @@ base64.workspace = true flate2.workspace = true futures-util.workspace = true hmac = "0.12.1" +itertools = "0.14.0" litellm-http.workspace = true litellm-migrate.workspace = true litellm-storage-clickhouse.workspace = true diff --git a/litellm-rust/crates/traces-clickhouse/query/spend_batch.sql b/litellm-rust/crates/traces-clickhouse/query/spend_batch.sql new file mode 100644 index 00000000000..658169dbe34 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/spend_batch.sql @@ -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_` 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} diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql b/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql new file mode 100644 index 00000000000..679edfdec2e --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/trace_list_span_batch.sql @@ -0,0 +1,31 @@ +SELECT * FROM ( +SELECT o.TraceId AS trace_id, o.SpanId AS span_id, o.ParentSpanId AS parent_span_id, o.SpanName AS name, + o.ObservationType AS type, toUInt8(o.WrapperCandidate) AS wrapper_candidate, o.AgentName AS agent, + o.Framework AS framework, o.StatusCode AS status, + substringUTF8(o.StatusMessage, 1, 128) AS status_message, + lengthUTF8(o.StatusMessage) > 128 AS error_truncated, + toUnixTimestamp64Nano(o.Timestamp) AS start_ns, o.Duration AS duration_ns, + o.ServiceName AS service, o.InputPreview AS input_preview, o.Model AS model, + o.InputTokens AS input_tokens, o.OutputTokens AS output_tokens, + o.LiteLLMRequestId AS litellm_request_id, + o.CallKeys AS call_keys, o.CallEvidence AS call_evidence, + -- Rows written before ToolCallId keep the call id only in their attributes. + if(o.ToolCallId != '' OR o.ObservationType != 'tool', o.ToolCallId, + coalesce(nullIf(o.SpanAttributes['gen_ai.tool.call.id'], ''), nullIf(o.SpanAttributes['tool.id'], ''), '')) + AS tool_call_id, + o.UserId AS user_id, o.TeamId AS team_id, o.ApiKeyHash AS api_key_hash +FROM otel_traces AS o +WHERE o.Timestamp >= fromUnixTimestamp64Milli({start_ms:Int64}) + AND o.Timestamp < fromUnixTimestamp64Milli({end_ms:Int64}) + AND ({all_teams:UInt8} = 1 + OR ({user_id:String} != '' AND o.UserId = {user_id:String}) + OR has({team_ids:Array(String)}, o.TeamId)) + AND hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) IN {trace_refs:Array(String)} + AND o.EngineReceivedMs <= {snapshot_ms:UInt64} +ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage +LIMIT 1 BY o.TeamId, o.ApiKeyHash, o.TraceId, o.SpanId + +) +WHERE (team_id, api_key_hash, trace_id, span_id) > ({after_team:String}, {after_key:String}, {after_trace:String}, {after_span:String}) +ORDER BY team_id, api_key_hash, trace_id, span_id +LIMIT {page_size:UInt32} diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql b/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql new file mode 100644 index 00000000000..967ef2fcf1f --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/query/trace_span_batch.sql @@ -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} diff --git a/litellm-rust/crates/traces-clickhouse/src/error.rs b/litellm-rust/crates/traces-clickhouse/src/error.rs index ae39fce25b2..fdbc3029a5d 100644 --- a/litellm-rust/crates/traces-clickhouse/src/error.rs +++ b/litellm-rust/crates/traces-clickhouse/src/error.rs @@ -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")] diff --git a/litellm-rust/crates/traces-clickhouse/src/lib.rs b/litellm-rust/crates/traces-clickhouse/src/lib.rs index fc67df4eba9..d83708d27f1 100644 --- a/litellm-rust/crates/traces-clickhouse/src/lib.rs +++ b/litellm-rust/crates/traces-clickhouse/src/lib.rs @@ -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, }; diff --git a/litellm-rust/crates/traces-clickhouse/src/reads.rs b/litellm-rust/crates/traces-clickhouse/src/reads.rs index 68c44efbde0..4d338e17f07 100644 --- a/litellm-rust/crates/traces-clickhouse/src/reads.rs +++ b/litellm-rust/crates/traces-clickhouse/src/reads.rs @@ -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>> = LazyLock::new(|| { + Cache::builder() + .max_capacity((2 * crate::span_batches::MAX_GRAPH_BYTES) as u64) + .weigher(|_: &String, trace: &Arc| { + 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(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::(client, connection, ¶ms).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 { + 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 = fetch::(client, connection, ¶ms) - .await? - .into_iter() - .map(|row| row.0) - .collect(); + let page: Vec = loop { + match fetch::(client, connection, ¶ms).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if params.0.limit > 1 => { + params.0.limit /= 2; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => break result?.into_iter().map(|row| row.0).collect(), + } + }; let next_cursor = page .last() - .filter(|_| page.len() == 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::>() + .await? + .into_iter() + .flatten() + .collect(); + Ok(TracePage { data, next_cursor }) +} + +async fn list_summaries( + client: &Client, + connection: &Connection, + access: &ReadAccessParams, + runs: &[contracts::ListTracesRow], +) -> Result, Error> { + let (Some(start_ms), Some(end_ms)) = ( + runs.iter().map(|row| row.start_ms).min(), + runs.iter() + .map(|row| row.start_ms.saturating_add(row.duration_ms)) + .max(), ) else { - return Ok(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 = - fetch::(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> = - 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::>(); + 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 = fetch::(client, connection, ¶ms) - .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, 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, diff --git a/litellm-rust/crates/traces-clickhouse/src/span_batches.rs b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs new file mode 100644 index 00000000000..d66ad9506da --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/src/span_batches.rs @@ -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 { + 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, 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::(client, connection, ¶meters).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) + if parameters.page_size > 1 => + { + parameters.page_size /= 2; + continue; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => result?, + }; + let complete = page.len() < parameters.page_size as usize; + if let Some(last) = page.last() { + parameters.after_span_id.clone_from(&last.0.span_id); + } + for row in page { + budget.record(&row)?; + spans.push(row.0); + } + if complete { + spans.sort_by_key(|row| row.start_ns); + return Ok(spans); + } + parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE); + } +} + +#[derive(Serialize)] +struct ListParameters { + #[serde(flatten)] + runs: crate::query::named::TracePageSpansParams, + after_team: String, + after_key: String, + after_trace: String, + after_span: String, + page_size: u32, + snapshot_ms: u64, +} + +struct ListSpanBatch; + +impl Query for ListSpanBatch { + type Params = ListParameters; + type Row = TraceSpansRow; + + const SQL: &'static str = include_str!("../query/trace_list_span_batch.sql"); +} + +pub(crate) async fn read_list_spans( + client: &Client, + connection: &Connection, + runs: crate::query::named::TracePageSpansParams, +) -> Result, Error> { + let parameters = ListParameters { + runs, + after_team: String::new(), + after_key: String::new(), + after_trace: String::new(), + after_span: String::new(), + page_size: PAGE_SIZE, + snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64, + }; + let pages = stream::try_unfold( + (Some(parameters), ReadBudget::default()), + |(parameters, budget)| async move { + let Some(parameters) = parameters else { + return Ok(None); + }; + let page = match fetch::(client, connection, ¶meters).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) + if parameters.page_size > 1 => + { + let retry = ListParameters { + page_size: parameters.page_size / 2, + ..parameters + }; + return Ok(Some((Vec::new(), (Some(retry), budget)))); + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => result?, + }; + let next = page + .last() + .filter(|_| page.len() == parameters.page_size as usize) + .map(|last| ListParameters { + after_team: last.0.team_id.clone(), + after_key: last.0.api_key_hash.clone(), + after_trace: last.0.trace_id.clone(), + after_span: last.0.span_id.clone(), + page_size: (parameters.page_size * 2).min(PAGE_SIZE), + ..parameters + }); + let next_budget = page.iter().try_fold(budget, |budget, row| { + let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?; + budget.checked_add(bytes.len(), 1) + })?; + Ok(Some((page, (next, next_budget)))) + }, + ) + .try_collect::>() + .await?; + Ok(pages + .into_iter() + .flatten() + .map(|row| row.0) + .sorted_by_key(|row| row.start_ns) + .collect()) +} + +#[derive(Serialize)] +struct SpendParameters { + #[serde(flatten)] + 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, 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::(client, connection, ¶meters).await { + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) + if parameters.page_size > 1 => + { + parameters.page_size /= 2; + continue; + } + Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => { + return Err(Error::ReadTooLarge); + } + result => result?, + }; + let complete = page.len() < parameters.page_size as usize; + if let Some(last) = page.last() { + parameters.has_cursor = 1; + parameters.after_team.clone_from(&last.0.team_id); + parameters.after_ms = last.0.start_ms; + parameters.after_id.clone_from(&last.0.request_id); + } + for row in page { + budget.record(&row)?; + rows.push(row.0); + } + if complete { + return Ok(rows); + } + parameters.page_size = (parameters.page_size * 2).min(PAGE_SIZE); + } +} + +#[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); + } +} diff --git a/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs b/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs index 49499ebfa39..73b5125929a 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs @@ -174,7 +174,7 @@ async fn admin_sql_enforces_result_row_limit( matches!( result, Err(Error::Storage( - litellm_storage_clickhouse::Error::QueryFailed(_) + litellm_storage_clickhouse::Error::ResponseTooLarge )) ), "{result:?}" diff --git a/litellm-rust/crates/traces-clickhouse/tests/reads.rs b/litellm-rust/crates/traces-clickhouse/tests/reads.rs new file mode 100644 index 00000000000..41820065d84 --- /dev/null +++ b/litellm-rust/crates/traces-clickhouse/tests/reads.rs @@ -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, + #[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, + #[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::>(); + 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::>(); + 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::()?; + assert!( + read_queries > 0 && read_queries < runs, + "{read_queries} span queries for {runs} runs" + ); + if costed { + let overlapping = client + .post(writer.url().clone()) + .body(format!( + "WITH spend_reads AS ( + SELECT query_start_time_microseconds AS started, event_time_microseconds AS finished + FROM system.query_log + WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' + AND query LIKE '%FROM spend_logs FINAL%' AND query NOT LIKE '%system.query_log%' + ), events AS ( + SELECT started AS at, 1 AS delta FROM spend_reads + UNION ALL SELECT finished AS at, -1 AS delta FROM spend_reads + ) + SELECT max(active) FROM ( + SELECT sum(delta) OVER (ORDER BY at, delta ROWS UNBOUNDED PRECEDING) AS active + FROM events + )" + )) + .send() + .await? + .error_for_status()? + .text() + .await? + .trim() + .parse::()?; + assert!( + (2..=4).contains(&overlapping), + "{overlapping} simultaneous spend reads for {runs} runs" + ); + } + } + for summary in &page.data { + assert_eq!(summary.span_count, steps as u64); + assert_eq!( + 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::>() + ); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive( + #[future(awt)] seeded_database: TestResult, +) -> 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::>(); + 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::>(); + while let Some(current) = cursor { + let next = get_trace_page( + client, + &reader, + &access, + &summary.trace_id, + &summary.trace_ref, + Some(¤t), + 1, + ) + .await? + .ok_or("missing next page")?; + 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, +) -> 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(()) +} diff --git a/litellm-rust/crates/traces/src/resolve/view.rs b/litellm-rust/crates/traces/src/resolve/view.rs index 51145e6bbbf..9c1c3756856 100644 --- a/litellm-rust/crates/traces/src/resolve/view.rs +++ b/litellm-rust/crates/traces/src/resolve/view.rs @@ -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(), diff --git a/litellm-rust/crates/traces/src/view.rs b/litellm-rust/crates/traces/src/view.rs index b7a67ac6822..a864c740526 100644 --- a/litellm-rust/crates/traces/src/view.rs +++ b/litellm-rust/crates/traces/src/view.rs @@ -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, @@ -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, @@ -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, pub spans: Vec, + #[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))] + pub next_cursor: Option, } #[macro_rules_attribute::apply(response_type)] diff --git a/litellm/__init__.py b/litellm/__init__.py index b0761da7f7c..b6428b51bfb 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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", } ) diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 41a7ef1ab64..e358636c105 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -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, diff --git a/litellm/decisions/__init__.py b/litellm/decisions/__init__.py new file mode 100644 index 00000000000..40bddd200a9 --- /dev/null +++ b/litellm/decisions/__init__.py @@ -0,0 +1,3 @@ +from litellm.decisions.main import adecisions, decisions + +__all__ = ["adecisions", "decisions"] diff --git a/litellm/decisions/main.py b/litellm/decisions/main.py new file mode 100644 index 00000000000..da037f1d8cb --- /dev/null +++ b/litellm/decisions/main.py @@ -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"] diff --git a/litellm/harness/__init__.py b/litellm/harness/__init__.py index 79322cdc3e5..213591c5949 100644 --- a/litellm/harness/__init__.py +++ b/litellm/harness/__init__.py @@ -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", diff --git a/litellm/harness/handlers/__init__.py b/litellm/harness/handlers/__init__.py index 4b32d7ca649..5b374a942fa 100644 --- a/litellm/harness/handlers/__init__.py +++ b/litellm/harness/handlers/__init__.py @@ -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}") diff --git a/litellm/harness/handlers/tool_loop_handler.py b/litellm/harness/handlers/tool_loop_handler.py new file mode 100644 index 00000000000..33efcba4ff7 --- /dev/null +++ b/litellm/harness/handlers/tool_loop_handler.py @@ -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") diff --git a/litellm/harness/options.py b/litellm/harness/options.py index 2e3ec5d0fd5..b0014fea393 100644 --- a/litellm/harness/options.py +++ b/litellm/harness/options.py @@ -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 diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py index 1cb8fbcbca2..cdcf251af51 100644 --- a/litellm/harness/sync.py +++ b/litellm/harness/sync.py @@ -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 diff --git a/litellm/harness/types.py b/litellm/harness/types.py index 88ad7f711ea..475568a9979 100644 --- a/litellm/harness/types.py +++ b/litellm/harness/types.py @@ -27,6 +27,7 @@ class Harness(Enum): CODEX = "codex" OPENCODE = "opencode" DEEPAGENTS = "deepagents" + TOOL_LOOP = "tool_loop" def require_harness(harness: object) -> Harness: diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 2cc8e035ebd..c27df51ade2 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -290,6 +290,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_Config", "LiteLLM_SpendLogs", "LiteLLM_BudgetWindowSpend", + "LiteLLM_BackgroundInteractionSettlement", "LiteLLM_ErrorLogs", "LiteLLM_UserNotifications", "LiteLLM_TeamMembership", diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index a0f027cd58f..41965404351 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -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), + } + ) + ), } diff --git a/litellm/litellm_core_utils/health_check_utils.py b/litellm/litellm_core_utils/health_check_utils.py index ae56ae8f899..7fe2d830f1e 100644 --- a/litellm/litellm_core_utils/health_check_utils.py +++ b/litellm/litellm_core_utils/health_check_utils.py @@ -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.""" diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 26c02bb0243..5a3a17f338c 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -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 diff --git a/litellm/llms/base_llm/decisions/__init__.py b/litellm/llms/base_llm/decisions/__init__.py new file mode 100644 index 00000000000..c18ac9b00f2 --- /dev/null +++ b/litellm/llms/base_llm/decisions/__init__.py @@ -0,0 +1,3 @@ +from .transformation import DecisionsProviderConfig, JevCompatibleDecisionsEndpoint + +__all__ = ["DecisionsProviderConfig", "JevCompatibleDecisionsEndpoint"] diff --git a/litellm/llms/base_llm/decisions/transformation.py b/litellm/llms/base_llm/decisions/transformation.py new file mode 100644 index 00000000000..d4fcea24793 --- /dev/null +++ b/litellm/llms/base_llm/decisions/transformation.py @@ -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: ... diff --git a/litellm/llms/base_llm/harness/utils.py b/litellm/llms/base_llm/harness/utils.py index 7c0278d5e36..8dcf784c90f 100644 --- a/litellm/llms/base_llm/harness/utils.py +++ b/litellm/llms/base_llm/harness/utils.py @@ -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: diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 12d04a08392..333bce0b967 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -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/`` 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: diff --git a/litellm/llms/cloudflare/decisions/transformation.py b/litellm/llms/cloudflare/decisions/transformation.py new file mode 100644 index 00000000000..6e8b2999778 --- /dev/null +++ b/litellm/llms/cloudflare/decisions/transformation.py @@ -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() diff --git a/litellm/llms/deepagents/harness/transformation.py b/litellm/llms/deepagents/harness/transformation.py index 82b2c044eac..14f87f0753b 100644 --- a/litellm/llms/deepagents/harness/transformation.py +++ b/litellm/llms/deepagents/harness/transformation.py @@ -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 diff --git a/litellm/llms/openrouter/decisions/transformation.py b/litellm/llms/openrouter/decisions/transformation.py new file mode 100644 index 00000000000..7a7466b1239 --- /dev/null +++ b/litellm/llms/openrouter/decisions/transformation.py @@ -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", +) diff --git a/litellm/llms/perplexity/decisions/transformation.py b/litellm/llms/perplexity/decisions/transformation.py new file mode 100644 index 00000000000..69a4753f4a3 --- /dev/null +++ b/litellm/llms/perplexity/decisions/transformation.py @@ -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", +) diff --git a/litellm/llms/strands_decider/decisions/transformation.py b/litellm/llms/strands_decider/decisions/transformation.py new file mode 100644 index 00000000000..265afb2b148 --- /dev/null +++ b/litellm/llms/strands_decider/decisions/transformation.py @@ -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, +) diff --git a/litellm/llms/tool_loop/__init__.py b/litellm/llms/tool_loop/__init__.py new file mode 100644 index 00000000000..f163f698841 --- /dev/null +++ b/litellm/llms/tool_loop/__init__.py @@ -0,0 +1 @@ +"""In-process tool loop harness.""" diff --git a/litellm/llms/tool_loop/harness/__init__.py b/litellm/llms/tool_loop/harness/__init__.py new file mode 100644 index 00000000000..49ca50b61c2 --- /dev/null +++ b/litellm/llms/tool_loop/harness/__init__.py @@ -0,0 +1 @@ +"""Tool Loop harness configuration.""" diff --git a/litellm/llms/tool_loop/harness/transformation.py b/litellm/llms/tool_loop/harness/transformation.py new file mode 100644 index 00000000000..16461916f1a --- /dev/null +++ b/litellm/llms/tool_loop/harness/transformation.py @@ -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=") diff --git a/litellm/llms/typesafe/decisions/transformation.py b/litellm/llms/typesafe/decisions/transformation.py new file mode 100644 index 00000000000..17fb24ce443 --- /dev/null +++ b/litellm/llms/typesafe/decisions/transformation.py @@ -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", +) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index ea383ef4c11..3f7f5194514 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 54d757d75aa..cb6d47ca4c8 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -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. diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 9eaf6c7e9ed..4cf6db6f257 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -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" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 6d5551d69f1..0671651ea82 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -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", diff --git a/litellm/proxy/agent_endpoints/auth/managed_authorization.py b/litellm/proxy/agent_endpoints/auth/managed_authorization.py index 5b3930b3299..5d8287227a5 100644 --- a/litellm/proxy/agent_endpoints/auth/managed_authorization.py +++ b/litellm/proxy/agent_endpoints/auth/managed_authorization.py @@ -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", } diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 78c3f53c44f..c85169f0ba5 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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", diff --git a/litellm/proxy/decisions_endpoints/__init__.py b/litellm/proxy/decisions_endpoints/__init__.py new file mode 100644 index 00000000000..ea9b7835485 --- /dev/null +++ b/litellm/proxy/decisions_endpoints/__init__.py @@ -0,0 +1 @@ +__all__ = () diff --git a/litellm/proxy/decisions_endpoints/endpoints.py b/litellm/proxy/decisions_endpoints/endpoints.py new file mode 100644 index 00000000000..7dba64ee791 --- /dev/null +++ b/litellm/proxy/decisions_endpoints/endpoints.py @@ -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, + ) diff --git a/litellm/proxy/lens/analysis.py b/litellm/proxy/lens/analysis.py index 001489f3123..7286cec9ba8 100644 --- a/litellm/proxy/lens/analysis.py +++ b/litellm/proxy/lens/analysis.py @@ -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))), ) diff --git a/litellm/proxy/lens/billing.py b/litellm/proxy/lens/billing.py index 8c1c691b87f..9adc0c0fe7b 100644 --- a/litellm/proxy/lens/billing.py +++ b/litellm/proxy/lens/billing.py @@ -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: diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index dc0ae8c985c..a20258853c0 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -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) diff --git a/litellm/proxy/lens/inference.py b/litellm/proxy/lens/inference.py index b8a9d7754ae..7208504cd1f 100644 --- a/litellm/proxy/lens/inference.py +++ b/litellm/proxy/lens/inference.py @@ -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: diff --git a/litellm/proxy/lens/models.py b/litellm/proxy/lens/models.py index 7add39e41be..d967d18d0d1 100644 --- a/litellm/proxy/lens/models.py +++ b/litellm/proxy/lens/models.py @@ -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) diff --git a/litellm/proxy/lens/prompts/investigate.md b/litellm/proxy/lens/prompts/investigate.md index edf795de462..ae3dcfba6ec 100644 --- a/litellm/proxy/lens/prompts/investigate.md +++ b/litellm/proxy/lens/prompts/investigate.md @@ -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. diff --git a/litellm/proxy/lens/prompts/review.md b/litellm/proxy/lens/prompts/review.md index 727c9ad55ed..d1a5e590dfd 100644 --- a/litellm/proxy/lens/prompts/review.md +++ b/litellm/proxy/lens/prompts/review.md @@ -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. diff --git a/litellm/proxy/lens/sources.py b/litellm/proxy/lens/sources.py index 9dbe635e348..9af0f3679b6 100644 --- a/litellm/proxy/lens/sources.py +++ b/litellm/proxy/lens/sources.py @@ -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 diff --git a/litellm/proxy/lens/worker.py b/litellm/proxy/lens/worker.py index 051b4a09392..f9e746489d9 100644 --- a/litellm/proxy/lens/worker.py +++ b/litellm/proxy/lens/worker.py @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index a5b414e5a18..f5ba7f9e877 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 7da09ddcb68..49da000155b 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -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", diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index 50c6e80b234..1563e4b5b55 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -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 diff --git a/litellm/router.py b/litellm/router.py index 3662d1f43eb..b93a4abdf03 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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", diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index e7ecec4df0f..c146a6eac92 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -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 diff --git a/litellm/rust_bridge/trace/generated/types.py b/litellm/rust_bridge/trace/generated/types.py index e1c09ffe281..127e86e9160 100644 --- a/litellm/rust_bridge/trace/generated/types.py +++ b/litellm/rust_bridge/trace/generated/types.py @@ -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): diff --git a/litellm/rust_bridge/trace/storage.py b/litellm/rust_bridge/trace/storage.py index 7e7c519365d..a80e41a5aa3 100644 --- a/litellm/rust_bridge/trace/storage.py +++ b/litellm/rust_bridge/trace/storage.py @@ -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: diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 8df2af4f064..cf6bf8feb01 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -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) diff --git a/litellm/types/decisions.py b/litellm/types/decisions.py new file mode 100644 index 00000000000..80db26f201d --- /dev/null +++ b/litellm/types/decisions.py @@ -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) diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 6d4137f07b5..dbf0d5e7eca 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -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" diff --git a/litellm/utils.py b/litellm/utils.py index 9200844a2e3..d72588e2b00 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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 diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index ea383ef4c11..3f7f5194514 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index eb27d3fe810..7ffaacdb3aa 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -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", diff --git a/scripts/lens_dev.sh b/scripts/lens_dev.sh new file mode 100755 index 00000000000..7ab4b9bb433 --- /dev/null +++ b/scripts/lens_dev.sh @@ -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 </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 diff --git a/scripts/trace_codegen/schemas/traces/Trace.json b/scripts/trace_codegen/schemas/traces/Trace.json index ee02937f556..dc66f77412c 100644 --- a/scripts/trace_codegen/schemas/traces/Trace.json +++ b/scripts/trace_codegen/schemas/traces/Trace.json @@ -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" diff --git a/scripts/trace_codegen/schemas/traces/TracePage.json b/scripts/trace_codegen/schemas/traces/TracePage.json index 72b2c2b2d95..dfa02ba0d39 100644 --- a/scripts/trace_codegen/schemas/traces/TracePage.json +++ b/scripts/trace_codegen/schemas/traces/TracePage.json @@ -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", diff --git a/tests/code_coverage_tests/check_provider_folders_documented.py b/tests/code_coverage_tests/check_provider_folders_documented.py index 08fbde3d979..6f58b393cd1 100644 --- a/tests/code_coverage_tests/check_provider_folders_documented.py +++ b/tests/code_coverage_tests/check_provider_folders_documented.py @@ -33,6 +33,7 @@ EXCLUDED_FOLDERS = { "codex", "opencode", "deepagents", + "tool_loop", } diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 978ac2ec092..0d199770181 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -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 diff --git a/tests/integration/cost_calculation/cost_tracking_case.py b/tests/integration/cost_calculation/cost_tracking_case.py index ac3fcd33d2e..8772141236e 100644 --- a/tests/integration/cost_calculation/cost_tracking_case.py +++ b/tests/integration/cost_calculation/cost_tracking_case.py @@ -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" diff --git a/tests/integration/cost_calculation/cost_tracking_cases.json b/tests/integration/cost_calculation/cost_tracking_cases.json index 09dfa66012f..8242c8683b5 100644 --- a/tests/integration/cost_calculation/cost_tracking_cases.json +++ b/tests/integration/cost_calculation/cost_tracking_cases.json @@ -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", diff --git a/tests/integration/cost_calculation/test_cost_tracking.py b/tests/integration/cost_calculation/test_cost_tracking.py index 82878634677..cfe0a99ef4b 100644 --- a/tests/integration/cost_calculation/test_cost_tracking.py +++ b/tests/integration/cost_calculation/test_cost_tracking.py @@ -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)) diff --git a/tests/integration/management/test_model_health_check.py b/tests/integration/management/test_model_health_check.py index 7d1b1cc2ea5..17bd1c2ebfe 100644 --- a/tests/integration/management/test_model_health_check.py +++ b/tests/integration/management/test_model_health_check.py @@ -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?"}}, + }, + ) + ] diff --git a/tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py b/tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py new file mode 100644 index 00000000000..107839ed320 --- /dev/null +++ b/tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py @@ -0,0 +1,1572 @@ +import asyncio +import base64 +import json +import os +import re +import signal +import socket +import threading +import time +import uuid +from collections.abc import Callable, Iterable, Iterator, Mapping +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import unquote, urlsplit + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually, gateway_from_environment +from integration._support.database import read_rows +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with + +_KIMI_US: Final = "us.moonshotai.kimi-k3" +_KIMI_GLOBAL: Final = "global.moonshotai.kimi-k3" +_KIMI_BASE: Final = "moonshotai.kimi-k3" +_CLAUDE: Final = "us.anthropic.claude-sonnet-5" +_GPT_OSS: Final = "openai.gpt-oss-120b-1:0" +_LLAMA: Final = "us.meta.llama4-maverick-17b-instruct-v1:0" +_PROFILE_ARN_PREFIX: Final = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/" +_YAML_FLAGGED_ARN: Final = f"{_PROFILE_ARN_PREFIX}yaml-flagged-kimi" +_YAML_PLAIN_ARN: Final = f"{_PROFILE_ARN_PREFIX}yaml-plain-kimi" +_KNOWN_MODEL_IDS: Final = frozenset({_KIMI_US, _KIMI_GLOBAL, _KIMI_BASE, _CLAUDE, _GPT_OSS, _LLAMA}) +_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +_ANSWER: Final = "kimi cache point control" +_REJECTION: Final = "This model doesn't support the cachePoint field. Remove cachePoint and try again." +_UNKNOWN_MODEL: Final = "The provided model identifier is invalid." +_CACHE_USAGE_MARKER: Final = "[peer:cache-usage]" +_REJECT_MARKER: Final = "[peer:reject]" +_SLOW_MARKER: Final = "[peer:slow]" +_HOLD_MARKER: Final = "[peer:hold]" +_TARGET: Final = re.compile( + r"^/model/(?P.+)/(?Pconverse|converse-stream|invoke|invoke-with-response-stream)$" +) +_CONVERSE_LIKE: Final = "converse_like" +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_SIGNING_KEY: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_EPHEMERAL: Final[dict[str, JsonValue]] = {"type": "ephemeral"} +_AWS: Final[dict[str, JsonValue]] = { + "aws_access_key_id": "AKIASCRIPTEDPROVIDER", + "aws_secret_access_key": "scripted-secret", + "aws_region_name": "us-east-1", +} +_EXTRA: Final[dict[str, JsonValue]] = {"num_retries": 0, "cache": {"no-cache": True}} +_SYSTEM_TEXT: Final = "You are terse." +_TOOL_PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} + +_NAMES: Final = MappingProxyType( + { + "kimi-us": f"bedrock/{_KIMI_US}", + "kimi-global-converse": f"bedrock/converse/{_KIMI_GLOBAL}", + "kimi-base": f"bedrock/{_KIMI_BASE}", + "kimi-regional": f"bedrock/us-east-1/{_KIMI_US}", + "claude-control": f"bedrock/{_CLAUDE}", + "gpt-oss-control": f"bedrock/{_GPT_OSS}", + "llama-control": f"bedrock/{_LLAMA}", + "arn-yaml-flagged": f"bedrock/{_YAML_FLAGGED_ARN}", + "arn-yaml-plain": f"bedrock/{_YAML_PLAIN_ARN}", + "kimi-inject-message": f"bedrock/{_KIMI_US}", + "kimi-inject-tool": f"bedrock/{_KIMI_US}", + "claude-inject-message": f"bedrock/{_CLAUDE}", + "bedrock/*": "bedrock/*", + } +) +_MODEL_INFO: Final = MappingProxyType({"arn-yaml-flagged": {"supports_prompt_cache_breakpoint": False}}) +_INJECTION: Final = MappingProxyType( + { + "kimi-inject-message": [{"location": "message", "role": "system"}], + "claude-inject-message": [{"location": "message", "role": "system"}], + "kimi-inject-tool": [{"location": "tool_config"}], + } +) + +Endpoint = Literal["chat", "messages", "responses"] +Marker = Literal["system", "user", "tool"] +_ALL_MARKERS: Final = frozenset[Marker]({"system", "user", "tool"}) +_USER_ONLY: Final = frozenset[Marker]({"user"}) +_SYSTEM_AND_USER: Final = frozenset[Marker]({"system", "user"}) +_NO_MARKERS: Final = frozenset[Marker]() + + +def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes: + return _aws_event_frame(event_type, payload, "sc", "u") + + +def _usage(cache_read: int = 0, cache_write: int = 0) -> dict[str, JsonValue]: + cached: Final[dict[str, JsonValue]] = { + **({"cacheReadInputTokens": cache_read} if cache_read else {}), + **({"cacheWriteInputTokens": cache_write} if cache_write else {}), + } + return {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15 + cache_read + cache_write, **cached} + + +def _converse_reply(usage: Mapping[str, JsonValue]) -> bytes: + return json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": _ANSWER}]}}, + "stopReason": "end_turn", + "usage": dict(usage), + "metrics": {"latencyMs": 1}, + } + ).encode() + + +def _text_parts(pieces: int) -> tuple[str, ...]: + words: Final = _ANSWER.split(" ") + assert pieces in (1, len(words)), pieces + if pieces == 1: + return (_ANSWER,) + return tuple(word if index == len(words) - 1 else f"{word} " for index, word in enumerate(words)) + + +def _stream_frames(usage: Mapping[str, JsonValue], pieces: int = 1) -> tuple[bytes, ...]: + deltas: Final = tuple( + _frame("contentBlockDelta", {"delta": {"text": part}, "contentBlockIndex": 0}) for part in _text_parts(pieces) + ) + return ( + _frame("messageStart", {"role": "assistant"}), + *deltas, + _frame("contentBlockStop", {"contentBlockIndex": 0}), + _frame("messageStop", {"stopReason": "end_turn"}), + _frame("metadata", {"usage": dict(usage)}), + ) + + +def _known(model: str) -> bool: + return model in _KNOWN_MODEL_IDS or "application-inference-profile" in model + + +def _error(message: str) -> bytes: + return json.dumps({"message": message}).encode() + + +def _anthropic_message(usage: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return { + "id": f"msg_{uuid.uuid4().hex}", + "type": "message", + "role": "assistant", + "model": _CLAUDE, + "content": [{"type": "text", "text": _ANSWER}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": usage["inputTokens"], "output_tokens": usage["outputTokens"]}, + } + + +def _invoke_chunk(event: Mapping[str, JsonValue]) -> bytes: + encoded: Final = base64.b64encode(json.dumps(event, separators=(",", ":")).encode()).decode() + return _frame("chunk", {"bytes": encoded}) + + +def _invoke_stream(usage: Mapping[str, JsonValue]) -> bytes: + message: Final = { + **_anthropic_message(usage), + "content": [], + "usage": {"input_tokens": usage["inputTokens"], "output_tokens": 1}, + } + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "message_start", "message": message}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _ANSWER}}, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": usage["outputTokens"]}, + }, + {"type": "message_stop"}, + ) + return b"".join(_invoke_chunk(event) for event in events) + + +@dataclass(frozen=True, slots=True) +class _Target: + model: str + action: str + + +def _target(raw: str) -> _Target: + path: Final = unquote(raw) + if path == "/": + return _Target(_CONVERSE_LIKE, "converse") + found: Final = _TARGET.match(path) + assert found is not None, raw + return _Target(found["model"], found["action"]) + + +def _peer(hold: threading.Event | None = None, held: SimpleQueue[str] | None = None) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + target: Final = _target(request.target) + streaming: Final = target.action in ("converse-stream", "invoke-with-response-stream") + text: Final = request.body.decode(errors="replace") + if target.model != _CONVERSE_LIKE and not _known(target.model): + return Reply(status=400, body=_error(_UNKNOWN_MODEL)) + if _REJECT_MARKER in text and not streaming: + return Reply(status=400, body=_error(_REJECTION)) + if _REJECT_MARKER in text: + return Reply(body=_frame("validationException", {"message": _REJECTION}), content_type=_EVENT_STREAM) + usage: Final = _usage(100, 50) if _CACHE_USAGE_MARKER in text else _usage() + if target.action == "invoke": + return Reply(body=json.dumps(_anthropic_message(usage)).encode()) + if target.action == "invoke-with-response-stream": + return Reply(body=_invoke_stream(usage), content_type=_EVENT_STREAM) + if not streaming: + return Reply(body=_converse_reply(usage)) + if _HOLD_MARKER in text and hold is not None: + if held is not None: + held.put(target.model) + return Reply(content_type=_EVENT_STREAM, chunks=_stream_frames(usage), gate_after_first=hold) + if _SLOW_MARKER in text: + return Reply(content_type=_EVENT_STREAM, chunks=_stream_frames(usage, 4), pause_between_chunks=0.2) + return Reply(body=b"".join(_stream_frames(usage)), content_type=_EVENT_STREAM) + + return respond + + +@dataclass(frozen=True, slots=True) +class _Received: + model: str + action: str + streaming: bool + body: dict[str, JsonValue] + cache_points: int + cache_controls: int + + +def _message_blocks(messages: JsonValue) -> Iterator[JsonValue]: + if not isinstance(messages, list): + return + for message in messages: + content: Final = message.get("content") if isinstance(message, dict) else None + if isinstance(content, list): + yield from content + + +def _tool_blocks(body: Mapping[str, JsonValue]) -> tuple[JsonValue, ...]: + tool_config: Final = body.get("toolConfig") + tools: Final = tool_config.get("tools") if isinstance(tool_config, dict) else None + return tuple(tools) if isinstance(tools, list) else () + + +def _cache_points(body: Mapping[str, JsonValue]) -> int: + system: Final = body.get("system") + blocks: Final = ( + *(system if isinstance(system, list) else ()), + *_message_blocks(body.get("messages")), + *_tool_blocks(body), + ) + return sum(1 for block in blocks if isinstance(block, dict) and "cachePoint" in block) + + +def _cache_controls(value: JsonValue) -> int: + if isinstance(value, dict): + return sum(_cache_controls(item) for item in value.values()) + (1 if "cache_control" in value else 0) + if isinstance(value, list): + return sum(_cache_controls(item) for item in value) + return 0 + + +def _parse(request: Request) -> _Received: + target: Final = _target(request.target) + body: Final = _JSON.validate_python(json.loads(request.body)) + streaming: Final = target.action in ("converse-stream", "invoke-with-response-stream") + return _Received(target.model, target.action, streaming, body, _cache_points(body), _cache_controls(body)) + + +def _received(wire: Wire) -> tuple[_Received, ...]: + return tuple(_parse(request) for request in wire.drain()) + + +def _only_received(wire: Wire) -> _Received: + (received,) = _received(wire) + return received + + +def _prompt(*markers: str) -> str: + return " ".join((f"Reply with the control sentence {uuid.uuid4().hex}", *markers)) + + +def _text_block(text: str, cached: bool, kind: str = "text") -> dict[str, JsonValue]: + return {"type": kind, "text": text, **({"cache_control": dict(_EPHEMERAL)} if cached else {})} + + +def _chat_tools(cached: bool) -> dict[str, JsonValue]: + return { + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Weather for a city", + "parameters": _TOOL_PARAMETERS, + }, + **({"cache_control": dict(_EPHEMERAL)} if cached else {}), + } + ], + } + + +def _anthropic_tools(cached: bool) -> dict[str, JsonValue]: + return { + "tools": [ + { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": _TOOL_PARAMETERS, + **({"cache_control": dict(_EPHEMERAL)} if cached else {}), + } + ] + } + + +def _chat_body( + model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False, with_tools: bool = False +) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 16, + "stream": stream, + "messages": [ + {"role": "system", "content": [_text_block(_SYSTEM_TEXT, "system" in markers)]}, + {"role": "user", "content": [_text_block(prompt, "user" in markers)]}, + ], + **(_chat_tools("tool" in markers) if with_tools or "tool" in markers else {}), + **_EXTRA, + } + + +def _messages_body( + model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False, with_tools: bool = False +) -> dict[str, JsonValue]: + return { + "model": model, + "max_tokens": 16, + "stream": stream, + "system": [_text_block(_SYSTEM_TEXT, "system" in markers)], + "messages": [{"role": "user", "content": [_text_block(prompt, "user" in markers)]}], + **(_anthropic_tools("tool" in markers) if with_tools or "tool" in markers else {}), + **_EXTRA, + } + + +def _responses_body( + model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False +) -> dict[str, JsonValue]: + return { + "model": model, + "max_output_tokens": 16, + "stream": stream, + "input": [{"role": "user", "content": [_text_block(prompt, "user" in markers, "input_text")]}], + **_EXTRA, + } + + +def _body( + endpoint: Endpoint, model: str, prompt: str, markers: frozenset[Marker], *, stream: bool = False +) -> dict[str, JsonValue]: + match endpoint: + case "chat": + return _chat_body(model, prompt, markers, stream=stream) + case "messages": + return _messages_body(model, prompt, markers, stream=stream) + case "responses": + return _responses_body(model, prompt, markers, stream=stream) + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +@dataclass(frozen=True, slots=True) +class _Outcome: + status: int + call_id: str + response_id: str + text: str + headers: Mapping[str, str] + raw: str + + +def _sse_payloads(lines: Iterable[str]) -> tuple[dict[str, JsonValue], ...]: + return tuple(json.loads(line[6:]) for line in lines if line.startswith("data: ") and line != "data: [DONE]") + + +def _first_choice(chunk: Mapping[str, JsonValue]) -> dict[str, JsonValue] | None: + choices: Final = chunk.get("choices") + return _JSON.validate_python(choices[0]) if isinstance(choices, list) and choices else None + + +def _chat_stream_text(chunks: Iterable[dict[str, JsonValue]]) -> str: + choices: Final = tuple(choice for choice in map(_first_choice, chunks) if choice is not None) + return "".join(str(_JSON.validate_python(choice["delta"]).get("content") or "") for choice in choices) + + +def _chat_stream_id(chunks: Iterable[dict[str, JsonValue]]) -> str: + (identity,) = {str(chunk["id"]) for chunk in chunks if "id" in chunk} + return identity + + +def _message_stream_id(payloads: Iterable[dict[str, JsonValue]]) -> str: + (started,) = tuple(payload for payload in payloads if payload.get("type") == "message_start") + return str(_JSON.validate_python(started["message"])["id"]) + + +def _messages_stream_text(payloads: Iterable[dict[str, JsonValue]]) -> str: + deltas: Final = tuple(payload for payload in payloads if payload.get("type") == "content_block_delta") + return "".join(str(_JSON.validate_python(delta["delta"]).get("text") or "") for delta in deltas) + + +def _responses_stream_text(events: Iterable[dict[str, JsonValue]]) -> str: + return "".join( + str(event.get("delta") or "") for event in events if event.get("type") == "response.output_text.delta" + ) + + +def _inner_response_id(identity: str) -> str: + managed: Final = decrypt_if_encrypted_with(identity.removeprefix("resp_"), _SIGNING_KEY) + assert managed is not None, identity + issued: Final = managed.split(";", 1)[0].rsplit("response_id:", 1)[1] + decoded: Final = base64.b64decode(issued.removeprefix("resp_")).decode() + return decoded.rsplit("response_id:", 1)[1] + + +def _completed_response_id(events: Iterable[dict[str, JsonValue]]) -> str: + (completed,) = tuple(event for event in events if event.get("type") == "response.completed") + return _inner_response_id(str(_JSON.validate_python(completed["response"])["id"])) + + +def _chat_text(body: Mapping[str, JsonValue]) -> str: + choices: Final = body.get("choices") + assert isinstance(choices, list) and choices, body + return str(_JSON.validate_python(_JSON.validate_python(choices[0])["message"]).get("content") or "") + + +def _messages_text(body: Mapping[str, JsonValue]) -> str: + content: Final = body.get("content") + assert isinstance(content, list), body + return "".join(str(block.get("text") or "") for block in content if isinstance(block, dict)) + + +def _responses_text(body: Mapping[str, JsonValue]) -> str: + output: Final = body.get("output") + assert isinstance(output, list), body + return "".join( + str(part.get("text") or "") + for item in output + if isinstance(item, dict) + for part in ( + item.get("content") if isinstance(item.get("content"), list) else () + ) # comprehension-ok: nested response items + if isinstance(part, dict) + ) + + +def _outcome_of(endpoint: Endpoint, stream: bool, response: httpx.Response, lines: tuple[str, ...]) -> _Outcome: + call_id: Final = response.headers.get("x-litellm-call-id", "") + raw: Final = "\n".join(lines) + if response.status_code != 200: + return _Outcome(response.status_code, call_id, "", "", response.headers, raw) + if stream: + payloads: Final = _sse_payloads(lines) + match endpoint: + case "chat": + return _Outcome( + 200, call_id, _chat_stream_id(payloads), _chat_stream_text(payloads), response.headers, raw + ) + case "messages": + return _Outcome( + 200, call_id, _message_stream_id(payloads), _messages_stream_text(payloads), response.headers, raw + ) + case "responses": + return _Outcome( + 200, + call_id, + _completed_response_id(payloads), + _responses_stream_text(payloads), + response.headers, + raw, + ) + body: Final = _JSON.validate_python(json.loads(raw)) + match endpoint: + case "chat": + return _Outcome(200, call_id, str(body["id"]), _chat_text(body), response.headers, raw) + case "messages": + return _Outcome(200, call_id, str(body["id"]), _messages_text(body), response.headers, raw) + case "responses": + return _Outcome( + 200, call_id, _inner_response_id(str(body["id"])), _responses_text(body), response.headers, raw + ) + + +def _send(gateway: Gateway, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None) -> _Outcome: + stream: Final = body.get("stream") is True + headers: Final = {"Authorization": f"Bearer {gateway.key if key is None else key}"} + with gateway.client.stream("POST", _path(endpoint), json=body, headers=headers, timeout=60) as response: + lines: Final = tuple(line for line in response.iter_lines() if line) + return _outcome_of(endpoint, stream, response, lines) + + +def _spend_rows(request_ids: frozenset[str], *, expected: int, seconds: float = 90) -> tuple[dict[str, JsonValue], ...]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, status, model, spend FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(string_to_array(%s, %s))', + (",".join(sorted(request_ids)), ","), + ), + lambda found: len(found) >= expected, + seconds=seconds, + ) + return tuple(rows) + + +def _success_row(request_id: str) -> dict[str, JsonValue]: + assert request_id, "No id to look the spend row up by" + (row,) = _spend_rows(frozenset({request_id}), expected=1) + assert row["status"] == "success", row + return row + + +def _failure_row(call_id: str) -> dict[str, JsonValue]: + assert call_id, "No call id to look the failure row up by" + (row,) = _spend_rows(frozenset({call_id}), expected=1) + assert row["status"] == "failure", row + return row + + +def _assert_answered(outcome: _Outcome) -> None: + assert outcome.status == 200, (outcome.status, outcome.raw) + assert outcome.text == _ANSWER, outcome.raw + + +def _assert_cache_points(received: _Received, emits: bool, model: str) -> None: + assert received.model == model, (received.model, model) + assert (received.cache_points > 0) == emits, (emits, received.body) + + +def _head_strips() -> bool: + return os.environ.get("INTEGRATION_LEG", "head") == "head" + + +def _owned_config(wire: Wire, directory: Path, *, names: Iterable[str] = tuple(_NAMES)) -> Path: + base: Final = _JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + config: Final[dict[str, JsonValue]] = { + **base, + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": _NAMES[name], + "api_base": wire.url, + "num_retries": 0, + **_AWS, + **({"cache_control_injection_points": _INJECTION[name]} if name in _INJECTION else {}), + }, + **({"model_info": dict(_MODEL_INFO[name])} if name in _MODEL_INFO else {}), + } + for name in names + ], + "router_settings": {**_JSON.validate_python(base["router_settings"]), "num_retries": 0}, + } + path: Final = directory / f"kimi-k3-cache-point-{uuid.uuid4().hex}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@dataclass(frozen=True, slots=True) +class _Rig: + gateway: Gateway + wire: Wire + owned: OwnedProxy + + +@pytest.fixture(scope="module") +def rig(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Rig]: + directory: Final = tmp_path_factory.mktemp("kimi-k3-cache-point") + with gateway_from_environment() as environment, wire_server(_peer()) as wire: + config: Final = _owned_config(wire, directory) + with owned_proxy_process(environment, directory, {}, config=config, workers=2) as owned: + eventually( + lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda count: count == 2, seconds=60 + ) + wire.drain() + yield _Rig(owned.gateway, wire, owned) + + +def _observe(rig: _Rig, endpoint: Endpoint, body: Mapping[str, JsonValue]) -> tuple[_Outcome, _Received]: + rig.wire.drain() + outcome: Final = _send(rig.gateway, endpoint, body) + return outcome, _only_received(rig.wire) + + +_KIMI_FORMS: Final = ("kimi-us", "kimi-global-converse", "kimi-base", "kimi-regional", "arn-yaml-flagged") +_EMITTING_CONTROLS: Final = ("claude-control", "arn-yaml-plain") +_SILENT_CONTROLS: Final = ("gpt-oss-control", "llama-control") + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("name", _KIMI_FORMS, ids=tuple(f"m-{name}" for name in _KIMI_FORMS)) +def test_m01_to_m05_every_kimi_k3_form_sends_converse_without_a_cache_point(rig: _Rig, name: str) -> None: + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + assert received.body.get("toolConfig") is not None, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("name", _EMITTING_CONTROLS, ids=tuple(f"m-{name}" for name in _EMITTING_CONTROLS)) +def test_m08_m09_models_that_take_cache_points_still_get_all_three(rig: _Rig, name: str) -> None: + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 3, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("name", _SILENT_CONTROLS, ids=tuple(f"m-{name}" for name in _SILENT_CONTROLS)) +def test_m10_m11_models_without_caching_never_got_cache_points(rig: _Rig, name: str) -> None: + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m14_a_wildcard_deployment_resolves_the_caller_model_before_deciding(rig: _Rig) -> None: + kimi, kimi_received = _observe(rig, "chat", _chat_body(f"bedrock/{_KIMI_US}", _prompt(), _SYSTEM_AND_USER)) + _assert_answered(kimi) + assert kimi_received.model == _KIMI_US, kimi_received.model + assert kimi_received.cache_points == 0, kimi_received.body + claude, claude_received = _observe(rig, "chat", _chat_body(f"bedrock/{_CLAUDE}", _prompt(), _SYSTEM_AND_USER)) + _assert_answered(claude) + assert claude_received.model == _CLAUDE, claude_received.model + assert claude_received.cache_points == 2, claude_received.body + _success_row(kimi.response_id) + _success_row(claude.response_id) + + +_RAW_CELLS: Final = ( + ("chat", False), + ("chat", True), + ("messages", False), + ("messages", True), + ("responses", False), + ("responses", True), +) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize( + ("endpoint", "stream"), + _RAW_CELLS, + ids=tuple(f"e-{endpoint}-{'stream' if stream else 'plain'}" for endpoint, stream in _RAW_CELLS), +) +def test_e01_to_e06_every_endpoint_reaches_kimi_without_a_cache_point( + rig: _Rig, endpoint: Endpoint, stream: bool +) -> None: + outcome, received = _observe(rig, endpoint, _body(endpoint, "kimi-us", _prompt(), _SYSTEM_AND_USER, stream=stream)) + _assert_answered(outcome) + assert received.streaming == stream, received + assert received.model == _KIMI_US, received.model + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize( + ("endpoint", "stream"), + _RAW_CELLS, + ids=tuple(f"e-{endpoint}-{'stream' if stream else 'plain'}" for endpoint, stream in _RAW_CELLS), +) +def test_e07_to_e12_every_endpoint_still_sends_claude_its_cache_points( + rig: _Rig, endpoint: Endpoint, stream: bool +) -> None: + outcome, received = _observe( + rig, endpoint, _body(endpoint, "claude-control", _prompt(), _SYSTEM_AND_USER, stream=stream) + ) + _assert_answered(outcome) + assert received.model == _CLAUDE, received.model + assert received.streaming == stream, received + if endpoint == "messages": + assert received.action.startswith("invoke"), received.action + assert received.cache_controls == 2 and received.cache_points == 0, received.body + else: + assert received.action.startswith("converse"), received.action + expected: Final = 1 if endpoint == "responses" else 2 + assert received.cache_points == expected, (endpoint, received.body) + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize( + ("endpoint", "stream"), + (("messages", False), ("messages", True)), + ids=("e-messages-arn-plain", "e-messages-arn-stream"), +) +def test_e13_e14_messages_keeps_cache_control_for_an_arn_and_the_flag_decides( + rig: _Rig, endpoint: Endpoint, stream: bool +) -> None: + flagged, flagged_received = _observe( + rig, endpoint, _body(endpoint, "arn-yaml-flagged", _prompt(), _SYSTEM_AND_USER, stream=stream) + ) + _assert_answered(flagged) + assert flagged_received.model == _YAML_FLAGGED_ARN, flagged_received.model + assert flagged_received.cache_points == 0, flagged_received.body + plain, plain_received = _observe( + rig, endpoint, _body(endpoint, "arn-yaml-plain", _prompt(), _SYSTEM_AND_USER, stream=stream) + ) + _assert_answered(plain) + assert plain_received.model == _YAML_PLAIN_ARN, plain_received.model + assert plain_received.cache_points == 2, plain_received.body + _success_row(flagged.response_id) + _success_row(plain.response_id) + + +def _openai_client(rig: _Rig) -> openai.OpenAI: + return openai.OpenAI(base_url=f"{_proxy_url(rig.gateway)}/v1", api_key=rig.gateway.key, max_retries=0) + + +def _async_openai_client(rig: _Rig) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=f"{_proxy_url(rig.gateway)}/v1", api_key=rig.gateway.key, max_retries=0) + + +def _anthropic_client(rig: _Rig) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=_proxy_url(rig.gateway), api_key=rig.gateway.key, max_retries=0) + + +def _async_anthropic_client(rig: _Rig) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=_proxy_url(rig.gateway), api_key=rig.gateway.key, max_retries=0) + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _sdk_messages(prompt: str) -> list[dict[str, JsonValue]]: + return [ + {"role": "system", "content": [_text_block(_SYSTEM_TEXT, True)]}, + {"role": "user", "content": [_text_block(prompt, True)]}, + ] + + +@pytest.mark.timeout(600) +def test_k01_openai_sdk_sync_chat_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + completion: Final = _openai_client(rig).chat.completions.create( + model="kimi-us", + messages=_sdk_messages(_prompt()), # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_tokens=16, + extra_body=_EXTRA, + ) + assert completion.choices[0].message.content == _ANSWER, completion + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and not received.streaming, received + assert received.cache_points == 0, received.body + _success_row(completion.id) + + +@pytest.mark.timeout(600) +async def test_k02_openai_sdk_async_chat_stream_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + stream: Final = await _async_openai_client(rig).chat.completions.create( + model="kimi-us", + messages=_sdk_messages(_prompt()), # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_tokens=16, + stream=True, + extra_body=_EXTRA, + ) + chunks: Final = [chunk async for chunk in stream] + text: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + assert text == _ANSWER, chunks + (identity,) = {chunk.id for chunk in chunks} + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and received.streaming, received + assert received.cache_points == 0, received.body + _success_row(identity) + + +_ANTHROPIC_SDK_TARGETS: Final = (("kimi-us", _KIMI_US), ("arn-yaml-flagged", _YAML_FLAGGED_ARN)) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("name", "model_id"), _ANTHROPIC_SDK_TARGETS, ids=("k03-kimi", "k03-flagged-arn")) +def test_k03_anthropic_sdk_sync_messages_reaches_the_model_without_a_cache_point( + rig: _Rig, name: str, model_id: str +) -> None: + rig.wire.drain() + message: Final = _anthropic_client(rig).messages.create( + model=name, + max_tokens=16, + system=[{"type": "text", "text": _SYSTEM_TEXT, "cache_control": {"type": "ephemeral"}}], + messages=[ + {"role": "user", "content": [{"type": "text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}]} + ], + extra_body=_EXTRA, + ) + assert "".join(block.text for block in message.content if block.type == "text") == _ANSWER, message + received: Final = _only_received(rig.wire) + assert received.model == model_id and not received.streaming, received + assert received.cache_points == 0, received.body + _success_row(message.id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("name", "model_id"), _ANTHROPIC_SDK_TARGETS, ids=("k04-kimi", "k04-flagged-arn")) +async def test_k04_anthropic_sdk_async_stream_reaches_the_model_without_a_cache_point( + rig: _Rig, name: str, model_id: str +) -> None: + rig.wire.drain() + async with _async_anthropic_client(rig).messages.stream( + model=name, + max_tokens=16, + system=[{"type": "text", "text": _SYSTEM_TEXT, "cache_control": {"type": "ephemeral"}}], + messages=[ + {"role": "user", "content": [{"type": "text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}]} + ], + extra_body=_EXTRA, + ) as stream: + final: Final = await stream.get_final_message() + assert "".join(block.text for block in final.content if block.type == "text") == _ANSWER, final + received: Final = _only_received(rig.wire) + assert received.model == model_id and received.streaming, received + assert received.cache_points == 0, received.body + _success_row(final.id) + + +@pytest.mark.timeout(600) +def test_k05_openai_sdk_sync_responses_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + response: Final = _openai_client(rig).responses.create( + model="kimi-us", + input=[ + { + "role": "user", + "content": [{"type": "input_text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}], + } + ], # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_output_tokens=16, + extra_body=_EXTRA, + ) + assert response.output_text == _ANSWER, response + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and not received.streaming, received + assert received.cache_points == 0, received.body + _success_row(_inner_response_id(response.id)) + + +@pytest.mark.timeout(600) +async def test_k06_openai_sdk_async_responses_stream_reaches_kimi_without_a_cache_point(rig: _Rig) -> None: + rig.wire.drain() + stream: Final = await _async_openai_client(rig).responses.create( + model="kimi-us", + input=[ + { + "role": "user", + "content": [{"type": "input_text", "text": _prompt(), "cache_control": {"type": "ephemeral"}}], + } + ], # pyright: ignore[reportArgumentType] # cache_control rides as an extra key + max_output_tokens=16, + stream=True, + extra_body=_EXTRA, + ) + completed: Final = [event.response async for event in stream if event.type == "response.completed"] + assert len(completed) == 1 and completed[0].output_text == _ANSWER, completed + received: Final = _only_received(rig.wire) + assert received.model == _KIMI_US and received.streaming, received + assert received.cache_points == 0, received.body + _success_row(_inner_response_id(completed[0].id)) + + +_LOCATIONS: Final = ( + ("system", frozenset[Marker]({"system"})), + ("user", frozenset[Marker]({"user"})), + ("tool", frozenset[Marker]({"tool"})), +) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("location", "markers"), _LOCATIONS, ids=tuple(f"l-{location}" for location, _ in _LOCATIONS)) +def test_l01_to_l03_each_marker_location_is_dropped_for_kimi_and_kept_for_claude( + rig: _Rig, location: str, markers: frozenset[Marker] +) -> None: + kimi, kimi_received = _observe(rig, "chat", _chat_body("kimi-us", _prompt(), markers, with_tools=True)) + _assert_answered(kimi) + assert kimi_received.cache_points == 0, (location, kimi_received.body) + claude, claude_received = _observe(rig, "chat", _chat_body("claude-control", _prompt(), markers, with_tools=True)) + _assert_answered(claude) + assert claude_received.cache_points == 1, (location, claude_received.body) + _success_row(kimi.response_id) + _success_row(claude.response_id) + + +@pytest.mark.timeout(600) +def test_l04_l05_gateway_injection_points_are_dropped_for_kimi_and_kept_for_claude(rig: _Rig) -> None: + kimi_message, kimi_message_received = _observe( + rig, "chat", _chat_body("kimi-inject-message", _prompt(), _NO_MARKERS, with_tools=True) + ) + _assert_answered(kimi_message) + assert kimi_message_received.cache_points == 0, kimi_message_received.body + kimi_tool, kimi_tool_received = _observe( + rig, "chat", _chat_body("kimi-inject-tool", _prompt(), _NO_MARKERS, with_tools=True) + ) + _assert_answered(kimi_tool) + assert kimi_tool_received.cache_points == 0, kimi_tool_received.body + claude, claude_received = _observe( + rig, "chat", _chat_body("claude-inject-message", _prompt(), _NO_MARKERS, with_tools=True) + ) + _assert_answered(claude) + assert claude_received.cache_points == 1, claude_received.body + system: Final = claude_received.body.get("system") + assert isinstance(system, list) and "cachePoint" in _JSON.validate_python(system[-1]), claude_received.body + for outcome in (kimi_message, kimi_tool, claude): + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_l06_a_request_without_markers_was_never_touched(rig: _Rig) -> None: + for name in ("kimi-us", "claude-control"): + outcome, received = _observe(rig, "chat", _chat_body(name, _prompt(), _NO_MARKERS, with_tools=True)) + _assert_answered(outcome) + assert received.cache_points == 0, (name, received.body) + _success_row(outcome.response_id) + + +def _rates(gateway: Gateway, name: str) -> dict[str, float]: + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list), entries + (entry,) = tuple( + _JSON.validate_python(item) for item in entries if _JSON.validate_python(item)["model_name"] == name + ) + info: Final = _JSON.validate_python(entry["model_info"]) + return { + key: float(str(info[key])) + for key in ( + "input_cost_per_token", + "output_cost_per_token", + "cache_read_input_token_cost", + "cache_creation_input_token_cost", + ) + } + + +@pytest.mark.timeout(600) +def test_p01_kimi_cache_usage_is_still_priced_from_its_own_row(rig: _Rig) -> None: + outcome, received = _observe(rig, "chat", _chat_body("kimi-us", _prompt(_CACHE_USAGE_MARKER), _NO_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + body: Final = _JSON.validate_python(json.loads(outcome.raw)) + usage: Final = _JSON.validate_python(body["usage"]) + assert usage["prompt_tokens"] == 161 and usage["completion_tokens"] == 4, usage + details: Final = _JSON.validate_python(usage["prompt_tokens_details"]) + assert details["cached_tokens"] == 100, usage + assert usage["cache_creation_input_tokens"] == 50, usage + rates: Final = _rates(rig.gateway, "kimi-us") + expected: Final = ( + 11 * rates["input_cost_per_token"] + + 100 * rates["cache_read_input_token_cost"] + + 50 * rates["cache_creation_input_token_cost"] + + 4 * rates["output_cost_per_token"] + ) + assert abs(float(outcome.headers["x-litellm-response-cost"]) - expected) < 1e-12, (outcome.headers, rates) + row: Final = _success_row(outcome.response_id) + assert abs(float(str(row["spend"])) - expected) < 1e-12, (row, rates) + + +@pytest.mark.timeout(600) +def test_p02_a_plain_kimi_reply_is_priced_from_its_own_row(rig: _Rig) -> None: + outcome, received = _observe(rig, "chat", _chat_body("kimi-us", _prompt(), _NO_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 0, received.body + rates: Final = _rates(rig.gateway, "kimi-us") + expected: Final = 11 * rates["input_cost_per_token"] + 4 * rates["output_cost_per_token"] + assert abs(float(outcome.headers["x-litellm-response-cost"]) - expected) < 1e-12, (outcome.headers, rates) + row: Final = _success_row(outcome.response_id) + assert abs(float(str(row["spend"])) - expected) < 1e-12, (row, rates) + + +@pytest.mark.timeout(600) +def test_s07_a_5kb_caller_model_under_the_wildcard_is_a_bounded_4xx_and_liveliness_stays_up(rig: _Rig) -> None: + rig.wire.drain() + started: Final = time.monotonic() + outcome: Final = _send(rig.gateway, "chat", _chat_body(f"bedrock/{'a' * 5120}", _prompt(), _SYSTEM_AND_USER)) + elapsed: Final = time.monotonic() - started + liveliness: Final = rig.gateway.client.get("/health/liveliness", timeout=5) + assert liveliness.status_code == 200, liveliness.text + assert outcome.status in (400, 404), (outcome.status, outcome.raw[:300]) + error: Final = _JSON.validate_python(_JSON.validate_python(json.loads(outcome.raw))["error"]) + assert "a" * 5120 in str(error["message"]), outcome.raw[:300] + assert elapsed < 10, elapsed + assert rig.wire.drain() == () + _failure_row(outcome.call_id) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("stream", (False, True), ids=("s08-plain", "s08-stream")) +def test_s08_a_peer_rejection_on_kimi_is_a_400_with_the_message_and_a_failure_row(rig: _Rig, stream: bool) -> None: + rig.wire.drain() + outcome: Final = _send( + rig.gateway, "chat", _chat_body("kimi-us", _prompt(_REJECT_MARKER), _NO_MARKERS, stream=stream) + ) + received: Final = _only_received(rig.wire) + assert received.cache_points == 0, received.body + assert outcome.status == 400, (outcome.status, outcome.raw) + assert _REJECTION in outcome.raw.replace('\\"', '"'), outcome.raw + _failure_row(outcome.call_id) + + +@pytest.mark.timeout(600) +def test_s09_an_unauthenticated_kimi_request_never_reaches_the_peer(rig: _Rig) -> None: + rig.wire.drain() + outcome: Final = _send( + rig.gateway, "chat", _chat_body("kimi-us", _prompt(), _SYSTEM_AND_USER), key="sk-integration-bogus" + ) + assert outcome.status == 401, outcome.raw + assert rig.wire.drain() == () + + +@pytest.mark.timeout(600) +def test_x02_two_identical_uncached_requests_are_two_peer_calls_and_two_rows(rig: _Rig) -> None: + body: Final = _chat_body("kimi-us", _prompt(), _NO_MARKERS) + rig.wire.drain() + first: Final = _send(rig.gateway, "chat", body) + second: Final = _send(rig.gateway, "chat", body) + _assert_answered(first) + _assert_answered(second) + assert first.response_id != second.response_id, (first.response_id, second.response_id) + received: Final = _received(rig.wire) + assert len(received) == 2 and all(item.cache_points == 0 for item in received), received + rows: Final = _spend_rows(frozenset({first.response_id, second.response_id}), expected=2) + assert {str(row["request_id"]) for row in rows} == {first.response_id, second.response_id}, rows + + +@pytest.mark.timeout(600) +def test_x03_a_response_cache_hit_still_answers_after_fewer_peer_calls_than_sends(rig: _Rig) -> None: + body: Final = {key: value for key, value in _chat_body("kimi-us", _prompt(), _NO_MARKERS).items() if key != "cache"} + rig.wire.drain() + first: Final = _send(rig.gateway, "chat", body) + _assert_answered(first) + sends: Final[list[_Outcome]] = [] + + def resend() -> _Outcome: + served: Final = _send(rig.gateway, "chat", body) + sends.append(served) + return served + + hit: Final = eventually( + resend, lambda served: served.status == 200 and served.response_id == first.response_id, seconds=15 + ) + _assert_answered(hit) + received: Final = _received(rig.wire) + assert all(item.cache_points == 0 for item in received), received + assert len(received) == len(sends), (len(received), len(sends)) + _success_row(first.response_id) + + +@pytest.mark.timeout(600) +def test_x05_an_arn_whose_profile_name_starts_with_openai_dot_is_classified_by_the_pre_existing_family_rule( + gateway: Gateway, +) -> None: + arn: Final = f"{_PROFILE_ARN_PREFIX}openai.custom-{uuid.uuid4().hex[:8]}" + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}") + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == arn, received.model + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +def _deployment( + gateway: Gateway, + scenario: Scenario, + wire: Wire, + model: str, + model_info: Mapping[str, JsonValue] | None = None, + *, + cleanup: bool = True, +) -> str: + name: Final = f"kimi-audit-{uuid.uuid4().hex}" + created: Final = gateway.post( + "/model/new", + { + "model_name": name, + "litellm_params": {"model": model, "api_base": wire.url, "num_retries": 0, **_AWS}, + "model_info": dict(model_info) if model_info is not None else {}, + }, + ) + identity: Final = str(_JSON.validate_python(created["model_info"])["id"]) + if cleanup: + scenario.cleanups.callback(gateway.post, "/model/delete", {"id": identity}) + _settled(gateway, name, wire) + return name + + +def _settled(gateway: Gateway, name: str, wire: Wire) -> None: + eventually( + lambda: tuple( + gateway.request("POST", "/v1/chat/completions", _chat_body(name, _prompt(), _NO_MARKERS)).status_code + for _ in range(12) + ), + lambda codes: all(code == 200 for code in codes), + seconds=60, + ) + wire.drain() + + +def _observe_at( + gateway: Gateway, wire: Wire, endpoint: Endpoint, body: Mapping[str, JsonValue] +) -> tuple[_Outcome, _Received]: + wire.drain() + outcome: Final = _send(gateway, endpoint, body) + return outcome, _only_received(wire) + + +def _arn(label: str) -> str: + return f"{_PROFILE_ARN_PREFIX}{label}-{uuid.uuid4().hex[:12]}" + + +_ROUTES: Final = ("", "converse/", "converse_like/") + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize("route", _ROUTES, ids=("m05-plain-route", "m06-converse-route", "m07-converse-like-route")) +def test_m05_to_m07_a_deployment_flag_false_on_an_arn_strips_cache_points_on_every_route( + gateway: Gateway, route: str +) -> None: + arn: Final = _arn("flagged") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{route}{arn}", {"supports_prompt_cache_breakpoint": False} + ) + expected_model: Final = _CONVERSE_LIKE if route == "converse_like/" else arn + for endpoint in ("chat", "messages", "responses"): + outcome, received = _observe_at(gateway, wire, endpoint, _body(endpoint, name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == expected_model, (endpoint, received.model) + assert received.action == "converse", (endpoint, received.action) + assert received.cache_points == 0, (endpoint, received.body) + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m08_an_arn_without_the_flag_still_gets_cache_points(gateway: Gateway) -> None: + arn: Final = _arn("plain") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}") + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _ALL_MARKERS)) + _assert_answered(outcome) + assert received.cache_points == 3, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m12_a_deployment_flag_true_on_kimi_overrides_the_cost_map_row(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{_KIMI_BASE}", {"supports_prompt_cache_breakpoint": True} + ) + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == _KIMI_BASE, received.model + assert received.cache_points == 2, received.body + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_m13_a_null_deployment_flag_on_kimi_falls_back_to_the_cost_map_row(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{_KIMI_US}", {"supports_prompt_cache_breakpoint": None} + ) + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.model == _KIMI_US, received.model + assert received.cache_points == 0, received.body + _success_row(outcome.response_id) + + +_ODD_FLAGS: Final[tuple[tuple[str, JsonValue], ...]] = ( + ("s01-string-true", "true"), + ("s02-int-one", 1), + ("s03-int-zero", 0), + ("s04-empty-list", []), + ("s05-empty-string", ""), + ("s06-5kb-string", "x" * 5120), +) + + +@pytest.mark.timeout(600) +@pytest.mark.parametrize(("label", "flag"), _ODD_FLAGS, ids=tuple(label for label, _ in _ODD_FLAGS)) +def test_s01_to_s06_an_odd_typed_flag_is_read_as_false_never_as_true( + gateway: Gateway, label: str, flag: JsonValue +) -> None: + arn: Final = _arn(label) + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": flag}) + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.cache_points == 0, (label, received.body) + _success_row(outcome.response_id) + + +@pytest.mark.timeout(600) +def test_s10_s11_two_deployments_of_one_arn_share_the_flag_and_a_delete_leaves_it_in_place(gateway: Gateway) -> None: + arn: Final = _arn("shared") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + flagged: Final = _deployment( + gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": False}, cleanup=False + ) + plain: Final = _deployment(gateway, scenario, wire, f"bedrock/{arn}") + for name in (flagged, plain): + outcome, received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(outcome) + assert received.cache_points == 0, (name, received.body) + _success_row(outcome.response_id) + flagged_identity: Final = _model_id(flagged) + gateway.post("/model/delete", {"id": flagged_identity}) + eventually( + lambda: tuple( + gateway.request("POST", "/v1/chat/completions", _chat_body(flagged, _prompt(), _NO_MARKERS)).status_code + for _ in range(12) + ), + lambda codes: all(code != 200 for code in codes), + seconds=60, + ) + wire.drain() + after, after_received = _observe_at(gateway, wire, "chat", _chat_body(plain, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(after) + assert after_received.cache_points == 0, after_received.body + _success_row(after.response_id) + + +def _model_id(name: str) -> str: + (row,) = read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_name=%s', (name,)) + return str(row["model_id"]) + + +def _stored_flag(identity: str) -> JsonValue: + (row,) = read_rows('SELECT model_info FROM "LiteLLM_ProxyModelTable" WHERE model_id=%s', (identity,)) + stored: Final = row["model_info"] + info: Final = _JSON.validate_python(stored if isinstance(stored, dict) else json.loads(str(stored))) + return info.get("supports_prompt_cache_breakpoint") + + +def _six_cache_point_counts(gateway: Gateway, wire: Wire, name: str) -> tuple[int, ...]: + wire.drain() + outcomes: Final = tuple(_send(gateway, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) for _ in range(6)) + assert all(item.status == 200 for item in outcomes), outcomes + return tuple(item.cache_points for item in _received(wire)) + + +@pytest.mark.timeout(600) +def test_x01_flipping_the_flag_to_true_through_a_patch_update_turns_cache_points_back_on(gateway: Gateway) -> None: + arn: Final = _arn("flip") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": False} + ) + before, before_received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(before) + assert before_received.cache_points == 0, before_received.body + identity: Final = _model_id(name) + patched: Final = gateway.request( + "PATCH", f"/model/{identity}/update", {"model_info": {"supports_prompt_cache_breakpoint": True}} + ) + assert patched.status_code == 200, patched.text + assert _stored_flag(identity) is True + points: Final = eventually( + lambda: _six_cache_point_counts(gateway, wire, name), + lambda counts: len(counts) == 6 and all(count == 2 for count in counts), + seconds=60, + ) + assert points == (2,) * 6, points + _success_row(before.response_id) + + +@pytest.mark.timeout(600) +def test_x06_the_legacy_post_update_answers_200_and_leaves_the_stored_flag_alone(gateway: Gateway) -> None: + arn: Final = _arn("legacy") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + name: Final = _deployment( + gateway, scenario, wire, f"bedrock/{arn}", {"supports_prompt_cache_breakpoint": False} + ) + before, before_received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(before) + identity: Final = _model_id(name) + updated: Final = gateway.request( + "POST", + "/model/update", + { + "model_name": name, + "litellm_params": {"model": f"bedrock/{arn}", "api_base": wire.url, "num_retries": 0, **_AWS}, + "model_info": {"id": identity, "supports_prompt_cache_breakpoint": True}, + }, + ) + assert updated.status_code == 200, updated.text + assert _stored_flag(identity) is False + _settled(gateway, name, wire) + after, after_received = _observe_at(gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER)) + _assert_answered(after) + assert after_received.cache_points == before_received.cache_points, (before_received.body, after_received.body) + _success_row(before.response_id) + _success_row(after.response_id) + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + index: int + stream: bool + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + outcome: _Outcome + + +def _burst_calls(count: int) -> tuple[_Call, ...]: + endpoints: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") + return tuple(_Call(endpoints[index % 3], index, index % 2 == 0) for index in range(count)) + + +async def _send_async(client: httpx.AsyncClient, key: str, model: str, call: _Call, prompt: str) -> _Served: + body: Final = _body(call.endpoint, model, f"{prompt} {call.index}", _SYSTEM_AND_USER, stream=call.stream) + async with client.stream( + "POST", _path(call.endpoint), json=body, headers={"Authorization": f"Bearer {key}"} + ) as response: + raw: Final = (await response.aread()).decode() + lines: Final = tuple(line for line in raw.splitlines() if line) + return _Served(call, _outcome_of(call.endpoint, call.stream, response, lines)) + + +async def _burst( + base_url: str, + key: str, + model: str, + calls: tuple[_Call, ...], + prompt: str, + *, + 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_async(client, key, model, call, prompt) 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_rows_once(served: tuple[_Served, ...]) -> None: + ids: Final = frozenset(item.outcome.response_id for item in served) + assert len(ids) == len(served), ids + rows: Final = _spend_rows(ids, expected=len(ids)) + assert sorted(str(row["request_id"]) for row in rows) == sorted(ids), rows + assert all(row["status"] == "success" for row in rows), rows + + +def _reserved_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +@pytest.mark.timeout(900) +async def test_c01_a_peer_outage_mid_traffic_fails_cleanly_and_recovery_lands_every_id_once( + gateway: Gateway, tmp_path: Path +) -> None: + port: Final = _reserved_port() + prompt: Final = _prompt() + with wire_server(_peer(), port=port) as first_peer: + config: Final = _owned_config(first_peer, tmp_path, names=("kimi-us", "claude-control")) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate = owned.gateway # rebind-ok: the same name covers the restarted proxy below + url = _proxy_url(candidate) # rebind-ok: the same name covers the restarted proxy below + served: Final = await _burst(url, candidate.key, "kimi-us", _burst_calls(30), prompt) + assert len(served) == 30 + for item in served: + _assert_answered(item.outcome) + received: Final = _received(first_peer) + assert len(received) == 30 and all(item.cache_points == 0 for item in received), received + _assert_rows_once(served) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate = owned.gateway + url = _proxy_url(candidate) + first_peer_down: Final = await _burst(url, candidate.key, "kimi-us", _burst_calls(30), prompt) + assert len(first_peer_down) == 30 + for item in first_peer_down: + assert item.outcome.status >= 500, (item.call, item.outcome.status, item.outcome.raw) + liveliness: Final = candidate.client.get("/health/liveliness", timeout=5) + assert liveliness.status_code == 200, liveliness.text + failed_ids: Final = frozenset( + item.outcome.call_id + for item in first_peer_down + if item.call.endpoint != "messages" and item.outcome.call_id + ) + assert len(failed_ids) == 20, failed_ids + failure_rows: Final = _spend_rows(failed_ids, expected=20) + assert all(row["status"] == "failure" for row in failure_rows), failure_rows + with wire_server(_peer(), port=port) as second_peer: + recovered: Final = await _burst(url, candidate.key, "kimi-us", _burst_calls(30), prompt) + assert len(recovered) == 30 + for item in recovered: + _assert_answered(item.outcome) + again: Final = _received(second_peer) + assert len(again) == 30 and all(item.cache_points == 0 for item in again), again + _assert_rows_once(recovered) + control: Final = await _burst(url, candidate.key, "claude-control", _burst_calls(6), prompt) + assert all(item.outcome.status == 200 for item in control), control + control_received: Final = _received(second_peer) + assert len(control_received) == 6, control_received + assert all(item.cache_points + item.cache_controls > 0 for item in control_received), control_received + + +def _open_peer_connections(pid: int, peer_url: str) -> int: + port: Final = urlsplit(peer_url).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 + ) + + +async def _held_burst(url: str, key: str, count: int, prompt: str) -> tuple[_Served, ...]: + calls: Final = tuple(_Call("chat", index, True) for index in range(count)) + return await _burst(url, key, "kimi-us", calls, f"{prompt} {_HOLD_MARKER}", tolerate_transport_errors=True) + + +@pytest.mark.timeout(900) +async def test_c02_a_worker_killed_mid_traffic_leaves_the_survivor_and_the_respawn_stripping( + gateway: Gateway, tmp_path: Path +) -> None: + hold: Final = threading.Event() + held: Final[SimpleQueue[str]] = SimpleQueue() + prompt: Final = _prompt() + with wire_server(_peer(hold, held)) as wire: + config: Final = _owned_config(wire, tmp_path, names=("kimi-us",)) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + url: Final = _proxy_url(candidate) + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=60, + ) + burst: Final = asyncio.create_task(_held_burst(url, candidate.key, 20, prompt)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_peer_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__) + psutil.Process(victim_pid).send_signal(signal.SIGKILL) + hold.set() + served: Final = await burst + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered(item.outcome) + received: Final = _received(wire) + assert len(received) == 20 and all(item.cache_points == 0 for item in received), received + respawned: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 3, + seconds=60, + )[-1] + hold.clear() + for _ in range(20): + for _ in range(held.qsize()): + held.get_nowait() + follow_up_task: Final = asyncio.create_task(_held_burst(url, candidate.key, 6, prompt)) + await asyncio.to_thread(eventually, held.qsize, lambda size: size == 6, 60) + on_respawn: Final = _open_peer_connections(respawned, wire.url) + hold.set() + follow_up: Final = await follow_up_task + hold.clear() + assert len(follow_up) == 6, follow_up + for item in follow_up: + _assert_answered(item.outcome) + later: Final = _received(wire) + assert len(later) == 6 and all(item.cache_points == 0 for item in later), later + if on_respawn > 0: + break + else: + raise AssertionError("The respawned worker never took a held stream") + _assert_rows_once(served) + + +@pytest.mark.timeout(900) +def test_c03_a_restart_re_registers_a_stored_deployment_flag_from_the_database( + gateway: Gateway, tmp_path: Path +) -> None: + arn: Final = _arn("restart") + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + config: Final = _owned_config(wire, tmp_path, names=("claude-control",)) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + name: Final = _deployment( + owned.gateway, + scenario, + wire, + f"bedrock/{arn}", + {"supports_prompt_cache_breakpoint": False}, + cleanup=False, + ) + before, before_received = _observe_at( + owned.gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER) + ) + _assert_answered(before) + assert before_received.cache_points == 0, before_received.body + _success_row(before.response_id) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as restarted: + _settled(restarted.gateway, name, wire) + after, after_received = _observe_at( + restarted.gateway, wire, "chat", _chat_body(name, _prompt(), _SYSTEM_AND_USER) + ) + _assert_answered(after) + assert after_received.model == arn, after_received.model + assert after_received.cache_points == 0, after_received.body + _success_row(after.response_id) + control, control_received = _observe_at( + restarted.gateway, wire, "chat", _chat_body("claude-control", _prompt(), _SYSTEM_AND_USER) + ) + _assert_answered(control) + assert control_received.cache_points == 2, control_received.body + restarted.gateway.post("/model/delete", {"id": _model_id(name)}) + + +@pytest.mark.timeout(900) +async def test_c04_a_slow_peer_under_ten_concurrent_streams_completes_every_call_once(rig: _Rig) -> None: + rig.wire.drain() + calls: Final = tuple(_Call("chat", index, True) for index in range(10)) + started: Final = time.monotonic() + served: Final = await _burst( + _proxy_url(rig.gateway), rig.gateway.key, "kimi-us", calls, f"{_prompt()} {_SLOW_MARKER}" + ) + elapsed: Final = time.monotonic() - started + assert len(served) == 10 + for item in served: + _assert_answered(item.outcome) + received: Final = _received(rig.wire) + assert len(received) == 10 and all(item.cache_points == 0 for item in received), received + assert elapsed < 30, elapsed + _assert_rows_once(served) diff --git a/tests/integration/providers/test_decisions_chaos.py b/tests/integration/providers/test_decisions_chaos.py new file mode 100644 index 00000000000..2e88cdbbc7f --- /dev/null +++ b/tests/integration/providers/test_decisions_chaos.py @@ -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}, + } diff --git a/tests/integration/providers/test_decisions_wire.py b/tests/integration/providers/test_decisions_wire.py new file mode 100644 index 00000000000..9878559cb07 --- /dev/null +++ b/tests/integration/providers/test_decisions_wire.py @@ -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 diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml index a3d29e42120..085adea81ac 100644 --- a/tests/integration/proxy_config.yaml +++ b/tests/integration/proxy_config.yaml @@ -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 diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index 77e1421675a..1dade952dd8 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -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) diff --git a/tests/test_litellm/tracing/test_receiver.py b/tests/test_litellm/tracing/test_receiver.py index 66b971e8e63..70282fcf694 100644 --- a/tests/test_litellm/tracing/test_receiver.py +++ b/tests/test_litellm/tracing/test_receiver.py @@ -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 diff --git a/tests/unit/decisions/__init__.py b/tests/unit/decisions/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/decisions/test_main.py b/tests/unit/decisions/test_main.py new file mode 100644 index 00000000000..106710d328f --- /dev/null +++ b/tests/unit/decisions/test_main.py @@ -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"] diff --git a/tests/unit/harness/handlers/test_tool_loop_handler.py b/tests/unit/harness/handlers/test_tool_loop_handler.py new file mode 100644 index 00000000000..2ca265cda67 --- /dev/null +++ b/tests/unit/harness/handlers/test_tool_loop_handler.py @@ -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" diff --git a/tests/unit/harness/test_init.py b/tests/unit/harness/test_init.py index effdc238249..75ec3b60c74 100644 --- a/tests/unit/harness/test_init.py +++ b/tests/unit/harness/test_init.py @@ -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 diff --git a/tests/unit/harness/test_types.py b/tests/unit/harness/test_types.py index 036d6dd625c..e59c4e983fb 100644 --- a/tests/unit/harness/test_types.py +++ b/tests/unit/harness/test_types.py @@ -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) diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 47c4576f91f..941e44feb26 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -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 diff --git a/tests/unit/llms/azure/test_azure_common_utils.py b/tests/unit/llms/azure/test_azure_common_utils.py index caf941ebd19..5f254a3bc6c 100644 --- a/tests/unit/llms/azure/test_azure_common_utils.py +++ b/tests/unit/llms/azure/test_azure_common_utils.py @@ -470,6 +470,7 @@ def test_default_max_retries_env_var_reaches_azure_sdk_client(): "allm_passthrough_route", "llm_passthrough_route", "asearch", + "adecisions", "avector_store_create", "avector_store_search", "acreate_skill", diff --git a/tests/unit/llms/bedrock/chat/test_converse_transformation.py b/tests/unit/llms/bedrock/chat/test_converse_transformation.py index 07c54eee395..e4a50317190 100644 --- a/tests/unit/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/unit/llms/bedrock/chat/test_converse_transformation.py @@ -5769,15 +5769,20 @@ def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model(): pytest.param("global.openai.gpt-6-astra", False, id="openai-family-implicit-caching-only"), pytest.param("openai.gpt-oss-120b-1:0", False, id="openai-gpt-oss"), pytest.param("us.openai.gpt-99-unmapped", False, id="unmapped-openai-family-still-suppressed"), + pytest.param("us.moonshotai.kimi-k3", False, id="kimi-k3-prices-cached-tokens-but-rejects-cachepoint"), + pytest.param("global.moonshotai.kimi-k3", False, id="kimi-k3-global-profile"), + pytest.param("us-east-1/us.moonshotai.kimi-k3", False, id="kimi-k3-regional-route-resolves-through-profile"), ], ) def test_cache_points_emitted_only_for_models_that_support_prompt_caching(model, expects_cache_points, monkeypatch): """Bedrock rejects cachePoint blocks for models without prompt caching support - ("You invoked an unsupported model or your request did not allow prompt caching"), - and clients like Claude Code attach cache_control to every request, so a map-known - model without the capability must not receive them. Unmapped ids (application - inference profile ARNs, models newer than the map) keep emitting so existing - caching setups never silently degrade.""" + ("You invoked an unsupported model or your request did not allow prompt caching") + and for models that price cached tokens yet take the marker only on their native + endpoints ("This model doesn't support the cachePoint field", Kimi K3), and clients + like Claude Code attach cache_control to every request, so a map-known model without + the capability must not receive them on system, message, or tool blocks. Unmapped ids + (application inference profile ARNs, models newer than the map) keep emitting so + existing caching setups never silently degrade.""" monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) @@ -5787,14 +5792,25 @@ def test_cache_points_emitted_only_for_models_that_support_prompt_caching(model, {"role": "system", "content": [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]}, {"role": "user", "content": [{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}]}, ], - optional_params={}, + optional_params={ + "tools": [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object", "properties": {}}}, + "cache_control": {"type": "ephemeral"}, + } + ] + }, litellm_params={}, headers={}, ) - assert ("cachePoint" in json.dumps(body)) is expects_cache_points + assert ("cachePoint" in json.dumps(body["system"])) is expects_cache_points + assert ("cachePoint" in json.dumps(body["messages"])) is expects_cache_points + assert ("cachePoint" in json.dumps(body["toolConfig"])) is expects_cache_points assert body["system"][0]["text"] == "sys" assert body["messages"][0]["content"][0]["text"] == "hi" + assert body["toolConfig"]["tools"][0]["toolSpec"]["name"] == "get_weather" def test_tool_config_cachepoint_not_placed_or_credited_for_model_without_prompt_caching(monkeypatch): diff --git a/tests/unit/llms/bedrock/test_bedrock_common_utils.py b/tests/unit/llms/bedrock/test_bedrock_common_utils.py index 22e7d354be7..d4c182dc952 100644 --- a/tests/unit/llms/bedrock/test_bedrock_common_utils.py +++ b/tests/unit/llms/bedrock/test_bedrock_common_utils.py @@ -479,6 +479,75 @@ def test_capability_lookups_fall_back_to_base_model_when_regional_entry_lacks_fi assert bedrock_converse_supports_parallel_tool_use_config(regional) is True +@pytest.mark.parametrize( + ("entry", "expected"), + [ + pytest.param( + {"supports_prompt_caching": True, "supports_prompt_cache_breakpoint": False}, + False, + id="priced-cached-tokens-but-rejects-the-explicit-marker", + ), + pytest.param( + {"supports_prompt_caching": False, "supports_prompt_cache_breakpoint": True}, + True, + id="explicit-marker-flag-wins-over-the-caching-flag", + ), + pytest.param({"supports_prompt_caching": True}, True, id="caching-flag-alone-keeps-emitting"), + pytest.param({"supports_prompt_caching": False}, False, id="no-caching-and-no-marker-flag"), + ], +) +def test_bedrock_model_accepts_cache_points_prefers_the_explicit_breakpoint_flag(monkeypatch, entry, expected): + import litellm + from litellm.llms.bedrock.common_utils import bedrock_model_accepts_cache_points + + base = "vendor.breakpoint-flag-test" + monkeypatch.setitem(litellm.model_cost, f"us.{base}", {"input_cost_per_token": 1e-06}) + monkeypatch.setitem(litellm.model_cost, base, entry) + + assert bedrock_model_accepts_cache_points(f"us.{base}") is expected + + +@pytest.mark.parametrize("model", ["moonshotai.kimi-k3", "us.moonshotai.kimi-k3", "global.moonshotai.kimi-k3"]) +def test_kimi_k3_keeps_cached_token_pricing_while_refusing_converse_cache_points(model, local_model_cost_map): + import litellm + from litellm.llms.bedrock.common_utils import bedrock_model_accepts_cache_points + + assert bedrock_model_accepts_cache_points(model) is False + assert litellm.utils.supports_prompt_caching(model=model, custom_llm_provider="bedrock") is True + assert litellm.model_cost[model]["cache_read_input_token_cost"] > 0 + + +def test_deployment_model_info_breakpoint_flag_covers_an_unmapped_arn(local_model_cost_map): + from litellm import Router + from litellm.llms.bedrock.common_utils import bedrock_model_accepts_cache_points + + flagged_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/flagged" + unflagged_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/unflagged" + converse_arn = "arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/converse" + Router( + model_list=[ + { + "model_name": "kimi-k3-profile-converse", + "litellm_params": {"model": f"bedrock/converse/{converse_arn}", "aws_region_name": "us-east-1"}, + "model_info": {"supports_prompt_cache_breakpoint": False}, + }, + { + "model_name": "kimi-k3-profile", + "litellm_params": {"model": f"bedrock/{flagged_arn}", "aws_region_name": "us-east-1"}, + "model_info": {"supports_prompt_cache_breakpoint": False}, + }, + { + "model_name": "kimi-k3-profile-unflagged", + "litellm_params": {"model": f"bedrock/{unflagged_arn}", "aws_region_name": "us-east-1"}, + }, + ] + ) + + assert bedrock_model_accepts_cache_points(flagged_arn) is False + assert bedrock_model_accepts_cache_points(converse_arn) is False + assert bedrock_model_accepts_cache_points(unflagged_arn) is True + + def test_merge_bedrock_aws_request_params_strips_caller_identity_when_deployment_has_static_credentials(): from litellm.llms.bedrock.common_utils import merge_bedrock_aws_request_params diff --git a/tests/unit/llms/tool_loop/__init__.py b/tests/unit/llms/tool_loop/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/tool_loop/harness/__init__.py b/tests/unit/llms/tool_loop/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/tool_loop/harness/test_transformation.py b/tests/unit/llms/tool_loop/harness/test_transformation.py new file mode 100644 index 00000000000..11bce41bb56 --- /dev/null +++ b/tests/unit/llms/tool_loop/harness/test_transformation.py @@ -0,0 +1,154 @@ +"""Tests for Tool Loop schemas and completion routing.""" + +from __future__ import annotations + +from collections.abc import Callable +from pathlib import Path +from typing import Final, Literal + +import pytest +from pydantic import BaseModel + +from litellm import sandbox +from litellm.harness.context import GatewayTarget, SessionContext +from litellm.harness.options import ToolLoopOptions +from litellm.harness.types import Harness +from litellm.llms.tool_loop.harness.transformation import ( + ToolLoopHarnessConfig, + completion_kwargs, + function_tool, +) + + +def make_context( + tmp_path: Path, + *, + model: str | None = "anthropic/claude", + gateway: GatewayTarget | None = None, + api_key: str | None = None, + api_base: str | None = None, + options: ToolLoopOptions | None = None, + output: type[BaseModel] | None = None, +) -> SessionContext: + return SessionContext( + harness=Harness.TOOL_LOOP, + sandbox=sandbox.local(tmp_path), + session_id="tool-loop-transform", + model=model, + gateway=gateway, + api_key=api_key, + api_base=api_base, + options=options, + output=output, + ) + + +def search( + query: str, + limit: int = 5, + state: Literal["open", "closed"] = "open", +) -> str: + """Search records.""" + return query + + +def test_function_tool_schema_has_required_defaulted_and_literal_fields() -> None: + specification: Final = function_tool(search).spec + schema: Final = specification["function"]["parameters"] + + assert schema["required"] == ["query"] + assert schema["properties"]["query"] == {"title": "Query", "type": "string"} + assert schema["properties"]["limit"] == {"default": 5, "title": "Limit", "type": "integer"} + assert schema["properties"]["state"] == { + "default": "open", + "enum": ["open", "closed"], + "title": "State", + "type": "string", + } + assert specification["function"]["description"] == "Search records." + + +def test_function_schema_rejects_unknown_arguments() -> None: + with pytest.raises(ValueError, match="Extra inputs are not permitted"): + function_tool(search).args_model.model_validate({"query": "owner", "unknown": "value"}) + + +def variadic_positional(*args: int) -> int: + return len(args) + + +def variadic_keyword(**kwargs: int) -> int: + return len(kwargs) + + +@pytest.mark.parametrize("fn", [variadic_positional, variadic_keyword]) +def test_variadic_tools_are_rejected(fn: Callable[..., object]) -> None: + with pytest.raises(ValueError, match="variadic parameters"): + function_tool(fn) + + +def test_sdk_routing_overrides_completion_kwargs(tmp_path: Path) -> None: + ctx: Final = make_context( + tmp_path, + api_key="provided-key", + api_base="https://provider", + options=ToolLoopOptions( + completion_kwargs={ + "model": "wrong-model", + "api_key": "wrong-key", + "api_base": "https://wrong", + "temperature": 0.2, + } + ), + ) + + kwargs: Final = completion_kwargs(ctx) + + assert kwargs == { + "model": "anthropic/claude", + "api_key": "provided-key", + "api_base": "https://provider", + "temperature": 0.2, + } + + +def test_gateway_routing_and_response_format_override_options(tmp_path: Path) -> None: + class OutputModel(BaseModel): + pass + + gateway: Final = GatewayTarget(api_base="https://gateway", api_key="virtual-key") + ctx: Final = make_context( + tmp_path, + gateway=gateway, + options=ToolLoopOptions( + completion_kwargs={ + "model": "wrong-model", + "api_key": "wrong-key", + "api_base": "https://wrong", + "response_format": "wrong-format", + } + ), + output=OutputModel, + ) + + kwargs: Final = completion_kwargs(ctx) + + assert kwargs == { + "model": "litellm_proxy/anthropic/claude", + "api_base": "https://gateway", + "api_key": "virtual-key", + "extra_headers": {"x-litellm-tags": "harness,tool_loop"}, + "response_format": OutputModel, + } + + +def test_configuration_requires_model_and_declares_capabilities(tmp_path: Path) -> None: + config: Final = ToolLoopHarnessConfig() + assert config.uses_model_endpoint is False + assert config.capabilities.structured_output + assert config.capabilities.tool_approval + assert config.capabilities.history + assert config.capabilities.custom_tools + assert config.capabilities.permission_modes == frozenset({"ask", "full"}) + with pytest.raises(ValueError, match=r"Harness\.TOOL_LOOP needs model="): + config.validate_environment(make_context(tmp_path, model=None)) diff --git a/tests/unit/proxy/decisions_endpoints/__init__.py b/tests/unit/proxy/decisions_endpoints/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/decisions_endpoints/test_endpoints.py b/tests/unit/proxy/decisions_endpoints/test_endpoints.py new file mode 100644 index 00000000000..6b1ac9e3404 --- /dev/null +++ b/tests/unit/proxy/decisions_endpoints/test_endpoints.py @@ -0,0 +1,326 @@ +from __future__ import annotations + +import asyncio +import json +from collections.abc import AsyncGenerator, Iterator, Mapping +from contextlib import asynccontextmanager +from pathlib import Path +from typing import Final + +import pytest +import respx +from fastapi import FastAPI +from fastapi.routing import APIRoute +from fastapi.testclient import TestClient +from starlette.routing import Match + +import litellm +from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeature, attach_lazy_features +from litellm.proxy.decisions_endpoints.endpoints import decisions +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import SafeRouteAdder +from litellm.proxy.proxy_server import ( + app, + cleanup_router_config_variables, + initialize, +) + +_INPUT_TOKENS: Final[int] = 367 +_OUTPUT_TOKENS: Final[int] = 3 +_RESPONSE: Final[Mapping[str, object]] = { + "model": "pplx-decider-v1-27b", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + }, + "usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS}, +} +_STRANDS_RESPONSE: Final[Mapping[str, object]] = { + "model": "strands-decider-2B-hobson-v19", + "answers": { + "is_defect": {"type": "noul", "noul": 0.9}, + }, + "usage": {"input_tokens": 216, "output_tokens": 3}, + "latency_ms": 3722.17, +} +_REQUEST: Final[Mapping[str, object]] = { + "model": "decider", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, +} + + +@pytest.fixture +def client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + monkeypatch.setenv("OPENAI_API_KEY", "fake-openai-key") + monkeypatch.setenv("OPENAI_API_BASE", "https://fake-openai.example") + monkeypatch.setenv("REDIS_HOST", "localhost") + cleanup_router_config_variables() + config_path: Final = Path(__file__).parents[1] / "test_configs" / "test_config_no_auth.yaml" + asyncio.run(initialize(config=str(config_path), debug=True)) + monkeypatch.setattr( + litellm.proxy.proxy_server, + "llm_router", + litellm.Router( + model_list=[ + { + "model_name": "decider", + "litellm_params": { + "model": "perplexity/pplx-decider-v1-27b", + "api_key": "test-key", + }, + } + ] + ), + ) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + yield TestClient(app) + litellm.in_memory_llm_clients_cache.flush_cache() + + +@pytest.mark.parametrize("endpoint", ("/v1/decisions", "/decisions")) +def test_proxy_decisions_route_returns_answers_and_cost( + client: TestClient, + respx_mock: respx.MockRouter, + endpoint: str, +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post(endpoint, json=_REQUEST) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert "_hidden_params" not in response.json() + 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 float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected_cost) + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "pplx-decider-v1-27b", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + assert upstream.calls[0].request.headers["authorization"] == "Bearer test-key" + + +def test_proxy_decisions_dispatches_typesafe_deployment( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + router: Final = litellm.Router( + model_list=[ + { + "model_name": "jev", + "litellm_params": { + "model": "typesafe/jev-latest", + "api_key": "k", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE) + + response: Final = client.post( + "/v1/decisions", + json={ + "model": "jev", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "jev-latest", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + + +def test_proxy_decisions_sends_the_env_key_to_the_deployment_api_base( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("PERPLEXITYAI_API_KEY", "server-key") + monkeypatch.delenv("PERPLEXITY_API_KEY", raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "decider", + "litellm_params": { + "model": "perplexity/pplx-decider-v1-27b", + "api_base": "https://egress.example/perplexity", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post("https://egress.example/perplexity/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post("/v1/decisions", json=_REQUEST) + + assert response.status_code == 200, response.text + assert upstream.call_count == 1 + assert upstream.calls[0].request.headers["authorization"] == "Bearer server-key" + + +def test_proxy_decisions_unknown_model_is_a_client_error( + client: TestClient, + respx_mock: respx.MockRouter, +) -> None: + response: Final = client.post( + "/v1/decisions", + json={ + "model": "missing-model", + "state": "review", + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert 400 <= response.status_code < 500, response.text + assert len(respx_mock.calls) == 0 + + +@pytest.mark.parametrize( + "request_body", + ( + { + "model": "decider", + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + { + "model": "decider", + "state": {"source": "proxy-test"}, + }, + ), + ids=("missing_state", "missing_questions"), +) +def test_proxy_decisions_missing_required_field_is_a_client_error( + client: TestClient, + respx_mock: respx.MockRouter, + request_body: Mapping[str, object], +) -> None: + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post("/v1/decisions", json=request_body) + + assert response.status_code == 400, response.text + assert not upstream.called + + +def test_proxy_decisions_dispatches_strands_decider( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False) + monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False) + router: Final = litellm.Router( + model_list=[ + { + "model_name": "strands", + "litellm_params": { + "model": "strands_decider/strands-decider-2B-hobson-v19", + "api_base": "https://strands.example", + }, + } + ] + ) + monkeypatch.setattr(litellm.proxy.proxy_server, "llm_router", router) + upstream: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE) + + response: Final = client.post( + "/v1/decisions", + json={ + "model": "strands", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + }, + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _STRANDS_RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content) == { + "model": "strands-decider-2B-hobson-v19", + "state": {"source": "proxy-test"}, + "questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}}, + } + assert "authorization" not in upstream.calls[0].request.headers + + +def test_proxy_decisions_without_model_uses_the_proxy_default_model( + client: TestClient, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm.proxy.proxy_server, "user_model", "decider") + upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE) + + response: Final = client.post( + "/v1/decisions", json={key: value for key, value in _REQUEST.items() if key != "model"} + ) + + assert response.status_code == 200, response.text + assert response.json()["answers"] == _RESPONSE["answers"] + assert upstream.called + assert json.loads(upstream.calls[0].request.content)["model"] == "pplx-decider-v1-27b" + + +def _decisions_feature() -> LazyFeature: + return next(feature for feature in LAZY_FEATURES if feature.name == "decisions") + + +def _serving_endpoint(bare: FastAPI, path: str) -> object: + scope: Final = {"type": "http", "method": "POST", "path": path, "root_path": "", "query_string": b"", "headers": ()} + return next( + route.endpoint for route in bare.routes if isinstance(route, APIRoute) and route.matches(scope)[0] is Match.FULL + ) + + +def test_a_config_pass_through_at_v1_decisions_keeps_its_route_and_the_native_api_serves_decisions( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv("LITELLM_DISABLE_LAZY_ROUTES", raising=False) + + async def pass_through() -> dict[str, str]: + return {"served_by": "pass-through"} + + bare: Final = FastAPI() + attach_lazy_features(bare, (_decisions_feature(),)) + SafeRouteAdder.add_api_route_if_not_exists(bare, "/v1/decisions", pass_through, ["POST"]) + with TestClient(bare) as client: + assert client.post("/v1/decisions", json={"model": "gpt-6-luna"}).json() == {"served_by": "pass-through"} + assert _serving_endpoint(bare, "/v1/decisions") is pass_through + assert _serving_endpoint(bare, "/decisions") is decisions + + +def test_with_lazy_routes_disabled_a_config_pass_through_at_v1_decisions_still_wins( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv("LITELLM_DISABLE_LAZY_ROUTES", "true") + + async def pass_through() -> dict[str, str]: + return {"served_by": "pass-through"} + + @asynccontextmanager + async def loads_the_config(app_: FastAPI) -> AsyncGenerator[None]: + assert SafeRouteAdder.add_api_route_if_not_exists(app_, "/v1/decisions", pass_through, ["POST"]), ( + "the native route registered at startup must not block the config pass-through" + ) + yield + + bare: Final = FastAPI(lifespan=loads_the_config) + attach_lazy_features(bare, (_decisions_feature(),)) + with TestClient(bare) as client: + assert client.post("/v1/decisions", json={"model": "gpt-6-luna"}).json() == {"served_by": "pass-through"} + assert _serving_endpoint(bare, "/v1/decisions") is pass_through + assert _serving_endpoint(bare, "/decisions") is decisions diff --git a/tests/unit/proxy/lens/test_analysis.py b/tests/unit/proxy/lens/test_analysis.py index d710bcce937..5231000e14b 100644 --- a/tests/unit/proxy/lens/test_analysis.py +++ b/tests/unit/proxy/lens/test_analysis.py @@ -398,23 +398,20 @@ async def test_investigator_keeps_final_outcome_ahead_of_repeated_model_history( "quote, check_id, accepted", [("timeout", "retries", True), ("invented quote", "retries", False), ("timeout", "unknown", False)], ) -async def test_oversized_model_evidence_is_retried_and_quotes_still_verified( +async def test_many_model_citations_are_accepted_but_quotes_are_still_verified( quote: str, check_id: str, accepted: bool ) -> None: execution: Final = Execution( id="run1", source="traces", trace_id="t", team_id="alpha", name="review", start_time="", span_count=1 ) part: Final = TracePart(execution_id="run1", span_id="span", name="tool", kind="tool", content="timeout") - attempts: Final = iter((8, 1)) + attempts: Final = iter((8,)) async def read(_execution_id: str, _cursor: str, _offset: int) -> ExecutionContent: return ExecutionContent(execution=execution, parts=(part,)) async def model(request: ModelRequest) -> ModelResult: count: Final = next(attempts) - if count == 1: - assert "validation errors" in request.prompt - assert '"max_length":6' in request.prompt evidence: Final = Evidence(execution_id="run1", span_id="span", quote=quote).model_dump_json() return ModelResult( content='{"observations":[{"check_id":"' @@ -434,9 +431,7 @@ async def test_oversized_model_evidence_is_retried_and_quotes_still_verified( @pytest.mark.asyncio async def test_invalid_model_output_has_only_one_repair_attempt() -> None: - from pydantic import ValidationError - - from litellm.proxy.lens.analysis import Extraction, structured_response + from litellm.proxy.lens.analysis import AnalysisResponseError, Extraction, structured_response attempts: Final = iter((1, 2)) @@ -444,7 +439,9 @@ async def test_invalid_model_output_has_only_one_repair_attempt() -> None: assert next(attempts, None) is not None, "Model repair exceeded its retry limit" return ModelResult(content="not JSON", cost=0) - with pytest.raises(ValidationError): + with pytest.raises( + AnalysisResponseError, match="Reading executions failed: Extraction response invalid after 2 attempts" + ): await structured_response(ModelRequest(purpose="extract", prompt="Extract observations"), Extraction, model) assert next(attempts, None) is None @@ -505,16 +502,18 @@ async def test_investigator_can_cite_a_later_page_or_offset(later_span: str) -> draft: Final = finding("run1").model_copy( update={"evidence": (Evidence(execution_id="run1", span_id=later_span, quote="timeout"),)} ) - decisions: Final = iter(("read", "submit")) + offsets: Final = iter((8000, 16000, None)) async def model(request: ModelRequest) -> ModelResult: - if next(decisions) == "read": - return ModelResult(content='{"action":"read","execution_id":"run1","offset":8000}', cost=0) + offset: Final = next(offsets) + if offset is not None: + return ModelResult(content=json.dumps({"action": "read", "execution_id": "run1", "offset": offset}), cost=0) + assert json.loads(request.prompt)["must_decide"] is False assert '"content": "timeout"' in request.prompt return ModelResult(content='{"action":"submit","finding":' + draft.model_dump_json() + "}", cost=0) async def read(execution_id: str, _cursor: str, offset: int) -> ExecutionContent: - assert execution_id == "run1" and offset == 8000 + assert execution_id == "run1" and offset in (8000, 16000) return ExecutionContent(execution=execution, parts=(later,)) claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) @@ -963,6 +962,7 @@ async def test_invalid_candidate_response_preserves_other_findings_and_reports_i ) assert tuple(result.finding for result in results if result.finding is not None) == (finding("run"),) assert sum(result.finding is None for result in results) == 1 + assert "[json_invalid]" in next(result.error for result in results if result.finding is None) assert max(counts.get_nowait() for _ in range(counts.qsize())) == 1 @@ -996,3 +996,203 @@ async def test_investigator_keeps_the_issue_brief() -> None: ) assert result.finding is not None assert result.finding.brief == draft.brief + + +@pytest.mark.asyncio +@pytest.mark.parametrize("finish_reason", (None, "length", "content_filter")) +async def test_grouping_failure_keeps_validation_details_without_model_content(finish_reason: str | None) -> None: + from litellm.proxy.lens.analysis import AnalysisResponseError, Clusters, structured_response + + async def model(_request: ModelRequest) -> ModelResult: + return ModelResult.model_validate( + {"content": '{"candidates":[{"title":"private trace"}]}', "cost": 0, "finish_reason": finish_reason} + ) + + with pytest.raises(AnalysisResponseError) as caught: + await structured_response(ModelRequest(purpose="cluster", prompt="private evidence"), Clusters, model) + message: Final = str(caught.value) + assert message.startswith("Grouping observations failed: Clusters response invalid after 2 attempts.") + assert "candidates.0.check_id: Field required [missing]" in message + assert "private" not in message + if finish_reason: + assert f"finish_reason={finish_reason}" in message + else: + assert "truncated" not in message + + +@pytest.mark.asyncio +async def test_truncated_but_valid_json_is_repaired_before_accepting_findings() -> None: + from litellm.proxy.lens.analysis import Clusters, structured_response + + outputs: Final = iter( + ( + ModelResult(content='{"candidates":[]}', cost=0, finish_reason="length"), + ModelResult(content='{"candidates":[]}', cost=0), + ) + ) + + async def model(_request: ModelRequest) -> ModelResult: + return next(outputs) + + assert await structured_response(ModelRequest(purpose="cluster", prompt="group"), Clusters, model) == Clusters() + assert next(outputs, None) is None + + +@pytest.mark.asyncio +async def test_large_context_and_long_verified_quotes_do_not_silently_end_investigation() -> None: + from litellm.proxy.lens.models import FindingDraft, LensSettings + + context: Final = "Read all recorded evidence. " * 5000 + long_quote: Final = "timeout detail " * 200 + execution: Final = Execution( + id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=1 + ) + part: Final = TracePart(execution_id="run", span_id="span", name="tool", kind="tool", content=long_quote) + reviewed: Final = Examined(execution=execution, observations=(), parts=(part,), partial=False, cannot_assess=False) + expected: Final = FindingDraft.model_validate( + { + **finding("run").model_dump(), + "description": "Recorded failure detail. " * 300, + "evidence": [{"execution_id": "run", "span_id": "span", "quote": long_quote}], + } + ) + settings: Final = LensSettings.model_validate({**lens().settings.model_dump(), "context": context}) + claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job", settings=settings).jobs[0], findings=()) + + async def model(request: ModelRequest) -> ModelResult: + assert json.loads(request.prompt)["context"] == context + return ModelResult(content=json.dumps({"action": "submit", "finding": expected.model_dump()}), cost=0) + + async def read(_execution_id: str, _cursor: str, _offset: int) -> ExecutionContent: + pytest.fail("Already supplied evidence should not require a read") + + result: Final = await investigate( + claim, + Candidate(check_id="retries", title="Failure", hypothesis="Retry failed", execution_ids=("run",)), + (reviewed,), + read, + model, + ) + assert result.finding == expected + + +@pytest.mark.asyncio +async def test_reviewer_can_read_every_offset_of_a_long_span_before_deciding() -> None: + execution: Final = Execution( + id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=1 + ) + original: Final = "trace evidence! " * 16000 + "late verified failure" + offsets: Final = SimpleQueue[int]() + seen: Final = SimpleQueue[str]() + + async def read(_execution_id: str, _cursor: str, offset: int) -> ExecutionContent: + offsets.put(offset) + content: Final = ( + "Preview; read for complete content" if offset == 0 else original[offset - 1 : offset - 1 + 8000] + ) + return ExecutionContent( + execution=execution, + parts=( + TracePart( + execution_id="run", + span_id="span", + name="agent", + kind="agent", + content=content, + truncated=offset == 0 or offset - 1 + 8000 < len(original), + ), + ), + ) + + async def model(request: ModelRequest) -> ModelResult: + payload: Final = json.loads(request.prompt) + read_count: Final = payload["completed_read_count"] + if read_count: + seen.put(payload["read_evidence"][0]["content"]) + if read_count * 8000 < len(original): + return ModelResult( + content=json.dumps({"reads": [{"span_id": "span", "offset": 1 + read_count * 8000}]}), cost=0 + ) + return ModelResult( + content=json.dumps( + { + "observations": [ + { + "check_id": "retries", + "summary": "Late failure", + "evidence": [{"execution_id": "run", "span_id": "span", "quote": "late verified failure"}], + } + ] + } + ), + cost=0, + ) + + claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) + result: Final = await extract(claim, execution, read, model) + assert "".join(seen.get_nowait() for _ in range(seen.qsize())) == original + assert tuple(offsets.get_nowait() for _ in range(offsets.qsize())) == (0, *range(1, len(original) + 1, 8000)) + assert result.observations[0].evidence[0].quote == "late verified failure" + assert not result.cannot_assess + + +@pytest.mark.asyncio +async def test_investigator_can_read_all_evidence_pages_across_successive_span_batches() -> None: + execution: Final = Execution( + id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=80 + ) + parts: Final = tuple( + TracePart( + execution_id="run", + span_id=f"span{i:03}", + parent_span_id="root", + name=f"Step {i}", + kind="tool", + content="recorded evidence " * 400 + ("timeout" if i == 79 else "complete"), + ) + for i in range(80) + ) + seen: Final = SimpleQueue[str]() + read_cursors: Final = SimpleQueue[str]() + expected: Final = finding("run").model_copy( + update={"evidence": (Evidence(execution_id="run", span_id="span079", quote="timeout"),)} + ) + + async def read(_identity: str, cursor: str, _offset: int) -> ExecutionContent: + read_cursors.put(cursor) + assert cursor in ("", "span039") + return ExecutionContent( + execution=execution, + parts=parts[:40] if not cursor else parts[40:], + next_cursor="span039" if not cursor else None, + ) + + async def model(request: ModelRequest) -> ModelResult: + payload: Final = json.loads(request.prompt) + if not payload["completed_read_count"]: + return ModelResult(content=json.dumps({"action": "read", "execution_id": "run"}), cost=0) + for part in payload["evidence"]: + seen.put(part["span_id"]) + if payload["evidence_page"] + 1 < payload["evidence_pages"]: + return ModelResult(content=json.dumps({"action": "evidence", "page": payload["evidence_page"] + 1}), cost=0) + if payload["last_read"]["next_cursor"]: + return ModelResult( + content=json.dumps( + {"action": "read", "execution_id": "run", "cursor": payload["last_read"]["next_cursor"]} + ), + cost=0, + ) + return ModelResult(content=json.dumps({"action": "submit", "finding": expected.model_dump()}), cost=0) + + claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) + result: Final = await investigate( + claim, + Candidate(check_id="retries", title="Failure", hypothesis="Failure", execution_ids=("run",)), + (Examined(execution=execution, observations=(), parts=(), partial=False, cannot_assess=False),), + read, + model, + ) + assert result.finding == expected + assert result.error == "" + assert tuple(seen.get_nowait() for _ in range(seen.qsize())) == tuple(p.span_id for p in parts) + assert tuple(read_cursors.get_nowait() for _ in range(read_cursors.qsize())) == ("", "span039") diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index a1441b34ffa..40a6fd55044 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -165,3 +165,45 @@ def test_regular_keys_cannot_read_lens_results(role: LitellmUserRoles | None) -> with pytest.raises(HTTPException) as error: user_scope(auth) assert error.value.status_code == 403 + + +@pytest.mark.parametrize("provider", (False, True)) +def test_model_errors_reach_worker_with_status_and_redacted_provider_message(provider: bool) -> None: + import httpx + + from litellm.proxy._types import ProxyException + from litellm.proxy.lens.endpoints import model_failure + from litellm.proxy.lens.worker import failure_message + + message: Final = "Token rate limit exceeded. api_key=secret-example-value-123456 Retry in 60 seconds." + error: Final = model_failure( + ProxyException(message, "rate_limit_error", None, 429, headers={"retry-after": "60"}) + if provider + else HTTPException(429, message, headers={"retry-after": "60"}) + ) + request: Final = httpx.Request("POST", "https://proxy.test/lens/worker/lens/run/model") + response: Final = httpx.Response(error.status_code, json={"detail": error.detail}, request=request) + with pytest.raises(httpx.HTTPStatusError) as caught: + response.raise_for_status() + diagnostic: Final = failure_message(caught.value) + assert diagnostic.startswith("Model request failed (HTTP 429):") + assert "Token rate limit exceeded." in diagnostic + assert "Retry in 60 seconds." in diagnostic + assert "secret-example" not in diagnostic + assert error.headers == {"retry-after": "60"} + + +@pytest.mark.asyncio +async def test_preview_reports_calendar_overflow_as_a_validation_error() -> None: + from datetime import datetime, timezone + + from litellm.proxy.lens.endpoints import Preview, preview_sample + + body: Final = Preview( + settings=LensSettings(name="Calendar regression", model="analysis", context="Read recorded activity"), + as_of=datetime.min.replace(tzinfo=timezone.utc), + ) + with pytest.raises(HTTPException) as error: + await preview_sample(body, UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), None) + assert error.value.status_code == 422 + assert "supported calendar range" in error.value.detail diff --git a/tests/unit/proxy/lens/test_inference.py b/tests/unit/proxy/lens/test_inference.py index 3b69d624a7e..e249693c4d5 100644 --- a/tests/unit/proxy/lens/test_inference.py +++ b/tests/unit/proxy/lens/test_inference.py @@ -15,6 +15,7 @@ def test_missing_optional_price_tiers_use_base_rates(monkeypatch: pytest.MonkeyP "openai/lens-base-rate-test": { "litellm_provider": "openai", "mode": "chat", + "max_output_tokens": 16384, "input_cost_per_token": 0.001, "output_cost_per_token": 0.002, "input_cost_per_token_above_200k_tokens": None, @@ -27,7 +28,10 @@ def test_missing_optional_price_tiers_use_base_rates(monkeypatch: pytest.MonkeyP deployment: Final = Deployment(litellm_params=DeploymentParams(model="openai/lens-base-rate-test")) explicit: Final = Deployment( litellm_params=DeploymentParams( - model="openai/lens-base-rate-test", input_cost_per_token=0.001, output_cost_per_token=0.002 + model="openai/lens-base-rate-test", + input_cost_per_token=0.001, + output_cost_per_token=0.002, + max_tokens=16384, ) ) assert quote((deployment,), "Answer the question") == quote((explicit,), "Answer the question") @@ -45,7 +49,7 @@ def test_unpriced_model_requires_explicit_rates() -> None: def test_custom_priced_model_charges_reported_tokens() -> None: deployment: Final = Deployment( litellm_params=DeploymentParams( - model="openai/lens-test", input_cost_per_token=0.001, output_cost_per_token=0.002 + model="openai/lens-test", input_cost_per_token=0.001, output_cost_per_token=0.002, max_tokens=16384 ) ) response: Final = ModelResponse( @@ -53,3 +57,66 @@ def test_custom_priced_model_charges_reported_tokens() -> None: ) assert completion_charge((deployment,), response, 10) == pytest.approx(0.04) assert quote((deployment,), "hello") > 0.04 + + +@pytest.mark.parametrize("capacity", (8192, 65536, 128000)) +def test_output_allowance_and_budget_follow_the_models_capacity(capacity: int, monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.lens.inference import output_tokens + + monkeypatch.setattr(litellm, "model_cost", {**litellm.model_cost}) + litellm.register_model( + model_cost={ + "openai/lens-capacity-test": { + "litellm_provider": "openai", + "mode": "chat", + "max_output_tokens": capacity, + "input_cost_per_token": 0, + "output_cost_per_token": 0.001, + } + } + ) + deployment: Final = Deployment(litellm_params=DeploymentParams(model="openai/lens-capacity-test")) + assert output_tokens(deployment) == capacity + assert quote((deployment,), "Review") == pytest.approx(capacity * 0.001) + + +def test_explicit_deployment_output_setting_is_respected() -> None: + from litellm.proxy.lens.inference import output_tokens + + deployment: Final = Deployment(litellm_params=DeploymentParams(model="custom/model", max_tokens=32000)) + assert output_tokens(deployment) == 32000 + + +def test_shared_context_capacity_leaves_room_for_the_entire_prompt(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.lens.inference import output_tokens + + monkeypatch.setattr(litellm, "model_cost", {**litellm.model_cost}) + litellm.register_model( + model_cost={ + "openai/lens-shared-context": { + "litellm_provider": "openai", + "mode": "chat", + "max_output_tokens": 8192, + "max_input_tokens": 8192, + "input_cost_per_token": 0, + "output_cost_per_token": 0.001, + } + } + ) + deployment: Final = Deployment(litellm_params=DeploymentParams(model="openai/lens-shared-context")) + short: Final = output_tokens(deployment, "Review this trace") + long: Final = output_tokens(deployment, "Review this trace " * 500) + assert 0 < long < short < output_tokens(deployment) + assert quote((deployment,), "Review this trace " * 500) == pytest.approx(long * 0.001) + + +def test_unknown_model_capacity_requires_explicit_operator_metadata() -> None: + from litellm.proxy.lens.inference import ModelCapacity, output_tokens + + params: Final = DeploymentParams(model="openai/lens-unknown-capacity") + with pytest.raises(HTTPException) as error: + output_tokens(Deployment(litellm_params=params)) + assert error.value.status_code == 400 + assert "model_info.max_output_tokens" in error.value.detail + configured: Final = Deployment(litellm_params=params, model_info=ModelCapacity(max_output_tokens=32000)) + assert output_tokens(configured) == 32000 diff --git a/tests/unit/proxy/lens/test_sources.py b/tests/unit/proxy/lens/test_sources.py index bae06f4afac..81bae7a0091 100644 --- a/tests/unit/proxy/lens/test_sources.py +++ b/tests/unit/proxy/lens/test_sources.py @@ -57,7 +57,10 @@ async def test_sample_never_returns_authentication_attributes() -> None: reader: Final = SourceReader(StorageResponse()) sample: Final = await reader.sample(Scope(team_id="alpha"), lens().settings, 1, 2) - assert sample.executions[0].metadata == (MetadataFilter(key="environment", value="production"),) + assert sample.executions[0].metadata == ( + MetadataFilter(key="environment", value="production"), + MetadataFilter(key="oversized", value="x" * 501), + ) assert "opaque-oauth-bearer" not in sample.model_dump_json() assert sample.eligible == 1 diff --git a/tests/unit/proxy/lens/test_state.py b/tests/unit/proxy/lens/test_state.py index c4220b7dd6d..fac43d6e850 100644 --- a/tests/unit/proxy/lens/test_state.py +++ b/tests/unit/proxy/lens/test_state.py @@ -103,7 +103,6 @@ def test_behavior_description_is_sufficient_without_separate_checks() -> None: ("sample_size", 0), ("concurrency", 0), ("lookback_hours", 0), - ("lookback_hours", 8761), ), ) def test_invalid_selection_and_parallelism_are_rejected(field: str, value: int) -> None: @@ -217,7 +216,7 @@ def test_custom_schedule_does_not_overlap_an_active_scan(interval: int) -> None: assert queue_job(running, NOW + timedelta(minutes=interval), "second") is running -@pytest.mark.parametrize("interval", (0, -1, 10081, 1.5)) +@pytest.mark.parametrize("interval", (0, -1, 1.5)) def test_invalid_schedule_is_rejected(interval: float) -> None: from pydantic import ValidationError @@ -286,3 +285,13 @@ def test_legacy_finding_identity_preserves_feedback_only_for_same_kind_and_check separate: Final = merge_finding(reviewed, other, 2, NOW) assert separate.id != legacy_id assert separate.status == "open" and separate.reason == "" + + +@pytest.mark.parametrize("field", ("lookback_hours", "interval_minutes")) +def test_calendar_overflow_is_rejected_without_the_old_history_and_interval_caps(field: str) -> None: + from pydantic import ValidationError + + accepted: Final = LensSettings.model_validate({**lens().settings.model_dump(), field: 100000}) + assert getattr(accepted, field) == 100000 + with pytest.raises(ValidationError, match="supported calendar range"): + LensSettings.model_validate({**lens().settings.model_dump(), field: 10**30}) diff --git a/tests/unit/proxy/lens/test_worker.py b/tests/unit/proxy/lens/test_worker.py index 0e64b0f10c8..17bac711e39 100644 --- a/tests/unit/proxy/lens/test_worker.py +++ b/tests/unit/proxy/lens/test_worker.py @@ -1,3 +1,4 @@ +import asyncio from queue import SimpleQueue from typing import Final @@ -89,7 +90,8 @@ async def test_incompatible_claim_reports_failure_instead_of_leaving_the_investi ) -> None: claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) payload: Final = claim.model_dump(mode="json") | { - "job": claim.job.model_dump(mode="json") | { + "job": claim.job.model_dump(mode="json") + | { "settings": claim.job.settings.model_dump() | {"future_setting": "private content"}, }, } @@ -197,3 +199,221 @@ def test_connection_timeout_and_invalid_response_have_distinct_private_diagnosti assert "connect to the proxy" in failure_message(httpx.ConnectError("private hostname")) assert "timed out" in failure_message(httpx.ReadTimeout("private prompt")) assert "structured JSON" in failure_message(ValueError("private model response")) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "purpose,stage,schema", + ( + ("extract", "Reading executions", "TraceReview"), + ("cluster", "Grouping observations", "Clusters"), + ("investigate", "Checking original evidence", "Decision"), + ), +) +async def test_worker_saves_validation_errors_from_every_analysis_stage(purpose: str, stage: str, schema: str) -> None: + import json + + claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) + execution: Final = Execution( + id="run", source="traces", trace_id="trace", team_id="alpha", name="review", start_time="", span_count=1 + ) + sample: Final = Sample(executions=(execution,), eligible=1) + content: Final = ExecutionContent( + execution=execution, + parts=(TracePart(execution_id="run", span_id="span", name="lead", kind="agent", content="Tool timeout"),), + ) + saved: Final = SimpleQueue[Result]() + attempts: Final = SimpleQueue[str]() + + def handle(request: httpx.Request) -> httpx.Response: + match request.url.path.rsplit("/", 1)[-1]: + case "claim": + return httpx.Response(200, json=claim.model_dump(mode="json")) + case "sample": + return httpx.Response(200, json=sample.model_dump(mode="json")) + case "content": + return httpx.Response(200, json=content.model_dump(mode="json")) + case "model": + body: Final = ModelRequest.model_validate_json(request.content) + if body.purpose == purpose: + attempts.put(body.purpose) + return httpx.Response( + 200, + json={"content": '{"candidates":[', "cost": 0.01}, + headers={"x-litellm-lens-finish-reason": "length"}, + ) + if body.purpose == "cluster": + return httpx.Response( + 200, + json={ + "content": json.dumps({"candidates": json.loads(body.prompt)["candidates"]}), + "cost": 0.01, + }, + ) + return httpx.Response( + 200, + json={ + "content": json.dumps( + { + "observations": [ + { + "check_id": claim.job.settings.analysis_checks[0].id, + "summary": "Tool timeout", + "evidence": [ + {"execution_id": "r0", "span_id": "span", "quote": "Tool timeout"} + ], + } + ] + } + ), + "cost": 0.01, + }, + ) + case "progress": + return httpx.Response(200, json=True) + case "result": + saved.put(Result.model_validate_json(request.content)) + return httpx.Response(200, json=True) + case _: + pytest.fail(f"Unexpected worker request: {request.url.path}") + + async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client: + assert await LensWorker(client).run_once() + message: Final = saved.get_nowait().error + assert message.startswith(f"{stage} failed: {schema} response invalid after 2 attempts.") + assert "finish_reason=length" in message + assert "EOF while parsing" in message and "[json_invalid]" in message + assert attempts.qsize() == 2 and saved.empty() + + +def test_response_validation_diagnostics_omit_input_values_and_unexpected_field_names() -> None: + with pytest.raises(ValidationError) as caught: + ModelResult.model_validate({"content": "private trace", "cost": "private token", "private field": "secret"}) + message: Final = failure_message(caught.value) + assert "Invalid ModelResult response" in message + assert "cost:" in message and "[float_parsing]" in message + assert "[extra_forbidden]" in message + assert "private" not in message and "secret" not in message + + +@pytest.mark.asyncio +@pytest.mark.parametrize("heartbeat_status", (401, 403, 409)) +async def test_losing_the_lease_interrupts_an_in_flight_model_request(heartbeat_status: int) -> None: + claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) + execution: Final = Execution( + id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=1 + ) + started: Final = asyncio.Event() + cancelled: Final = asyncio.Event() + never: Final = asyncio.Event() + saved: Final = SimpleQueue[Result]() + + async def heartbeat_wait(_seconds: float) -> None: + await started.wait() + + async def handle(request: httpx.Request) -> httpx.Response: + match request.url.path.rsplit("/", 1)[-1]: + case "claim": + return httpx.Response(200, json=claim.model_dump(mode="json")) + case "sample": + return httpx.Response(200, json=Sample(executions=(execution,), eligible=1).model_dump()) + case "content": + return httpx.Response( + 200, + json=ExecutionContent( + execution=execution, + parts=( + TracePart(execution_id="run", span_id="span", name="step", kind="tool", content="evidence"), + ), + ).model_dump(), + ) + case "model": + assert request.extensions["timeout"] == {"connect": 13, "read": None, "write": 13, "pool": 13} + started.set() + try: + await never.wait() + finally: + cancelled.set() + pytest.fail("The cancelled model request must not finish") + case "heartbeat": + return httpx.Response(heartbeat_status) + case "progress": + return httpx.Response(200, json=True) + case "result": + saved.put(Result.model_validate_json(request.content)) + return httpx.Response(409) + case _: + pytest.fail(f"Unexpected worker request: {request.url.path}") + + async with httpx.AsyncClient( + base_url="https://proxy.test", transport=httpx.MockTransport(handle), timeout=13 + ) as client: + assert await LensWorker(client, heartbeat_wait=heartbeat_wait).run_once() + assert cancelled.is_set() + assert f"HTTP {heartbeat_status}" in saved.get_nowait().error + assert saved.empty() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", (429, 500, 502, 503, 504, "connection", "timeout")) +async def test_transient_heartbeat_failure_recovers_without_cancelling_analysis(failure: int | str) -> None: + claim: Final = Claim(lens_id="lens", job=queue_job(lens(), NOW, "job").jobs[0], findings=()) + execution: Final = Execution( + id="run", source="traces", trace_id="t", team_id="", name="task", start_time="", span_count=1 + ) + started: Final = asyncio.Event() + recovered: Final = asyncio.Event() + never: Final = asyncio.Event() + attempts: Final = SimpleQueue[str]() + saved: Final = SimpleQueue[Result]() + + async def heartbeat_wait(_seconds: float) -> None: + await started.wait() + if attempts.qsize() >= 2: + await never.wait() + + async def handle(request: httpx.Request) -> httpx.Response: + match request.url.path.rsplit("/", 1)[-1]: + case "claim": + return httpx.Response(200, json=claim.model_dump(mode="json")) + case "sample": + return httpx.Response(200, json=Sample(executions=(execution,), eligible=1).model_dump()) + case "content": + return httpx.Response( + 200, + json=ExecutionContent( + execution=execution, + parts=( + TracePart(execution_id="run", span_id="span", name="step", kind="tool", content="evidence"), + ), + ).model_dump(), + ) + case "model": + started.set() + await recovered.wait() + return httpx.Response(200, json={"content": '{"observations":[],"cannot_assess":false}', "cost": 0.01}) + case "heartbeat": + attempts.put(request.url.path) + if attempts.qsize() == 1: + if failure == "connection": + raise httpx.ConnectError("temporary connection failure", request=request) + if failure == "timeout": + raise httpx.ReadTimeout("temporary response timeout", request=request) + assert isinstance(failure, int) + return httpx.Response(failure) + recovered.set() + return httpx.Response(200, json=True) + case "progress": + return httpx.Response(200, json=True) + case "result": + saved.put(Result.model_validate_json(request.content)) + return httpx.Response(200, json=True) + case _: + pytest.fail(f"Unexpected worker request: {request.url.path}") + + async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client: + assert await LensWorker(client, heartbeat_wait=heartbeat_wait).run_once() + result: Final = saved.get_nowait() + assert result.error == "" + assert result.coverage.screened == 1 and result.coverage.unassessable == 0 + assert attempts.qsize() == 2 and saved.empty() diff --git a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py index c6c81c14b16..31321254d94 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -9,15 +9,16 @@ from collections.abc import Callable, Mapping from contextlib import ExitStack, contextmanager from dataclasses import dataclass from io import BytesIO -from types import MappingProxyType, SimpleNamespace +from types import MappingProxyType, ModuleType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest import respx -from fastapi import HTTPException, Request, Response, UploadFile +from fastapi import APIRouter, FastAPI, HTTPException, Request, Response, UploadFile from fastapi.responses import StreamingResponse +from fastapi.testclient import TestClient from pydantic import TypeAdapter, ValidationError from starlette.datastructures import FormData, Headers, QueryParams from starlette.datastructures import UploadFile as StarletteUploadFile @@ -27,12 +28,14 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy._lazy_features import LazyFeature, attach_lazy_features from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, HttpPassThroughEndpointHelpers, InitPassThroughEndpointHelpers, + SafeRouteAdder, _registered_pass_through_routes, _truncate_upstream_error_body, _with_trace_context, @@ -8319,3 +8322,31 @@ async def test_pass_throughs_stay_open_while_a_db_sync_reads_the_database(tmp_pa await sync assert (config_during_sync.status_code, db_during_sync.status_code) == (200, 200) + + +def _lazy_feature(monkeypatch: pytest.MonkeyPatch, name: str, path: str) -> LazyFeature: + async def served() -> dict[str, str]: + return {"served_by": name} + + router: Final = APIRouter() + router.add_api_route(path, served, methods=["POST"]) + module: Final = ModuleType(f"tests.unit.proxy.pass_through_endpoints.lazy_fixture_{name}") + module.router = router # pyright: ignore[reportAttributeAccessIssue] # fixture module built at test time + monkeypatch.setitem(sys.modules, module.__name__, module) + return LazyFeature(name=name, module_path=module.__name__, path_prefixes=(path,)) + + +def test_a_pass_through_added_after_a_lazy_feature_loaded_takes_over_its_path(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("LITELLM_DISABLE_LAZY_ROUTES", raising=False) + + async def pass_through() -> dict[str, str]: + return {"served_by": "pass-through"} + + app: Final = FastAPI() + attach_lazy_features(app, (_lazy_feature(monkeypatch, "decider", "/v1/decider"),)) + with TestClient(app) as client: + assert client.post("/v1/decider").json() == {"served_by": "decider"} + assert SafeRouteAdder.add_api_route_if_not_exists(app, "/v1/decider", pass_through, ["POST"]) + assert client.post("/v1/decider").json() == {"served_by": "pass-through"} + assert not SafeRouteAdder.add_api_route_if_not_exists(app, "/v1/decider", pass_through, ["POST"]) + assert client.post("/v1/decider").json() == {"served_by": "pass-through"} diff --git a/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py b/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py index 44a75c362e5..5da08e6af75 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py +++ b/tests/unit/proxy/pass_through_endpoints/test_passthrough_auth_default.py @@ -19,19 +19,23 @@ defaults to ``True`` so a config dict (raw, not Pydantic) without an ``auth`` key still requires authentication. """ -from unittest.mock import AsyncMock, MagicMock +from typing import Final +from unittest.mock import MagicMock import pytest from fastapi import FastAPI +from fastapi.routing import APIRoute +from fastapi.testclient import TestClient -from litellm.proxy._types import PassThroughGenericEndpoint +from litellm.proxy._types import PassThroughGenericEndpoint, ProxyException from litellm.proxy.auth.user_api_key_auth import ( check_api_key_for_custom_headers_or_pass_through_endpoints, ) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( _register_pass_through_endpoint, ) +from litellm.proxy.proxy_server import openai_exception_handler def test_passthrough_auth_defaults_to_true(): @@ -57,26 +61,28 @@ def test_passthrough_auth_can_still_be_explicitly_disabled(): @pytest.mark.asyncio -async def test_register_passthrough_with_auth_true_works_for_oss(monkeypatch): - # Regression: setting ``auth: true`` used to raise at startup - # unless ``premium_user`` was True, leaving OSS with no safe - # configuration. - app = MagicMock(spec=FastAPI) - visited: set = set() +async def test_register_passthrough_with_auth_true_works_for_oss(monkeypatch: pytest.MonkeyPatch) -> None: + app: Final = FastAPI(exception_handlers={ProxyException: openai_exception_handler}) + visited: Final[set[str]] = set() + monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "sk-passthrough-test") - endpoint = PassThroughGenericEndpoint( + endpoint: Final = PassThroughGenericEndpoint( path="/forwarder", target="https://example.com", auth=True, ) - # Should not raise; OSS premium_user=False is allowed to use auth=True. await _register_pass_through_endpoint( endpoint=endpoint, app=app, premium_user=False, visited_endpoints=visited, ) + assert [route.path for route in app.routes if isinstance(route, APIRoute)] == ["/forwarder"] + with TestClient(app) as client: + response: Final = client.get(endpoint.path) + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "auth_error" @pytest.mark.asyncio diff --git a/tests/unit/proxy/public_endpoints/test_public_endpoints.py b/tests/unit/proxy/public_endpoints/test_public_endpoints.py index be309a67d58..448ad9c712e 100644 --- a/tests/unit/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/unit/proxy/public_endpoints/test_public_endpoints.py @@ -436,10 +436,12 @@ ADD_MODEL_UNLISTED_PROVIDERS: Final = frozenset( "sagemaker_nova", "scaleway", "stability", + "strands_decider", "synthetic", "tensormesh", "text-completion-inception", "transcribe", + "typesafe", "valkey", "xiaomi_mimo", "zai", diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index 16d02865e92..cff67d1f57d 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -266,7 +266,7 @@ def test_get_trace_404_and_200(client, receiver): response = client.get("/v1/traces/t1") assert response.status_code == 200 assert response.json() == TRACE_RESPONSE - receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "") + receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "", None, None) def test_get_span_404_and_200(client, receiver): @@ -278,10 +278,45 @@ def test_get_span_404_and_200(client, receiver): receiver.get_span.assert_awaited_with("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "") -def test_trace_detail_passes_scoped_reference(client, receiver): +@pytest.mark.parametrize("suffix,cursor,page_size", [("", None, None), ("&cursor=next&page_size=200", "next", 200)]) +def test_trace_detail_passes_scoped_reference(client, receiver, suffix, cursor, page_size): receiver.get_trace.return_value = TRACE_RESPONSE - assert client.get("/v1/traces/t1?trace_ref=run-one").status_code == 200 - receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one") + assert client.get(f"/v1/traces/t1?trace_ref=run-one{suffix}").status_code == 200 + receiver.get_trace.assert_awaited_with( + "t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one", cursor, page_size + ) + + +@pytest.mark.parametrize( + "path,method", + ( + ("/v1/traces", "list_traces"), + ("/v1/traces/t1", "get_trace"), + ("/v1/traces/t1/spans/s1", "get_span"), + ("/v1/traces/t1/spans/s1/error", "get_span_error"), + ), +) +@pytest.mark.parametrize( + "error,status,message", + ( + (RuntimeError("private database details"), 503, "Traces are temporarily unavailable. Please try again."), + (OverflowError("private query details"), 413, "Trace is too large for this view. Use a filtered trace query."), + ), +) +def test_read_failures_are_actionable_without_exposing_database_details( + client: TestClient, receiver: MagicMock, path: str, method: str, error: Exception, status: int, message: str +) -> None: + getattr(receiver, method).side_effect = error + response: Final = client.get(path) + assert response.status_code == status + assert response.json() == {"detail": message} + + +@pytest.mark.parametrize("query", ("page_size=0", "page_size=501", "cursor=" + "x" * 513)) +def test_trace_page_rejects_unbounded_parameters(client: TestClient, receiver: MagicMock, query: str) -> None: + response: Final = client.get(f"/v1/traces/t1?{query}") + assert response.status_code == 422 + receiver.get_trace.assert_not_awaited() def test_invalid_export_and_cursor_are_client_errors(client, receiver): @@ -370,7 +405,9 @@ def test_injected_receiver_ingests_with_the_authenticated_tenant(client: TestCli storage: Final = MagicMock(spec=ClickHouseStorage) storage.ingest = AsyncMock(return_value=1) client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: TraceReceiver(storage) - response: Final = client.post("/v1/traces", content=b'{"resourceSpans": []}', headers={"content-type": "application/json"}) + response: Final = client.post( + "/v1/traces", content=b'{"resourceSpans": []}', headers={"content-type": "application/json"} + ) assert response.status_code == 200, response.text assert response.json() == {} storage.ingest.assert_awaited_once_with( @@ -538,9 +575,7 @@ def test_sql_and_help_use_authenticated_scope( result: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) assert result.status_code == 200, result.text assert result.json() == SQL_ENVELOPE - receiver.storage.query_sql.assert_awaited_once_with( - "SELECT * FROM otel_traces", expected_scope, "test-secret" - ) + receiver.storage.query_sql.assert_awaited_once_with("SELECT * FROM otel_traces", expected_scope, "test-secret") help_result: Final = client.get("/v1/traces/query/help") assert help_result.status_code == 200, help_result.text assert help_result.json() == QUERY_HELP @@ -692,7 +727,7 @@ class _NativeConfig: class _NativeReturningHelp(ModuleType): - def __init__(self, help_payload: Mapping[str, object]) -> None: + def __init__(self, help_payload: Mapping[str, object], trace_payload: Mapping[str, object] | None = None) -> None: super().__init__("native_traces") class Storage: @@ -702,12 +737,31 @@ class _NativeReturningHelp(ModuleType): async def query_help(self, scope: AllQueryScope, secret: str) -> Mapping[str, object]: return help_payload + get_trace = AsyncMock(return_value=trace_payload) + + self.trace_read: Final = Storage.get_trace self.NativeTraceConfig: Final = _NativeConfig self.NativeTraceStorage: Final = Storage self.trace_encode_error: Final = bytes self.trace_span_rows: Final = list +@pytest.mark.parametrize("cursor,page_size", ((None, None), ("next", 200))) +async def test_storage_preserves_page_cursor_and_normalizes_native_trace_data( + monkeypatch: pytest.MonkeyPatch, cursor: str | None, page_size: int | None +) -> None: + native: Final = _NativeReturningHelp(QUERY_HELP, {**TRACE_RESPONSE, "next_cursor": "more"}) + monkeypatch.setattr(loader, "_cached_bridge", native) + storage: Final = ClickHouseStorage(TraceStorageConfig("http://clickhouse:8123")) + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "owner", "team_ids": ()} + trace: Final = await storage.get_trace("t1", scope, "run", cursor, page_size) + assert trace is not None + assert trace["next_cursor"] == "more" + assert trace["spans"] == () + assert trace["summary"]["span_count"] == 0 + native.trace_read.assert_awaited_once_with("t1", scope, "run", cursor, page_size) + + async def test_storage_validates_the_native_query_help_value(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(loader, "_cached_bridge", _NativeReturningHelp(QUERY_HELP)) storage: Final = ClickHouseStorage(TraceStorageConfig("http://clickhouse:8123")) diff --git a/tests/unit/test_lens_dev.py b/tests/unit/test_lens_dev.py new file mode 100644 index 00000000000..e14b27fd27e --- /dev/null +++ b/tests/unit/test_lens_dev.py @@ -0,0 +1,151 @@ +import os +import subprocess +import sys +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[2] +SCRIPT = ROOT / "scripts" / "lens_dev.sh" + +# Fake curl: answers the worker-token check with $CLAIM_STATUS, and key/generate and +# workers/register with a JSON "token". Every call is appended to $CURL_LOG. +FAKE_CURL = """#!/bin/sh +echo "$@" >> "$CURL_LOG" +case "$*" in + *worker/claim*) printf '%s' "$CLAIM_STATUS" ;; + */key/generate*) printf '{"token": "%064d"}' 0 ;; + */lens/workers/register*) printf '{"token": "lens-fresh"}' ;; +esac +""" + + +def _run(tmp_path: Path, snippet: str, **env: str) -> subprocess.CompletedProcess[str]: + bin_dir = tmp_path / "bin" + bin_dir.mkdir(exist_ok=True) + curl = bin_dir / "curl" + curl.write_text(FAKE_CURL) + curl.chmod(0o755) + state = tmp_path / "state" + state.mkdir(exist_ok=True) + return subprocess.run( + ["bash", "-c", f'source "{SCRIPT}"\n{snippet}'], + capture_output=True, + text=True, + env={ + "PATH": f"{bin_dir}{os.pathsep}/usr/bin{os.pathsep}/bin", + "HOME": str(tmp_path), + "LENS_DEV_STATE_DIR": str(state), + "LENS_DEV_PYTHON": sys.executable, + "CURL_LOG": str(tmp_path / "curl.log"), + "CLAIM_STATUS": "409", + **env, + }, + ) + + +def _curl_calls(tmp_path: Path) -> str: + log = tmp_path / "curl.log" + return log.read_text() if log.exists() else "" + + +def test_missing_worker_token_registers_a_worker(tmp_path): + proc = _run(tmp_path, "ensure_worker_token") + assert proc.returncode == 0, proc.stderr + assert (tmp_path / "state" / "worker_token").read_text().strip() == "lens-fresh" + assert oct((tmp_path / "state" / "worker_token").stat().st_mode & 0o777) == "0o600" + assert f'"analysis_key_id": "{0:064d}"' in _curl_calls(tmp_path) + + +def test_accepted_worker_token_is_reused(tmp_path): + (tmp_path / "state").mkdir() + (tmp_path / "state" / "worker_token").write_text("lens-saved\n") + proc = _run(tmp_path, "ensure_worker_token", CLAIM_STATUS="409") + assert proc.returncode == 0, proc.stderr + assert "reusing worker token" in proc.stdout + assert (tmp_path / "state" / "worker_token").read_text().strip() == "lens-saved" + assert "/lens/workers/register" not in _curl_calls(tmp_path) + + +def test_rejected_worker_token_is_replaced(tmp_path): + (tmp_path / "state").mkdir() + (tmp_path / "state" / "worker_token").write_text("lens-revoked\n") + proc = _run(tmp_path, "ensure_worker_token", CLAIM_STATUS="401") + assert proc.returncode == 0, proc.stderr + assert "was rejected" in proc.stdout + assert (tmp_path / "state" / "worker_token").read_text().strip() == "lens-fresh" + + +def test_unexpected_token_check_status_fails(tmp_path): + (tmp_path / "state").mkdir() + (tmp_path / "state" / "worker_token").write_text("lens-saved\n") + proc = _run(tmp_path, "ensure_worker_token", CLAIM_STATUS="500") + assert proc.returncode == 1 + assert "unexpected HTTP 500" in proc.stderr + + +def test_default_master_key_is_random_and_stable(tmp_path): + first = _run(tmp_path, 'load_master_key; echo "$master_key"') + second = _run(tmp_path, 'load_master_key; echo "$master_key"') + assert first.returncode == 0, first.stderr + key = first.stdout.strip() + assert key.startswith("sk-") and len(key) == 51 and key != "sk-1234" + assert second.stdout.strip() == key + assert oct((tmp_path / "state" / "master_key").stat().st_mode & 0o777) == "0o600" + + +def test_master_key_override_wins(tmp_path): + proc = _run(tmp_path, 'load_master_key; echo "$master_key"', LENS_DEV_MASTER_KEY="sk-mine") + assert proc.stdout.strip() == "sk-mine" + assert not (tmp_path / "state" / "master_key").exists() + + +def test_proxy_env_drops_inherited_redis_and_base_urls(tmp_path): + proc = _run( + tmp_path, + 'master_key=sk-strong; proxy_env "export OPENAI_API_KEY=from-dotenv"; env', + REDIS_HOST="redis.example", + REDIS_PORT="6379", + REDIS_PASSWORD="secret", + ANTHROPIC_BASE_URL="http://elsewhere", + OPENAI_BASE_URL="http://elsewhere", + ) + assert proc.returncode == 0, proc.stderr + names = {line.split("=", 1)[0] for line in proc.stdout.splitlines()} + assert not {n for n in names if n.startswith("REDIS_")} + assert not names & {"ANTHROPIC_BASE_URL", "OPENAI_BASE_URL"} + assert "OPENAI_API_KEY=from-dotenv" in proc.stdout + assert "LITELLM_MODE=PRODUCTION" in proc.stdout + assert "UI_PASSWORD=sk-strong" in proc.stdout + assert "LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY" not in names + + +def test_proxy_env_permits_the_weak_key_only_when_chosen(tmp_path): + proc = _run(tmp_path, 'master_key=sk-1234; proxy_env ""; env') + assert "LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true" in proc.stdout + + +def test_external_database_url_never_starts_compose_postgres(tmp_path): + docker = tmp_path / "bin" / "docker" + proc = _run( + tmp_path, + f"listening() {{ return 1; }}\n" + f"printf '#!/bin/sh\\necho \"$@\" > {tmp_path}/docker.log\\n' > {docker}; chmod +x {docker}\n" + "ensure_services", + LENS_DEV_DATABASE_URL="postgresql://elsewhere/db?schema=public", + ) + assert proc.returncode == 0, proc.stderr + assert (tmp_path / "docker.log").read_text().split()[-2:] == ["--wait", "clickhouse"] + + +def test_cleanup_kills_child_process_trees(tmp_path): + proc = _run( + tmp_path, + "set -m\n" + "(sleep 300 & wait) & pids+=($!)\n" + "child=$!; sleep 0.3\n" + "cleanup\n" + "sleep 0.3\n" + 'pgrep -g "$child" >/dev/null && echo LEFTOVER || echo CLEAN', + ) + assert proc.returncode == 0, proc.stderr + assert "lens-dev: stopping" in proc.stdout + assert proc.stdout.strip().endswith("CLEAN") diff --git a/ui/litellm-dashboard/src/app/globals.css b/ui/litellm-dashboard/src/app/globals.css index e1ff921648f..72d304c0dc7 100644 --- a/ui/litellm-dashboard/src/app/globals.css +++ b/ui/litellm-dashboard/src/app/globals.css @@ -473,3 +473,12 @@ [data-slot="dialog-content"][data-nested-dialog-open] { visibility: hidden; } + +@keyframes lens-sweep { + from { + left: -6rem; + } + to { + left: 100%; + } +} diff --git a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx index ab9a723e57a..bd0c3c07fc4 100644 --- a/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx +++ b/ui/litellm-dashboard/src/components/Navbar/UserDropdown/UserDropdown.tsx @@ -204,7 +204,7 @@ const UserDropdown: React.FC = ({ onLogout, variant = "navbar aria-label="Toggle hide bouncing icon" /> - {canUseLiteAdmin && ( + {premiumUser === true && canUseLiteAdmin && (
Hide LiteAdmin = ({ onLogout, colla />
))} - {canUseLiteAdmin && ( + {premiumUser === true && canUseLiteAdmin && (
Hide LiteAdmin onDemo(activeTab) : undefined; return ( -
-
-

-

-
-
+
{demo && } (demo ? setDemoTab(value as Tab) : void setTab(value as Tab))} - className="min-h-0 flex-1 gap-4" + className="min-h-0 flex-1 gap-2" > - - - Traces - - - Investigations - - - +
+
+

+

+ + + Traces + + + Investigations + + +
+
+
+ expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); expect(apiClient.request).not.toHaveBeenCalled(); }); + +it("shows the actual saved failure and run context without opening backend logs", async () => { + testQueryClient.clear(); + const error = + "Grouping observations failed: Clusters response invalid after 2 attempts.\n" + + "candidates.0.check_id: Field required [missing]"; + const job = { ...lens.jobs[0], id: "failed-run", status: "failed" as const, stage: "Failed", error, findings: [] }; + vi.mocked(apiClient.get).mockImplementation(async (path) => { + if (path === "/lens") return { lenses: [{ ...lens, jobs: [job] }], workers: [], tracing_enabled: true }; + if (path === "/lens/lens/runs") return [job]; + if (path === "/lens/lens/runs/failed-run") return job; + return { data: [] }; + }); + renderWithProviders(); + const failure = within(await screen.findByRole("alert")); + expect(failure.getByLabelText("Investigation error")).toHaveTextContent(error.replaceAll("\n", " ")); + expect(failure.getByText("failed-run")).toBeVisible(); + expect(failure.getByText(job.settings.model)).toBeVisible(); + expect(failure.queryByText(/find the error in proxy and worker logs/)).not.toBeInTheDocument(); +}); diff --git a/ui/litellm-dashboard/src/components/lens/investigations/detail/InvestigationFailure.tsx b/ui/litellm-dashboard/src/components/lens/investigations/detail/InvestigationFailure.tsx index 4fa906a1b7f..ac6dc65babb 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/detail/InvestigationFailure.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/detail/InvestigationFailure.tsx @@ -17,9 +17,11 @@ export function InvestigationFailure({ job, connected, className, ...props }: In className={cn("space-y-2 rounded-md border border-destructive/20 p-3 text-sm", className)} >

This investigation did not finish

-

{job.error}

-
- Troubleshooting details +
+        {job.error}
+      
+
+ Run details
Run:
@@ -38,10 +40,6 @@ export function InvestigationFailure({ job, connected, className, ...props }: In
{runTime(job.created_at)}
-

- Use the run ID to find the error in proxy and worker logs. Check the worker key's model permissions and - budget before retrying. -

); diff --git a/ui/litellm-dashboard/src/components/lens/setup/MatchingActivity.tsx b/ui/litellm-dashboard/src/components/lens/setup/MatchingActivity.tsx index da0ca45860e..38971aa19af 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/MatchingActivity.tsx +++ b/ui/litellm-dashboard/src/components/lens/setup/MatchingActivity.tsx @@ -50,7 +50,7 @@ export function MatchingActivity({ return () => clearTimeout(timer); }, [serialized]); const historyHours = selection.lookback_hours ?? 24; - const validWindow = Number.isInteger(historyHours) && historyHours >= 1 && historyHours <= 8760; + const validWindow = Number.isInteger(historyHours) && historyHours >= 1; const percent = scope.sample_percent ?? 100; const cap = scope.sample_size; const validCap = cap == null || (Number.isInteger(cap) && cap > 0); diff --git a/ui/litellm-dashboard/src/components/lens/setup/MonitoringDialog.tsx b/ui/litellm-dashboard/src/components/lens/setup/MonitoringDialog.tsx index b283673dcf6..9ddac5dd281 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/MonitoringDialog.tsx +++ b/ui/litellm-dashboard/src/components/lens/setup/MonitoringDialog.tsx @@ -15,7 +15,7 @@ import { DurationInput } from "@/components/shared/DurationInput"; import { type Settings } from "../model/types"; const monitoringSchema = z.object({ - interval_minutes: z.number().int().min(1).max(10080), + interval_minutes: z.number().int().min(1), }); export function MonitoringDialog({ @@ -62,13 +62,7 @@ export function MonitoringDialog({ control={control} name="interval_minutes" render={({ field }) => ( - + )} /> {formState.errors.interval_minutes?.message && ( diff --git a/ui/litellm-dashboard/src/components/lens/setup/fields/SampleFields.tsx b/ui/litellm-dashboard/src/components/lens/setup/fields/SampleFields.tsx index 5af37dadca5..305a28af34b 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/fields/SampleFields.tsx +++ b/ui/litellm-dashboard/src/components/lens/setup/fields/SampleFields.tsx @@ -18,13 +18,7 @@ export function SampleFields() { control={control} name="selection.lookback_hours" render={({ field }) => ( - + )} /> {errors.selection?.lookback_hours?.message && ( diff --git a/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.test.ts b/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.test.ts index 95e16a12a8a..7b034dcbe2c 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.test.ts +++ b/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.test.ts @@ -38,11 +38,11 @@ describe("investigation validation", () => { ); }); - it.each([0, 0.5, 8761, NaN, Infinity])("rejects an invalid history window: %s", (lookback_hours) => { + it.each([0, 0.5, NaN, Infinity])("rejects an invalid history window: %s", (lookback_hours) => { expectIssue( { selection: { lookback_hours } }, ["selection", "lookback_hours"], - "Choose a time range between 1 hour and 365 days", + "Choose a time range of at least 1 hour", ); }); @@ -62,10 +62,13 @@ describe("investigation validation", () => { ); }); - it.each([1, 8760])("accepts the history boundary %s with fractional sampling and no maximum", (lookback_hours) => { - const result = parse({ selection: { lookback_hours, sample_percent: 0.01, sample_size: null } }); - expect(result.success).toBe(true); - }); + it.each([1, 8760, 8761, 100000])( + "accepts the history window %s with fractional sampling and no maximum", + (lookback_hours) => { + const result = parse({ selection: { lookback_hours, sample_percent: 0.01, sample_size: null } }); + expect(result.success).toBe(true); + }, + ); it("requires individual runs only on the Run step", () => { expectIssue( @@ -89,8 +92,8 @@ describe("investigation validation", () => { }); it("validates budget and repeat interval with field-specific paths", () => { - expectIssue({ budget: Infinity }, ["budget"], "Choose a monthly limit greater than zero and up to 100000"); - expectIssue({ repeat: true, interval: 1.5 }, ["interval"], "Choose a repeat interval between 1 and 10080 minutes"); + expectIssue({ budget: Infinity }, ["budget"], "Choose a monthly limit greater than zero"); + expectIssue({ repeat: true, interval: 1.5 }, ["interval"], "Choose a repeat interval of at least 1 minute"); expect(parse({ repeat: false, interval: 0 }).success).toBe(true); expect(parse({ repeat: true, interval: 10080 }).success).toBe(true); }); diff --git a/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.ts b/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.ts index 0db08dbb836..2d3efc9abd2 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.ts +++ b/ui/litellm-dashboard/src/components/lens/setup/investigationSchema.ts @@ -64,10 +64,10 @@ function validateManualSelection(draft: InvestigationDraft, ctx: z.RefinementCtx function validateSampleWindow(draft: InvestigationDraft, ctx: z.RefinementCtx) { const selection = draft.selection; const hours = selection.lookback_hours ?? 24; - if (!Number.isInteger(hours) || hours < 1 || hours > 8760) { + if (!Number.isInteger(hours) || hours < 1) { ctx.addIssue({ code: "custom", - message: "Choose a time range between 1 hour and 365 days", + message: "Choose a time range of at least 1 hour", path: ["selection", "lookback_hours"], }); } @@ -89,19 +89,19 @@ function validateSampleWindow(draft: InvestigationDraft, ctx: z.RefinementCtx) { } function validateBudgetAndSchedule(draft: InvestigationDraft, ctx: z.RefinementCtx) { - if (!Number.isFinite(draft.budget) || draft.budget <= 0 || draft.budget > 100000) { + if (!Number.isFinite(draft.budget) || draft.budget <= 0) { ctx.addIssue({ code: "custom", - message: "Choose a monthly limit greater than zero and up to 100000", + message: "Choose a monthly limit greater than zero", path: ["budget"], }); } - const intervalOutOfRange = draft.interval < 1 || draft.interval > 10080; + const intervalOutOfRange = draft.interval < 1; const intervalInvalid = !Number.isInteger(draft.interval) || intervalOutOfRange; if (draft.repeat && intervalInvalid) { ctx.addIssue({ code: "custom", - message: "Choose a repeat interval between 1 and 10080 minutes", + message: "Choose a repeat interval of at least 1 minute", path: ["interval"], }); } @@ -208,7 +208,7 @@ export function investigationSettings( return { ...initial, ...draft.selection, - name: draft.name.trim() || suggestedName.slice(0, 100), + name: draft.name.trim() || suggestedName, context: draft.context, model, monthly_budget: draft.budget, diff --git a/ui/litellm-dashboard/src/components/lens/setup/steps/ExpectationsStep.tsx b/ui/litellm-dashboard/src/components/lens/setup/steps/ExpectationsStep.tsx index 3dc696aee8a..128d58b8126 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/steps/ExpectationsStep.tsx +++ b/ui/litellm-dashboard/src/components/lens/setup/steps/ExpectationsStep.tsx @@ -22,7 +22,6 @@ export function ExpectationsStep() { What should the agent be doing?