chore: merge upstream main into reasoning history

This commit is contained in:
jibanez-staticduo 2026-10-01 10:52:35 +02:00
commit 9f6bfb26ac
No known key found for this signature in database
198 changed files with 8672 additions and 957 deletions

View file

@ -36,11 +36,17 @@ jobs:
- name: Build Lens worker
run: docker build -f deploy/lens/Dockerfile -t lens-worker:${{ github.sha }} .
- name: Verify standalone imports with a read-only filesystem
run: >-
docker run --rm --network none --read-only --cap-drop ALL
--security-opt no-new-privileges --entrypoint python
lens-worker:${{ github.sha }}
-c 'import os; import engine.worker; assert os.getuid() == 65532'
run: |
docker run --rm --network none --read-only --cap-drop ALL \
--security-opt no-new-privileges --entrypoint python \
lens-worker:${{ github.sha }} -c '
import os
import engine.worker
from engine.trace_store import trace_store
assert os.getuid() == 65532
with trace_store() as store:
assert store.count() == 0
'
- name: Publish versioned Lens worker
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
env:

View file

@ -1,6 +1,7 @@
FROM python:3.12-slim
WORKDIR /app
RUN pip install --no-cache-dir httpx==0.28.1 pydantic==2.11.7
COPY litellm/proxy/engine/__init__.py litellm/proxy/engine/models.py litellm/proxy/engine/analysis.py litellm/proxy/engine/worker.py /app/engine/
COPY litellm/proxy/engine/__init__.py litellm/proxy/engine/models.py litellm/proxy/engine/trace_store.py litellm/proxy/engine/analysis.py litellm/proxy/engine/worker.py /app/engine/
VOLUME /tmp
USER 65532:65532
CMD ["python", "-m", "engine.worker"]

View file

@ -22,29 +22,35 @@ Developers can build locally with `LENS_WORKER_IMAGE=litellm-lens-worker:local d
The worker needs outbound HTTPS access to LiteLLM. It needs no inbound ports, provider keys, direct database access, or GPU. The proxy calls your selected model through its configured router; trace content reaches that model provider. Use a model with JSON output support and known token prices. One worker handles one scan at a time and can serve multiple lenses. For more throughput, start another worker with a separate credential
V1 setup, manual runs, feedback, and worker credentials are restricted to proxy administrators. Admin viewers can inspect results. Worker credentials can serve the administrator’s lenses. Revoke it in the connection dialog when retiring a worker. Redeploy the worker alongside proxy upgrades so their API versions match
V1 setup, manual runs, feedback, and worker credentials are restricted to proxy administrators. Proxy-admin viewers can inspect results. Regular user and team keys cannot access the Lens API. Worker credentials can serve the administrator’s lenses. Revoke it in the connection dialog when retiring a worker. Redeploy the worker alongside proxy upgrades so their API versions match
## Configure a lens
Choose agent runs, individual LLM requests, or both. The matching-activity preview updates as you choose an application (the recorded OpenTelemetry service.name) or, for request activity, a LiteLLM model group and add metadata conditions. It shows run names, timestamps, and trace IDs; open a run to inspect its original steps before starting analysis. Suggestions come from up to 100 recent executions and may not include every recorded attribute. You can enter other exact keys and values. Leave service and filters blank for all activity your account can access. Filters are exact key/value matches, combined with AND. Trace filters match span or resource attributes on the same span. Request filters match logged metadata, including caller metadata stored under `requester_metadata`; `tag=value` matches request tags. `swarm=research` works only if your instrumentation records that attribute
Write a few questions, give context about a successful run, choose a model, and set the monthly limit and sample size. Choose an initial history window from 1 hour to 30 days, in hours or days. Creation queues the first scan over that window. New lenses run once by default; opt into background monitoring for a custom interval from 1 minute to 7 days, entered in minutes, hours, or days. **Analyze now** checks activity since the last successful scan; **Recheck the last 24 hours** revisits recent history. The runs API accepts `lookback_hours` from 1 to 720 for other historical windows
Describe how the agent should behave and optionally add specific checks. Select the lookback window, team and metadata, then choose the percentage to review and an optional maximum. **100% with no maximum selects every matching run**. The preview pages through all matching activity and lets you select particular runs. Percentage sampling uses a stable hash order, rounds up, and applies the optional maximum after the percentage
Pausing stops future scheduled scans; cancel the active scan separately if needed. The worker polls every 10 seconds; creating a lens or clicking Analyze now queues a scan, and due schedules are queued when the worker polls. Scans for the same lens never overlap, and its next interval starts after completion. Closing the browser does not stop the worker. Configuration edits apply to the next scan. A running scan retains its settings and selected execution IDs across retries
Choose your analysis model, parallelism and monthly budget. Parallelism controls simultaneous model calls, not the number of runs selected. New lenses run once by default. Turn on monitoring to repeat the same setup at a custom interval. **Run now** uses the same saved settings immediately, including the same lookback window and sampling. Every scan recalculates the window, so overlapping windows can review the same activity again. Duplicate a lens when you want a separate investigation without changing an existing monitor
Pausing stops future scheduled scans; cancel the active scan separately if needed. The worker polls every 10 seconds; creating a lens or clicking Run now queues a scan, and due schedules are queued when the worker polls. Scans for the same lens never overlap, and its next interval starts after completion. Closing the browser does not stop the worker. Configuration edits apply to the next scan. A running scan retains its settings and selected execution IDs across retries
## Read the results
Needs attention shows issues, highest priority first. Patterns contains useful trends and successful behavior that may not need a fix. Each finding starts with a short explanation and a next step when useful. Expand the limitations for uncertainty and counterexamples. Evidence is grouped by run and collapsed until you need it; each quote opens the original step
The Runs tab lists the actual sample frozen for the latest scan. Linked-run counts on findings include cited counterexamples, so they are not failure counts. The Scans tab shows history and coverage. Existing findings retain their original wording; the shorter summaries apply to new analysis
Use the batch selector or Scans tab to reopen previous results. Each batch keeps its own findings, settings, selected runs, coverage and cost. Older batches created before snapshot support remain available through accumulated findings. The Runs tab lists the selected batch's sample and can filter per-run observations, including runs without an observed issue and runs with insufficient evidence. These observations precede the final evidence investigation. Linked-run counts on findings include cited counterexamples, so they are not failure counts
Choose **This is expected** and explain why to teach later scans about acceptable behavior. Feedback is kept with the lens and included in subsequent reviews. It does not alter historical evidence or exempt different problems
## What a scan does
The proxy selects newly received or updated executions with a two-minute settling period and a five-minute overlap. Older rows without receipt timestamps use execution end time. Overlapping scans do not increment a finding's occurrence count for the same execution ID
The proxy selects executions received or updated within the configured lookback window, with a two-minute settling period. Older rows without receipt timestamps use execution end time. Overlapping scans do not increment a finding's occurrence count for the same execution ID
A trace is spans sharing a trace ID within one team, not an automatically reconstructed conversation session. Requests are individual LLM calls. When both sources are enabled, requests correlated to a recorded span by response ID are excluded to reduce double counting
The worker screens a deterministic sample, at most the configured 1–500 executions. For each execution it reads up to 160 spans, with 8,000 characters per span section, and splits these into model calls. It consolidates observations across batches, then investigates at most 10 candidate patterns using up to five model turns each. The dashboard shows these three stages, completed work counts, and elapsed time; progress is based on the selected sample, not every eligible execution. The investigator can read more original content from the selected executions. It has no shell, browsing, code-editing, or production-action tools
The worker reviews the selected executions in parallel. It pages through their recorded spans and gives the first reviewer a catalog, task and outcome excerpts. The reviewer can read more original content to resolve uncertainties. Large catalogs and groups of observations are processed in bounded context windows, with every page available. Grouping retains supporting run IDs in code, so a pattern occurring thousands of times does not require a model to repeat thousands of IDs. Candidate investigators can page through supporting observations, other runs and original evidence
There is no fixed total run, span, candidate or investigation-turn cutoff. Repeated or empty evidence requests stop a stalled investigation. Context windows, the configured budget, available model capacity and recorded evidence still bound practical work. The dashboard reports completed work and gaps. The investigator has no shell, browsing, code-editing or production-action tools
Each model response must match a bounded JSON schema. A malformed response gets one repair attempt through the same budget controls; repeated invalid output fails the scan. Both the worker and proxy validate quoted evidence. Findings retain exact quotes and open the source trace or request. Resolve a finding after a fix, or dismiss it with a reason. A resolved finding reopens when new execution IDs support the same pattern; dismissed findings remain dismissed
@ -52,8 +58,48 @@ Coverage distinguishes eligible, sampled, reviewed, partial, and unassessable ex
## Operations and limits
PostgreSQL stores configurations, findings and the latest 50 jobs. Workers claim jobs with optimistic concurrency and a five-minute lease, renewed every 30 seconds. A disconnected job can be reclaimed up to three times. Cancellation stops subsequent work; a model call already in flight may finish and incur cost
PostgreSQL stores configurations, findings and all scan history, returned in pages of 50 jobs. Workers claim jobs with optimistic concurrency and a five-minute lease, renewed every 30 seconds. A disconnected job can be reclaimed up to three times. Cancellation stops subsequent work; a model call already in flight may finish and incur cost
Before every model call, Lens reserves a conservative amount against the monthly lens budget. Successful calls reconcile to reported cost where pricing is available. Interrupted calls retain their reservation because the provider may have charged. A scan stops when the next reservation would exceed the limit, so it can stop with some budget remaining. Lens budgets are separate from virtual-key budgets; analysis calls use the proxy router directly
V1 requires ClickHouse for both sources. It does not reconstruct sessions from unrelated trace IDs, guarantee exhaustive reviews, cache all per-execution observations across scans, or automatically fix agent code. Trace contents can change as late spans arrive, even though a job's selected IDs are fixed. Findings should be reviewed by a person before acting on them
## API access
The UI and API use the same scan lifecycle. Authenticate with a proxy administrator credential for writes, or a proxy-admin viewer credential for reads. Worker credentials are only for worker operations
```bash
curl "$LITELLM_URL/engine" -H "Authorization: Bearer $LITELLM_API_KEY" \
-H 'Content-Type: application/json' -d '{
"name": "Research quality", "model": "your-model-alias",
"context": "Answer the requested question using cited, retrieved evidence.",
"source": "traces", "lookback_hours": 24,
"sample_percent": 100, "sample_size": null, "concurrency": 8,
"enabled": true, "interval_minutes": 1440, "monthly_budget": 50
}'
curl "$LITELLM_URL/engine/$LENS_ID/runs" -X POST \
-H "Authorization: Bearer $LITELLM_API_KEY" -H 'Content-Type: application/json' -d '{}'
curl "$LITELLM_URL/engine/$LENS_ID/runs?offset=0" -H "Authorization: Bearer $LITELLM_API_KEY"
curl "$LITELLM_URL/engine/$LENS_ID/runs/$BATCH_ID" -H "Authorization: Bearer $LITELLM_API_KEY"
```
Creation queues the first batch. Posting to `/engine/{id}/runs` queues another, or returns the existing active batch. The run response contains its ID under `jobs[0].id`. Poll the batch URL for status, findings and assessments. List responses omit large result payloads; request a batch to retrieve them. Supply an optional complete `settings` object on the runs POST for a one-off override; the saved lens stays unchanged. Selection accepts `team_id`, exact `filters`, and opaque `execution_ids` returned by `/engine/preview/sample`. Preview accepts `offset` and `as_of` to keep the time window fixed while paging. Feedback uses `PATCH /engine/{id}/findings/{finding_id}` with `status` and `reason`
## Quality evaluation
Run the checked-in cases against a configured real model. Expected labels are used only for scoring, never passed to the model. Dev and held-out cases include missing outcomes, failed tools, recovery, handoffs, unsupported claims, repeated work, long evidence and prompt injection. The background option adds clean arithmetic traces to test rare-issue discovery at scale; those repeated synthetic cases do not establish accuracy on every production workload
```bash
python -m tests.proxy_behavior.lens.evaluate --api-base "$LITELLM_URL" \
--model your-model-alias --split all --background 1000 --concurrency 16 \
--output /tmp/lens-quality.json
```
Set `LITELLM_API_KEY` privately. This makes paid model calls. Inspect missed and unexpected per-run labels, final findings and coverage; do not equate a passing dataset with guaranteed detection on arbitrary traces
The worker uses temporary disk space for trace content while reviewing it, and removes those files after each review. Its Docker image supplies a writable temporary volume while keeping the application filesystem read-only
To check that accepted behavior stays accepted without hiding new problems, run the evaluator with `--dataset tests/proxy_behavior/lens/feedback_cases.json`. Reports include elapsed time, model call count, reported cost when the proxy provides it, missed checks, unexpected checks, and inconclusive candidates

View file

@ -1,6 +1,6 @@
services:
lens-worker:
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:47445afedfb6de2ae37a3a246ea1c939196bfd365436a880ab96ecf5f42b2342}
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:40fdb82113dd4474cb6e833cf28552487d87c8baf61693a1c3fc2863b7968c6a}
environment:
LITELLM_URL: ${LITELLM_URL:?Set the URL reachable from this container}
LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:?Create a worker credential in the Lens UI}

View file

@ -0,0 +1,7 @@
CREATE TABLE IF NOT EXISTS "LiteLLM_EngineRun" (
"id" TEXT NOT NULL PRIMARY KEY,
"engine_id" TEXT NOT NULL,
"created_at" TIMESTAMP(3) NOT NULL,
"data" JSONB NOT NULL
);
CREATE INDEX IF NOT EXISTS "LiteLLM_EngineRun_engine_id_created_at_idx" ON "LiteLLM_EngineRun"("engine_id", "created_at");

View file

@ -1901,6 +1901,15 @@ model LiteLLM_Engine {
data Json
}
model LiteLLM_EngineRun {
id String @id
engine_id String
created_at DateTime
data Json
@@index([engine_id, created_at])
}
model LiteLLM_EngineWorker {
id String @id
token_hash String @unique

View file

@ -1,10 +1,17 @@
WITH greatest(toInt64({offset:UInt32})-1,1) AS content_offset,
(value, budget) -> if(lengthUTF8(value) <= budget, value,
concat(substringUTF8(value, 1, intDiv(budget, 3)), '\n[... content omitted ...]\n',
substringUTF8(value, -(budget - intDiv(budget, 3))))) AS excerpt
SELECT * FROM (
SELECT SpanId AS span_id, ParentSpanId AS parent_span_id, SpanName AS name,
ObservationType AS kind,
substringUTF8(concat('Input: ',Input,'\nOutput: ',Output,'\nStatus: ',StatusCode,' ',StatusMessage),
{offset:UInt32},8000) AS content,
if({offset:UInt32}=1 AND lengthUTF8(concat('Input: ',Input,'\nOutput: ',Output,'\nStatus: ',StatusCode,' ',StatusMessage))>8000,
concat('Input: ',excerpt(Input,2000),'\nOutput: ',excerpt(Output,5000),
'\nStatus: ',StatusCode,' ',excerpt(StatusMessage,500)),
substringUTF8(concat('Input: ',Input,'\nOutput: ',Output,'\nStatus: ',StatusCode,' ',StatusMessage),
content_offset,8000)) AS content,
lengthUTF8(concat('Input: ',Input,'\nOutput: ',Output,'\nStatus: ',StatusCode,' ',StatusMessage))
>= {offset:UInt32}+8000 AS truncated
>= content_offset+8000 AS truncated
FROM otel_traces WHERE {source:String}='traces'
AND ({all_teams:UInt8}=1 OR TeamId={team:String})
AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String})
@ -15,10 +22,12 @@ SELECT * FROM (
UNION ALL
SELECT * FROM (
SELECT request_id AS span_id, '' AS parent_span_id, model AS name, 'llm' AS kind,
substringUTF8(concat('Input: ',messages,'\nOutput: ',response,'\nError: ',error_str),
{offset:UInt32},8000) AS content,
if({offset:UInt32}=1 AND lengthUTF8(concat('Input: ',messages,'\nOutput: ',response,'\nError: ',error_str))>8000,
concat('Input: ',excerpt(messages,2000),'\nOutput: ',excerpt(response,5000),'\nError: ',excerpt(error_str,500)),
substringUTF8(concat('Input: ',messages,'\nOutput: ',response,'\nError: ',error_str),
content_offset,8000)) AS content,
lengthUTF8(concat('Input: ',messages,'\nOutput: ',response,'\nError: ',error_str))
>= {offset:UInt32}+8000 AS truncated
>= content_offset+8000 AS truncated
FROM spend_logs FINAL WHERE {source:String}='requests'
AND ({all_teams:UInt8}=1 OR team_id={team:String})
AND ({key_hash:String}='' OR api_key={key_hash:String})

View file

@ -1,4 +1,12 @@
SELECT *, count() OVER () AS eligible FROM (
WITH concat(leftPad(toString(cityHash64(concat(source,team_id,trace_ref,trace_id))),20,'0'),
hex(concat(source,char(0),team_id,char(0),trace_ref,char(0),trace_id))) AS selection_key
SELECT *, selection_key FROM (
SELECT *, if({sample_cap:UInt64}=0, ceiling(eligible*{sample_percent:Float64}/100),
least(toFloat64({sample_cap:UInt64}),ceiling(eligible*{sample_percent:Float64}/100))) AS selected
FROM (
SELECT *, count() OVER () AS eligible,
row_number() OVER (ORDER BY selection_key) AS position
FROM (
SELECT 'traces' AS source, TraceId AS trace_id, TeamId AS team_id, hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) AS trace_ref,
coalesce(nullIf(argMin(ResourceAttributes['run.name'], Timestamp), ''),
argMin(SpanName, Timestamp)) AS name, toString(min(Timestamp)) AS start_time,
@ -47,4 +55,11 @@ SELECT *, count() OVER () AS eligible FROM (
AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) AND LiteLLMRequestId!=''
))
)
ORDER BY cityHash64(concat(source,team_id,trace_id)) LIMIT {limit:UInt32}
WHERE ({selected_team:String}='' OR team_id={selected_team:String})
AND (empty({execution_ids:Array(String)}) OR has({execution_ids:Array(String)},
concat(source,char(0),team_id,char(0),if(trace_ref='',trace_id,trace_ref))))
)
)
WHERE ({preview:UInt8}=1 OR position <= selected)
AND selection_key > {after:String}
ORDER BY selection_key LIMIT {limit:UInt32} OFFSET {offset:UInt64}

View file

@ -561,6 +561,13 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate(
Parameter::Strings(vec!["release".into()]),
),
("limit".into(), Parameter::Integer(10)),
("offset".into(), Parameter::Integer(0)),
("after".into(), Parameter::Text(String::new())),
("sample_percent".into(), Parameter::Text("100".into())),
("sample_cap".into(), Parameter::Integer(0)),
("preview".into(), Parameter::Integer(0)),
("selected_team".into(), Parameter::Text(String::new())),
("execution_ids".into(), Parameter::Strings(vec![])),
]);
let sample: serde_json::Value = serde_json::from_str(
&execute_read(
@ -649,6 +656,13 @@ async fn lens_request_sample_does_not_trust_caller_tags(
("filter_keys".into(), Parameter::Strings(vec![])),
("filter_values".into(), Parameter::Strings(vec![])),
("limit".into(), Parameter::Integer(10)),
("offset".into(), Parameter::Integer(0)),
("after".into(), Parameter::Text(String::new())),
("sample_percent".into(), Parameter::Text("100".into())),
("sample_cap".into(), Parameter::Integer(0)),
("preview".into(), Parameter::Integer(0)),
("selected_team".into(), Parameter::Text(String::new())),
("execution_ids".into(), Parameter::Strings(vec![])),
]);
let sample: serde_json::Value = serde_json::from_str(
&execute_read(
@ -664,3 +678,160 @@ async fn lens_request_sample_does_not_trust_caller_tags(
assert_eq!(rows[0]["trace_id"], "external");
Ok(())
}
#[rstest]
#[case::changing("100", 0, 0, 1001, 100, true)]
#[case::all("100", 0, 0, 1001, 100, false)]
#[case::percentage("10", 0, 0, 101, 100, false)]
#[case::capped("100", 25, 0, 25, 100, false)]
#[case::preview("10", 25, 1, 1001, 100, false)]
#[tokio::test]
async fn lens_selection_pages_without_losing_or_repeating_runs(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
#[case] percent: &str,
#[case] cap: i64,
#[case] preview: i64,
#[case] expected: usize,
#[case] page_size: usize,
#[case] changing: bool,
) -> TestResult {
use litellm_traces::LensQuery;
let database = database?;
ensure_schema(
&database.client,
&Connection::writer(&database.url)?,
"trace_test",
7,
14,
)
.await?;
execute_write(&database, "INSERT INTO trace_test.spend_logs (request_id,team_id,start_time,end_time) SELECT toString(number),'team',now64(3)-INTERVAL 5 MINUTE,now64(3)-INTERVAL 5 MINUTE FROM numbers(1001)").await?;
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
let end = time::OffsetDateTime::now_utc().unix_timestamp() * 1000 + 60000;
let mut seen = std::collections::BTreeSet::new();
let mut cursor = String::new();
let step = if page_size == 0 { expected } else { page_size };
for offset in (0..expected).step_by(step) {
let parameters = BTreeMap::from([
("source".into(), Parameter::Text("requests".into())),
("all_teams".into(), Parameter::Integer(0)),
("team".into(), Parameter::Text("team".into())),
("key_hash".into(), Parameter::Text(String::new())),
("start".into(), Parameter::Integer(0)),
("end".into(), Parameter::Integer(end)),
("service".into(), Parameter::Text(String::new())),
("filter_keys".into(), Parameter::Strings(vec![])),
("filter_values".into(), Parameter::Strings(vec![])),
("limit".into(), Parameter::Integer(page_size as i64)),
(
"offset".into(),
Parameter::Integer(if changing { 0 } else { offset as i64 }),
),
("after".into(), Parameter::Text(cursor.clone())),
("sample_percent".into(), Parameter::Text(percent.into())),
("sample_cap".into(), Parameter::Integer(cap)),
("preview".into(), Parameter::Integer(preview)),
("selected_team".into(), Parameter::Text(String::new())),
("execution_ids".into(), Parameter::Strings(vec![])),
]);
let body = execute_read(
&database.client,
&connection,
LensQuery::Sample.sql(),
&parameters,
)
.await?;
let json: serde_json::Value = serde_json::from_str(&body)?;
let rows = json["data"].as_array().expect("sample rows");
assert_eq!(rows.len(), step.min(expected - offset));
for row in rows {
assert_eq!(
row["eligible"],
if changing && offset > 0 { 1000 } else { 1001 }
);
assert!(seen.insert(row["trace_id"].as_str().expect("run id").to_owned()));
}
if changing {
cursor = rows.last().expect("last run")["selection_key"]
.as_str()
.expect("selection key")
.to_owned();
if offset == 0 {
let removed = rows[0]["trace_id"].as_str().expect("request id");
execute_write(&database, &format!("ALTER TABLE trace_test.spend_logs DELETE WHERE request_id='{removed}' SETTINGS mutations_sync=1")).await?;
}
}
}
assert_eq!(seen.len(), expected);
Ok(())
}
#[rstest]
#[case::short(100)]
#[case::boundary(7970)]
#[case::long(16000)]
#[tokio::test]
async fn lens_content_keeps_output_visible_after_long_input(
#[future(awt)] database: TestResult<ClickHouseDatabase>,
#[case] input_length: usize,
) -> TestResult {
use litellm_traces::LensQuery;
let database = database?;
ensure_schema(
&database.client,
&Connection::writer(&database.url)?,
"trace_test",
7,
14,
)
.await?;
insert_rows(&database, "spend_logs", vec![serde_json::from_value(serde_json::json!({
"request_id": "request", "team_id": "team", "start_time": time::OffsetDateTime::now_utc().unix_timestamp()*1000, "end_time": time::OffsetDateTime::now_utc().unix_timestamp()*1000, "messages": "x".repeat(input_length), "response": "Delivered result"
}))?]).await?;
let connection = Connection::configured(&database.url, "trace_test", "default", "")?;
let mut parameters = BTreeMap::from([
("source".into(), Parameter::Text("requests".into())),
("all_teams".into(), Parameter::Integer(0)),
("team".into(), Parameter::Text("team".into())),
("record_team".into(), Parameter::Text("team".into())),
("key_hash".into(), Parameter::Text(String::new())),
("trace_ref".into(), Parameter::Text(String::new())),
("id".into(), Parameter::Text("request".into())),
("cursor".into(), Parameter::Text(String::new())),
("offset".into(), Parameter::Integer(1)),
]);
let body = execute_read(
&database.client,
&connection,
LensQuery::Content.sql(),
&parameters,
)
.await?;
let json: serde_json::Value = serde_json::from_str(&body)?;
let text = json["data"][0]["content"].as_str().expect("content");
assert!(text.contains("Output: Delivered result"));
assert!(text.len() <= 8000);
assert_eq!(
json["data"][0]["truncated"],
u8::from(input_length + "Input: \nOutput: Delivered result\nError: ".len() > 8000)
);
let original = format!(
"Input: {}\nOutput: Delivered result\nError: ",
"x".repeat(input_length)
);
let mut recovered = String::new();
for offset in (2..original.len() + 2).step_by(8000) {
parameters.insert("offset".into(), Parameter::Integer(offset as i64));
let body = execute_read(
&database.client,
&connection,
LensQuery::Content.sql(),
&parameters,
)
.await?;
let page: serde_json::Value = serde_json::from_str(&body)?;
recovered.push_str(page["data"][0]["content"].as_str().expect("content"));
}
assert_eq!(recovered, original);
Ok(())
}

View file

@ -8,7 +8,6 @@
# Thank you users! We ❤️ you! - Krrish & Ishaan
import ast
import asyncio
import hashlib
import json
import logging
@ -19,7 +18,7 @@ from enum import Enum
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
import litellm
from litellm._logging import verbose_logger
@ -31,11 +30,10 @@ from litellm.types.utils import EmbeddingResponse, is_litellm_owned_kwarg
from .azure_blob_cache import AzureBlobCache
from .base_cache import BaseCache
from .disk_cache import DiskCache
from .dual_cache import DualCache
from .dual_cache import DualCache # noqa: F401 # re-exported, callers import DualCache from litellm.caching.caching
from .gcs_cache import GCSCache
from .in_memory_cache import InMemoryCache
from .qdrant_semantic_cache import QdrantSemanticCache
from .redis_batch import active_post_call_redis_batch
from .redis_cache import RedisCache, log_redis_failure
from .redis_cluster_cache import RedisClusterCache
from .redis_semantic_cache import RedisSemanticCache
@ -61,6 +59,9 @@ def _native_response(result: object) -> object:
return result
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def print_verbose(print_statement):
try:
verbose_logger.debug(print_statement)
@ -70,15 +71,6 @@ def print_verbose(print_statement):
pass
def _ttl_seconds(raw: object) -> int | None:
if not isinstance(raw, (int, float, str)):
return None
try:
return int(raw)
except ValueError:
return None
class CacheMode(str, Enum):
default_on = "default_on"
default_off = "default_off"
@ -401,7 +393,9 @@ class Cache:
param_value = kwargs[param]
cache_key += f"{param}: {param_value}"
nested_litellm_params: Final = kwargs.get("litellm_params") or MappingProxyType({})
nested_litellm_params: Final = _LITELLM_PARAMS_ADAPTER.validate_python(
kwargs.get("litellm_params") or MappingProxyType({})
)
forward_reasoning_content: Final = kwargs.get(
"forward_reasoning_content", nested_litellm_params.get("forward_reasoning_content")
)
@ -782,8 +776,6 @@ class Cache:
await self.batch_cache_write(result, **kwargs)
else:
cache_key, cached_data, kwargs = self._add_cache_logic(result=result, **kwargs)
if await self._defer_set_to_post_call_batch(cache_key, cached_data, kwargs, dynamic_cache_object):
return
if dynamic_cache_object is not None:
await dynamic_cache_object.async_set_cache(cache_key, cached_data, **kwargs)
else:
@ -791,39 +783,6 @@ class Cache:
except Exception as e:
self._log_add_cache_failure(e)
async def _defer_set_to_post_call_batch(
self,
cache_key: str,
cached_data: object,
kwargs: Mapping[str, object],
dynamic_cache_object: BaseCache | None,
) -> bool:
"""A plain SET on the Redis response cache rides the request's post-call pipeline with the counters,
instead of its own round trip. Anything with SET options keeps the direct path."""
if kwargs.get("nx"):
return False
ttl: Final = _ttl_seconds(kwargs.get("ttl"))
if isinstance(dynamic_cache_object, DualCache):
deferred: Final = await dynamic_cache_object.async_set_cache_post_call(cache_key, cached_data, ttl)
if deferred is None:
return False
deferred.on_settled(self._log_deferred_add_cache_failure)
return True
if dynamic_cache_object is not None or not isinstance(self.cache, RedisCache):
return False
batch: Final = active_post_call_redis_batch(self.cache)
if batch is None:
return False
batch.set(cache_key, cached_data, ttl).on_settled(self._log_deferred_add_cache_failure)
return True
def _log_deferred_add_cache_failure(self, future: asyncio.Future[None]) -> None:
if future.cancelled():
return
failure: Final = future.exception()
if isinstance(failure, Exception):
self._log_add_cache_failure(failure)
def _convert_to_cached_embedding(
self,
embedding_response: Any,

View file

@ -525,12 +525,6 @@ class DualCache(BaseCache):
batch: Final = None if self.redis_cache is None else active_request_redis_batch(self.redis_cache)
return None if batch is None else await self._set_on_batch(batch, key, value, ttl)
async def async_set_cache_post_call(self, key: str, value: object, ttl: float | None) -> BatchResult[None] | None:
"""Memory now, the Redis SET on the request's post-call pipeline; None when no pipeline is open, so the
caller takes its direct path."""
batch: Final = None if self.redis_cache is None else active_post_call_redis_batch(self.redis_cache)
return None if batch is None else await self._set_on_batch(batch, key, value, ttl)
async def async_delete_cache_pre_call(self, key: str) -> BatchResult[None] | None:
"""Memory now, the Redis DEL on the request's pipeline; None when no pipeline is open, so the caller
takes its direct path."""

View file

@ -961,6 +961,7 @@ openai_compatible_endpoints: Final[list] = [
"https://api.meta.ai/v1",
"https://api.sailresearch.com/v1",
"https://api.cognition.ai/v1",
"https://api.cortecs.ai/v1",
"https://api.scx.ai/v1",
"https://api.prisminference.com/v1",
"https://gigachat.devices.sberbank.ru/api/v1",
@ -1035,6 +1036,7 @@ openai_compatible_providers: Final[list] = [
"darkbloom",
"meta", # Meta Model API (Muse Spark) - JSON-configured provider
"cognition",
"cortecs",
"scx-ai",
"prism",
"sail",

View file

@ -8,6 +8,8 @@ from datetime import datetime
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, get_args
import httpx
from litellm._logging import verbose_logger
from litellm.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
@ -176,6 +178,8 @@ class CustomGuardrail(CustomLogger):
records_own_guardrail_information: ClassVar[bool] = False
timeout: float | httpx.Timeout | None = None
def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks
super().__init_subclass__(**kwargs)
own_apply_guardrail: Final[object] = cls.__dict__.get("apply_guardrail")
@ -201,6 +205,7 @@ class CustomGuardrail(CustomLogger):
run_in_parallel: bool = False,
scan_raw_request: bool = False,
only_scan_new_messages: bool = False,
timeout: float | None = None,
**kwargs,
):
"""
@ -229,6 +234,8 @@ class CustomGuardrail(CustomLogger):
guardrails: any data this guardrail returns is discarded, matching run_in_parallel's
contract, since applying its mutations on top of a stale snapshot would silently
undo whatever later guardrails already did to the live request.
timeout: Per-request timeout in seconds for the guardrail provider's API call. When
None, the guardrail keeps whatever default its HTTP handler or SDK already uses.
"""
self.guardrail_name = guardrail_name
self.supported_event_hooks = supported_event_hooks
@ -246,6 +253,8 @@ class CustomGuardrail(CustomLogger):
self.run_in_parallel: bool = run_in_parallel
self.scan_raw_request: bool = scan_raw_request
self.only_scan_new_messages: bool = only_scan_new_messages
if timeout is not None:
self.timeout = timeout
if supported_event_hooks:
## validate event_hook is in supported_event_hooks

View file

@ -1120,6 +1120,7 @@ class RubrikLogger(CustomGuardrail, CustomBatchLogger):
endpoint,
json=dict(payload),
headers=dict(self._headers),
timeout=self.timeout,
)
http_response.raise_for_status()
result: Final[_ModerationResponse | None] = http_response.json()

View file

@ -1,5 +1,9 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from pydantic import TypeAdapter
import litellm
from litellm import verbose_logger
@ -9,6 +13,8 @@ from ...litellm_core_utils.get_llm_provider_logic import (
)
from ...types.router import LiteLLM_Params
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
def _api_base_without_login(provider: str) -> str | None:
if provider == "github_copilot":
@ -50,9 +56,11 @@ def get_api_base(model: str, optional_params: dict | LiteLLM_Params) -> str | No
if isinstance(optional_params, LiteLLM_Params):
_optional_params = optional_params
elif "model" in optional_params:
_optional_params = LiteLLM_Params(**optional_params)
_optional_params = LiteLLM_Params.model_validate(optional_params)
else: # prevent needing to copy and pop the dict
_optional_params = LiteLLM_Params(model=model, **optional_params) # convert to pydantic object
_optional_params = LiteLLM_Params.model_validate(
_LITELLM_PARAMS_ADAPTER.validate_python(MappingProxyType({"model": model, **optional_params}))
) # convert to pydantic object
except Exception:
return None
# get llm provider

View file

@ -10,6 +10,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast, overload
from urllib.parse import urlparse
import httpx
from pydantic import TypeAdapter
import litellm
from litellm.constants import OPENAI_SYSTEM_MESSAGES_FIRST_PROVIDERS
@ -75,6 +76,7 @@ else:
_NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
_LITELLM_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
@ -491,16 +493,17 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
Returns:
dict: The transformed request. Sent as the body of the API call.
"""
transport_params: Final = _LITELLM_PARAMS_ADAPTER.validate_python(litellm_params)
request_messages: Final = (
normalize_reasoning_content(messages)
if litellm_params.get("custom_llm_provider") == "openai"
if transport_params.get("custom_llm_provider") == "openai"
and should_normalize_reasoning_content(
litellm_params.get("reasoning_content_field"), model=model, provider="openai"
transport_params.get("reasoning_content_field"), model=model, provider="openai"
)
else messages
)
messages = self._transform_messages(
messages=self._prompt_cache_ordered_messages(request_messages, litellm_params), model=model
messages=self._prompt_cache_ordered_messages(request_messages, transport_params), model=model
)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
@ -530,16 +533,17 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
litellm_params: dict,
headers: dict,
) -> dict:
transport_params: Final = _LITELLM_PARAMS_ADAPTER.validate_python(litellm_params)
request_messages: Final = (
normalize_reasoning_content(messages)
if litellm_params.get("custom_llm_provider") == "openai"
if transport_params.get("custom_llm_provider") == "openai"
and should_normalize_reasoning_content(
litellm_params.get("reasoning_content_field"), model=model, provider="openai"
transport_params.get("reasoning_content_field"), model=model, provider="openai"
)
else messages
)
transformed_messages = await self._transform_messages(
messages=self._prompt_cache_ordered_messages(request_messages, litellm_params), model=model, is_async=True
messages=self._prompt_cache_ordered_messages(request_messages, transport_params), model=model, is_async=True
)
if not self._should_preserve_cache_control_for_endpoint(
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")

View file

@ -180,6 +180,12 @@
"api_key_env": "COGNITION_API_KEY",
"api_base_env": "COGNITION_API_BASE"
},
"cortecs": {
"base_url": "https://api.cortecs.ai/v1",
"api_key_env": "CORTECS_API_KEY",
"api_base_env": "CORTECS_API_BASE",
"supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/messages"]
},
"pinstripes": {
"base_url": "https://pinstripes.io/v1",
"api_key_env": "PINSTRIPES_API_KEY",

View file

@ -14870,6 +14870,7 @@
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"cache_read_input_token_cost_above_200k_tokens_batches": 3e-07,
"cache_read_input_token_cost_batches": 1.5e-07,
"deprecation_date": "2026-11-30",
"input_cost_per_token_above_200k_tokens_batches": 3e-06,
"input_cost_per_token_batches": 1.5e-06,
"litellm_provider": "anthropic",
@ -14912,6 +14913,7 @@
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"cache_read_input_token_cost_above_200k_tokens_batches": 3e-07,
"cache_read_input_token_cost_batches": 1.5e-07,
"deprecation_date": "2026-11-30",
"input_cost_per_token_above_200k_tokens_batches": 3e-06,
"input_cost_per_token_batches": 1.5e-06,
"litellm_provider": "anthropic",
@ -64530,13 +64532,16 @@
},
"fireworks_ai/accounts/fireworks/models/inkling": {
"cache_read_input_token_cost": 1.7e-07,
"cache_read_input_token_cost_priority": 1.7e-07,
"input_cost_per_token": 1e-06,
"input_cost_per_token_priority": 1e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_tokens": 1048576,
"mode": "chat",
"output_cost_per_token": 4.05e-06,
"source": "https://fireworks.ai/models/fireworks/inkling",
"output_cost_per_token_priority": 4.05e-06,
"source": "https://api.fireworks.ai/v1/serverless/models?format=nested",
"supports_function_calling": true,
"supports_response_schema": true,
"supports_tool_choice": true,
@ -77203,8 +77208,8 @@
"input_cost_per_token": 3e-07,
"litellm_provider": "baseten",
"max_input_tokens": 1048576,
"max_output_tokens": 32768,
"max_tokens": 32768,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 1.2e-06,
"source": "https://inference.baseten.co/v1/models",

View file

@ -618,6 +618,24 @@
"interactions": true
}
},
"cortecs": {
"display_name": "Cortecs (`cortecs`)",
"url": "https://docs.litellm.ai/docs/providers/cortecs",
"endpoints": {
"chat_completions": true,
"messages": true,
"responses": true,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false,
"interactions": false
}
},
"custom": {
"display_name": "Custom (`custom`)",
"url": "https://docs.litellm.ai/docs/providers/custom_llm_server",

View file

@ -524,6 +524,7 @@ class LiteLLMRoutes(enum.Enum):
"/engine",
"/engine/{engine_id}",
"/engine/{engine_id}/runs",
"/engine/{engine_id}/runs/{job_id}",
"/engine/{engine_id}/executions/{execution_id}",
"/engine/{engine_id}/cancel",
"/engine/{engine_id}/findings/{finding_id}",

View file

@ -78,6 +78,13 @@ def _build_dual_source_keys() -> Mapping[tuple[Section, str], KeyRule]:
DUAL_SOURCE_KEYS: Final[Mapping[tuple[Section, str], KeyRule]] = _build_dual_source_keys()
RESOURCE_LIST_KEYS: Final[frozenset[tuple[Section, str]]] = frozenset({("general_settings", "pass_through_endpoints")})
def is_resource_list(section: Section, key: str) -> bool:
return (section, key) in RESOURCE_LIST_KEYS
def rule_for(section: Section, key: str) -> KeyRule:
return DUAL_SOURCE_KEYS.get((section, key), DUAL_SOURCE_KEYS[(section, "*")])

View file

@ -13,6 +13,7 @@ from litellm.proxy.config_resolvers.settings_rules import (
Resolved,
Section,
SettingValue,
is_resource_list,
resolve,
rule_for,
)
@ -49,8 +50,13 @@ class SettingsStore(MutableMapping[str, JsonValue]):
self._deleted_runtime_keys: frozenset[str] = frozenset()
def load_yaml(self, mapping: Mapping[str, JsonValue]) -> None:
self._yaml_values = MappingProxyType(dict(mapping))
self._clear_runtime()
self._yaml_values = MappingProxyType(
{key: value for key, value in mapping.items() if not is_resource_list(self._section, key)}
)
self._runtime_values = MappingProxyType(
{key: value for key, value in self._runtime_values.items() if is_resource_list(self._section, key)}
)
self._deleted_runtime_keys = frozenset()
def config_value(self, key: str) -> JsonValue:
return self._yaml_values.get(key)
@ -136,10 +142,6 @@ class SettingsStore(MutableMapping[str, JsonValue]):
def __bool__(self) -> bool:
return any(True for _ in self)
def _clear_runtime(self) -> None:
self._runtime_values = _EMPTY_VALUES
self._deleted_runtime_keys = frozenset()
def _clear_runtime_keys(self, keys: frozenset[str]) -> None:
stale: Final = frozenset(key for key in keys if not self.owned_by_config(key))
if not stale:
@ -160,6 +162,9 @@ class SettingsStore(MutableMapping[str, JsonValue]):
)
)
def db_value(self, key: str) -> SettingValue:
return self._db_value(key) if is_resource_list(self._section, key) else ABSENT
def _db_value(self, key: str) -> SettingValue:
rule: Final = rule_for(self._section, key)
return self._database_rows.get(rule.db_row, _EMPTY_VALUES).get(key, ABSENT)

File diff suppressed because it is too large Load diff

View file

@ -8,7 +8,7 @@ from uuid import uuid4
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, Field, TypeAdapter
from pydantic import AwareDatetime, BaseModel, Field, TypeAdapter
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -35,7 +35,15 @@ from litellm.proxy.engine.models import (
)
from litellm.proxy.engine.repository import EngineRepository, WriterDatabase
from litellm.proxy.engine.sources import SourceReader, parse_execution
from litellm.proxy.engine.state import can_access, claim_job, current_job, merge_finding, queue_job, replace_job
from litellm.proxy.engine.state import (
can_access,
claim_job,
current_job,
merge_finding,
queue_job,
replace_job,
snapshot_finding,
)
router: Final = APIRouter(prefix="/engine", tags=["Lens"]) # mutable-ok: FastAPI requires list
_bearer: Final = HTTPBearer()
@ -61,11 +69,7 @@ def user_scope(auth: UserAPIKeyAuth, write: bool = False) -> Scope:
raise HTTPException(403, "Only proxy admins can configure or run Lens")
if auth.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY):
return Scope(all_teams=True)
if auth.team_id:
return Scope(team_id=auth.team_id)
if auth.token:
return Scope(api_key_hash=auth.token)
raise HTTPException(403, "A team or API key is required")
raise HTTPException(403, "Lens requires proxy administrator access")
async def get_engine(engine_id: str, scope: Scope) -> Engine:
@ -106,9 +110,20 @@ def required(engine: Engine | None) -> Engine:
return engine
def validate_selection(settings: EngineSettings) -> None:
for identity in settings.execution_ids:
try:
source, _, _, _ = parse_execution(identity)
if source not in ("traces", "requests"):
raise ValueError("Unsupported source")
except ValueError:
raise HTTPException(422, "Choose execution IDs returned by the activity preview")
def validate_model(settings: EngineSettings, auth: UserAPIKeyAuth) -> None:
from litellm.proxy.proxy_server import llm_router
validate_selection(settings)
if llm_router is None or settings.model not in llm_router.get_model_names(team_id=auth.team_id):
raise HTTPException(400, "Choose a model configured on this LiteLLM instance")
allowed_models: Final = TypeAdapter(tuple[str, ...]).validate_python(auth.model_dump().get("models") or ())
@ -171,9 +186,36 @@ async def update_engine(engine_id: str, settings: EngineSettings, auth: Auth) ->
@router.post("/{engine_id}/runs", response_model=Engine)
async def run_engine(engine_id: str, body: RunRequest, auth: Auth) -> Engine:
await get_engine(engine_id, user_scope(auth, write=True))
if body.settings is not None:
validate_model(body.settings, auth)
now: Final = datetime.now(timezone.utc)
job_id: Final = str(uuid4())
return required(await repository().update(engine_id, lambda e: queue_job(e, now, job_id, body.lookback_hours)))
return required(
await repository().update(engine_id, lambda e: queue_job(e, now, job_id, body.lookback_hours, body.settings))
)
@router.get("/{engine_id}", response_model=Engine)
async def read_engine(engine_id: str, auth: Auth) -> Engine:
return await get_engine(engine_id, user_scope(auth))
@router.get("/{engine_id}/runs", response_model=tuple[Job, ...])
async def list_runs(engine_id: str, auth: Auth, offset: int = Query(default=0, ge=0)) -> tuple[Job, ...]:
await get_engine(engine_id, user_scope(auth))
return tuple(
j.model_copy(update=MappingProxyType({"sample": None, "findings": None, "assessments": ()}))
for j in await repository().jobs(engine_id, offset)
)
@router.get("/{engine_id}/runs/{job_id}", response_model=Job)
async def read_run(engine_id: str, job_id: str, auth: Auth) -> Job:
await get_engine(engine_id, user_scope(auth))
job: Final = await repository().job(engine_id, job_id)
if job is None:
raise HTTPException(404, "Investigation not found")
return job
@router.post("/{engine_id}/cancel", response_model=Engine)
@ -215,18 +257,23 @@ async def update_finding(engine_id: str, finding_id: str, body: FindingUpdate, a
class Preview(BaseModel):
as_of: AwareDatetime | None = None
offset: int = Field(default=0, ge=0)
settings: EngineSettings
lookback_hours: int = Field(default=24, ge=1, le=720)
@router.post("/preview/sample", response_model=Sample)
async def preview_sample(body: Preview, auth: Auth) -> Sample:
now: Final = datetime.now(timezone.utc)
validate_selection(body.settings)
now: Final = min(body.as_of or datetime.now(timezone.utc), datetime.now(timezone.utc))
return await source_reader().sample(
user_scope(auth),
body.settings,
int((now - timedelta(hours=body.lookback_hours)).timestamp() * 1000),
int((now - timedelta(minutes=2)).timestamp() * 1000),
offset=body.offset,
preview=True,
)
@ -256,7 +303,9 @@ async def revoke_worker(worker_id: str, auth: Auth) -> bool:
@router.post("/worker/claim", response_model=Claim | None)
async def claim(worker: WorkerAuth) -> Claim | None:
async def claim(worker: WorkerAuth, protocol_version: int = 1) -> Claim | None:
if protocol_version != 2:
raise HTTPException(409, "Upgrade the Lens worker using the current Connect worker command")
now: Final = datetime.now(timezone.utc)
await repository().heartbeat(worker.id, now.isoformat())
for candidate in await repository().engines():
@ -295,9 +344,24 @@ async def sample(engine_id: str, job_id: str, worker: WorkerAuth) -> Sample:
engine, job = await assigned(engine_id, job_id, worker)
if job.sample is not None:
return job.sample
selected: Final = await source_reader().sample(
engine.scope, job.settings, int(job.start.timestamp() * 1000), int(job.end.timestamp() * 1000)
)
pages: list[Sample] = [] # mutable-ok: freeze selection after stable cursor traversal
cursor = "" # rebind-ok: advance by immutable identity, never by shifting row positions
while True:
page = await source_reader().sample(
engine.scope,
job.settings,
int(job.start.timestamp() * 1000),
int(job.end.timestamp() * 1000),
cursor=cursor,
)
pages.append(page)
if not page.next_cursor or sum(len(p.executions) for p in pages) >= pages[0].selected:
break
cursor = page.next_cursor
executions: Final = tuple(
execution for p in pages for execution in p.executions
) # comprehension-ok: flatten query pages
selected: Final = Sample(executions=executions, eligible=pages[0].eligible, selected=len(executions))
def freeze(e: Engine) -> Engine:
active: Final = current_job(e)
@ -323,7 +387,7 @@ async def content(
execution_id: str,
worker: WorkerAuth,
cursor: str = "",
offset: int = Query(default=0, ge=0, le=1000000),
offset: int = Query(default=0, ge=0),
) -> ExecutionContent:
engine, job = await assigned(engine_id, job_id, worker)
selected: Final = job.sample or Sample(executions=(), eligible=0)
@ -351,7 +415,13 @@ async def result(engine_id: str, job_id: str, body: Result, worker: WorkerAuth)
now: Final = datetime.now(timezone.utc)
selected: Final = job.sample or Sample(executions=(), eligible=0)
allowed: Final = frozenset(e.id for e in selected.executions)
check_ids: Final = frozenset(c.id for c in job.settings.checks if c.enabled)
if len(frozenset(a.execution_id for a in body.assessments)) != len(body.assessments):
raise HTTPException(422, "Each run must have one assessment")
if any(a.execution_id not in allowed for a in body.assessments):
raise HTTPException(422, "Assessment references a run outside this job")
check_ids: Final = frozenset(c.id for c in job.settings.analysis_checks)
if any(not check_ids.issuperset((*a.issue_checks, *a.pattern_checks)) for a in body.assessments):
raise HTTPException(422, "Assessment references an unknown check")
if any(
f.check_id not in check_ids or any(e.execution_id not in allowed for e in f.evidence) for f in body.findings
):
@ -376,6 +446,8 @@ async def result(engine_id: str, job_id: str, body: Result, worker: WorkerAuth)
"finished_at": now,
"coverage": active.coverage if body.error else body.coverage,
"error": body.error,
"assessments": body.assessments,
"findings": tuple(snapshot_finding(e, f, job.revision, now) for f in body.findings),
}
)
),
@ -435,7 +507,7 @@ async def validate_finding(engine: Engine, selected: Sample, finding: FindingDra
@router.get("/{engine_id}/executions/{execution_id}", response_model=ExecutionContent)
async def evidence_content(
engine_id: str, execution_id: str, auth: Auth, cursor: str = "", offset: int = Query(default=0, ge=0, le=1000000)
engine_id: str, execution_id: str, auth: Auth, cursor: str = "", offset: int = Query(default=0, ge=0)
) -> ExecutionContent:
engine: Final = await get_engine(engine_id, user_scope(auth))
try:

View file

@ -1,5 +1,5 @@
from datetime import datetime
from typing import Literal
from typing import Final, Literal
from pydantic import BaseModel, ConfigDict, Field, model_validator
@ -32,24 +32,47 @@ class EngineSettings(Record):
lookback_hours: int = Field(default=24, ge=1, le=720)
service: str = Field(default="", max_length=200)
filters: tuple[MetadataFilter, ...] = Field(default=(), max_length=8)
checks: tuple[Check, ...] = Field(min_length=1, max_length=12)
checks: tuple[Check, ...] = ()
model: str = Field(min_length=1, max_length=200)
enabled: bool = True
interval_minutes: int = Field(default=15, ge=1, le=10080)
sample_size: int = Field(default=100, ge=1, le=500)
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=20, gt=0, le=100000, allow_inf_nan=False)
@model_validator(mode="after")
def unique_checks(self) -> "EngineSettings":
if len(frozenset(c.id for c in self.checks)) != len(self.checks):
raise ValueError("Each check must have a unique ID")
if not self.context.strip() and not any(c.enabled for c in self.checks):
raise ValueError("Describe expected behavior or add an enabled check")
if any(c.id == "expected_behavior" for c in self.checks):
raise ValueError("expected_behavior is reserved for the behavior description")
return self
@property
def analysis_checks(self) -> tuple[Check, ...]:
behavior: Final = (
(
Check(
id="expected_behavior",
instruction="Identify deviations from the expected behavior described in context.",
),
)
if self.context.strip()
else ()
)
return (*behavior, *(c for c in self.checks if c.enabled))
class Evidence(Record):
execution_id: str
span_id: str
quote: str = Field(min_length=1, max_length=1000)
role: Literal["support", "counterexample"] = "support"
class FindingDraft(Record):
@ -79,6 +102,7 @@ class Coverage(Record):
selected: int = 0
screened: int = 0
investigated: int = 0
inconclusive: int = 0
grouping_batches: int = 0
grouped_batches: int = 0
candidates: int = 0
@ -120,6 +144,16 @@ class ExecutionContent(Record):
class Sample(Record):
executions: tuple[Execution, ...]
eligible: int
selected: int = 0
next_offset: int | None = None
next_cursor: str | None = None
class RunAssessment(Record):
execution_id: str
issue_checks: tuple[str, ...] = ()
pattern_checks: tuple[str, ...] = ()
cannot_assess: bool = False
class Job(Record):
@ -139,6 +173,8 @@ class Job(Record):
error: str = ""
sample: Sample | None = None
cost: float = 0
findings: tuple[Finding, ...] | None = None
assessments: tuple[RunAssessment, ...] = ()
class Engine(Record):
@ -176,6 +212,7 @@ class EngineList(Record):
class RunRequest(Record):
settings: EngineSettings | None = None
lookback_hours: int | None = Field(default=None, ge=1, le=720)
@ -196,7 +233,8 @@ class Progress(Record):
class Result(Record):
findings: tuple[FindingDraft, ...] = Field(default=(), max_length=30)
assessments: tuple[RunAssessment, ...] = ()
findings: tuple[FindingDraft, ...] = ()
coverage: Coverage
error: str = Field(default="", max_length=1000)

View file

@ -5,7 +5,7 @@ from typing import Final, Protocol
from pydantic import BaseModel, JsonValue, TypeAdapter
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.engine.models import Engine, Worker
from litellm.proxy.engine.models import Engine, Job, Worker
class Database(Protocol):
@ -60,13 +60,54 @@ class EngineRepository:
if candidate == previous:
return True, previous
updated: Final = candidate.model_copy(update=MappingProxyType({"version": previous.version + 1}))
count: Final = await self.db.execute_raw(
'UPDATE "LiteLLM_Engine" SET data=$1::jsonb, version=version+1 WHERE id=$2 AND version=$3',
updated.model_dump_json(),
engine_id,
previous.version,
rows: Final = _ROWS.validate_python(
await self.db.query_raw(
"""WITH previous AS MATERIALIZED (
SELECT data FROM "LiteLLM_Engine" WHERE id=$2 AND version=$3 FOR UPDATE
), updated AS (
UPDATE "LiteLLM_Engine" SET data=$1::jsonb, version=version+1
WHERE id=$2 AND version=$3 AND EXISTS (SELECT 1 FROM previous) RETURNING id
)
, archived AS (INSERT INTO "LiteLLM_EngineRun" (id, engine_id, created_at, data)
SELECT job->>'id', $2, (job->>'created_at')::timestamp, job
FROM previous, jsonb_array_elements(previous.data->'jobs') AS job
WHERE EXISTS (SELECT 1 FROM updated)
AND NOT EXISTS (SELECT 1 FROM jsonb_array_elements(($1::jsonb)->'jobs') AS retained
WHERE retained->>'id'=job->>'id')
ON CONFLICT (id) DO NOTHING)
SELECT to_jsonb(count(*)) AS data FROM updated""",
updated.model_dump_json(),
engine_id,
previous.version,
)
)
return bool(count), updated
return bool(rows and rows[0].data == 1), updated
async def jobs(self, engine_id: str, offset: int = 0) -> tuple[Job, ...]:
rows: Final = _ROWS.validate_python(
await self.db.query_raw(
"""SELECT data FROM (
SELECT data FROM "LiteLLM_EngineRun" WHERE engine_id=$1
UNION ALL
SELECT jsonb_array_elements(data->'jobs') AS data FROM "LiteLLM_Engine" WHERE id=$1
) AS jobs ORDER BY data->>'created_at' DESC, data->>'id' DESC LIMIT 50 OFFSET $2""",
engine_id,
offset,
)
)
return tuple(Job.model_validate(row.data) for row in rows)
async def job(self, engine_id: str, job_id: str) -> Job | None:
rows: Final = _ROWS.validate_python(
await self.db.query_raw(
"""SELECT data FROM "LiteLLM_EngineRun" WHERE engine_id=$1 AND id=$2
UNION ALL SELECT job AS data FROM "LiteLLM_Engine", jsonb_array_elements(data->'jobs') AS job
WHERE id=$1 AND job->>'id'=$2 LIMIT 1""",
engine_id,
job_id,
)
)
return Job.model_validate(rows[0].data) if rows else None
async def workers(self) -> tuple[Worker, ...]:
rows: Final = _ROWS.validate_python(await self.db.query_raw('SELECT data FROM "LiteLLM_EngineWorker"'))

View file

@ -25,6 +25,7 @@ class Storage(Protocol):
class ExecutionRow(BaseModel):
selection_key: str = ""
source: Literal["traces", "requests"]
trace_id: str
trace_ref: str = ""
@ -34,6 +35,7 @@ class ExecutionRow(BaseModel):
span_count: int
root_seen: int
eligible: int
selected: int = 0
service: str = ""
attributes: tuple[tuple[str, str], ...] = ()
@ -79,11 +81,26 @@ def parameters(scope: Scope, filters: tuple[MetadataFilter, ...]) -> Mapping[str
)
def selection_id(value: str) -> str:
source, team, trace_id, trace_ref = parse_execution(value)
return "\0".join((source, team, trace_ref or trace_id))
class SourceReader:
def __init__(self, storage: Storage) -> None:
self.storage: Final = storage
async def sample(self, scope: Scope, settings: EngineSettings, start: int, end: int) -> Sample:
async def sample(
self,
scope: Scope,
settings: EngineSettings,
start: int,
end: int,
offset: int = 0,
page_size: int = 100,
preview: bool = False,
cursor: str = "",
) -> Sample:
params: Final = MappingProxyType(
{
**parameters(scope, settings.filters),
@ -91,12 +108,26 @@ class SourceReader:
"start": start,
"end": end,
"service": settings.service,
"limit": settings.sample_size,
"limit": page_size,
"offset": offset,
"after": cursor,
"sample_percent": str(settings.sample_percent),
"sample_cap": settings.sample_size or 0,
"preview": int(preview),
"selected_team": settings.team_id,
"execution_ids": tuple(selection_id(value) for value in settings.execution_ids),
}
)
rows: Final = _ROWS.validate_python(await self.storage.lens_sample(params))
return Sample(
eligible=rows[0].eligible if rows else 0,
selected=rows[0].selected if rows else 0,
next_cursor=rows[-1].selection_key if len(rows) == page_size else None,
next_offset=(
offset + len(rows)
if page_size and rows and offset + len(rows) < (rows[0].eligible if preview else rows[0].selected)
else None
),
executions=tuple(
Execution(
id=execution_id(row.source, row.team_id, row.trace_id, row.trace_ref),

View file

@ -3,7 +3,7 @@ from datetime import datetime, timedelta
from types import MappingProxyType
from typing import Final
from litellm.proxy.engine.models import Engine, Finding, FindingDraft, Job, Scope, Worker
from litellm.proxy.engine.models import Engine, EngineSettings, Finding, FindingDraft, Job, Scope, Worker
def can_access(viewer: Scope, target: Scope) -> bool:
@ -24,23 +24,25 @@ def replace_job(engine: Engine, job: Job) -> Engine:
)
def queue_job(engine: Engine, now: datetime, job_id: str, lookback_hours: int | None = None) -> Engine:
def queue_job(
engine: Engine,
now: datetime,
job_id: str,
lookback_hours: int | None = None,
settings: EngineSettings | None = None,
) -> Engine:
if current_job(engine):
return engine
start: Final = (
now - timedelta(hours=lookback_hours)
if lookback_hours is not None
else (engine.last_scan_at or now - timedelta(hours=engine.settings.lookback_hours)) - timedelta(minutes=5)
)
selected: Final = settings or engine.settings
job: Final = Job(
id=job_id,
created_at=now,
start=start,
start=now - timedelta(hours=lookback_hours if lookback_hours is not None else selected.lookback_hours),
end=now - timedelta(minutes=2),
settings=engine.settings,
settings=selected,
revision=engine.revision,
)
return engine.model_copy(update=MappingProxyType({"jobs": (job, *engine.jobs[:49])}))
return engine.model_copy(update=MappingProxyType({"jobs": (job,)}))
def claim_job(engine: Engine, worker: Worker, now: datetime) -> Engine:
@ -91,7 +93,7 @@ def renew_budget(engine: Engine, now: datetime) -> Engine:
def merge_finding(engine: Engine, draft: FindingDraft, revision: int, now: datetime) -> Finding:
identity: Final = hashlib.sha256(f"{engine.id}:{draft.check_id}:{draft.title.lower()}".encode()).hexdigest()[:24]
previous: Final = next((f for f in engine.findings if f.id == (draft.existing_finding_id or identity)), None)
occurrences: Final = tuple(sorted(frozenset(e.execution_id for e in draft.evidence)))
occurrences: Final = tuple(sorted(frozenset(e.execution_id for e in draft.evidence if e.role == "support")))
if previous is None:
return Finding(
title=draft.title,
@ -124,3 +126,19 @@ def merge_finding(engine: Engine, draft: FindingDraft, revision: int, now: datet
}
)
)
def snapshot_finding(engine: Engine, draft: FindingDraft, revision: int, now: datetime) -> Finding:
merged: Final = merge_finding(engine, draft, revision, now)
return Finding.model_validate(
MappingProxyType(
{
**merged.model_dump(),
**draft.model_dump(),
"revision": revision,
"first_seen": now,
"last_seen": now,
"occurrences": tuple(sorted(frozenset(e.execution_id for e in draft.evidence if e.role == "support"))),
}
)
)

View file

@ -0,0 +1,103 @@
import json
import sqlite3
from collections.abc import Generator, Iterator
from contextlib import contextmanager
from tempfile import TemporaryDirectory
from typing import Final
from pydantic import TypeAdapter
from .models import Evidence, TracePart
_ROW: Final = TypeAdapter(tuple[str])
_OPTIONAL_ROW: Final = TypeAdapter(tuple[str] | None)
_COUNT: Final = TypeAdapter(tuple[int])
class TraceStore:
def __init__(self, connection: sqlite3.Connection) -> None:
self.connection: Final = connection
connection.execute("CREATE TABLE spans (span_id TEXT PRIMARY KEY, body TEXT NOT NULL)")
connection.execute("CREATE TABLE reads (span_id TEXT, body TEXT, UNIQUE(span_id, body))")
def add(self, parts: tuple[TracePart, ...]) -> None:
self.connection.executemany(
"INSERT OR REPLACE INTO spans VALUES (?, ?)",
((part.span_id, part.model_dump_json()) for part in parts),
)
def add_reads(self, parts: tuple[TracePart, ...]) -> None:
self.connection.executemany(
"INSERT OR IGNORE INTO reads VALUES (?, ?)",
((part.span_id, part.model_dump_json()) for part in parts),
)
def evidence(self, evidence: Evidence) -> TracePart | None:
rows: Final = self.connection.execute(
"SELECT body FROM spans WHERE span_id=? UNION ALL SELECT body FROM reads WHERE span_id=?",
(evidence.span_id, evidence.span_id),
)
for row in map(_ROW.validate_python, rows):
part = TracePart.model_validate_json(row[0])
if part.execution_id == evidence.execution_id and any(
evidence.quote in segment for segment in part.content.split("\n[... content omitted ...]\n")
):
return part
return None
def parts(self) -> Iterator[TracePart]:
for row in map(_ROW.validate_python, self.connection.execute("SELECT body FROM spans ORDER BY span_id")):
yield TracePart.model_validate_json(row[0])
def get(self, span_id: str) -> TracePart | None:
row: Final = _OPTIONAL_ROW.validate_python(
self.connection.execute("SELECT body FROM spans WHERE span_id=?", (span_id,)).fetchone()
)
return TracePart.model_validate_json(row[0]) if row else None
def previous(self, span_id: str) -> str:
row: Final = _OPTIONAL_ROW.validate_python(
self.connection.execute(
"SELECT span_id FROM spans WHERE span_id < ? ORDER BY span_id DESC LIMIT 1", (span_id,)
).fetchone()
)
return row[0] if row else ""
def count(self) -> int:
return _COUNT.validate_python(self.connection.execute("SELECT count(*) FROM spans").fetchone())[0]
def catalogs(self, root_count: int) -> Iterator[tuple[tuple[str, str, str, str, str], ...]]:
rows: list[tuple[str, str, str, str, str]] = [] # mutable-ok: one bounded catalog window
size = 0 # rebind-ok: track the current window's serialized size
for part in self.parts():
row = (part.span_id, part.parent_span_id, part.name, part.kind, overview_content(part, root_count))
width = len(json.dumps(row))
if rows and size + width > 24000:
yield tuple(rows)
rows.clear()
size = 0
rows.append(row)
size += width
if rows:
yield tuple(rows)
def overview_content(part: TracePart, root_count: int) -> str:
limit: Final = max(160, min(2000, 12000 // max(root_count, 1))) if not part.parent_span_id else 160
if len(part.content) <= limit:
return part.content
return (
part.content[: limit // 3]
+ "\n[... preview omitted; read this span for evidence ...]\n"
+ part.content[-(limit * 2 // 3) :]
)
@contextmanager
def trace_store() -> Generator[TraceStore]:
with TemporaryDirectory(prefix="lens-trace-") as directory:
connection: Final = sqlite3.connect(f"{directory}/trace.sqlite")
try:
yield TraceStore(connection)
finally:
connection.close()

View file

@ -1,6 +1,7 @@
import asyncio
import logging
import os
from collections.abc import Awaitable, Callable
from contextlib import suppress
from types import MappingProxyType
from typing import Final
@ -14,11 +15,31 @@ logger: Final = logging.getLogger("litellm.engine.worker")
class EngineWorker:
def __init__(self, client: httpx.AsyncClient) -> None:
def __init__(self, client: httpx.AsyncClient, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep) -> None:
self.client: Final = client
self.sleep: Final = sleep
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())
result.raise_for_status()
return ModelResult.model_validate(result.json())
except (httpx.TransportError, httpx.HTTPStatusError) as exc:
retryable: Final = not isinstance(exc, httpx.HTTPStatusError) or exc.response.status_code in (
429,
502,
503,
504,
)
if not retryable or attempt >= 2:
raise
await self.sleep(2**attempt)
return await self.model_request(path, body, attempt + 1)
async def run_once(self) -> bool:
response: Final = await self.client.post("/engine/worker/claim")
response: Final = await self.client.post(
"/engine/worker/claim", params=MappingProxyType({"protocol_version": 2})
)
response.raise_for_status()
if response.json() is None:
return False
@ -26,9 +47,7 @@ class EngineWorker:
prefix: Final = f"/engine/worker/{claim.engine_id}/{claim.job.id}"
async def model(body: ModelRequest) -> ModelResult:
result: Final = await self.client.post(prefix + "/model", json=body.model_dump())
result.raise_for_status()
return ModelResult.model_validate(result.json())
return await self.model_request(prefix + "/model", body)
async def read(execution_id: str, cursor: str, offset: int) -> ExecutionContent:
result: Final = await self.client.get(

View file

@ -19,6 +19,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
inspect_embeddings=litellm_params.inspect_embeddings,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_aim_callback)

View file

@ -181,6 +181,7 @@ class AimGuardrail(CustomGuardrail):
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": self._build_aim_inspection_messages(data)},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[AimAnalyzeResponse] = response.json()
@ -285,6 +286,7 @@ class AimGuardrail(CustomGuardrail):
"messages": self._build_aim_inspection_messages(request_data)
+ [{"role": "assistant", "content": output}]
},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[AimAnalyzeResponse] = response.json()

View file

@ -18,6 +18,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_alice_guardrail_callback)

View file

@ -227,6 +227,7 @@ class AliceGuardrail(CustomGuardrail):
"Content-Type": "application/json",
"af-api-key": self.alice_api_key,
},
timeout=self.timeout,
)
response.raise_for_status()
body = response.json()

View file

@ -17,6 +17,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_aporia_callback)

View file

@ -123,6 +123,7 @@ class AporiaGuardrail(CustomGuardrail):
"X-APORIA-API-KEY": self.aporia_api_key,
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("Aporia AI response: %s", response.text)
if response.status_code == 200:

View file

@ -1,5 +1,8 @@
import re
from typing import TYPE_CHECKING, Any, Final
from collections.abc import Mapping
from typing import Any, Final, cast
import httpx
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -9,9 +12,9 @@ from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
if TYPE_CHECKING:
from litellm.types.llms.openai import AllMessageValues
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import AllMessageValues, ResponseInputParam
from litellm.types.utils import CallTypes, CallTypesLiteral
# Azure Content Safety APIs have a 10,000 character limit per request.
AZURE_CONTENT_SAFETY_MAX_TEXT_LENGTH: Final = 10000
@ -23,6 +26,8 @@ AZURE_CONTENT_SAFETY_TEXT_RECORD_LENGTH: Final = 1000
AZURE_CONTENT_SAFETY_DEFAULT_API_VERSION: Final = "2024-09-01"
JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES: Final = "v1"
_RESPONSES_API_CALL_TYPES: Final = frozenset({CallTypes.responses, CallTypes.aresponses})
def resolve_content_safety_api_version(configured: str | None) -> str:
if not configured or configured == JAVELIN_API_VERSION_STORED_BY_OLDER_RELEASES:
@ -49,6 +54,7 @@ class AzureGuardrailBase:
# (typically CustomGuardrail).
super().__init__(**kwargs)
self.timeout: float | httpx.Timeout | None
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.api_key = api_key
self.api_base = api_base
@ -77,6 +83,7 @@ class AzureGuardrailBase:
url=url,
headers=headers,
json=request_body,
timeout=self.timeout,
)
response_json: Final[dict[str, Any]] = response.json()
verbose_proxy_logger.debug("Azure Content Safety response [%s]: %s", endpoint_path, response_json)
@ -131,16 +138,15 @@ class AzureGuardrailBase:
return chunks
def get_user_prompt(self, messages: list["AllMessageValues"]) -> str | None:
"""
Get the last consecutive block of messages from the user.
def get_user_prompt_from_request(self, data: Mapping[str, object], call_type: CallTypesLiteral) -> str | None:
if call_type in _RESPONSES_API_CALL_TYPES:
responses_input: Final = data.get("input")
if not isinstance(responses_input, (str, list)):
return None
validated_input: Final = cast(ResponseInputParam, responses_input) # cast-ok: narrowed to str | list
return get_last_user_message(ResponsesAPIRequestUtils.responses_input_to_chat_messages(validated_input))
Example:
messages = [
{"role": "user", "content": "Hello, how are you?"},
{"role": "assistant", "content": "I'm good, thank you!"},
{"role": "user", "content": "What is the weather in Tokyo?"},
]
get_user_prompt(messages) -> "What is the weather in Tokyo?"
"""
return get_last_user_message(messages)
messages: Final = data.get("messages")
if not isinstance(messages, list):
return None
return get_last_user_message(cast(list[AllMessageValues], messages)) # cast-ok: narrowed to list

View file

@ -33,7 +33,6 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import LitellmParams
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_prompt_shield import (
AzurePromptShieldGuardrailResponse,
)
@ -250,11 +249,7 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai
"Azure Prompt Shield: Running pre-call prompt scan, on call_type: %s",
call_type,
)
new_messages: Final[list[AllMessageValues] | None] = data.get("messages")
if new_messages is None:
verbose_proxy_logger.warning("Azure Prompt Shield: not running guardrail. No messages in data")
return data
user_prompt: Final = self.get_user_prompt(new_messages)
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
if user_prompt:
verbose_proxy_logger.debug("Azure Prompt Shield: User prompt: %s", user_prompt)

View file

@ -21,7 +21,6 @@ from .base import AzureGuardrailBase
if TYPE_CHECKING:
from litellm.caching.caching import DualCache
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.azure.azure_text_moderation import (
AzureTextModerationGuardrailResponse,
)
@ -232,14 +231,10 @@ class AzureContentSafetyTextModerationGuardrail(AzureGuardrailBase, CustomGuardr
"Azure Text Moderation: Running pre-call prompt scan, on call_type: %s",
call_type,
)
new_messages: Final[list[AllMessageValues] | None] = data.get("messages")
if new_messages is None:
verbose_proxy_logger.warning("Azure Text Moderation: not running guardrail. No messages in data")
return data
user_prompt: Final = self.get_user_prompt(new_messages)
user_prompt: Final = self.get_user_prompt_from_request(data, call_type)
if user_prompt:
verbose_proxy_logger.info("Azure Text Moderation: User prompt: %s", user_prompt)
verbose_proxy_logger.debug("Azure Text Moderation: User prompt: %s", user_prompt)
await self.async_make_request(
text=user_prompt,
)

View file

@ -1787,6 +1787,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
url=prepared_request.url,
data=prepared_request.body,
headers=prepared_request.headers,
timeout=self.timeout,
)
except HTTPException:
# Propagate HTTPException (e.g. from non-200 path) as-is

View file

@ -22,6 +22,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
default_on=litellm_params.default_on,
inspect_embeddings=litellm_params.inspect_embeddings,
ssl_verify=getattr(litellm_params, "ssl_verify", None),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_cato_callback)

View file

@ -305,6 +305,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
f"{self.api_base}/fw/v1/analyze",
headers=headers,
json={"messages": self._inspection_messages(data)},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_CatoAnalyzeResponse] = response.json()
@ -445,6 +446,7 @@ class CatoNetworksGuardrail(CustomGuardrail):
litellm_call_id=call_id,
),
json={"messages": inspection_messages + [{"role": "assistant", "content": output}]},
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_CatoAnalyzeResponse] = response.json()

View file

@ -214,8 +214,6 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail):
else:
env_timeout: Final = os.environ.get("CISCO_AI_DEFENSE_TIMEOUT")
resolved_timeout = self._coerce_timeout(env_timeout) if env_timeout is not None else None
self.timeout: float = resolved_timeout if resolved_timeout is not None else DEFAULT_TIMEOUT_SECONDS
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
# Register broadly; runtime filtering happens in ``_surface_matches``.
@ -224,6 +222,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail):
supported_event_hooks=list(self.get_supported_event_hooks()),
**kwargs,
)
self.timeout = resolved_timeout if resolved_timeout is not None else DEFAULT_TIMEOUT_SECONDS
self._warn_if_mode_surface_mismatch(kwargs.get("event_hook"))

View file

@ -59,6 +59,7 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) ->
event_hook=_coerce_event_hook(litellm_params.mode),
default_on=litellm_params.default_on or False,
unreachable_fallback=litellm_params.unreachable_fallback,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] # callback manager is untyped
_callback

View file

@ -520,6 +520,7 @@ class CompresrGuardrail(CustomGuardrail):
dynamic_min_ratio: float | None = None,
dynamic_max_ratio: float | None = None,
compression_params: dict[str, object] | None = None,
timeout: float | None = None,
):
raw_api_base: Final = (api_base or get_secret_str("COMPRESR_API_BASE") or DEFAULT_API_BASE).rstrip("/")
self.compresr_api_base = _validate_api_base(raw_api_base)
@ -583,6 +584,7 @@ class CompresrGuardrail(CustomGuardrail):
guardrail_name=guardrail_name,
event_hook=event_hook,
default_on=default_on,
timeout=timeout,
)
def _should_bypass(self, request_data: dict) -> bool:
@ -755,7 +757,7 @@ class CompresrGuardrail(CustomGuardrail):
url=url,
json=payload,
headers=self._request_headers(),
timeout=_COMPRESS_TIMEOUT_SECONDS,
timeout=self.timeout if self.timeout is not None else _COMPRESS_TIMEOUT_SECONDS,
)
except asyncio.CancelledError:
raise

View file

@ -27,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_buffer_release_on_scan=streaming_params.streaming_buffer_release_on_scan,
streaming_end_of_stream_only=streaming_params.streaming_end_of_stream_only,
streaming_sampling_rate=streaming_params.streaming_sampling_rate,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_crowdstrike_aidr_callback)

View file

@ -355,7 +355,9 @@ class CrowdStrikeAIDRHandler(CustomGuardrail):
"CrowdStrike AIDR Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload
)
response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
response: Final = await self.async_handler.post(
url=endpoint, json=payload, headers=headers, timeout=self.timeout
)
assert response is not None
response.raise_for_status()

View file

@ -20,6 +20,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_deepkeep_guardrail_callback)

View file

@ -393,6 +393,7 @@ class DeepKeepGuardrail(CustomGuardrail):
url=self.api_base,
json=guardrail_request,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()

View file

@ -17,6 +17,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_dynamoai_callback)

View file

@ -130,6 +130,7 @@ class DynamoAIGuardrails(CustomGuardrail):
url=self.api_url,
json=dict(payload),
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final = response.json()

View file

@ -24,6 +24,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
block_on_violation=litellm_params.block_on_violation,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_enkryptai_callback)

View file

@ -123,6 +123,7 @@ class EnkryptAIGuardrails(CustomGuardrail):
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final = response.json()

View file

@ -39,6 +39,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -477,6 +477,7 @@ class GenericGuardrailAPI(CustomGuardrail):
url=self.api_base,
json=guardrail_request.model_dump(mode="json"),
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()

View file

@ -24,6 +24,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
default_on=litellm_params.default_on,
guard_name=litellm_params.guard_name,
guardrails_ai_api_input_format=getattr(litellm_params, "guardrails_ai_api_input_format", "llmOutput"),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_guardrails_ai_callback)

View file

@ -80,6 +80,7 @@ class GuardrailsAI(CustomGuardrail):
headers={
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
_json_response: Final = GuardrailsAIResponse(**response.json())
@ -117,6 +118,7 @@ class GuardrailsAI(CustomGuardrail):
headers={
"Content-Type": "application/json",
},
timeout=self.timeout,
)
verbose_proxy_logger.debug("guardrails_ai response: %s", response)
if response.status_code == 400:

View file

@ -508,7 +508,6 @@ class HeadroomGuardrail(CustomGuardrail):
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
"fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
)
self.timeout: httpx.Timeout = self._resolve_timeout(timeout)
self.ccr_retrieval = ccr_retrieval
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
@ -520,6 +519,7 @@ class HeadroomGuardrail(CustomGuardrail):
default_on=default_on,
supported_event_hooks=list(self.get_supported_event_hooks()),
)
self.timeout = self._resolve_timeout(timeout)
def _should_bypass(self, request_data: dict) -> bool:
psr: Final = request_data.get("proxy_server_request")

View file

@ -25,6 +25,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
else:
_hiddenlayer_callback = HiddenlayerGuardrailV2(
@ -35,6 +36,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_hiddenlayer_callback)

View file

@ -243,15 +243,19 @@ class HiddenlayerGuardrail(CustomGuardrail):
if not self.hiddenlayer_client_secret:
raise RuntimeError("`api_key` cannot be None when using the SaaS version of HiddenLayer.")
ctor_timeout: Final = kwargs.get("timeout")
auth_timeout: Final = ctor_timeout if isinstance(ctor_timeout, (int, float)) else _AUTH_TIMEOUT_SECONDS
self.jwt_token = _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self.refresh_jwt_func = lambda: _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self._http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
@ -382,6 +386,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
f"{self.api_base}/detection/v1/interactions",
json=data,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
result: _HiddenlayerResponse = _interaction_body(response)
@ -403,6 +408,7 @@ class HiddenlayerGuardrail(CustomGuardrail):
f"{self.api_base}/detection/v1/interactions",
json=data,
headers=headers,
timeout=self.timeout,
)
else:
raise e
@ -447,15 +453,19 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
if not self.hiddenlayer_client_secret:
raise RuntimeError("`api_key` cannot be None when using the SaaS version of HiddenLayer.")
ctor_timeout: Final = kwargs.get("timeout")
auth_timeout: Final = ctor_timeout if isinstance(ctor_timeout, (int, float)) else _AUTH_TIMEOUT_SECONDS
self.jwt_token = _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self.refresh_jwt_func = lambda: _get_jwt(
auth_url=auth_url,
api_id=self.hiddenlayer_client_id,
api_key=self.hiddenlayer_client_secret,
timeout=auth_timeout,
)
self._http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
@ -584,6 +594,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
f"{self.api_base}/{path}",
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
@ -604,6 +615,7 @@ class HiddenlayerGuardrailV2(CustomGuardrail):
f"{self.api_base}/{path}",
json=payload,
headers=headers,
timeout=self.timeout,
)
else:
raise e

View file

@ -49,6 +49,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
verify_ssl=verify_ssl,
default_on=litellm_params.default_on,
event_hook=litellm_params.mode,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(ibm_guardrail)

View file

@ -140,6 +140,7 @@ class IBMGuardrailDetector(CustomGuardrail):
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final[list[list[IBMDetectorDetection]]] = response.json()
@ -231,6 +232,7 @@ class IBMGuardrailDetector(CustomGuardrail):
url=self.api_url,
json=payload,
headers=headers,
timeout=self.timeout,
)
response.raise_for_status()
response_json: Final[IBMDetectorResponseOrchestrator] = response.json()

View file

@ -27,6 +27,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
config=litellm_params.config,
metadata=litellm_params.metadata,
application=litellm_params.application,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_javelin_callback)

View file

@ -111,6 +111,7 @@ class JavelinGuardrail(CustomGuardrail):
url=url,
headers=headers,
json=dict(request),
timeout=self.timeout,
)
verbose_proxy_logger.debug("Javelin Guardrail: Javelin guard API response: %s", response.json())
response_data: Final = response.json()

View file

@ -250,6 +250,7 @@ class lakeraAI_Moderation(CustomGuardrail):
"Authorization": "Bearer " + self.lakera_api_key,
"Content-Type": "application/json",
},
timeout=self.timeout,
)
except httpx.HTTPStatusError as e:
raise Exception(e.response.text)

View file

@ -402,6 +402,7 @@ class LakeraAIGuardrail(CustomGuardrail):
url=f"{self.api_base}/v2/guard",
headers={"Authorization": f"Bearer {self.lakera_api_key}"},
json=request,
timeout=self.timeout,
)
verbose_proxy_logger.debug("Lakera AI v2 guard response: %s", response.json())
lakera_response = LakeraAIResponse(**response.json())

View file

@ -19,6 +19,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
conversation_id=litellm_params.lasso_conversation_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_lasso_callback)

View file

@ -814,7 +814,7 @@ class LassoGuardrail(CustomGuardrail):
url=url,
headers=headers,
json=payload,
timeout=10.0,
timeout=self.timeout if self.timeout is not None else 10.0,
)
response.raise_for_status()
return response.json()

View file

@ -62,6 +62,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
debug_headers=_get("debug_headers") or False,
# FR-10: configurable scopes
allowed_scopes=_get("allowed_scopes"),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(signer)
return signer

View file

@ -76,6 +76,7 @@ import time
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Optional
import httpx
import jwt
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric import rsa
@ -173,7 +174,7 @@ def _compute_kid(public_key: RSAPublicKey) -> str:
return hashlib.sha256(der_bytes).hexdigest()[:16]
async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
async def _fetch_jwks(jwks_uri: str, timeout: float | httpx.Timeout | None = None) -> Sequence[Mapping[str, object]]:
"""
Fetch and cache a JWKS from the given URI.
@ -192,7 +193,7 @@ async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
)
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"})
resp: Final = await client.get(jwks_uri, headers={"Accept": "application/json"}, timeout=timeout)
resp.raise_for_status()
jwks_body: Final[Mapping[str, Sequence[Mapping[str, object]]]] = resp.json()
fetched_keys: Final = jwks_body.get("keys", [])
@ -200,7 +201,9 @@ async def _fetch_jwks(jwks_uri: str) -> Sequence[Mapping[str, object]]:
return fetched_keys
async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument:
async def _fetch_oidc_discovery(
discovery_uri: str, timeout: float | httpx.Timeout | None = None
) -> _OIDCDiscoveryDocument:
"""Fetch an OIDC discovery document and return its parsed JSON."""
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
@ -208,7 +211,7 @@ async def _fetch_oidc_discovery(discovery_uri: str) -> _OIDCDiscoveryDocument:
)
client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"})
resp: Final = await client.get(discovery_uri, headers={"Accept": "application/json"}, timeout=timeout)
resp.raise_for_status()
document: Final[_OIDCDiscoveryDocument] = resp.json()
return document
@ -417,7 +420,7 @@ class MCPJWTSigner(CustomGuardrail):
now: Final = time.time()
cache_expired: Final = (now - self._oidc_discovery_fetched_at) >= self._OIDC_DISCOVERY_TTL
if (self._oidc_discovery_doc is None or cache_expired) and self.access_token_discovery_uri:
doc: Final = await _fetch_oidc_discovery(self.access_token_discovery_uri)
doc: Final = await _fetch_oidc_discovery(self.access_token_discovery_uri, timeout=self.timeout)
if "jwks_uri" in doc:
self._oidc_discovery_doc = doc
self._oidc_discovery_fetched_at = now
@ -440,7 +443,7 @@ class MCPJWTSigner(CustomGuardrail):
f"at {self.access_token_discovery_uri!r} has no 'jwks_uri'."
)
jwks_keys: Final = await _fetch_jwks(jwks_uri)
jwks_keys: Final = await _fetch_jwks(jwks_uri, timeout=self.timeout)
# Only read `kid` from the unverified header — never `alg`.
# Reading `alg` from an attacker-controlled header enables algorithm
@ -511,6 +514,7 @@ class MCPJWTSigner(CustomGuardrail):
self.token_introspection_endpoint,
data={"token": token},
headers={"Accept": "application/json"},
timeout=self.timeout,
)
resp.raise_for_status()
result: Final[dict[str, object]] = resp.json()

View file

@ -38,6 +38,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
user_id_field=str(getattr(litellm_params, "user_id_field", None) or "user_id"),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(purview_guardrail)

View file

@ -5,6 +5,7 @@ from collections import OrderedDict
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final
import httpx
from typing_extensions import NotRequired, TypedDict
from litellm._logging import verbose_proxy_logger
@ -56,6 +57,7 @@ class PurviewGuardrailBase:
# (typically CustomGuardrail).
super().__init__(**kwargs)
self.timeout: float | httpx.Timeout | None
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.tenant_id = tenant_id
self.client_id = client_id
@ -107,6 +109,7 @@ class PurviewGuardrailBase:
url=url,
data=data,
headers={"Content-Type": "application/x-www-form-urlencoded"},
timeout=self.timeout,
)
response.raise_for_status()
token_data: Final[GraphTokenResponse] = response.json()
@ -143,7 +146,7 @@ class PurviewGuardrailBase:
headers.update(extra_headers)
verbose_proxy_logger.debug("Purview Graph POST %s", url)
response: Final = await self.async_handler.post(url=url, headers=headers, json=json_body)
response: Final = await self.async_handler.post(url=url, headers=headers, json=json_body, timeout=self.timeout)
response.raise_for_status()
response_json: Final[dict[str, object]] = response.json()
response_headers: Final = dict(response.headers)

View file

@ -28,6 +28,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
fail_on_error=litellm_params.fail_on_error,
skip_unscannable_attachments=litellm_params.skip_unscannable_attachments,
sanitize_error_detail=litellm_params.sanitize_error_detail,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_model_armor_callback)

View file

@ -337,6 +337,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase):
url=url,
json=body,
headers=headers,
timeout=self.timeout,
)
except httpx.HTTPStatusError as e:
detail = self._build_api_error_detail(e.response.status_code, e.response.text)

View file

@ -28,6 +28,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
anonymize_input=litellm_params.anonymize_input,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_noma_callback)
@ -47,6 +48,7 @@ def initialize_guardrail_v2(litellm_params: "LitellmParams", guardrail: "Guardra
block_failures=litellm_params.block_failures,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_noma_v2_callback)

View file

@ -751,6 +751,7 @@ class NomaGuardrail(CustomGuardrail):
"requestId": llm_request_id,
},
},
timeout=self.timeout,
)
response.raise_for_status()

View file

@ -220,6 +220,7 @@ class NomaV2Guardrail(CustomGuardrail):
url=endpoint,
headers=headers,
json=sanitized_payload,
timeout=self.timeout,
)
verbose_proxy_logger.debug(
"Noma v2 AIDR response: status_code=%s body=%s",

View file

@ -16,6 +16,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_onyx_callback)

View file

@ -116,6 +116,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail):
"Content-Type": "application/json",
},
json=request_body,
timeout=self.timeout,
)
verbose_proxy_logger.debug("OpenAI Moderation guard response: %s", response.json())

View file

@ -29,6 +29,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
post_checkpoint_id=post_checkpoint_id,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_ovalix_callback)

View file

@ -194,7 +194,7 @@ class OvalixGuardrail(CustomGuardrail):
"data_type": "TEXT",
"data": {"content": content},
}
response: Final = await self._async_handler.post(url, headers=headers, json=payload)
response: Final = await self._async_handler.post(url, headers=headers, json=payload, timeout=self.timeout)
response.raise_for_status()
return response.json()

View file

@ -23,6 +23,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
api_key=litellm_params.api_key,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_pangea_callback)

View file

@ -131,7 +131,9 @@ class PangeaHandler(CustomGuardrail):
"Pangea Guardrail (%s): Calling endpoint %s with payload: %s", hook_name, endpoint, payload
)
response: Final = await self.async_handler.post(url=endpoint, json=payload, headers=headers)
response: Final = await self.async_handler.post(
url=endpoint, json=payload, headers=headers, timeout=self.timeout
)
response.raise_for_status()
result: Final = response.json()

View file

@ -11,6 +11,8 @@ import os
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol
from urllib.parse import quote
import httpx
# Third-party imports
from fastapi import HTTPException
from typing_extensions import NotRequired, ReadOnly, TypedDict
@ -66,7 +68,7 @@ class _PillarProtectHTTPClient(Protocol):
url: str,
headers: dict[str, str],
json: dict[str, object],
timeout: float,
timeout: float | httpx.Timeout | None,
) -> _PillarProtectHTTPResponse: ...
@ -284,7 +286,12 @@ class PillarGuardrail(CustomGuardrail):
verbose_proxy_logger.debug("Pillar Guardrail: Initialized with fallback_on_error: %s", self.fallback_on_error)
# Set timeout with graceful fallback on invalid configuration
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=list(self.get_supported_event_hooks()),
**kwargs,
)
if timeout is not None:
self.timeout = timeout
else:
@ -298,12 +305,6 @@ class PillarGuardrail(CustomGuardrail):
)
self.timeout = self.DEFAULT_TIMEOUT
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=list(self.get_supported_event_hooks()),
**kwargs,
)
# =========================================================================
# PUBLIC HOOK METHODS (Main Interface)
# =========================================================================

View file

@ -460,6 +460,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
analyze_url,
json=analyze_payload,
headers={"Accept": "application/json"},
timeout=(
aiohttp.ClientTimeout(total=self.timeout)
if isinstance(self.timeout, (int, float))
else aiohttp.client.DEFAULT_TIMEOUT
),
) as response:
# Validate HTTP status
if response.status >= 400:
@ -745,6 +750,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
anonymize_url,
json=anonymize_payload,
headers={"Accept": "application/json"},
timeout=(
aiohttp.ClientTimeout(total=self.timeout)
if isinstance(self.timeout, (int, float))
else aiohttp.client.DEFAULT_TIMEOUT
),
) as response:
if response.status >= 400:
error_body = await response.text()

View file

@ -23,6 +23,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_transform_mode=getattr(litellm_params, "streaming_transform_mode", None),
file_sanitization_fail_open=getattr(litellm_params, "file_sanitization_fail_open", None),
block_on_file_modify=getattr(litellm_params, "block_on_file_modify", None),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_prompt_security_callback)

View file

@ -290,6 +290,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/protect",
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_ProtectResponse] = response.json()
@ -407,6 +408,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/protect",
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
res: Final[_ProtectResponse] = response.json()
@ -522,6 +524,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/sanitizeFile",
headers=headers,
files=files,
timeout=self.timeout,
)
upload_response.raise_for_status()
upload_result: Final[_SanitizeUploadResponse] = upload_response.json()
@ -552,6 +555,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
f"{self.api_base}/api/sanitizeFile",
headers=headers,
params={"jobId": job_id},
timeout=self.timeout,
)
poll_response.raise_for_status()
result: _SanitizeStatusResponse = poll_response.json()

View file

@ -24,6 +24,7 @@ def initialize_guardrail(
),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(
_cb,

View file

@ -168,7 +168,7 @@ class PromptGuardGuardrail(CustomGuardrail):
"Content-Type": "application/json",
},
json=payload,
timeout=10.0,
timeout=self.timeout if self.timeout is not None else 10.0,
)
response.raise_for_status()
view: Final[PromptGuardHTTPView] = {"guard_response": response.json()}

View file

@ -18,6 +18,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
default_on=litellm_params.default_on,
additional_provider_specific_params=litellm_params.additional_provider_specific_params,
extra_headers=getattr(litellm_params, "extra_headers", None),
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_instance)

View file

@ -26,6 +26,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_qualifire_callback)

View file

@ -378,6 +378,7 @@ class QualifireGuardrail(CustomGuardrail):
url=url,
headers=headers,
json=payload,
timeout=self.timeout,
)
response.raise_for_status()
result: Final = response.json()

View file

@ -33,6 +33,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
unreachable_fallback=litellm_params.unreachable_fallback,
event_hook=_event_hook_from_mode(litellm_params.mode),
default_on=litellm_params.default_on or False,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(_repelloai_callback)

View file

@ -148,6 +148,7 @@ class RepelloAIGuardrail(CustomGuardrail):
guardrail_name: str | None = None,
event_hook: (GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None) = None,
default_on: bool = False,
timeout: float | None = None,
):
self.repelloai_api_key = api_key or get_secret_str("ARGUS_API_KEY") or get_secret_str("REPELLOAI_API_KEY") or ""
if not self.repelloai_api_key:
@ -176,6 +177,7 @@ class RepelloAIGuardrail(CustomGuardrail):
event_hook=event_hook,
default_on=default_on,
supported_event_hooks=list(self.get_supported_event_hooks()),
timeout=timeout,
)
async def _call_analyze(
@ -201,6 +203,7 @@ class RepelloAIGuardrail(CustomGuardrail):
url=endpoint,
headers={"X-API-Key": self.repelloai_api_key},
json=request,
timeout=self.timeout,
)
self._raise_for_config_error(response)
response.raise_for_status()

View file

@ -30,6 +30,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
timeout=litellm_params.timeout,
)
litellm.logging_callback_manager.add_litellm_callback(rubrik_callback)

View file

@ -85,8 +85,6 @@ class SingulrGuardrail(CustomGuardrail):
else:
self.block_on_error = block_on_error
self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
@ -101,6 +99,7 @@ class SingulrGuardrail(CustomGuardrail):
]
super().__init__(**kwargs)
self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:

View file

@ -909,7 +909,6 @@ class StraikerGuardrail(CustomGuardrail):
max_size_in_memory=V3_BLOCKED_TURN_MEMORY, default_ttl=V3_BLOCKED_TURN_TTL_SECONDS
)
self.source = source
self.timeout = float(timeout)
self.max_retries = max(0, int(max_retries))
self.initial_backoff = max(0.0, float(initial_backoff))
self.max_backoff = max(self.initial_backoff, float(max_backoff))
@ -928,7 +927,8 @@ class StraikerGuardrail(CustomGuardrail):
)
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
super().__init__(**kwargs)
super().__init__(**kwargs) # pyright: ignore[reportArgumentType] # kwargs splat carries object-typed values
self.timeout = float(timeout)
self.configured_modes = _configured_modes(self.event_hook)

View file

@ -55,6 +55,7 @@ def initialize_guardrail(litellm_params: LitellmParams, guardrail: Guardrail) ->
guardrail_name=guardrail["guardrail_name"],
event_hook=_coerce_event_hook(litellm_params.mode),
default_on=litellm_params.default_on or False,
timeout=litellm_params.timeout,
unreachable_fallback=(
litellm_params.unreachable_fallback if "unreachable_fallback" in litellm_params.model_fields_set else None
),

View file

@ -161,6 +161,7 @@ class TypeSafeGuardrail(CustomGuardrail):
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None = None,
default_on: bool = False,
async_handler: AsyncHTTPHandler | None = None,
timeout: float | None = None,
) -> None:
raw_api_base: Final = (api_base or get_secret_str("TYPESAFE_API_BASE") or DEFAULT_API_BASE).rstrip("/")
self.typesafe_api_base = raw_api_base
@ -188,6 +189,7 @@ class TypeSafeGuardrail(CustomGuardrail):
guardrail_name=guardrail_name,
event_hook=event_hook,
default_on=default_on,
timeout=timeout,
)
def _handle_failure(self, error: str, log_detail: dict[str, object]) -> None:
@ -271,7 +273,7 @@ class TypeSafeGuardrail(CustomGuardrail):
"Authorization": f"Bearer {self.typesafe_api_key}",
"Content-Type": "application/json",
},
timeout=_JEV_TIMEOUT_SECONDS,
timeout=self.timeout if self.timeout is not None else _JEV_TIMEOUT_SECONDS,
)
except asyncio.CancelledError:
raise

View file

@ -85,7 +85,7 @@ class _AsyncPostHandler(Protocol):
url: str,
headers: dict[str, str],
json: _AnalyzePayload,
timeout: httpx.Timeout,
timeout: float | httpx.Timeout | None,
) -> Awaitable[httpx.Response]: ...
@ -122,10 +122,6 @@ class VigilGuardGuardrail(CustomGuardrail):
fallback: Final = (unreachable_fallback or "fail_closed").lower()
self.unreachable_fallback: _FallbackMode = "fail_open" if fallback == "fail_open" else "fail_closed"
self.timeout: httpx.Timeout = (
_DEFAULT_VIGIL_TIMEOUT if timeout is None else httpx.Timeout(timeout, connect=min(timeout, 5.0))
)
self.async_handler: _AsyncPostHandler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback,
)
@ -137,6 +133,8 @@ class VigilGuardGuardrail(CustomGuardrail):
super().__init__(**forwarded)
self.timeout = _DEFAULT_VIGIL_TIMEOUT if timeout is None else httpx.Timeout(timeout, connect=min(timeout, 5.0))
@staticmethod
def get_config_model() -> type["GuardrailConfigModel"] | None:
from litellm.types.proxy.guardrails.guardrail_hooks.vigil_guard import (

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