mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin/main' into litellm_otel_v2_team_capture_message_content
This commit is contained in:
commit
c3810ea755
149 changed files with 9770 additions and 492 deletions
3
.github/workflows/test-unit.yml
vendored
3
.github/workflows/test-unit.yml
vendored
|
|
@ -61,7 +61,7 @@ jobs:
|
|||
|
||||
- shard: core-utils
|
||||
artifact-name: core-utils
|
||||
test-path: ""
|
||||
test-path: tests/unit/decisions
|
||||
unit-flag: core-utils
|
||||
workers: 2
|
||||
reruns: 1
|
||||
|
|
@ -141,6 +141,7 @@ jobs:
|
|||
artifact-name: proxy-endpoints
|
||||
test-path: >-
|
||||
tests/unit/proxy/analytics_endpoints
|
||||
tests/unit/proxy/decisions_endpoints
|
||||
tests/unit/proxy/management_endpoints
|
||||
tests/unit/proxy/list_api
|
||||
tests/unit/proxy/memory
|
||||
|
|
|
|||
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -151,3 +151,6 @@ litellm.log
|
|||
|
||||
.coverage-rust
|
||||
coverage-rust.xml
|
||||
|
||||
# make lens-dev worker token, generated config and logs
|
||||
.lens-dev/
|
||||
|
|
|
|||
6
Makefile
6
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/
|
||||
|
||||
|
|
|
|||
|
|
@ -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) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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/",
|
||||
|
|
|
|||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -4478,6 +4478,7 @@ dependencies = [
|
|||
"flate2",
|
||||
"futures-util",
|
||||
"hmac 0.12.1",
|
||||
"itertools 0.14.0",
|
||||
"jsonschema",
|
||||
"litellm-http",
|
||||
"litellm-migrate",
|
||||
|
|
|
|||
|
|
@ -35,13 +35,14 @@ fn map_error_ref(error: &Error) -> PyErr {
|
|||
use litellm_storage_clickhouse::Error as StorageError;
|
||||
|
||||
match error {
|
||||
Error::Decode(litellm_traces::Error::TooLarge) | Error::InsertTooLarge => {
|
||||
PyOverflowError::new_err(error.to_string())
|
||||
}
|
||||
Error::Decode(litellm_traces::Error::TooLarge)
|
||||
| Error::InsertTooLarge
|
||||
| Error::ReadTooLarge => PyOverflowError::new_err(error.to_string()),
|
||||
Error::InvalidRow
|
||||
| Error::InvalidTable
|
||||
| Error::InvalidCursor(_)
|
||||
| Error::AmbiguousTrace
|
||||
| Error::TraceChanged
|
||||
| Error::Decode(_)
|
||||
| Error::InvalidSchema
|
||||
| Error::InvalidQuery
|
||||
|
|
@ -233,26 +234,44 @@ impl NativeTraceStorage {
|
|||
)
|
||||
}
|
||||
|
||||
#[pyo3(signature = (trace_id, scope, trace_ref, cursor=None, page_size=None))]
|
||||
fn get_trace<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
trace_id: String,
|
||||
#[pyo3(from_py_with = litellm_host_python::from_py_argument)] scope: ReadAccessParams,
|
||||
trace_ref: String,
|
||||
cursor: Option<String>,
|
||||
page_size: Option<u32>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let client = crate::http::host_client(py, ClientVariant::NoRedirect)?;
|
||||
let connection = self.config.storage().reader().clone();
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
litellm_traces_clickhouse::get_trace(
|
||||
&client,
|
||||
&connection,
|
||||
&scope,
|
||||
&trace_id,
|
||||
&trace_ref,
|
||||
)
|
||||
.await
|
||||
if let Some(page_size) = page_size {
|
||||
litellm_traces_clickhouse::get_trace_page(
|
||||
&client,
|
||||
&connection,
|
||||
&scope,
|
||||
&trace_id,
|
||||
&trace_ref,
|
||||
cursor.as_deref(),
|
||||
page_size,
|
||||
)
|
||||
.await
|
||||
} else if cursor.is_some() {
|
||||
Err(Error::InvalidParameters)
|
||||
} else {
|
||||
litellm_traces_clickhouse::get_trace(
|
||||
&client,
|
||||
&connection,
|
||||
&scope,
|
||||
&trace_id,
|
||||
&trace_ref,
|
||||
)
|
||||
.await
|
||||
}
|
||||
},
|
||||
map_error,
|
||||
)
|
||||
|
|
@ -465,6 +484,8 @@ mod tests {
|
|||
#[case::invalid_export(Error::Decode(litellm_traces::Error::InvalidPayload), "ValueError")]
|
||||
#[case::cursor(Error::InvalidCursor("trace"), "ValueError")]
|
||||
#[case::ambiguous(Error::AmbiguousTrace, "ValueError")]
|
||||
#[case::changed_snapshot(Error::TraceChanged, "ValueError")]
|
||||
#[case::read_budget(Error::ReadTooLarge, "OverflowError")]
|
||||
fn trace_read_and_ingest_failures_preserve_public_exception_types(
|
||||
#[case] error: Error,
|
||||
#[case] exception_name: &str,
|
||||
|
|
|
|||
|
|
@ -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()));
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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)));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
27
litellm-rust/crates/traces-clickhouse/query/spend_batch.sql
Normal file
27
litellm-rust/crates/traces-clickhouse/query/spend_batch.sql
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
SELECT * FROM (
|
||||
SELECT request_id, response_id, upstream_response_id, trace_id, span_id, team_id, api_key, user, spend,
|
||||
toUnixTimestamp64Milli(start_time) AS start_ms
|
||||
FROM (
|
||||
SELECT *,
|
||||
-- A chat request served through the Responses API returns the upstream `resp_` id to the
|
||||
-- client but logs LiteLLM's managed `resp_<base64>` id, which embeds it.
|
||||
if(startsWith(response_id, 'resp_'),
|
||||
extract(tryBase64Decode(substring(response_id, 6)), 'response_id:([^;]+)'),
|
||||
'') AS upstream_response_id
|
||||
FROM spend_logs FINAL
|
||||
WHERE start_time >= fromUnixTimestamp64Milli({start_ms:Int64})
|
||||
AND start_time < fromUnixTimestamp64Milli({end_ms:Int64})
|
||||
AND ({all_teams:UInt8} = 1
|
||||
OR ({user_id:String} != '' AND user = {user_id:String})
|
||||
OR has({team_ids:Array(String)}, team_id))
|
||||
)
|
||||
WHERE response_id IN {response_ids:Array(String)}
|
||||
OR upstream_response_id IN {response_ids:Array(String)}
|
||||
OR request_id IN {request_ids:Array(String)}
|
||||
OR (trace_id != '' AND trace_id IN {trace_ids:Array(String)})
|
||||
ORDER BY start_time DESC
|
||||
)
|
||||
WHERE {has_cursor:UInt8} = 0
|
||||
OR (team_id, start_ms, request_id) > ({after_team:String}, {after_ms:Int64}, {after_id:String})
|
||||
ORDER BY team_id, start_ms, request_id
|
||||
LIMIT {page_size:UInt32}
|
||||
|
|
@ -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}
|
||||
|
|
@ -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}
|
||||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,28 +1,59 @@
|
|||
//! Scoped trace reads: the trace list, one trace resolved with its spend, and span payloads.
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::{Arc, LazyLock};
|
||||
use std::time::Duration;
|
||||
|
||||
use base64::{Engine, engine::general_purpose::URL_SAFE};
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use itertools::Itertools;
|
||||
use litellm_http::Client;
|
||||
use litellm_storage_clickhouse::fetch;
|
||||
use litellm_storage_clickhouse::{Query, fetch};
|
||||
use litellm_traces::{
|
||||
SpanDetail, SpanErrorPage, SpendLookup, Trace, TracePage, listed_summary,
|
||||
query::named as contracts, resolve_trace, to_ui_content,
|
||||
};
|
||||
use moka::future::Cache;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use crate::{
|
||||
Connection, Error,
|
||||
query::named::{
|
||||
ListTraces, ListTracesParams, ReadAccessParams, SpanDetail as SpanDetailQuery,
|
||||
SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIds, SpendByResponseIdsParams,
|
||||
TraceIdentity, TraceIdentityParams, TracePageSpans, TracePageSpansParams, TraceSpans,
|
||||
TraceSpansParams,
|
||||
ListTracesParams, ListTracesRow, ReadAccessParams, SpanDetail as SpanDetailQuery,
|
||||
SpanDetailParams, SpanError, SpanErrorParams, SpendByResponseIdsParams, TraceIdentity,
|
||||
TraceIdentityParams, TracePageSpansParams, TraceSpansParams,
|
||||
},
|
||||
};
|
||||
|
||||
struct RunCandidates;
|
||||
|
||||
impl Query for RunCandidates {
|
||||
type Params = ListTracesParams;
|
||||
type Row = ListTracesRow;
|
||||
const SQL: &'static str = concat!(
|
||||
"SELECT * EXCEPT (request_ids), [] AS request_ids FROM (",
|
||||
include_str!("../query/list_traces.sql"),
|
||||
") ORDER BY start_ms DESC, trace_ref DESC"
|
||||
);
|
||||
}
|
||||
|
||||
// Cursor pages share a bounded snapshot so advancing does not resolve the whole graph again.
|
||||
static TRACE_SNAPSHOTS: LazyLock<Cache<String, Arc<Trace>>> = LazyLock::new(|| {
|
||||
Cache::builder()
|
||||
.max_capacity((2 * crate::span_batches::MAX_GRAPH_BYTES) as u64)
|
||||
.weigher(|_: &String, trace: &Arc<Trace>| {
|
||||
serde_json::to_vec(trace.as_ref())
|
||||
.ok()
|
||||
.and_then(|bytes| u32::try_from(bytes.len().saturating_mul(2)).ok())
|
||||
.unwrap_or(u32::MAX)
|
||||
})
|
||||
.time_to_live(Duration::from_secs(120))
|
||||
.build()
|
||||
});
|
||||
|
||||
const NANOS_PER_MS: i64 = 1_000_000;
|
||||
const SPEND_WINDOW_MS: i64 = 30 * 60 * 1000;
|
||||
const SPEND_CONCURRENCY: usize = 4;
|
||||
|
||||
fn encode_cursor<T: Serialize>(position: &T) -> String {
|
||||
URL_SAFE.encode(serde_json::to_vec(position).unwrap_or_default())
|
||||
|
|
@ -121,8 +152,8 @@ async fn spend(
|
|||
start_ms: start_ns.div_euclid(NANOS_PER_MS) - SPEND_WINDOW_MS,
|
||||
end_ms: end_ns.div_euclid(NANOS_PER_MS) + SPEND_WINDOW_MS,
|
||||
});
|
||||
match fetch::<SpendByResponseIds>(client, connection, ¶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<TracePage, Error> {
|
||||
if limit == 0 {
|
||||
return Err(Error::InvalidParameters);
|
||||
}
|
||||
let (cursor_ms, cursor_trace_id) = trace_position(cursor)?;
|
||||
let params = ListTracesParams::from(contracts::ListTracesParams {
|
||||
let mut params = ListTracesParams::from(contracts::ListTracesParams {
|
||||
access: access.clone(),
|
||||
start_ms,
|
||||
end_ms,
|
||||
cursor_ms,
|
||||
cursor_trace_id,
|
||||
limit,
|
||||
limit: limit.min(500),
|
||||
});
|
||||
let page: Vec<contracts::ListTracesRow> = fetch::<ListTraces>(client, connection, ¶ms)
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|row| row.0)
|
||||
.collect();
|
||||
let page: Vec<contracts::ListTracesRow> = loop {
|
||||
match fetch::<RunCandidates>(client, connection, ¶ms).await {
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) if params.0.limit > 1 => {
|
||||
params.0.limit /= 2;
|
||||
}
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
result => break result?.into_iter().map(|row| row.0).collect(),
|
||||
}
|
||||
};
|
||||
let next_cursor = page
|
||||
.last()
|
||||
.filter(|_| page.len() == limit as usize)
|
||||
.filter(|_| page.len() == params.0.limit as usize)
|
||||
.map(|last| encode_cursor(&(last.start_ms, &last.trace_ref)));
|
||||
let (Some(page_start), Some(page_end)) = (
|
||||
page.iter().map(|row| row.start_ms).min(),
|
||||
page.iter().map(|row| row.start_ms + row.duration_ms).max(),
|
||||
let data = stream::iter(page.chunks(16))
|
||||
.then(|batch| list_summaries(client, connection, access, batch))
|
||||
.try_collect::<Vec<_>>()
|
||||
.await?
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.collect();
|
||||
Ok(TracePage { data, next_cursor })
|
||||
}
|
||||
|
||||
async fn list_summaries(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
access: &ReadAccessParams,
|
||||
runs: &[contracts::ListTracesRow],
|
||||
) -> Result<Vec<litellm_traces::TraceSummary>, Error> {
|
||||
let (Some(start_ms), Some(end_ms)) = (
|
||||
runs.iter().map(|row| row.start_ms).min(),
|
||||
runs.iter()
|
||||
.map(|row| row.start_ms.saturating_add(row.duration_ms))
|
||||
.max(),
|
||||
) else {
|
||||
return Ok(TracePage {
|
||||
data: Vec::new(),
|
||||
next_cursor,
|
||||
});
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
let span_params = TracePageSpansParams::from(contracts::TracePageSpansParams {
|
||||
let params = TracePageSpansParams::from(contracts::TracePageSpansParams {
|
||||
access: access.clone(),
|
||||
trace_refs: page.iter().map(|row| row.trace_ref.clone()).collect(),
|
||||
start_ms: page_start,
|
||||
end_ms: page_end + 1,
|
||||
trace_refs: runs.iter().map(|row| row.trace_ref.clone()).collect(),
|
||||
start_ms,
|
||||
end_ms: end_ms.saturating_add(1),
|
||||
});
|
||||
let span_rows: Vec<contracts::TraceSpansRow> =
|
||||
fetch::<TracePageSpans>(client, connection, &span_params)
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|row| row.0)
|
||||
.collect();
|
||||
let spend_rows = spend(client, connection, access, &span_rows).await;
|
||||
let mut by_trace: HashMap<(String, String, String), Vec<contracts::TraceSpansRow>> =
|
||||
HashMap::new();
|
||||
for span in span_rows {
|
||||
let key = (
|
||||
let spans = match crate::span_batches::read_list_spans(client, connection, params).await {
|
||||
Ok(spans) => spans,
|
||||
Err(Error::ReadTooLarge) => {
|
||||
return stream::iter(runs)
|
||||
.then(|row| async move {
|
||||
match get_trace(client, connection, access, &row.trace_id, &row.trace_ref).await
|
||||
{
|
||||
Ok(trace) => {
|
||||
Ok(trace.map_or_else(|| listed_summary(row), |trace| trace.summary))
|
||||
}
|
||||
Err(Error::ReadTooLarge) => Ok(listed_summary(row)),
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
})
|
||||
.try_collect()
|
||||
.await;
|
||||
}
|
||||
Err(error) => return Err(error),
|
||||
};
|
||||
let by_trace = spans.into_iter().into_group_map_by(|span| {
|
||||
(
|
||||
span.team_id.clone(),
|
||||
span.api_key_hash.clone(),
|
||||
span.trace_id.clone(),
|
||||
);
|
||||
by_trace.entry(key).or_default().push(span);
|
||||
}
|
||||
let data = page
|
||||
)
|
||||
});
|
||||
let summaries = runs
|
||||
.iter()
|
||||
.map(|row| {
|
||||
let spans = by_trace
|
||||
|
|
@ -200,11 +264,17 @@ pub async fn list_traces(
|
|||
))
|
||||
.map(Vec::as_slice)
|
||||
.unwrap_or_default();
|
||||
resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows)
|
||||
.map_or_else(|| listed_summary(row), |trace| trace.summary)
|
||||
async move {
|
||||
let spend_rows = spend(client, connection, access, spans).await;
|
||||
resolve_trace(&row.trace_id, &row.trace_ref, spans, &spend_rows)
|
||||
.map_or_else(|| listed_summary(row), |trace| trace.summary)
|
||||
}
|
||||
})
|
||||
.collect();
|
||||
Ok(TracePage { data, next_cursor })
|
||||
.collect::<Vec<_>>();
|
||||
Ok(stream::iter(summaries)
|
||||
.buffered(SPEND_CONCURRENCY)
|
||||
.collect()
|
||||
.await)
|
||||
}
|
||||
|
||||
pub async fn get_trace(
|
||||
|
|
@ -222,11 +292,7 @@ pub async fn get_trace(
|
|||
trace_id: trace_id.to_owned(),
|
||||
trace_ref: trace_ref.clone(),
|
||||
};
|
||||
let rows: Vec<contracts::TraceSpansRow> = fetch::<TraceSpans>(client, connection, ¶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<Option<Trace>, Error> {
|
||||
if !(1..=500).contains(&page_size) {
|
||||
return Err(Error::InvalidParameters);
|
||||
}
|
||||
let Some(trace_ref) = reference(client, connection, access, trace_id, trace_ref).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let position = match cursor {
|
||||
Some(cursor) => {
|
||||
let position: SpanPosition = decode_cursor(cursor, "span")?;
|
||||
if position.trace_ref != trace_ref || position.snapshot_ms == 0 {
|
||||
return Err(Error::InvalidCursor("span"));
|
||||
}
|
||||
position
|
||||
}
|
||||
None => SpanPosition {
|
||||
trace_ref: trace_ref.clone(),
|
||||
snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000)
|
||||
as u64,
|
||||
offset: 0,
|
||||
version: String::new(),
|
||||
},
|
||||
};
|
||||
let key_bytes = serde_json::to_vec(&(
|
||||
connection.url().as_str(),
|
||||
access,
|
||||
trace_id,
|
||||
&trace_ref,
|
||||
position.snapshot_ms,
|
||||
))
|
||||
.map_err(|_| Error::InvalidParameters)?;
|
||||
let key = format!("{:x}", Sha256::digest(key_bytes));
|
||||
let snapshot = if let Some(trace) = TRACE_SNAPSHOTS.get(&key).await {
|
||||
trace
|
||||
} else {
|
||||
let params = TraceSpansParams {
|
||||
access: access.clone(),
|
||||
trace_id: trace_id.to_owned(),
|
||||
trace_ref: trace_ref.clone(),
|
||||
};
|
||||
let rows =
|
||||
crate::span_batches::read_spans(client, connection, params, position.snapshot_ms)
|
||||
.await?;
|
||||
let spend_rows = spend(client, connection, access, &rows).await;
|
||||
let Some(trace) = resolve_trace(trace_id, &trace_ref, &rows, &spend_rows) else {
|
||||
return Ok(None);
|
||||
};
|
||||
if serde_json::to_vec(&trace)
|
||||
.map_err(|_| Error::InvalidResponse)?
|
||||
.len()
|
||||
> crate::span_batches::MAX_GRAPH_BYTES
|
||||
{
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
let trace = Arc::new(trace);
|
||||
TRACE_SNAPSHOTS.insert(key, Arc::clone(&trace)).await;
|
||||
trace
|
||||
};
|
||||
let span_ids: Vec<&str> = snapshot
|
||||
.spans
|
||||
.iter()
|
||||
.map(|span| span.span_id.as_str())
|
||||
.collect();
|
||||
let version = format!(
|
||||
"{:x}",
|
||||
Sha256::digest(serde_json::to_vec(&span_ids).map_err(|_| Error::InvalidResponse)?)
|
||||
);
|
||||
if cursor.is_some() && position.version != version {
|
||||
return Err(Error::TraceChanged);
|
||||
}
|
||||
let mut trace = Trace {
|
||||
summary: snapshot.summary.clone(),
|
||||
agents: snapshot.agents.clone(),
|
||||
spans: Vec::new(),
|
||||
next_cursor: None,
|
||||
};
|
||||
if position.offset > snapshot.spans.len() {
|
||||
return Err(Error::InvalidCursor("span"));
|
||||
}
|
||||
let end = position
|
||||
.offset
|
||||
.saturating_add(page_size as usize)
|
||||
.min(snapshot.spans.len());
|
||||
trace.next_cursor = (end < snapshot.spans.len()).then(|| {
|
||||
encode_cursor(&SpanPosition {
|
||||
offset: end,
|
||||
version: version.clone(),
|
||||
..position
|
||||
})
|
||||
});
|
||||
trace.spans = snapshot.spans[position.offset..end].to_vec();
|
||||
while serde_json::to_vec(&trace)
|
||||
.map_err(|_| Error::InvalidResponse)?
|
||||
.len()
|
||||
> litellm_storage_clickhouse::READ_LIMITS.response_bytes
|
||||
{
|
||||
if trace.spans.len() <= 1 {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
trace.spans.truncate(trace.spans.len() / 2);
|
||||
trace.next_cursor = Some(encode_cursor(&SpanPosition {
|
||||
trace_ref: trace_ref.clone(),
|
||||
snapshot_ms: position.snapshot_ms,
|
||||
offset: position.offset + trace.spans.len(),
|
||||
version: version.clone(),
|
||||
}));
|
||||
}
|
||||
Ok(Some(trace))
|
||||
}
|
||||
|
||||
pub async fn get_span(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
|
|
|
|||
270
litellm-rust/crates/traces-clickhouse/src/span_batches.rs
Normal file
270
litellm-rust/crates/traces-clickhouse/src/span_batches.rs
Normal file
|
|
@ -0,0 +1,270 @@
|
|||
use futures_util::{TryStreamExt, stream};
|
||||
use itertools::Itertools;
|
||||
use litellm_http::Client;
|
||||
use litellm_storage_clickhouse::{Query, fetch};
|
||||
use litellm_traces::query::named as contracts;
|
||||
use serde::Serialize;
|
||||
|
||||
use crate::{Connection, Error, query::named::TraceSpansRow};
|
||||
|
||||
const PAGE_SIZE: u32 = 256;
|
||||
pub(crate) const MAX_GRAPH_BYTES: usize = 64 * 1024 * 1024;
|
||||
const MAX_GRAPH_SPANS: usize = 100_000;
|
||||
|
||||
#[derive(Default)]
|
||||
struct ReadBudget {
|
||||
bytes: usize,
|
||||
rows: usize,
|
||||
}
|
||||
|
||||
impl ReadBudget {
|
||||
fn checked_add(&self, bytes: usize, rows: usize) -> Result<Self, Error> {
|
||||
let next = Self {
|
||||
bytes: self.bytes.saturating_add(bytes),
|
||||
rows: self.rows.saturating_add(rows),
|
||||
};
|
||||
if next.bytes > MAX_GRAPH_BYTES || next.rows > MAX_GRAPH_SPANS {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
Ok(next)
|
||||
}
|
||||
|
||||
fn record(&mut self, row: &impl Serialize) -> Result<(), Error> {
|
||||
let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?;
|
||||
*self = self.checked_add(bytes.len(), 1)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct Parameters {
|
||||
#[serde(flatten)]
|
||||
trace: contracts::TraceSpansParams,
|
||||
after_span_id: String,
|
||||
page_size: u32,
|
||||
snapshot_ms: u64,
|
||||
}
|
||||
|
||||
struct SpanBatch;
|
||||
|
||||
impl Query for SpanBatch {
|
||||
type Params = Parameters;
|
||||
type Row = TraceSpansRow;
|
||||
|
||||
const SQL: &'static str = include_str!("../query/trace_span_batch.sql");
|
||||
}
|
||||
|
||||
pub(crate) async fn read_spans(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
trace: contracts::TraceSpansParams,
|
||||
snapshot_ms: u64,
|
||||
) -> Result<Vec<contracts::TraceSpansRow>, Error> {
|
||||
let mut parameters = Parameters {
|
||||
trace,
|
||||
after_span_id: String::new(),
|
||||
page_size: PAGE_SIZE,
|
||||
snapshot_ms,
|
||||
};
|
||||
let mut spans = Vec::new();
|
||||
let mut budget = ReadBudget::default();
|
||||
loop {
|
||||
let page = match fetch::<SpanBatch>(client, connection, ¶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<Vec<contracts::TraceSpansRow>, Error> {
|
||||
let parameters = ListParameters {
|
||||
runs,
|
||||
after_team: String::new(),
|
||||
after_key: String::new(),
|
||||
after_trace: String::new(),
|
||||
after_span: String::new(),
|
||||
page_size: PAGE_SIZE,
|
||||
snapshot_ms: (time::OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64,
|
||||
};
|
||||
let pages = stream::try_unfold(
|
||||
(Some(parameters), ReadBudget::default()),
|
||||
|(parameters, budget)| async move {
|
||||
let Some(parameters) = parameters else {
|
||||
return Ok(None);
|
||||
};
|
||||
let page = match fetch::<ListSpanBatch>(client, connection, ¶meters).await {
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge)
|
||||
if parameters.page_size > 1 =>
|
||||
{
|
||||
let retry = ListParameters {
|
||||
page_size: parameters.page_size / 2,
|
||||
..parameters
|
||||
};
|
||||
return Ok(Some((Vec::new(), (Some(retry), budget))));
|
||||
}
|
||||
Err(litellm_storage_clickhouse::Error::ResponseTooLarge) => {
|
||||
return Err(Error::ReadTooLarge);
|
||||
}
|
||||
result => result?,
|
||||
};
|
||||
let next = page
|
||||
.last()
|
||||
.filter(|_| page.len() == parameters.page_size as usize)
|
||||
.map(|last| ListParameters {
|
||||
after_team: last.0.team_id.clone(),
|
||||
after_key: last.0.api_key_hash.clone(),
|
||||
after_trace: last.0.trace_id.clone(),
|
||||
after_span: last.0.span_id.clone(),
|
||||
page_size: (parameters.page_size * 2).min(PAGE_SIZE),
|
||||
..parameters
|
||||
});
|
||||
let next_budget = page.iter().try_fold(budget, |budget, row| {
|
||||
let bytes = serde_json::to_vec(row).map_err(|_| Error::InvalidResponse)?;
|
||||
budget.checked_add(bytes.len(), 1)
|
||||
})?;
|
||||
Ok(Some((page, (next, next_budget))))
|
||||
},
|
||||
)
|
||||
.try_collect::<Vec<_>>()
|
||||
.await?;
|
||||
Ok(pages
|
||||
.into_iter()
|
||||
.flatten()
|
||||
.map(|row| row.0)
|
||||
.sorted_by_key(|row| row.start_ns)
|
||||
.collect())
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct SpendParameters {
|
||||
#[serde(flatten)]
|
||||
lookup: crate::query::named::SpendByResponseIdsParams,
|
||||
has_cursor: u8,
|
||||
after_team: String,
|
||||
after_ms: i64,
|
||||
after_id: String,
|
||||
page_size: u32,
|
||||
}
|
||||
|
||||
struct SpendBatch;
|
||||
|
||||
impl Query for SpendBatch {
|
||||
type Params = SpendParameters;
|
||||
type Row = crate::query::named::SpendByResponseIdsRow;
|
||||
|
||||
const SQL: &'static str = include_str!("../query/spend_batch.sql");
|
||||
}
|
||||
|
||||
pub(crate) async fn read_spend(
|
||||
client: &Client,
|
||||
connection: &Connection,
|
||||
lookup: crate::query::named::SpendByResponseIdsParams,
|
||||
) -> Result<Vec<contracts::SpendByResponseIdsRow>, Error> {
|
||||
let mut parameters = SpendParameters {
|
||||
lookup,
|
||||
has_cursor: 0,
|
||||
after_team: String::new(),
|
||||
after_ms: 0,
|
||||
after_id: String::new(),
|
||||
page_size: PAGE_SIZE,
|
||||
};
|
||||
let mut rows = Vec::new();
|
||||
let mut budget = ReadBudget::default();
|
||||
loop {
|
||||
let page = match fetch::<SpendBatch>(client, connection, ¶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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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:?}"
|
||||
|
|
|
|||
595
litellm-rust/crates/traces-clickhouse/tests/reads.rs
Normal file
595
litellm-rust/crates/traces-clickhouse/tests/reads.rs
Normal file
|
|
@ -0,0 +1,595 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use litellm_traces::query::named::ReadAccessParams;
|
||||
use litellm_traces_clickhouse::{
|
||||
Connection, InsertTable, QueryScope, get_trace, get_trace_page, insert_rows, list_traces,
|
||||
};
|
||||
use rstest::rstest;
|
||||
use serde_json::json;
|
||||
|
||||
#[path = "queries/support.rs"]
|
||||
mod fixtures;
|
||||
mod support;
|
||||
|
||||
use fixtures::{DATABASE, SeededDatabase, migrated_database, seeded_database};
|
||||
use support::TestResult;
|
||||
|
||||
#[rstest]
|
||||
#[case::api_key("key-a", "")]
|
||||
#[case::user("", "user-a")]
|
||||
#[tokio::test]
|
||||
async fn list_costs_match_each_run_when_response_ids_are_reused(
|
||||
#[future(awt)] migrated_database: TestResult<SeededDatabase>,
|
||||
#[case] api_key: &str,
|
||||
#[case] user_id: &str,
|
||||
) -> TestResult {
|
||||
let fixture = migrated_database?;
|
||||
let client = &fixture.database.client;
|
||||
let writer = Connection::writer(&fixture.database.url)?;
|
||||
let runs = [
|
||||
("earlier-run", 1_790_000_000_000_i64, 0.25),
|
||||
("later-run", 1_790_007_200_000_i64, 0.75),
|
||||
];
|
||||
insert_rows(
|
||||
client,
|
||||
&writer,
|
||||
DATABASE,
|
||||
InsertTable::OtelTraces,
|
||||
runs.iter()
|
||||
.map(|(trace_id, start_ms, _)| {
|
||||
BTreeMap::from([
|
||||
("Timestamp".into(), json!(start_ms * 1_000_000)),
|
||||
("TraceId".into(), json!(trace_id)),
|
||||
("SpanId".into(), json!("llm-span")),
|
||||
("ObservationType".into(), json!("llm")),
|
||||
("TeamId".into(), json!("team-a")),
|
||||
("ApiKeyHash".into(), json!(api_key)),
|
||||
("UserId".into(), json!(user_id)),
|
||||
("Duration".into(), json!(1_000_000)),
|
||||
("LiteLLMRequestId".into(), json!("reused-response")),
|
||||
])
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
.await?;
|
||||
insert_rows(
|
||||
client,
|
||||
&writer,
|
||||
DATABASE,
|
||||
InsertTable::SpendLogs,
|
||||
runs.iter()
|
||||
.map(|(trace_id, start_ms, cost)| {
|
||||
BTreeMap::from([
|
||||
("request_id".into(), json!(format!("request-{trace_id}"))),
|
||||
("response_id".into(), json!("reused-response")),
|
||||
("team_id".into(), json!("team-a")),
|
||||
("api_key".into(), json!(api_key)),
|
||||
("user".into(), json!(user_id)),
|
||||
("start_time".into(), json!(start_ms)),
|
||||
("end_time".into(), json!(start_ms + 1)),
|
||||
("spend".into(), json!(cost)),
|
||||
])
|
||||
})
|
||||
.collect(),
|
||||
)
|
||||
.await?;
|
||||
let reader = fixture
|
||||
.readers
|
||||
.connection(client, &QueryScope::All, "fixture-secret")
|
||||
.await?;
|
||||
let access = ReadAccessParams {
|
||||
all_teams: false,
|
||||
user_id: user_id.into(),
|
||||
team_ids: vec!["team-a".into()],
|
||||
};
|
||||
let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
|
||||
assert_eq!(page.data.len(), runs.len());
|
||||
for (trace_id, _, cost) in runs {
|
||||
let summary = page
|
||||
.data
|
||||
.iter()
|
||||
.find(|summary| summary.trace_id == trace_id)
|
||||
.ok_or("missing run")?;
|
||||
let detail = get_trace(client, &reader, &access, trace_id, &summary.trace_ref)
|
||||
.await?
|
||||
.ok_or("missing trace")?;
|
||||
assert_eq!(detail.summary.spend, Some(cost));
|
||||
assert_eq!(summary.spend, detail.summary.spend, "{trace_id}");
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::many_runs(50, 21, 0, false)]
|
||||
#[case::one_large_run(1, 1100, 0, false)]
|
||||
#[case::large_rows(1, 280, 20_000, false)]
|
||||
#[case::large_cached_snapshot(1, 280, 140_000, false)]
|
||||
#[case::many_costs(1, 1101, 0, true)]
|
||||
#[case::many_costed_runs(500, 2, 0, true)]
|
||||
#[tokio::test]
|
||||
async fn large_runs_remain_complete_under_default_reader_limits(
|
||||
#[future(awt)] migrated_database: TestResult<SeededDatabase>,
|
||||
#[case] runs: usize,
|
||||
#[case] steps: usize,
|
||||
#[case] name_bytes: usize,
|
||||
#[case] costed: bool,
|
||||
) -> TestResult {
|
||||
let fixture = migrated_database?;
|
||||
let client = &fixture.database.client;
|
||||
let writer = Connection::writer(&fixture.database.url)?;
|
||||
let rows = (0..runs)
|
||||
.flat_map(|run| {
|
||||
(0..steps).map(move |step| {
|
||||
BTreeMap::from([
|
||||
(
|
||||
"Timestamp".into(),
|
||||
json!(1_790_000_000_000_000_000_i64 + step as i64),
|
||||
),
|
||||
("TraceId".into(), json!(format!("trace-{run:04}"))),
|
||||
("SpanId".into(), json!(format!("span-{step:04}"))),
|
||||
(
|
||||
"ParentSpanId".into(),
|
||||
json!(if step == 0 { "" } else { "span-0000" }),
|
||||
),
|
||||
(
|
||||
"SpanName".into(),
|
||||
json!(if name_bytes == 0 {
|
||||
format!("step-{step}")
|
||||
} else {
|
||||
"x".repeat(name_bytes)
|
||||
}),
|
||||
),
|
||||
(
|
||||
"ObservationType".into(),
|
||||
json!(if step == 0 {
|
||||
"agent"
|
||||
} else if costed {
|
||||
"llm"
|
||||
} else {
|
||||
"tool"
|
||||
}),
|
||||
),
|
||||
("TeamId".into(), json!("team-a")),
|
||||
("ApiKeyHash".into(), json!("key-a")),
|
||||
("Duration".into(), json!(1000)),
|
||||
(
|
||||
"LiteLLMRequestId".into(),
|
||||
json!(if costed && step > 0 {
|
||||
format!("response-{step}")
|
||||
} else {
|
||||
String::new()
|
||||
}),
|
||||
),
|
||||
])
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
for chunk in rows.chunks(100) {
|
||||
insert_rows(
|
||||
client,
|
||||
&writer,
|
||||
DATABASE,
|
||||
InsertTable::OtelTraces,
|
||||
chunk.to_vec(),
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
if costed {
|
||||
let costs = (1..steps)
|
||||
.map(|step| {
|
||||
BTreeMap::from([
|
||||
("request_id".into(), json!(format!("request-{step}"))),
|
||||
("response_id".into(), json!(format!("response-{step}"))),
|
||||
("team_id".into(), json!("team-a")),
|
||||
("api_key".into(), json!("key-a")),
|
||||
("start_time".into(), json!(1_790_000_000_000_i64)),
|
||||
("end_time".into(), json!(1_790_000_000_001_i64)),
|
||||
("spend".into(), json!(0.25)),
|
||||
])
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
insert_rows(client, &writer, DATABASE, InsertTable::SpendLogs, costs).await?;
|
||||
}
|
||||
let reader = fixture
|
||||
.readers
|
||||
.connection(client, &QueryScope::All, "fixture-secret")
|
||||
.await?;
|
||||
let access = ReadAccessParams {
|
||||
all_teams: false,
|
||||
user_id: String::new(),
|
||||
team_ids: vec!["team-a".into()],
|
||||
};
|
||||
let page = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 500).await?;
|
||||
assert_eq!(page.data.len(), runs);
|
||||
assert!(
|
||||
page.data
|
||||
.windows(2)
|
||||
.all(|runs| runs[0].trace_ref > runs[1].trace_ref)
|
||||
);
|
||||
if runs > 1 {
|
||||
client
|
||||
.post(writer.url().clone())
|
||||
.body("SYSTEM FLUSH LOGS")
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
let read_queries = client.post(writer.url().clone()).body(format!(
|
||||
"SELECT count() FROM system.query_log WHERE type = 'QueryFinish' AND current_database = '{DATABASE}' AND query LIKE '%FROM otel_traces AS o%' AND query NOT LIKE '%system.query_log%'"
|
||||
)).send().await?.error_for_status()?.text().await?;
|
||||
let read_queries = read_queries.trim().parse::<usize>()?;
|
||||
assert!(
|
||||
read_queries > 0 && read_queries < runs,
|
||||
"{read_queries} span queries for {runs} runs"
|
||||
);
|
||||
if costed {
|
||||
let overlapping = client
|
||||
.post(writer.url().clone())
|
||||
.body(format!(
|
||||
"WITH spend_reads AS (
|
||||
SELECT query_start_time_microseconds AS started, event_time_microseconds AS finished
|
||||
FROM system.query_log
|
||||
WHERE type = 'QueryFinish' AND current_database = '{DATABASE}'
|
||||
AND query LIKE '%FROM spend_logs FINAL%' AND query NOT LIKE '%system.query_log%'
|
||||
), events AS (
|
||||
SELECT started AS at, 1 AS delta FROM spend_reads
|
||||
UNION ALL SELECT finished AS at, -1 AS delta FROM spend_reads
|
||||
)
|
||||
SELECT max(active) FROM (
|
||||
SELECT sum(delta) OVER (ORDER BY at, delta ROWS UNBOUNDED PRECEDING) AS active
|
||||
FROM events
|
||||
)"
|
||||
))
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?
|
||||
.text()
|
||||
.await?
|
||||
.trim()
|
||||
.parse::<usize>()?;
|
||||
assert!(
|
||||
(2..=4).contains(&overlapping),
|
||||
"{overlapping} simultaneous spend reads for {runs} runs"
|
||||
);
|
||||
}
|
||||
}
|
||||
for summary in &page.data {
|
||||
assert_eq!(summary.span_count, steps as u64);
|
||||
assert_eq!(
|
||||
if costed {
|
||||
summary.llm_calls
|
||||
} else {
|
||||
summary.tool_calls
|
||||
},
|
||||
(steps - 1) as u64
|
||||
);
|
||||
if costed {
|
||||
assert_eq!(summary.spend, Some((steps - 1) as f64 * 0.25));
|
||||
}
|
||||
}
|
||||
let trace_ref = &page
|
||||
.data
|
||||
.iter()
|
||||
.find(|run| run.trace_id == "trace-0000")
|
||||
.ok_or("missing run")?
|
||||
.trace_ref;
|
||||
let detail = get_trace(client, &reader, &access, "trace-0000", trace_ref)
|
||||
.await?
|
||||
.ok_or("missing trace")?;
|
||||
assert_eq!(detail.spans.len(), steps);
|
||||
assert_eq!(detail.spans[0].span_id, "span-0000");
|
||||
assert_eq!(
|
||||
detail.spans[steps - 1].span_id,
|
||||
format!("span-{:04}", steps - 1)
|
||||
);
|
||||
assert_eq!(
|
||||
if costed {
|
||||
detail.summary.llm_calls
|
||||
} else {
|
||||
detail.summary.tool_calls
|
||||
},
|
||||
(steps - 1) as u64
|
||||
);
|
||||
let denied = ReadAccessParams {
|
||||
team_ids: vec!["other-team".into()],
|
||||
..access.clone()
|
||||
};
|
||||
assert!(
|
||||
get_trace(client, &reader, &denied, "trace-0000", trace_ref)
|
||||
.await?
|
||||
.is_none()
|
||||
);
|
||||
let mut cursor = None;
|
||||
let mut ids = Vec::new();
|
||||
loop {
|
||||
let page = get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
"trace-0000",
|
||||
trace_ref,
|
||||
cursor.as_deref(),
|
||||
200,
|
||||
)
|
||||
.await?
|
||||
.ok_or("missing page")?;
|
||||
assert_eq!(page.summary, detail.summary);
|
||||
assert!(page.spans.len() <= 200);
|
||||
assert!(
|
||||
serde_json::to_vec(&page)?.len()
|
||||
<= litellm_storage_clickhouse::READ_LIMITS.response_bytes
|
||||
);
|
||||
if ids.is_empty() {
|
||||
assert!(
|
||||
get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&denied,
|
||||
"trace-0000",
|
||||
trace_ref,
|
||||
page.next_cursor.as_deref(),
|
||||
200,
|
||||
)
|
||||
.await?
|
||||
.is_none()
|
||||
);
|
||||
client
|
||||
.post(writer.url().clone())
|
||||
.body(format!("TRUNCATE TABLE {DATABASE}.otel_traces"))
|
||||
.send()
|
||||
.await?
|
||||
.error_for_status()?;
|
||||
}
|
||||
ids.extend(page.spans.into_iter().map(|span| span.span_id));
|
||||
cursor = page.next_cursor;
|
||||
if cursor.is_none() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
assert_eq!(
|
||||
ids,
|
||||
detail
|
||||
.spans
|
||||
.iter()
|
||||
.map(|span| span.span_id.clone())
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[tokio::test]
|
||||
async fn cursor_pages_keep_a_tenant_scoped_snapshot_when_more_spans_arrive(
|
||||
#[future(awt)] seeded_database: TestResult<SeededDatabase>,
|
||||
) -> TestResult {
|
||||
let fixture = seeded_database?;
|
||||
let client = &fixture.database.client;
|
||||
let reader = fixture
|
||||
.readers
|
||||
.connection(client, &QueryScope::All, "fixture-secret")
|
||||
.await?;
|
||||
let access = ReadAccessParams {
|
||||
all_teams: true,
|
||||
user_id: String::new(),
|
||||
team_ids: Vec::new(),
|
||||
};
|
||||
let listed = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 10).await?;
|
||||
let summary = listed
|
||||
.data
|
||||
.iter()
|
||||
.find(|summary| summary.span_count == 3)
|
||||
.ok_or("missing fixture")?;
|
||||
let first = get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
None,
|
||||
1,
|
||||
)
|
||||
.await?
|
||||
.ok_or("missing first page")?;
|
||||
let original_ids = get_trace(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
)
|
||||
.await?
|
||||
.ok_or("missing trace")?
|
||||
.spans
|
||||
.into_iter()
|
||||
.map(|span| span.span_id)
|
||||
.collect::<Vec<_>>();
|
||||
let writer = Connection::writer(&fixture.database.url)?;
|
||||
insert_rows(
|
||||
client,
|
||||
&writer,
|
||||
DATABASE,
|
||||
InsertTable::OtelTraces,
|
||||
vec![BTreeMap::from([
|
||||
("Timestamp".into(), json!(1_790_000_000_000_000_000_i64)),
|
||||
("TraceId".into(), json!(summary.trace_id)),
|
||||
("SpanId".into(), json!("late-span")),
|
||||
("ParentSpanId".into(), json!(first.spans[0].span_id)),
|
||||
("TeamId".into(), json!("team-a")),
|
||||
("ApiKeyHash".into(), json!("key-a")),
|
||||
("EngineReceivedMs".into(), json!(u64::MAX / 2)),
|
||||
])],
|
||||
)
|
||||
.await?;
|
||||
let denied = ReadAccessParams {
|
||||
all_teams: false,
|
||||
user_id: String::new(),
|
||||
team_ids: vec!["not-this-team".into()],
|
||||
};
|
||||
assert!(
|
||||
get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&denied,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
first.next_cursor.as_deref(),
|
||||
1
|
||||
)
|
||||
.await?
|
||||
.is_none()
|
||||
);
|
||||
let first_cursor = first.next_cursor.clone();
|
||||
let mut cursor = first.next_cursor;
|
||||
let mut ids = first
|
||||
.spans
|
||||
.into_iter()
|
||||
.map(|span| span.span_id)
|
||||
.collect::<Vec<_>>();
|
||||
while let Some(current) = cursor {
|
||||
let next = get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&summary.trace_id,
|
||||
&summary.trace_ref,
|
||||
Some(¤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<SeededDatabase>,
|
||||
) -> TestResult {
|
||||
let fixture = seeded_database?;
|
||||
let client = &fixture.database.client;
|
||||
let reader = fixture
|
||||
.readers
|
||||
.connection(client, &QueryScope::All, "fixture-secret")
|
||||
.await?;
|
||||
let access = ReadAccessParams {
|
||||
all_teams: true,
|
||||
user_id: String::new(),
|
||||
team_ids: Vec::new(),
|
||||
};
|
||||
let before = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
|
||||
let run = before
|
||||
.data
|
||||
.iter()
|
||||
.find(|run| run.span_count == 3)
|
||||
.ok_or("missing fixture")?;
|
||||
let writer = Connection::writer(&fixture.database.url)?;
|
||||
insert_rows(
|
||||
client,
|
||||
&writer,
|
||||
DATABASE,
|
||||
InsertTable::OtelTraces,
|
||||
vec![BTreeMap::from([
|
||||
("Timestamp".into(), json!(1_790_000_000_000_000_000_i64)),
|
||||
("TraceId".into(), json!(run.trace_id)),
|
||||
("SpanId".into(), json!("oversized-child")),
|
||||
("ParentSpanId".into(), json!("0101010101010101")),
|
||||
(
|
||||
"SpanName".into(),
|
||||
json!("x".repeat(litellm_storage_clickhouse::READ_LIMITS.response_bytes + 1)),
|
||||
),
|
||||
("ObservationType".into(), json!("tool")),
|
||||
("TeamId".into(), json!("team-a")),
|
||||
("ApiKeyHash".into(), json!("key-a")),
|
||||
])],
|
||||
)
|
||||
.await?;
|
||||
let after = list_traces(client, &reader, &access, 0, 2_000_000_000_000, None, 50).await?;
|
||||
assert_eq!(after.data.len(), before.data.len());
|
||||
let limited = after
|
||||
.data
|
||||
.iter()
|
||||
.find(|item| item.trace_ref == run.trace_ref)
|
||||
.ok_or("missing run")?;
|
||||
assert!(limited.resolution_limited);
|
||||
assert_eq!(limited.span_count, 4);
|
||||
assert!(
|
||||
after
|
||||
.data
|
||||
.iter()
|
||||
.filter(|item| item.trace_ref != run.trace_ref)
|
||||
.all(|item| !item.resolution_limited)
|
||||
);
|
||||
assert!(matches!(
|
||||
get_trace_page(
|
||||
client,
|
||||
&reader,
|
||||
&access,
|
||||
&run.trace_id,
|
||||
&run.trace_ref,
|
||||
None,
|
||||
200
|
||||
)
|
||||
.await,
|
||||
Err(litellm_traces_clickhouse::Error::ReadTooLarge)
|
||||
));
|
||||
Ok(())
|
||||
}
|
||||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ pub enum SpanStatus {
|
|||
}
|
||||
|
||||
#[macro_rules_attribute::apply(response_type)]
|
||||
#[derive(Debug, PartialEq)]
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct Span {
|
||||
pub span_id: String,
|
||||
pub parent_span_id: Option<String>,
|
||||
|
|
@ -41,7 +41,7 @@ pub struct Span {
|
|||
|
||||
/// One distinct agent in a trace: 200 invocations of `researcher` are one node.
|
||||
#[macro_rules_attribute::apply(response_type)]
|
||||
#[derive(Debug, PartialEq)]
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct AgentNode {
|
||||
pub name: String,
|
||||
pub parent_agent: Option<String>,
|
||||
|
|
@ -53,8 +53,10 @@ pub struct AgentNode {
|
|||
}
|
||||
|
||||
#[macro_rules_attribute::apply(response_type)]
|
||||
#[derive(Debug, PartialEq)]
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct TraceSummary {
|
||||
#[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))]
|
||||
pub resolution_limited: bool,
|
||||
pub trace_id: String,
|
||||
#[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))]
|
||||
pub trace_ref: String,
|
||||
|
|
@ -81,11 +83,13 @@ pub struct TraceSummary {
|
|||
}
|
||||
|
||||
#[macro_rules_attribute::apply(response_type)]
|
||||
#[derive(Debug, PartialEq)]
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct Trace {
|
||||
pub summary: TraceSummary,
|
||||
pub agents: Vec<AgentNode>,
|
||||
pub spans: Vec<Span>,
|
||||
#[cfg_attr(feature = "schema", schemars(extend("x-python-optional" = true)))]
|
||||
pub next_cursor: Option<String>,
|
||||
}
|
||||
|
||||
#[macro_rules_attribute::apply(response_type)]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
3
litellm/decisions/__init__.py
Normal file
3
litellm/decisions/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.decisions.main import adecisions, decisions
|
||||
|
||||
__all__ = ["adecisions", "decisions"]
|
||||
299
litellm/decisions/main.py
Normal file
299
litellm/decisions/main.py
Normal file
|
|
@ -0,0 +1,299 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass, field
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.decisions.transformation import DecisionsProviderConfig
|
||||
from litellm.llms.cloudflare.decisions.transformation import CLOUDFLARE_DECISIONS_ENDPOINT
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client, get_async_httpx_client
|
||||
from litellm.llms.openrouter.decisions.transformation import OPENROUTER_DECISIONS_ENDPOINT
|
||||
from litellm.llms.perplexity.decisions.transformation import PERPLEXITY_DECISIONS_ENDPOINT
|
||||
from litellm.llms.strands_decider.decisions.transformation import STRANDS_DECIDER_DECISIONS_ENDPOINT
|
||||
from litellm.llms.typesafe.decisions.transformation import TYPESAFE_DECISIONS_ENDPOINT
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.decisions import (
|
||||
DecisionQuestion,
|
||||
DecisionsJSON,
|
||||
DecisionsRequest,
|
||||
DecisionsResponse,
|
||||
)
|
||||
from litellm.utils import client
|
||||
|
||||
DECISIONS_ENDPOINTS: Final[Mapping[str, DecisionsProviderConfig]] = MappingProxyType(
|
||||
{
|
||||
"perplexity": PERPLEXITY_DECISIONS_ENDPOINT,
|
||||
"typesafe": TYPESAFE_DECISIONS_ENDPOINT,
|
||||
"openrouter": OPENROUTER_DECISIONS_ENDPOINT,
|
||||
"cloudflare": CLOUDFLARE_DECISIONS_ENDPOINT,
|
||||
"strands_decider": STRANDS_DECIDER_DECISIONS_ENDPOINT,
|
||||
}
|
||||
)
|
||||
|
||||
_DECISIONS_REQUEST_ADAPTER: Final[TypeAdapter[DecisionsRequest]] = TypeAdapter(DecisionsRequest)
|
||||
_DECISIONS_PAYLOAD_ADAPTER: Final[TypeAdapter[object]] = TypeAdapter(object)
|
||||
_DECISIONS_RESPONSE_ADAPTER: Final[TypeAdapter[DecisionsResponse]] = TypeAdapter(DecisionsResponse)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, repr=False)
|
||||
class _PreparedDecisionsRequest:
|
||||
config: DecisionsProviderConfig
|
||||
provider: str
|
||||
upstream_model: str
|
||||
url: str
|
||||
api_key: str | None = field(repr=False)
|
||||
headers: Mapping[str, str] = field(repr=False)
|
||||
body: Mapping[str, object] = field(repr=False)
|
||||
|
||||
|
||||
def _resolve_provider_model(model: str, custom_llm_provider: str | None) -> tuple[str, str]:
|
||||
provider: Final = model.partition("/")[0] if custom_llm_provider is None else custom_llm_provider
|
||||
if provider not in DECISIONS_ENDPOINTS:
|
||||
supported: Final = ", ".join(DECISIONS_ENDPOINTS)
|
||||
raise litellm.BadRequestError(
|
||||
message=f"Unknown Decisions provider '{provider}'. Supported providers: {supported}",
|
||||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
upstream_model: Final = model.removeprefix(f"{provider}/") if model.startswith(f"{provider}/") else model
|
||||
if not upstream_model:
|
||||
raise litellm.BadRequestError(
|
||||
message="A model name is required for the Decisions API",
|
||||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
return provider, upstream_model
|
||||
|
||||
|
||||
def _resolve_api_key(
|
||||
*,
|
||||
provider: str,
|
||||
model: str,
|
||||
endpoint: DecisionsProviderConfig,
|
||||
api_key: str | None,
|
||||
) -> str | None:
|
||||
if api_key is not None:
|
||||
return api_key
|
||||
|
||||
server_api_key: Final = next(
|
||||
(key for key in (get_secret_str(name) for name in endpoint.api_key_env) if key),
|
||||
None,
|
||||
)
|
||||
if server_api_key is None:
|
||||
if not endpoint.api_key_required:
|
||||
return None
|
||||
raise litellm.AuthenticationError(
|
||||
message=f"Missing API key for Decisions provider '{provider}'",
|
||||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
|
||||
return server_api_key
|
||||
|
||||
|
||||
def _prepare_request(
|
||||
*,
|
||||
model: str,
|
||||
state: DecisionsJSON,
|
||||
questions: Mapping[str, DecisionQuestion | Mapping[str, object]],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
) -> _PreparedDecisionsRequest:
|
||||
provider, upstream_model = _resolve_provider_model(model, custom_llm_provider)
|
||||
try:
|
||||
validated_request: Final = _DECISIONS_REQUEST_ADAPTER.validate_python(
|
||||
{"model": model, "state": state, "questions": questions}
|
||||
)
|
||||
except ValidationError as error:
|
||||
raise litellm.BadRequestError(
|
||||
message=f"Invalid Decisions request: {error}",
|
||||
model=model,
|
||||
llm_provider=provider,
|
||||
) from error
|
||||
|
||||
endpoint: Final = DECISIONS_ENDPOINTS[provider]
|
||||
env_api_base: Final = get_secret_str(endpoint.api_base_env)
|
||||
default_api_base: Final = endpoint.default_api_base()
|
||||
resolved_api_base: Final = api_base or env_api_base or default_api_base
|
||||
if resolved_api_base is None:
|
||||
raise litellm.BadRequestError(
|
||||
message=endpoint.missing_api_base_message(provider),
|
||||
model=model,
|
||||
llm_provider=provider,
|
||||
)
|
||||
|
||||
resolved_api_key: Final = _resolve_api_key(
|
||||
provider=provider,
|
||||
model=model,
|
||||
endpoint=endpoint,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
canonical_model: Final = endpoint.canonical_model(upstream_model)
|
||||
outbound_headers: Final = MappingProxyType(
|
||||
{
|
||||
**{
|
||||
name: value
|
||||
for name, value in (extra_headers or {}).items()
|
||||
if name.lower() not in {"authorization", "content-type"}
|
||||
},
|
||||
**({"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key is not None else {}),
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
)
|
||||
body: Final = MappingProxyType(
|
||||
{
|
||||
"model": endpoint.request_model(canonical_model),
|
||||
"state": validated_request.state,
|
||||
"questions": {
|
||||
name: question.model_dump(mode="json", exclude_none=True)
|
||||
for name, question in validated_request.questions.items()
|
||||
},
|
||||
}
|
||||
)
|
||||
return _PreparedDecisionsRequest(
|
||||
config=endpoint,
|
||||
provider=provider,
|
||||
upstream_model=canonical_model,
|
||||
url=endpoint.endpoint_url(resolved_api_base, canonical_model),
|
||||
api_key=resolved_api_key,
|
||||
headers=outbound_headers,
|
||||
body=body,
|
||||
)
|
||||
|
||||
|
||||
def _log_request(
|
||||
prepared: _PreparedDecisionsRequest,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> LiteLLMLoggingObj | None:
|
||||
logging_obj: Final = kwargs.get("litellm_logging_obj")
|
||||
if not isinstance(logging_obj, LiteLLMLoggingObj):
|
||||
return None
|
||||
logging_obj.update_from_kwargs(
|
||||
kwargs=dict(kwargs),
|
||||
model=prepared.upstream_model,
|
||||
litellm_params={
|
||||
"litellm_call_id": kwargs.get("litellm_call_id"),
|
||||
"api_base": prepared.url,
|
||||
},
|
||||
custom_llm_provider=prepared.provider,
|
||||
)
|
||||
request_body: Final = dict(prepared.body)
|
||||
request_headers: Final = dict(prepared.headers)
|
||||
logging_obj.pre_call(
|
||||
input=request_body,
|
||||
api_key=prepared.api_key,
|
||||
model=prepared.upstream_model,
|
||||
additional_args={
|
||||
"api_base": prepared.url,
|
||||
"complete_input_dict": request_body,
|
||||
"headers": request_headers,
|
||||
},
|
||||
)
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _parse_response(
|
||||
response: httpx.Response,
|
||||
prepared: _PreparedDecisionsRequest,
|
||||
) -> DecisionsResponse:
|
||||
response.raise_for_status()
|
||||
payload: Final[object] = _DECISIONS_PAYLOAD_ADAPTER.validate_json(response.content)
|
||||
result: Final = _DECISIONS_RESPONSE_ADAPTER.validate_python(prepared.config.unwrap_response(payload))
|
||||
result._hidden_params.update(
|
||||
{
|
||||
"model": f"{prepared.provider}/{prepared.upstream_model}",
|
||||
"custom_llm_provider": prepared.provider,
|
||||
"provider_response_model": f"{prepared.provider}/{prepared.upstream_model}",
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _map_upstream_exception(error: Exception, prepared: _PreparedDecisionsRequest) -> Exception:
|
||||
return litellm.exception_type(
|
||||
model=f"{prepared.provider}/{prepared.upstream_model}",
|
||||
custom_llm_provider=prepared.provider,
|
||||
original_exception=error,
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
async def adecisions(
|
||||
model: str,
|
||||
state: DecisionsJSON,
|
||||
questions: Mapping[str, DecisionQuestion | Mapping[str, object]],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
**kwargs: object,
|
||||
) -> DecisionsResponse:
|
||||
prepared: Final = _prepare_request(
|
||||
model=model,
|
||||
state=state,
|
||||
questions=questions,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
logging_obj: Final = _log_request(prepared, kwargs)
|
||||
try:
|
||||
handler: Final = get_async_httpx_client(llm_provider=prepared.provider)
|
||||
response: Final = await handler.post(
|
||||
prepared.url,
|
||||
json=dict(prepared.body),
|
||||
headers=dict(prepared.headers),
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return _parse_response(response=response, prepared=prepared)
|
||||
except Exception as error:
|
||||
raise _map_upstream_exception(error, prepared) from error
|
||||
|
||||
|
||||
@client
|
||||
def decisions(
|
||||
model: str,
|
||||
state: DecisionsJSON,
|
||||
questions: Mapping[str, DecisionQuestion | Mapping[str, object]],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
timeout: float | httpx.Timeout | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
extra_headers: Mapping[str, str] | None = None,
|
||||
**kwargs: object,
|
||||
) -> DecisionsResponse:
|
||||
prepared: Final = _prepare_request(
|
||||
model=model,
|
||||
state=state,
|
||||
questions=questions,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
)
|
||||
logging_obj: Final = _log_request(prepared, kwargs)
|
||||
try:
|
||||
handler: Final = _get_httpx_client()
|
||||
response: Final = handler.post(
|
||||
prepared.url,
|
||||
json=dict(prepared.body),
|
||||
headers=dict(prepared.headers),
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return _parse_response(response=response, prepared=prepared)
|
||||
except Exception as error:
|
||||
raise _map_upstream_exception(error, prepared) from error
|
||||
|
||||
|
||||
__all__ = ["DECISIONS_ENDPOINTS", "adecisions", "decisions"]
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
||||
|
||||
|
|
|
|||
345
litellm/harness/handlers/tool_loop_handler.py
Normal file
345
litellm/harness/handlers/tool_loop_handler.py
Normal file
|
|
@ -0,0 +1,345 @@
|
|||
"""In-process tool-calling loop handler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import inspect
|
||||
import json
|
||||
import math
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol, TypeAlias
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.errors import CapabilityUnsupported
|
||||
from litellm.harness.handlers.base import BaseHarnessHandler
|
||||
from litellm.harness.types import Approval, Event, Reasoning, Text, ToolCall, ToolResult
|
||||
from litellm.llms.base_llm.harness.transformation import HarnessTurnError
|
||||
from litellm.llms.tool_loop.harness.transformation import (
|
||||
TOOL_LOOP_MAX_MODEL_CALLS,
|
||||
FunctionTool,
|
||||
ToolLoopHarnessConfig,
|
||||
completion_kwargs,
|
||||
function_tool,
|
||||
)
|
||||
from litellm.types.completion import ChatCompletionMessageParam
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionMessageCustomToolCall,
|
||||
ChatCompletionMessageToolCall,
|
||||
ChatCompletionToolParam,
|
||||
ModelResponse,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _FunctionToolCall:
|
||||
id: str
|
||||
name: str
|
||||
arguments: str
|
||||
|
||||
|
||||
AsyncCompletion: TypeAlias = Callable[..., Awaitable[ModelResponse]]
|
||||
_ARGUMENTS_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_MAPPING_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_MODEL_RESPONSE_ADAPTER: Final = TypeAdapter(ModelResponse)
|
||||
_RAW_DECODE_ADAPTER: Final = TypeAdapter(tuple[object, int])
|
||||
_JSON_DECODER: Final = json.JSONDecoder()
|
||||
_HISTORY_ADAPTER: Final = TypeAdapter(list[dict[str, object]])
|
||||
|
||||
|
||||
class _Usage(BaseModel):
|
||||
model_config = ConfigDict(from_attributes=True)
|
||||
|
||||
prompt_tokens: int | None = None
|
||||
completion_tokens: int | None = None
|
||||
|
||||
|
||||
_USAGE_ADAPTER: Final = TypeAdapter(_Usage)
|
||||
|
||||
|
||||
class _AwaitableObject(Protocol):
|
||||
def __await__(self) -> Generator[object, None, object]: ...
|
||||
|
||||
|
||||
async def _await_tool_result(result: _AwaitableObject) -> object:
|
||||
return await result
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ToolOutcome:
|
||||
output: str
|
||||
is_error: bool
|
||||
|
||||
|
||||
def _normalize_tool_call(
|
||||
call: ChatCompletionMessageToolCall | ChatCompletionMessageCustomToolCall,
|
||||
) -> _FunctionToolCall:
|
||||
if isinstance(call, ChatCompletionMessageToolCall):
|
||||
return _FunctionToolCall(
|
||||
id=call.id,
|
||||
name=call.function.name or "",
|
||||
arguments=call.function.arguments,
|
||||
)
|
||||
return _FunctionToolCall(id=call.id, name=call.custom.name, arguments=call.custom.input)
|
||||
|
||||
|
||||
def _parse_arguments(raw: str) -> tuple[dict[str, object], str | None]: # mutable-ok: event input uses a dict
|
||||
text: Final = raw.strip()
|
||||
try:
|
||||
raw_decoded: object = _JSON_DECODER.raw_decode(text)
|
||||
except json.JSONDecodeError as error:
|
||||
return {}, f"{type(error).__name__}: {error}"
|
||||
parsed, end = _RAW_DECODE_ADAPTER.validate_python(raw_decoded)
|
||||
if text[end:].strip():
|
||||
return {}, "JSONDecodeError: Extra data after tool arguments"
|
||||
try:
|
||||
arguments: Final = _ARGUMENTS_ADAPTER.validate_python(parsed)
|
||||
except ValidationError as error:
|
||||
if not isinstance(parsed, dict):
|
||||
return {}, "ValueError: tool arguments must be a JSON object"
|
||||
return {}, f"{type(error).__name__}: {error}"
|
||||
return arguments, None
|
||||
|
||||
|
||||
def _cost_value(value: object) -> float | None:
|
||||
if isinstance(value, bool):
|
||||
return None
|
||||
if not isinstance(value, str | int | float):
|
||||
return None
|
||||
try:
|
||||
cost: Final = float(value)
|
||||
except (OverflowError, TypeError, ValueError):
|
||||
return None
|
||||
return cost if math.isfinite(cost) else None
|
||||
|
||||
|
||||
def _as_mapping(value: object) -> Mapping[str, object] | None:
|
||||
try:
|
||||
return _MAPPING_ADAPTER.validate_python(value)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _response_cost(response: ModelResponse) -> float:
|
||||
hidden_params_value: Final[object] = getattr(response, "_hidden_params", {})
|
||||
hidden_params: Final = _as_mapping(hidden_params_value)
|
||||
if hidden_params is not None:
|
||||
additional_headers: Final = _as_mapping(hidden_params.get("additional_headers"))
|
||||
if additional_headers is not None:
|
||||
header_cost: Final = _cost_value(additional_headers.get("llm_provider-x-litellm-response-cost"))
|
||||
if header_cost is not None:
|
||||
return header_cost
|
||||
hidden_cost: Final = _cost_value(hidden_params.get("response_cost"))
|
||||
if hidden_cost is not None:
|
||||
return hidden_cost
|
||||
try:
|
||||
calculated_cost: Final = litellm.completion_cost(completion_response=response)
|
||||
except Exception:
|
||||
return 0.0
|
||||
return _cost_value(calculated_cost) or 0.0
|
||||
|
||||
|
||||
async def _approval_error(approval: Approval | None) -> str | None:
|
||||
if approval is None:
|
||||
return None
|
||||
allowed, reason = await approval.wait()
|
||||
return None if allowed else f"denied: {reason}"
|
||||
|
||||
|
||||
async def _tool_outcome(
|
||||
tool: FunctionTool | None,
|
||||
tool_name: str,
|
||||
arguments: dict[str, object],
|
||||
parse_error: str | None,
|
||||
approval_error: str | None,
|
||||
) -> _ToolOutcome:
|
||||
if parse_error is not None:
|
||||
return _ToolOutcome(output=parse_error, is_error=True)
|
||||
if approval_error is not None:
|
||||
return _ToolOutcome(output=approval_error, is_error=True)
|
||||
if tool is None:
|
||||
return _ToolOutcome(output=f"ValueError: unknown tool {tool_name!r}", is_error=True)
|
||||
try:
|
||||
result: Final = await _execute_tool(tool, arguments)
|
||||
output: Final = result if isinstance(result, str) else json.dumps(result, default=str)
|
||||
return _ToolOutcome(output=output, is_error=False)
|
||||
except Exception as error:
|
||||
return _ToolOutcome(output=f"{type(error).__name__}: {error}", is_error=True)
|
||||
|
||||
|
||||
def _record_usage(ctx: SessionContext, response: ModelResponse) -> None:
|
||||
try:
|
||||
usage_value: Final[object] = getattr(response, "usage", None)
|
||||
usage: Final = _USAGE_ADAPTER.validate_python(usage_value)
|
||||
input_tokens: Final = usage.prompt_tokens or 0
|
||||
output_tokens: Final = usage.completion_tokens or 0
|
||||
ctx.calls += 1 # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.input_tokens += input_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.output_tokens += output_tokens # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
ctx.cost += _response_cost(response) # rebind-ok: SessionContext is the runtime's per-session usage sink
|
||||
except Exception:
|
||||
return
|
||||
|
||||
|
||||
async def _execute_tool(tool: FunctionTool, arguments: Mapping[str, object]) -> object:
|
||||
validated_model: Final = tool.args_model.model_validate(arguments)
|
||||
values_object: Final[object] = validated_model.model_dump()
|
||||
validated: Final = _MAPPING_ADAPTER.validate_python(values_object)
|
||||
parameters: Final = tuple(inspect.signature(tool.fn).parameters.values())
|
||||
positional_args: Final = tuple(
|
||||
validated[parameter.name] for parameter in parameters if parameter.kind is inspect.Parameter.POSITIONAL_ONLY
|
||||
)
|
||||
keyword_args: Final[dict[str, object]] = { # mutable-ok: tool calls need keyword arguments
|
||||
parameter.name: validated[parameter.name]
|
||||
for parameter in parameters
|
||||
if parameter.kind is not inspect.Parameter.POSITIONAL_ONLY
|
||||
}
|
||||
if inspect.iscoroutinefunction(tool.fn):
|
||||
async_result: Final[object] = tool.fn(*positional_args, **keyword_args)
|
||||
if inspect.isawaitable(async_result):
|
||||
return await _await_tool_result(async_result)
|
||||
return async_result
|
||||
sync_result: Final[object] = await asyncio.to_thread(tool.fn, *positional_args, **keyword_args)
|
||||
if inspect.isawaitable(sync_result):
|
||||
return await _await_tool_result(sync_result)
|
||||
return sync_result
|
||||
|
||||
|
||||
async def _default_acompletion(**kwargs: object) -> ModelResponse: # kwargs-ok: provider-specific completion options
|
||||
response: Final[object] = await litellm.acompletion(**kwargs)
|
||||
return _MODEL_RESPONSE_ADAPTER.validate_python(response)
|
||||
|
||||
|
||||
class ToolLoopHandler(BaseHarnessHandler):
|
||||
def __init__(
|
||||
self,
|
||||
config: ToolLoopHarnessConfig,
|
||||
acompletion: AsyncCompletion | None = None,
|
||||
) -> None:
|
||||
super().__init__(config) # pyright: ignore[reportUnknownMemberType] # base handler config is unparameterized
|
||||
self._config = config
|
||||
self._acompletion = acompletion if acompletion is not None else _default_acompletion
|
||||
self._messages: tuple[ChatCompletionMessageParam, ...] = ()
|
||||
self._tools: Mapping[str, FunctionTool] = MappingProxyType({})
|
||||
self._tool_specs: tuple[ChatCompletionToolParam, ...] = ()
|
||||
self._completion_kwargs: Mapping[str, object] = MappingProxyType({})
|
||||
|
||||
async def start(self, ctx: SessionContext) -> None:
|
||||
self._config.validate_environment(ctx)
|
||||
tools: Final = tuple(function_tool(fn) for fn in ctx.tools)
|
||||
if len({tool.name for tool in tools}) != len(tools):
|
||||
raise ValueError("Harness.TOOL_LOOP tool names must be unique")
|
||||
self._tools = MappingProxyType({tool.name: tool for tool in tools})
|
||||
self._tool_specs = tuple(tool.spec for tool in tools)
|
||||
self._completion_kwargs = MappingProxyType(completion_kwargs(ctx))
|
||||
if not self._messages and ctx.instructions:
|
||||
self._messages = ({"role": "system", "content": ctx.instructions},)
|
||||
|
||||
def native_session_id(self) -> str | None:
|
||||
return None
|
||||
|
||||
async def resume(self, ctx: SessionContext, native_session_id: str) -> None:
|
||||
raise CapabilityUnsupported("Harness.TOOL_LOOP does not support resume")
|
||||
|
||||
async def history(
|
||||
self, ctx: SessionContext
|
||||
) -> list[dict[str, object]]: # mutable-ok: public API returns copied message dictionaries
|
||||
history_object: Final[object] = copy.deepcopy(list(self._messages))
|
||||
return _HISTORY_ADAPTER.validate_python(history_object) # pyright: ignore[reportIncompatibleMethodOverride] # base history uses Any
|
||||
|
||||
async def stop(self, ctx: SessionContext) -> None:
|
||||
return None
|
||||
|
||||
async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]:
|
||||
ctx.final_text = "" # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.output_json = None # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
user_message: Final[ChatCompletionMessageParam] = {"role": "user", "content": prompt}
|
||||
self._messages = (*self._messages, user_message)
|
||||
for _ in range(TOOL_LOOP_MAX_MODEL_CALLS):
|
||||
messages: list[ChatCompletionMessageParam] = copy.deepcopy( # mutable-ok: acompletion takes list messages
|
||||
list(self._messages)
|
||||
)
|
||||
tool_specs: list[ChatCompletionToolParam] = copy.deepcopy( # mutable-ok: acompletion takes tool list
|
||||
list(self._tool_specs)
|
||||
)
|
||||
request_kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
key: value for key, value in self._completion_kwargs.items() if key not in {"messages", "tools"}
|
||||
}
|
||||
kwargs: dict[str, object] = { # mutable-ok: acompletion takes keyword arguments
|
||||
**request_kwargs,
|
||||
"messages": messages,
|
||||
**({"tools": tool_specs} if tool_specs else {}),
|
||||
}
|
||||
response = await self._acompletion(**kwargs)
|
||||
_record_usage(ctx, response)
|
||||
message = response.choices[0].message
|
||||
reasoning_value: object = getattr(message, "reasoning_content", None)
|
||||
reasoning = reasoning_value if isinstance(reasoning_value, str) else None
|
||||
content = message.content
|
||||
if reasoning:
|
||||
yield Reasoning(reasoning)
|
||||
if content:
|
||||
yield Text(content)
|
||||
tool_calls = message.tool_calls or ()
|
||||
if not tool_calls:
|
||||
final_text = content or ""
|
||||
ctx.final_text = final_text # rebind-ok: SessionContext is the runtime's per-turn result sink
|
||||
ctx.output_json = content if ctx.output is not None else None # rebind-ok: per-turn output sink
|
||||
final_message: ChatCompletionMessageParam = {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
}
|
||||
self._messages = (*self._messages, final_message)
|
||||
return
|
||||
normalized_calls = tuple(_normalize_tool_call(call) for call in tool_calls)
|
||||
assistant_message: ChatCompletionMessageParam = {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": call.id,
|
||||
"type": "function",
|
||||
"function": {"name": call.name, "arguments": call.arguments},
|
||||
}
|
||||
for call in normalized_calls
|
||||
],
|
||||
}
|
||||
self._messages = (*self._messages, assistant_message)
|
||||
for call in normalized_calls:
|
||||
arguments, parse_error = _parse_arguments(call.arguments)
|
||||
yield ToolCall(
|
||||
id=call.id,
|
||||
name=call.name,
|
||||
native_name=call.name,
|
||||
input=arguments,
|
||||
builtin=False,
|
||||
)
|
||||
approval = (
|
||||
Approval(tool=call.name, input=arguments)
|
||||
if parse_error is None and ctx.permissions == "ask"
|
||||
else None
|
||||
)
|
||||
if approval is not None:
|
||||
yield approval
|
||||
approval_error = await _approval_error(approval)
|
||||
tool = self._tools.get(call.name)
|
||||
outcome = await _tool_outcome(
|
||||
tool,
|
||||
call.name,
|
||||
arguments,
|
||||
parse_error,
|
||||
approval_error,
|
||||
)
|
||||
yield ToolResult(id=call.id, output=outcome.output, is_error=outcome.is_error)
|
||||
tool_message: ChatCompletionMessageParam = {
|
||||
"role": "tool",
|
||||
"tool_call_id": call.id,
|
||||
"content": outcome.output,
|
||||
}
|
||||
self._messages = (*self._messages, tool_message)
|
||||
raise HarnessTurnError(f"Harness.TOOL_LOOP exceeded {TOOL_LOOP_MAX_MODEL_CALLS} model calls in one turn")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -27,6 +27,7 @@ class Harness(Enum):
|
|||
CODEX = "codex"
|
||||
OPENCODE = "opencode"
|
||||
DEEPAGENTS = "deepagents"
|
||||
TOOL_LOOP = "tool_loop"
|
||||
|
||||
|
||||
def require_harness(harness: object) -> Harness:
|
||||
|
|
|
|||
|
|
@ -290,6 +290,7 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset(
|
|||
"LiteLLM_Config",
|
||||
"LiteLLM_SpendLogs",
|
||||
"LiteLLM_BudgetWindowSpend",
|
||||
"LiteLLM_BackgroundInteractionSettlement",
|
||||
"LiteLLM_ErrorLogs",
|
||||
"LiteLLM_UserNotifications",
|
||||
"LiteLLM_TeamMembership",
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
)
|
||||
),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
3
litellm/llms/base_llm/decisions/__init__.py
Normal file
3
litellm/llms/base_llm/decisions/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .transformation import DecisionsProviderConfig, JevCompatibleDecisionsEndpoint
|
||||
|
||||
__all__ = ["DecisionsProviderConfig", "JevCompatibleDecisionsEndpoint"]
|
||||
52
litellm/llms/base_llm/decisions/transformation.py
Normal file
52
litellm/llms/base_llm/decisions/transformation.py
Normal file
|
|
@ -0,0 +1,52 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class JevCompatibleDecisionsEndpoint:
|
||||
default_api_base_value: str | None
|
||||
path: str
|
||||
api_key_env: tuple[str, ...]
|
||||
api_base_env: str
|
||||
api_key_required: bool = True
|
||||
|
||||
def default_api_base(self) -> str | None:
|
||||
return self.default_api_base_value
|
||||
|
||||
def missing_api_base_message(self, provider: str) -> str:
|
||||
return f"api_base is required for Decisions provider '{provider}'"
|
||||
|
||||
def canonical_model(self, model: str) -> str:
|
||||
return model
|
||||
|
||||
def request_model(self, model: str) -> str:
|
||||
return model
|
||||
|
||||
def endpoint_url(self, api_base: str, model: str) -> str:
|
||||
return f"{api_base.rstrip('/').removesuffix('/v1')}{self.path}"
|
||||
|
||||
def unwrap_response(self, payload: object) -> object:
|
||||
return payload
|
||||
|
||||
|
||||
class DecisionsProviderConfig(Protocol):
|
||||
@property
|
||||
def api_key_env(self) -> tuple[str, ...]: ...
|
||||
|
||||
@property
|
||||
def api_base_env(self) -> str: ...
|
||||
|
||||
@property
|
||||
def api_key_required(self) -> bool: ...
|
||||
|
||||
def default_api_base(self) -> str | None: ...
|
||||
|
||||
def missing_api_base_message(self, provider: str) -> str: ...
|
||||
|
||||
def canonical_model(self, model: str) -> str: ...
|
||||
|
||||
def request_model(self, model: str) -> str: ...
|
||||
|
||||
def endpoint_url(self, api_base: str, model: str) -> str: ...
|
||||
|
||||
def unwrap_response(self, payload: object) -> object: ...
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1125,23 +1125,42 @@ def bedrock_model_accepts_cache_points(model: str | None) -> bool:
|
|||
``cachePoint`` blocks. Bedrock rejects requests carrying cachePoint blocks for
|
||||
models without prompt caching support ("You invoked an unsupported model or your
|
||||
request did not allow prompt caching"), so a model whose cost-map entry does not declare
|
||||
``supports_prompt_caching`` must not receive them. A model absent from the map
|
||||
(an application inference profile ARN, a model newer than the map) keeps emitting
|
||||
so existing caching setups never silently degrade. ``litellm.utils.supports_prompt_caching``
|
||||
is not reusable here: it returns False for unmapped models, the opposite polarity.
|
||||
``supports_prompt_caching`` must not receive them. An explicit
|
||||
``supports_prompt_cache_breakpoint`` on the entry wins over that flag: a model can price
|
||||
cached tokens through implicit caching yet reject the marker on Converse ("This model
|
||||
doesn't support the cachePoint field", Kimi K3). The router registers a deployment's
|
||||
``model_info`` under ``bedrock/<model>`` as configured, route prefix included, while the
|
||||
Converse transformation sees the model with ``converse/`` or ``converse_like/`` already
|
||||
stripped, so every registration form is read. That flag set there covers an application
|
||||
inference profile ARN or a model newer than the map, while only the map decides whether
|
||||
a model is known: absent a map entry the model keeps emitting so existing caching setups
|
||||
never silently degrade. ``litellm.utils.supports_prompt_caching`` is not reusable here:
|
||||
it returns False for unmapped models, the opposite polarity.
|
||||
"""
|
||||
if model is None:
|
||||
return True
|
||||
if _OPENAI_FAMILY_MODEL_RE.search(model):
|
||||
return False
|
||||
entries: Final = tuple(
|
||||
entry
|
||||
for candidate in (model, get_bedrock_base_model(model))
|
||||
if (entry := litellm.model_cost.get(candidate)) is not None
|
||||
map_keys: Final = (model, get_bedrock_base_model(model))
|
||||
registered_keys: Final = tuple(f"bedrock/{route}{model}" for route in ("", "converse/", "converse_like/"))
|
||||
explicit_marker_support: Final = next(
|
||||
(
|
||||
entry.get("supports_prompt_cache_breakpoint") is True
|
||||
for key in (*registered_keys, *map_keys)
|
||||
if (entry := litellm.model_cost.get(key)) is not None
|
||||
and entry.get("supports_prompt_cache_breakpoint") is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not entries:
|
||||
if explicit_marker_support is not None:
|
||||
return explicit_marker_support
|
||||
if not any(key in litellm.model_cost for key in map_keys):
|
||||
return True
|
||||
return any(entry.get("supports_prompt_caching") is True for entry in entries)
|
||||
return any(
|
||||
entry.get("supports_prompt_caching") is True
|
||||
for key in map_keys
|
||||
if (entry := litellm.model_cost.get(key)) is not None
|
||||
)
|
||||
|
||||
|
||||
def bedrock_supports_tool_search(model: str) -> bool:
|
||||
|
|
|
|||
58
litellm/llms/cloudflare/decisions/transformation.py
Normal file
58
litellm/llms/cloudflare/decisions/transformation.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.secret_managers.main import (
|
||||
get_secret_str,
|
||||
normalize_nonempty_secret_str,
|
||||
)
|
||||
|
||||
_RESPONSE_MAPPING_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CloudflareDecisionsEndpoint:
|
||||
api_key_env: tuple[str, ...] = ("CLOUDFLARE_API_KEY",)
|
||||
api_base_env: str = "CLOUDFLARE_API_BASE"
|
||||
api_key_required: bool = True
|
||||
|
||||
def default_api_base(self) -> str | None:
|
||||
account_id: Final = normalize_nonempty_secret_str(get_secret_str("CLOUDFLARE_ACCOUNT_ID"))
|
||||
if account_id is None:
|
||||
return None
|
||||
return f"https://api.cloudflare.com/client/v4/accounts/{account_id}/ai/run"
|
||||
|
||||
def missing_api_base_message(self, provider: str) -> str:
|
||||
return "Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID or pass api_base explicitly"
|
||||
|
||||
def canonical_model(self, model: str) -> str:
|
||||
if model.startswith("@cf/"):
|
||||
return model
|
||||
return f"@cf/cloudflare/{model}"
|
||||
|
||||
def request_model(self, model: str) -> str:
|
||||
return model.rsplit("/", maxsplit=1)[-1]
|
||||
|
||||
def endpoint_url(self, api_base: str, model: str) -> str:
|
||||
normalized_api_base: Final = api_base.rstrip("/")
|
||||
if normalized_api_base.endswith("/ai/v1"):
|
||||
return f"{normalized_api_base.removesuffix('/ai/v1')}/ai/run/{model}"
|
||||
if normalized_api_base.endswith("/ai/run"):
|
||||
return f"{normalized_api_base}/{model}"
|
||||
return f"{normalized_api_base}/ai/run/{model}"
|
||||
|
||||
def unwrap_response(self, payload: object) -> object:
|
||||
if not isinstance(payload, Mapping):
|
||||
return payload
|
||||
response_mapping: Final = _RESPONSE_MAPPING_ADAPTER.validate_python(payload)
|
||||
if "answers" in response_mapping:
|
||||
return payload
|
||||
result: Final = response_mapping.get("result")
|
||||
if isinstance(result, Mapping):
|
||||
return result
|
||||
return payload
|
||||
|
||||
|
||||
CLOUDFLARE_DECISIONS_ENDPOINT: Final[CloudflareDecisionsEndpoint] = CloudflareDecisionsEndpoint()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
10
litellm/llms/openrouter/decisions/transformation.py
Normal file
10
litellm/llms/openrouter/decisions/transformation.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
|
||||
|
||||
OPENROUTER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
|
||||
default_api_base_value="https://openrouter.ai/api",
|
||||
path="/alpha/decisions",
|
||||
api_key_env=("OPENROUTER_API_KEY",),
|
||||
api_base_env="OPENROUTER_API_BASE",
|
||||
)
|
||||
10
litellm/llms/perplexity/decisions/transformation.py
Normal file
10
litellm/llms/perplexity/decisions/transformation.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
|
||||
|
||||
PERPLEXITY_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
|
||||
default_api_base_value="https://api.perplexity.ai",
|
||||
path="/v1/decisions",
|
||||
api_key_env=("PERPLEXITYAI_API_KEY", "PERPLEXITY_API_KEY"),
|
||||
api_base_env="PERPLEXITY_API_BASE",
|
||||
)
|
||||
11
litellm/llms/strands_decider/decisions/transformation.py
Normal file
11
litellm/llms/strands_decider/decisions/transformation.py
Normal file
|
|
@ -0,0 +1,11 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
|
||||
|
||||
STRANDS_DECIDER_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
|
||||
default_api_base_value=None,
|
||||
path="/v1/systemone",
|
||||
api_key_env=("STRANDS_DECIDER_API_KEY",),
|
||||
api_base_env="STRANDS_DECIDER_API_BASE",
|
||||
api_key_required=False,
|
||||
)
|
||||
1
litellm/llms/tool_loop/__init__.py
Normal file
1
litellm/llms/tool_loop/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""In-process tool loop harness."""
|
||||
1
litellm/llms/tool_loop/harness/__init__.py
Normal file
1
litellm/llms/tool_loop/harness/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""Tool Loop harness configuration."""
|
||||
119
litellm/llms/tool_loop/harness/transformation.py
Normal file
119
litellm/llms/tool_loop/harness/transformation.py
Normal file
|
|
@ -0,0 +1,119 @@
|
|||
"""Configuration and tool-schema helpers for the in-process Tool Loop harness."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, create_model
|
||||
|
||||
from litellm.harness.context import SessionContext
|
||||
from litellm.harness.options import ToolLoopOptions
|
||||
from litellm.harness.types import Capabilities, Harness
|
||||
from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig
|
||||
from litellm.llms.base_llm.harness.utils import gateway_headers
|
||||
from litellm.types.utils import ChatCompletionToolParam
|
||||
|
||||
TOOL_LOOP_MAX_MODEL_CALLS: Final = 100
|
||||
_ANNOTATIONS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
_OBJECT_ADAPTER: Final = TypeAdapter(object)
|
||||
_MODEL_FACTORY: Final[Callable[..., type[BaseModel]]] = create_model
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FunctionTool:
|
||||
name: str
|
||||
fn: Callable[..., object]
|
||||
args_model: type[BaseModel]
|
||||
spec: ChatCompletionToolParam
|
||||
|
||||
|
||||
def _field_definition(
|
||||
parameter: inspect.Parameter,
|
||||
annotations: Mapping[str, object],
|
||||
) -> tuple[object, object]:
|
||||
annotation: Final = annotations.get(parameter.name, object)
|
||||
raw_default: Final[object] = parameter.default # pyright: ignore[reportAny] # inspect exposes defaults as Any
|
||||
if raw_default is inspect.Parameter.empty:
|
||||
return annotation, ...
|
||||
default: Final = _OBJECT_ADAPTER.validate_python(raw_default)
|
||||
return annotation, default
|
||||
|
||||
|
||||
def function_tool(fn: Callable[..., object]) -> FunctionTool:
|
||||
signature: Final = inspect.signature(fn)
|
||||
parameters: Final = tuple(signature.parameters.values())
|
||||
raw_annotations: Final[object] = inspect.get_annotations(fn, eval_str=True)
|
||||
annotations: Final = _ANNOTATIONS_ADAPTER.validate_python(raw_annotations)
|
||||
if any(
|
||||
parameter.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD) for parameter in parameters
|
||||
):
|
||||
raise ValueError(f"Tool {fn.__name__} cannot use variadic parameters")
|
||||
|
||||
fields: Final = MappingProxyType(
|
||||
{parameter.name: _field_definition(parameter, annotations) for parameter in parameters}
|
||||
)
|
||||
args_model: Final[type[BaseModel]] = _MODEL_FACTORY(
|
||||
f"{fn.__name__}_args",
|
||||
__config__=ConfigDict(extra="forbid"),
|
||||
**fields, # pyright: ignore[reportCallIssue, reportArgumentType] # Pydantic creates fields dynamically
|
||||
)
|
||||
spec: Final[ChatCompletionToolParam] = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": fn.__name__,
|
||||
"description": inspect.getdoc(fn) or "",
|
||||
"parameters": args_model.model_json_schema(),
|
||||
},
|
||||
}
|
||||
return FunctionTool(name=fn.__name__, fn=fn, args_model=args_model, spec=spec)
|
||||
|
||||
|
||||
def _routing_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
||||
if not ctx.model:
|
||||
raise ValueError("Harness.TOOL_LOOP needs model=")
|
||||
if ctx.gateway is not None:
|
||||
return {
|
||||
"model": f"litellm_proxy/{ctx.model}",
|
||||
"api_base": ctx.gateway.api_base,
|
||||
"api_key": ctx.gateway.api_key,
|
||||
"extra_headers": gateway_headers(ctx),
|
||||
}
|
||||
return {
|
||||
"model": ctx.model,
|
||||
**({"api_key": ctx.api_key} if ctx.api_key is not None else {}),
|
||||
**({"api_base": ctx.api_base} if ctx.api_base is not None else {}),
|
||||
}
|
||||
|
||||
|
||||
def completion_kwargs(ctx: SessionContext) -> Mapping[str, object]:
|
||||
options: Final = ToolLoopHarnessConfig().get_options(ctx)
|
||||
routing: Final = _routing_kwargs(ctx)
|
||||
kwargs: Final[Mapping[str, object]] = MappingProxyType({**options.completion_kwargs, **routing})
|
||||
if ctx.output is None:
|
||||
return kwargs
|
||||
return {**kwargs, "response_format": ctx.output}
|
||||
|
||||
|
||||
class ToolLoopHarnessConfig(BaseHarnessConfig[ToolLoopOptions]):
|
||||
harness = Harness.TOOL_LOOP
|
||||
options_type = ToolLoopOptions
|
||||
uses_model_endpoint = False
|
||||
capabilities = Capabilities(
|
||||
structured_output=True,
|
||||
tool_approval=True,
|
||||
tool_filtering=False,
|
||||
history=True,
|
||||
custom_tools=True,
|
||||
skills=False,
|
||||
resume=False,
|
||||
permission_modes=frozenset({"ask", "full"}),
|
||||
)
|
||||
|
||||
def validate_environment(self, ctx: SessionContext) -> None:
|
||||
self.get_options(ctx)
|
||||
if not ctx.model:
|
||||
raise ValueError("Harness.TOOL_LOOP needs model=")
|
||||
10
litellm/llms/typesafe/decisions/transformation.py
Normal file
10
litellm/llms/typesafe/decisions/transformation.py
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
from typing import Final
|
||||
|
||||
from litellm.llms.base_llm.decisions.transformation import JevCompatibleDecisionsEndpoint
|
||||
|
||||
TYPESAFE_DECISIONS_ENDPOINT: Final[JevCompatibleDecisionsEndpoint] = JevCompatibleDecisionsEndpoint(
|
||||
default_api_base_value="https://api.typesafe.ai",
|
||||
path="/v1/systemone",
|
||||
api_key_env=("TYPESAFE_API_KEY",),
|
||||
api_base_env="TYPESAFE_API_BASE",
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
1
litellm/proxy/decisions_endpoints/__init__.py
Normal file
1
litellm/proxy/decisions_endpoints/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
__all__ = ()
|
||||
102
litellm/proxy/decisions_endpoints/endpoints.py
Normal file
102
litellm/proxy/decisions_endpoints/endpoints.py
Normal file
|
|
@ -0,0 +1,102 @@
|
|||
from typing import Annotated, Final
|
||||
|
||||
from fastapi import APIRouter, Depends, Request, Response
|
||||
from fastapi.responses import ORJSONResponse # pyright: ignore[reportDeprecated] # required endpoint contract
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm.exceptions import BadRequestError
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.types.decisions import DecisionsRequestBody
|
||||
|
||||
router: Final = APIRouter()
|
||||
_REQUEST_DATA_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
|
||||
_DECISIONS_REQUEST_BODY_ADAPTER: Final[TypeAdapter[DecisionsRequestBody]] = TypeAdapter(DecisionsRequestBody)
|
||||
_GENERAL_SETTINGS_ADAPTER: Final[TypeAdapter[dict[str, object]]] = TypeAdapter(dict[str, object])
|
||||
_OPTIONAL_STRING_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
|
||||
_OPTIONAL_FLOAT_ADAPTER: Final[TypeAdapter[float | None]] = TypeAdapter(float | None)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/decisions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract
|
||||
tags=["decisions"],
|
||||
)
|
||||
@router.post(
|
||||
"/decisions",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_class=ORJSONResponse, # pyright: ignore[reportDeprecated] # required endpoint contract
|
||||
tags=["decisions"],
|
||||
)
|
||||
async def decisions(
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
|
||||
):
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings as proxy_general_settings,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
user_max_tokens,
|
||||
user_request_timeout,
|
||||
version,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
user_api_base as proxy_user_api_base,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
user_model as proxy_user_model,
|
||||
)
|
||||
from litellm.proxy.proxy_server import (
|
||||
user_temperature as proxy_user_temperature,
|
||||
)
|
||||
|
||||
data: Final = _REQUEST_DATA_ADAPTER.validate_json(await request.body())
|
||||
general_settings: Final = _GENERAL_SETTINGS_ADAPTER.validate_python(proxy_general_settings)
|
||||
user_api_base: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_api_base)
|
||||
user_model: Final = _OPTIONAL_STRING_ADAPTER.validate_python(proxy_user_model)
|
||||
user_temperature: Final = _OPTIONAL_FLOAT_ADAPTER.validate_python(proxy_user_temperature)
|
||||
processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
try:
|
||||
_DECISIONS_REQUEST_BODY_ADAPTER.validate_python(data)
|
||||
return await processor.base_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
route_type="adecisions",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=None,
|
||||
model=None,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except ValidationError as error:
|
||||
bad_request_error: Final = BadRequestError(
|
||||
message=f"Invalid Decisions request: {error}",
|
||||
model=str(data.get("model", "")),
|
||||
llm_provider="",
|
||||
)
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=bad_request_error,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
except Exception as error:
|
||||
raise await processor._handle_llm_api_exception(
|
||||
e=error,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
version=version,
|
||||
)
|
||||
|
|
@ -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))),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -99,8 +99,15 @@ class TraceReceiver:
|
|||
async def list_traces(self, scope: TraceScope, start_ms: int, end_ms: int, cursor: str | None = None) -> TracePage:
|
||||
return await self.storage.list_traces(scope, start_ms, end_ms, cursor, AGENT_TRACING_LIST_PAGE_SIZE)
|
||||
|
||||
async def get_trace(self, trace_id: str, scope: TraceScope, trace_ref: str = "") -> Trace | None:
|
||||
return await self.storage.get_trace(trace_id, scope, trace_ref)
|
||||
async def get_trace(
|
||||
self,
|
||||
trace_id: str,
|
||||
scope: TraceScope,
|
||||
trace_ref: str = "",
|
||||
cursor: str | None = None,
|
||||
page_size: int | None = None,
|
||||
) -> Trace | None:
|
||||
return await self.storage.get_trace(trace_id, scope, trace_ref, cursor, page_size)
|
||||
|
||||
async def get_span(self, trace_id: str, span_id: str, scope: TraceScope, trace_ref: str = "") -> SpanDetail | None:
|
||||
return await self.storage.get_span(trace_id, span_id, scope, trace_ref)
|
||||
|
|
|
|||
123
litellm/types/decisions.py
Normal file
123
litellm/types/decisions.py
Normal file
|
|
@ -0,0 +1,123 @@
|
|||
from collections.abc import Mapping, Sequence
|
||||
from typing import Annotated, Literal, TypeAlias
|
||||
|
||||
from pydantic import ConfigDict, Field, PrivateAttr, model_validator, with_config
|
||||
from typing_extensions import ReadOnly, Required, TypedDict
|
||||
|
||||
from litellm.types.llms.base import LiteLLMPydanticObjectBase
|
||||
|
||||
DecisionsJSON: TypeAlias = str | Mapping[str, object] | Sequence[object]
|
||||
NoulCriteria: TypeAlias = Mapping[Literal["true", "false"], DecisionsJSON | None]
|
||||
|
||||
|
||||
class NoulQuestion(LiteLLMPydanticObjectBase):
|
||||
type: Literal["noul"]
|
||||
instructions: DecisionsJSON | None = None
|
||||
criteria: NoulCriteria | None = None
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def require_instructions_or_criteria(self) -> "NoulQuestion":
|
||||
if self.instructions is None and self.criteria is None:
|
||||
raise ValueError("A noul question requires instructions or criteria")
|
||||
return self
|
||||
|
||||
|
||||
class ChoiceQuestion(LiteLLMPydanticObjectBase):
|
||||
type: Literal["choice"]
|
||||
instructions: DecisionsJSON | None = None
|
||||
criteria: Annotated[Mapping[str, DecisionsJSON | None], Field(min_length=1, max_length=255)]
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
class ScoreQuestion(LiteLLMPydanticObjectBase):
|
||||
type: Literal["score"]
|
||||
instructions: DecisionsJSON | None = None
|
||||
criteria: Annotated[Sequence[DecisionsJSON], Field(min_length=1, max_length=10)]
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
DecisionQuestion: TypeAlias = Annotated[
|
||||
NoulQuestion | ChoiceQuestion | ScoreQuestion,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
DecisionQuestionMap: TypeAlias = Annotated[
|
||||
Mapping[Annotated[str, Field(min_length=1)], DecisionQuestion],
|
||||
Field(min_length=1, max_length=128),
|
||||
]
|
||||
|
||||
|
||||
class DecisionsRequestBody(LiteLLMPydanticObjectBase):
|
||||
state: DecisionsJSON
|
||||
questions: DecisionQuestionMap
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
class DecisionsRequest(DecisionsRequestBody):
|
||||
model: str
|
||||
|
||||
|
||||
@with_config(ConfigDict(extra="allow"))
|
||||
class DecisionsCallParams(TypedDict, total=False):
|
||||
model: Required[ReadOnly[str]]
|
||||
state: Required[ReadOnly[DecisionsJSON]]
|
||||
questions: Required[ReadOnly[DecisionQuestionMap]]
|
||||
api_key: ReadOnly[str | None]
|
||||
api_base: ReadOnly[str | None]
|
||||
timeout: ReadOnly[float | None]
|
||||
custom_llm_provider: ReadOnly[str | None]
|
||||
extra_headers: ReadOnly[Mapping[str, str] | None]
|
||||
|
||||
|
||||
class NoulAnswer(LiteLLMPydanticObjectBase):
|
||||
type: Literal["noul"]
|
||||
noul: float
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
class ChoiceAnswer(LiteLLMPydanticObjectBase):
|
||||
type: Literal["choice"]
|
||||
choice: str
|
||||
confidence: float
|
||||
probabilities: Mapping[str, float]
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
class ScoreAnswer(LiteLLMPydanticObjectBase):
|
||||
type: Literal["score"]
|
||||
score: float
|
||||
confidence: float
|
||||
legend: Mapping[str, DecisionsJSON]
|
||||
probabilities: Mapping[str, float]
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
DecisionAnswer: TypeAlias = Annotated[
|
||||
NoulAnswer | ChoiceAnswer | ScoreAnswer,
|
||||
Field(discriminator="type"),
|
||||
]
|
||||
|
||||
|
||||
class DecisionsUsage(LiteLLMPydanticObjectBase):
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
|
||||
class DecisionsResponse(LiteLLMPydanticObjectBase):
|
||||
model: str | None = None
|
||||
answers: Mapping[str, DecisionAnswer]
|
||||
usage: DecisionsUsage | None = None
|
||||
|
||||
model_config = ConfigDict(extra="allow", frozen=True)
|
||||
|
||||
_hidden_params: dict[str, object] = PrivateAttr(default_factory=dict)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -2566,6 +2566,20 @@
|
|||
"image_variations": true
|
||||
}
|
||||
},
|
||||
"typesafe": {
|
||||
"display_name": "TypeSafe (`typesafe`)",
|
||||
"url": "https://docs.typesafe.ai/models",
|
||||
"endpoints": {
|
||||
"systemone": true
|
||||
}
|
||||
},
|
||||
"strands_decider": {
|
||||
"display_name": "Strands Decider (`strands_decider`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers",
|
||||
"endpoints": {
|
||||
"systemone": true
|
||||
}
|
||||
},
|
||||
"tavily": {
|
||||
"display_name": "Tavily (`tavily`)",
|
||||
"url": "https://docs.litellm.ai/docs/search/tavily",
|
||||
|
|
|
|||
268
scripts/lens_dev.sh
Executable file
268
scripts/lens_dev.sh
Executable file
|
|
@ -0,0 +1,268 @@
|
|||
#!/usr/bin/env bash
|
||||
# One-command Lens local dev loop: proxy + Lens worker + hot-reload dashboard.
|
||||
#
|
||||
# LENS_DEV_PROXY_PORT proxy port (default 4000)
|
||||
# LENS_DEV_UI_PORT next dev port (default 3000)
|
||||
# LENS_DEV_MASTER_KEY master key, also the admin UI password
|
||||
# (default: random, generated once into .lens-dev/master_key)
|
||||
# LENS_DEV_CONFIG proxy config to use instead of the generated one
|
||||
# LENS_DEV_DATABASE_URL Postgres URL (default: the tracing stack's litellm DB on :15432)
|
||||
# LENS_DEV_REBUILD_RUST=1 rebuild the Rust bridge even if it imports
|
||||
#
|
||||
# State (master key, worker token, generated config, logs) lives in .lens-dev/ (gitignored).
|
||||
set -euo pipefail
|
||||
|
||||
repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
|
||||
proxy_port="${LENS_DEV_PROXY_PORT:-4000}"
|
||||
ui_port="${LENS_DEV_UI_PORT:-3000}"
|
||||
state_dir="${LENS_DEV_STATE_DIR:-$repo_root/.lens-dev}"
|
||||
log_dir="$state_dir/logs"
|
||||
token_file="$state_dir/worker_token"
|
||||
key_file="$state_dir/master_key"
|
||||
proxy_url="http://localhost:$proxy_port"
|
||||
py="${LENS_DEV_PYTHON:-$repo_root/.venv/bin/python}"
|
||||
database_url="${LENS_DEV_DATABASE_URL:-postgresql://litellm:litellm@127.0.0.1:15432/litellm}"
|
||||
clickhouse_url=http://default:local-tracing@127.0.0.1:18123
|
||||
master_key=""
|
||||
pids=()
|
||||
|
||||
die() { echo "lens-dev: $*" >&2; exit 1; }
|
||||
listening() { lsof -nP -iTCP:"$1" -sTCP:LISTEN >/dev/null 2>&1; }
|
||||
|
||||
# A fixed key would let anyone who can reach the proxy sign in as admin, so default to a
|
||||
# random key generated once per checkout and kept next to the worker token.
|
||||
load_master_key() {
|
||||
if [ -n "${LENS_DEV_MASTER_KEY:-}" ]; then
|
||||
master_key="$LENS_DEV_MASTER_KEY"
|
||||
return
|
||||
fi
|
||||
if [ ! -s "$key_file" ]; then
|
||||
(umask 077 && printf 'sk-%s\n' "$(openssl rand -hex 24)" > "$key_file")
|
||||
fi
|
||||
master_key="$(cat "$key_file")"
|
||||
}
|
||||
|
||||
# Only reuse a listener on 15432/18123 if it accepts the tracing stack's credentials;
|
||||
# start the compose service when nothing is listening; fail if something else is.
|
||||
# A LENS_DEV_DATABASE_URL is left to the proxy, which may use Prisma-only URL params.
|
||||
ensure_services() {
|
||||
local services=()
|
||||
if [ -z "${LENS_DEV_DATABASE_URL:-}" ] && listening 15432; then
|
||||
"$py" -c 'import sys, psycopg; psycopg.connect(sys.argv[1], connect_timeout=5).close()' "$database_url" 2>/dev/null \
|
||||
|| die "port 15432 is taken by something that isn't the tracing Postgres (litellm/litellm)"
|
||||
elif [ -z "${LENS_DEV_DATABASE_URL:-}" ]; then
|
||||
services+=(db)
|
||||
fi
|
||||
if listening 18123; then
|
||||
[ "$(curl -fsS --max-time 5 "$clickhouse_url/?query=SELECT%201" 2>/dev/null)" = 1 ] \
|
||||
|| die "port 18123 is taken by something that isn't the tracing ClickHouse (default/local-tracing)"
|
||||
else
|
||||
services+=(clickhouse)
|
||||
fi
|
||||
if [ "${#services[@]}" -gt 0 ]; then
|
||||
docker compose -f docker/docker-compose.tracing.yml up -d --wait "${services[@]}"
|
||||
else
|
||||
echo "lens-dev: reusing running Postgres and ClickHouse"
|
||||
fi
|
||||
}
|
||||
|
||||
write_default_config() {
|
||||
cat > "$1" <<'EOF'
|
||||
model_list:
|
||||
- model_name: gpt-4.1-mini
|
||||
litellm_params:
|
||||
model: openai/gpt-4.1-mini
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
store_prompts_in_spend_logs: true
|
||||
tracing:
|
||||
store:
|
||||
type: clickhouse
|
||||
url: os.environ/CLICKHOUSE_URL
|
||||
retention_days: 14
|
||||
EOF
|
||||
}
|
||||
|
||||
# litellm's implicit load_dotenv() walks up from a worktree into the parent checkout's
|
||||
# .env and picks up REDIS_* / UI_* from there. LITELLM_MODE=PRODUCTION turns that off;
|
||||
# this prints export lines for the same .env minus those vars, so provider keys still load.
|
||||
dotenv_exports() {
|
||||
"$py" - <<'PY'
|
||||
import os, re, shlex
|
||||
from dotenv import dotenv_values, find_dotenv
|
||||
|
||||
path = find_dotenv(usecwd=True)
|
||||
skip = re.compile(r"REDIS_.*|UI_USERNAME|UI_PASSWORD|LITELLM_MODE|ANTHROPIC_BASE_URL|ANTHROPIC_AUTH_TOKEN|ANTHROPIC_CUSTOM_HEADERS|OPENAI_BASE_URL|OPENAI_API_BASE")
|
||||
for key, value in (dotenv_values(path) if path else {}).items():
|
||||
if value is not None and key not in os.environ and re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", key) and not skip.fullmatch(key):
|
||||
print(f"export {key}={shlex.quote(value)}")
|
||||
PY
|
||||
}
|
||||
|
||||
# Run in the proxy's subshell: drop inherited settings that would point it at someone
|
||||
# else's services, then set the local stack's.
|
||||
proxy_env() {
|
||||
local var
|
||||
# Claude Code and similar tools export these; provider calls would go to them.
|
||||
unset ANTHROPIC_BASE_URL ANTHROPIC_AUTH_TOKEN ANTHROPIC_CUSTOM_HEADERS OPENAI_BASE_URL OPENAI_API_BASE
|
||||
for var in $(compgen -e | grep '^REDIS_' || true); do unset "$var"; done
|
||||
eval "$1"
|
||||
export LITELLM_MODE=PRODUCTION
|
||||
export LITELLM_MASTER_KEY="$master_key"
|
||||
if [ "$master_key" = sk-1234 ]; then export LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true; fi
|
||||
export LITELLM_SALT_KEY=sk-local-tracing-salt-key
|
||||
export DATABASE_URL="$database_url"
|
||||
export STORE_MODEL_IN_DB=True
|
||||
export CLICKHOUSE_URL="$clickhouse_url"
|
||||
export CLICKHOUSE_DATABASE=litellm
|
||||
export LITELLM_LOCAL_MODEL_COST_MAP=True
|
||||
export PROXY_BASE_URL="$proxy_url"
|
||||
export UI_USERNAME=admin
|
||||
export UI_PASSWORD="$master_key"
|
||||
}
|
||||
|
||||
# POST JSON as the admin and print one field of the response; dies with the body on failure.
|
||||
admin_post() {
|
||||
local body
|
||||
body="$(curl -sS --fail-with-body "$proxy_url$1" -H "Authorization: Bearer $master_key" \
|
||||
-H "Content-Type: application/json" -d "$2")" || die "POST $1 failed: $body"
|
||||
"$py" -c 'import json, sys; print(json.loads(sys.argv[1])[sys.argv[2]])' "$body" "$3"
|
||||
}
|
||||
|
||||
register_worker() {
|
||||
local key_hash worker_token
|
||||
key_hash="$(admin_post /key/generate "{\"key_alias\": \"lens-dev-$(date +%s)\"}" token)"
|
||||
worker_token="$(admin_post /lens/workers/register "{\"name\": \"lens-dev\", \"analysis_key_id\": \"$key_hash\"}" token)"
|
||||
(umask 077 && printf '%s\n' "$worker_token" > "$token_file")
|
||||
echo "lens-dev: registered a new Lens worker (token in $token_file)"
|
||||
}
|
||||
|
||||
# Auth runs before the handler, so protocol_version=1 answers 401 for a bad token and
|
||||
# 409 for a good one without claiming a job.
|
||||
ensure_worker_token() {
|
||||
local status
|
||||
if [ ! -s "$token_file" ]; then
|
||||
register_worker
|
||||
return
|
||||
fi
|
||||
status="$(curl -sS -o /dev/null -w '%{http_code}' -X POST "$proxy_url/lens/worker/claim?protocol_version=1" \
|
||||
-H "Authorization: Bearer $(cat "$token_file")")"
|
||||
case "$status" in
|
||||
409) echo "lens-dev: reusing worker token from $token_file" ;;
|
||||
401) echo "lens-dev: stored worker token was rejected"; register_worker ;;
|
||||
*) die "unexpected HTTP $status checking the worker token" ;;
|
||||
esac
|
||||
}
|
||||
|
||||
wait_for_proxy() {
|
||||
local proxy_pid="$1"
|
||||
echo "lens-dev: waiting for the proxy (log: $log_dir/proxy.log)"
|
||||
for _ in $(seq 1 300); do
|
||||
kill -0 "$proxy_pid" 2>/dev/null || die "proxy exited; see $log_dir/proxy.log"
|
||||
curl -fsS "$proxy_url/health/readiness" -H "Authorization: Bearer $master_key" >/dev/null 2>&1 && return
|
||||
sleep 1
|
||||
done
|
||||
die "proxy not ready after 300s; see $log_dir/proxy.log"
|
||||
}
|
||||
|
||||
# Children run in their own process groups (set -m), so killing -pid takes their trees too.
|
||||
cleanup() {
|
||||
local alive pid
|
||||
trap - EXIT INT TERM
|
||||
[ "${#pids[@]}" -gt 0 ] || return 0
|
||||
echo "lens-dev: stopping"
|
||||
for pid in "${pids[@]}"; do kill -TERM -- "-$pid" 2>/dev/null || true; done
|
||||
for _ in $(seq 1 20); do
|
||||
alive=0
|
||||
for pid in "${pids[@]}"; do kill -0 "$pid" 2>/dev/null && alive=1; done
|
||||
[ "$alive" = 0 ] && break
|
||||
sleep 0.5
|
||||
done
|
||||
for pid in "${pids[@]}"; do kill -KILL -- "-$pid" 2>/dev/null || true; done
|
||||
}
|
||||
|
||||
main() {
|
||||
local config_file exports proxy_pid pid key_hint
|
||||
if [ -n "${LENS_DEV_CONFIG:-}" ]; then
|
||||
[ -f "$LENS_DEV_CONFIG" ] || die "LENS_DEV_CONFIG not found: $LENS_DEV_CONFIG"
|
||||
config_file="$(cd "$(dirname "$LENS_DEV_CONFIG")" && pwd)/$(basename "$LENS_DEV_CONFIG")"
|
||||
fi
|
||||
cd "$repo_root"
|
||||
|
||||
listening "$proxy_port" && die "port $proxy_port is in use; set LENS_DEV_PROXY_PORT"
|
||||
listening "$ui_port" && die "port $ui_port is in use; set LENS_DEV_UI_PORT"
|
||||
[ "$proxy_port" != "$ui_port" ] || die "proxy and UI ports must differ"
|
||||
mkdir -p "$log_dir"
|
||||
load_master_key
|
||||
|
||||
uv sync --inexact --frozen --extra proxy --group proxy-dev --no-install-project
|
||||
ensure_services
|
||||
"$py" scripts/prisma_generate_if_needed.py
|
||||
|
||||
if [ "${LENS_DEV_REBUILD_RUST:-0}" = "1" ] || ! "$py" -c "import litellm.rust_bridge._native" >/dev/null 2>&1; then
|
||||
echo "lens-dev: building the Rust bridge (litellm.rust_bridge._native); the ClickHouse trace store uses it"
|
||||
VIRTUAL_ENV="$repo_root/.venv" uvx --from maturin==1.15.0 maturin develop \
|
||||
--release --manifest-path litellm-rust/crates/python-bridge/Cargo.toml --features extension-module
|
||||
fi
|
||||
|
||||
if [ ! -x ui/litellm-dashboard/node_modules/.bin/next ]; then
|
||||
(cd ui/litellm-dashboard && "$repo_root/scripts/with_dashboard_node.sh" npm ci)
|
||||
fi
|
||||
|
||||
if [ -z "${config_file:-}" ]; then
|
||||
config_file="$state_dir/config.yaml"
|
||||
write_default_config "$config_file"
|
||||
fi
|
||||
exports="$(dotenv_exports)"
|
||||
|
||||
trap cleanup EXIT
|
||||
trap 'exit 130' INT TERM
|
||||
set -m
|
||||
|
||||
(
|
||||
proxy_env "$exports"
|
||||
exec "$py" litellm/proxy/proxy_cli.py --config "$config_file" --host 127.0.0.1 --port "$proxy_port"
|
||||
) < /dev/null > "$log_dir/proxy.log" 2>&1 &
|
||||
proxy_pid=$!
|
||||
pids+=("$proxy_pid")
|
||||
|
||||
(
|
||||
cd ui/litellm-dashboard
|
||||
NEXT_PUBLIC_BASE_URL="$proxy_url" exec "$repo_root/scripts/with_dashboard_node.sh" npx next dev -p "$ui_port"
|
||||
) < /dev/null > "$log_dir/ui.log" 2>&1 &
|
||||
pids+=("$!")
|
||||
|
||||
wait_for_proxy "$proxy_pid"
|
||||
ensure_worker_token
|
||||
|
||||
LITELLM_MODE=PRODUCTION LITELLM_URL="$proxy_url" LENS_WORKER_TOKEN="$(cat "$token_file")" \
|
||||
"$py" -c "import asyncio, logging; from litellm.proxy.lens.worker import main; logging.basicConfig(level=logging.INFO); asyncio.run(main())" \
|
||||
< /dev/null > "$log_dir/worker.log" 2>&1 &
|
||||
pids+=("$!")
|
||||
|
||||
key_hint="password in $key_file"
|
||||
[ -z "${LENS_DEV_MASTER_KEY:-}" ] || key_hint="password from LENS_DEV_MASTER_KEY"
|
||||
cat <<EOF
|
||||
|
||||
Lens dev is up. Ctrl-C stops everything.
|
||||
Log in: $proxy_url/ui/login (admin / $key_hint)
|
||||
Lens: http://localhost:$ui_port/lens
|
||||
Logs: $log_dir/proxy.log
|
||||
$log_dir/worker.log
|
||||
$log_dir/ui.log
|
||||
Restart (Ctrl-C, make lens-dev) to pick up backend or worker edits; the UI hot-reloads.
|
||||
EOF
|
||||
|
||||
while :; do
|
||||
for pid in "${pids[@]}"; do
|
||||
kill -0 "$pid" 2>/dev/null || die "a child process (pid $pid) exited; check the logs above"
|
||||
done
|
||||
sleep 2
|
||||
done
|
||||
}
|
||||
|
||||
# Sourcing (tests) only defines the functions.
|
||||
if [ "${BASH_SOURCE[0]}" = "$0" ]; then
|
||||
main "$@"
|
||||
fi
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -33,6 +33,7 @@ EXCLUDED_FOLDERS = {
|
|||
"codex",
|
||||
"opencode",
|
||||
"deepagents",
|
||||
"tool_loop",
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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?"}},
|
||||
},
|
||||
)
|
||||
]
|
||||
|
|
|
|||
1572
tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py
Normal file
1572
tests/integration/providers/test_bedrock_kimi_k3_cache_point_wire.py
Normal file
File diff suppressed because it is too large
Load diff
262
tests/integration/providers/test_decisions_chaos.py
Normal file
262
tests/integration/providers/test_decisions_chaos.py
Normal file
|
|
@ -0,0 +1,262 @@
|
|||
import asyncio
|
||||
import json
|
||||
import re
|
||||
import signal
|
||||
import threading
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from queue import SimpleQueue
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
import psutil
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "pplx-decider-v1-27b"
|
||||
_CONFIG_MODEL: Final = "decisions-chaos"
|
||||
_API_KEY: Final = "synthetic-decisions-key"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
|
||||
_QUESTIONS: Final[dict[str, JsonValue]] = {"fine": {"type": "noul", "instructions": "Is the state fine?"}}
|
||||
_ROUTES: Final = ("/v1/decisions", "/decisions")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Call:
|
||||
route: str
|
||||
marker: str
|
||||
fail: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Served:
|
||||
call: _Call
|
||||
status: int
|
||||
call_id: str
|
||||
model_group: str
|
||||
text: str
|
||||
|
||||
|
||||
def _calls(count: int, *, fail: bool) -> tuple[_Call, ...]:
|
||||
return tuple(
|
||||
_Call(
|
||||
route=_ROUTES[index % len(_ROUTES)],
|
||||
marker=f"{'fail' if fail and index % 2 else 'ok'}-{uuid.uuid4().hex}",
|
||||
fail=fail and index % 2 == 1,
|
||||
)
|
||||
for index in range(count)
|
||||
)
|
||||
|
||||
|
||||
def _marker_of(request: Request) -> str:
|
||||
state: Final = _JSON_OBJECT.validate_json(request.body)["state"]
|
||||
assert isinstance(state, str), request.body
|
||||
return state
|
||||
|
||||
|
||||
def _reply(request: Request) -> Reply:
|
||||
marker: Final = _marker_of(request)
|
||||
if marker.startswith("fail-"):
|
||||
return Reply(status=500, body=json.dumps({"error": {"message": f"scripted outage {marker}"}}).encode())
|
||||
answer: Final = {
|
||||
"model": f"model-{marker}",
|
||||
"answers": {"fine": {"type": "noul", "noul": 0.5}},
|
||||
"usage": {"input_tokens": 12, "output_tokens": 1},
|
||||
}
|
||||
return Reply(body=json.dumps(answer).encode())
|
||||
|
||||
|
||||
async def _send(client: httpx.AsyncClient, key: str, model: str | None, call: _Call) -> _Served:
|
||||
body: Final[dict[str, JsonValue]] = {
|
||||
**({"model": model} if model is not None else {}),
|
||||
"state": call.marker,
|
||||
"questions": _QUESTIONS,
|
||||
"num_retries": 0,
|
||||
}
|
||||
response: Final = await client.post(call.route, json=body, headers={"Authorization": f"Bearer {key}"})
|
||||
return _Served(
|
||||
call=call,
|
||||
status=response.status_code,
|
||||
call_id=response.headers.get("x-litellm-call-id", ""),
|
||||
model_group=response.headers.get("x-litellm-model-group", ""),
|
||||
text=response.text,
|
||||
)
|
||||
|
||||
|
||||
async def _burst(
|
||||
base_url: str, key: str, model: str | None, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
|
||||
) -> tuple[_Served, ...]:
|
||||
async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client:
|
||||
results: Final = await asyncio.gather(
|
||||
*(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors
|
||||
)
|
||||
for result in results:
|
||||
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
|
||||
return tuple(result for result in results if isinstance(result, _Served))
|
||||
|
||||
|
||||
def _assert_served_its_own(served: _Served) -> None:
|
||||
assert served.call_id, served.text
|
||||
if served.call.fail:
|
||||
assert served.status == 500, (served.status, served.text)
|
||||
assert served.call.marker in served.text, served.text
|
||||
return
|
||||
assert served.status == 200, (served.status, served.text)
|
||||
assert _JSON_OBJECT.validate_json(served.text)["model"] == f"model-{served.call.marker}", served.text
|
||||
|
||||
|
||||
def _statuses_by_call_id(call_ids: tuple[str, ...]) -> dict[str, JsonValue]:
|
||||
placeholders: Final = ", ".join("%s" for _ in call_ids)
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
f'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id IN ({placeholders})', call_ids
|
||||
),
|
||||
lambda found: len(found) >= len(call_ids),
|
||||
seconds=70,
|
||||
)
|
||||
assert len(rows) == len(call_ids), rows
|
||||
return {str(row["request_id"]): row["status"] for row in rows}
|
||||
|
||||
|
||||
def _expected_statuses(served: tuple[_Served, ...]) -> dict[str, JsonValue]:
|
||||
return {item.call_id: "failure" if item.call.fail else "success" for item in served}
|
||||
|
||||
|
||||
def _health(gateway: Gateway, model: str) -> tuple[int, int]:
|
||||
health: Final = gateway.request("GET", "/health", params={"model": model})
|
||||
assert health.status_code in (200, 503), health.text
|
||||
report: Final = health.json()
|
||||
return (report["healthy_count"], report["unhealthy_count"])
|
||||
|
||||
|
||||
def _chaos_config(wire: Wire, tmp_path: Path) -> Path:
|
||||
config: Final = {
|
||||
**_JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())),
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": _CONFIG_MODEL,
|
||||
"litellm_params": {"model": f"perplexity/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY},
|
||||
}
|
||||
],
|
||||
}
|
||||
path: Final = tmp_path / "decisions-chaos.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]:
|
||||
text: Final = log.read_text()
|
||||
return tuple(int(pid) for pid in _STARTED_WORKER.findall(text)), text.count("Application startup complete.")
|
||||
|
||||
|
||||
def _open_upstream_connections(pid: int, upstream: str) -> int:
|
||||
port: Final = urlsplit(upstream).port
|
||||
return sum(
|
||||
1
|
||||
for connection in psutil.Process(pid).net_connections(kind="tcp")
|
||||
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
async def test_burst_over_both_routes_bills_each_call_once_with_its_own_status(gateway: Gateway) -> None:
|
||||
calls: Final = _calls(30, fail=True)
|
||||
with wire_server(_reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"perplexity/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls)
|
||||
assert len(served) == 30
|
||||
for item in served:
|
||||
_assert_served_its_own(item)
|
||||
assert len({item.call_id for item in served}) == 30
|
||||
assert _statuses_by_call_id(tuple(item.call_id for item in served)) == _expected_statuses(served)
|
||||
received: Final = wire.drain()
|
||||
assert sorted(_marker_of(request) for request in received) == sorted(call.marker for call in calls)
|
||||
assert {request.target for request in received} == {"/v1/decisions"}, received
|
||||
|
||||
|
||||
@pytest.mark.timeout(420)
|
||||
async def test_worker_sigkill_mid_burst_leaves_the_sibling_serving_the_default_model(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
calls: Final = _calls(20, fail=False)
|
||||
release: Final = threading.Event()
|
||||
held_markers: Final[SimpleQueue[str]] = SimpleQueue()
|
||||
|
||||
def held(request: Request) -> Reply:
|
||||
held_markers.put(_marker_of(request))
|
||||
assert release.wait(timeout=60), "The burst was never released"
|
||||
return _reply(request)
|
||||
|
||||
with wire_server(held) as wire:
|
||||
path: Final = _chaos_config(wire, tmp_path)
|
||||
with owned_proxy_process(
|
||||
gateway, tmp_path, {}, config=path, workers=2, extra_arguments=("--model", _CONFIG_MODEL)
|
||||
) as owned:
|
||||
candidate: Final = owned.gateway
|
||||
base_url: Final = str(candidate.client.base_url)
|
||||
workers, _ = eventually(
|
||||
lambda: _worker_startups(owned.log),
|
||||
lambda found: len(found[0]) == 2 and found[1] == 2,
|
||||
seconds=120,
|
||||
)
|
||||
burst: Final = asyncio.create_task(
|
||||
_burst(base_url, candidate.key, None, calls, tolerate_transport_errors=True)
|
||||
)
|
||||
await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60)
|
||||
held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers})
|
||||
assert sum(held_by.values()) == 20, held_by
|
||||
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
|
||||
victim: Final = psutil.Process(victim_pid)
|
||||
victim.suspend()
|
||||
victim.send_signal(signal.SIGKILL)
|
||||
release.set()
|
||||
served: Final = await burst
|
||||
assert len(served) == held_by[survivor_pid], (held_by, len(served))
|
||||
follow_up: Final = _Call(route="/decisions", marker=f"ok-{uuid.uuid4().hex}", fail=False)
|
||||
(answered,) = await _burst(base_url, candidate.key, None, (follow_up,))
|
||||
await asyncio.to_thread(
|
||||
eventually,
|
||||
lambda: _worker_startups(owned.log),
|
||||
lambda found: len(found[0]) == 3 and found[1] == 3,
|
||||
180,
|
||||
)
|
||||
for item in (*served, answered):
|
||||
_assert_served_its_own(item)
|
||||
assert item.model_group == _CONFIG_MODEL, item.model_group
|
||||
call_ids: Final = tuple(item.call_id for item in (*served, answered))
|
||||
assert set(_statuses_by_call_id(call_ids).values()) == {"success"}
|
||||
assert len({_marker_of(request) for request in wire.drain()}) == 21
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
async def test_upstream_outage_fails_its_calls_and_recovery_on_the_same_port_restores_them(gateway: Gateway) -> None:
|
||||
base_url: Final = str(gateway.client.base_url)
|
||||
with gateway.scenario() as scenario:
|
||||
with wire_server(_reply) as wire:
|
||||
port: Final = urlsplit(wire.url).port
|
||||
assert port is not None
|
||||
model: Final = scenario.model(model=f"perplexity/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
before: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False))
|
||||
assert _health(gateway, model) == (1, 0)
|
||||
during: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False))
|
||||
assert [item.status for item in during] == [500] * 5, [item.text for item in during]
|
||||
assert _health(gateway, model) == (0, 1)
|
||||
with wire_server(_reply, port=port):
|
||||
after: Final = await _burst(base_url, gateway.key, model, _calls(5, fail=False))
|
||||
assert _health(gateway, model) == (1, 0)
|
||||
for item in (*before, *after):
|
||||
_assert_served_its_own(item)
|
||||
statuses: Final = _statuses_by_call_id(tuple(item.call_id for item in (*before, *during, *after)))
|
||||
assert statuses == {
|
||||
**{item.call_id: "success" for item in (*before, *after)},
|
||||
**{item.call_id: "failure" for item in during},
|
||||
}
|
||||
456
tests/integration/providers/test_decisions_wire.py
Normal file
456
tests/integration/providers/test_decisions_wire.py
Normal file
|
|
@ -0,0 +1,456 @@
|
|||
import json
|
||||
import math
|
||||
import socket
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import yaml
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.process import owned_proxy_process
|
||||
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
|
||||
from integration.cost_calculation.cost_tracking_case import JsonResponse
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
import litellm
|
||||
|
||||
_API_KEY: Final = "synthetic-decisions-key"
|
||||
_ENV_KEY: Final = "synthetic-decisions-env-key"
|
||||
_PASS_THROUGH_MODEL: Final = "gpt-6-luna"
|
||||
_PASS_THROUGH_AUTHORIZATION: Final = "Bearer customer-held-upstream-key"
|
||||
_PASS_THROUGH_NEIGHBOUR: Final = "decisions-beside-a-pass-through"
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_USAGE: Final[dict[str, JsonValue]] = {"input_tokens": 367, "output_tokens": 3}
|
||||
_STATE: Final[dict[str, JsonValue]] = {"ticket": "The export job hangs at 99%", "component": "billing"}
|
||||
_QUESTIONS: Final[dict[str, JsonValue]] = {
|
||||
"defect": {"type": "noul", "instructions": "Is this a defect?"},
|
||||
"severity": {"type": "choice", "criteria": {"low": "cosmetic", "high": "blocks users"}, "weight": 2},
|
||||
"confidence": {"type": "score", "instructions": "How sure are you?", "criteria": ["unsure", "sure"]},
|
||||
}
|
||||
_ANSWERS: Final[dict[str, JsonValue]] = {
|
||||
"defect": {"type": "noul", "noul": 0.93},
|
||||
"severity": {"type": "choice", "choice": "high", "confidence": 0.8, "probabilities": {"low": 0.2, "high": 0.8}},
|
||||
"confidence": {
|
||||
"type": "score",
|
||||
"score": 1.0,
|
||||
"confidence": 0.7,
|
||||
"legend": {"0": "unsure", "1": "sure"},
|
||||
"probabilities": {"0": 0.3, "1": 0.7},
|
||||
},
|
||||
}
|
||||
_CHAT_BODY: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": "hi"}]}
|
||||
_CHAT_REPLY: Final[dict[str, JsonValue]] = {
|
||||
"id": "chatcmpl-decisions-parity",
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "pplx-decider-v1-27b",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 3, "completion_tokens": 1, "total_tokens": 4},
|
||||
}
|
||||
_SPEND_QUERY: Final = (
|
||||
"SELECT spend, status, call_type, model_group, custom_llm_provider, api_base, prompt_tokens, completion_tokens, "
|
||||
'request_tags FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Provider:
|
||||
name: str
|
||||
model: str
|
||||
path: str
|
||||
body_model: str
|
||||
api_key: str | None
|
||||
wraps_result: bool
|
||||
cost_map_key: str | None
|
||||
|
||||
|
||||
_PROVIDERS: Final = (
|
||||
_Provider(
|
||||
"perplexity",
|
||||
"perplexity/pplx-decider-v1-27b",
|
||||
"/v1/decisions",
|
||||
"pplx-decider-v1-27b",
|
||||
_API_KEY,
|
||||
False,
|
||||
"perplexity/pplx-decider-v1-27b",
|
||||
),
|
||||
_Provider("typesafe", "typesafe/jev-1.13.0", "/v1/systemone", "jev-1.13.0", _API_KEY, False, "typesafe/jev-1.13.0"),
|
||||
_Provider(
|
||||
"openrouter",
|
||||
"openrouter/typesafe/jev-1.13",
|
||||
"/alpha/decisions",
|
||||
"typesafe/jev-1.13",
|
||||
_API_KEY,
|
||||
False,
|
||||
"openrouter/typesafe/jev-1.13",
|
||||
),
|
||||
_Provider(
|
||||
"strands_decider", "strands_decider/systemone-decider", "/v1/systemone", "systemone-decider", None, False, None
|
||||
),
|
||||
_Provider(
|
||||
"cloudflare",
|
||||
"cloudflare/clef",
|
||||
"/ai/run/@cf/cloudflare/clef",
|
||||
"clef",
|
||||
_API_KEY,
|
||||
True,
|
||||
"cloudflare/@cf/cloudflare/clef",
|
||||
),
|
||||
)
|
||||
_PERPLEXITY: Final = _PROVIDERS[0]
|
||||
_INVALID_BODIES: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = (
|
||||
("missing questions", {"state": _STATE}),
|
||||
("missing state", {"questions": _QUESTIONS}),
|
||||
("numeric state", {"state": 5, "questions": _QUESTIONS}),
|
||||
("empty questions", {"state": _STATE, "questions": {}}),
|
||||
("noul without instructions or criteria", {"state": _STATE, "questions": {"q": {"type": "noul"}}}),
|
||||
("choice without criteria", {"state": _STATE, "questions": {"q": {"type": "choice", "criteria": {}}}}),
|
||||
(
|
||||
"score with eleven criteria",
|
||||
{"state": _STATE, "questions": {"q": {"type": "score", "criteria": [f"level-{index}" for index in range(11)]}}},
|
||||
),
|
||||
("unknown question type", {"state": _STATE, "questions": {"q": {"type": "ranking", "criteria": ["a"]}}}),
|
||||
)
|
||||
|
||||
|
||||
def _number(value: JsonValue) -> float:
|
||||
assert isinstance(value, (int, float)) and not isinstance(value, bool), value
|
||||
return float(value)
|
||||
|
||||
|
||||
def _expected_spend(cost_map_key: str | None) -> float:
|
||||
if cost_map_key is None:
|
||||
return 0.0
|
||||
prices: Final = object_value(json.loads(Path("model_prices_and_context_window.json").read_text())[cost_map_key])
|
||||
return _number(_USAGE["input_tokens"]) * _number(prices["input_cost_per_token"]) + _number(
|
||||
_USAGE["output_tokens"]
|
||||
) * _number(prices["output_cost_per_token"])
|
||||
|
||||
|
||||
def _answer_body(provider: _Provider) -> dict[str, JsonValue]:
|
||||
answer: Final[dict[str, JsonValue]] = {"model": provider.body_model, "answers": _ANSWERS, "usage": _USAGE}
|
||||
return {"result": answer, "success": True} if provider.wraps_result else answer
|
||||
|
||||
|
||||
def _register(scenario: Scenario, body: dict[str, JsonValue], *, status: int = 200) -> ScenarioHandle:
|
||||
handle: Final = register_scenario(
|
||||
f"decisions-{uuid.uuid4().hex[:12]}", JsonResponse(content_type="application/json", body=body, status=status)
|
||||
)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
return handle
|
||||
|
||||
|
||||
def _deployment(scenario: Scenario, handle: ScenarioHandle, provider: _Provider) -> str:
|
||||
return scenario.model(model=provider.model, api_base=handle.api_base(), api_key=provider.api_key)
|
||||
|
||||
|
||||
def _decide(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response:
|
||||
return gateway.request(
|
||||
"POST", "/v1/decisions", {"model": model, "state": _STATE, "questions": _QUESTIONS, **extra}, key=key
|
||||
)
|
||||
|
||||
|
||||
def _chat(gateway: Gateway, model: str, *, key: str | None = None, **extra: JsonValue) -> httpx.Response:
|
||||
return gateway.request("POST", "/v1/chat/completions", {"model": model, **_CHAT_BODY, **extra}, key=key)
|
||||
|
||||
|
||||
def _observed_requests(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]:
|
||||
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
|
||||
return tuple(map(object_value, upstream.get("/__observations").json()["requests"]))
|
||||
|
||||
|
||||
def _calls_to(requests: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> list[dict[str, JsonValue]]:
|
||||
return [request for request in requests if string_value(request["path"]).startswith(f"/{handle.scenario_id}/")]
|
||||
|
||||
|
||||
def _upstream_calls(gateway: Gateway, handle: ScenarioHandle) -> list[dict[str, JsonValue]]:
|
||||
return _calls_to(_observed_requests(gateway), handle)
|
||||
|
||||
|
||||
def _spend_row(call_id: str) -> dict[str, JsonValue]:
|
||||
rows: Final = eventually(lambda: read_rows(_SPEND_QUERY, (call_id,)), lambda found: len(found) == 1, seconds=70)
|
||||
return rows[0]
|
||||
|
||||
|
||||
def _free_closed_port() -> int:
|
||||
with socket.socket() as probe:
|
||||
probe.bind(("127.0.0.1", 0))
|
||||
return probe.getsockname()[1]
|
||||
|
||||
|
||||
def _pass_through_config(directory: Path, pass_through_target: str, native_api_base: str) -> Path:
|
||||
base: Final = _JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
|
||||
config: Final = {
|
||||
**base,
|
||||
"general_settings": {
|
||||
**object_value(base["general_settings"]),
|
||||
"pass_through_endpoints": [
|
||||
{
|
||||
"path": "/v1/decisions",
|
||||
"target": pass_through_target,
|
||||
"headers": {"Authorization": _PASS_THROUGH_AUTHORIZATION},
|
||||
}
|
||||
],
|
||||
},
|
||||
"model_list": [
|
||||
{
|
||||
"model_name": _PASS_THROUGH_NEIGHBOUR,
|
||||
"litellm_params": {"model": _PERPLEXITY.model, "api_base": native_api_base, "api_key": _API_KEY},
|
||||
}
|
||||
],
|
||||
}
|
||||
path: Final = directory / "decisions-pass-through.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
return path
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", _PROVIDERS, ids=lambda provider: provider.name)
|
||||
def test_each_provider_gets_its_own_path_key_and_body_and_is_billed_from_the_cost_map(
|
||||
gateway: Gateway, provider: _Provider
|
||||
) -> None:
|
||||
expected_spend: Final = _expected_spend(provider.cost_map_key)
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(provider))
|
||||
model: Final = _deployment(scenario, handle, provider)
|
||||
response: Final = _decide(gateway, model)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json() == {"model": provider.body_model, "answers": _ANSWERS, "usage": _USAGE}
|
||||
assert response.headers["x-litellm-model-group"] == model
|
||||
assert math.isclose(float(response.headers.get("x-litellm-response-cost", "0")), expected_spend, rel_tol=1e-9)
|
||||
(call,) = _upstream_calls(gateway, handle)
|
||||
assert call["path"] == f"/{handle.scenario_id}{provider.path}"
|
||||
assert call["authorization"] == (f"Bearer {provider.api_key}" if provider.api_key else "")
|
||||
assert call["body"] == {"model": provider.body_model, "state": _STATE, "questions": _QUESTIONS}
|
||||
row: Final = _spend_row(response.headers["x-litellm-call-id"])
|
||||
assert (
|
||||
row["status"],
|
||||
row["call_type"],
|
||||
row["custom_llm_provider"],
|
||||
row["model_group"],
|
||||
row["api_base"],
|
||||
row["prompt_tokens"],
|
||||
row["completion_tokens"],
|
||||
) == ("success", "adecisions", provider.name, model, f"{handle.api_base()}{provider.path}", 367, 3)
|
||||
assert math.isclose(_number(row["spend"]), expected_spend, rel_tol=1e-9), row
|
||||
|
||||
|
||||
def test_repeated_identical_requests_each_reach_the_upstream_and_are_each_billed(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
responses: Final = tuple(_decide(gateway, model) for _ in range(2))
|
||||
assert [response.status_code for response in responses] == [200, 200], [r.text for r in responses]
|
||||
call_ids: Final = tuple(response.headers["x-litellm-call-id"] for response in responses)
|
||||
assert len(set(call_ids)) == 2, call_ids
|
||||
assert len(_upstream_calls(gateway, handle)) == 2
|
||||
for call_id in call_ids:
|
||||
assert _spend_row(call_id)["status"] == "success"
|
||||
|
||||
|
||||
async def test_sdk_sync_and_async_clients_send_the_same_request(gateway: Gateway) -> None:
|
||||
provider: Final = _PROVIDERS[1]
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(provider))
|
||||
synchronous: Final = litellm.decisions(
|
||||
model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY
|
||||
)
|
||||
asynchronous: Final = await litellm.adecisions(
|
||||
model=provider.model, state=_STATE, questions=_QUESTIONS, api_base=handle.api_base(), api_key=_API_KEY
|
||||
)
|
||||
for response in (synchronous, asynchronous):
|
||||
assert response.model_dump(mode="json") == {
|
||||
"model": provider.body_model,
|
||||
"answers": _ANSWERS,
|
||||
"usage": _USAGE,
|
||||
}
|
||||
calls: Final = _upstream_calls(gateway, handle)
|
||||
assert len(calls) == 2, calls
|
||||
for call in calls:
|
||||
assert call["path"] == f"/{handle.scenario_id}{provider.path}"
|
||||
assert call["authorization"] == f"Bearer {_API_KEY}"
|
||||
assert call["body"] == {"model": provider.body_model, "state": _STATE, "questions": _QUESTIONS}
|
||||
|
||||
|
||||
def test_gateway_only_fields_stay_at_the_gateway_and_tags_reach_the_spend_log(gateway: Gateway) -> None:
|
||||
tag: Final = f"decisions-audit-{uuid.uuid4().hex[:8]}"
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
response: Final = _decide(
|
||||
gateway, model, user="auditor", num_retries=0, temperature=0.2, metadata={"tags": [tag]}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
(call,) = _upstream_calls(gateway, handle)
|
||||
assert call["body"] == {"model": _PERPLEXITY.body_model, "state": _STATE, "questions": _QUESTIONS}
|
||||
row: Final = _spend_row(response.headers["x-litellm-call-id"])
|
||||
tags: Final = row["request_tags"]
|
||||
assert isinstance(tags, list) and tag in tags, row
|
||||
|
||||
|
||||
def test_invalid_bodies_are_refused_at_the_gateway_without_an_upstream_call(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
for label, body in _INVALID_BODIES:
|
||||
response: Final = gateway.request("POST", "/v1/decisions", {"model": model, **body})
|
||||
assert response.status_code == 400, (label, response.text)
|
||||
assert "Invalid Decisions request" in response.text, (label, response.text)
|
||||
assert _upstream_calls(gateway, handle) == []
|
||||
|
||||
|
||||
def test_unknown_model_is_refused_like_chat(gateway: Gateway) -> None:
|
||||
model: Final = f"missing-{uuid.uuid4().hex}"
|
||||
decisions: Final = _decide(gateway, model)
|
||||
chat: Final = _chat(gateway, model)
|
||||
assert 400 <= decisions.status_code < 500, decisions.text
|
||||
assert decisions.status_code == chat.status_code, (decisions.text, chat.text)
|
||||
|
||||
|
||||
def test_key_checks_match_chat(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
anonymous: Final = gateway.client.post(
|
||||
"/v1/decisions", json={"model": model, "state": _STATE, "questions": _QUESTIONS}
|
||||
)
|
||||
assert anonymous.status_code == 401, anonymous.text
|
||||
restricted: Final = scenario.key(models=[f"other-{uuid.uuid4().hex}"])
|
||||
refused: Final = _decide(gateway, model, key=restricted)
|
||||
assert 400 <= refused.status_code < 500, refused.text
|
||||
assert refused.status_code == _chat(gateway, model, key=restricted).status_code, refused.text
|
||||
assert _upstream_calls(gateway, handle) == []
|
||||
spender: Final = scenario.key(max_budget=1e-06)
|
||||
first: Final = _decide(gateway, model, key=spender)
|
||||
assert first.status_code == 200, first.text
|
||||
blocked: Final = eventually(
|
||||
lambda: _decide(gateway, model, key=spender), lambda response: response.status_code != 200, seconds=70
|
||||
)
|
||||
assert 400 <= blocked.status_code < 500, blocked.text
|
||||
assert blocked.status_code == _chat(gateway, model, key=spender).status_code, blocked.text
|
||||
|
||||
|
||||
def test_request_body_api_base_is_refused_like_chat_without_an_upstream_call(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
decisions: Final = _decide(gateway, model, api_base=f"http://127.0.0.1:{_free_closed_port()}")
|
||||
chat: Final = _chat(gateway, model, api_base=f"http://127.0.0.1:{_free_closed_port()}")
|
||||
assert 400 <= decisions.status_code < 500, decisions.text
|
||||
assert decisions.status_code == chat.status_code, (decisions.text, chat.text)
|
||||
assert _upstream_calls(gateway, handle) == []
|
||||
|
||||
|
||||
def test_a_deployment_without_a_key_sends_the_provider_env_key_to_its_configured_api_base(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
model: Final = scenario.model(model=_PERPLEXITY.model, api_base=handle.api_base(), api_key=None)
|
||||
response: Final = _decide(gateway, model)
|
||||
assert response.status_code == 200, response.text
|
||||
(call,) = _upstream_calls(gateway, handle)
|
||||
assert (call["path"], call["authorization"]) == (f"/{handle.scenario_id}/v1/decisions", f"Bearer {_ENV_KEY}")
|
||||
|
||||
|
||||
def test_a_deployment_opted_into_client_api_base_sends_decisions_and_chat_to_the_body_api_base(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
configured: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
decisions_target: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
chat_target: Final = _register(scenario, _CHAT_REPLY)
|
||||
model: Final = scenario.model(
|
||||
model=_PERPLEXITY.model,
|
||||
api_base=configured.api_base(),
|
||||
api_key=_API_KEY,
|
||||
configurable_clientside_auth_params=["api_base"],
|
||||
)
|
||||
decisions: Final = _decide(gateway, model, api_base=decisions_target.api_base())
|
||||
chat: Final = _chat(gateway, model, api_base=chat_target.api_base())
|
||||
assert decisions.status_code == 200, decisions.text
|
||||
assert chat.status_code == 200, chat.text
|
||||
observed: Final = _observed_requests(gateway)
|
||||
assert [call["path"] for call in _calls_to(observed, decisions_target)] == [
|
||||
f"/{decisions_target.scenario_id}/v1/decisions"
|
||||
]
|
||||
assert [call["path"] for call in _calls_to(observed, chat_target)] == [
|
||||
f"/{chat_target.scenario_id}/chat/completions"
|
||||
]
|
||||
assert _calls_to(observed, configured) == []
|
||||
|
||||
|
||||
def test_a_config_pass_through_at_v1_decisions_keeps_answering_and_the_native_api_serves_decisions(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
pass_through_target: Final = _register(
|
||||
scenario, {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE}
|
||||
)
|
||||
native_target: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
config: Final = _pass_through_config(
|
||||
tmp_path, f"{pass_through_target.api_base()}/v1/decisions", native_target.api_base()
|
||||
)
|
||||
with owned_proxy_process(gateway, tmp_path, {}, config=config) as owned:
|
||||
through: Final = _decide(owned.gateway, _PASS_THROUGH_MODEL)
|
||||
native: Final = owned.gateway.request(
|
||||
"POST", "/decisions", {"model": _PASS_THROUGH_NEIGHBOUR, "state": _STATE, "questions": _QUESTIONS}
|
||||
)
|
||||
assert through.status_code == 200, through.text
|
||||
assert through.json() == {"model": _PASS_THROUGH_MODEL, "answers": _ANSWERS, "usage": _USAGE}
|
||||
observed: Final = _observed_requests(gateway)
|
||||
(forwarded,) = _calls_to(observed, pass_through_target)
|
||||
assert (forwarded["path"], forwarded["authorization"], object_value(forwarded["body"])["model"]) == (
|
||||
f"/{pass_through_target.scenario_id}/v1/decisions",
|
||||
_PASS_THROUGH_AUTHORIZATION,
|
||||
_PASS_THROUGH_MODEL,
|
||||
)
|
||||
assert native.status_code == 200, native.text
|
||||
assert [call["path"] for call in _calls_to(observed, native_target)] == [
|
||||
f"/{native_target.scenario_id}/v1/decisions"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("status", (401, 429, 500))
|
||||
def test_upstream_errors_keep_their_status_and_log_an_unbilled_failure(gateway: Gateway, status: int) -> None:
|
||||
marker: Final = f"scripted-{status}-{uuid.uuid4().hex[:8]}"
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, {"error": {"message": marker}}, status=status)
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
response: Final = _decide(gateway, model, num_retries=0)
|
||||
assert response.status_code == status, response.text
|
||||
assert marker in response.text
|
||||
row: Final = _spend_row(response.headers["x-litellm-call-id"])
|
||||
assert (row["status"], row["call_type"], row["model_group"], _number(row["spend"])) == (
|
||||
"failure",
|
||||
"adecisions",
|
||||
model,
|
||||
0.0,
|
||||
)
|
||||
|
||||
|
||||
def test_upstream_success_without_answers_is_a_gateway_side_server_error(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, {"model": _PERPLEXITY.body_model, "usage": _USAGE})
|
||||
model: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
response: Final = _decide(gateway, model, num_retries=0)
|
||||
assert 500 <= response.status_code < 600, response.text
|
||||
assert "answers" in response.text
|
||||
assert _spend_row(response.headers["x-litellm-call-id"])["status"] == "failure"
|
||||
|
||||
|
||||
def test_unreachable_upstream_fails_only_its_own_deployment(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
handle: Final = _register(scenario, _answer_body(_PERPLEXITY))
|
||||
healthy: Final = _deployment(scenario, handle, _PERPLEXITY)
|
||||
dead: Final = scenario.model(
|
||||
model=_PERPLEXITY.model, api_base=f"http://127.0.0.1:{_free_closed_port()}", api_key=_API_KEY
|
||||
)
|
||||
failed: Final = _decide(gateway, dead, num_retries=0)
|
||||
assert 500 <= failed.status_code < 600, failed.text
|
||||
assert _spend_row(failed.headers["x-litellm-call-id"])["status"] == "failure"
|
||||
served: Final = _decide(gateway, healthy)
|
||||
assert served.status_code == 200, served.text
|
||||
assert len(_upstream_calls(gateway, handle)) == 1
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
tests/unit/decisions/__init__.py
Normal file
0
tests/unit/decisions/__init__.py
Normal file
592
tests/unit/decisions/test_main.py
Normal file
592
tests/unit/decisions/test_main.py
Normal file
|
|
@ -0,0 +1,592 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.types.decisions import (
|
||||
ChoiceAnswer,
|
||||
DecisionsResponse,
|
||||
DecisionsUsage,
|
||||
NoulAnswer,
|
||||
ScoreAnswer,
|
||||
)
|
||||
|
||||
_QUESTIONS: Final[Mapping[str, object]] = MappingProxyType(
|
||||
{
|
||||
"is_defect": {"type": "noul", "instructions": "Is this a defect?", "provider_field": "kept"},
|
||||
"sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}},
|
||||
"severity": {"type": "score", "criteria": ["none", "low", "high"]},
|
||||
}
|
||||
)
|
||||
_INPUT_TOKENS: Final[int] = 367
|
||||
_OUTPUT_TOKENS: Final[int] = 3
|
||||
_RESPONSE: Final[Mapping[str, object]] = {
|
||||
"model": "jev-1.13",
|
||||
"answers": {
|
||||
"is_defect": {"type": "noul", "noul": 0.9},
|
||||
"sentiment": {
|
||||
"type": "choice",
|
||||
"choice": "positive",
|
||||
"confidence": 0.8,
|
||||
"probabilities": {"positive": 0.8, "negative": 0.2},
|
||||
},
|
||||
"severity": {
|
||||
"type": "score",
|
||||
"score": 1,
|
||||
"confidence": 0.7,
|
||||
"legend": {"0": "none", "1": "low", "2": "high"},
|
||||
"probabilities": {"0": 0.1, "1": 0.8, "2": 0.1},
|
||||
},
|
||||
},
|
||||
"usage": {"input_tokens": _INPUT_TOKENS, "output_tokens": _OUTPUT_TOKENS},
|
||||
}
|
||||
_STRANDS_RESPONSE: Final[Mapping[str, object]] = {
|
||||
"model": "strands-decider-2B-hobson-v19",
|
||||
"answers": {
|
||||
"severity": {
|
||||
"type": "score",
|
||||
"score": 1,
|
||||
"confidence": 0.7,
|
||||
"legend": {"0": "none", "1": "low", "2": "high"},
|
||||
"probabilities": {"0": 0.1, "1": 0.8, "2": 0.1},
|
||||
}
|
||||
},
|
||||
"usage": {"input_tokens": 216, "output_tokens": 3},
|
||||
"latency_ms": 3722.17,
|
||||
}
|
||||
_PROVIDERS: Final[tuple[tuple[str, str, str, str], ...]] = (
|
||||
(
|
||||
"perplexity",
|
||||
"perplexity/pplx-decider-v1-27b",
|
||||
"https://api.perplexity.ai/v1/decisions",
|
||||
"pplx-decider-v1-27b",
|
||||
),
|
||||
("typesafe", "typesafe/jev-1.13", "https://api.typesafe.ai/v1/systemone", "jev-1.13"),
|
||||
(
|
||||
"openrouter",
|
||||
"openrouter/typesafe/jev-1.13",
|
||||
"https://openrouter.ai/api/alpha/decisions",
|
||||
"typesafe/jev-1.13",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class _RecordingLogger(CustomLogger):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self.standard_logging_object: Mapping[str, object] | None = None
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: object,
|
||||
end_time: object,
|
||||
) -> None:
|
||||
standard_logging_object: Final = kwargs.get("standard_logging_object")
|
||||
if isinstance(standard_logging_object, dict):
|
||||
self.standard_logging_object = standard_logging_object
|
||||
|
||||
|
||||
async def _drain_logging_worker() -> None:
|
||||
await asyncio.sleep(0)
|
||||
GLOBAL_LOGGING_WORKER.start()
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), timeout=10.0)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _httpx_transport(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(("provider", "model", "url", "upstream_model"), _PROVIDERS)
|
||||
async def test_adecisions_sends_the_provider_wire_contract(
|
||||
provider: str,
|
||||
model: str,
|
||||
url: str,
|
||||
upstream_model: str,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
route: Final = respx_mock.post(url).respond(json=_RESPONSE)
|
||||
|
||||
response: Final = await litellm.adecisions(
|
||||
model=model,
|
||||
state={"source": "unit-test"},
|
||||
questions=_QUESTIONS,
|
||||
api_key="caller-key",
|
||||
extra_headers={
|
||||
"x-request-tag": "decisions-test",
|
||||
"AUTHORIZATION": "attacker-key",
|
||||
"Content-Type": "text/plain",
|
||||
},
|
||||
internal_kwarg="must-not-leak",
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert len(respx_mock.calls) == 1
|
||||
request: Final = respx_mock.calls[0].request
|
||||
assert request.headers["authorization"] == "Bearer caller-key"
|
||||
assert request.headers["content-type"] == "application/json"
|
||||
assert request.headers["x-request-tag"] == "decisions-test"
|
||||
assert json.loads(request.content) == {
|
||||
"model": upstream_model,
|
||||
"state": {"source": "unit-test"},
|
||||
"questions": {
|
||||
"is_defect": {
|
||||
"type": "noul",
|
||||
"instructions": "Is this a defect?",
|
||||
"provider_field": "kept",
|
||||
},
|
||||
"sentiment": {"type": "choice", "criteria": {"positive": None, "negative": "unhappy"}},
|
||||
"severity": {"type": "score", "criteria": ["none", "low", "high"]},
|
||||
},
|
||||
}
|
||||
assert isinstance(response.answers["is_defect"], NoulAnswer)
|
||||
assert isinstance(response.answers["sentiment"], ChoiceAnswer)
|
||||
assert isinstance(response.answers["severity"], ScoreAnswer)
|
||||
assert response._hidden_params["custom_llm_provider"] == provider
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_dispatches_typesafe_decisions_without_api_base(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.delenv("TYPESAFE_API_KEY", raising=False)
|
||||
monkeypatch.delenv("TYPESAFE_API_BASE", raising=False)
|
||||
provider_resolution: Final = litellm.get_llm_provider("typesafe/jev-latest")
|
||||
|
||||
assert provider_resolution[:2] == ("jev-latest", "typesafe")
|
||||
|
||||
router: Final = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "jev",
|
||||
"litellm_params": {
|
||||
"model": "typesafe/jev-latest",
|
||||
"api_key": "k",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
upstream: Final = respx_mock.post("https://api.typesafe.ai/v1/systemone").respond(json=_RESPONSE)
|
||||
|
||||
response: Final = await router.adecisions(
|
||||
model="jev",
|
||||
state="router-test",
|
||||
questions={
|
||||
"sentiment": {
|
||||
"type": "choice",
|
||||
"criteria": {"positive": None, "negative": "unhappy"},
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
assert upstream.called
|
||||
assert len(respx_mock.calls) == 1
|
||||
assert json.loads(respx_mock.calls[0].request.content) == {
|
||||
"model": "jev-latest",
|
||||
"state": "router-test",
|
||||
"questions": {
|
||||
"sentiment": {
|
||||
"type": "choice",
|
||||
"criteria": {"positive": None, "negative": "unhappy"},
|
||||
}
|
||||
},
|
||||
}
|
||||
assert respx_mock.calls[0].request.headers["authorization"] == "Bearer k"
|
||||
assert isinstance(response.answers["sentiment"], ChoiceAnswer)
|
||||
assert response.answers["sentiment"].choice == "positive"
|
||||
|
||||
|
||||
def test_decisions_uses_the_same_wire_contract_for_sync_calls(respx_mock: respx.MockRouter) -> None:
|
||||
route: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE)
|
||||
|
||||
response: Final = litellm.decisions(
|
||||
model="perplexity/pplx-decider-v1-27b",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
api_key="caller-key",
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert response.model == "jev-1.13"
|
||||
|
||||
|
||||
def test_openrouter_response_keeps_provider_fields(respx_mock: respx.MockRouter) -> None:
|
||||
payload: Final = {
|
||||
**_RESPONSE,
|
||||
"id": "decision-1",
|
||||
"provider": "typesafe",
|
||||
"usage": {**_RESPONSE["usage"], "cost": 0.25},
|
||||
}
|
||||
respx_mock.post("https://openrouter.ai/api/alpha/decisions").respond(json=payload)
|
||||
|
||||
response: Final = litellm.decisions(
|
||||
model="openrouter/typesafe/jev-1.13",
|
||||
state="review",
|
||||
questions=_QUESTIONS,
|
||||
api_key="caller-key",
|
||||
)
|
||||
|
||||
assert response.model_extra["id"] == "decision-1"
|
||||
assert response.model_extra["provider"] == "typesafe"
|
||||
assert response.usage is not None
|
||||
assert response.usage.model_extra["cost"] == 0.25
|
||||
|
||||
|
||||
def test_decisions_cost_uses_litellm_token_pricing() -> None:
|
||||
response: Final = DecisionsResponse(
|
||||
model="pplx-decider-v1-27b",
|
||||
answers={},
|
||||
usage=DecisionsUsage(input_tokens=_INPUT_TOKENS, output_tokens=_OUTPUT_TOKENS),
|
||||
)
|
||||
response._hidden_params = {
|
||||
"model": "perplexity/pplx-decider-v1-27b",
|
||||
"custom_llm_provider": "perplexity",
|
||||
}
|
||||
|
||||
cost: Final = litellm.completion_cost(completion_response=response)
|
||||
perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"]
|
||||
expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float(
|
||||
perplexity_cost["output_cost_per_token"]
|
||||
)
|
||||
|
||||
assert expected_cost > 0
|
||||
assert cost == pytest.approx(expected_cost)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_decisions_cost_is_in_standard_logging_object(respx_mock: respx.MockRouter) -> None:
|
||||
respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(json=_RESPONSE)
|
||||
recording_logger: Final = _RecordingLogger()
|
||||
original_callbacks: Final = litellm.callbacks
|
||||
litellm.callbacks = [recording_logger]
|
||||
|
||||
try:
|
||||
await litellm.adecisions(
|
||||
model="perplexity/pplx-decider-v1-27b",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
api_key="caller-key",
|
||||
)
|
||||
await _drain_logging_worker()
|
||||
finally:
|
||||
litellm.callbacks = original_callbacks
|
||||
|
||||
assert recording_logger.standard_logging_object is not None
|
||||
perplexity_cost: Final = litellm.model_cost["perplexity/pplx-decider-v1-27b"]
|
||||
expected_cost: Final = _INPUT_TOKENS * float(perplexity_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float(
|
||||
perplexity_cost["output_cost_per_token"]
|
||||
)
|
||||
|
||||
assert expected_cost > 0
|
||||
assert recording_logger.standard_logging_object["response_cost"] == pytest.approx(expected_cost)
|
||||
assert recording_logger.standard_logging_object["prompt_tokens"] == _INPUT_TOKENS
|
||||
assert recording_logger.standard_logging_object["completion_tokens"] == _OUTPUT_TOKENS
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unknown_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None:
|
||||
with pytest.raises(litellm.BadRequestError, match="Supported providers"):
|
||||
await litellm.adecisions(
|
||||
model="unknown/jev-1.13",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
api_key="caller-key",
|
||||
)
|
||||
|
||||
assert len(respx_mock.calls) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_custom_provider_is_rejected_before_http(respx_mock: respx.MockRouter) -> None:
|
||||
with pytest.raises(litellm.BadRequestError, match="Supported providers"):
|
||||
await litellm.adecisions(
|
||||
model="perplexity/pplx-decider-v1-27b",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
api_key="caller-key",
|
||||
custom_llm_provider="",
|
||||
)
|
||||
|
||||
assert len(respx_mock.calls) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_invalid_question_is_rejected_before_http(respx_mock: respx.MockRouter) -> None:
|
||||
with pytest.raises(litellm.BadRequestError, match="Invalid Decisions request"):
|
||||
await litellm.adecisions(
|
||||
model="perplexity/pplx-decider-v1-27b",
|
||||
state="review",
|
||||
questions={"sentiment": {"type": "choice"}},
|
||||
api_key="caller-key",
|
||||
)
|
||||
|
||||
assert len(respx_mock.calls) == 0
|
||||
|
||||
|
||||
def test_upstream_bad_request_maps_to_litellm_error(respx_mock: respx.MockRouter) -> None:
|
||||
respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(
|
||||
status_code=400,
|
||||
json={"error": {"message": "invalid decision"}},
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError):
|
||||
litellm.decisions(
|
||||
model="perplexity/pplx-decider-v1-27b",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
api_key="caller-key",
|
||||
)
|
||||
|
||||
|
||||
def test_server_key_is_sent_to_an_explicit_api_base(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setenv("PERPLEXITYAI_API_KEY", "server-key")
|
||||
monkeypatch.delenv("PERPLEXITY_API_KEY", raising=False)
|
||||
route: Final = respx_mock.post("https://egress.example/perplexity/v1/decisions").respond(json=_RESPONSE)
|
||||
|
||||
litellm.decisions(
|
||||
model="perplexity/pplx-decider-v1-27b",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
api_base="https://egress.example/perplexity",
|
||||
)
|
||||
|
||||
assert route.call_count == 1
|
||||
assert route.calls[0].request.headers["authorization"] == "Bearer server-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", ("cloudflare/clef", "cloudflare/@cf/cloudflare/clef"))
|
||||
@pytest.mark.parametrize("wrapped", (False, True))
|
||||
async def test_cloudflare_clef_resolves_model_and_response_envelope(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
model: str,
|
||||
wrapped: bool,
|
||||
) -> None:
|
||||
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
|
||||
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
|
||||
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
|
||||
response_body: Final[Mapping[str, object]] = (
|
||||
{"result": _RESPONSE, "success": True, "errors": [], "messages": []} if wrapped else _RESPONSE
|
||||
)
|
||||
route: Final = respx_mock.post(
|
||||
"https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef"
|
||||
).respond(json=response_body)
|
||||
|
||||
response: Final = await litellm.adecisions(
|
||||
model=model,
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
)
|
||||
|
||||
assert route.called
|
||||
request: Final = respx_mock.calls[0].request
|
||||
assert request.headers["authorization"] == "Bearer cloudflare-key"
|
||||
assert json.loads(request.content) == {
|
||||
"model": "clef",
|
||||
"state": "review",
|
||||
"questions": {"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
}
|
||||
assert response.answers == DecisionsResponse.model_validate(_RESPONSE).answers
|
||||
assert response._hidden_params["model"] == "cloudflare/@cf/cloudflare/clef"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloudflare_clef_flash_uses_flash_endpoint_and_request_model(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
|
||||
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
|
||||
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
|
||||
route: Final = respx_mock.post(
|
||||
"https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef-flash"
|
||||
).respond(json=_RESPONSE)
|
||||
|
||||
await litellm.adecisions(
|
||||
model="cloudflare/clef-flash",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert json.loads(respx_mock.calls[0].request.content)["model"] == "clef-flash"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloudflare_api_base_from_env_uses_workers_ai_run_path(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setenv("CLOUDFLARE_API_BASE", "https://api.cloudflare.com/client/v4/accounts/acct/ai/v1")
|
||||
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
|
||||
monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
|
||||
route: Final = respx_mock.post(
|
||||
"https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef"
|
||||
).respond(json=_RESPONSE)
|
||||
|
||||
await litellm.adecisions(
|
||||
model="cloudflare/clef",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
)
|
||||
|
||||
assert route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloudflare_requires_account_id_or_api_base_before_http(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.delenv("CLOUDFLARE_ACCOUNT_ID", raising=False)
|
||||
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
|
||||
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="Missing CLOUDFLARE_ACCOUNT_ID - set CLOUDFLARE_ACCOUNT_ID"):
|
||||
await litellm.adecisions(
|
||||
model="cloudflare/clef",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
)
|
||||
|
||||
assert len(respx_mock.calls) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloudflare_clef_cost_uses_the_model_cost_map(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct")
|
||||
monkeypatch.setenv("CLOUDFLARE_API_KEY", "cloudflare-key")
|
||||
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
|
||||
respx_mock.post("https://api.cloudflare.com/client/v4/accounts/acct/ai/run/@cf/cloudflare/clef").respond(
|
||||
json=_RESPONSE
|
||||
)
|
||||
|
||||
response: Final = await litellm.adecisions(
|
||||
model="cloudflare/clef",
|
||||
state="review",
|
||||
questions={"is_defect": {"type": "noul", "instructions": "Is this a defect?"}},
|
||||
)
|
||||
|
||||
cost: Final = litellm.completion_cost(completion_response=response)
|
||||
clef_cost: Final = litellm.model_cost["cloudflare/@cf/cloudflare/clef"]
|
||||
expected_cost: Final = _INPUT_TOKENS * float(clef_cost["input_cost_per_token"]) + _OUTPUT_TOKENS * float(
|
||||
clef_cost["output_cost_per_token"]
|
||||
)
|
||||
|
||||
assert expected_cost > 0
|
||||
assert cost == pytest.approx(expected_cost)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strands_decider_requires_api_base_before_http(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="api_base is required"):
|
||||
await litellm.adecisions(
|
||||
model="strands_decider/strands-decider-2B-hobson-v19",
|
||||
state="review",
|
||||
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
|
||||
)
|
||||
|
||||
assert len(respx_mock.calls) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strands_decider_without_key_preserves_response_extras(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
|
||||
route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE)
|
||||
|
||||
response: Final = await litellm.adecisions(
|
||||
model="strands_decider/strands-decider-2B-hobson-v19",
|
||||
state="review",
|
||||
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
|
||||
api_base="https://strands.example",
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert "authorization" not in respx_mock.calls[0].request.headers
|
||||
assert response.model_extra["latency_ms"] == _STRANDS_RESPONSE["latency_ms"]
|
||||
severity: Final = response.answers["severity"]
|
||||
assert isinstance(severity, ScoreAnswer)
|
||||
assert severity.legend == {"0": "none", "1": "low", "2": "high"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strands_decider_uses_key_from_matching_environment_base(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setenv("STRANDS_DECIDER_API_BASE", "https://strands.example")
|
||||
monkeypatch.setenv("STRANDS_DECIDER_API_KEY", "strands-key")
|
||||
route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE)
|
||||
|
||||
await litellm.adecisions(
|
||||
model="strands_decider/strands-decider-2B-hobson-v19",
|
||||
state="review",
|
||||
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
|
||||
api_base="https://strands.example",
|
||||
)
|
||||
|
||||
assert route.called
|
||||
assert respx_mock.calls[0].request.headers["authorization"] == "Bearer strands-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_strands_decider_provider_resolution_and_router_dispatch(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
|
||||
provider_resolution: Final = litellm.get_llm_provider("strands_decider/strands-decider-2B-hobson-v19")
|
||||
router: Final = litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "strands",
|
||||
"litellm_params": {
|
||||
"model": "strands_decider/strands-decider-2B-hobson-v19",
|
||||
"api_base": "https://strands.example",
|
||||
},
|
||||
}
|
||||
]
|
||||
)
|
||||
route: Final = respx_mock.post("https://strands.example/v1/systemone").respond(json=_STRANDS_RESPONSE)
|
||||
|
||||
response: Final = await router.adecisions(
|
||||
model="strands",
|
||||
state="review",
|
||||
questions={"severity": {"type": "score", "criteria": ["none", "low", "high"]}},
|
||||
)
|
||||
|
||||
assert provider_resolution[:2] == ("strands-decider-2B-hobson-v19", "strands_decider")
|
||||
assert route.called
|
||||
assert response.model == _STRANDS_RESPONSE["model"]
|
||||
476
tests/unit/harness/handlers/test_tool_loop_handler.py
Normal file
476
tests/unit/harness/handlers/test_tool_loop_handler.py
Normal file
|
|
@ -0,0 +1,476 @@
|
|||
"""Tests for the in-process Tool Loop handler."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable, Iterator, Mapping
|
||||
from pathlib import Path
|
||||
from typing import Final, Literal
|
||||
|
||||
import litellm
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm import sandbox
|
||||
from litellm.harness.context import GatewayTarget, SessionContext
|
||||
from litellm.harness.handlers.tool_loop_handler import ToolLoopHandler
|
||||
from litellm.harness.options import ToolLoopOptions
|
||||
from litellm.harness.types import Approval, Event, Harness, Text, ToolCall, ToolResult
|
||||
from litellm.llms.base_llm.harness.transformation import HarnessTurnError
|
||||
from litellm.llms.tool_loop.harness.transformation import ToolLoopHarnessConfig
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
class ScriptedCompletion:
|
||||
def __init__(self, responses: tuple[ModelResponse, ...]) -> None:
|
||||
self.responses: Iterator[ModelResponse] = iter(responses)
|
||||
self.calls: list[dict[str, object]] = [] # mutable-ok: captures injected completion requests
|
||||
|
||||
async def __call__(self, **kwargs: object) -> ModelResponse:
|
||||
self.calls.append(dict(kwargs))
|
||||
return next(self.responses)
|
||||
|
||||
|
||||
def model_response(
|
||||
*,
|
||||
content: str | None = None,
|
||||
tool_calls: tuple[dict[str, object], ...] = (),
|
||||
prompt_tokens: int = 0,
|
||||
completion_tokens: int = 0,
|
||||
hidden_params: dict[str, object] | None = None,
|
||||
) -> ModelResponse:
|
||||
response: Final = ModelResponse(
|
||||
model="gpt-test",
|
||||
choices=[
|
||||
{
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": content,
|
||||
"tool_calls": list(tool_calls) if tool_calls else None,
|
||||
}
|
||||
}
|
||||
],
|
||||
usage={
|
||||
"prompt_tokens": prompt_tokens,
|
||||
"completion_tokens": completion_tokens,
|
||||
"total_tokens": prompt_tokens + completion_tokens,
|
||||
},
|
||||
)
|
||||
if hidden_params is not None:
|
||||
response._hidden_params = hidden_params
|
||||
return response
|
||||
|
||||
|
||||
def function_call(name: str, arguments: str, call_id: str = "call-1") -> dict[str, object]:
|
||||
return {
|
||||
"id": call_id,
|
||||
"type": "function",
|
||||
"function": {"name": name, "arguments": arguments},
|
||||
}
|
||||
|
||||
|
||||
def make_context(
|
||||
tmp_path: Path,
|
||||
*,
|
||||
model: str | None = "gpt-4o-mini",
|
||||
gateway: GatewayTarget | None = None,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
instructions: str | None = None,
|
||||
tools: tuple[Callable[..., object], ...] = (),
|
||||
permissions: Literal["ask", "full"] = "full",
|
||||
output: type[BaseModel] | None = None,
|
||||
metadata: Mapping[str, object] | None = None,
|
||||
options: ToolLoopOptions | None = None,
|
||||
) -> SessionContext:
|
||||
return SessionContext(
|
||||
harness=Harness.TOOL_LOOP,
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
session_id="tool-loop-session",
|
||||
model=model,
|
||||
gateway=gateway,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
instructions=instructions,
|
||||
tools=tools,
|
||||
permissions=permissions,
|
||||
output=output,
|
||||
metadata={} if metadata is None else metadata,
|
||||
options=options,
|
||||
)
|
||||
|
||||
|
||||
def make_handler(completion: ScriptedCompletion) -> ToolLoopHandler:
|
||||
return ToolLoopHandler(ToolLoopHarnessConfig(), acompletion=completion)
|
||||
|
||||
|
||||
async def run_turn(
|
||||
handler: ToolLoopHandler,
|
||||
ctx: SessionContext,
|
||||
prompt: str,
|
||||
allow: bool = True,
|
||||
) -> tuple[Event, ...]:
|
||||
events: list[Event] = [] # mutable-ok: gathers this async turn for assertions
|
||||
async for event in handler.turn(ctx, prompt):
|
||||
events.append(event)
|
||||
if isinstance(event, Approval):
|
||||
event.allow() if allow else event.deny("not approved")
|
||||
return tuple(events)
|
||||
|
||||
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
|
||||
def duplicate_add() -> Callable[..., object]:
|
||||
def add(a: int, b: int) -> int:
|
||||
return a + b
|
||||
|
||||
return add
|
||||
|
||||
|
||||
async def test_tool_round_trip_appends_assistant_and_tool_messages(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
|
||||
model_response(content="The sum is 5"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, instructions="Use tools when needed", tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add two and three")
|
||||
|
||||
assert events == (
|
||||
ToolCall(
|
||||
id="call-1",
|
||||
name="add",
|
||||
native_name="add",
|
||||
input={"a": 2, "b": 3},
|
||||
builtin=False,
|
||||
),
|
||||
ToolResult(id="call-1", output="5", is_error=False),
|
||||
Text(delta="The sum is 5"),
|
||||
)
|
||||
assert ctx.final_text == "The sum is 5"
|
||||
assert completion.calls[1]["messages"] == [
|
||||
{"role": "system", "content": "Use tools when needed"},
|
||||
{"role": "user", "content": "Add two and three"},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "call-1",
|
||||
"type": "function",
|
||||
"function": {"name": "add", "arguments": '{"a": 2, "b": 3}'},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"role": "tool", "tool_call_id": "call-1", "content": "5"},
|
||||
]
|
||||
|
||||
|
||||
async def test_multiple_tool_calls_run_in_order(tmp_path: Path) -> None:
|
||||
values: list[int] = [] # mutable-ok: records the order of calls from the injected model response
|
||||
|
||||
def record(value: int) -> int:
|
||||
values.append(value)
|
||||
return value
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(
|
||||
tool_calls=(
|
||||
function_call("record", '{"value": 1}', "call-1"),
|
||||
function_call("record", '{"value": 2}', "call-2"),
|
||||
)
|
||||
),
|
||||
model_response(content="Recorded"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(record,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Record both")
|
||||
|
||||
assert values == [1, 2]
|
||||
assert tuple(event for event in events if isinstance(event, ToolResult)) == (
|
||||
ToolResult(id="call-1", output="1", is_error=False),
|
||||
ToolResult(id="call-2", output="2", is_error=False),
|
||||
)
|
||||
|
||||
|
||||
async def test_tool_exception_is_returned_to_model_and_loop_continues(tmp_path: Path) -> None:
|
||||
def fail() -> str:
|
||||
raise RuntimeError("tool failed")
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("fail", "{}"),)),
|
||||
model_response(content="Recovered"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(fail,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Run fail")
|
||||
|
||||
assert ToolResult(id="call-1", output="RuntimeError: tool failed", is_error=True) in events
|
||||
assert completion.calls[1]["messages"][-1] == {
|
||||
"role": "tool",
|
||||
"tool_call_id": "call-1",
|
||||
"content": "RuntimeError: tool failed",
|
||||
}
|
||||
assert ctx.final_text == "Recovered"
|
||||
|
||||
|
||||
async def test_unknown_tool_is_returned_and_tools_are_omitted_when_empty(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("missing", "{}"),)),
|
||||
model_response(content="Unknown tool handled"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path)
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Call a missing tool")
|
||||
|
||||
assert "tools" not in completion.calls[0]
|
||||
assert (
|
||||
ToolResult(
|
||||
id="call-1",
|
||||
output="ValueError: unknown tool 'missing'",
|
||||
is_error=True,
|
||||
)
|
||||
in events
|
||||
)
|
||||
assert ctx.final_text == "Unknown tool handled"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"arguments",
|
||||
['{"a": 2}', '{"a": 2, "b": 3, "extra": 4}'],
|
||||
)
|
||||
async def test_invalid_tool_arguments_return_validation_error(
|
||||
tmp_path: Path,
|
||||
arguments: str,
|
||||
) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", arguments),)),
|
||||
model_response(content="Arguments were invalid"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add")
|
||||
result: Final = next(event for event in events if isinstance(event, ToolResult))
|
||||
|
||||
assert result.is_error
|
||||
assert result.output.startswith("ValidationError:")
|
||||
assert completion.calls[1]["messages"][-1]["content"] == result.output
|
||||
|
||||
|
||||
async def test_malformed_tool_arguments_return_error_without_calling_tool(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", "{"),)),
|
||||
model_response(content="Malformed arguments"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add")
|
||||
result: Final = next(event for event in events if isinstance(event, ToolResult))
|
||||
|
||||
assert result.is_error
|
||||
assert result.output.startswith("JSONDecodeError:")
|
||||
assert completion.calls[1]["messages"][-1]["content"] == result.output
|
||||
|
||||
|
||||
async def test_ask_permission_denial_skips_tool_and_returns_reason(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
|
||||
model_response(content="Denied"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask")
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add", allow=False)
|
||||
|
||||
assert any(isinstance(event, Approval) for event in events)
|
||||
assert ToolResult(id="call-1", output="denied: not approved", is_error=True) in events
|
||||
assert completion.calls[1]["messages"][-1]["content"] == "denied: not approved"
|
||||
|
||||
|
||||
async def test_ask_permission_allow_runs_tool(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("add", '{"a": 2, "b": 3}'),)),
|
||||
model_response(content="Allowed"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,), permissions="ask")
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Add")
|
||||
|
||||
assert any(isinstance(event, Approval) for event in events)
|
||||
assert ToolResult(id="call-1", output="5", is_error=False) in events
|
||||
|
||||
|
||||
async def test_async_tool_is_awaited(tmp_path: Path) -> None:
|
||||
async def multiply(a: int, b: int) -> int:
|
||||
return a * b
|
||||
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(tool_calls=(function_call("multiply", '{"a": 3, "b": 4}'),)),
|
||||
model_response(content="12"),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(multiply,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
events: Final = await run_turn(handler, ctx, "Multiply")
|
||||
|
||||
assert ToolResult(id="call-1", output="12", is_error=False) in events
|
||||
|
||||
|
||||
class Answer(BaseModel):
|
||||
value: int
|
||||
|
||||
|
||||
async def test_structured_output_is_forwarded_and_retained(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion((model_response(content='{"value": 7}'),))
|
||||
ctx: Final = make_context(tmp_path, output=Answer)
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Return a value")
|
||||
|
||||
assert completion.calls[0]["response_format"] is Answer
|
||||
assert ctx.output_json == '{"value": 7}'
|
||||
assert ctx.final_text == '{"value": 7}'
|
||||
|
||||
|
||||
async def test_gateway_routing_uses_proxy_model_and_tool_loop_tag(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion((model_response(content="done"),))
|
||||
gateway: Final = GatewayTarget(api_base="http://gateway", api_key="sk-virtual")
|
||||
ctx: Final = make_context(tmp_path, gateway=gateway, metadata={"team": "test"})
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Hi")
|
||||
|
||||
assert completion.calls[0]["model"] == "litellm_proxy/gpt-4o-mini"
|
||||
assert completion.calls[0]["api_base"] == "http://gateway"
|
||||
assert completion.calls[0]["api_key"] == "sk-virtual"
|
||||
headers: Final = completion.calls[0]["extra_headers"]
|
||||
assert isinstance(headers, dict)
|
||||
assert headers["x-litellm-tags"] == "harness,tool_loop"
|
||||
assert '"team": "test"' in headers["x-litellm-spend-logs-metadata"]
|
||||
|
||||
|
||||
async def test_usage_and_cost_accumulate_across_model_calls(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion(
|
||||
(
|
||||
model_response(
|
||||
tool_calls=(function_call("add", '{"a": 2, "b": 3}'),),
|
||||
prompt_tokens=10,
|
||||
completion_tokens=4,
|
||||
hidden_params={
|
||||
"additional_headers": {"llm_provider-x-litellm-response-cost": "0.4"},
|
||||
"response_cost": 0.1,
|
||||
},
|
||||
),
|
||||
model_response(
|
||||
content="done",
|
||||
prompt_tokens=20,
|
||||
completion_tokens=5,
|
||||
hidden_params={"response_cost": 0.2},
|
||||
),
|
||||
)
|
||||
)
|
||||
ctx: Final = make_context(tmp_path, tools=(add,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "Add")
|
||||
|
||||
assert ctx.calls == 2
|
||||
assert ctx.input_tokens == 30
|
||||
assert ctx.output_tokens == 9
|
||||
assert ctx.cost == pytest.approx(0.6)
|
||||
|
||||
|
||||
async def test_history_survives_stop_and_start(tmp_path: Path) -> None:
|
||||
completion: Final = ScriptedCompletion((model_response(content="first"), model_response(content="second")))
|
||||
ctx: Final = make_context(tmp_path, instructions="Keep answers concise")
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
await run_turn(handler, ctx, "first prompt")
|
||||
await handler.stop(ctx)
|
||||
await handler.start(ctx)
|
||||
|
||||
await run_turn(handler, ctx, "second prompt")
|
||||
|
||||
assert completion.calls[1]["messages"] == [
|
||||
{"role": "system", "content": "Keep answers concise"},
|
||||
{"role": "user", "content": "first prompt"},
|
||||
{"role": "assistant", "content": "first"},
|
||||
{"role": "user", "content": "second prompt"},
|
||||
]
|
||||
history: Final = await handler.history(ctx)
|
||||
history[0]["content"] = "changed"
|
||||
assert (await handler.history(ctx))[0]["content"] == "Keep answers concise"
|
||||
|
||||
|
||||
async def test_duplicate_tool_names_are_rejected(tmp_path: Path) -> None:
|
||||
ctx: Final = make_context(tmp_path, tools=(add, duplicate_add()))
|
||||
handler: Final = make_handler(ScriptedCompletion(()))
|
||||
|
||||
with pytest.raises(ValueError, match="tool names must be unique"):
|
||||
await handler.start(ctx)
|
||||
|
||||
|
||||
async def test_model_call_limit_raises_harness_turn_error(tmp_path: Path) -> None:
|
||||
repeating_response: Final = model_response(tool_calls=(function_call("ping", "{}"),))
|
||||
completion: Final = ScriptedCompletion((repeating_response,) * 100)
|
||||
|
||||
def ping() -> str:
|
||||
return "pong"
|
||||
|
||||
ctx: Final = make_context(tmp_path, tools=(ping,))
|
||||
handler: Final = make_handler(completion)
|
||||
await handler.start(ctx)
|
||||
|
||||
with pytest.raises(HarnessTurnError, match="exceeded 100 model calls"):
|
||||
await run_turn(handler, ctx, "Ping repeatedly")
|
||||
|
||||
|
||||
async def test_public_aagent_uses_tool_loop_with_mock_response(tmp_path: Path) -> None:
|
||||
result: Final = await litellm.aagent(
|
||||
Harness.TOOL_LOOP,
|
||||
"Say done",
|
||||
sandbox=sandbox.local(tmp_path),
|
||||
model="gpt-4o-mini",
|
||||
options=ToolLoopOptions(completion_kwargs={"mock_response": "done"}),
|
||||
)
|
||||
|
||||
assert result.text == "done"
|
||||
assert result.stop_reason == "done"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Test health check helper functions"""
|
||||
|
||||
import json
|
||||
import socket
|
||||
import struct
|
||||
import zlib
|
||||
|
|
@ -8,6 +9,7 @@ from typing import Final
|
|||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME
|
||||
|
|
@ -548,3 +550,100 @@ def test_ocr_health_check_document_raises_without_the_extension():
|
|||
_ocr_health_check_document(model="mistral/mistral-ocr-latest", custom_llm_provider="mistral")
|
||||
finally:
|
||||
NATIVE_OCR_HEALTH_CHECK_DOCUMENT.reset()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "upstream_url"),
|
||||
(
|
||||
("perplexity/pplx-decider-v1-27b", "https://api.perplexity.ai/v1/decisions"),
|
||||
("cloudflare/clef", "https://api.cloudflare.com/client/v4/accounts/acct-1/ai/run/@cf/cloudflare/clef"),
|
||||
),
|
||||
)
|
||||
async def test_ahealth_check_probes_evaluation_models_through_the_decisions_api(
|
||||
model: str,
|
||||
upstream_url: str,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setenv("CLOUDFLARE_ACCOUNT_ID", "acct-1")
|
||||
monkeypatch.delenv("CLOUDFLARE_API_BASE", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post(upstream_url).respond(
|
||||
json={
|
||||
"model": model,
|
||||
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
|
||||
"usage": {"input_tokens": 12, "output_tokens": 1},
|
||||
}
|
||||
)
|
||||
|
||||
result: Final = await ahealth_check({"model": model, "api_key": "sk-test"}, mode=None)
|
||||
|
||||
assert "error" not in result, result
|
||||
assert upstream.called
|
||||
sent: Final = json.loads(upstream.calls[0].request.content)
|
||||
assert sent["state"] == "health check"
|
||||
assert sent["questions"]["reachable"]["type"] == "noul"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_evaluation_uses_configured_probe_state_and_questions(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("https://api.perplexity.ai/v1/decisions").respond(
|
||||
json={
|
||||
"model": "perplexity/pplx-decider-v1-27b",
|
||||
"answers": {"ok": {"type": "noul", "noul": 1.0}},
|
||||
"usage": {"input_tokens": 12, "output_tokens": 1},
|
||||
}
|
||||
)
|
||||
|
||||
result: Final = await ahealth_check(
|
||||
{
|
||||
"model": "perplexity/pplx-decider-v1-27b",
|
||||
"api_key": "sk-test",
|
||||
"state": "custom probe",
|
||||
"questions": {"ok": {"type": "noul", "instructions": "Is it ok?"}},
|
||||
},
|
||||
mode=None,
|
||||
)
|
||||
|
||||
assert "error" not in result, result
|
||||
assert upstream.called
|
||||
sent: Final = json.loads(upstream.calls[0].request.content)
|
||||
assert sent["state"] == "custom probe"
|
||||
assert set(sent["questions"]) == {"ok"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_probes_strands_through_decisions_without_mode(
|
||||
local_model_cost_map: None,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
respx_mock: respx.MockRouter,
|
||||
) -> None:
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_KEY", raising=False)
|
||||
monkeypatch.delenv("STRANDS_DECIDER_API_BASE", raising=False)
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
upstream: Final = respx_mock.post("http://strands.local:8080/v1/systemone").respond(
|
||||
json={
|
||||
"model": "strands-decider-2B-hobson-v19",
|
||||
"answers": {"reachable": {"type": "noul", "noul": 1.0}},
|
||||
"usage": {"input_tokens": 12, "output_tokens": 1},
|
||||
}
|
||||
)
|
||||
|
||||
result: Final = await ahealth_check(
|
||||
{
|
||||
"model": "strands_decider/strands-decider-2B-hobson-v19",
|
||||
"api_base": "http://strands.local:8080",
|
||||
},
|
||||
mode=None,
|
||||
)
|
||||
|
||||
assert "error" not in result, result
|
||||
assert upstream.called
|
||||
assert "authorization" not in upstream.calls[0].request.headers
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue