mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge remote-tracking branch 'origin/main' into litellm_passthrough_error_preview_stream
This commit is contained in:
commit
58d02d5f84
789 changed files with 12839 additions and 6061 deletions
|
|
@ -7,6 +7,9 @@ legacy_flags=(
|
|||
caching-local
|
||||
enterprise-package
|
||||
enterprise-routing
|
||||
integrations
|
||||
llm-other-providers
|
||||
llm-vertex-ai
|
||||
mcp-integration
|
||||
misc
|
||||
proxy-db-auth-checks
|
||||
|
|
@ -50,6 +53,9 @@ legacy_paths() {
|
|||
echo tests/unit/enterprise/proxy/test_file_deletion_blocking.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_access_check.py
|
||||
echo tests/unit/enterprise/proxy/test_managed_files_hook.py ;;
|
||||
integrations) echo tests/unit/integrations ;;
|
||||
llm-other-providers) find tests/unit/llms -name 'test_*.py' -not -path 'tests/unit/llms/vertex_ai/*' ;;
|
||||
llm-vertex-ai) echo tests/unit/llms/vertex_ai ;;
|
||||
mcp-integration)
|
||||
echo tests/unit/experimental_mcp_client
|
||||
echo tests/unit/proxy/_experimental/mcp_server
|
||||
|
|
@ -71,6 +77,7 @@ legacy_paths() {
|
|||
echo tests/unit/messages
|
||||
echo tests/unit/rag
|
||||
echo tests/unit/rerank_api
|
||||
echo tests/unit/secret_managers
|
||||
echo tests/unit/vector_stores
|
||||
echo tests/unit/videos ;;
|
||||
proxy-db-auth-checks)
|
||||
|
|
|
|||
|
|
@ -354,6 +354,28 @@ workflows:
|
|||
- proxy-db-endpoints-and-responses
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-llm-vertex-ai
|
||||
flag: llm-vertex-ai
|
||||
shards: 2
|
||||
workers: 1
|
||||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-llm-other-providers
|
||||
flag: llm-other-providers
|
||||
shards: 3
|
||||
reruns: 2
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-integrations
|
||||
flag: integrations
|
||||
shards: 2
|
||||
reruns: 3
|
||||
base_ref: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.base.ref or "" >>
|
||||
pull_request_url: << pipeline.event.name == "pull_request" and pipeline.event.github.pull_request.url or "" >>
|
||||
- unit:
|
||||
name: unit-misc
|
||||
flag: misc
|
||||
|
|
|
|||
6
.github/merge-smoke-tests.json
vendored
6
.github/merge-smoke-tests.json
vendored
|
|
@ -1,8 +1,8 @@
|
|||
{
|
||||
"cases": {
|
||||
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"CHAT-JSON": "tests/unit/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/unit/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
|
||||
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
|
||||
"COST-EXPLICIT": "tests/unit/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
|
|
|
|||
6
.github/workflows/test-unit.yml
vendored
6
.github/workflows/test-unit.yml
vendored
|
|
@ -80,7 +80,8 @@ jobs:
|
|||
|
||||
- shard: integrations
|
||||
artifact-name: integrations
|
||||
test-path: "tests/test_litellm/integrations"
|
||||
test-path: ""
|
||||
unit-flag: integrations
|
||||
workers: 2
|
||||
reruns: 3
|
||||
timeout-minutes: 20
|
||||
|
|
@ -89,6 +90,7 @@ jobs:
|
|||
- shard: Vertex AI
|
||||
artifact-name: llm-vertex-ai
|
||||
test-path: "tests/test_litellm/llms/vertex_ai"
|
||||
unit-flag: llm-vertex-ai
|
||||
workers: 1
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -97,6 +99,7 @@ jobs:
|
|||
- shard: All Other Providers
|
||||
artifact-name: llm-other-providers
|
||||
test-path: "tests/test_litellm/llms --ignore=tests/test_litellm/llms/vertex_ai"
|
||||
unit-flag: llm-other-providers
|
||||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
|
|
@ -105,7 +108,6 @@ jobs:
|
|||
- shard: misc
|
||||
artifact-name: misc
|
||||
test-path: >-
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/ocr
|
||||
tests/test_litellm/passthrough
|
||||
|
|
|
|||
6
Makefile
6
Makefile
|
|
@ -314,7 +314,7 @@ test-unit: install-test-deps
|
|||
|
||||
# Matrix test targets (matching CI workflow groups)
|
||||
test-unit-llms: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/llms --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/llms --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-proxy-guardrails: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/proxy/guardrails tests/test_litellm/proxy/management_endpoints tests/test_litellm/proxy/management_helpers --tb=short -vv -n 4 --durations=20
|
||||
|
|
@ -326,13 +326,13 @@ test-unit-proxy-misc: install-test-deps
|
|||
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-integrations: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/unit/integrations --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-core-utils: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/litellm_core_utils --tb=short -vv -n 2 --durations=20
|
||||
|
||||
test-unit-other: install-test-deps
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/test_litellm/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
$(UV_RUN) pytest tests/test_litellm/caching tests/test_litellm/responses tests/unit/secret_managers tests/unit/vector_stores tests/unit/a2a_protocol tests/test_litellm/anthropic_interface tests/unit/completion_extras tests/unit/containers tests/unit/enterprise tests/unit/experimental_mcp_client tests/unit/google_genai tests/unit/images tests/unit/interactions tests/test_litellm/interactions tests/test_litellm/passthrough tests/test_litellm/router_strategy tests/test_litellm/router_utils tests/unit/types --tb=short -vv -n 4 --durations=20
|
||||
|
||||
test-unit-root: install-test-deps
|
||||
$(UV_RUN) pytest tests/unit/test_*.py tests/test_litellm/test_*.py --tb=short -vv -n 4 --durations=20
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
|
||||
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/test_litellm/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard
|
||||
|
||||
The first eleven rows need only `callbacks: ["prometheus"]`. The last three rows and the `litellm_admission_*` panels are emitted by other subsystems and stay empty until those are on: the service callback row needs `service_callback: ["prometheus_system"]` in `litellm_settings`, the circuit breaker row needs a Redis cache, the cleanup row needs spend log retention, and admission control needs its middleware enabled. Within the base rows, many panels only fill in once the matching feature is in use: budgets need keys, teams, users or orgs with `max_budget` set, cache panels need caching on, guardrail and MCP panels need those features configured, deployment health needs the router with more than one deployment or a failure to record, and `litellm_in_flight_requests` needs traffic at scrape time. An empty panel for a feature you do not use is expected
|
||||
|
||||
|
|
|
|||
16
litellm-rust/Cargo.lock
generated
16
litellm-rust/Cargo.lock
generated
|
|
@ -3166,10 +3166,11 @@ dependencies = [
|
|||
"aws-smithy-types",
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"proptest",
|
||||
"rstest",
|
||||
"sse-stream",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tokio-util",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -5468,19 +5469,6 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "sse-stream"
|
||||
version = "0.2.6"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "c25ac7aff0abd1dbc474536e40416e1102c7dd9bfba0b9861c6d357f835dcfb4"
|
||||
dependencies = [
|
||||
"bytes",
|
||||
"futures-util",
|
||||
"http-body 1.1.0",
|
||||
"http-body-util",
|
||||
"pin-project-lite",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "stable_deref_trait"
|
||||
version = "1.2.1"
|
||||
|
|
|
|||
|
|
@ -398,7 +398,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
|
|||
api_key: Some("sk".into()),
|
||||
api_base: Some(upstream.uri()),
|
||||
shaping: MessagesShaping {
|
||||
capabilities: capabilities.clone(),
|
||||
capabilities,
|
||||
drop_params,
|
||||
..MessagesShaping::default()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -8,16 +8,17 @@ repository.workspace = true
|
|||
[features]
|
||||
default = ["aws", "sse"]
|
||||
aws = ["dep:aws-smithy-eventstream", "dep:aws-smithy-types"]
|
||||
sse = ["dep:sse-stream"]
|
||||
sse = []
|
||||
|
||||
[dependencies]
|
||||
aws-smithy-eventstream = { version = "=0.61.4", optional = true }
|
||||
aws-smithy-types = { version = "1.6.1", optional = true }
|
||||
bytes = "1"
|
||||
futures-util.workspace = true
|
||||
sse-stream = { version = "=0.2.6", optional = true }
|
||||
thiserror.workspace = true
|
||||
tokio-util = { version = "0.7", features = ["codec", "io"] }
|
||||
|
||||
[dev-dependencies]
|
||||
proptest.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,66 +1,47 @@
|
|||
use bytes::{Buf, Bytes, BytesMut};
|
||||
use futures_util::{Stream, StreamExt};
|
||||
use aws_smithy_eventstream::frame::{read_message_from, write_message_to};
|
||||
pub use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
|
||||
use bytes::BytesMut;
|
||||
use tokio_util::codec::{Decoder, Encoder};
|
||||
|
||||
use aws_smithy_eventstream::frame::read_message_from;
|
||||
use aws_smithy_types::event_stream::Header;
|
||||
|
||||
use crate::{Error, Framer};
|
||||
use crate::EventStreamError;
|
||||
|
||||
const MIN_FRAME_BYTES: usize = 16;
|
||||
const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct AwsEventStreamFrame {
|
||||
pub headers: Vec<Header>,
|
||||
pub payload: Bytes,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct AwsEventStreamFramer;
|
||||
pub struct AwsEventStreamCodec;
|
||||
|
||||
impl Framer for AwsEventStreamFramer {
|
||||
type Frame = AwsEventStreamFrame;
|
||||
impl Decoder for AwsEventStreamCodec {
|
||||
type Item = Message;
|
||||
type Error = EventStreamError;
|
||||
|
||||
fn frame<S, B, E>(self, input: S) -> impl Stream<Item = Result<Self::Frame, Error>> + Send
|
||||
where
|
||||
S: Stream<Item = Result<B, E>> + Send,
|
||||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
futures_util::stream::try_unfold(
|
||||
(Box::pin(input), BytesMut::new()),
|
||||
|(mut input, mut buffer)| async move {
|
||||
loop {
|
||||
if buffer.len() >= 4 {
|
||||
let length = (&buffer[..4]).get_u32() as usize;
|
||||
if !(16..=MAX_FRAME_BYTES).contains(&length) {
|
||||
return Err(Error::InvalidLength(length));
|
||||
}
|
||||
if buffer.len() >= length {
|
||||
let raw = buffer.split_to(length).freeze();
|
||||
let message = read_message_from(raw)?;
|
||||
let frame = AwsEventStreamFrame {
|
||||
headers: message.headers().to_vec(),
|
||||
payload: message.payload().clone(),
|
||||
};
|
||||
return Ok(Some((frame, (input, buffer))));
|
||||
}
|
||||
}
|
||||
match input.next().await {
|
||||
Some(Ok(mut chunk)) => {
|
||||
while chunk.has_remaining() {
|
||||
let bytes = chunk.chunk();
|
||||
buffer.extend_from_slice(bytes);
|
||||
let length = bytes.len();
|
||||
chunk.advance(length);
|
||||
}
|
||||
}
|
||||
Some(Err(error)) => return Err(Error::Body(Box::new(error))),
|
||||
None if buffer.is_empty() => return Ok(None),
|
||||
None => return Err(Error::Truncated),
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
.fuse()
|
||||
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Message>, EventStreamError> {
|
||||
let Some(prefix) = src.first_chunk::<4>() else {
|
||||
return Ok(None);
|
||||
};
|
||||
let length = u32::from_be_bytes(*prefix) as usize;
|
||||
if !(MIN_FRAME_BYTES..=MAX_FRAME_BYTES).contains(&length) {
|
||||
return Err(EventStreamError::InvalidLength(length));
|
||||
}
|
||||
if src.len() < length {
|
||||
return Ok(None);
|
||||
}
|
||||
Ok(Some(read_message_from(src.split_to(length).freeze())?))
|
||||
}
|
||||
|
||||
fn decode_eof(&mut self, src: &mut BytesMut) -> Result<Option<Message>, EventStreamError> {
|
||||
match self.decode(src)? {
|
||||
Some(message) => Ok(Some(message)),
|
||||
None if src.is_empty() => Ok(None),
|
||||
None => Err(EventStreamError::Truncated),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Encoder<Message> for AwsEventStreamCodec {
|
||||
type Error = EventStreamError;
|
||||
|
||||
fn encode(&mut self, message: Message, dst: &mut BytesMut) -> Result<(), EventStreamError> {
|
||||
Ok(write_message_to(&message, dst)?)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,17 +1,21 @@
|
|||
#[cfg(feature = "sse")]
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[cfg(feature = "sse")]
|
||||
#[error("SSE framing failed: {0}")]
|
||||
Sse(#[from] sse_stream::Error),
|
||||
#[cfg(feature = "aws")]
|
||||
#[error("AWS EventStream framing failed: {0}")]
|
||||
Aws(#[from] aws_smithy_eventstream::error::Error),
|
||||
pub enum SseError {
|
||||
#[error("body stream failed: {0}")]
|
||||
Body(#[source] Box<dyn std::error::Error + Send + Sync>),
|
||||
#[cfg(feature = "aws")]
|
||||
Body(#[from] std::io::Error),
|
||||
#[error("SSE field is not UTF-8: {0}")]
|
||||
InvalidUtf8(#[from] std::str::Utf8Error),
|
||||
}
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum EventStreamError {
|
||||
#[error("body stream failed: {0}")]
|
||||
Body(#[from] std::io::Error),
|
||||
#[error("invalid AWS EventStream frame length: {0}")]
|
||||
InvalidLength(usize),
|
||||
#[cfg(feature = "aws")]
|
||||
#[error("truncated AWS EventStream frame")]
|
||||
Truncated,
|
||||
#[error("malformed AWS EventStream frame: {0}")]
|
||||
Malformed(#[from] aws_smithy_eventstream::error::Error),
|
||||
}
|
||||
|
|
|
|||
21
litellm-rust/crates/framer/src/framed.rs
Normal file
21
litellm-rust/crates/framer/src/framed.rs
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
use std::io;
|
||||
|
||||
use bytes::Buf;
|
||||
use futures_util::{Stream, StreamExt, TryStreamExt};
|
||||
use tokio_util::{
|
||||
codec::{Decoder, FramedRead},
|
||||
io::StreamReader,
|
||||
};
|
||||
|
||||
pub fn frames<S, B, E, D>(
|
||||
input: S,
|
||||
codec: D,
|
||||
) -> impl Stream<Item = Result<D::Item, D::Error>> + Send
|
||||
where
|
||||
S: Stream<Item = Result<B, E>> + Send,
|
||||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
D: Decoder + Send,
|
||||
{
|
||||
FramedRead::new(StreamReader::new(input.map_err(io::Error::other)), codec).fuse()
|
||||
}
|
||||
|
|
@ -1,8 +1,8 @@
|
|||
mod error;
|
||||
mod framer;
|
||||
mod framed;
|
||||
|
||||
pub use error::*;
|
||||
pub use framer::*;
|
||||
pub use framed::frames;
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
pub mod aws_event_stream;
|
||||
|
|
|
|||
|
|
@ -1,43 +1,170 @@
|
|||
use futures_util::{Stream, StreamExt};
|
||||
use std::str;
|
||||
|
||||
use crate::{Error, Framer};
|
||||
use bytes::{Buf, BufMut, BytesMut};
|
||||
use tokio_util::codec::{Decoder, Encoder};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub struct SseFrame {
|
||||
use crate::SseError;
|
||||
|
||||
const BOM: &[u8] = b"\xEF\xBB\xBF";
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq)]
|
||||
pub struct SseEvent {
|
||||
pub event: Option<String>,
|
||||
pub data: Option<String>,
|
||||
pub data: String,
|
||||
pub id: Option<String>,
|
||||
pub retry: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
pub struct SseFramer;
|
||||
pub struct SseCodec {
|
||||
past_bom: bool,
|
||||
}
|
||||
|
||||
impl Framer for SseFramer {
|
||||
type Frame = SseFrame;
|
||||
impl Decoder for SseCodec {
|
||||
type Item = SseEvent;
|
||||
type Error = SseError;
|
||||
|
||||
fn frame<S, B, E>(self, input: S) -> impl Stream<Item = Result<SseFrame, Error>> + Send
|
||||
where
|
||||
S: Stream<Item = Result<B, E>> + Send,
|
||||
B: bytes::Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
let frames = Box::pin(sse_stream::SseStream::from_bytes_stream(input));
|
||||
futures_util::stream::try_unfold(frames, |mut frames| async move {
|
||||
let Some(frame) = frames.next().await else {
|
||||
return Ok(None);
|
||||
};
|
||||
let frame = frame?;
|
||||
Ok(Some((
|
||||
SseFrame {
|
||||
event: frame.event,
|
||||
data: frame.data,
|
||||
id: frame.id,
|
||||
retry: frame.retry,
|
||||
},
|
||||
frames,
|
||||
)))
|
||||
})
|
||||
.fuse()
|
||||
fn decode(&mut self, src: &mut BytesMut) -> Result<Option<SseEvent>, SseError> {
|
||||
if !self.skip_bom(src) {
|
||||
return Ok(None);
|
||||
}
|
||||
while let Some(end) = block_end(src) {
|
||||
let block = src.split_to(end);
|
||||
let pending = lines(&block)
|
||||
.map(|(line, _)| line)
|
||||
.take_while(|line| !line.is_empty())
|
||||
.try_fold(Pending::default(), Pending::apply)?;
|
||||
if let Some(event) = pending.dispatch() {
|
||||
return Ok(Some(event));
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn decode_eof(&mut self, _pending: &mut BytesMut) -> Result<Option<SseEvent>, SseError> {
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
|
||||
impl SseCodec {
|
||||
fn skip_bom(&mut self, src: &mut BytesMut) -> bool {
|
||||
if self.past_bom {
|
||||
return true;
|
||||
}
|
||||
if src.starts_with(BOM) {
|
||||
src.advance(BOM.len());
|
||||
} else if BOM.starts_with(src) {
|
||||
return false;
|
||||
}
|
||||
self.past_bom = true;
|
||||
true
|
||||
}
|
||||
}
|
||||
|
||||
fn block_end(bytes: &[u8]) -> Option<usize> {
|
||||
lines(bytes)
|
||||
.find(|(line, _)| line.is_empty())
|
||||
.map(|(_, end)| end)
|
||||
}
|
||||
|
||||
fn lines(bytes: &[u8]) -> impl Iterator<Item = (&[u8], usize)> {
|
||||
let mut cursor: usize = 0;
|
||||
std::iter::from_fn(move || {
|
||||
let rest = &bytes[cursor..];
|
||||
let end = rest.iter().position(|byte| matches!(byte, b'\n' | b'\r'))?;
|
||||
cursor += end + terminator_len(&rest[end..]);
|
||||
Some((&rest[..end], cursor))
|
||||
})
|
||||
}
|
||||
|
||||
fn terminator_len(terminated: &[u8]) -> usize {
|
||||
match terminated {
|
||||
[b'\r', b'\n', ..] => 2,
|
||||
_ => 1,
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Default)]
|
||||
struct Pending {
|
||||
event: Option<String>,
|
||||
data: Option<String>,
|
||||
id: Option<String>,
|
||||
retry: Option<u64>,
|
||||
}
|
||||
|
||||
impl Pending {
|
||||
fn apply(self, line: &[u8]) -> Result<Self, SseError> {
|
||||
let (name, value) = split_field(line);
|
||||
Ok(match name {
|
||||
b"event" => Self {
|
||||
event: Some(str::from_utf8(value)?.to_owned()),
|
||||
..self
|
||||
},
|
||||
b"data" => Self {
|
||||
data: Some(append_data(self.data, str::from_utf8(value)?)),
|
||||
..self
|
||||
},
|
||||
b"id" if !value.contains(&0) => Self {
|
||||
id: Some(str::from_utf8(value)?.to_owned()),
|
||||
..self
|
||||
},
|
||||
b"retry" => Self {
|
||||
retry: parse_retry(value).or(self.retry),
|
||||
..self
|
||||
},
|
||||
_ => self,
|
||||
})
|
||||
}
|
||||
|
||||
fn dispatch(self) -> Option<SseEvent> {
|
||||
Some(SseEvent {
|
||||
event: self.event,
|
||||
data: self.data?,
|
||||
id: self.id,
|
||||
retry: self.retry,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn split_field(line: &[u8]) -> (&[u8], &[u8]) {
|
||||
let Some(colon) = line.iter().position(|byte| *byte == b':') else {
|
||||
return (line, &[]);
|
||||
};
|
||||
let value = &line[colon + 1..];
|
||||
(&line[..colon], value.strip_prefix(b" ").unwrap_or(value))
|
||||
}
|
||||
|
||||
fn append_data(buffer: Option<String>, line: &str) -> String {
|
||||
match buffer {
|
||||
Some(existing) => format!("{existing}\n{line}"),
|
||||
None => line.to_owned(),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_retry(value: &[u8]) -> Option<u64> {
|
||||
if !value.iter().all(u8::is_ascii_digit) {
|
||||
return None;
|
||||
}
|
||||
str::from_utf8(value).ok()?.parse().ok()
|
||||
}
|
||||
|
||||
impl Encoder<SseEvent> for SseCodec {
|
||||
type Error = SseError;
|
||||
|
||||
fn encode(&mut self, event: SseEvent, dst: &mut BytesMut) -> Result<(), SseError> {
|
||||
if let Some(name) = event.event {
|
||||
dst.put_slice(format!("event: {name}\n").as_bytes());
|
||||
}
|
||||
for line in event.data.split('\n') {
|
||||
dst.put_slice(format!("data: {line}\n").as_bytes());
|
||||
}
|
||||
if let Some(id) = event.id {
|
||||
dst.put_slice(format!("id: {id}\n").as_bytes());
|
||||
}
|
||||
if let Some(retry) = event.retry {
|
||||
dst.put_slice(format!("retry: {retry}\n").as_bytes());
|
||||
}
|
||||
dst.put_u8(b'\n');
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,89 +4,174 @@ mod support;
|
|||
|
||||
use std::io;
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer};
|
||||
use litellm_framing::{Error, Framer};
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_framing::{
|
||||
EventStreamError,
|
||||
aws_event_stream::{AwsEventStreamCodec, Header, HeaderValue, Message},
|
||||
frames,
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
use rstest::{fixture, rstest};
|
||||
use support::{body_cause, cut_at, encode_all, every, input, runtime};
|
||||
|
||||
use support::encode;
|
||||
|
||||
async fn collect_aws(bytes: &[u8], chunk_size: usize) -> Result<Vec<AwsEventStreamFrame>, Error> {
|
||||
AwsEventStreamFramer
|
||||
.frame(futures_util::stream::iter(
|
||||
bytes.chunks(chunk_size).map(Ok::<_, io::Error>),
|
||||
))
|
||||
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<Message>, EventStreamError> {
|
||||
frames(input(pieces), AwsEventStreamCodec)
|
||||
.try_collect()
|
||||
.await
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn two_frames() -> Vec<u8> {
|
||||
[encode(b"\xff\x00"), encode(b"second")].concat()
|
||||
fn message(payload: &[u8]) -> Message {
|
||||
Message::new(Bytes::copy_from_slice(payload))
|
||||
.add_header(Header::new(
|
||||
":event-type",
|
||||
HeaderValue::String("payload".into()),
|
||||
))
|
||||
.add_header(Header::new("sequence", HeaderValue::Int32(7)))
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn payload_frame() -> Vec<u8> {
|
||||
encode(b"payload")
|
||||
encode_all(AwsEventStreamCodec, [message(b"payload")])
|
||||
}
|
||||
|
||||
fn header_value() -> impl Strategy<Value = HeaderValue> {
|
||||
prop_oneof![
|
||||
"[a-z]{0,8}".prop_map(|text| HeaderValue::String(text.into())),
|
||||
any::<i32>().prop_map(HeaderValue::Int32),
|
||||
any::<bool>().prop_map(HeaderValue::Bool),
|
||||
proptest::collection::vec(any::<u8>(), 0..8)
|
||||
.prop_map(|bytes| HeaderValue::ByteArray(bytes.into())),
|
||||
]
|
||||
}
|
||||
|
||||
fn arbitrary_message() -> impl Strategy<Value = Message> {
|
||||
(
|
||||
proptest::collection::vec(("[a-z:-]{1,12}", header_value()), 0..3),
|
||||
proptest::collection::vec(any::<u8>(), 0..32),
|
||||
)
|
||||
.prop_map(|(headers, payload)| {
|
||||
headers.into_iter().fold(
|
||||
Message::new(Bytes::from(payload)),
|
||||
|message, (name, value)| message.add_header(Header::new(name, value)),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#[test]
|
||||
fn any_messages_survive_a_round_trip_through_any_cuts(
|
||||
messages in proptest::collection::vec(arbitrary_message(), 1..4),
|
||||
cuts in proptest::collection::vec(0_usize..512, 0..4),
|
||||
) {
|
||||
let wire = encode_all(AwsEventStreamCodec, messages.clone());
|
||||
let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap();
|
||||
prop_assert_eq!(decoded, messages);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(1)]
|
||||
#[case(3)]
|
||||
#[case(12)]
|
||||
#[case(usize::MAX)]
|
||||
#[case::prelude_crc(8)]
|
||||
#[case::message_crc(usize::MAX)]
|
||||
#[tokio::test]
|
||||
async fn fragmented_and_coalesced_frames_preserve_typed_headers_and_binary_payloads(
|
||||
two_frames: Vec<u8>,
|
||||
#[case] chunk_size: usize,
|
||||
) {
|
||||
let chunk_size = chunk_size.min(two_frames.len());
|
||||
let frames = collect_aws(&two_frames, chunk_size).await.unwrap();
|
||||
assert_eq!(frames.len(), 2);
|
||||
assert_eq!(frames[0].payload, &b"\xff\x00"[..]);
|
||||
assert_eq!(frames[1].payload, "second");
|
||||
assert_eq!(
|
||||
frames[0].headers[0].value().as_string().unwrap().as_str(),
|
||||
"payload"
|
||||
);
|
||||
assert_eq!(frames[0].headers[1].value().as_int32(), Ok(7));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(8)]
|
||||
#[case(0)]
|
||||
#[tokio::test]
|
||||
async fn rejects_corrupt_crcs(payload_frame: Vec<u8>, #[case] index: usize) {
|
||||
let corrupt_index = if index == 0 {
|
||||
payload_frame.len() - 1
|
||||
} else {
|
||||
index
|
||||
};
|
||||
async fn a_corrupt_crc_is_malformed(payload_frame: Vec<u8>, #[case] index: usize) {
|
||||
let mut corrupt = payload_frame;
|
||||
corrupt[corrupt_index] ^= 1;
|
||||
assert!(matches!(collect_aws(&corrupt, 3).await, Err(Error::Aws(_))));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(0_u32)]
|
||||
#[case(15)]
|
||||
#[case(u32::MAX)]
|
||||
#[tokio::test]
|
||||
async fn rejects_invalid_lengths(#[case] length: u32) {
|
||||
let flipped = index.min(corrupt.len() - 1);
|
||||
corrupt[flipped] ^= 1;
|
||||
assert!(matches!(
|
||||
collect_aws(&length.to_be_bytes(), 1).await,
|
||||
Err(Error::InvalidLength(_))
|
||||
collect(every(&corrupt, 3)).await,
|
||||
Err(EventStreamError::Malformed(_))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(1)]
|
||||
#[case(3)]
|
||||
#[case(5)]
|
||||
#[case::zero(0)]
|
||||
#[case::below_minimum(15)]
|
||||
#[case::above_maximum(16 * 1024 * 1024 + 1)]
|
||||
#[case::u32_max(u32::MAX)]
|
||||
#[tokio::test]
|
||||
async fn rejects_truncation(payload_frame: Vec<u8>, #[case] end: usize) {
|
||||
async fn a_length_outside_the_frame_bounds_fails_before_buffering(#[case] length: u32) {
|
||||
assert!(matches!(
|
||||
collect_aws(&payload_frame[..end], 1).await,
|
||||
Err(Error::Truncated)
|
||||
collect(every(&length.to_be_bytes(), 1)).await,
|
||||
Err(EventStreamError::InvalidLength(seen)) if seen == length as usize
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::before_the_length(1)]
|
||||
#[case::inside_the_prelude(5)]
|
||||
#[case::one_byte_short(usize::MAX)]
|
||||
#[tokio::test]
|
||||
async fn eof_inside_a_frame_is_truncation(payload_frame: Vec<u8>, #[case] end: usize) {
|
||||
let end = end.min(payload_frame.len() - 1);
|
||||
assert!(matches!(
|
||||
collect(every(&payload_frame[..end], 1)).await,
|
||||
Err(EventStreamError::Truncated)
|
||||
));
|
||||
}
|
||||
|
||||
const FRAME_OVERHEAD_BYTES: usize = 16;
|
||||
const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024;
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_frame_at_exactly_the_maximum_length_decodes() {
|
||||
let largest = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES]);
|
||||
let wire = encode_all(AwsEventStreamCodec, [largest.clone()]);
|
||||
assert_eq!(wire.len(), MAX_FRAME_BYTES);
|
||||
assert_eq!(collect(every(&wire, 1 << 20)).await.unwrap(), vec![largest]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_frame_one_byte_over_the_maximum_length_is_rejected_by_its_prelude() {
|
||||
let oversized = Message::new(vec![0xAB; MAX_FRAME_BYTES - FRAME_OVERHEAD_BYTES + 1]);
|
||||
let wire = encode_all(AwsEventStreamCodec, [oversized]);
|
||||
assert!(matches!(
|
||||
collect(every(&wire[..4], 1)).await,
|
||||
Err(EventStreamError::InvalidLength(length)) if length == MAX_FRAME_BYTES + 1
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn an_empty_body_yields_nothing() {
|
||||
assert_eq!(collect(vec![]).await.unwrap(), vec![]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_complete_frame_precedes_a_truncated_following_frame() {
|
||||
let wire = encode_all(AwsEventStreamCodec, [message(b"first"), message(b"second")]);
|
||||
let mut messages = Box::pin(frames(
|
||||
input(every(&wire[..wire.len() - 1], 3)),
|
||||
AwsEventStreamCodec,
|
||||
));
|
||||
|
||||
assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first"));
|
||||
assert!(matches!(
|
||||
messages.next().await,
|
||||
Some(Err(EventStreamError::Truncated))
|
||||
));
|
||||
assert!(messages.next().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn a_body_error_after_a_complete_frame_preserves_its_cause() {
|
||||
let first = encode_all(AwsEventStreamCodec, [message(b"first")]);
|
||||
let mut messages = Box::pin(frames(
|
||||
stream::iter([
|
||||
Ok(cut_at(&first, [5])[0].clone()),
|
||||
Ok(cut_at(&first, [5])[1].clone()),
|
||||
Ok(Bytes::from_static(b"\0\0\0")),
|
||||
Err(io::Error::new(io::ErrorKind::ConnectionReset, "reset")),
|
||||
]),
|
||||
AwsEventStreamCodec,
|
||||
));
|
||||
|
||||
assert_eq!(messages.next().await.unwrap().unwrap(), message(b"first"));
|
||||
let Some(Err(EventStreamError::Body(body))) = messages.next().await else {
|
||||
panic!("the body error surfaces");
|
||||
};
|
||||
assert_eq!(
|
||||
body_cause::<io::Error>(&body).unwrap().kind(),
|
||||
io::ErrorKind::ConnectionReset
|
||||
);
|
||||
assert!(messages.next().await.is_none());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,28 +2,64 @@
|
|||
|
||||
mod support;
|
||||
|
||||
use std::io;
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_framing::{
|
||||
EventStreamError, SseError,
|
||||
aws_event_stream::{AwsEventStreamCodec, Message},
|
||||
frames,
|
||||
sse::{SseCodec, SseEvent},
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
use support::{body_cause, cut_at, encode_all, every, input, runtime};
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use litellm_framing::Framer;
|
||||
use litellm_framing::aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer};
|
||||
use litellm_framing::sse::SseFramer;
|
||||
fn delta(data: &str) -> SseEvent {
|
||||
SseEvent {
|
||||
event: Some("delta".into()),
|
||||
data: data.into(),
|
||||
id: Some("7".into()),
|
||||
retry: None,
|
||||
}
|
||||
}
|
||||
|
||||
use support::encode;
|
||||
fn envelopes(payloads: Vec<Bytes>) -> Vec<u8> {
|
||||
encode_all(AwsEventStreamCodec, payloads.into_iter().map(Message::new))
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#[test]
|
||||
fn an_sse_event_cut_anywhere_across_envelopes_is_reassembled(cut in 0_usize..64, chunk in 1_usize..8) {
|
||||
let sse = encode_all(SseCodec::default(), [delta("hello")]);
|
||||
let wire = envelopes(cut_at(&sse, [cut.min(sse.len())]));
|
||||
let events = runtime().block_on(async {
|
||||
let payloads = frames(input(every(&wire, chunk)), AwsEventStreamCodec)
|
||||
.map_ok(|message| message.payload().clone());
|
||||
frames(payloads, SseCodec::default()).try_collect::<Vec<_>>().await
|
||||
})
|
||||
.unwrap();
|
||||
prop_assert_eq!(events, vec![delta("hello")]);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn hosting_payloads_feed_the_same_sse_framer_across_envelope_boundaries() {
|
||||
let bytes = [encode(b"event: delta\ndata: hel"), encode(b"lo\nid: 7\n\n")].concat();
|
||||
let envelopes = AwsEventStreamFramer.frame(futures_util::stream::iter(
|
||||
bytes.chunks(3).map(Ok::<_, io::Error>),
|
||||
async fn a_truncated_envelope_after_an_sse_event_keeps_the_event_and_its_cause() {
|
||||
let complete = encode_all(SseCodec::default(), [delta("complete")]);
|
||||
let incomplete = encode_all(SseCodec::default(), [delta("incomplete")]);
|
||||
let wire = envelopes(vec![complete.into(), incomplete.into()]);
|
||||
let payloads = frames(
|
||||
input(every(&wire[..wire.len() - 1], 3)),
|
||||
AwsEventStreamCodec,
|
||||
)
|
||||
.map_ok(|message| message.payload().clone());
|
||||
let mut events = Box::pin(frames(payloads, SseCodec::default()));
|
||||
|
||||
assert_eq!(events.next().await.unwrap().unwrap(), delta("complete"));
|
||||
let Some(Err(SseError::Body(body))) = events.next().await else {
|
||||
panic!("the envelope error surfaces through the SSE layer");
|
||||
};
|
||||
assert!(matches!(
|
||||
body_cause::<EventStreamError>(&body),
|
||||
Some(EventStreamError::Truncated)
|
||||
));
|
||||
let frames = SseFramer
|
||||
.frame(envelopes.map_ok(|frame: AwsEventStreamFrame| frame.payload))
|
||||
.try_collect::<Vec<_>>()
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(frames.len(), 1);
|
||||
assert_eq!(frames[0].event.as_deref(), Some("delta"));
|
||||
assert_eq!(frames[0].data.as_deref(), Some("hello"));
|
||||
assert_eq!(frames[0].id.as_deref(), Some("7"));
|
||||
assert!(events.next().await.is_none());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,67 +1,169 @@
|
|||
#![cfg(feature = "sse")]
|
||||
|
||||
mod support;
|
||||
|
||||
use std::io;
|
||||
|
||||
use futures_util::{StreamExt, TryStreamExt};
|
||||
use litellm_framing::sse::{SseFrame, SseFramer};
|
||||
use litellm_framing::{Error, Framer};
|
||||
use bytes::Bytes;
|
||||
use futures_util::{StreamExt, TryStreamExt, stream};
|
||||
use litellm_framing::{
|
||||
SseError, frames,
|
||||
sse::{SseCodec, SseEvent},
|
||||
};
|
||||
use proptest::prelude::*;
|
||||
use rstest::rstest;
|
||||
use support::{body_cause, cut_at, encode_all, every, input, runtime};
|
||||
|
||||
async fn collect_sse(chunks: &[&[u8]]) -> Result<Vec<SseFrame>, Error> {
|
||||
SseFramer
|
||||
.frame(futures_util::stream::iter(
|
||||
chunks.iter().copied().map(Ok::<_, io::Error>),
|
||||
))
|
||||
async fn collect(pieces: Vec<Bytes>) -> Result<Vec<SseEvent>, SseError> {
|
||||
frames(input(pieces), SseCodec::default())
|
||||
.try_collect()
|
||||
.await
|
||||
}
|
||||
|
||||
fn event(name: Option<&str>, data: &str) -> SseEvent {
|
||||
SseEvent {
|
||||
event: name.map(str::to_owned),
|
||||
data: data.to_owned(),
|
||||
id: None,
|
||||
retry: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn sse_event() -> impl Strategy<Value = SseEvent> {
|
||||
(
|
||||
proptest::option::of("[^\r\n\0]{0,8}"),
|
||||
"[^\r\0]{0,16}",
|
||||
proptest::option::of("[^\r\n\0]{0,8}"),
|
||||
proptest::option::of(any::<u64>()),
|
||||
)
|
||||
.prop_map(|(event, data, id, retry)| SseEvent {
|
||||
event,
|
||||
data,
|
||||
id,
|
||||
retry,
|
||||
})
|
||||
}
|
||||
|
||||
fn terminators() -> impl Strategy<Value = &'static [u8]> {
|
||||
prop_oneof![Just(&b"\n"[..]), Just(&b"\r\n"[..]), Just(&b"\r"[..])]
|
||||
}
|
||||
|
||||
proptest! {
|
||||
#[test]
|
||||
fn any_events_survive_a_round_trip_through_any_terminator_and_any_cuts(
|
||||
events in proptest::collection::vec(sse_event(), 1..4),
|
||||
terminator in terminators(),
|
||||
cuts in proptest::collection::vec(0_usize..256, 0..4),
|
||||
bom in any::<bool>(),
|
||||
) {
|
||||
let lf_wire = encode_all(SseCodec::default(), events.clone());
|
||||
let body: Vec<u8> = lf_wire
|
||||
.iter()
|
||||
.flat_map(|byte| if *byte == b'\n' { terminator.to_vec() } else { vec![*byte] })
|
||||
.collect();
|
||||
let wire = if bom { [&b"\xEF\xBB\xBF"[..], &body].concat() } else { body };
|
||||
let decoded = runtime().block_on(collect(cut_at(&wire, cuts))).unwrap();
|
||||
prop_assert_eq!(decoded, events);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(
|
||||
&[&b":ping\r\nevent: delta\r\nid: 7\r\nretry: 10\r\ndata: \xe2"[..], &b"\x82"[..], &b"\xac\r"[..], &b"\ndata: next\r\n\r"[..], &b"\ndata: [DONE]\n\n"[..]],
|
||||
vec![
|
||||
SseFrame {
|
||||
event: Some("delta".into()),
|
||||
data: Some("€\nnext".into()),
|
||||
id: Some("7".into()),
|
||||
retry: Some(10),
|
||||
},
|
||||
SseFrame {
|
||||
event: None,
|
||||
data: Some("[DONE]".into()),
|
||||
id: None,
|
||||
retry: None,
|
||||
},
|
||||
]
|
||||
)]
|
||||
#[case::comment(b":ping\ndata: x\n\n")]
|
||||
#[case::unknown_field(b"vendor: 1\ndata: x\n\n")]
|
||||
#[case::field_without_colon(b"garbage\ndata: x\n\n")]
|
||||
#[case::retry_with_non_digits(b"retry: soon\ndata: x\n\n")]
|
||||
#[case::retry_with_a_sign(b"retry: +5\ndata: x\n\n")]
|
||||
#[case::retry_without_a_value(b"retry:\ndata: x\n\n")]
|
||||
#[case::id_with_nul(b"id: a\0b\ndata: x\n\n")]
|
||||
#[tokio::test]
|
||||
async fn fragmented_utf8_crlf_and_multiline_data_retain_metadata_and_sentinel(
|
||||
#[case] chunks: &[&[u8]],
|
||||
#[case] expected: Vec<SseFrame>,
|
||||
async fn lines_the_spec_ignores_do_not_change_the_event(#[case] wire: &[u8]) {
|
||||
assert_eq!(
|
||||
collect(every(wire, 1)).await.unwrap(),
|
||||
vec![event(None, "x")]
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::no_data_at_all(b"event: ping\nid: 1\n\ndata: x\n\n", vec![event(None, "x")])]
|
||||
#[case::empty_data_field(b"data:\n\n", vec![event(None, "")])]
|
||||
#[case::one_leading_space_stripped(b"data: x\n\n", vec![event(None, " x")])]
|
||||
#[case::multiline_data(b"data: a\ndata: b\ndata:\n\n", vec![event(None, "a\nb\n")])]
|
||||
#[case::last_event_name_wins(b"event: a\nevent: b\ndata: x\n\n", vec![event(Some("b"), "x")])]
|
||||
#[case::last_retry_wins(b"retry: 1\nretry: 2\ndata: x\n\n", vec![SseEvent { retry: Some(2), ..event(None, "x") }])]
|
||||
#[case::split_utf8_across_lines_is_not_joined(b"data: \xe2\x82\xac\ndata: \xe2\x82\xac\n\n", vec![event(None, "€\n€")])]
|
||||
#[tokio::test]
|
||||
async fn dispatch_follows_the_data_buffer(#[case] wire: &[u8], #[case] expected: Vec<SseEvent>) {
|
||||
assert_eq!(collect(every(wire, 1)).await.unwrap(), expected);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::unterminated_single(b"data: partial\n", vec![])]
|
||||
#[case::unterminated_tail_after_complete(b"data: complete\n\ndata: unfinished\n", vec![event(None, "complete")])]
|
||||
#[case::lone_cr_terminates_at_eof(b"data: x\r\r", vec![event(None, "x")])]
|
||||
#[case::lone_cr_line_then_eof(b"data: x\r", vec![])]
|
||||
#[tokio::test]
|
||||
async fn eof_dispatches_only_terminated_events(
|
||||
#[case] wire: &[u8],
|
||||
#[case] expected: Vec<SseEvent>,
|
||||
) {
|
||||
assert_eq!(collect_sse(chunks).await.unwrap(), expected);
|
||||
assert_eq!(
|
||||
collect(vec![Bytes::copy_from_slice(wire)]).await.unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::inside_the_first_line(vec![&b"data: a\r"[..], &b"\ndata: b\r\n\r\n"[..]])]
|
||||
#[case::inside_the_blank_line(vec![&b"data: a\r\ndata: b\r\n\r"[..], &b"\n"[..]])]
|
||||
#[tokio::test]
|
||||
async fn a_crlf_split_across_chunks_is_one_terminator(#[case] pieces: Vec<&[u8]>) {
|
||||
let pieces = pieces.into_iter().map(Bytes::copy_from_slice).collect();
|
||||
assert_eq!(collect(pieces).await.unwrap(), vec![event(None, "a\nb")]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn eof_does_not_dispatch_an_unterminated_frame() {
|
||||
assert!(collect_sse(&[b"data: partial\n"]).await.unwrap().is_empty());
|
||||
async fn a_bom_is_stripped_only_at_the_start_of_the_stream() {
|
||||
let wire = b"\xEF\xBB\xBFdata: a\n\n\xEF\xBB\xBFdata: b\ndata: c\n\n";
|
||||
let decoded = collect(every(wire, 2)).await.unwrap();
|
||||
assert_eq!(decoded, vec![event(None, "a"), event(None, "c")]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_utf8_in_a_field_fails_after_earlier_events_and_terminates() {
|
||||
let mut events = Box::pin(frames(
|
||||
input(every(b"data: ok\n\ndata: \xff\n\n", 3)),
|
||||
SseCodec::default(),
|
||||
));
|
||||
|
||||
assert_eq!(events.next().await.unwrap().unwrap(), event(None, "ok"));
|
||||
assert!(matches!(
|
||||
events.next().await,
|
||||
Some(Err(SseError::InvalidUtf8(_)))
|
||||
));
|
||||
assert!(events.next().await.is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(io::ErrorKind::ConnectionReset)]
|
||||
#[case(io::ErrorKind::UnexpectedEof)]
|
||||
#[tokio::test]
|
||||
async fn framing_errors_terminate_and_preserve_input_error_causes(#[case] kind: io::ErrorKind) {
|
||||
let mut frames = Box::pin(SseFramer.frame(futures_util::stream::iter([
|
||||
Err(io::Error::new(kind, "reset")),
|
||||
Ok(&b"data: later\n\n"[..]),
|
||||
])));
|
||||
let error = frames.next().await.unwrap().unwrap_err();
|
||||
assert!(matches!(
|
||||
error,
|
||||
Error::Sse(sse_stream::Error::Body(ref cause))
|
||||
if cause.downcast_ref::<io::Error>().unwrap().kind() == kind
|
||||
async fn a_body_error_keeps_earlier_events_and_its_cause_then_terminates(
|
||||
#[case] kind: io::ErrorKind,
|
||||
) {
|
||||
let mut events = Box::pin(frames(
|
||||
stream::iter([
|
||||
Ok(&b"data: first\n\ndata: partial"[..]),
|
||||
Err(io::Error::new(kind, "reset")),
|
||||
Ok(&b"\n\n"[..]),
|
||||
]),
|
||||
SseCodec::default(),
|
||||
));
|
||||
assert!(frames.next().await.is_none());
|
||||
assert!(frames.next().await.is_none());
|
||||
|
||||
assert_eq!(events.next().await.unwrap().unwrap(), event(None, "first"));
|
||||
let Some(Err(SseError::Body(body))) = events.next().await else {
|
||||
panic!("the body error surfaces");
|
||||
};
|
||||
assert_eq!(body_cause::<io::Error>(&body).unwrap().kind(), kind);
|
||||
assert!(events.next().await.is_none());
|
||||
assert!(events.next().await.is_none());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,15 +1,57 @@
|
|||
use aws_smithy_eventstream::frame::write_message_to;
|
||||
use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
|
||||
use bytes::Bytes;
|
||||
#![allow(dead_code)]
|
||||
|
||||
pub fn encode(payload: &'static [u8]) -> Vec<u8> {
|
||||
let message = Message::new(Bytes::from_static(payload))
|
||||
.add_header(Header::new(
|
||||
":event-type",
|
||||
HeaderValue::String("payload".into()),
|
||||
))
|
||||
.add_header(Header::new("sequence", HeaderValue::Int32(7)));
|
||||
let mut bytes = Vec::new();
|
||||
write_message_to(&message, &mut bytes).unwrap();
|
||||
bytes
|
||||
use std::{error::Error, io};
|
||||
|
||||
use bytes::{Bytes, BytesMut};
|
||||
use futures_util::{Stream, stream};
|
||||
use tokio_util::codec::Encoder;
|
||||
|
||||
pub fn encode_all<C, I>(mut codec: C, items: impl IntoIterator<Item = I>) -> Vec<u8>
|
||||
where
|
||||
C: Encoder<I>,
|
||||
C::Error: std::fmt::Debug,
|
||||
{
|
||||
let mut wire = BytesMut::new();
|
||||
for item in items {
|
||||
codec.encode(item, &mut wire).unwrap();
|
||||
}
|
||||
wire.to_vec()
|
||||
}
|
||||
|
||||
pub fn cut_at(bytes: &[u8], offsets: impl IntoIterator<Item = usize>) -> Vec<Bytes> {
|
||||
let mut sorted: Vec<usize> = offsets
|
||||
.into_iter()
|
||||
.filter(|offset| *offset <= bytes.len())
|
||||
.collect();
|
||||
sorted.sort_unstable();
|
||||
sorted.dedup();
|
||||
let bounds = std::iter::once(0)
|
||||
.chain(sorted)
|
||||
.chain(std::iter::once(bytes.len()))
|
||||
.collect::<Vec<_>>();
|
||||
bounds
|
||||
.windows(2)
|
||||
.map(|pair| Bytes::copy_from_slice(&bytes[pair[0]..pair[1]]))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn every(bytes: &[u8], size: usize) -> Vec<Bytes> {
|
||||
bytes
|
||||
.chunks(size.max(1))
|
||||
.map(Bytes::copy_from_slice)
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub fn input(pieces: Vec<Bytes>) -> impl Stream<Item = Result<Bytes, io::Error>> + Send {
|
||||
stream::iter(pieces.into_iter().map(Ok))
|
||||
}
|
||||
|
||||
pub fn body_cause<T: Error + 'static>(body: &io::Error) -> Option<&T> {
|
||||
body.get_ref()?.downcast_ref::<T>()
|
||||
}
|
||||
|
||||
pub fn runtime() -> tokio::runtime::Runtime {
|
||||
tokio::runtime::Builder::new_current_thread()
|
||||
.build()
|
||||
.unwrap()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,9 +2,9 @@ use base64::Engine;
|
|||
use bytes::Buf;
|
||||
use futures_util::{Stream, StreamExt};
|
||||
use litellm_framing::{
|
||||
Framer,
|
||||
aws_event_stream::{AwsEventStreamFrame, AwsEventStreamFramer},
|
||||
sse::{SseFrame, SseFramer},
|
||||
aws_event_stream::{AwsEventStreamCodec, Message},
|
||||
frames,
|
||||
sse::{SseCodec, SseEvent},
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
|
@ -13,8 +13,6 @@ use serde_json::{Map, Value};
|
|||
pub enum Error {
|
||||
#[error("stream framing failed: {0}")]
|
||||
StreamFraming(String),
|
||||
#[error("Anthropic SSE frame has no data")]
|
||||
MissingStreamData,
|
||||
#[error("Anthropic stream event is invalid: {0}")]
|
||||
InvalidStreamEvent(String),
|
||||
#[error("Bedrock event payload is invalid: {0}")]
|
||||
|
|
@ -165,15 +163,14 @@ struct BedrockChunkPayload {
|
|||
bytes: String,
|
||||
}
|
||||
|
||||
pub fn decode_anthropic_sse_frame(frame: SseFrame) -> Result<AnthropicMessagesStreamEvent, Error> {
|
||||
let data = frame.data.ok_or(Error::MissingStreamData)?;
|
||||
serde_json::from_str(&data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
|
||||
pub fn decode_anthropic_sse_frame(event: SseEvent) -> Result<AnthropicMessagesStreamEvent, Error> {
|
||||
serde_json::from_str(&event.data).map_err(|error| Error::InvalidStreamEvent(error.to_string()))
|
||||
}
|
||||
|
||||
pub fn decode_bedrock_anthropic_frame(
|
||||
frame: AwsEventStreamFrame,
|
||||
message: Message,
|
||||
) -> Result<AnthropicMessagesStreamEvent, Error> {
|
||||
let payload: BedrockChunkPayload = serde_json::from_slice(&frame.payload)
|
||||
let payload: BedrockChunkPayload = serde_json::from_slice(message.payload())
|
||||
.map_err(|error| Error::InvalidBedrockPayload(error.to_string()))?;
|
||||
let event = base64::engine::general_purpose::STANDARD
|
||||
.decode(payload.bytes)
|
||||
|
|
@ -189,9 +186,8 @@ where
|
|||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
SseFramer.frame(input).map(|frame| {
|
||||
let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?;
|
||||
decode_anthropic_sse_frame(frame)
|
||||
frames(input, SseCodec::default()).map(|event| {
|
||||
decode_anthropic_sse_frame(event.map_err(|error| Error::StreamFraming(error.to_string()))?)
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -203,9 +199,10 @@ where
|
|||
B: Buf + Send,
|
||||
E: std::error::Error + Send + Sync + 'static,
|
||||
{
|
||||
AwsEventStreamFramer.frame(input).map(|frame| {
|
||||
let frame = frame.map_err(|error| Error::StreamFraming(error.to_string()))?;
|
||||
decode_bedrock_anthropic_frame(frame)
|
||||
frames(input, AwsEventStreamCodec).map(|message| {
|
||||
decode_bedrock_anthropic_frame(
|
||||
message.map_err(|error| Error::StreamFraming(error.to_string()))?,
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -247,12 +244,10 @@ mod tests {
|
|||
|
||||
#[test]
|
||||
fn decodes_citations_delta_events() {
|
||||
let event = decode_anthropic_sse_frame(SseFrame {
|
||||
let event = decode_anthropic_sse_frame(SseEvent {
|
||||
event: Some("content_block_delta".into()),
|
||||
data: Some(
|
||||
r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#
|
||||
.into(),
|
||||
),
|
||||
data: r#"{"type":"content_block_delta","index":0,"delta":{"type":"citations_delta","citation":{"type":"char_location"}}}"#
|
||||
.into(),
|
||||
id: None,
|
||||
retry: None,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -118,7 +118,7 @@ pub(super) struct ParityCase {
|
|||
pub(super) fn parity_cases() -> Vec<ParityCase> {
|
||||
serde_json::from_str(include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json"
|
||||
"/../../../tests/unit/secret_managers/hashicorp_vault_parity.json"
|
||||
)))
|
||||
.unwrap()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -32,7 +32,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
|
||||
`_SecretManagerRuntime` is a private implementation detail, not a replacement SDK class. Its async methods return Futures; public `async def` methods retain lazy coroutine creation and `asyncio.create_task` support. Passing the same names and arguments is insufficient to claim parity until the remaining return-value, error, cache and configuration differences above are closed
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py)
|
||||
## [tests/unit/secret_managers/test_aws_secret_manager_replication.py](../../../tests/unit/secret_managers/test_aws_secret_manager_replication.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -48,7 +48,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_replicate_secret_http_error_raises` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
| `test_replicate_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py)
|
||||
## [tests/unit/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/unit/secret_managers/test_aws_secret_manager_rotation.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -59,7 +59,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_write_secret_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
| `test_write_secret_to_live_existing_name_still_fails_without_overwriting` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py)
|
||||
## [tests/unit/secret_managers/test_aws_secret_manager_v2.py](../../../tests/unit/secret_managers/test_aws_secret_manager_v2.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -70,14 +70,14 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_prepare_request_explicit_bedrock_runtime_endpoint_param_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) |
|
||||
| `test_prepare_request_env_bedrock_runtime_endpoint_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_base_secret_manager.py](../../../tests/test_litellm/secret_managers/test_base_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_base_secret_manager.py](../../../tests/unit/secret_managers/test_base_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
| `test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks` | [names_reject_path_traversal_and_control_characters](../secrets-types/tests/rotation.rs) |
|
||||
| `test_raise_if_unsafe_secret_name_allows_legitimate_aliases` | [names_allow_safe_values](../secrets-types/tests/rotation.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_custom_secret_manager.py](../../../tests/test_litellm/secret_managers/test_custom_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_custom_secret_manager.py](../../../tests/unit/secret_managers/test_custom_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -89,7 +89,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_custom_secret_manager_integration_with_litellm` | [manager_strings_are_coerced_like_literal_eval](../secrets/tests/resolution.rs) |
|
||||
| `test_minimal_custom_secret_manager` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_cyberark_secret_manager.py](../../../tests/test_litellm/secret_managers/test_cyberark_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_cyberark_secret_manager.py](../../../tests/unit/secret_managers/test_cyberark_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -97,7 +97,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_async_write_matches_parity_fixture` | [writes_match_python_parity_fixture](../secrets-cyberark/tests/secret_manager/writes.rs) |
|
||||
| `test_missing_credentials_raise_value_error` | [new_validates_credentials_before_license_and_configuration](../secrets-cyberark/tests/secret_manager/configuration.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py)
|
||||
## [tests/unit/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/unit/secret_managers/test_get_azure_ad_token_provider.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -115,7 +115,7 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate |
|
||||
| `test_get_azure_ad_token_provider_defaults_to_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py](../../../tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py)
|
||||
## [tests/unit/secret_managers/test_hashicorp_secret_manager.py](../../../tests/unit/secret_managers/test_hashicorp_secret_manager.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
@ -130,13 +130,13 @@ The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalo
|
|||
| `test_tls_login_uses_login_namespace` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) |
|
||||
| `test_configuration_matches_native_parity_fixture` | [configuration_matches_python_parity_fixture](../secrets-hashicorp/tests/secret_manager/configuration.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_secret_manager_handler.py](../../../tests/test_litellm/secret_managers/test_secret_manager_handler.py)
|
||||
## [tests/unit/secret_managers/test_secret_manager_handler.py](../../../tests/unit/secret_managers/test_secret_manager_handler.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
| `test_azure_key_vault_matches_rust_parity_fixture` | [parity_fixture_matches_python_backend_contract](../secrets-azure/tests/key_vault.rs) |
|
||||
|
||||
## [tests/test_litellm/secret_managers/test_secret_managers_main.py](../../../tests/test_litellm/secret_managers/test_secret_managers_main.py)
|
||||
## [tests/unit/secret_managers/test_secret_managers_main.py](../../../tests/unit/secret_managers/test_secret_managers_main.py)
|
||||
|
||||
| Python test | Rust coverage or boundary |
|
||||
| --- | --- |
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.types.integrations.datadog import DatadogInitParams
|
|||
from litellm.types.integrations.newrelic import NewRelicInitParams
|
||||
from litellm.litellm_core_utils.core_helpers import drop_params_env_flag
|
||||
from litellm.types.integrations.pointfive import PointFiveInitParams
|
||||
from litellm.types.integrations.zerobus import ZerobusInitParams
|
||||
from litellm._logging import (
|
||||
set_verbose,
|
||||
_turn_on_debug,
|
||||
|
|
@ -157,6 +158,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"deepeval",
|
||||
"s3_v2",
|
||||
"pointfive",
|
||||
"zerobus",
|
||||
"aws_sqs",
|
||||
"vector_store_pre_call_hook",
|
||||
"dotprompt",
|
||||
|
|
@ -442,6 +444,7 @@ datadog_llm_observability_params: Optional[Union[DatadogLLMObsInitParams, Dict]]
|
|||
datadog_params: Optional[Union[DatadogInitParams, Dict]] = None
|
||||
newrelic_params: Optional[Union[NewRelicInitParams, Dict]] = None
|
||||
pointfive_params: Optional[Union[PointFiveInitParams, Mapping[str, object]]] = None
|
||||
zerobus_params: Optional[Union[ZerobusInitParams, Mapping[str, object]]] = None
|
||||
aws_sqs_callback_params: Optional[Dict] = None
|
||||
generic_logger_headers: Optional[Dict] = None
|
||||
default_key_generate_params: Optional[Dict] = None
|
||||
|
|
|
|||
|
|
@ -2,7 +2,9 @@
|
|||
|
||||
from .exception_mapping_utils import (
|
||||
ANTHROPIC_ERROR_TYPE_MAP,
|
||||
AnthropicErrorSseFrame,
|
||||
AnthropicExceptionMapping,
|
||||
anthropic_error_sse_frame,
|
||||
)
|
||||
from .exceptions import (
|
||||
AnthropicErrorDetail,
|
||||
|
|
@ -14,6 +16,8 @@ __all__ = [
|
|||
"ANTHROPIC_ERROR_TYPE_MAP",
|
||||
"AnthropicErrorDetail",
|
||||
"AnthropicErrorResponse",
|
||||
"AnthropicErrorSseFrame",
|
||||
"AnthropicErrorType",
|
||||
"AnthropicExceptionMapping",
|
||||
"anthropic_error_sse_frame",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -4,11 +4,12 @@ Utilities for mapping exceptions to Anthropic error format.
|
|||
Similar to litellm/litellm_core_utils/exception_mapping_utils.py but for Anthropic response format.
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
|
||||
from .exceptions import AnthropicErrorResponse, AnthropicErrorType
|
||||
from .exceptions import AnthropicErrorDetail, AnthropicErrorResponse, AnthropicErrorType
|
||||
|
||||
# HTTP status code -> Anthropic error type
|
||||
# Source: https://docs.anthropic.com/en/api/errors
|
||||
|
|
@ -166,3 +167,36 @@ class AnthropicExceptionMapping:
|
|||
message=message,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
|
||||
class AnthropicErrorSseFrame(str):
|
||||
"""One `event: error` frame, for a stream that fails once the response headers are out.
|
||||
|
||||
Anthropic clients pick stream events by the `event:` name, so a frame carrying only a `data:`
|
||||
line is skipped and the failure never reaches the caller. The frame remembers the status and
|
||||
body it was built from, so a stream that fails before its first byte can still answer as a
|
||||
JSON error with that exact status instead of a 200 that only says `api_error`
|
||||
"""
|
||||
|
||||
status_code: int
|
||||
error_response: AnthropicErrorResponse
|
||||
|
||||
def __new__(cls, status_code: int, error_response: AnthropicErrorResponse) -> "AnthropicErrorSseFrame":
|
||||
frame: Final = super().__new__(cls, f"event: error\ndata: {json.dumps(error_response)}\n\n")
|
||||
frame.status_code = status_code
|
||||
frame.error_response = error_response
|
||||
return frame
|
||||
|
||||
def json_body(self, call_id: str | None) -> AnthropicErrorResponse:
|
||||
if call_id is None:
|
||||
return self.error_response
|
||||
detail: Final[AnthropicErrorDetail] = {**self.error_response["error"], "litellm_call_id": call_id}
|
||||
body: Final[AnthropicErrorResponse] = {**self.error_response, "error": detail}
|
||||
return body
|
||||
|
||||
|
||||
def anthropic_error_sse_frame(status_code: int, raw_message: str) -> AnthropicErrorSseFrame:
|
||||
return AnthropicErrorSseFrame(
|
||||
status_code,
|
||||
AnthropicExceptionMapping.transform_to_anthropic_error(status_code=status_code, raw_message=raw_message),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -706,6 +706,10 @@ def _get_batch_job_usage_from_response_body(
|
|||
if ResponseAPILoggingUtils._is_response_api_usage(_usage_dict):
|
||||
return ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(_usage_dict)
|
||||
usage: Final[Usage] = Usage(**_usage_dict)
|
||||
if custom_llm_provider == "xai":
|
||||
from litellm.llms.xai.chat.transformation import XAIChatConfig
|
||||
|
||||
XAIChatConfig.fold_reasoning_tokens_into_completion(usage)
|
||||
return usage
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -31,6 +31,7 @@ from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
|||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.openai.openai import OpenAIBatchesAPI
|
||||
from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction
|
||||
from litellm.llms.xai.batches.handler import XAIBatchesHandler
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import (
|
||||
CancelBatchRequest,
|
||||
|
|
@ -59,6 +60,7 @@ openai_batches_instance: Final = OpenAIBatchesAPI()
|
|||
azure_batches_instance: Final = AzureBatchesAPI()
|
||||
vertex_ai_batches_instance: Final = VertexAIBatchPrediction(gcs_bucket_name="")
|
||||
anthropic_batches_instance: Final = AnthropicBatchesHandler()
|
||||
xai_batches_instance: Final = XAIBatchesHandler()
|
||||
base_llm_http_handler = BaseLLMHTTPHandler()
|
||||
#################################################
|
||||
|
||||
|
|
@ -105,10 +107,22 @@ def _resolve_timeout(
|
|||
@client
|
||||
async def acreate_batch(
|
||||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"],
|
||||
endpoint: Literal[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/completions",
|
||||
"/v1/responses",
|
||||
"/v1/ocr",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/videos/generations",
|
||||
"/v1/videos",
|
||||
"/v1/videos/edits",
|
||||
"/v1/videos/extensions",
|
||||
],
|
||||
input_file_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -157,10 +171,22 @@ async def acreate_batch(
|
|||
@client
|
||||
def create_batch(
|
||||
completion_window: Literal["24h"],
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"],
|
||||
endpoint: Literal[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/completions",
|
||||
"/v1/responses",
|
||||
"/v1/ocr",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/videos/generations",
|
||||
"/v1/videos",
|
||||
"/v1/videos/edits",
|
||||
"/v1/videos/extensions",
|
||||
],
|
||||
input_file_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -243,6 +269,14 @@ def create_batch(
|
|||
model=model,
|
||||
)
|
||||
return response
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.create_batch(
|
||||
_is_async=_is_async,
|
||||
create_batch_data=_create_batch_request,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
|
|
@ -345,7 +379,7 @@ def create_batch(
|
|||
async def aretrieve_batch(
|
||||
batch_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -393,10 +427,18 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
_retrieve_batch_request: RetrieveBatchRequest,
|
||||
_is_async: bool,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai"
|
||||
] = "openai",
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
):
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.retrieve_batch(
|
||||
_is_async=_is_async,
|
||||
batch_id=batch_id,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
|
|
@ -518,7 +560,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
|
|||
def retrieve_batch(
|
||||
batch_id: str,
|
||||
custom_llm_provider: Literal[
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral"
|
||||
"openai", "azure", "vertex_ai", "bedrock", "hosted_vllm", "litellm_proxy", "anthropic", "mistral", "xai"
|
||||
] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
|
|
@ -741,6 +783,15 @@ def list_batches(
|
|||
timeout = 600.0
|
||||
|
||||
_is_async: Final = kwargs.pop("alist_batches", False) is True
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.list_batches(
|
||||
_is_async=_is_async,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
after=after,
|
||||
limit=limit,
|
||||
)
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
|
||||
api_base = (
|
||||
|
|
@ -837,7 +888,7 @@ def list_batches(
|
|||
async def acancel_batch(
|
||||
batch_id: str,
|
||||
model: str | None = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
extra_body: dict[str, str] | None = None,
|
||||
|
|
@ -883,7 +934,7 @@ async def acancel_batch(
|
|||
def cancel_batch(
|
||||
batch_id: str,
|
||||
model: str | None = None,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy"] | str = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "litellm_proxy", "xai"] | str = "openai",
|
||||
metadata: dict[str, str] | None = None,
|
||||
extra_headers: dict[str, str] | None = None,
|
||||
extra_body: dict[str, str] | None = None,
|
||||
|
|
@ -933,6 +984,14 @@ def cancel_batch(
|
|||
)
|
||||
|
||||
_is_async: Final = kwargs.pop("acancel_batch", False) is True
|
||||
if custom_llm_provider == LlmProviders.XAI.value:
|
||||
return xai_batches_instance.cancel_batch(
|
||||
_is_async=_is_async,
|
||||
batch_id=batch_id,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
api_base: str | None = None
|
||||
if custom_llm_provider in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS:
|
||||
api_base = (
|
||||
|
|
|
|||
|
|
@ -1952,6 +1952,15 @@ SENTRY_DENYLIST: Final = [
|
|||
"auth_token",
|
||||
"jwt_token",
|
||||
"private_key",
|
||||
"authorization",
|
||||
"api-key",
|
||||
"x-api-key",
|
||||
"x-goog-api-key",
|
||||
"ocp-apim-subscription-key",
|
||||
"x-litellm-api-key",
|
||||
"x-mcp-auth",
|
||||
"cookie",
|
||||
"set-cookie",
|
||||
"SLACK_WEBHOOK_URL",
|
||||
"ALERTING_WEBHOOK_URL",
|
||||
"webhook_url",
|
||||
|
|
@ -1974,6 +1983,12 @@ SENTRY_DENYLIST: Final = [
|
|||
]
|
||||
SENTRY_PII_DENYLIST: Final = [
|
||||
"user_id",
|
||||
"user_email",
|
||||
"end_user_id",
|
||||
"user_api_key_hash",
|
||||
"user_api_key_user_id",
|
||||
"user_api_key_user_email",
|
||||
"user_api_key_end_user_id",
|
||||
"email",
|
||||
"phone",
|
||||
"address",
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import os
|
|||
from collections.abc import AsyncIterator, Awaitable, Callable, Generator, Sequence
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from functools import partial
|
||||
from importlib.metadata import version
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeAlias, TypeVar, cast
|
||||
|
||||
|
|
@ -34,8 +35,16 @@ _TransportContext: TypeAlias = AbstractAsyncContextManager[_TransportStreams]
|
|||
from mcp.types import (
|
||||
METHOD_NOT_FOUND,
|
||||
REQUEST_TIMEOUT,
|
||||
ClientCapabilities,
|
||||
ElicitationCapability,
|
||||
FormElicitationCapability,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
Implementation,
|
||||
InitializedNotification,
|
||||
InitializeRequest,
|
||||
InitializeRequestParams,
|
||||
InitializeResult,
|
||||
InputRequiredResult,
|
||||
ListPromptsResult,
|
||||
ListResourcesResult,
|
||||
|
|
@ -44,12 +53,14 @@ from mcp.types import (
|
|||
PaginatedResult,
|
||||
Prompt,
|
||||
ResourceTemplate,
|
||||
SamplingCapability,
|
||||
ServerNotification,
|
||||
UrlElicitationCapability,
|
||||
)
|
||||
from mcp.types import CallToolRequestParams as MCPCallToolRequestParams
|
||||
from mcp.types import CallToolResult as MCPCallToolResult
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl
|
||||
from pydantic import AnyUrl, TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -64,11 +75,13 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_er
|
|||
from litellm.proxy._experimental.mcp_server.result_conversion import error_text_result
|
||||
from litellm.types.llms.custom_http import VerifyTypes
|
||||
from litellm.types.mcp import (
|
||||
MCP_LEGACY_VERSIONS,
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
MCPStdioConfig,
|
||||
MCPTransport,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
credential_redirect_hook,
|
||||
has_header,
|
||||
without_header,
|
||||
|
|
@ -386,7 +399,9 @@ class MCPClient:
|
|||
sampling_callback: Callable | None = None,
|
||||
elicitation_callback: Callable | None = None,
|
||||
logging_callback: Callable | None = None,
|
||||
protocol_version: MCPUpstreamProtocol = "auto",
|
||||
):
|
||||
self.protocol_version: MCPUpstreamProtocol = TypeAdapter(MCPUpstreamProtocol).validate_python(protocol_version)
|
||||
self.server_url: str = server_url
|
||||
self.transport_type: MCPTransport = transport_type
|
||||
self.auth_type: MCPAuthType = auth_type
|
||||
|
|
@ -525,6 +540,35 @@ class MCPClient:
|
|||
|
||||
return safe_env
|
||||
|
||||
async def _initialize_session(self, session: ClientSession) -> InitializeResult:
|
||||
if self.protocol_version == "auto":
|
||||
automatic: Final = await session.initialize()
|
||||
if automatic.protocol_version not in MCP_LEGACY_VERSIONS:
|
||||
raise MCPError(code=-32022, message="Upstream selected an unsupported MCP protocol version")
|
||||
return automatic
|
||||
result: Final = await session.send_request(
|
||||
InitializeRequest(
|
||||
params=InitializeRequestParams(
|
||||
protocol_version=self.protocol_version,
|
||||
client_info=Implementation(name="litellm", version=version("litellm")),
|
||||
capabilities=ClientCapabilities(
|
||||
sampling=SamplingCapability() if self._sampling_callback is not None else None,
|
||||
elicitation=ElicitationCapability(
|
||||
form=FormElicitationCapability(), url=UrlElicitationCapability()
|
||||
)
|
||||
if self._elicitation_callback is not None
|
||||
else None,
|
||||
),
|
||||
)
|
||||
),
|
||||
InitializeResult,
|
||||
)
|
||||
if result.protocol_version != self.protocol_version:
|
||||
raise MCPError(code=-32022, message="Upstream did not accept the configured MCP protocol version")
|
||||
session.adopt(result)
|
||||
await session.send_notification(InitializedNotification())
|
||||
return result
|
||||
|
||||
async def _execute_session_operation(
|
||||
self,
|
||||
transport_ctx: _TransportContext,
|
||||
|
|
@ -579,7 +623,7 @@ class MCPClient:
|
|||
)
|
||||
session: Final = await session_ctx.__aenter__()
|
||||
try:
|
||||
init_result: Final = await session.initialize()
|
||||
init_result: Final = await self._initialize_session(session)
|
||||
instructions: Final = getattr(init_result, "instructions", None)
|
||||
self._last_initialize_instructions = (
|
||||
instructions.strip() or None if isinstance(instructions, str) else None
|
||||
|
|
|
|||
|
|
@ -28,12 +28,15 @@ FileCreateProvider = Literal[
|
|||
"manus",
|
||||
"anthropic",
|
||||
"mistral",
|
||||
"xai",
|
||||
]
|
||||
FileRetrieveProvider = Literal[
|
||||
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral"
|
||||
"openai", "azure", "gemini", "vertex_ai", "hosted_vllm", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
|
||||
]
|
||||
FileDeleteProvider = Literal["openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral"]
|
||||
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral"]
|
||||
FileDeleteProvider = Literal[
|
||||
"openai", "azure", "gemini", "bedrock", "litellm_proxy", "manus", "anthropic", "mistral", "xai"
|
||||
]
|
||||
FileListProvider = Literal["openai", "azure", "litellm_proxy", "manus", "anthropic", "mistral", "xai"]
|
||||
import litellm
|
||||
from litellm import get_secret_str
|
||||
from litellm.files.streaming import FileContentStreamingResponse
|
||||
|
|
@ -49,6 +52,8 @@ from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
|||
from litellm.llms.openai.common_utils import get_openai_credentials
|
||||
from litellm.llms.openai.openai import FileDeleted, FileObject, OpenAIFilesAPI
|
||||
from litellm.llms.vertex_ai.files.handler import VertexAIFilesHandler
|
||||
from litellm.llms.xai.batches.handler import XAIBatchesHandler
|
||||
from litellm.llms.xai.batches.transformation import is_xai_batch_results_id
|
||||
from litellm.types.llms.openai import (
|
||||
CreateFileRequest,
|
||||
FileContentRequest,
|
||||
|
|
@ -103,6 +108,7 @@ openai_files_instance: Final = OpenAIFilesAPI()
|
|||
azure_files_instance: Final = AzureOpenAIFilesAPI()
|
||||
vertex_ai_files_instance: Final = VertexAIFilesHandler()
|
||||
bedrock_files_instance: Final = BedrockFilesHandler()
|
||||
xai_batch_results_instance: Final = XAIBatchesHandler()
|
||||
#################################################
|
||||
|
||||
|
||||
|
|
@ -920,6 +926,15 @@ def file_content(
|
|||
client=client,
|
||||
)
|
||||
|
||||
if custom_llm_provider == LlmProviders.XAI.value and is_xai_batch_results_id(file_id):
|
||||
return xai_batch_results_instance.batch_results_content(
|
||||
_is_async=_is_async,
|
||||
batch_id=file_id,
|
||||
api_base=optional_params.api_base,
|
||||
api_key=optional_params.api_key,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
# Check if provider has a custom files config (e.g., Anthropic, Manus)
|
||||
provider_config: Final = ProviderConfigManager.get_provider_files_config(
|
||||
model="",
|
||||
|
|
|
|||
|
|
@ -406,6 +406,45 @@
|
|||
},
|
||||
"description": "PointFive Logging Integration"
|
||||
},
|
||||
{
|
||||
"id": "zerobus",
|
||||
"displayName": "Databricks Zerobus",
|
||||
"logo": "databricks.svg",
|
||||
"supports_key_team_logging": false,
|
||||
"dynamic_params": {
|
||||
"ZEROBUS_WORKSPACE_URL": {
|
||||
"type": "text",
|
||||
"ui_name": "Workspace URL",
|
||||
"description": "Databricks workspace URL, e.g. https://dbc-a1b2c3d4-e5f6.cloud.databricks.com",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_SERVER_ENDPOINT": {
|
||||
"type": "text",
|
||||
"ui_name": "Zerobus Endpoint",
|
||||
"description": "Zerobus ingest endpoint, e.g. https://<workspace-id>.zerobus.<region>.cloud.databricks.com",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_CLIENT_ID": {
|
||||
"type": "text",
|
||||
"ui_name": "Service Principal Client ID",
|
||||
"description": "OAuth client id of a service principal with USE CATALOG, USE SCHEMA, SELECT and MODIFY on the table",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_CLIENT_SECRET": {
|
||||
"type": "password",
|
||||
"ui_name": "Service Principal Client Secret",
|
||||
"description": "OAuth client secret of the service principal",
|
||||
"required": true
|
||||
},
|
||||
"ZEROBUS_TABLE_NAME": {
|
||||
"type": "text",
|
||||
"ui_name": "Table",
|
||||
"description": "Fully qualified Unity Catalog table, catalog.schema.table, created with the LiteLLM trace schema",
|
||||
"required": true
|
||||
}
|
||||
},
|
||||
"description": "Databricks Zerobus Ingest Logging Integration"
|
||||
},
|
||||
{
|
||||
"id": "s3",
|
||||
"displayName": "S3",
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ litellm/integrations/levo/
|
|||
|
||||
## Testing
|
||||
|
||||
See the test files in `tests/test_litellm/integrations/levo/`:
|
||||
See the test files in `tests/unit/integrations/levo/`:
|
||||
- `test_levo.py`: Unit tests for configuration
|
||||
- `test_levo_integration.py`: Integration tests for callback registration
|
||||
|
||||
|
|
|
|||
5
litellm/integrations/zerobus/__init__.py
Normal file
5
litellm/integrations/zerobus/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""Databricks Zerobus logging integration for LiteLLM."""
|
||||
|
||||
from litellm.integrations.zerobus.logger import ZerobusLogger
|
||||
|
||||
__all__ = ("ZerobusLogger",)
|
||||
161
litellm/integrations/zerobus/client.py
Normal file
161
litellm/integrations/zerobus/client.py
Normal file
|
|
@ -0,0 +1,161 @@
|
|||
"""
|
||||
Writes rows to a Unity Catalog table through the Zerobus Ingest REST API.
|
||||
|
||||
Zerobus only accepts a Databricks OAuth token minted for its own resource and scoped to
|
||||
the target table's privileges, so the client mints that token itself with the service
|
||||
principal's client credentials and reuses it until shortly before it expires.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.types.integrations.zerobus import (
|
||||
RETRYABLE_INGEST_STATUS_CODES,
|
||||
TOKEN_REFRESH_LEEWAY_SECONDS,
|
||||
ZerobusAccessToken,
|
||||
ZerobusConnection,
|
||||
ZerobusIngestFailure,
|
||||
)
|
||||
|
||||
TOKEN_PATH: Final = "/oidc/v1/token"
|
||||
OAUTH_SCOPE: Final = "all-apis"
|
||||
|
||||
|
||||
class _TokenResponse(BaseModel):
|
||||
access_token: str
|
||||
expires_in: float = 3600
|
||||
|
||||
|
||||
class ZerobusIngestError(Exception):
|
||||
"""A batch could not be written and the failure is worth retrying."""
|
||||
|
||||
|
||||
def zerobus_resource(workspace_id: str) -> str:
|
||||
return f"api://databricks/workspaces/{workspace_id}/zerobusDirectWriteApi"
|
||||
|
||||
|
||||
def authorization_details(table_name: str) -> str:
|
||||
"""The Unity Catalog privileges Zerobus requires the token to carry, as the token endpoint expects them."""
|
||||
catalog, schema, _table = table_name.split(".", 2)
|
||||
return json.dumps(
|
||||
(
|
||||
{
|
||||
"type": "unity_catalog_privileges",
|
||||
"privileges": ("USE CATALOG",),
|
||||
"object_type": "CATALOG",
|
||||
"object_full_path": catalog,
|
||||
},
|
||||
{
|
||||
"type": "unity_catalog_privileges",
|
||||
"privileges": ("USE SCHEMA",),
|
||||
"object_type": "SCHEMA",
|
||||
"object_full_path": f"{catalog}.{schema}",
|
||||
},
|
||||
{
|
||||
"type": "unity_catalog_privileges",
|
||||
"privileges": ("SELECT", "MODIFY"),
|
||||
"object_type": "TABLE",
|
||||
"object_full_path": table_name,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def insert_url(connection: ZerobusConnection) -> str:
|
||||
return f"{connection.server_endpoint.rstrip('/')}/zerobus/v1/tables/{connection.table_name}/insert"
|
||||
|
||||
|
||||
def token_url(connection: ZerobusConnection) -> str:
|
||||
return f"{connection.workspace_url.rstrip('/')}{TOKEN_PATH}"
|
||||
|
||||
|
||||
def _basic_auth(client_id: str, client_secret: str) -> str:
|
||||
return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
|
||||
|
||||
|
||||
def _status_failure(what: str, error: httpx.HTTPStatusError) -> ZerobusIngestFailure:
|
||||
status: Final = error.response.status_code
|
||||
return ZerobusIngestFailure(
|
||||
detail=f"{what} returned {status}: {error.response.text}"[:500],
|
||||
retryable=status in RETRYABLE_INGEST_STATUS_CODES,
|
||||
)
|
||||
|
||||
|
||||
class ZerobusIngestClient:
|
||||
def __init__(
|
||||
self,
|
||||
connection: ZerobusConnection,
|
||||
http_client: AsyncHTTPHandler,
|
||||
clock: Callable[[], float] = time.time,
|
||||
) -> None:
|
||||
self.connection: Final = connection
|
||||
self.http_client: Final = http_client
|
||||
self.clock: Final = clock
|
||||
self._token: ZerobusAccessToken | None = None
|
||||
self._token_lock: Final = asyncio.Lock()
|
||||
|
||||
async def insert(self, rows: Sequence[Mapping[str, object]]) -> ZerobusIngestFailure | None:
|
||||
"""Write ``rows`` as one request. ``None`` means Zerobus accepted every row."""
|
||||
token: Final = await self.access_token()
|
||||
if isinstance(token, ZerobusIngestFailure):
|
||||
return token
|
||||
try:
|
||||
await self.http_client.post(
|
||||
insert_url(self.connection),
|
||||
content=json.dumps([dict(row) for row in rows]).encode(),
|
||||
headers={"Content-Type": "application/json", "Authorization": f"Bearer {token.value}"},
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
if error.response.status_code == 401:
|
||||
self._token = None
|
||||
return ZerobusIngestFailure(detail="insert returned 401, token discarded", retryable=True)
|
||||
return _status_failure("insert", error)
|
||||
except (httpx.HTTPError, litellm.Timeout) as error:
|
||||
return ZerobusIngestFailure(detail=f"insert failed: {error}", retryable=True)
|
||||
return None
|
||||
|
||||
async def access_token(self) -> ZerobusAccessToken | ZerobusIngestFailure:
|
||||
"""The cached token while it has more than the leeway left, otherwise a fresh one."""
|
||||
async with self._token_lock:
|
||||
cached: Final = self._token
|
||||
if cached is not None and cached.expires_at - self.clock() > TOKEN_REFRESH_LEEWAY_SECONDS:
|
||||
return cached
|
||||
minted: Final = await self._mint_token()
|
||||
if isinstance(minted, ZerobusAccessToken):
|
||||
self._token = minted
|
||||
return minted
|
||||
|
||||
async def _mint_token(self) -> ZerobusAccessToken | ZerobusIngestFailure:
|
||||
connection: Final = self.connection
|
||||
try:
|
||||
response: Final = await self.http_client.post(
|
||||
token_url(connection),
|
||||
data={
|
||||
"grant_type": "client_credentials",
|
||||
"scope": OAUTH_SCOPE,
|
||||
"resource": zerobus_resource(connection.workspace_id),
|
||||
"authorization_details": authorization_details(connection.table_name),
|
||||
},
|
||||
headers={
|
||||
"Content-Type": "application/x-www-form-urlencoded",
|
||||
"Authorization": _basic_auth(connection.client_id, connection.client_secret),
|
||||
},
|
||||
)
|
||||
except httpx.HTTPStatusError as error:
|
||||
return _status_failure("token request", error)
|
||||
except (httpx.HTTPError, litellm.Timeout) as error:
|
||||
return ZerobusIngestFailure(detail=f"token request failed: {error}", retryable=True)
|
||||
try:
|
||||
parsed: Final = _TokenResponse.model_validate_json(response.text)
|
||||
except ValidationError as error:
|
||||
return ZerobusIngestFailure(detail=f"token response was not understood: {error}", retryable=False)
|
||||
return ZerobusAccessToken(value=parsed.access_token, expires_at=self.clock() + parsed.expires_in)
|
||||
230
litellm/integrations/zerobus/logger.py
Normal file
230
litellm/integrations/zerobus/logger.py
Normal file
|
|
@ -0,0 +1,230 @@
|
|||
"""Databricks Zerobus logging integration."""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.custom_batch_logger import CustomBatchLogger
|
||||
from litellm.integrations.zerobus.client import ZerobusIngestClient, ZerobusIngestError
|
||||
from litellm.integrations.zerobus.row import trace_row
|
||||
from litellm.litellm_core_utils.redact_messages import (
|
||||
redacted_standard_logging_payload,
|
||||
should_redact_message_logging,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client, httpxSpecialProvider
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.integrations.zerobus import ZerobusConnection, ZerobusInitParams
|
||||
|
||||
_ENV_REFERENCE_PREFIX: Final = "os.environ/"
|
||||
|
||||
|
||||
def _resolved_secret(value: str | None) -> str | None:
|
||||
"""Resolve a config value that may name a secret; an unset ``os.environ/NAME`` stays unresolved."""
|
||||
if value is None:
|
||||
return None
|
||||
resolved: Final = get_secret_str(value)
|
||||
if resolved:
|
||||
return resolved
|
||||
return None if value.startswith(_ENV_REFERENCE_PREFIX) else value
|
||||
|
||||
|
||||
def _configured_params() -> ZerobusInitParams:
|
||||
configured: Final = litellm.zerobus_params
|
||||
if isinstance(configured, ZerobusInitParams):
|
||||
return configured
|
||||
if isinstance(configured, Mapping):
|
||||
return ZerobusInitParams.model_validate(configured)
|
||||
return ZerobusInitParams()
|
||||
|
||||
|
||||
def _setting(configured: str | None, env_var: str) -> str:
|
||||
"""Prefer the configured value, falling back to the environment the proxy UI writes."""
|
||||
value: Final = _resolved_secret(configured) or get_secret_str(env_var)
|
||||
if not value:
|
||||
raise ValueError(
|
||||
f"zerobus logging requires {env_var}. Set it in the environment, or "
|
||||
f"litellm_settings.zerobus_params.{env_var.removeprefix('ZEROBUS_').lower()} in config.yaml"
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def _workspace_id(server_endpoint: str) -> str:
|
||||
"""The Zerobus endpoint is ``https://<workspace_id>.zerobus.<region>.<cloud>``, so the id is its first label."""
|
||||
host: Final = urlsplit(server_endpoint).hostname or ""
|
||||
workspace_id: Final = host.split(".", 1)[0]
|
||||
if not workspace_id.isdigit():
|
||||
raise ValueError(
|
||||
f"ZEROBUS_SERVER_ENDPOINT {server_endpoint!r} does not look like "
|
||||
"https://<workspace_id>.zerobus.<region>.cloud.databricks.com"
|
||||
)
|
||||
return workspace_id
|
||||
|
||||
|
||||
def _table_name(configured: str | None) -> str:
|
||||
table_name: Final = _setting(configured, "ZEROBUS_TABLE_NAME")
|
||||
if table_name.count(".") != 2:
|
||||
raise ValueError(f"ZEROBUS_TABLE_NAME {table_name!r} must be fully qualified as catalog.schema.table")
|
||||
return table_name
|
||||
|
||||
|
||||
def connection_for(params: ZerobusInitParams) -> ZerobusConnection:
|
||||
"""The connection configured right now, so a UI edit takes effect without a restart."""
|
||||
server_endpoint: Final = _setting(params.server_endpoint, "ZEROBUS_SERVER_ENDPOINT")
|
||||
return ZerobusConnection(
|
||||
workspace_url=_setting(params.workspace_url, "ZEROBUS_WORKSPACE_URL"),
|
||||
workspace_id=_workspace_id(server_endpoint),
|
||||
server_endpoint=server_endpoint,
|
||||
client_id=_setting(params.client_id, "ZEROBUS_CLIENT_ID"),
|
||||
client_secret=_setting(params.client_secret, "ZEROBUS_CLIENT_SECRET"),
|
||||
table_name=_table_name(params.table_name),
|
||||
)
|
||||
|
||||
|
||||
class ZerobusLogger(CustomBatchLogger):
|
||||
preserve_events_added_during_flush = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
params: ZerobusInitParams | None = None,
|
||||
client: ZerobusIngestClient | None = None,
|
||||
start_periodic_flush: bool = True,
|
||||
) -> None:
|
||||
resolved: Final = params if params is not None else _configured_params()
|
||||
self.params: Final = resolved
|
||||
self.given_client: Final = client
|
||||
self._cached_client: ZerobusIngestClient | None = None
|
||||
if client is None:
|
||||
connection_for(resolved)
|
||||
super().__init__(
|
||||
flush_lock=asyncio.Lock(),
|
||||
batch_size=resolved.batch_size,
|
||||
flush_interval=resolved.flush_interval,
|
||||
turn_off_message_logging=bool(resolved.turn_off_message_logging),
|
||||
)
|
||||
self._flushing: bool = False
|
||||
self._batch_flush_task: asyncio.Task[None] | None = None
|
||||
self._periodic_flush_task: asyncio.Task[None] | None = (
|
||||
self._start_periodic_flush_task() if start_periodic_flush else None
|
||||
)
|
||||
|
||||
@property
|
||||
def client(self) -> ZerobusIngestClient:
|
||||
"""A client for the current connection, kept while the connection is unchanged so its token is reused."""
|
||||
if self.given_client is not None:
|
||||
return self.given_client
|
||||
connection: Final = connection_for(self.params)
|
||||
cached: Final = self._cached_client
|
||||
if cached is not None and cached.connection == connection:
|
||||
return cached
|
||||
fresh: Final = ZerobusIngestClient(
|
||||
connection=connection,
|
||||
http_client=get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback),
|
||||
)
|
||||
self._cached_client = fresh
|
||||
return fresh
|
||||
|
||||
def _start_periodic_flush_task(self) -> asyncio.Task[None] | None:
|
||||
try:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return None
|
||||
return loop.create_task(self.periodic_flush())
|
||||
|
||||
def _start_batch_flush_task(self) -> None:
|
||||
if self._batch_flush_task is not None and not self._batch_flush_task.done():
|
||||
return
|
||||
try:
|
||||
loop: Final = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return
|
||||
self._batch_flush_task = loop.create_task(self.flush_queue(skip_if_flushing=True))
|
||||
|
||||
def _flush_task_is_alive(self) -> bool:
|
||||
task: Final = self._periodic_flush_task
|
||||
return task is not None and not task.done() and not task.get_loop().is_closed()
|
||||
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
await self._enqueue(kwargs)
|
||||
|
||||
async def async_log_failure_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
) -> None:
|
||||
await self._enqueue(kwargs)
|
||||
|
||||
async def _enqueue(self, kwargs: Mapping[str, object]) -> None:
|
||||
try:
|
||||
if not self._flush_task_is_alive():
|
||||
self._periodic_flush_task = self._start_periodic_flush_task()
|
||||
|
||||
payload: Final = self._payload_for(kwargs)
|
||||
if payload is None:
|
||||
verbose_logger.debug("zerobus: event carried no standard_logging_object, skipping")
|
||||
return
|
||||
|
||||
if self._flushing and len(self.log_queue) >= self.max_queue_size:
|
||||
verbose_logger.warning("zerobus: queue at %s rows during a flush, dropped a row", self.max_queue_size)
|
||||
return
|
||||
|
||||
self.log_queue.append(trace_row(payload))
|
||||
self._drop_overflow()
|
||||
if len(self.log_queue) >= self.batch_size:
|
||||
self._start_batch_flush_task()
|
||||
except Exception: # noqa: BLE001 # logging must never break the request path
|
||||
verbose_logger.exception("zerobus: failed to queue an event")
|
||||
|
||||
def _payload_for(self, kwargs: Mapping[str, object]) -> Mapping[str, object] | None:
|
||||
"""The payload to buffer, redacted the way the framework redacts the success path."""
|
||||
details: Final = self.redact_standard_logging_payload_from_model_call_details(
|
||||
dict(kwargs) # mutable-ok: both framework helpers take the call details as a dict
|
||||
)
|
||||
payload: Final = details.get("standard_logging_object")
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
if should_redact_message_logging(details):
|
||||
return redacted_standard_logging_payload(payload)
|
||||
return payload
|
||||
|
||||
def _drop_overflow(self) -> None:
|
||||
"""Trim the oldest rows, except mid flush when the in-flight batch is the head of the queue."""
|
||||
if self._flushing:
|
||||
return
|
||||
overflow: Final = len(self.log_queue) - self.max_queue_size
|
||||
if overflow <= 0:
|
||||
return
|
||||
del self.log_queue[:overflow]
|
||||
verbose_logger.warning("zerobus: queue over %s rows, dropped %s oldest", self.max_queue_size, overflow)
|
||||
|
||||
async def flush_queue(self, skip_if_flushing: bool = False) -> None:
|
||||
if skip_if_flushing and self._flushing:
|
||||
return
|
||||
self._flushing = True
|
||||
try:
|
||||
await super().flush_queue()
|
||||
finally:
|
||||
self._flushing = False
|
||||
|
||||
async def async_send_batch(self) -> None:
|
||||
"""A retryable failure propagates so the rows are kept; a permanent one drops them so the queue moves on."""
|
||||
rows: Final = tuple(self.log_queue)
|
||||
if not rows:
|
||||
return
|
||||
failure: Final = await self.client.insert(rows)
|
||||
if failure is None:
|
||||
return
|
||||
if failure.retryable:
|
||||
raise ZerobusIngestError(failure.detail)
|
||||
verbose_logger.error("zerobus: dropping %s rows, %s", len(rows), failure.detail)
|
||||
156
litellm/integrations/zerobus/row.py
Normal file
156
litellm/integrations/zerobus/row.py
Normal file
|
|
@ -0,0 +1,156 @@
|
|||
"""
|
||||
Shape of one Delta table row per LiteLLM request.
|
||||
|
||||
Zerobus validates every record against the target table and rejects unknown columns, so
|
||||
the row is a fixed set of scalar columns for filtering plus JSON-encoded ``VARIANT``
|
||||
columns for anything nested. ``create_table_sql`` renders the matching DDL.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
|
||||
TRACE_TABLE_COLUMNS: Final[Mapping[str, str]] = MappingProxyType(
|
||||
{
|
||||
"id": "STRING",
|
||||
"trace_id": "STRING",
|
||||
"session_id": "STRING",
|
||||
"litellm_call_id": "STRING",
|
||||
"call_type": "STRING",
|
||||
"status": "STRING",
|
||||
"model": "STRING",
|
||||
"model_group": "STRING",
|
||||
"model_id": "STRING",
|
||||
"custom_llm_provider": "STRING",
|
||||
"api_base": "STRING",
|
||||
"stream": "BOOLEAN",
|
||||
"cache_hit": "BOOLEAN",
|
||||
"start_time": "TIMESTAMP",
|
||||
"end_time": "TIMESTAMP",
|
||||
"completion_start_time": "TIMESTAMP",
|
||||
"response_time": "DOUBLE",
|
||||
"prompt_tokens": "LONG",
|
||||
"completion_tokens": "LONG",
|
||||
"total_tokens": "LONG",
|
||||
"response_cost": "DOUBLE",
|
||||
"saved_cache_cost": "DOUBLE",
|
||||
"api_key_hash": "STRING",
|
||||
"api_key_alias": "STRING",
|
||||
"team_id": "STRING",
|
||||
"team_alias": "STRING",
|
||||
"user_id": "STRING",
|
||||
"org_id": "STRING",
|
||||
"end_user": "STRING",
|
||||
"requester_ip_address": "STRING",
|
||||
"user_agent": "STRING",
|
||||
"request_tags": "VARIANT",
|
||||
"messages": "VARIANT",
|
||||
"response": "VARIANT",
|
||||
"error_str": "STRING",
|
||||
"error_information": "VARIANT",
|
||||
"metadata": "VARIANT",
|
||||
"model_parameters": "VARIANT",
|
||||
"hidden_params": "VARIANT",
|
||||
"guardrail_information": "VARIANT",
|
||||
"cost_breakdown": "VARIANT",
|
||||
}
|
||||
)
|
||||
|
||||
_MICROSECONDS: Final = 1_000_000
|
||||
|
||||
|
||||
def create_table_sql(table_name: str) -> str:
|
||||
columns: Final = ",\n".join(f" {name} {delta_type}" for name, delta_type in TRACE_TABLE_COLUMNS.items())
|
||||
return f"CREATE TABLE {table_name} (\n{columns}\n);"
|
||||
|
||||
|
||||
def _text(payload: Mapping[str, object], key: str) -> str | None:
|
||||
value: Final = payload.get(key)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def _flag(payload: Mapping[str, object], key: str) -> bool | None:
|
||||
value: Final = payload.get(key)
|
||||
return value if isinstance(value, bool) else None
|
||||
|
||||
|
||||
def _number(payload: Mapping[str, object], key: str) -> float | None:
|
||||
value: Final = payload.get(key)
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)):
|
||||
return None
|
||||
return float(value)
|
||||
|
||||
|
||||
def _count(payload: Mapping[str, object], key: str) -> int | None:
|
||||
value: Final = _number(payload, key)
|
||||
return None if value is None else int(value)
|
||||
|
||||
|
||||
def _timestamp_micros(payload: Mapping[str, object], key: str) -> int | None:
|
||||
"""Delta ``TIMESTAMP`` over Zerobus is epoch microseconds; LiteLLM keeps epoch seconds."""
|
||||
seconds: Final = _number(payload, key)
|
||||
if seconds is None or seconds <= 0:
|
||||
return None
|
||||
return int(seconds * _MICROSECONDS)
|
||||
|
||||
|
||||
def _json(payload: Mapping[str, object], key: str) -> str | None:
|
||||
value: Final = payload.get(key)
|
||||
return None if value is None else safe_dumps(value)
|
||||
|
||||
|
||||
def _metadata(payload: Mapping[str, object]) -> Mapping[str, object]:
|
||||
value: Final = payload.get("metadata")
|
||||
return value if isinstance(value, Mapping) else MappingProxyType({})
|
||||
|
||||
|
||||
def trace_row(payload: Mapping[str, object]) -> Mapping[str, object]:
|
||||
"""One ``TRACE_TABLE_COLUMNS`` row for a ``StandardLoggingPayload``."""
|
||||
metadata: Final = _metadata(payload)
|
||||
return MappingProxyType(
|
||||
{
|
||||
"id": _text(payload, "id"),
|
||||
"trace_id": _text(payload, "trace_id"),
|
||||
"session_id": _text(payload, "session_id"),
|
||||
"litellm_call_id": _text(payload, "litellm_call_id"),
|
||||
"call_type": _text(payload, "call_type"),
|
||||
"status": _text(payload, "status"),
|
||||
"model": _text(payload, "model"),
|
||||
"model_group": _text(payload, "model_group"),
|
||||
"model_id": _text(payload, "model_id"),
|
||||
"custom_llm_provider": _text(payload, "custom_llm_provider"),
|
||||
"api_base": _text(payload, "api_base"),
|
||||
"stream": _flag(payload, "stream"),
|
||||
"cache_hit": _flag(payload, "cache_hit"),
|
||||
"start_time": _timestamp_micros(payload, "startTime"),
|
||||
"end_time": _timestamp_micros(payload, "endTime"),
|
||||
"completion_start_time": _timestamp_micros(payload, "completionStartTime"),
|
||||
"response_time": _number(payload, "response_time"),
|
||||
"prompt_tokens": _count(payload, "prompt_tokens"),
|
||||
"completion_tokens": _count(payload, "completion_tokens"),
|
||||
"total_tokens": _count(payload, "total_tokens"),
|
||||
"response_cost": _number(payload, "response_cost"),
|
||||
"saved_cache_cost": _number(payload, "saved_cache_cost"),
|
||||
"api_key_hash": _text(metadata, "user_api_key_hash"),
|
||||
"api_key_alias": _text(metadata, "user_api_key_alias"),
|
||||
"team_id": _text(metadata, "user_api_key_team_id"),
|
||||
"team_alias": _text(metadata, "user_api_key_team_alias"),
|
||||
"user_id": _text(metadata, "user_api_key_user_id"),
|
||||
"org_id": _text(metadata, "user_api_key_org_id"),
|
||||
"end_user": _text(payload, "end_user"),
|
||||
"requester_ip_address": _text(payload, "requester_ip_address"),
|
||||
"user_agent": _text(payload, "user_agent"),
|
||||
"request_tags": _json(payload, "request_tags"),
|
||||
"messages": _json(payload, "messages"),
|
||||
"response": _json(payload, "response"),
|
||||
"error_str": _text(payload, "error_str"),
|
||||
"error_information": _json(payload, "error_information"),
|
||||
"metadata": _json(payload, "metadata"),
|
||||
"model_parameters": _json(payload, "model_parameters"),
|
||||
"hidden_params": _json(payload, "hidden_params"),
|
||||
"guardrail_information": _json(payload, "guardrail_information"),
|
||||
"cost_breakdown": _json(payload, "cost_breakdown"),
|
||||
}
|
||||
)
|
||||
|
|
@ -52,6 +52,7 @@ from litellm.integrations.vantage.vantage_logger import VantageLogger
|
|||
from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import (
|
||||
VectorStorePreCallHook,
|
||||
)
|
||||
from litellm.integrations.zerobus import ZerobusLogger
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter import _PROXY_DynamicRateLimitHandler
|
||||
from litellm.proxy.hooks.dynamic_rate_limiter_v3 import _PROXY_DynamicRateLimitHandlerV3
|
||||
|
||||
|
|
@ -97,6 +98,7 @@ class CustomLoggerRegistry:
|
|||
"deepeval": DeepEvalLogger,
|
||||
"s3_v2": S3Logger,
|
||||
"pointfive": PointFiveLogger,
|
||||
"zerobus": ZerobusLogger,
|
||||
"aws_sqs": SQSLogger,
|
||||
"dynamic_rate_limiter": _PROXY_DynamicRateLimitHandler,
|
||||
"dynamic_rate_limiter_v3": _PROXY_DynamicRateLimitHandlerV3,
|
||||
|
|
|
|||
|
|
@ -114,10 +114,9 @@ class HealthCheckHelpers:
|
|||
"""
|
||||
Health check for batch mode.
|
||||
|
||||
Calls list_batches for providers that support it (openai, hosted_vllm, azure,
|
||||
vertex_ai). For all other providers (e.g. bedrock) the batch API surface doesn't
|
||||
include list_batches, so we fall back to acompletion to verify connectivity and
|
||||
credential validity instead.
|
||||
Calls list_batches for providers that support it. For all other providers (e.g. bedrock)
|
||||
the batch API surface doesn't include list_batches, so we fall back to acompletion to
|
||||
verify connectivity and credential validity instead.
|
||||
"""
|
||||
import litellm
|
||||
|
||||
|
|
@ -132,10 +131,9 @@ class HealthCheckHelpers:
|
|||
litellm_params={"api_base": api_base} if api_base else None,
|
||||
)
|
||||
|
||||
if custom_llm_provider in LIST_BATCHES_SUPPORTED_PROVIDERS:
|
||||
return await litellm.alist_batches(**filtered_model_params)
|
||||
else:
|
||||
if custom_llm_provider not in LIST_BATCHES_SUPPORTED_PROVIDERS:
|
||||
return await litellm.acompletion(**model_params)
|
||||
return await litellm.alist_batches(**{**filtered_model_params, "custom_llm_provider": custom_llm_provider})
|
||||
|
||||
@staticmethod
|
||||
async def _image_edit_health_check(edit_request: Callable[[], Awaitable["ImageResponse"]]) -> "ImageResponse":
|
||||
|
|
|
|||
|
|
@ -43,8 +43,6 @@ from litellm.constants import (
|
|||
DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT,
|
||||
EMPTY_MAPPING,
|
||||
PROVIDER_REQUEST_ID_HEADERS,
|
||||
SENTRY_DENYLIST,
|
||||
SENTRY_PII_DENYLIST,
|
||||
)
|
||||
from litellm.cost_calculator import (
|
||||
RealtimeAPITokenUsageProcessor,
|
||||
|
|
@ -213,6 +211,7 @@ from ..integrations.s3 import S3Logger
|
|||
from ..integrations.s3_v2 import S3Logger as S3V2Logger
|
||||
from ..integrations.supabase import Supabase
|
||||
from ..integrations.traceloop import TraceloopLogger
|
||||
from ..integrations.zerobus import ZerobusLogger
|
||||
from .exception_mapping_utils import _get_response_headers
|
||||
from .initialize_dynamic_callback_params import (
|
||||
get_trusted_callback_params,
|
||||
|
|
@ -380,9 +379,12 @@ _DEPLOYMENT_PRICING_KEYS: Final = (
|
|||
"output_cost_per_token",
|
||||
"input_cost_per_token_batches",
|
||||
"output_cost_per_token_batches",
|
||||
"input_cost_per_token_above_200k_tokens_batches",
|
||||
"input_cost_per_token_above_272k_tokens_batches",
|
||||
"output_cost_per_token_above_200k_tokens_batches",
|
||||
"output_cost_per_token_above_272k_tokens_batches",
|
||||
"cache_read_input_token_cost_batches",
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches",
|
||||
"cache_read_input_token_cost_above_272k_tokens_batches",
|
||||
"cache_creation_input_token_cost_batches",
|
||||
"cache_creation_input_token_cost_above_272k_tokens_batches",
|
||||
|
|
@ -4215,6 +4217,9 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
json_mode=False,
|
||||
litellm_params={},
|
||||
)
|
||||
elif result is None:
|
||||
verbose_logger.warning("LiteLLM: the anthropic_messages stream assembled no response, logging an empty one")
|
||||
return litellm.ModelResponse(model=self.model)
|
||||
else:
|
||||
from litellm.types.llms.anthropic import AnthropicResponse
|
||||
|
||||
|
|
@ -4423,21 +4428,10 @@ def set_callbacks(callback_list, function_id=None):
|
|||
print_verbose("Package 'sentry_sdk' is missing. Installing it...")
|
||||
subprocess.check_call([sys.executable, "-m", "pip", "install", "sentry_sdk"])
|
||||
import sentry_sdk
|
||||
from sentry_sdk.scrubber import EventScrubber
|
||||
from litellm.litellm_core_utils.sentry_scrubbing import build_sentry_init_options
|
||||
|
||||
sentry_sdk_instance = sentry_sdk
|
||||
sentry_trace_rate = os.environ.get("SENTRY_API_TRACE_RATE", "1.0")
|
||||
sentry_sample_rate = (
|
||||
os.environ.get("SENTRY_API_SAMPLE_RATE") if "SENTRY_API_SAMPLE_RATE" in os.environ else "1.0"
|
||||
)
|
||||
sentry_sdk_instance.init(
|
||||
dsn=os.environ.get("SENTRY_DSN"),
|
||||
traces_sample_rate=float(sentry_trace_rate),
|
||||
sample_rate=float(sentry_sample_rate if sentry_sample_rate else 1.0),
|
||||
send_default_pii=False, # Prevent sending Personal Identifiable Information
|
||||
event_scrubber=EventScrubber(denylist=SENTRY_DENYLIST, pii_denylist=SENTRY_PII_DENYLIST),
|
||||
environment=os.environ.get("SENTRY_ENVIRONMENT", "production"),
|
||||
)
|
||||
sentry_sdk_instance.init(**build_sentry_init_options(os.environ))
|
||||
capture_exception = sentry_sdk_instance.capture_exception
|
||||
add_breadcrumb = sentry_sdk_instance.add_breadcrumb
|
||||
elif callback == "slack":
|
||||
|
|
@ -4660,6 +4654,14 @@ def _init_custom_logger_compatible_class(
|
|||
_pointfive_logger: Final = PointFiveLogger()
|
||||
_in_memory_loggers.append(_pointfive_logger)
|
||||
return _pointfive_logger
|
||||
elif logging_integration == "zerobus":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ZerobusLogger):
|
||||
return callback
|
||||
|
||||
_zerobus_logger: Final = ZerobusLogger()
|
||||
_in_memory_loggers.append(_zerobus_logger)
|
||||
return _zerobus_logger
|
||||
elif logging_integration == "aws_sqs":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SQSLogger):
|
||||
|
|
@ -5352,6 +5354,10 @@ def get_custom_logger_compatible_class(
|
|||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PointFiveLogger):
|
||||
return callback
|
||||
elif logging_integration == "zerobus":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, ZerobusLogger):
|
||||
return callback
|
||||
elif logging_integration == "aws_sqs":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, SQSLogger):
|
||||
|
|
|
|||
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal file
152
litellm/litellm_core_utils/sentry_scrubbing.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from functools import reduce
|
||||
from typing import TYPE_CHECKING, Final, TypeAlias, cast
|
||||
|
||||
from pydantic import JsonValue
|
||||
from sentry_sdk.scrubber import DEFAULT_DENYLIST, DEFAULT_PII_DENYLIST, EventScrubber
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.constants import (
|
||||
LENGTH_OF_LITELLM_GENERATED_KEY,
|
||||
MINIMUM_CUSTOM_KEY_LENGTH,
|
||||
SENTRY_DENYLIST,
|
||||
SENTRY_PII_DENYLIST,
|
||||
)
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sentry_sdk.types import Event, Hint
|
||||
|
||||
EventScrubFn: TypeAlias = "Callable[[Event, Hint], Event]"
|
||||
JsonPath: TypeAlias = tuple[str, ...]
|
||||
|
||||
FILTERED: Final = "[Filtered]"
|
||||
SEND_DEFAULT_PII_ENV: Final = "SENTRY_SEND_DEFAULT_PII"
|
||||
SECRET_FIELD_NAMES: Final = tuple(DEFAULT_DENYLIST) + tuple(SENTRY_DENYLIST)
|
||||
PII_FIELD_NAMES: Final = tuple(DEFAULT_PII_DENYLIST) + tuple(SENTRY_PII_DENYLIST)
|
||||
|
||||
KEY_PREFIX: Final = "sk-"
|
||||
|
||||
|
||||
def build_key_pattern(custom_key_minimum: int, generated_key_bytes: int) -> re.Pattern[str]:
|
||||
generated_suffix_length: Final = (generated_key_bytes * 4 + 2) // 3
|
||||
floor: Final = min(custom_key_minimum - len(KEY_PREFIX), generated_suffix_length)
|
||||
return re.compile(rf"{KEY_PREFIX}[A-Za-z0-9_-]{{{floor},}}")
|
||||
|
||||
|
||||
LITELLM_KEY_PATTERN: Final = build_key_pattern(MINIMUM_CUSTOM_KEY_LENGTH, LENGTH_OF_LITELLM_GENERATED_KEY)
|
||||
SOURCE_CONTEXT_KEYS: Final = frozenset({"pre_context", "context_line", "post_context"})
|
||||
STACK_FRAME_PATHS: Final = frozenset(
|
||||
{
|
||||
("exception", "values", "*", "stacktrace", "frames", "*"),
|
||||
("threads", "values", "*", "stacktrace", "frames", "*"),
|
||||
("stacktrace", "frames", "*"),
|
||||
}
|
||||
)
|
||||
MAX_SCRUB_DEPTH: Final = 64
|
||||
EMAIL_PATTERN: Final = re.compile(r"[A-Za-z0-9._%+-]+@[A-Za-z0-9-]+(?:\.[A-Za-z0-9-]+)*\.[A-Za-z]{2,}")
|
||||
SHA256_HEX_PATTERN: Final = re.compile(r"(?<![0-9A-Za-z])[0-9a-f]{64}(?![0-9A-Za-z])")
|
||||
QUOTED_VALUE: Final = r"'(?:[^'\\]|\\.)*'|\"(?:[^\"\\]|\\.)*\""
|
||||
BRACKET_ATOM: Final = rf"(?:{QUOTED_VALUE})|[^\[\]{{}}()'\"]"
|
||||
NESTED_BRACKET_LEVELS: Final = 3
|
||||
BRACKETED_VALUE: Final = reduce(
|
||||
lambda inner, _: rf"[\[{{(](?:{BRACKET_ATOM}|{inner})*[\]}})]",
|
||||
range(NESTED_BRACKET_LEVELS),
|
||||
rf"[\[{{(](?:{BRACKET_ATOM})*[\]}})]",
|
||||
)
|
||||
BARE_VALUE: Final = r"(?!None(?![0-9A-Za-z_]))[^,)\]}\s]+"
|
||||
|
||||
|
||||
class SentryInitOptions(TypedDict):
|
||||
dsn: ReadOnly[str | None]
|
||||
traces_sample_rate: ReadOnly[float]
|
||||
sample_rate: ReadOnly[float]
|
||||
send_default_pii: ReadOnly[bool]
|
||||
event_scrubber: ReadOnly[EventScrubber]
|
||||
before_send: ReadOnly[EventScrubFn]
|
||||
before_send_transaction: ReadOnly[EventScrubFn]
|
||||
environment: ReadOnly[str]
|
||||
|
||||
|
||||
def build_repr_field_pattern(field_names: Sequence[str]) -> re.Pattern[str]:
|
||||
names: Final = "|".join(re.escape(name) for name in field_names)
|
||||
return re.compile(
|
||||
rf"(?P<field>(?<![0-9A-Za-z_])(?:{names})=|['\"](?:{names})['\"]:\s*)(?P<value>{QUOTED_VALUE}|{BRACKETED_VALUE}|{BARE_VALUE})",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
|
||||
def build_string_scrubber(send_default_pii: bool) -> Callable[[str], str]:
|
||||
field_names: Final = SECRET_FIELD_NAMES if send_default_pii else SECRET_FIELD_NAMES + PII_FIELD_NAMES
|
||||
field_pattern: Final = build_repr_field_pattern(field_names)
|
||||
value_patterns: Final = (
|
||||
(LITELLM_KEY_PATTERN,) if send_default_pii else (LITELLM_KEY_PATTERN, EMAIL_PATTERN, SHA256_HEX_PATTERN)
|
||||
)
|
||||
|
||||
def scrub(text: str) -> str:
|
||||
fields_scrubbed: Final = field_pattern.sub(_filtered_field, text)
|
||||
return _substitute_all(value_patterns, fields_scrubbed)
|
||||
|
||||
return scrub
|
||||
|
||||
|
||||
def _filtered_field(match: re.Match[str]) -> str:
|
||||
quote: Final = '"' if match.group("value").startswith('"') else "'"
|
||||
return f"{match.group('field')}{quote}{FILTERED}{quote}"
|
||||
|
||||
|
||||
def _substitute_all(patterns: Sequence[re.Pattern[str]], text: str) -> str:
|
||||
return reduce(lambda scrubbed, pattern: pattern.sub(FILTERED, scrubbed), patterns, text)
|
||||
|
||||
|
||||
def scrub_json_strings(value: JsonValue, scrub: Callable[[str], str], path: JsonPath = ()) -> JsonValue:
|
||||
if len(path) > MAX_SCRUB_DEPTH:
|
||||
return FILTERED
|
||||
if isinstance(value, str):
|
||||
return scrub(value)
|
||||
if isinstance(value, dict):
|
||||
unscrubbed_keys: Final = SOURCE_CONTEXT_KEYS if path in STACK_FRAME_PATHS else frozenset[str]()
|
||||
return { # mutable-ok: JSON object
|
||||
key: item if key in unscrubbed_keys else scrub_json_strings(item, scrub, (*path, key))
|
||||
for key, item in value.items()
|
||||
}
|
||||
if isinstance(value, list):
|
||||
return [scrub_json_strings(item, scrub, (*path, "*")) for item in value] # mutable-ok: JSON array
|
||||
return value
|
||||
|
||||
|
||||
def build_event_scrubber(send_default_pii: bool) -> EventScrubFn:
|
||||
scrub: Final = build_string_scrubber(send_default_pii)
|
||||
|
||||
def scrub_event(event: Event, _hint: Hint) -> Event:
|
||||
json_event: Final = cast("JsonValue", event) # cast-ok: [LIT006] the SDK serialized the event to JSON already
|
||||
return cast("Event", scrub_json_strings(json_event, scrub)) # cast-ok: [LIT006] same JSON shape going back
|
||||
|
||||
return scrub_event
|
||||
|
||||
|
||||
def send_default_pii_from_env(env: Mapping[str, str]) -> bool:
|
||||
return str_to_bool(env.get(SEND_DEFAULT_PII_ENV)) is True
|
||||
|
||||
|
||||
def build_sentry_init_options(env: Mapping[str, str]) -> SentryInitOptions:
|
||||
send_default_pii: Final = send_default_pii_from_env(env)
|
||||
scrub_event: Final = build_event_scrubber(send_default_pii)
|
||||
return SentryInitOptions(
|
||||
dsn=env.get("SENTRY_DSN"),
|
||||
traces_sample_rate=float(env.get("SENTRY_API_TRACE_RATE") or "1.0"),
|
||||
sample_rate=float(env.get("SENTRY_API_SAMPLE_RATE") or "1.0"),
|
||||
send_default_pii=send_default_pii,
|
||||
event_scrubber=EventScrubber(
|
||||
denylist=list(SECRET_FIELD_NAMES), # mutable-ok: EventScrubber appends pii_denylist onto denylist in place
|
||||
pii_denylist=list(PII_FIELD_NAMES), # mutable-ok: EventScrubber takes List[str]
|
||||
recursive=True,
|
||||
send_default_pii=send_default_pii,
|
||||
),
|
||||
before_send=scrub_event,
|
||||
before_send_transaction=scrub_event,
|
||||
environment=env.get("SENTRY_ENVIRONMENT", "production"),
|
||||
)
|
||||
|
|
@ -68,15 +68,11 @@ def _error_status_and_message(exc: Exception) -> tuple[int, str]:
|
|||
|
||||
def _mid_stream_error_sse_event(exc: Exception) -> bytes:
|
||||
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
|
||||
AnthropicExceptionMapping,
|
||||
anthropic_error_sse_frame,
|
||||
)
|
||||
|
||||
status_code, message = _error_status_and_message(exc)
|
||||
error_response = AnthropicExceptionMapping.transform_to_anthropic_error(
|
||||
status_code=status_code,
|
||||
raw_message=message,
|
||||
)
|
||||
return f"event: error\ndata: {json.dumps(error_response)}\n\n".encode()
|
||||
return anthropic_error_sse_frame(status_code=status_code, raw_message=message).encode()
|
||||
|
||||
|
||||
def _delta_payload_field(delta_type: StreamingContentBlockDeltaType) -> str:
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ from litellm.litellm_core_utils.core_helpers import process_response_headers
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import INCOMPLETE_STREAM_ERROR_MESSAGE
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
PassThroughEndpointLogging,
|
||||
)
|
||||
|
|
@ -28,11 +29,6 @@ GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
|
|||
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
|
||||
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
|
||||
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
|
||||
"Provider stream ended before emitting a message_stop event; "
|
||||
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
|
||||
)
|
||||
|
||||
|
||||
def _is_message_stop_chunk(chunk: object) -> bool:
|
||||
if isinstance(chunk, dict):
|
||||
|
|
|
|||
|
|
@ -15,6 +15,12 @@ if TYPE_CHECKING:
|
|||
from litellm.exceptions import ContentPolicyViolationError
|
||||
|
||||
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
|
||||
"Provider stream ended before emitting a message_stop event; "
|
||||
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
|
||||
)
|
||||
|
||||
|
||||
def get_safeguard_refusal_stop_details(response: object) -> Mapping[str, Any] | None:
|
||||
"""
|
||||
Return the ``stop_details`` of an Anthropic Messages response refused by a
|
||||
|
|
|
|||
|
|
@ -2,20 +2,25 @@
|
|||
## Translates OpenAI call to Anthropic `/v1/messages` format
|
||||
import asyncio
|
||||
import json
|
||||
import traceback
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, field_validator
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._logging import redact_internal_details_from_client_message
|
||||
from litellm._uuid import uuid
|
||||
from litellm.exceptions import MidStreamFallbackError
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
encrypted_reasoning_signature,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE,
|
||||
refusal_stop_details,
|
||||
responses_output_refusal_text,
|
||||
)
|
||||
from litellm.responses.streaming_iterator import stream_error_status_and_message
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicUsage
|
||||
|
||||
from .transformation import (
|
||||
|
|
@ -27,6 +32,72 @@ if TYPE_CHECKING:
|
|||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject
|
||||
|
||||
|
||||
class _UpstreamFailure(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
status_code: int | None = None
|
||||
message: str | None = None
|
||||
|
||||
@field_validator("status_code", mode="before")
|
||||
@classmethod
|
||||
def http_error_status_or_none(cls, value: object) -> int | None:
|
||||
candidate: Final = (
|
||||
value
|
||||
if isinstance(value, int) and not isinstance(value, bool)
|
||||
else int(value)
|
||||
if isinstance(value, str) and value.isdecimal()
|
||||
else None
|
||||
)
|
||||
return candidate if candidate is not None and 400 <= candidate <= 599 else None
|
||||
|
||||
@field_validator("message", mode="before")
|
||||
@classmethod
|
||||
def str_or_none(cls, value: object) -> str | None:
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
class _FailedResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, from_attributes=True)
|
||||
|
||||
error: object | None = None
|
||||
|
||||
|
||||
class _FailedResponseEvent(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, from_attributes=True)
|
||||
|
||||
response: _FailedResponse | None = None
|
||||
|
||||
|
||||
def _original_failure(exception: Exception) -> Exception:
|
||||
failure = exception # rebind-ok: walks the MidStreamFallbackError chain down to the provider failure
|
||||
while isinstance(failure, MidStreamFallbackError) and failure.original_exception is not None:
|
||||
failure = failure.original_exception
|
||||
return failure
|
||||
|
||||
|
||||
def _failure_status_and_message(exception: Exception) -> tuple[int, str]:
|
||||
original: Final = _original_failure(exception)
|
||||
failure: Final = _UpstreamFailure.model_validate(
|
||||
{"status_code": getattr(original, "status_code", None), "message": getattr(original, "message", None)}
|
||||
)
|
||||
status_code: Final = failure.status_code if failure.status_code is not None else 500
|
||||
message: Final = failure.message or str(original) or INCOMPLETE_STREAM_ERROR_MESSAGE
|
||||
return status_code, message
|
||||
|
||||
|
||||
def _anthropic_error_chunk(status_code: int, message: str) -> dict[str, object]:
|
||||
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
|
||||
AnthropicExceptionMapping,
|
||||
)
|
||||
|
||||
return dict(
|
||||
AnthropicExceptionMapping.transform_to_anthropic_error(
|
||||
status_code=status_code,
|
||||
raw_message=redact_internal_details_from_client_message(message),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class AnthropicResponsesStreamWrapper:
|
||||
"""
|
||||
Wraps a Responses API streaming iterator and re-emits events in Anthropic SSE format.
|
||||
|
|
@ -40,6 +111,7 @@ class AnthropicResponsesStreamWrapper:
|
|||
response.function_call_arguments.delta -> content_block_delta (input_json_delta)
|
||||
response.output_item.done -> content_block_delta (signature_delta) + content_block_stop
|
||||
response.completed -> message_delta + message_stop
|
||||
response.failed -> error (the stream ends without message_stop)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
|
|
@ -60,6 +132,7 @@ class AnthropicResponsesStreamWrapper:
|
|||
self._pending_tool_ids: dict[str, str] = {} # item_id -> call_id / name accumulator
|
||||
self._sent_message_start = False
|
||||
self._sent_message_stop = False
|
||||
self._stream_failed = False
|
||||
self._chunk_queue: deque[dict[str, object]] = deque()
|
||||
self._refusal_text: str = ""
|
||||
self._sync_responses_iterator: Iterator[object] | None = None
|
||||
|
|
@ -293,10 +366,23 @@ class AnthropicResponsesStreamWrapper:
|
|||
)
|
||||
return
|
||||
|
||||
if event_type == "response.failed":
|
||||
failed: Final = _FailedResponseEvent.model_validate(event)
|
||||
status_code, message = stream_error_status_and_message(
|
||||
failed.response.error if failed.response is not None else None
|
||||
)
|
||||
verbose_logger.error(
|
||||
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed (%s): %s",
|
||||
self.model,
|
||||
status_code,
|
||||
message,
|
||||
)
|
||||
self._fail_stream(status_code, message)
|
||||
return
|
||||
|
||||
# ---- response completed -> message_delta + message_stop ----
|
||||
if event_type in (
|
||||
"response.completed",
|
||||
"response.failed",
|
||||
"response.incomplete",
|
||||
):
|
||||
response_obj: Final = getattr(event, "response", None) or (
|
||||
|
|
@ -350,21 +436,24 @@ class AnthropicResponsesStreamWrapper:
|
|||
self._sent_message_stop = True
|
||||
return
|
||||
|
||||
def _fail_stream(self, status_code: int, message: str) -> None:
|
||||
self._stream_failed = True
|
||||
self._chunk_queue.append(_anthropic_error_chunk(status_code, message))
|
||||
|
||||
def __aiter__(self) -> "AnthropicResponsesStreamWrapper":
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> dict[str, object]:
|
||||
# Return any queued chunks first
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
if self._stream_failed:
|
||||
raise StopAsyncIteration
|
||||
|
||||
# Emit message_start if not yet done (fallback if response.created wasn't fired)
|
||||
if not self._sent_message_start:
|
||||
self._sent_message_start = True
|
||||
self._chunk_queue.append(self._make_message_start())
|
||||
return self._chunk_queue.popleft()
|
||||
|
||||
# Consume the upstream stream
|
||||
try:
|
||||
if hasattr(self.responses_stream, "__aiter__"):
|
||||
async for event in self.responses_stream:
|
||||
|
|
@ -382,10 +471,19 @@ class AnthropicResponsesStreamWrapper:
|
|||
return self._chunk_queue.popleft()
|
||||
except StopAsyncIteration:
|
||||
pass
|
||||
except Exception as e:
|
||||
verbose_logger.error("AnthropicResponsesStreamWrapper error: %s\n%s", e, traceback.format_exc())
|
||||
except Exception as e: # noqa: BLE001 # every upstream failure becomes a client error event
|
||||
verbose_logger.exception(
|
||||
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s failed", self.model
|
||||
)
|
||||
self._fail_stream(*_failure_status_and_message(e))
|
||||
|
||||
if not self._chunk_queue and not self._sent_message_stop and not self._stream_failed:
|
||||
verbose_logger.error(
|
||||
"AnthropicResponsesStreamWrapper: upstream Responses stream for %s ended without a terminal event",
|
||||
self.model,
|
||||
)
|
||||
self._fail_stream(500, INCOMPLETE_STREAM_ERROR_MESSAGE)
|
||||
|
||||
# Drain any remaining queued chunks
|
||||
if self._chunk_queue:
|
||||
return self._chunk_queue.popleft()
|
||||
|
||||
|
|
|
|||
|
|
@ -69,7 +69,9 @@ def make_sync_call(
|
|||
completion_stream: Any = MockResponseIterator(model_response=model_response, json_mode=json_mode)
|
||||
else:
|
||||
decoder: Final = AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
completion_stream = decoder.iter_bytes(
|
||||
response.iter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import types
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from typing import Final, cast
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Final, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
|
@ -51,7 +51,11 @@ from ..common_utils import (
|
|||
bedrock_tool_name_mappings: Final[InMemoryCache] = InMemoryCache(max_size_in_memory=50, default_ttl=600)
|
||||
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from botocore.eventstream import EventStreamMessage
|
||||
|
||||
converse_config: Final = AmazonConverseConfig()
|
||||
_STREAM_HEAD_BYTES: Final = 200
|
||||
NOVA_INVOKE_STREAM_EVENT_TYPES: Final = (
|
||||
"messageStart",
|
||||
"contentBlockStart",
|
||||
|
|
@ -162,6 +166,22 @@ class AmazonCohereChatConfig:
|
|||
return optional_params
|
||||
|
||||
|
||||
def _stream_decoder(
|
||||
bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None,
|
||||
*,
|
||||
model: str,
|
||||
json_mode: bool | None,
|
||||
sync_stream: bool,
|
||||
) -> "AWSEventStreamDecoder":
|
||||
if bedrock_invoke_provider == "anthropic":
|
||||
return AmazonAnthropicClaudeStreamDecoder(model=model, sync_stream=sync_stream, json_mode=json_mode)
|
||||
if bedrock_invoke_provider == "deepseek_r1":
|
||||
return AmazonDeepSeekR1StreamDecoder(model=model, sync_stream=sync_stream)
|
||||
if bedrock_invoke_provider == "moonshot":
|
||||
return AmazonOpenAICompatibleStreamDecoder(model=model, sync_stream=sync_stream)
|
||||
return AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
|
||||
|
||||
async def make_call(
|
||||
client: AsyncHTTPHandler | None,
|
||||
api_base: str,
|
||||
|
|
@ -218,28 +238,13 @@ async def make_call(
|
|||
completion_stream: MockResponseIterator | AsyncIterator[GChunk | ModelResponseStream | dict] = (
|
||||
MockResponseIterator(model_response=model_response, json_mode=json_mode)
|
||||
)
|
||||
elif bedrock_invoke_provider == "anthropic":
|
||||
decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "deepseek_r1":
|
||||
decoder = AmazonDeepSeekR1StreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "moonshot":
|
||||
decoder = AmazonOpenAICompatibleStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=False,
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
else:
|
||||
decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
completion_stream = decoder.aiter_bytes(response.aiter_bytes(chunk_size=stream_chunk_size))
|
||||
decoder: Final = _stream_decoder(
|
||||
bedrock_invoke_provider, model=model, json_mode=json_mode, sync_stream=False
|
||||
)
|
||||
completion_stream = decoder.aiter_bytes(
|
||||
response.aiter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -322,28 +327,13 @@ def make_sync_call(
|
|||
completion_stream: MockResponseIterator | Iterator[GChunk | ModelResponseStream | dict] = (
|
||||
MockResponseIterator(model_response=model_response, json_mode=json_mode)
|
||||
)
|
||||
elif bedrock_invoke_provider == "anthropic":
|
||||
decoder: AWSEventStreamDecoder = AmazonAnthropicClaudeStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "deepseek_r1":
|
||||
decoder = AmazonDeepSeekR1StreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
elif bedrock_invoke_provider == "moonshot":
|
||||
decoder = AmazonOpenAICompatibleStreamDecoder(
|
||||
model=model,
|
||||
sync_stream=True,
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
else:
|
||||
decoder = AWSEventStreamDecoder(model=model, json_mode=json_mode)
|
||||
completion_stream = decoder.iter_bytes(response.iter_bytes(chunk_size=stream_chunk_size))
|
||||
decoder: Final = _stream_decoder(
|
||||
bedrock_invoke_provider, model=model, json_mode=json_mode, sync_stream=True
|
||||
)
|
||||
completion_stream = decoder.iter_bytes(
|
||||
response.iter_bytes(chunk_size=stream_chunk_size), response_headers=response.headers
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
|
|
@ -370,6 +360,49 @@ def make_sync_call(
|
|||
raise BedrockError(status_code=500, message=str(e))
|
||||
|
||||
|
||||
def _response_header(response_headers: Mapping[str, str] | None, name: str) -> str | None:
|
||||
return None if response_headers is None else response_headers.get(name)
|
||||
|
||||
|
||||
class _EventStreamTally:
|
||||
def __init__(self) -> None:
|
||||
self.bytes_received = 0
|
||||
self.bytes_decoded = 0
|
||||
self.events = 0
|
||||
self.head = b""
|
||||
|
||||
def add_chunk(self, chunk: bytes) -> None:
|
||||
self.bytes_received += len(chunk)
|
||||
if len(self.head) < _STREAM_HEAD_BYTES:
|
||||
self.head = (self.head + chunk)[:_STREAM_HEAD_BYTES]
|
||||
|
||||
def add_event(self, event: "EventStreamMessage") -> None:
|
||||
self.events += 1
|
||||
self.bytes_decoded += event.prelude.total_length
|
||||
|
||||
def undecoded_stream_error(self, response_headers: Mapping[str, str] | None) -> BedrockError | None:
|
||||
undecoded: Final = self.bytes_received - self.bytes_decoded
|
||||
if self.events and not undecoded:
|
||||
return None
|
||||
detail: Final = (
|
||||
f"content-type={_response_header(response_headers, 'content-type')!r}, "
|
||||
f"x-amzn-requestid={_response_header(response_headers, 'x-amzn-requestid')!r}, "
|
||||
f"{self.bytes_received} bytes received"
|
||||
)
|
||||
if not self.events:
|
||||
return BedrockError(
|
||||
status_code=502,
|
||||
message=(
|
||||
"Bedrock answered the stream with HTTP 200 but its body decoded to no events "
|
||||
f"({detail}, first bytes={self.head!r})"
|
||||
),
|
||||
)
|
||||
return BedrockError(
|
||||
status_code=502,
|
||||
message=f"Bedrock stream ended with {undecoded} undecoded bytes after {self.events} events ({detail})",
|
||||
)
|
||||
|
||||
|
||||
class AWSEventStreamDecoder:
|
||||
def __init__(self, model: str, json_mode: bool | None = False) -> None:
|
||||
from botocore.parsers import EventStreamJSONParser
|
||||
|
|
@ -709,32 +742,48 @@ class AWSEventStreamDecoder:
|
|||
tool_use=None,
|
||||
)
|
||||
|
||||
def iter_bytes(self, iterator: Iterator[bytes]) -> Iterator[GChunk | ModelResponseStream | dict]:
|
||||
def iter_bytes(
|
||||
self, iterator: Iterator[bytes], *, response_headers: Mapping[str, str] | None = None
|
||||
) -> Iterator[GChunk | ModelResponseStream | dict]:
|
||||
"""Given an iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
event_stream_buffer: Final = EventStreamBuffer()
|
||||
tally: Final = _EventStreamTally()
|
||||
for chunk in iterator:
|
||||
event_stream_buffer.add_data(chunk)
|
||||
tally.add_chunk(chunk)
|
||||
for event in event_stream_buffer:
|
||||
tally.add_event(event)
|
||||
message = self._parse_message_from_event(event)
|
||||
if message:
|
||||
# sse_event = ServerSentEvent(data=message, event="completion")
|
||||
_data = json.loads(message)
|
||||
yield self._chunk_parser(chunk_data=_data)
|
||||
undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers)
|
||||
if undecoded_stream_error is not None:
|
||||
raise undecoded_stream_error
|
||||
|
||||
async def aiter_bytes(self, iterator: AsyncIterator[bytes]) -> AsyncIterator[GChunk | ModelResponseStream | dict]:
|
||||
async def aiter_bytes(
|
||||
self, iterator: AsyncIterator[bytes], *, response_headers: Mapping[str, str] | None = None
|
||||
) -> AsyncIterator[GChunk | ModelResponseStream | dict]:
|
||||
"""Given an async iterator that yields lines, iterate over it & yield every event encountered"""
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
|
||||
event_stream_buffer: Final = EventStreamBuffer()
|
||||
tally: Final = _EventStreamTally()
|
||||
async for chunk in iterator:
|
||||
event_stream_buffer.add_data(chunk)
|
||||
tally.add_chunk(chunk)
|
||||
for event in event_stream_buffer:
|
||||
tally.add_event(event)
|
||||
message = self._parse_message_from_event(event)
|
||||
if message:
|
||||
_data = json.loads(message)
|
||||
yield self._chunk_parser(chunk_data=_data)
|
||||
undecoded_stream_error: Final = tally.undecoded_stream_error(response_headers)
|
||||
if undecoded_stream_error is not None:
|
||||
raise undecoded_stream_error
|
||||
|
||||
def _parse_message_from_event(self, event) -> str | None:
|
||||
response_stream_shape: Final = get_bedrock_response_stream_shape()
|
||||
|
|
|
|||
|
|
@ -770,7 +770,9 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
aws_decoder: Final = AmazonAnthropicClaudeMessagesStreamDecoder(
|
||||
model=model,
|
||||
)
|
||||
completion_stream: Final = aws_decoder.aiter_bytes(httpx_response.aiter_bytes())
|
||||
completion_stream: Final = aws_decoder.aiter_bytes(
|
||||
httpx_response.aiter_bytes(), response_headers=httpx_response.headers
|
||||
)
|
||||
# Convert decoded Bedrock events to Server-Sent Events expected by Anthropic clients.
|
||||
return self.bedrock_sse_wrapper(
|
||||
completion_stream=completion_stream,
|
||||
|
|
|
|||
|
|
@ -53,14 +53,18 @@ class VertexAIFilesHandler(GCSBucketBase):
|
|||
|
||||
Sources them from the deployment's ``litellm_params`` (``gcs_bucket_name`` /
|
||||
``bucket_name`` and ``vertex_credentials``), mirroring the write path in
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the global
|
||||
``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT`` env vars. This lets Vertex batch
|
||||
run entirely at the model-group level, so output written to a per-model bucket is
|
||||
readable without setting the global env vars.
|
||||
``VertexAIFilesConfig._get_configured_bucket_name``, and falls back to the
|
||||
``GCS_BATCH_BUCKET_NAME`` then ``GCS_BUCKET_NAME`` / ``GCS_PATH_SERVICE_ACCOUNT``
|
||||
env vars. This lets Vertex batch run entirely at the model-group level, so output
|
||||
written to a per-model bucket is readable without setting the global env vars.
|
||||
"""
|
||||
params: Final[Mapping[str, object]] = litellm_params or {}
|
||||
bucket_candidate: Final = params.get("gcs_bucket_name") or params.get("bucket_name")
|
||||
configured_bucket_name = bucket_candidate if isinstance(bucket_candidate, str) else os.getenv("GCS_BUCKET_NAME")
|
||||
configured_bucket_name = (
|
||||
bucket_candidate
|
||||
if isinstance(bucket_candidate, str)
|
||||
else os.getenv("GCS_BATCH_BUCKET_NAME") or os.getenv("GCS_BUCKET_NAME")
|
||||
)
|
||||
|
||||
credentials: Final = params.get("vertex_credentials") or vertex_credentials
|
||||
if isinstance(credentials, dict):
|
||||
|
|
|
|||
|
|
@ -961,7 +961,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
|
|||
|
||||
def _get_configured_bucket_name(self, litellm_params: dict) -> str:
|
||||
bucket_name: Final = (
|
||||
litellm_params.get("gcs_bucket_name") or litellm_params.get("bucket_name") or os.getenv("GCS_BUCKET_NAME")
|
||||
litellm_params.get("gcs_bucket_name")
|
||||
or litellm_params.get("bucket_name")
|
||||
or os.getenv("GCS_BATCH_BUCKET_NAME")
|
||||
or os.getenv("GCS_BUCKET_NAME")
|
||||
)
|
||||
if not bucket_name:
|
||||
raise ValueError("GCS bucket_name is required")
|
||||
|
|
|
|||
|
|
@ -122,41 +122,26 @@ class VertexAIRAGIngestion(BaseRAGIngestion):
|
|||
"""
|
||||
import litellm
|
||||
|
||||
# Set GCS_BUCKET_NAME env var for litellm.files.create_file
|
||||
# The handler uses this to determine where to upload
|
||||
original_bucket: Final = os.environ.get("GCS_BUCKET_NAME")
|
||||
if self.gcs_bucket:
|
||||
os.environ["GCS_BUCKET_NAME"] = self.gcs_bucket
|
||||
file_tuple: Final = (filename, file_content, content_type)
|
||||
|
||||
try:
|
||||
# Create file tuple for litellm.files.acreate_file
|
||||
file_tuple: Final = (filename, file_content, content_type)
|
||||
verbose_logger.debug(
|
||||
"Uploading file to GCS via litellm.files.acreate_file: %s (bucket: %s)", filename, self.gcs_bucket
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Uploading file to GCS via litellm.files.acreate_file: %s (bucket: %s)", filename, self.gcs_bucket
|
||||
)
|
||||
response: Final = await litellm.acreate_file(
|
||||
file=file_tuple,
|
||||
purpose="assistants",
|
||||
custom_llm_provider="vertex_ai",
|
||||
gcs_bucket_name=self.gcs_bucket,
|
||||
vertex_project=self.vertex_project,
|
||||
vertex_location=self.vertex_location,
|
||||
vertex_credentials=self.vertex_credentials,
|
||||
)
|
||||
|
||||
# Upload to GCS using LiteLLM's file upload
|
||||
response: Final = await litellm.acreate_file(
|
||||
file=file_tuple,
|
||||
purpose="assistants", # Purpose for file storage
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex_project=self.vertex_project,
|
||||
vertex_location=self.vertex_location,
|
||||
vertex_credentials=self.vertex_credentials,
|
||||
)
|
||||
gcs_uri: Final = response.id
|
||||
verbose_logger.info("Uploaded file to GCS: %s", gcs_uri)
|
||||
|
||||
# The response.id should be the GCS URI
|
||||
gcs_uri: Final = response.id
|
||||
verbose_logger.info("Uploaded file to GCS: %s", gcs_uri)
|
||||
|
||||
return gcs_uri
|
||||
finally:
|
||||
# Restore original env var
|
||||
if original_bucket is not None:
|
||||
os.environ["GCS_BUCKET_NAME"] = original_bucket
|
||||
elif "GCS_BUCKET_NAME" in os.environ:
|
||||
del os.environ["GCS_BUCKET_NAME"]
|
||||
return gcs_uri
|
||||
|
||||
async def _import_file_to_corpus_via_sdk(
|
||||
self,
|
||||
|
|
@ -259,6 +244,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: str | None,
|
||||
chunks: list[str],
|
||||
embeddings: list[list[float]] | None,
|
||||
existing_file_id: str | None = None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
"""
|
||||
Store content in Vertex AI RAG corpus.
|
||||
|
|
@ -274,6 +260,7 @@ class VertexAIRAGIngestion(BaseRAGIngestion):
|
|||
content_type: MIME type
|
||||
chunks: Ignored - Vertex AI handles chunking
|
||||
embeddings: Ignored - Vertex AI handles embedding
|
||||
existing_file_id: Existing provider file ID, unsupported for Vertex AI RAG Engine
|
||||
|
||||
Returns:
|
||||
Tuple of (corpus_id, gcs_uri)
|
||||
|
|
|
|||
195
litellm/llms/xai/batches/handler.py
Normal file
195
litellm/llms/xai/batches/handler.py
Normal file
|
|
@ -0,0 +1,195 @@
|
|||
from collections.abc import Coroutine
|
||||
from itertools import chain
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
HTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.types.llms.openai import CreateBatchRequest, HttpxBinaryResponseContent
|
||||
from litellm.types.utils import LiteLLMBatch, LlmProviders
|
||||
|
||||
from .transformation import (
|
||||
XAI_RESULTS_PAGE_SIZE,
|
||||
OpenAIBatchListResponse,
|
||||
XAIBatch,
|
||||
XAIBatchList,
|
||||
XAIBatchResult,
|
||||
XAIBatchResultsPage,
|
||||
get_xai_auth_headers,
|
||||
raise_for_xai_status,
|
||||
results_to_openai_jsonl,
|
||||
to_create_batch_body,
|
||||
to_litellm_batch,
|
||||
to_openai_batch_list,
|
||||
xai_batches_url,
|
||||
)
|
||||
|
||||
_JSONL_CONTENT_TYPE: Final = ("content-type", "application/jsonl")
|
||||
|
||||
|
||||
class _PageParams(TypedDict):
|
||||
limit: ReadOnly[int]
|
||||
pagination_token: NotRequired[ReadOnly[str]]
|
||||
|
||||
|
||||
def _results_params(after: str | None, limit: int | None) -> dict[str, object]: # mutable-ok: httpx params
|
||||
if after is None:
|
||||
return dict(_PageParams(limit=limit or XAI_RESULTS_PAGE_SIZE)) # mutable-ok: httpx params
|
||||
return dict(_PageParams(limit=limit or XAI_RESULTS_PAGE_SIZE, pagination_token=after)) # mutable-ok: httpx params
|
||||
|
||||
|
||||
def _flatten(pages: list[XAIBatchResultsPage]) -> tuple[XAIBatchResult, ...]:
|
||||
return tuple(chain.from_iterable(page.results for page in pages))
|
||||
|
||||
|
||||
def _jsonl_response(url: str, results: tuple[XAIBatchResult, ...]) -> HttpxBinaryResponseContent:
|
||||
return HttpxBinaryResponseContent(
|
||||
response=httpx.Response(
|
||||
status_code=200,
|
||||
content=results_to_openai_jsonl(results),
|
||||
headers=(_JSONL_CONTENT_TYPE,),
|
||||
request=httpx.Request(method="GET", url=url),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class XAIBatchesHandler:
|
||||
def __init__(self, sync_client: HTTPHandler | None = None, async_client: AsyncHTTPHandler | None = None) -> None:
|
||||
self._sync_client = sync_client
|
||||
self._async_client = async_client
|
||||
|
||||
def _sync(self, timeout: float | httpx.Timeout) -> HTTPHandler:
|
||||
return self._sync_client or HTTPHandler(timeout=timeout)
|
||||
|
||||
def _async(self, timeout: float | httpx.Timeout) -> AsyncHTTPHandler:
|
||||
return self._async_client or get_async_httpx_client(
|
||||
llm_provider=LlmProviders.XAI,
|
||||
params={"timeout": timeout}, # mutable-ok: get_async_httpx_client takes a dict
|
||||
)
|
||||
|
||||
def create_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
create_batch_data: CreateBatchRequest,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
url: Final = xai_batches_url(api_base)
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
body: Final = dict(to_create_batch_body(create_batch_data)) # mutable-ok: httpx json body
|
||||
endpoint: Final = create_batch_data.get("endpoint") or "/v1/chat/completions"
|
||||
if _is_async:
|
||||
|
||||
async def _acreate() -> LiteLLMBatch:
|
||||
response: Final = await self._async(timeout).post(url, json=body, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()), endpoint)
|
||||
|
||||
return _acreate()
|
||||
response: Final = self._sync(timeout).post(url, json=body, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()), endpoint)
|
||||
|
||||
def retrieve_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
batch_id: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
url: Final = xai_batches_url(api_base, batch_id)
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
if _is_async:
|
||||
|
||||
async def _aretrieve() -> LiteLLMBatch:
|
||||
response: Final = await self._async(timeout).get(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
return _aretrieve()
|
||||
response: Final = self._sync(timeout).get(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
def cancel_batch(
|
||||
self,
|
||||
_is_async: bool,
|
||||
batch_id: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> LiteLLMBatch | Coroutine[None, None, LiteLLMBatch]:
|
||||
url: Final = xai_batches_url(api_base, batch_id, suffix=":cancel")
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
if _is_async:
|
||||
|
||||
async def _acancel() -> LiteLLMBatch:
|
||||
response: Final = await self._async(timeout).post(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
return _acancel()
|
||||
response: Final = self._sync(timeout).post(url, headers=headers, timeout=timeout)
|
||||
return to_litellm_batch(XAIBatch.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
def list_batches(
|
||||
self,
|
||||
_is_async: bool,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
after: str | None = None,
|
||||
limit: int | None = None,
|
||||
) -> OpenAIBatchListResponse | Coroutine[None, None, OpenAIBatchListResponse]:
|
||||
url: Final = xai_batches_url(api_base)
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
params: Final = _results_params(after, limit)
|
||||
if _is_async:
|
||||
|
||||
async def _alist() -> OpenAIBatchListResponse:
|
||||
response: Final = await self._async(timeout).get(url, params=params, headers=headers, timeout=timeout)
|
||||
return to_openai_batch_list(XAIBatchList.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
return _alist()
|
||||
response: Final = self._sync(timeout).get(url, params=params, headers=headers, timeout=timeout)
|
||||
return to_openai_batch_list(XAIBatchList.model_validate(raise_for_xai_status(response).json()))
|
||||
|
||||
def batch_results_content(
|
||||
self,
|
||||
_is_async: bool,
|
||||
batch_id: str,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
timeout: float | httpx.Timeout,
|
||||
) -> HttpxBinaryResponseContent | Coroutine[None, None, HttpxBinaryResponseContent]:
|
||||
url: Final = xai_batches_url(api_base, batch_id, suffix="/results")
|
||||
headers: Final = get_xai_auth_headers(api_key=api_key)
|
||||
if _is_async:
|
||||
|
||||
async def _aresults() -> HttpxBinaryResponseContent:
|
||||
client: Final = self._async(timeout)
|
||||
|
||||
async def _page(after: str | None) -> XAIBatchResultsPage:
|
||||
response: Final = await client.get(
|
||||
url, params=_results_params(after, None), headers=headers, timeout=timeout
|
||||
)
|
||||
return XAIBatchResultsPage.model_validate(raise_for_xai_status(response).json())
|
||||
|
||||
pages = [await _page(None)] # mutable-ok: page walk terminates on the cursor, not on a fixed count
|
||||
while pages[-1].pagination_token and pages[-1].results:
|
||||
pages.append(await _page(pages[-1].pagination_token))
|
||||
return _jsonl_response(url, _flatten(pages))
|
||||
|
||||
return _aresults()
|
||||
client: Final = self._sync(timeout)
|
||||
|
||||
def _page(after: str | None) -> XAIBatchResultsPage:
|
||||
response: Final = client.get(url, params=_results_params(after, None), headers=headers, timeout=timeout)
|
||||
return XAIBatchResultsPage.model_validate(raise_for_xai_status(response).json())
|
||||
|
||||
pages = [_page(None)] # mutable-ok: page walk terminates on the cursor, not on a fixed count
|
||||
while pages[-1].pagination_token and pages[-1].results:
|
||||
pages.append(_page(pages[-1].pagination_token))
|
||||
return _jsonl_response(url, _flatten(pages))
|
||||
278
litellm/llms/xai/batches/transformation.py
Normal file
278
litellm/llms/xai/batches/transformation.py
Normal file
|
|
@ -0,0 +1,278 @@
|
|||
"""
|
||||
xAI Batch API reference: https://docs.x.ai/developers/advanced-api-usage/batch-api
|
||||
|
||||
xAI batches carry request counters, not a status, and no output file: results are paged from
|
||||
``GET /v1/batches/{id}/results``, so LiteLLM hands back the batch id as ``output_file_id``.
|
||||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime, timezone
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import httpx
|
||||
from openai.types.batch import BatchRequestCounts
|
||||
from openai.types.batch import Errors as BatchErrors
|
||||
from openai.types.batch_error import BatchError
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm.constants import XAI_API_BASE
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.xai.common_utils import XAIModelInfo
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import CreateBatchRequest
|
||||
from litellm.types.utils import LiteLLMBatch
|
||||
|
||||
OpenAIBatchStatus: TypeAlias = Literal[
|
||||
"validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"
|
||||
]
|
||||
|
||||
XAI_BATCH_ID_PREFIX: Final = "batch_"
|
||||
XAI_RESULTS_PAGE_SIZE: Final = 1000
|
||||
DEFAULT_BATCH_NAME: Final = "litellm-batch"
|
||||
DEFAULT_BATCH_ENDPOINT: Final = "/v1/chat/completions"
|
||||
_EMPTY_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
class XAIBatchesError(BaseLLMException):
|
||||
pass
|
||||
|
||||
|
||||
def xai_batches_error(
|
||||
error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
|
||||
) -> XAIBatchesError:
|
||||
return XAIBatchesError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers if isinstance(headers, httpx.Headers) else httpx.Headers(tuple(headers.items())),
|
||||
)
|
||||
|
||||
|
||||
def raise_for_xai_status(response: httpx.Response) -> httpx.Response:
|
||||
if response.status_code >= 400:
|
||||
raise xai_batches_error(response.text, response.status_code, response.headers)
|
||||
return response
|
||||
|
||||
|
||||
def get_xai_api_base(api_base: str | None) -> str:
|
||||
resolved: Final = (api_base or get_secret_str("XAI_API_BASE") or XAI_API_BASE).rstrip("/")
|
||||
return resolved.removesuffix("/v1")
|
||||
|
||||
|
||||
def get_xai_auth_headers(
|
||||
headers: Mapping[str, str] = _EMPTY_HEADERS, api_key: str | None = None
|
||||
) -> dict[str, str]: # mutable-ok: BaseConfig.validate_environment contract returns dict
|
||||
resolved_key: Final = XAIModelInfo.get_api_key(api_key)
|
||||
if resolved_key is None:
|
||||
raise xai_batches_error(
|
||||
"Missing xAI API Key. Pass api_key, set litellm.xai_key or XAI_API_KEY", 401, _EMPTY_HEADERS
|
||||
)
|
||||
return dict(headers, Authorization=f"Bearer {resolved_key}") # mutable-ok: BaseConfig contract returns dict
|
||||
|
||||
|
||||
def xai_batches_url(api_base: str | None, batch_id: str | None = None, suffix: str = "") -> str:
|
||||
base: Final = f"{get_xai_api_base(api_base)}/v1/batches"
|
||||
if batch_id is None:
|
||||
return base
|
||||
return f"{base}/{encode_url_path_segment(batch_id, field_name='batch_id')}{suffix}"
|
||||
|
||||
|
||||
def is_xai_batch_results_id(file_id: str) -> bool:
|
||||
return file_id.startswith(XAI_BATCH_ID_PREFIX)
|
||||
|
||||
|
||||
class XAICreateBatchRequest(TypedDict):
|
||||
name: ReadOnly[str]
|
||||
input_file_id: NotRequired[ReadOnly[str]]
|
||||
|
||||
|
||||
class XAIBatchState(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
num_requests: int = 0
|
||||
num_pending: int = 0
|
||||
num_success: int = 0
|
||||
num_error: int = 0
|
||||
num_cancelled: int = 0
|
||||
|
||||
|
||||
class XAIBatch(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
batch_id: str
|
||||
name: str = ""
|
||||
create_time: str | None = None
|
||||
expire_time: str | None = None
|
||||
cancel_time: str | None = None
|
||||
cancel_by_xai_message: str | None = None
|
||||
state: XAIBatchState = XAIBatchState()
|
||||
input_file_id: str | None = None
|
||||
|
||||
|
||||
class XAIBatchList(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
batches: tuple[XAIBatch, ...] = ()
|
||||
pagination_token: str | None = None
|
||||
|
||||
|
||||
class XAIBatchResultError(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
code: int | str | None = None
|
||||
message: str = ""
|
||||
|
||||
|
||||
class XAIBatchResultData(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
response: Mapping[str, Mapping[str, object]] | None = None
|
||||
error: XAIBatchResultError | None = None
|
||||
|
||||
|
||||
class XAIBatchResult(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
batch_request_id: str
|
||||
batch_result: XAIBatchResultData = XAIBatchResultData()
|
||||
|
||||
|
||||
class XAIBatchResultsPage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
results: tuple[XAIBatchResult, ...] = ()
|
||||
pagination_token: str | None = None
|
||||
|
||||
|
||||
def _to_unix_timestamp(value: str | None) -> int | None:
|
||||
"""xAI returns RFC 3339 timestamps over gRPC but a bare ``YYYY-MM-DD`` over REST."""
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
parsed: Final = datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
except ValueError:
|
||||
return None
|
||||
return int((parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc)).timestamp())
|
||||
|
||||
|
||||
def xai_batch_status(batch: XAIBatch) -> OpenAIBatchStatus:
|
||||
"""xAI exposes counters, not a status. A batch xAI itself cancelled (input validation failed) is a failure,
|
||||
a caller-cancelled batch is cancelled, an empty batch is still validating its input file, and a batch
|
||||
with nothing pending has completed."""
|
||||
if batch.cancel_time is not None:
|
||||
return "failed" if batch.cancel_by_xai_message else "cancelled"
|
||||
if batch.state.num_requests == 0:
|
||||
return "validating"
|
||||
if batch.state.num_pending > 0:
|
||||
return "in_progress"
|
||||
return "completed"
|
||||
|
||||
|
||||
def to_litellm_batch(batch: XAIBatch, endpoint: str = DEFAULT_BATCH_ENDPOINT) -> LiteLLMBatch:
|
||||
status: Final = xai_batch_status(batch)
|
||||
created_at: Final = _to_unix_timestamp(batch.create_time)
|
||||
cancelled_at: Final = _to_unix_timestamp(batch.cancel_time)
|
||||
errors: Final = (
|
||||
BatchErrors(object="list", data=[BatchError(message=batch.cancel_by_xai_message)]) # mutable-ok: openai type
|
||||
if batch.cancel_by_xai_message
|
||||
else None
|
||||
)
|
||||
return LiteLLMBatch(
|
||||
id=batch.batch_id,
|
||||
object="batch",
|
||||
endpoint=endpoint,
|
||||
input_file_id=batch.input_file_id or "",
|
||||
completion_window="24h",
|
||||
status=status,
|
||||
created_at=created_at if created_at is not None else 0,
|
||||
expires_at=_to_unix_timestamp(batch.expire_time),
|
||||
failed_at=cancelled_at if status == "failed" else None,
|
||||
cancelled_at=cancelled_at if status == "cancelled" else None,
|
||||
output_file_id=batch.batch_id if status == "completed" else None,
|
||||
errors=errors,
|
||||
request_counts=BatchRequestCounts(
|
||||
total=batch.state.num_requests,
|
||||
completed=batch.state.num_success,
|
||||
failed=batch.state.num_error + batch.state.num_cancelled,
|
||||
),
|
||||
metadata={"name": batch.name} if batch.name else None, # mutable-ok: LiteLLMBatch.metadata is a dict
|
||||
)
|
||||
|
||||
|
||||
class OpenAIBatchListResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
object: Literal["list"] = "list"
|
||||
data: tuple[LiteLLMBatch, ...]
|
||||
first_id: str | None
|
||||
last_id: str | None
|
||||
has_more: bool
|
||||
next_page_token: str | None = None
|
||||
|
||||
|
||||
def to_openai_batch_list(page: XAIBatchList) -> OpenAIBatchListResponse:
|
||||
data: Final = tuple(to_litellm_batch(b) for b in page.batches)
|
||||
return OpenAIBatchListResponse(
|
||||
data=data,
|
||||
first_id=data[0].id if data else None,
|
||||
last_id=data[-1].id if data else None,
|
||||
has_more=bool(page.pagination_token),
|
||||
next_page_token=page.pagination_token or None,
|
||||
)
|
||||
|
||||
|
||||
def to_create_batch_body(create_batch_data: CreateBatchRequest) -> XAICreateBatchRequest:
|
||||
input_file_id: Final = create_batch_data.get("input_file_id")
|
||||
if not input_file_id:
|
||||
raise xai_batches_error("input_file_id is required to create an xAI batch", 400, _EMPTY_HEADERS)
|
||||
metadata: Final = create_batch_data.get("metadata")
|
||||
name: Final = metadata.get("name") if metadata else None
|
||||
return XAICreateBatchRequest(name=name or DEFAULT_BATCH_NAME, input_file_id=input_file_id)
|
||||
|
||||
|
||||
class OpenAIBatchOutputError(TypedDict):
|
||||
code: ReadOnly[str]
|
||||
message: ReadOnly[str]
|
||||
|
||||
|
||||
class OpenAIBatchOutputResponse(TypedDict):
|
||||
status_code: ReadOnly[int]
|
||||
request_id: ReadOnly[object]
|
||||
body: ReadOnly[Mapping[str, object]]
|
||||
|
||||
|
||||
class OpenAIBatchOutputLine(TypedDict):
|
||||
id: ReadOnly[str]
|
||||
custom_id: ReadOnly[str]
|
||||
response: ReadOnly[OpenAIBatchOutputResponse | None]
|
||||
error: ReadOnly[OpenAIBatchOutputError | None]
|
||||
|
||||
|
||||
def _result_to_openai_line(result: XAIBatchResult) -> OpenAIBatchOutputLine:
|
||||
"""One output JSONL line. xAI wraps the body in a one-key map named after the endpoint
|
||||
(``chat_get_completion``, ``responses``, ``image_generation``, ...); the value is the OpenAI body."""
|
||||
error: Final = result.batch_result.error
|
||||
response: Final = result.batch_result.response
|
||||
body: Final = next(iter(response.values()), None) if response else None
|
||||
if body is None:
|
||||
message: Final = error.message if error is not None else "xAI returned no response for this request"
|
||||
code: Final = str(error.code) if error is not None and error.code is not None else "request_failed"
|
||||
return OpenAIBatchOutputLine(
|
||||
id=f"batch_req_{result.batch_request_id}",
|
||||
custom_id=result.batch_request_id,
|
||||
response=None,
|
||||
error=OpenAIBatchOutputError(code=code, message=message),
|
||||
)
|
||||
return OpenAIBatchOutputLine(
|
||||
id=f"batch_req_{result.batch_request_id}",
|
||||
custom_id=result.batch_request_id,
|
||||
response=OpenAIBatchOutputResponse(status_code=200, request_id=body.get("id"), body=body),
|
||||
error=None,
|
||||
)
|
||||
|
||||
|
||||
def results_to_openai_jsonl(results: Sequence[XAIBatchResult]) -> bytes:
|
||||
return "".join(f"{json.dumps(_result_to_openai_line(r), ensure_ascii=False)}\n" for r in results).encode()
|
||||
|
|
@ -296,7 +296,7 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
except Exception as e:
|
||||
verbose_logger.debug("Error extracting X.AI web search usage: %s", e)
|
||||
|
||||
self._fold_reasoning_tokens_into_completion(response)
|
||||
self.fold_reasoning_tokens_into_completion(response)
|
||||
self._normalize_openai_compatible_usage_totals(getattr(response, "usage", None))
|
||||
restated_usage: Final = _usage_restated_from_xai_ticks(getattr(response, "usage", None))
|
||||
if restated_usage is not None:
|
||||
|
|
@ -304,7 +304,7 @@ class XAIChatConfig(OpenAIGPTConfig):
|
|||
return response
|
||||
|
||||
@staticmethod
|
||||
def _fold_reasoning_tokens_into_completion(
|
||||
def fold_reasoning_tokens_into_completion(
|
||||
target: ModelResponse | Usage | dict[str, Any] | None,
|
||||
) -> None:
|
||||
"""Reconcile xAI Usage to the OpenAI invariant.
|
||||
|
|
@ -426,7 +426,7 @@ class XAIChatCompletionStreamingHandler(OpenAIChatCompletionStreamingHandler):
|
|||
chunk["choices"] = [{"index": 0, "delta": {}, "finish_reason": None}]
|
||||
|
||||
if "usage" in chunk and chunk["usage"] is not None:
|
||||
XAIChatConfig._fold_reasoning_tokens_into_completion(chunk["usage"])
|
||||
XAIChatConfig.fold_reasoning_tokens_into_completion(chunk["usage"])
|
||||
XAIChatConfig._normalize_openai_compatible_usage_totals(chunk["usage"])
|
||||
|
||||
parsed_chunk: Final = super().chunk_parser(chunk)
|
||||
|
|
|
|||
247
litellm/llms/xai/files/transformation.py
Normal file
247
litellm/llms/xai/files/transformation.py
Normal file
|
|
@ -0,0 +1,247 @@
|
|||
"""
|
||||
xAI Files API reference: https://docs.x.ai/developers/rest-api-reference/inference/files
|
||||
|
||||
xAI stores ``purpose`` as an empty string; LiteLLM reports uploads as ``batch``, the only purpose xAI files serve.
|
||||
"""
|
||||
|
||||
import time
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.llms.base_llm.files.transformation import BaseFilesConfig, LiteLLMLoggingObj
|
||||
from litellm.types.llms.openai import (
|
||||
CreateFileRequest,
|
||||
FileContentRequest,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
OpenAIFileObject,
|
||||
OpenAIFilesPurpose,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
from ..batches.transformation import (
|
||||
get_xai_api_base,
|
||||
get_xai_auth_headers,
|
||||
raise_for_xai_status,
|
||||
xai_batches_error,
|
||||
)
|
||||
|
||||
_NO_QUERY_PARAMS: Final[dict[str, str]] = {} # mutable-ok: BaseFilesConfig request transforms return tuple[str, dict]
|
||||
_DEFAULT_PURPOSE: Final[OpenAIFilesPurpose] = "batch"
|
||||
|
||||
|
||||
class XAIMultipartUpload(TypedDict):
|
||||
file: ReadOnly[tuple[str, object, str]]
|
||||
purpose: ReadOnly[tuple[None, str]]
|
||||
|
||||
|
||||
class XAIFile(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
id: str
|
||||
bytes: int = 0
|
||||
created_at: int | None = None
|
||||
filename: str = ""
|
||||
purpose: str = ""
|
||||
expires_at: int | None = None
|
||||
|
||||
|
||||
class XAIFileList(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
data: tuple[XAIFile, ...] = ()
|
||||
pagination_token: str | None = None
|
||||
|
||||
|
||||
class XAIFileDeleted(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
|
||||
id: str
|
||||
deleted: bool = True
|
||||
|
||||
|
||||
def _to_openai_file_object(file: XAIFile) -> OpenAIFileObject:
|
||||
return OpenAIFileObject(
|
||||
id=file.id,
|
||||
bytes=file.bytes,
|
||||
created_at=file.created_at if file.created_at is not None else int(time.time()),
|
||||
filename=file.filename,
|
||||
object="file",
|
||||
purpose=_DEFAULT_PURPOSE,
|
||||
status="uploaded",
|
||||
expires_at=file.expires_at,
|
||||
)
|
||||
|
||||
|
||||
def _api_base_from(litellm_params: Mapping[str, object]) -> str:
|
||||
api_base: Final = litellm_params.get("api_base")
|
||||
return get_xai_api_base(api_base if isinstance(api_base, str) else None)
|
||||
|
||||
|
||||
class XAIFilesConfig(BaseFilesConfig):
|
||||
@property
|
||||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.XAI
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
api_key: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
stream: bool | None = None,
|
||||
) -> str:
|
||||
return f"{get_xai_api_base(api_base)}/v1/files"
|
||||
|
||||
def _file_url(self, file_id: str, litellm_params: Mapping[str, object], suffix: str = "") -> str:
|
||||
encoded_file_id: Final = encode_url_path_segment(file_id, field_name="file_id")
|
||||
return f"{_api_base_from(litellm_params)}/v1/files/{encoded_file_id}{suffix}"
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
|
||||
) -> BaseLLMException:
|
||||
return xai_batches_error(error_message, status_code, headers)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict[str, str]: # mutable-ok: BaseFilesConfig signature
|
||||
return get_xai_auth_headers(headers, api_key)
|
||||
|
||||
def get_supported_openai_params(
|
||||
self, model: str
|
||||
) -> list[OpenAICreateFileRequestOptionalParams]: # mutable-ok: BaseFilesConfig signature
|
||||
return ["purpose"] # mutable-ok: BaseFilesConfig signature
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: dict[str, object], # mutable-ok: BaseConfig signature, returned as-is
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict[str, object]: # mutable-ok: BaseConfig signature
|
||||
return optional_params
|
||||
|
||||
def transform_create_file_request(
|
||||
self,
|
||||
model: str,
|
||||
create_file_data: CreateFileRequest,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> dict[str, object]: # mutable-ok: BaseFilesConfig signature
|
||||
if "file" not in create_file_data:
|
||||
raise ValueError("File data is required")
|
||||
extracted: Final = extract_file_data(create_file_data["file"])
|
||||
filename: Final = extracted["filename"] or f"file_{int(time.time())}.jsonl"
|
||||
content_type: Final = extracted.get("content_type") or "application/octet-stream"
|
||||
upload: Final = XAIMultipartUpload(
|
||||
file=(filename, extracted["content"], content_type),
|
||||
purpose=(None, create_file_data.get("purpose") or _DEFAULT_PURPOSE),
|
||||
)
|
||||
return dict(upload) # mutable-ok: BaseFilesConfig signature
|
||||
|
||||
def transform_create_file_response(
|
||||
self,
|
||||
model: str | None,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> OpenAIFileObject:
|
||||
return _to_openai_file_object(XAIFile.model_validate(raise_for_xai_status(raw_response).json()))
|
||||
|
||||
def transform_retrieve_file_request(
|
||||
self,
|
||||
file_id: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
return self._file_url(file_id, litellm_params), _NO_QUERY_PARAMS
|
||||
|
||||
def transform_retrieve_file_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> OpenAIFileObject:
|
||||
return _to_openai_file_object(XAIFile.model_validate(raise_for_xai_status(raw_response).json()))
|
||||
|
||||
def transform_delete_file_request(
|
||||
self,
|
||||
file_id: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
return self._file_url(file_id, litellm_params), _NO_QUERY_PARAMS
|
||||
|
||||
def transform_delete_file_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> FileDeleted:
|
||||
deleted: Final = XAIFileDeleted.model_validate(raise_for_xai_status(raw_response).json())
|
||||
return FileDeleted(id=deleted.id, deleted=deleted.deleted, object="file")
|
||||
|
||||
def transform_list_files_request(
|
||||
self,
|
||||
purpose: str | None,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
return f"{_api_base_from(litellm_params)}/v1/files", _NO_QUERY_PARAMS
|
||||
|
||||
def transform_list_files_next_request(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]] | None: # mutable-ok: BaseFilesConfig signature
|
||||
page: Final = XAIFileList.model_validate(raw_response.json())
|
||||
if not page.pagination_token or not page.data:
|
||||
return None
|
||||
return f"{_api_base_from(litellm_params)}/v1/files", {"pagination_token": page.pagination_token}
|
||||
|
||||
def transform_list_files_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> list[OpenAIFileObject]: # mutable-ok: BaseFilesConfig signature
|
||||
return [ # mutable-ok: BaseFilesConfig signature
|
||||
_to_openai_file_object(f)
|
||||
for f in XAIFileList.model_validate(raise_for_xai_status(raw_response).json()).data
|
||||
]
|
||||
|
||||
def transform_file_content_request(
|
||||
self,
|
||||
file_content_request: FileContentRequest,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]: # mutable-ok: BaseFilesConfig signature
|
||||
file_id: Final = file_content_request.get("file_id")
|
||||
if file_id is None:
|
||||
raise ValueError("file_id is required to download file content")
|
||||
return self._file_url(file_id, litellm_params, suffix="/content"), _NO_QUERY_PARAMS
|
||||
|
||||
def transform_file_content_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> HttpxBinaryResponseContent:
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
|
@ -51436,13 +51436,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51450,8 +51453,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -51480,9 +51486,13 @@
|
|||
"xai/grok-4.3": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51490,6 +51500,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -51502,9 +51514,13 @@
|
|||
"xai/grok-4.3-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51512,6 +51528,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59483,13 +59501,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59497,20 +59518,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": false,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59519,8 +59546,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -62787,13 +62817,16 @@
|
|||
},
|
||||
"xai/grok-4.20": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62801,21 +62834,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62823,21 +62862,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62845,8 +62890,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -63067,13 +63115,16 @@
|
|||
},
|
||||
"xai/grok-4.20-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63081,20 +63132,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-non-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63102,20 +63159,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63127,20 +63190,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63152,8 +63221,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
|
|
@ -75459,13 +75531,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -75473,8 +75548,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from typing import Final
|
|||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.litellm_core_utils.request_timeout_resolver import get_configured_request_timeout
|
||||
|
||||
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS: Final = 600.0
|
||||
|
||||
_SECONDS: Final = TypeAdapter(float)
|
||||
|
|
@ -48,8 +50,8 @@ def resolve_llm_passthrough_timeout(
|
|||
Anthropic /v1/messages).
|
||||
|
||||
Non-streaming precedence: kwargs timeout/request_timeout -> litellm_params
|
||||
timeout/request_timeout -> router_timeout -> general_settings.pass_through_request_timeout
|
||||
-> 600s default.
|
||||
timeout/request_timeout -> router_timeout -> litellm.request_timeout (litellm_settings.request_timeout,
|
||||
when explicitly set) -> general_settings.pass_through_request_timeout -> 600s default.
|
||||
|
||||
Streaming (``kwargs["stream"]`` truthy) resolves ``stream_timeout`` at every level before
|
||||
any generic timeout, matching ``Router._get_stream_timeout`` on the completion route:
|
||||
|
|
@ -73,6 +75,7 @@ def resolve_llm_passthrough_timeout(
|
|||
deployment.get("timeout"),
|
||||
deployment.get("request_timeout"),
|
||||
router_timeout,
|
||||
get_configured_request_timeout(),
|
||||
)
|
||||
winner: Final = next((val for val in candidates if val is not None), None)
|
||||
return resolve_pass_through_request_timeout() if winner is None else _SECONDS.validate_python(winner)
|
||||
|
|
|
|||
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
145
litellm/proxy/_experimental/mcp_server/capabilities.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
from collections.abc import Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from itertools import product
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from mcp.server.context import CallNext, HandlerResult, ServerRequestContext
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.types import DiscoverResult, InitializeRequestParams, InitializeResult, ServerCapabilities
|
||||
from mcp_types.methods import CLIENT_REQUESTS
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS, LATEST_HANDSHAKE_VERSION
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.types.mcp import MCP_LEGACY_VERSIONS, MCPAdvertisedVersions, MCPLegacyVersion, MCPSpecVersion, MCPTransport
|
||||
|
||||
GATEWAY_OPERATIONS: Final = frozenset(
|
||||
{
|
||||
"tools/list",
|
||||
"tools/call",
|
||||
"prompts/list",
|
||||
"prompts/get",
|
||||
"resources/list",
|
||||
"resources/read",
|
||||
"resources/templates/list",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RevisionSupport:
|
||||
transports: frozenset[MCPTransport]
|
||||
operations: frozenset[str]
|
||||
results: frozenset[Literal["complete", "input_required"]]
|
||||
extensions: frozenset[str]
|
||||
completed: bool
|
||||
|
||||
|
||||
REVISION_SUPPORT: Final[Mapping[str, RevisionSupport]] = MappingProxyType(
|
||||
{
|
||||
version.value: RevisionSupport(
|
||||
transports=frozenset(MCPTransport)
|
||||
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
|
||||
else frozenset({MCPTransport.http, MCPTransport.stdio}),
|
||||
operations=frozenset(method for method in GATEWAY_OPERATIONS if (method, version.value) in CLIENT_REQUESTS),
|
||||
results=frozenset({"complete"})
|
||||
if version.value in HANDSHAKE_PROTOCOL_VERSIONS
|
||||
else frozenset({"complete", "input_required"}),
|
||||
extensions=frozenset(),
|
||||
completed=version.value in HANDSHAKE_PROTOCOL_VERSIONS,
|
||||
)
|
||||
for version in MCPSpecVersion
|
||||
}
|
||||
)
|
||||
_COMPLETED_REVISIONS: Final = tuple(version for version, support in REVISION_SUPPORT.items() if support.completed)
|
||||
TRANSLATION_PAIRS: Final = frozenset(product(_COMPLETED_REVISIONS, repeat=2))
|
||||
_ADVERTISED_VERSIONS: Final[TypeAdapter[tuple[MCPLegacyVersion, ...]]] = TypeAdapter(MCPAdvertisedVersions)
|
||||
|
||||
|
||||
def configured_versions() -> tuple[str, ...]:
|
||||
from litellm.proxy.proxy_server import general_settings_view
|
||||
|
||||
configured: Final = general_settings_view().get("mcp_advertised_versions")
|
||||
return _ADVERTISED_VERSIONS.validate_python(MCP_LEGACY_VERSIONS if configured is None else configured)
|
||||
|
||||
|
||||
def build_discovery(
|
||||
*,
|
||||
configured: tuple[str, ...],
|
||||
revision: str,
|
||||
transport: MCPTransport,
|
||||
authorized_operations: frozenset[str],
|
||||
upstream_versions: frozenset[str],
|
||||
capabilities: ServerCapabilities,
|
||||
client_extensions: frozenset[str] = frozenset(),
|
||||
upstream_extensions: frozenset[str] = frozenset(),
|
||||
instructions: str | None = None,
|
||||
) -> DiscoverResult:
|
||||
supported: Final = tuple(
|
||||
version
|
||||
for version, support in REVISION_SUPPORT.items()
|
||||
if version in configured and support.completed and transport in support.transports
|
||||
)
|
||||
revision_support: Final = REVISION_SUPPORT.get(revision)
|
||||
operations: Final[frozenset[str]] = (
|
||||
authorized_operations & revision_support.operations
|
||||
if revision in supported
|
||||
and revision_support is not None
|
||||
and any((revision, upstream) in TRANSLATION_PAIRS for upstream in upstream_versions)
|
||||
else frozenset()
|
||||
)
|
||||
extensions: Final[frozenset[str]] = (
|
||||
revision_support.extensions & client_extensions & upstream_extensions
|
||||
if operations and revision_support is not None
|
||||
else frozenset()
|
||||
)
|
||||
caller_capabilities: Final = capabilities.model_copy(deep=True)
|
||||
return DiscoverResult(
|
||||
supported_versions=list(supported),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=caller_capabilities.tools if {"tools/list", "tools/call"} <= operations else None,
|
||||
prompts=caller_capabilities.prompts if {"prompts/list", "prompts/get"} <= operations else None,
|
||||
resources=caller_capabilities.resources if {"resources/list", "resources/read"} <= operations else None,
|
||||
extensions={
|
||||
key: value for key, value in (caller_capabilities.extensions or {}).items() if key in extensions
|
||||
}
|
||||
or None,
|
||||
),
|
||||
instructions=instructions,
|
||||
cache_scope="private",
|
||||
ttl_ms=0,
|
||||
)
|
||||
|
||||
|
||||
class GatewayVersionPolicy:
|
||||
def __init__(self, versions: Callable[[], tuple[str, ...]] = configured_versions) -> None:
|
||||
self._versions = versions
|
||||
|
||||
async def __call__(self, ctx: ServerRequestContext[object, object], call_next: CallNext) -> HandlerResult:
|
||||
versions: Final = self._versions()
|
||||
requested: Final = (
|
||||
InitializeRequestParams.model_validate(ctx.params or {}).protocol_version
|
||||
if ctx.method == "initialize"
|
||||
else ctx.protocol_version
|
||||
)
|
||||
negotiated: Final = (
|
||||
(requested if requested in HANDSHAKE_PROTOCOL_VERSIONS else LATEST_HANDSHAKE_VERSION)
|
||||
if ctx.method == "initialize"
|
||||
else requested
|
||||
)
|
||||
if negotiated not in versions:
|
||||
raise MCPError(code=-32022, message="Unsupported MCP protocol version", data={"supported": list(versions)})
|
||||
result: Final = await call_next(ctx)
|
||||
if ctx.method != "initialize":
|
||||
return result
|
||||
initialized: Final = InitializeResult.model_validate(result)
|
||||
discovery: Final = build_discovery(
|
||||
configured=versions,
|
||||
revision=initialized.protocol_version,
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=frozenset(HANDSHAKE_PROTOCOL_VERSIONS),
|
||||
capabilities=initialized.capabilities,
|
||||
instructions=initialized.instructions,
|
||||
)
|
||||
return initialized.model_copy(update={"capabilities": discovery.capabilities})
|
||||
|
|
@ -28,6 +28,7 @@ class OperationContext:
|
|||
client_ip: str | None = None
|
||||
mcp_proxy_mode: bool = False
|
||||
wire_compat: WireCompat = WireCompat.LEGACY
|
||||
protocol_version: str | None = None
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
object.__setattr__(self, "_caller", copy_caller(self._caller))
|
||||
|
|
|
|||
|
|
@ -193,6 +193,7 @@ from litellm.types.mcp import (
|
|||
MCPAuth,
|
||||
MCPStdioConfig,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPUpstreamProtocol,
|
||||
has_header,
|
||||
without_header,
|
||||
)
|
||||
|
|
@ -340,6 +341,7 @@ class MCPServerConfig(TypedDict, total=False):
|
|||
whatever the admin wrote, and each read applies its own default."""
|
||||
|
||||
server_id: ReadOnly[str]
|
||||
protocol_version: ReadOnly[MCPUpstreamProtocol]
|
||||
alias: str
|
||||
description: str
|
||||
mcp_info: MCPInfo
|
||||
|
|
@ -2549,6 +2551,9 @@ class MCPServerManager:
|
|||
new_server = MCPServer(
|
||||
server_id=server_id,
|
||||
name=name_for_prefix,
|
||||
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
server_config.get("protocol_version", mcp_info.get("protocol_version", "auto"))
|
||||
),
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
spec_path=server_config.get("spec_path", None),
|
||||
|
|
@ -3109,6 +3114,9 @@ class MCPServerManager:
|
|||
new_server: Final = MCPServer(
|
||||
server_id=mcp_server.server_id,
|
||||
name=name_for_prefix,
|
||||
protocol_version=TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
_mcp_info.get("protocol_version", "auto")
|
||||
),
|
||||
alias=getattr(mcp_server, "alias", None),
|
||||
server_name=getattr(mcp_server, "server_name", None),
|
||||
url=mcp_server.url,
|
||||
|
|
@ -4145,6 +4153,7 @@ class MCPServerManager:
|
|||
cred_provider: UpstreamCredentialProvider | None = None,
|
||||
raw_headers: Mapping[str, str] | None = None,
|
||||
client_ip: str | None = None,
|
||||
protocol_version_override: MCPUpstreamProtocol | None = None,
|
||||
) -> MCPClient:
|
||||
"""
|
||||
Create an MCPClient instance for the given server.
|
||||
|
|
@ -4168,6 +4177,9 @@ class MCPServerManager:
|
|||
"""
|
||||
record_auth_resolution(server.server_id, AuthResolution.unresolved)
|
||||
resolved_server: Final = await self.ensure_oauth_metadata_discovered(server)
|
||||
protocol_version: Final = (
|
||||
protocol_version_override if protocol_version_override is not None else resolved_server.protocol_version
|
||||
)
|
||||
transport: Final = resolved_server.transport or MCPTransport.sse
|
||||
spec = None if transport == MCPTransport.stdio else _to_server_spec_fail_closed(resolved_server)
|
||||
provider: Final = cred_provider or self._cred_provider
|
||||
|
|
@ -4249,6 +4261,7 @@ class MCPServerManager:
|
|||
return MCPClient(
|
||||
server_url="", # Not used for stdio
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
timeout=(resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT),
|
||||
|
|
@ -4281,6 +4294,7 @@ class MCPServerManager:
|
|||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
timeout=(
|
||||
resolved_server.timeout if resolved_server.timeout is not None else MCP_CLIENT_TIMEOUT
|
||||
|
|
@ -4324,6 +4338,7 @@ class MCPServerManager:
|
|||
MCPClient(
|
||||
server_url=server_url,
|
||||
transport_type=transport,
|
||||
protocol_version=protocol_version,
|
||||
auth_type=resolved_server.auth_type,
|
||||
auth_value=auth_value,
|
||||
auth_header_name=auth_header_name,
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ from mcp.types import (
|
|||
CallToolRequest,
|
||||
CallToolRequestParams,
|
||||
CallToolResult,
|
||||
DiscoverRequest,
|
||||
DiscoverResult,
|
||||
GetPromptRequest,
|
||||
GetPromptRequestParams,
|
||||
GetPromptResult,
|
||||
|
|
@ -28,10 +30,14 @@ from mcp.types import (
|
|||
ListToolsResult,
|
||||
PaginatedRequestParams,
|
||||
Prompt,
|
||||
PromptsCapability,
|
||||
ReadResourceRequest,
|
||||
ReadResourceRequestParams,
|
||||
ResourcesCapability,
|
||||
ResourceTemplate,
|
||||
ServerCapabilities,
|
||||
TextContent,
|
||||
ToolsCapability,
|
||||
)
|
||||
from mcp.types import Tool as MCPTool
|
||||
from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter
|
||||
|
|
@ -51,6 +57,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
|
|||
cache_byok_credential,
|
||||
get_cached_byok_credential,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import (
|
||||
GATEWAY_OPERATIONS,
|
||||
build_discovery,
|
||||
configured_versions,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.contracts import (
|
||||
AuthorizedToolCall,
|
||||
OperationContext,
|
||||
|
|
@ -122,7 +133,9 @@ from litellm.proxy.litellm_pre_call_utils import (
|
|||
)
|
||||
from litellm.types.mcp import (
|
||||
DEFAULT_CREDENTIAL_HEADER,
|
||||
MCP_LEGACY_VERSIONS,
|
||||
MCPAuth,
|
||||
MCPTransport,
|
||||
without_header,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer
|
||||
|
|
@ -2657,7 +2670,11 @@ class _McpDeniedDetail(TypedDict):
|
|||
|
||||
|
||||
async def _execute_handle_list_tools(
|
||||
context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None
|
||||
context: OperationContext,
|
||||
params: PaginatedRequestParams,
|
||||
host_progress_callback: ProgressCallback | None = None,
|
||||
*,
|
||||
log_list_tools_to_spendlogs: bool = True,
|
||||
) -> ListToolsResult:
|
||||
try:
|
||||
(
|
||||
|
|
@ -2700,7 +2717,7 @@ async def _execute_handle_list_tools(
|
|||
mcp_server_auth_headers=mcp_server_auth_headers,
|
||||
oauth2_headers=oauth2_headers,
|
||||
raw_headers=raw_headers,
|
||||
log_list_tools_to_spendlogs=True,
|
||||
log_list_tools_to_spendlogs=log_list_tools_to_spendlogs,
|
||||
list_tools_log_source="mcp_protocol",
|
||||
client_ip=_client_ip,
|
||||
)
|
||||
|
|
@ -3065,6 +3082,7 @@ def prepare_context(
|
|||
client_ip: str | None = None,
|
||||
mcp_proxy_mode: bool = False,
|
||||
wire_compat: WireCompat = WireCompat.LEGACY,
|
||||
protocol_version: str | None = None,
|
||||
) -> OperationContext:
|
||||
return OperationContext(
|
||||
_caller=user_api_key_auth,
|
||||
|
|
@ -3076,11 +3094,13 @@ def prepare_context(
|
|||
client_ip=client_ip,
|
||||
mcp_proxy_mode=mcp_proxy_mode,
|
||||
wire_compat=wire_compat,
|
||||
protocol_version=protocol_version,
|
||||
)
|
||||
|
||||
|
||||
GatewayOperation: TypeAlias = (
|
||||
AuthorizedToolCall
|
||||
| DiscoverRequest
|
||||
| ListToolsRequest
|
||||
| CallToolRequest
|
||||
| ListPromptsRequest
|
||||
|
|
@ -3090,7 +3110,8 @@ GatewayOperation: TypeAlias = (
|
|||
| ReadResourceRequest
|
||||
)
|
||||
GatewayResult: TypeAlias = (
|
||||
ListToolsResult
|
||||
DiscoverResult
|
||||
| ListToolsResult
|
||||
| CallToolResult
|
||||
| InputRequiredResult
|
||||
| ListPromptsResult
|
||||
|
|
@ -3105,6 +3126,9 @@ class GatewayOperations:
|
|||
def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None:
|
||||
self._host_progress_callback = host_progress_callback
|
||||
|
||||
@overload
|
||||
async def execute(self, operation: DiscoverRequest, context: OperationContext) -> DiscoverResult: ...
|
||||
|
||||
@overload
|
||||
async def execute(
|
||||
self, operation: AuthorizedToolCall, context: OperationContext
|
||||
|
|
@ -3137,6 +3161,51 @@ class GatewayOperations:
|
|||
|
||||
async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult:
|
||||
match operation:
|
||||
case DiscoverRequest():
|
||||
listings: Final = (
|
||||
()
|
||||
if context.mcp_proxy_mode
|
||||
else (ListPromptsRequest(), ListResourcesRequest(), ListResourceTemplatesRequest())
|
||||
)
|
||||
tasks: Final = (
|
||||
asyncio.create_task(
|
||||
_execute_handle_list_tools(
|
||||
context,
|
||||
PaginatedRequestParams(),
|
||||
self._host_progress_callback,
|
||||
log_list_tools_to_spendlogs=False,
|
||||
)
|
||||
),
|
||||
*(asyncio.create_task(self.execute(listing, context)) for listing in listings),
|
||||
)
|
||||
try:
|
||||
results: Final = await asyncio.gather(*tasks)
|
||||
finally:
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
return build_discovery(
|
||||
configured=configured_versions(),
|
||||
revision=context.protocol_version or "2025-11-25",
|
||||
transport=MCPTransport.http,
|
||||
authorized_operations=GATEWAY_OPERATIONS,
|
||||
upstream_versions=frozenset(MCP_LEGACY_VERSIONS),
|
||||
capabilities=ServerCapabilities(
|
||||
tools=ToolsCapability()
|
||||
if any(isinstance(result, ListToolsResult) and result.tools for result in results)
|
||||
else None,
|
||||
prompts=PromptsCapability()
|
||||
if any(isinstance(result, ListPromptsResult) and result.prompts for result in results)
|
||||
else None,
|
||||
resources=ResourcesCapability()
|
||||
if any(
|
||||
(isinstance(result, ListResourcesResult) and result.resources)
|
||||
or (isinstance(result, ListResourceTemplatesResult) and result.resource_templates)
|
||||
for result in results
|
||||
)
|
||||
else None,
|
||||
),
|
||||
)
|
||||
case AuthorizedToolCall():
|
||||
auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth()
|
||||
return await _execute_mcp_tool(
|
||||
|
|
|
|||
|
|
@ -1375,7 +1375,16 @@ if MCP_AVAILABLE:
|
|||
and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY)
|
||||
else None
|
||||
)
|
||||
return _StagedServerTest(request=request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers)
|
||||
preview_request: Final = (
|
||||
request.model_copy(
|
||||
update={"mcp_info": {**(request.mcp_info or {}), "protocol_version": saved_server.protocol_version}}
|
||||
)
|
||||
if saved_server is not None and "protocol_version" not in (request.mcp_info or {})
|
||||
else request
|
||||
)
|
||||
return _StagedServerTest(
|
||||
request=preview_request, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers
|
||||
)
|
||||
|
||||
async def _list_tools_within(client: MCPClient, deadline: float) -> list[MCPTool] | None:
|
||||
with anyio.move_on_after(deadline):
|
||||
|
|
@ -1512,6 +1521,7 @@ if MCP_AVAILABLE:
|
|||
extra_headers=merged_headers,
|
||||
stdio_env=stdio_env,
|
||||
cred_provider=preview_cred_provider,
|
||||
protocol_version_override=server_model.protocol_version,
|
||||
)
|
||||
|
||||
return await operation(client)
|
||||
|
|
|
|||
|
|
@ -125,12 +125,14 @@ def unsupported_protocol_version(scope: Scope) -> str | None:
|
|||
``HANDSHAKE_PROTOCOL_VERSIONS`` to the modern single-exchange path, which
|
||||
bypasses litellm's session/auth model, so the ASGI entry rejects it.
|
||||
"""
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import configured_versions
|
||||
|
||||
headers: Final[Iterable[tuple[bytes, bytes]]] = scope.get("headers") or ()
|
||||
values: Final = tuple(
|
||||
raw.decode("latin-1").strip() for key, raw in headers if key.lower() == _MCP_PROTOCOL_VERSION_HEADER
|
||||
)
|
||||
for value in values:
|
||||
if value and value not in HANDSHAKE_PROTOCOL_VERSIONS:
|
||||
if value and value not in configured_versions():
|
||||
return value
|
||||
return None
|
||||
|
||||
|
|
@ -149,7 +151,10 @@ try:
|
|||
from mcp.server.session import ServerSession as _McpServerSession
|
||||
from mcp.types import (
|
||||
BlobResourceContents,
|
||||
DiscoverRequest,
|
||||
DiscoverResult,
|
||||
GetPromptResult,
|
||||
RequestParams,
|
||||
ResourceTemplate,
|
||||
TextResourceContents,
|
||||
)
|
||||
|
|
@ -526,11 +531,11 @@ if MCP_AVAILABLE:
|
|||
PaginatedRequestParams,
|
||||
ReadResourceRequestParams,
|
||||
)
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import (
|
||||
MCPAuthenticatedUser,
|
||||
)
|
||||
from litellm.proxy._experimental.mcp_server.capabilities import GatewayVersionPolicy, configured_versions
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
|
||||
MCPServerManager,
|
||||
global_mcp_server_manager,
|
||||
|
|
@ -585,6 +590,7 @@ if MCP_AVAILABLE:
|
|||
name=LITELLM_MCP_SERVER_NAME,
|
||||
version=LITELLM_MCP_SERVER_VERSION,
|
||||
)
|
||||
server.middleware.append(GatewayVersionPolicy())
|
||||
server.create_initialization_options = types.MethodType(_gateway_create_initialization_options, server)
|
||||
sse: Final[SseServerTransport] = SseServerTransport("/sse/messages")
|
||||
|
||||
|
|
@ -830,6 +836,7 @@ if MCP_AVAILABLE:
|
|||
client_ip,
|
||||
_mcp_proxy_mode.get(),
|
||||
wire_compat_for(ctx.protocol_version),
|
||||
ctx.protocol_version,
|
||||
)
|
||||
|
||||
async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult:
|
||||
|
|
@ -948,6 +955,11 @@ if MCP_AVAILABLE:
|
|||
ReadResourceRequest(params=params), context
|
||||
)
|
||||
|
||||
async def discover(ctx: ServerRequestContext, params: RequestParams) -> DiscoverResult:
|
||||
async with _legacy_operation_context(ctx, trace=False) as context:
|
||||
return await operations.GatewayOperations().execute(DiscoverRequest(params=params), context)
|
||||
|
||||
server.add_request_handler("server/discover", RequestParams, discover)
|
||||
server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools)
|
||||
server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call)
|
||||
server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts)
|
||||
|
|
@ -1954,7 +1966,7 @@ if MCP_AVAILABLE:
|
|||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
supported: Final = ", ".join(configured_versions())
|
||||
await JSONResponse(
|
||||
status_code=400,
|
||||
content={ # mutable-ok: JSON-RPC error payload
|
||||
|
|
@ -2299,7 +2311,7 @@ if MCP_AVAILABLE:
|
|||
reject_disallowed_mcp_origin(StarletteRequest(scope))
|
||||
bad_version: Final = unsupported_protocol_version(scope)
|
||||
if bad_version is not None:
|
||||
supported: Final = ", ".join(sorted(HANDSHAKE_PROTOCOL_VERSIONS))
|
||||
supported: Final = ", ".join(configured_versions())
|
||||
await JSONResponse(
|
||||
status_code=400,
|
||||
content={ # mutable-ok: JSON-RPC error payload
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ from litellm.types.llms.openai import (
|
|||
ResponsesAPIResponse,
|
||||
)
|
||||
from litellm.types.mcp import (
|
||||
MCPAdvertisedVersions,
|
||||
MCPAllowedClient,
|
||||
MCPAuth,
|
||||
MCPAuthType,
|
||||
|
|
@ -2998,6 +2999,11 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
None,
|
||||
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
|
||||
)
|
||||
mcp_advertised_versions: MCPAdvertisedVersions | None = Field(
|
||||
None,
|
||||
description="MCP revisions enabled by the gateway. Defaults to all completed legacy revisions. "
|
||||
"Modern protocol serving and Apps/Tasks remain disabled.",
|
||||
)
|
||||
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
|
||||
None,
|
||||
description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.",
|
||||
|
|
@ -4013,6 +4019,18 @@ class AllCallbacks(LiteLLMPydanticObjectBase):
|
|||
],
|
||||
)
|
||||
|
||||
zerobus: CallbackOnUI = CallbackOnUI(
|
||||
litellm_callback_name="zerobus",
|
||||
ui_callback_name="Databricks Zerobus",
|
||||
litellm_callback_params=[ # mutable-ok: the registry field is typed list
|
||||
"ZEROBUS_WORKSPACE_URL",
|
||||
"ZEROBUS_SERVER_ENDPOINT",
|
||||
"ZEROBUS_CLIENT_ID",
|
||||
"ZEROBUS_CLIENT_SECRET",
|
||||
"ZEROBUS_TABLE_NAME",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class HTTPExceptionErrorDetail(TypedDict):
|
||||
"""The `{"error": <message>}` shape most proxy endpoints raise as `HTTPException.detail`."""
|
||||
|
|
@ -5349,11 +5367,12 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
|
|||
default=False,
|
||||
description=(
|
||||
"When True, users whose JWT contains no team claims are authenticated "
|
||||
"using their database team memberships instead of receiving HTTP 403. "
|
||||
"Usage is attributed to the user's first resolvable DB team, or to the "
|
||||
"team specified via the x-litellm-team-id request header (validated "
|
||||
"against DB membership). Requires user_id_upsert=True so that user "
|
||||
"records exist before the fallback runs."
|
||||
"using their database team memberships instead of receiving HTTP 403, "
|
||||
"with usage attributed to the user's first resolvable DB team. Whether or "
|
||||
"not the JWT carries team claims, the x-litellm-team-id request header may "
|
||||
"select any team the user is a member of in the database (validated against "
|
||||
"DB membership); without the header the JWT team stays the default. Requires "
|
||||
"user_id_upsert=True so that user records exist before the fallback runs."
|
||||
),
|
||||
)
|
||||
issuers: list[JWTIssuerConfig] | None = Field(
|
||||
|
|
|
|||
|
|
@ -1930,12 +1930,12 @@ class JWTAuthManager:
|
|||
) -> HeaderTeam | None:
|
||||
"""
|
||||
The team named by x-litellm-team-id, which may carry a team id or a team
|
||||
alias. A value that is already an allowed team id (or, under the DB
|
||||
fallback, an existing team id) never costs an alias lookup; an alias is
|
||||
accepted only when the team it names would have been accepted by id.
|
||||
Under the DB fallback only a team row that is provably absent falls
|
||||
through to the alias lookup; a read that failed for any other reason
|
||||
keeps the membership denial the id path already gives.
|
||||
alias. A value that is already an allowed team id never costs a lookup;
|
||||
under the DB fallback any other value is accepted provisionally, by id
|
||||
or alias, for the membership check auth_builder runs later. Under the
|
||||
DB fallback only a team row that is provably absent falls through to
|
||||
the alias lookup; a read that failed for any other reason keeps the
|
||||
membership denial the id path already gives.
|
||||
|
||||
Raises:
|
||||
HTTPException: 403 when neither the value nor the team it aliases is
|
||||
|
|
@ -1948,7 +1948,11 @@ class JWTAuthManager:
|
|||
if not header_value:
|
||||
return None
|
||||
|
||||
if fallback_to_db_teams and not allowed_team_ids:
|
||||
if header_value in allowed_team_ids:
|
||||
verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value)
|
||||
return HeaderTeam(header_value=header_value, team_id=header_value)
|
||||
|
||||
if fallback_to_db_teams:
|
||||
try:
|
||||
await get_team_object(
|
||||
team_id=header_value,
|
||||
|
|
@ -1969,10 +1973,6 @@ class JWTAuthManager:
|
|||
JWTAuthManager._raise_header_team_membership_denial(header_value)
|
||||
return HeaderTeam(header_value=header_value, team_id=header_value)
|
||||
|
||||
if header_value in allowed_team_ids:
|
||||
verbose_proxy_logger.debug("Using team_id from x-litellm-team-id header: %s", header_value)
|
||||
return HeaderTeam(header_value=header_value, team_id=header_value)
|
||||
|
||||
team_id_by_alias: Final = await JWTAuthManager._team_id_by_alias(
|
||||
header_value, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj
|
||||
)
|
||||
|
|
@ -2353,9 +2353,9 @@ class JWTAuthManager:
|
|||
header_value: str,
|
||||
) -> None:
|
||||
"""
|
||||
A provisional team_id from the x-litellm-team-id header (accepted without
|
||||
JWT-team validation when the JWT carries no team claims) must exist in the
|
||||
user's DB team memberships before it becomes request context. The denial
|
||||
A provisional team_id from the x-litellm-team-id header (accepted under
|
||||
fallback_to_db_teams because it is outside the JWT's teams) must exist in
|
||||
the user's DB team memberships before it becomes request context. The denial
|
||||
names `header_value`, the id or alias the caller sent, not `team_id`.
|
||||
"""
|
||||
user_team_ids: Final = user_object.teams if user_object else []
|
||||
|
|
@ -2587,22 +2587,30 @@ class JWTAuthManager:
|
|||
if specific_team_id and not db_team_fallback:
|
||||
all_team_ids.add(specific_team_id)
|
||||
|
||||
header_db_fallback: Final = handler.litellm_jwtauth.fallback_to_db_teams and team_id is None
|
||||
|
||||
header_team: Final = await JWTAuthManager.resolve_team_from_header(
|
||||
request_headers=request_headers,
|
||||
allowed_team_ids=all_team_ids,
|
||||
fallback_to_db_teams=db_team_fallback,
|
||||
fallback_to_db_teams=header_db_fallback,
|
||||
prisma_client=prisma_client,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
provisional_header_team: Final = (
|
||||
header_team
|
||||
if header_team is not None and header_db_fallback and header_team.team_id not in all_team_ids
|
||||
else None
|
||||
)
|
||||
if header_team:
|
||||
team_id = header_team.team_id
|
||||
# A provisional header team (accepted only because the JWT carries no
|
||||
# team claims) is validated against DB membership further down; never
|
||||
# upsert it here or an attacker-supplied x-litellm-team-id would create
|
||||
# an orphaned team row before that check runs. A genuine membership team
|
||||
# already exists, so suppressing the upsert in that case costs nothing.
|
||||
# A provisional header team (accepted because it is outside the
|
||||
# JWT's teams under fallback_to_db_teams) is validated against DB
|
||||
# membership further down; never upsert it here or an
|
||||
# attacker-supplied x-litellm-team-id would create an orphaned team
|
||||
# row before that check runs. A genuine membership team already
|
||||
# exists, so suppressing the upsert in that case costs nothing.
|
||||
try:
|
||||
team_object = await get_team_object(
|
||||
team_id=team_id,
|
||||
|
|
@ -2610,10 +2618,10 @@ class JWTAuthManager:
|
|||
user_api_key_cache=user_api_key_cache,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=(team_id_upsert and not db_team_fallback),
|
||||
team_id_upsert=(team_id_upsert and provisional_header_team is None),
|
||||
)
|
||||
except HTTPException:
|
||||
if not db_team_fallback:
|
||||
if provisional_header_team is None:
|
||||
raise
|
||||
JWTAuthManager._raise_header_team_membership_denial(header_team.header_value)
|
||||
elif not team_id and not db_team_fallback:
|
||||
|
|
@ -2756,11 +2764,11 @@ class JWTAuthManager:
|
|||
proxy_logging_obj=proxy_logging_obj,
|
||||
team_id_upsert=team_id_upsert,
|
||||
)
|
||||
elif db_team_fallback and header_team is not None and team_id == header_team.team_id:
|
||||
elif provisional_header_team is not None and team_id == provisional_header_team.team_id:
|
||||
JWTAuthManager._validate_header_team_in_db_membership(
|
||||
team_id=team_id,
|
||||
user_object=user_object,
|
||||
header_value=header_team.header_value,
|
||||
header_value=provisional_header_team.header_value,
|
||||
)
|
||||
if not JWTAuthManager._is_team_route_allowed(
|
||||
route=route,
|
||||
|
|
@ -2770,7 +2778,7 @@ class JWTAuthManager:
|
|||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=(
|
||||
f"Team '{header_team.header_value}' (from x-litellm-team-id header) "
|
||||
f"Team '{provisional_header_team.header_value}' (from x-litellm-team-id header) "
|
||||
f"is not allowed to access route '{route}'."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ from starlette.types import Receive, Scope, Send
|
|||
import litellm
|
||||
from litellm._logging import redact_internal_details_from_client_message, verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.anthropic_interface.exceptions import AnthropicErrorSseFrame, anthropic_error_sse_frame
|
||||
from litellm.constants import (
|
||||
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE,
|
||||
DEFAULT_MAX_RECURSE_DEPTH,
|
||||
|
|
@ -102,8 +103,11 @@ from litellm.proxy.common_utils.openai_error_payload import (
|
|||
)
|
||||
from litellm.proxy.common_utils.sse_keepalive import (
|
||||
SSE_COMMENT_PING_BYTES,
|
||||
SSE_STREAM_START_TAIL,
|
||||
advance_sse_tail,
|
||||
coerce_keepalive_interval,
|
||||
resolve_ttft_keepalive_interval,
|
||||
seal_open_sse_frame,
|
||||
wrap_sse_stream_with_keepalive_pings,
|
||||
)
|
||||
from litellm.proxy.dd_span_tagger import DDSpanTagger
|
||||
|
|
@ -999,6 +1003,17 @@ async def create_response(
|
|||
first_chunk_value = await _buffer_first_chunk_honoring_disconnect(generator, request)
|
||||
resolved_headers: Final = await _resolve_stream_headers(headers, refresh_headers)
|
||||
|
||||
if isinstance(first_chunk_value, AnthropicErrorSseFrame):
|
||||
with contextlib.suppress(Exception):
|
||||
await generator.aclose()
|
||||
return JSONResponse(
|
||||
status_code=first_chunk_value.status_code,
|
||||
content=first_chunk_value.json_body(
|
||||
error_body_call_id(general_settings, resolved_headers.get(LITELLM_CALL_ID_HEADER))
|
||||
),
|
||||
headers=resolved_headers,
|
||||
)
|
||||
|
||||
if first_chunk_value is not None:
|
||||
try:
|
||||
error_code_from_chunk: Final = await _parse_event_data_for_error(first_chunk_value)
|
||||
|
|
@ -2781,15 +2796,6 @@ class ProxyBaseLLMRequestProcessing:
|
|||
request=request,
|
||||
)
|
||||
if route_type == "aresponses":
|
||||
# Streaming /v1/responses returns here without
|
||||
# reaching the non-streaming ownership tail below.
|
||||
# Wrap the SSE generator so container ownership is
|
||||
# written once the upstream iterator finishes
|
||||
# assembling ``completed_response`` — otherwise
|
||||
# code-interpreter containers created during the
|
||||
# stream stay unregistered and follow-up file API
|
||||
# calls 403. Covers the background-polling path
|
||||
# too, which loops ``body_iterator`` end-to-end.
|
||||
selected_data_generator = (
|
||||
ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership(
|
||||
original_stream_response=response,
|
||||
|
|
@ -3011,50 +3017,50 @@ class ProxyBaseLLMRequestProcessing:
|
|||
wrapped_generator: Any,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
):
|
||||
"""Forward SSE chunks, then record container ownership at stream end.
|
||||
"""Forward SSE chunks and record container ownership before the terminal chunk goes out.
|
||||
|
||||
Streaming ``/v1/responses`` short-circuits out of
|
||||
``base_process_llm_request`` before the non-streaming ownership
|
||||
tail runs, so without this wrap the
|
||||
``LiteLLM_ManagedObjectTable`` row for any container created
|
||||
during the stream is never written and follow-up file API calls
|
||||
return 403.
|
||||
tail runs. The OpenAI SDK closes the connection at ``data: [DONE]``
|
||||
and starlette cancels the body task on disconnect, so a write that
|
||||
waits for the generator to finish never lands. The iterator sets
|
||||
``completed_response`` before it hands over its terminal chunk, so
|
||||
the ``LiteLLM_ManagedObjectTable`` row is written the moment it
|
||||
appears, ahead of the chunk carrying ``response.completed``.
|
||||
"""
|
||||
try:
|
||||
async for chunk in wrapped_generator:
|
||||
async for chunk in wrapped_generator:
|
||||
completed_obj = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
|
||||
original_stream_response
|
||||
)
|
||||
if completed_obj is None:
|
||||
yield chunk
|
||||
finally:
|
||||
try:
|
||||
completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
|
||||
original_stream_response
|
||||
)
|
||||
if completed_obj is not None:
|
||||
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
|
||||
response=completed_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
else:
|
||||
# Silent skip caused #30210: the proxy's Router wrapper
|
||||
# of the responses streaming iterator wasn't propagating
|
||||
# ``completed_response``, so this hook recorded nothing
|
||||
# and follow-up /v1/containers/<id>/files calls 403'd
|
||||
# for non-admin keys with no proxy-side hint. Log a
|
||||
# warning so future regressions of the same shape
|
||||
# surface in operator logs.
|
||||
verbose_proxy_logger.warning(
|
||||
"Container ownership recording skipped on streaming "
|
||||
"/v1/responses: no completed_response on stream "
|
||||
"iterator %s. If this stream created any tool "
|
||||
"container (e.g. code_interpreter), follow-up "
|
||||
"/v1/containers/<id>/files calls will 403 for "
|
||||
"non-admin keys.",
|
||||
type(original_stream_response).__name__,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
"Container ownership recording failed after streaming responses call: %s",
|
||||
e,
|
||||
)
|
||||
continue
|
||||
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
|
||||
response=completed_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
yield chunk
|
||||
async for remaining_chunk in wrapped_generator:
|
||||
yield remaining_chunk
|
||||
return
|
||||
late_completed_obj: Final = ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
|
||||
original_stream_response
|
||||
)
|
||||
if late_completed_obj is not None:
|
||||
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
|
||||
response=late_completed_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return
|
||||
verbose_proxy_logger.warning(
|
||||
"Container ownership recording skipped on streaming "
|
||||
"/v1/responses: no completed_response on stream "
|
||||
"iterator %s. If this stream created any tool "
|
||||
"container (e.g. code_interpreter), follow-up "
|
||||
"/v1/containers/<id>/files calls will 403 for "
|
||||
"non-admin keys.",
|
||||
type(original_stream_response).__name__,
|
||||
)
|
||||
|
||||
async def base_passthrough_process_llm_request(
|
||||
self,
|
||||
|
|
@ -3861,6 +3867,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
serialize_error: StreamErrorSerializer,
|
||||
request: Request | None = None,
|
||||
flush_tail: Callable[[], bytes] | None = None,
|
||||
seal_open_frame: Callable[[bytes], str] | None = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Shared streaming data generator: runs proxy iterator hook, per-chunk hook,
|
||||
|
|
@ -3870,6 +3877,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
``flush_tail`` runs once after the upstream iterator completes cleanly and
|
||||
its non-empty result is yielded, so a serializer that buffers bytes across
|
||||
chunks can emit anything still held at end of stream.
|
||||
|
||||
``seal_open_frame`` is given the tail of what has been yielded when the
|
||||
error frame goes out, and what it returns is written first. A passthrough
|
||||
relays raw upstream bytes, so an upstream that hangs up mid-frame leaves the
|
||||
client inside an open frame, where an error frame would be swallowed or
|
||||
misparsed instead of raised.
|
||||
"""
|
||||
verbose_proxy_logger.debug("inside generator")
|
||||
# Resolve per-stream (not per-chunk) whether the heavy per-chunk path
|
||||
|
|
@ -3886,6 +3899,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
stream_completed = False
|
||||
client_disconnected = False
|
||||
delivered_chunk = False
|
||||
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the yielded bytes
|
||||
try:
|
||||
str_so_far = ""
|
||||
async for chunk in proxy_logging_obj.async_post_call_streaming_iterator_hook(
|
||||
|
|
@ -3931,7 +3945,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
# False and refunds. A keepalive ping carries no provider output,
|
||||
# so it must not suppress that refund.
|
||||
delivered_chunk = delivered_chunk or chunk != STREAM_SSE_KEEPALIVE_PING_BYTES
|
||||
yield serialize_chunk(chunk)
|
||||
serialized = serialize_chunk(chunk)
|
||||
recent_tail = advance_sse_tail(recent_tail, serialized)
|
||||
yield serialized
|
||||
held_tail: Final = flush_tail() if flush_tail is not None else b""
|
||||
if held_tail:
|
||||
yield serialize_chunk(held_tail)
|
||||
|
|
@ -3979,7 +3995,9 @@ class ProxyBaseLLMRequestProcessing:
|
|||
code=stream_error_status,
|
||||
)
|
||||
stream_completed = True
|
||||
yield serialize_error(proxy_exception)
|
||||
error_frame: Final = serialize_error(proxy_exception)
|
||||
seal: Final = "" if seal_open_frame is None else seal_open_frame(recent_tail)
|
||||
yield seal + error_frame if seal else error_frame
|
||||
finally:
|
||||
await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup(
|
||||
request=request,
|
||||
|
|
@ -4001,7 +4019,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
restamp_model: str | None = None,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""
|
||||
Anthropic /messages and Google /generateContent streaming data generator require SSE events.
|
||||
Anthropic /messages streaming data generator, which requires SSE events.
|
||||
|
||||
Returns the underlying ``async_streaming_data_generator`` configured with
|
||||
SSE serializers directly (rather than re-wrapping it in another
|
||||
|
|
@ -4019,11 +4037,13 @@ class ProxyBaseLLMRequestProcessing:
|
|||
request_data=request_data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
serialize_chunk=ProxyBaseLLMRequestProcessing._sse_chunk_serializer(restamper),
|
||||
serialize_error=lambda proxy_exc: (
|
||||
f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n"
|
||||
serialize_error=lambda proxy_exc: anthropic_error_sse_frame(
|
||||
status_code=error_status_code(proxy_exc, status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
raw_message=proxy_exc.message,
|
||||
),
|
||||
request=request,
|
||||
flush_tail=None if restamper is None else restamper.flush,
|
||||
seal_open_frame=seal_open_sse_frame,
|
||||
)
|
||||
|
||||
@overload
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ SSE_COMMENT_PING_BYTES: Final = SSE_COMMENT_PING.encode()
|
|||
# terminates a line with CRLF, LF or CR, so a blank line is any of these three.
|
||||
_SSE_FRAME_DELIMITERS: Final = (b"\r\n\r\n", b"\n\n", b"\r\r")
|
||||
_SSE_DELIMITER_LOOKBACK: Final = max(len(delimiter) for delimiter in _SSE_FRAME_DELIMITERS)
|
||||
_STREAM_START_TAIL: Final = b"\n\n"
|
||||
SSE_STREAM_START_TAIL: Final = b"\n\n"
|
||||
_SSE_MEDIA_TYPE: Final = "text/event-stream"
|
||||
|
||||
|
||||
|
|
@ -128,7 +128,7 @@ async def _keepalive_ping_byte_stream(
|
|||
# Seeded as a delimiter because a stream starts at a frame boundary, and kept
|
||||
# across chunks because a delimiter can be split between two transport reads,
|
||||
# which testing only the latest chunk would miss for the rest of the stream.
|
||||
recent_tail = _STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes
|
||||
recent_tail = SSE_STREAM_START_TAIL # rebind-ok: rolling window over the relayed bytes
|
||||
try:
|
||||
while True:
|
||||
await asyncio.wait((pending,), timeout=ping_interval_seconds)
|
||||
|
|
@ -155,6 +155,28 @@ async def _keepalive_ping_byte_stream(
|
|||
await stream.aclose()
|
||||
|
||||
|
||||
def advance_sse_tail(recent_tail: bytes, chunk: object) -> bytes:
|
||||
written: Final = _sse_tail_bytes(chunk)
|
||||
if not written:
|
||||
return recent_tail
|
||||
return (recent_tail + written)[-_SSE_DELIMITER_LOOKBACK:]
|
||||
|
||||
|
||||
def _sse_tail_bytes(chunk: object) -> bytes:
|
||||
if isinstance(chunk, bytes):
|
||||
return chunk[-_SSE_DELIMITER_LOOKBACK:]
|
||||
if isinstance(chunk, str):
|
||||
return chunk[-_SSE_DELIMITER_LOOKBACK:].encode()
|
||||
return b""
|
||||
|
||||
|
||||
def seal_open_sse_frame(recent_tail: bytes) -> str:
|
||||
if recent_tail.endswith(_SSE_FRAME_DELIMITERS):
|
||||
return ""
|
||||
line_break: Final = "" if recent_tail.endswith((b"\n", b"\r")) else "\n"
|
||||
return f"{line_break}{ANTHROPIC_PING_SSE_CHUNK}"
|
||||
|
||||
|
||||
def resolve_ttft_keepalive_interval(
|
||||
deployments: Iterable[Mapping[str, object]],
|
||||
global_interval: float | str | None,
|
||||
|
|
|
|||
|
|
@ -4618,6 +4618,13 @@ class GoogleSSOHandler:
|
|||
return result or {}
|
||||
|
||||
|
||||
def _raise_if_sso_debug_disabled() -> None:
|
||||
"""The debug routes run the browser-redirect SSO flow, so they cannot carry a
|
||||
bearer credential; an explicit opt-in flag is the only way to gate them."""
|
||||
if get_secret_bool("ENABLE_SSO_DEBUG") is not True:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not Found")
|
||||
|
||||
|
||||
@router.get("/sso/debug/login", tags=["experimental"], include_in_schema=False)
|
||||
async def debug_sso_login(request: Request):
|
||||
"""
|
||||
|
|
@ -4625,6 +4632,8 @@ async def debug_sso_login(request: Request):
|
|||
PROXY_BASE_URL should be the your deployed proxy endpoint, e.g. PROXY_BASE_URL="https://litellm-production-7002.up.railway.app/"
|
||||
Example:
|
||||
"""
|
||||
_raise_if_sso_debug_disabled()
|
||||
|
||||
from litellm.proxy.proxy_server import premium_user
|
||||
|
||||
microsoft_client_id: Final = os.getenv("MICROSOFT_CLIENT_ID", None)
|
||||
|
|
@ -4670,6 +4679,8 @@ async def debug_sso_callback(request: Request):
|
|||
"""
|
||||
Returns the OpenID object returned by the SSO provider
|
||||
"""
|
||||
_raise_if_sso_debug_disabled()
|
||||
|
||||
import json
|
||||
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence
|
|||
from datetime import datetime
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, overload, runtime_checkable
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, NoReturn, Protocol, overload, runtime_checkable
|
||||
|
||||
import httpx
|
||||
from openai._streaming import SSEDecoder
|
||||
|
|
@ -230,6 +230,11 @@ def _status_code_for_error_fields(error_type: str | None, error_code: str | None
|
|||
return next((status for status in map(_status_code_for_error_field, fields) if status is not None), 500)
|
||||
|
||||
|
||||
def stream_error_status_and_message(error_obj: object) -> tuple[int, str]:
|
||||
message, error_type, error_code = _error_event_fields(error_obj)
|
||||
return _status_code_for_error_fields(error_type, error_code), message
|
||||
|
||||
|
||||
def _map_stream_error_to_exception(error_obj: object, model: str, custom_llm_provider: str) -> Exception:
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
|
|
@ -260,6 +265,9 @@ def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool:
|
|||
return not isinstance(status_code, int) or status_code >= 500 or status_code == 429
|
||||
|
||||
|
||||
_PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"})
|
||||
|
||||
|
||||
class BaseResponsesAPIStreamingIterator:
|
||||
"""
|
||||
Base class for streaming iterators that process responses from the Responses API.
|
||||
|
|
@ -287,6 +295,7 @@ class BaseResponsesAPIStreamingIterator:
|
|||
self.start_time = getattr(logging_obj, "start_time", datetime.now())
|
||||
self._failure_handled = False # Track if failure handler has been called
|
||||
self._yielded_first_chunk = False
|
||||
self._output_started = False
|
||||
self._generated_content = ""
|
||||
self._generated_tool_arguments = ""
|
||||
self._completed_response_cached = False
|
||||
|
|
@ -874,6 +883,46 @@ class BaseResponsesAPIStreamingIterator:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
def _note_yielded_event(self, event: ResponsesAPIStreamingResponse) -> None:
|
||||
self._yielded_first_chunk = True
|
||||
if event.type not in _PRE_OUTPUT_LIFECYCLE_EVENT_TYPES:
|
||||
self._output_started = True
|
||||
|
||||
def _fallback_error(self, original: Exception) -> MidStreamFallbackError:
|
||||
return MidStreamFallbackError(
|
||||
message=str(original),
|
||||
model=self.model or "",
|
||||
llm_provider=self.custom_llm_provider or "",
|
||||
original_exception=original,
|
||||
generated_content="",
|
||||
is_pre_first_chunk=not self._yielded_first_chunk,
|
||||
)
|
||||
|
||||
def _stream_ended_early_error(self) -> litellm.APIConnectionError:
|
||||
return litellm.APIConnectionError(
|
||||
message=(
|
||||
f"{self.custom_llm_provider or 'provider'} closed the responses stream before any terminal event "
|
||||
"(response.completed, response.incomplete or response.failed)"
|
||||
),
|
||||
llm_provider=self.custom_llm_provider or "",
|
||||
model=self.model or "",
|
||||
)
|
||||
|
||||
def _raise_if_ended_without_terminal_event(self) -> None:
|
||||
if self.completed_response is not None:
|
||||
return
|
||||
error: Final = self._stream_ended_early_error()
|
||||
self._handle_failure(error)
|
||||
if self._output_started:
|
||||
raise error
|
||||
raise self._fallback_error(error) from error
|
||||
|
||||
def _raise_for_transport_error(self, error: httpx.ReadError | httpx.RemoteProtocolError) -> NoReturn:
|
||||
self._handle_failure(error)
|
||||
if self._output_started:
|
||||
raise error
|
||||
raise self._fallback_error(error) from error
|
||||
|
||||
|
||||
async def call_post_streaming_hooks_for_testing(
|
||||
iterator: object, chunk: ResponsesAPIStreamingResponse
|
||||
|
|
@ -929,12 +978,14 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
sse = await self.stream_iterator.__anext__()
|
||||
except StopAsyncIteration:
|
||||
self.finished = True
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopAsyncIteration
|
||||
|
||||
self._check_max_streaming_duration()
|
||||
result = self._process_chunk(sse.data)
|
||||
|
||||
if self.finished:
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopAsyncIteration
|
||||
elif result is not None:
|
||||
self._maybe_raise_for_error_event(result)
|
||||
|
|
@ -943,7 +994,7 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
result = await self._call_post_streaming_deployment_hook(
|
||||
chunk=result,
|
||||
)
|
||||
self._yielded_first_chunk = True
|
||||
self._note_yielded_event(result)
|
||||
return result
|
||||
# If result is None, continue the loop to get the next chunk
|
||||
|
||||
|
|
@ -952,10 +1003,9 @@ class ResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
raise
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
self.finished = True
|
||||
if self.completed_response is None:
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
raise StopAsyncIteration from e
|
||||
if self.completed_response is not None:
|
||||
raise StopAsyncIteration from e
|
||||
self._raise_for_transport_error(e)
|
||||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
|
|
@ -1011,12 +1061,14 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
sse = next(self.stream_iterator)
|
||||
except StopIteration:
|
||||
self.finished = True
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopIteration
|
||||
|
||||
self._check_max_streaming_duration()
|
||||
result = self._process_chunk(sse.data)
|
||||
|
||||
if self.finished:
|
||||
self._raise_if_ended_without_terminal_event()
|
||||
raise StopIteration
|
||||
elif result is not None:
|
||||
self._maybe_raise_for_error_event(result)
|
||||
|
|
@ -1025,7 +1077,7 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
async_function=self._call_post_streaming_deployment_hook,
|
||||
chunk=result,
|
||||
)
|
||||
self._yielded_first_chunk = True
|
||||
self._note_yielded_event(result)
|
||||
return result
|
||||
# If result is None, continue the loop to get the next chunk
|
||||
|
||||
|
|
@ -1034,10 +1086,9 @@ class SyncResponsesAPIStreamingIterator(BaseResponsesAPIStreamingIterator):
|
|||
raise
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
self.finished = True
|
||||
if self.completed_response is None:
|
||||
self._handle_failure(e)
|
||||
raise
|
||||
raise StopIteration from e
|
||||
if self.completed_response is not None:
|
||||
raise StopIteration from e
|
||||
self._raise_for_transport_error(e)
|
||||
except httpx.HTTPError as e:
|
||||
# Handle HTTP errors
|
||||
self.finished = True
|
||||
|
|
|
|||
|
|
@ -502,12 +502,12 @@ class RouterBudgetLimiting(CustomLogger):
|
|||
|
||||
response_cost: Final[float] = standard_logging_payload.get("response_cost", 0)
|
||||
model_id: Final[str] = str(standard_logging_payload.get("model_id", ""))
|
||||
custom_llm_provider: Final[str] = kwargs.get("litellm_params", {}).get("custom_llm_provider", None)
|
||||
if custom_llm_provider is None:
|
||||
raise ValueError("custom_llm_provider is required")
|
||||
custom_llm_provider: Final[str | None] = standard_logging_payload.get("custom_llm_provider")
|
||||
|
||||
budget_config: Final = self._get_budget_config_for_provider(custom_llm_provider)
|
||||
if budget_config:
|
||||
budget_config: Final = (
|
||||
self._get_budget_config_for_provider(custom_llm_provider) if custom_llm_provider is not None else None
|
||||
)
|
||||
if custom_llm_provider is not None and budget_config is not None:
|
||||
# increment spend for provider
|
||||
spend_key: Final = f"provider_spend:{custom_llm_provider}:{budget_config.budget_duration}"
|
||||
start_time_key: Final = f"provider_budget_start_time:{custom_llm_provider}"
|
||||
|
|
|
|||
53
litellm/types/integrations/zerobus.py
Normal file
53
litellm/types/integrations/zerobus.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
from dataclasses import dataclass, field
|
||||
from typing import Final
|
||||
|
||||
from pydantic import Field
|
||||
|
||||
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
|
||||
|
||||
RETRYABLE_INGEST_STATUS_CODES: Final = frozenset({408, 429, 500, 502, 503, 504})
|
||||
|
||||
TOKEN_REFRESH_LEEWAY_SECONDS: Final = 60
|
||||
|
||||
|
||||
class ZerobusInitParams(StandardCustomLoggerInitParams):
|
||||
"""
|
||||
Params for initializing a Databricks Zerobus logger on litellm.
|
||||
|
||||
Every connection field falls back to its ``ZEROBUS_*`` environment variable, which is
|
||||
what the proxy UI writes. ``table_name`` is the fully qualified ``catalog.schema.table``.
|
||||
"""
|
||||
|
||||
workspace_url: str | None = None
|
||||
server_endpoint: str | None = None
|
||||
client_id: str | None = None
|
||||
client_secret: str | None = None
|
||||
table_name: str | None = None
|
||||
batch_size: int = Field(default=100, gt=0)
|
||||
flush_interval: int = Field(default=10, gt=0)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ZerobusConnection:
|
||||
"""Everything needed to mint a token for one table and post rows to it."""
|
||||
|
||||
workspace_url: str
|
||||
workspace_id: str
|
||||
server_endpoint: str
|
||||
client_id: str
|
||||
client_secret: str = field(repr=False)
|
||||
table_name: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ZerobusAccessToken:
|
||||
value: str = field(repr=False)
|
||||
expires_at: float
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ZerobusIngestFailure:
|
||||
"""Why a batch could not be written, and whether a later attempt could still succeed."""
|
||||
|
||||
detail: str
|
||||
retryable: bool
|
||||
|
|
@ -522,7 +522,19 @@ class CreateBatchRequest(TypedDict, total=False):
|
|||
"""
|
||||
|
||||
completion_window: Literal["24h"]
|
||||
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions", "/v1/responses", "/v1/ocr"]
|
||||
endpoint: Literal[
|
||||
"/v1/chat/completions",
|
||||
"/v1/embeddings",
|
||||
"/v1/completions",
|
||||
"/v1/responses",
|
||||
"/v1/ocr",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/videos/generations",
|
||||
"/v1/videos",
|
||||
"/v1/videos/edits",
|
||||
"/v1/videos/extensions",
|
||||
]
|
||||
input_file_id: str
|
||||
metadata: dict[str, str] | None
|
||||
output_expires_after: FileExpiresAfter
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import enum
|
|||
import re
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal
|
||||
from typing import TYPE_CHECKING, Annotated, Any, Final, Literal
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
import httpx
|
||||
|
|
@ -34,6 +34,8 @@ class MCPSpecVersion(str, enum.Enum):
|
|||
nov_2024 = "2024-11-05"
|
||||
mar_2025 = "2025-03-26"
|
||||
jun_2025 = "2025-06-18"
|
||||
nov_2025 = "2025-11-25"
|
||||
jul_2026 = "2026-07-28"
|
||||
|
||||
|
||||
class MCPAuth(str, enum.Enum):
|
||||
|
|
@ -59,7 +61,17 @@ DEFAULT_SUBJECT_TOKEN_TYPE: Final = "urn:ietf:params:oauth:token-type:access_tok
|
|||
|
||||
# MCP Literals
|
||||
MCPTransportType = Literal[MCPTransport.sse, MCPTransport.http, MCPTransport.stdio]
|
||||
MCPSpecVersionType = Literal[MCPSpecVersion.nov_2024, MCPSpecVersion.mar_2025, MCPSpecVersion.jun_2025]
|
||||
MCPLegacyVersion = Literal["2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"]
|
||||
MCP_LEGACY_VERSIONS: Final[tuple[MCPLegacyVersion, ...]] = ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25")
|
||||
MCPUpstreamProtocol = MCPLegacyVersion | Literal["auto"]
|
||||
MCPAdvertisedVersions = Annotated[tuple[MCPLegacyVersion, ...], Field(min_length=1)]
|
||||
MCPSpecVersionType = Literal[
|
||||
MCPSpecVersion.nov_2024,
|
||||
MCPSpecVersion.mar_2025,
|
||||
MCPSpecVersion.jun_2025,
|
||||
MCPSpecVersion.nov_2025,
|
||||
MCPSpecVersion.jul_2026,
|
||||
]
|
||||
MCPAuthType = (
|
||||
Literal[
|
||||
MCPAuth.none,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
from datetime import datetime
|
||||
from typing import Any, Final, Literal
|
||||
from typing import Annotated, Any, Final, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
||||
from pydantic import AfterValidator, BaseModel, ConfigDict, Field, TypeAdapter, field_validator, model_validator
|
||||
from typing_extensions import Self
|
||||
|
||||
from litellm.types.mcp import (
|
||||
|
|
@ -10,11 +10,19 @@ from litellm.types.mcp import (
|
|||
MCPAuthType,
|
||||
MCPTokenEndpointAuthMethod,
|
||||
MCPTransportType,
|
||||
MCPUpstreamProtocol,
|
||||
normalize_upstream_header_name,
|
||||
)
|
||||
|
||||
|
||||
# MCPInfo now allows arbitrary additional fields for custom metadata
|
||||
MCPInfo = dict[str, Any]
|
||||
def _validate_mcp_protocol_metadata(value: dict[str, object]) -> dict[str, object]:
|
||||
if "protocol_version" in value:
|
||||
TypeAdapter(MCPUpstreamProtocol).validate_python(value["protocol_version"])
|
||||
return value
|
||||
|
||||
|
||||
MCPInfo = Annotated[dict[str, Any], AfterValidator(_validate_mcp_protocol_metadata)]
|
||||
|
||||
|
||||
class MCPOAuthMetadata(BaseModel):
|
||||
|
|
@ -66,6 +74,7 @@ class MCPServer(BaseModel):
|
|||
server_name: str | None = None
|
||||
url: str | None = None
|
||||
transport: MCPTransportType
|
||||
protocol_version: MCPUpstreamProtocol = "auto"
|
||||
spec_path: str | None = None
|
||||
auth_type: MCPAuthType | None = None
|
||||
authentication_token: str | None = None
|
||||
|
|
@ -246,6 +255,14 @@ class MCPServer(BaseModel):
|
|||
"""
|
||||
return self.oauth2_flow == "client_credentials"
|
||||
|
||||
@model_validator(mode="after")
|
||||
def resolve_protocol_version(self) -> Self:
|
||||
if "protocol_version" not in self.model_fields_set and self.mcp_info is not None:
|
||||
self.protocol_version = TypeAdapter(MCPUpstreamProtocol).validate_python(
|
||||
self.mcp_info.get("protocol_version", "auto")
|
||||
)
|
||||
return self
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_identity_binding_mode(self) -> Self:
|
||||
binding: Final = self.oauth_identity_binding
|
||||
|
|
|
|||
|
|
@ -345,6 +345,7 @@ class CredentialLiteLLMParams(BaseModel):
|
|||
|
||||
## OBJECT STORAGE (files / batches) ##
|
||||
gcs_bucket_name: str | None = None
|
||||
bucket_name: str | None = None
|
||||
|
||||
## AWS BEDROCK / SAGEMAKER ##
|
||||
aws_access_key_id: str | None = None
|
||||
|
|
|
|||
|
|
@ -299,6 +299,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
cache_read_input_token_cost_above_272k_tokens_flex: float | None
|
||||
cache_read_input_token_cost_above_512k_tokens: float | None
|
||||
cache_read_input_token_cost_batches: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
cache_read_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
cache_creation_input_token_cost_batches: ReadOnly[float | None]
|
||||
cache_creation_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
|
|
@ -327,8 +328,10 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False):
|
|||
input_cost_per_second: float | None # for OpenAI Speech models
|
||||
input_cost_per_token_batches: float | None
|
||||
input_cost_per_video_token_batches: ReadOnly[float | None]
|
||||
input_cost_per_token_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
input_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
output_cost_per_token_batches: float | None
|
||||
output_cost_per_token_above_200k_tokens_batches: ReadOnly[float | None]
|
||||
output_cost_per_token_above_272k_tokens_batches: ReadOnly[float | None]
|
||||
output_cost_per_token: Required[float | None]
|
||||
output_cost_per_token_flex: float | None # OpenAI flex service tier pricing
|
||||
|
|
@ -3729,6 +3732,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
cache_read_input_token_cost_above_272k_tokens_priority: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_flex: float | None = None
|
||||
cache_read_input_token_cost_batches: float | None = None
|
||||
cache_read_input_token_cost_above_200k_tokens_batches: float | None = None
|
||||
cache_read_input_token_cost_above_272k_tokens_batches: float | None = None
|
||||
cache_creation_input_token_cost_batches: float | None = None
|
||||
cache_creation_input_token_cost_above_272k_tokens_batches: float | None = None
|
||||
|
|
@ -3742,6 +3746,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
input_cost_per_token_above_200k_tokens_priority: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_priority: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_flex: float | None = None
|
||||
input_cost_per_token_above_200k_tokens_batches: float | None = None
|
||||
input_cost_per_token_above_272k_tokens_batches: float | None = None
|
||||
input_cost_per_query: float | None = None
|
||||
input_cost_per_image: float | None = None
|
||||
|
|
@ -3766,6 +3771,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams):
|
|||
output_cost_per_token_above_200k_tokens_priority: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_priority: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_flex: float | None = None
|
||||
output_cost_per_token_above_200k_tokens_batches: float | None = None
|
||||
output_cost_per_token_above_272k_tokens_batches: float | None = None
|
||||
output_cost_per_character_above_128k_tokens: float | None = None
|
||||
output_cost_per_image: float | None = None
|
||||
|
|
@ -4140,7 +4146,7 @@ FILE_CONTENT_STREAMING_PROVIDERS: Final[frozenset[str]] = frozenset(
|
|||
|
||||
LITELLM_EXECUTED_BATCH_PROVIDERS: Final[frozenset[str]] = frozenset({LlmProviders.HOSTED_VLLM.value})
|
||||
|
||||
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai"]
|
||||
ListBatchesSupportedProvider = Literal["openai", "azure", "hosted_vllm", "litellm_proxy", "vertex_ai", "xai"]
|
||||
|
||||
LIST_BATCHES_SUPPORTED_PROVIDERS: Final[frozenset[str]] = frozenset(get_args(ListBatchesSupportedProvider))
|
||||
|
||||
|
|
|
|||
|
|
@ -6160,6 +6160,9 @@ def _get_model_info_helper(
|
|||
cache_read_input_token_cost_priority=_model_info.get("cache_read_input_token_cost_priority", None),
|
||||
cache_read_input_token_cost_ultrafast=_model_info.get("cache_read_input_token_cost_ultrafast", None),
|
||||
cache_read_input_token_cost_batches=_model_info.get("cache_read_input_token_cost_batches"),
|
||||
cache_read_input_token_cost_above_200k_tokens_batches=_model_info.get(
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches"
|
||||
),
|
||||
cache_read_input_token_cost_above_272k_tokens_batches=_model_info.get(
|
||||
"cache_read_input_token_cost_above_272k_tokens_batches"
|
||||
),
|
||||
|
|
@ -6197,10 +6200,16 @@ def _get_model_info_helper(
|
|||
input_cost_per_video_per_second=_model_info.get("input_cost_per_video_per_second", None),
|
||||
input_cost_per_token_batches=_model_info.get("input_cost_per_token_batches"),
|
||||
input_cost_per_video_token_batches=_model_info.get("input_cost_per_video_token_batches", None),
|
||||
input_cost_per_token_above_200k_tokens_batches=_model_info.get(
|
||||
"input_cost_per_token_above_200k_tokens_batches"
|
||||
),
|
||||
input_cost_per_token_above_272k_tokens_batches=_model_info.get(
|
||||
"input_cost_per_token_above_272k_tokens_batches"
|
||||
),
|
||||
output_cost_per_token_batches=_model_info.get("output_cost_per_token_batches"),
|
||||
output_cost_per_token_above_200k_tokens_batches=_model_info.get(
|
||||
"output_cost_per_token_above_200k_tokens_batches"
|
||||
),
|
||||
output_cost_per_token_above_272k_tokens_batches=_model_info.get(
|
||||
"output_cost_per_token_above_272k_tokens_batches"
|
||||
),
|
||||
|
|
@ -9357,6 +9366,10 @@ class ProviderConfigManager:
|
|||
from litellm.llms.mistral.files.transformation import MistralFilesConfig
|
||||
|
||||
return MistralFilesConfig()
|
||||
elif LlmProviders.XAI == provider:
|
||||
from litellm.llms.xai.files.transformation import XAIFilesConfig
|
||||
|
||||
return XAIFilesConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -51436,13 +51436,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -51450,8 +51453,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -51480,9 +51486,13 @@
|
|||
"xai/grok-4.3": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51490,6 +51500,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -51502,9 +51514,13 @@
|
|||
"xai/grok-4.3-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
|
|
@ -51512,6 +51528,8 @@
|
|||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59483,13 +59501,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59497,20 +59518,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": false,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -59519,8 +59546,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true,
|
||||
"supported_endpoints": [
|
||||
|
|
@ -62787,13 +62817,16 @@
|
|||
},
|
||||
"xai/grok-4.20": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62801,21 +62834,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62823,21 +62862,27 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -62845,8 +62890,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
@ -63067,13 +63115,16 @@
|
|||
},
|
||||
"xai/grok-4.20-non-reasoning": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63081,20 +63132,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-non-reasoning-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_prompt_caching": true,
|
||||
|
|
@ -63102,20 +63159,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63127,20 +63190,26 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
"xai/grok-4.20-multi-agent-latest": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "responses",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supported_endpoints": [
|
||||
"/v1/responses"
|
||||
|
|
@ -63152,8 +63221,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_response_schema": true
|
||||
},
|
||||
|
|
@ -75459,13 +75531,16 @@
|
|||
},
|
||||
"xai/grok-4.20-0309": {
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_batches": 1.6e-07,
|
||||
"input_cost_per_token": 1.25e-06,
|
||||
"input_cost_per_token_batches": 1e-06,
|
||||
"litellm_provider": "xai",
|
||||
"max_input_tokens": 1000000,
|
||||
"max_output_tokens": 1000000,
|
||||
"max_tokens": 1000000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.5e-06,
|
||||
"output_cost_per_token_batches": 2e-06,
|
||||
"source": "https://api.x.ai/v1/language-models",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -75473,8 +75548,11 @@
|
|||
"supports_vision": true,
|
||||
"supports_web_search": true,
|
||||
"input_cost_per_token_above_200k_tokens": 2.5e-06,
|
||||
"input_cost_per_token_above_200k_tokens_batches": 2e-06,
|
||||
"output_cost_per_token_above_200k_tokens": 5e-06,
|
||||
"output_cost_per_token_above_200k_tokens_batches": 4e-06,
|
||||
"cache_read_input_token_cost_above_200k_tokens": 4e-07,
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": 3.2e-07,
|
||||
"input_cost_per_image_token": 1.25e-06,
|
||||
"supports_prompt_caching": true,
|
||||
"supports_response_schema": true
|
||||
|
|
|
|||
|
|
@ -170,6 +170,11 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"cache_read_input_token_cost_above_200k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
@ -351,6 +356,11 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"input_cost_per_token_above_200k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"input_cost_per_token_above_200k_tokens_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
@ -708,6 +718,11 @@
|
|||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"output_cost_per_token_above_200k_tokens_batches": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
"description": "Rate applied once the prompt exceeds the token threshold in the field name."
|
||||
},
|
||||
"output_cost_per_token_above_200k_tokens_priority": {
|
||||
"type": "number",
|
||||
"minimum": 0,
|
||||
|
|
|
|||
|
|
@ -249,6 +249,7 @@ proxy-dev = [
|
|||
"prisma==0.11.0",
|
||||
"hypercorn==0.17.3",
|
||||
"prometheus-client==0.20.0",
|
||||
"sentry-sdk==2.21.0",
|
||||
"opentelemetry-api==1.33.1",
|
||||
"opentelemetry-sdk==1.33.1",
|
||||
"opentelemetry-exporter-otlp==1.33.1",
|
||||
|
|
|
|||
|
|
@ -72,6 +72,7 @@ IGNORE_FUNCTIONS = [
|
|||
"_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible).
|
||||
"_replace_string_leaves", # bounded by the nesting depth of a safe_json_structure output (a finite JSON tree, no cycles possible).
|
||||
"_sort_processed_sets", # bounded by the nesting depth of the log-record extra it walks (a finite JSON tree, no cycles possible).
|
||||
"scrub_json_strings", # max depth set (MAX_SCRUB_DEPTH); fails closed by returning "[Filtered]" for anything nested past the cap.
|
||||
]
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -85,6 +85,7 @@
|
|||
- {id: llm.responses.azure_openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Azure OpenAI (smoke)"}
|
||||
- {id: llm.responses.azure_openai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Responses tool calls w/ Azure OpenAI"}
|
||||
- {id: llm.responses.azure_openai.code_interpreter.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: nonstream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "An implicit code_interpreter container on an Azure deployment that carries its own api_base must serve GET /v1/containers/{id}/files/{fid}/content by its native cntr_ id to a team service-account key, the customer's shape (#27921, #28990)", fail_before_fix: proven}
|
||||
- {id: llm.responses.azure_openai.code_interpreter.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: azure_openai, capability: code_interpreter, streaming: stream, assertions: [works], source: "llm_translation/test_containers_e2e.py", rationale: "A container created by a streamed /v1/responses code_interpreter call must serve /v1/containers/{id}/files to the same service-account key right after the OpenAI SDK closes at [DONE]; the ownership row used to be written after the stream and the disconnect cancelled it (LIT-8612)", fail_before_fix: proven}
|
||||
- {id: llm.chat_completions.together_ai.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning surfaces as reasoning_content (LIT-5960)"}
|
||||
- {id: llm.chat_completions.together_ai.thinking.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together reasoning deltas stream as reasoning_content"}
|
||||
- {id: llm.chat_completions.together_ai.thinking.nonstream.template_kwargs_forwarded, module: llm, tier: P1, subject_endpoint: chat_completions, route: together_ai, capability: thinking, streaming: nonstream, assertions: [template_kwargs_forwarded], source: "llm_translation/test_together_ai_e2e.py", rationale: "chat_template_kwargs reaches Together and turns thinking off"}
|
||||
|
|
@ -102,6 +103,8 @@
|
|||
- {id: llm.chat_completions.anthropic.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic over /chat/completions: cost header and spend row agree"}
|
||||
- {id: llm.chat_completions.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /chat/completions"}
|
||||
- {id: llm.messages.anthropic.multi_turn.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: anthropic, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "Anthropic tool result round trip over /v1/messages"}
|
||||
- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_event, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_event], source: "customer report", rationale: "An upstream that hangs up mid-stream must reach Anthropic clients as an event: error frame, not an OpenAI-shaped data-only error they silently drop"}
|
||||
- {id: llm.messages.anthropic.upstream_stream_failure.stream.error_status, module: llm, tier: P1, subject_endpoint: messages, route: anthropic, capability: upstream_stream_failure, streaming: stream, assertions: [error_status], source: "customer report", rationale: "An upstream that hangs up before its first byte must answer as a JSON error carrying its status, so Anthropic clients raise the status-specific error and retry on it instead of reading a 200 stream that only carries an error event"}
|
||||
- {id: llm.messages.openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI models served on the Anthropic Messages contract"}
|
||||
- {id: llm.messages.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages streams the Anthropic event grammar"}
|
||||
- {id: llm.messages.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: messages, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "llm_translation/test_conversational_matrix_e2e.py", rationale: "OpenAI over /v1/messages: cost header and spend row agree"}
|
||||
|
|
|
|||
|
|
@ -86,6 +86,7 @@ LlmCapability = Literal[
|
|||
"tool_search",
|
||||
"tool_search_history",
|
||||
"tool_use",
|
||||
"upstream_stream_failure",
|
||||
"vision",
|
||||
"web_search",
|
||||
"web_search_server_tool",
|
||||
|
|
|
|||
|
|
@ -48,7 +48,7 @@ most likely to silently break and the one a mock can't prove works.
|
|||
|----------|---------------|-----------|------------|-------------|--------|
|
||||
| Chat | live (spend suite) | live (spend suite) | gap | live | partial |
|
||||
| Embeddings | live (spend suite) | n/a | n/a | live | covered |
|
||||
| Responses (Azure code_interpreter container files) | live | gap | live | gap | partial |
|
||||
| Responses (Azure code_interpreter container files) | live | live | live | gap | partial |
|
||||
| Image / audio / rerank / realtime | - | - | - | - | gap |
|
||||
|
||||
## This suite's files
|
||||
|
|
@ -63,6 +63,7 @@ most likely to silently break and the one a mock can't prove works.
|
|||
| `test_anthropic_passthrough_tool_call_logs_cost` | anthropic native, tool call, cost |
|
||||
| `test_vertex_passthrough_via_managed_model_logs_cost` | vertex_ai native, non-stream, cost |
|
||||
| `test_service_account_key_reads_container_file_by_native_id` | azure responses code_interpreter, non-stream, native container id, service-account key |
|
||||
| `test_service_account_key_reads_container_file_created_by_a_streamed_response` | azure responses code_interpreter, stream, native container id, service-account key, upload right after `[DONE]` |
|
||||
|
||||
Vertex keeps the credential on the proxy like gemini/anthropic, but the deployment is
|
||||
added at runtime instead of declared in the gateway config: the test POSTs `/model/new`
|
||||
|
|
|
|||
|
|
@ -31,10 +31,11 @@ A proxy whose env carries ``AZURE_API_BASE`` for the same resource masks the
|
|||
second regression, since the global-credential fallback then reaches the
|
||||
container anyway.
|
||||
|
||||
The streaming variant is not here: a streamed ``/v1/responses`` writes the
|
||||
container ownership row only after the ``[DONE]`` frame, and the OpenAI SDK
|
||||
closes the connection at ``[DONE]``, so the write is cancelled and every
|
||||
follow-up container call 403s (LIT-8612). That cell comes with its fix.
|
||||
The streaming cell repeats the flow with ``stream=True`` and uploads right
|
||||
after the last event. The OpenAI SDK closes the connection at ``[DONE]``, so an
|
||||
ownership row written after the stream is cancelled with the body task and every
|
||||
follow-up container call 403s (LIT-8612); the row has to land before the
|
||||
``response.completed`` frame goes out.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -52,7 +53,7 @@ from lifecycle import ResourceManager
|
|||
from management.management_client import ManagementClient, build_client
|
||||
from models import KeyGenerateBody, KeyGenerateResponse, LiteLLMParamsBody, TeamNewBody, UserNewBody
|
||||
from openai import OpenAI
|
||||
from openai.types.responses import Response, ResponseCodeInterpreterToolCall
|
||||
from openai.types.responses import Response, ResponseCodeInterpreterToolCall, ResponseCompletedEvent
|
||||
from openai.types.responses.tool_param import CodeInterpreter
|
||||
from proxy_client import ProxyClient
|
||||
from sdk_clients import NO_PROXY_CACHE, SdkClients
|
||||
|
|
@ -120,6 +121,24 @@ def _response_with_code_interpreter(client: OpenAI, model: str) -> Response:
|
|||
)
|
||||
|
||||
|
||||
def _streamed_response_with_code_interpreter(client: OpenAI, model: str) -> Response:
|
||||
events: Final = tuple(
|
||||
client.with_options(timeout=CODE_INTERPRETER_TIMEOUT).responses.create(
|
||||
model=model,
|
||||
input=PROMPT,
|
||||
tools=[CODE_INTERPRETER],
|
||||
tool_choice="required",
|
||||
stream=True,
|
||||
extra_body=NO_PROXY_CACHE,
|
||||
)
|
||||
)
|
||||
assert events, "responses stream returned no events"
|
||||
assert isinstance(events[-1], ResponseCompletedEvent), (
|
||||
f"responses stream did not terminate with response.completed: {events[-1].type}"
|
||||
)
|
||||
return events[-1].response
|
||||
|
||||
|
||||
def _container_id(response: Response) -> str:
|
||||
calls: Final = tuple(item for item in response.output if isinstance(item, ResponseCodeInterpreterToolCall))
|
||||
assert calls, f"no code_interpreter_call in the responses output: {response.output!r}"
|
||||
|
|
@ -165,3 +184,17 @@ class TestAzureContainerFiles:
|
|||
f"container id is not the provider's own id: {native_id}"
|
||||
)
|
||||
_assert_file_round_trip(client, native_id, marker)
|
||||
|
||||
@pytest.mark.covers("llm.responses.azure_openai.code_interpreter.stream.works")
|
||||
def test_service_account_key_reads_container_file_created_by_a_streamed_response(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, sdk: SdkClients
|
||||
) -> None:
|
||||
marker: Final = unique_marker()
|
||||
model: Final = _register_two_azure_deployments(proxy, resources, marker)
|
||||
key: Final = _service_account_key(proxy, resources, build_client(proxy), marker, model)
|
||||
client: Final = sdk.openai(key)
|
||||
native_id: Final = _native_container_id(
|
||||
_container_id(_streamed_response_with_code_interpreter(client, model))
|
||||
)
|
||||
resources.defer(lambda: client.containers.delete(native_id, extra_query=AZURE_PROVIDER_QUERY))
|
||||
_assert_file_round_trip(client, native_id, marker)
|
||||
|
|
|
|||
|
|
@ -11,8 +11,11 @@ litellm-regression-tests/tests/test_inference_endpoints.py.
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import anthropic
|
||||
import pytest
|
||||
from anthropic import Anthropic
|
||||
from anthropic.types import (
|
||||
|
|
@ -30,12 +33,21 @@ from anthropic.types import (
|
|||
ToolParam,
|
||||
ToolUseBlock,
|
||||
)
|
||||
from e2e_config import STREAM_MIN_LEAD_SECONDS, provider_edge_base, provider_paces_stream, unique_marker
|
||||
from e2e_config import (
|
||||
PROVIDER_EDGE_ADVERTISE_HOST,
|
||||
PROVIDER_EDGE_BIND_HOST,
|
||||
STREAM_MIN_LEAD_SECONDS,
|
||||
provider_edge_base,
|
||||
provider_paces_stream,
|
||||
unique_marker,
|
||||
)
|
||||
from e2e_http import assert_client_error
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatMessage, LiteLLMParamsBody, SpendLogRow
|
||||
from models import AnthropicErrorEvent, AnthropicMessagesBody, ChatMessage, LiteLLMParamsBody, SpendLogRow
|
||||
from provider_edge import EDGE_MOUNTS, LiveEdge, RunningEdge, StreamCut, start_provider_edge
|
||||
from provider_edge_bedrock import bedrock_signer
|
||||
from proxy_client import ProxyClient
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError
|
||||
from sdk_clients import NO_PROXY_CACHE, SdkClients, response_header
|
||||
|
||||
pytestmark = [pytest.mark.e2e, pytest.mark.replayable]
|
||||
|
|
@ -385,3 +397,245 @@ class TestOpenAIMessagesToolContinuation:
|
|||
)
|
||||
assert _text(continuation).strip() == receipt, "continuation did not consume the correlated tool result"
|
||||
assert all(not isinstance(block, ToolUseBlock) for block in continuation.content)
|
||||
|
||||
|
||||
BEDROCK_BACKEND: Final = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
|
||||
BEDROCK_EDGE_REGION: Final = "us-east-1"
|
||||
_STREAM_FAILURE_PROMPT: Final = "Count from 1 to 100, one number per line."
|
||||
_FRAME_PAYLOAD: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
_AT_FRAME_BOUNDARY: Final = StreamCut(after_content=True)
|
||||
_MID_FRAME: Final = StreamCut(after_content=True, mid_chunk=True)
|
||||
_BEFORE_FIRST_BYTE: Final = StreamCut(after_content=False)
|
||||
|
||||
type _CutRegistration = Callable[[ProxyClient, ResourceManager, StreamCut], tuple[str, str]]
|
||||
|
||||
|
||||
def _cut_edge(backend: LiveEdge, mount: str) -> RunningEdge:
|
||||
return start_provider_edge(
|
||||
backend,
|
||||
mounts=MappingProxyType({mount: EDGE_MOUNTS[mount]}),
|
||||
bind_host=PROVIDER_EDGE_BIND_HOST,
|
||||
advertise_host=PROVIDER_EDGE_ADVERTISE_HOST,
|
||||
)
|
||||
|
||||
|
||||
def _register_cut_bedrock(proxy: ProxyClient, resources: ResourceManager, cut: StreamCut) -> tuple[str, str]:
|
||||
mount: Final = f"bedrock/{BEDROCK_EDGE_REGION}"
|
||||
edge: Final = _cut_edge(LiveEdge(cut=cut, sign=bedrock_signer(BEDROCK_EDGE_REGION)), mount)
|
||||
resources.defer(edge.shutdown)
|
||||
return _register(
|
||||
proxy,
|
||||
resources,
|
||||
LiteLLMParamsBody(
|
||||
model=BEDROCK_BACKEND,
|
||||
api_base=edge.edge.api_base(mount),
|
||||
aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID",
|
||||
aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY",
|
||||
aws_region_name=BEDROCK_EDGE_REGION,
|
||||
),
|
||||
prefix="e2e-messages-cut",
|
||||
)
|
||||
|
||||
|
||||
def _register_cut_anthropic(proxy: ProxyClient, resources: ResourceManager, cut: StreamCut) -> tuple[str, str]:
|
||||
edge: Final = _cut_edge(LiveEdge(cut=cut), "anthropic")
|
||||
resources.defer(edge.shutdown)
|
||||
return _register(
|
||||
proxy,
|
||||
resources,
|
||||
LiteLLMParamsBody(
|
||||
model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY", api_base=edge.edge.api_base("anthropic")
|
||||
),
|
||||
prefix="e2e-messages-cut",
|
||||
)
|
||||
|
||||
|
||||
_DROPPED_UPSTREAMS: Final[tuple[tuple[str, _CutRegistration, StreamCut], ...]] = (
|
||||
("bedrock_at_a_frame_boundary", _register_cut_bedrock, _AT_FRAME_BOUNDARY),
|
||||
("anthropic_at_a_frame_boundary", _register_cut_anthropic, _AT_FRAME_BOUNDARY),
|
||||
("anthropic_mid_frame", _register_cut_anthropic, _MID_FRAME),
|
||||
)
|
||||
_DROPPED_BEFORE_FIRST_BYTE: Final[tuple[tuple[str, _CutRegistration, StreamCut], ...]] = (
|
||||
("bedrock_before_the_first_byte", _register_cut_bedrock, _BEFORE_FIRST_BYTE),
|
||||
("anthropic_before_the_first_byte", _register_cut_anthropic, _BEFORE_FIRST_BYTE),
|
||||
)
|
||||
|
||||
|
||||
def _payload(frame: str) -> JsonValue | None:
|
||||
try:
|
||||
return _FRAME_PAYLOAD.validate_json(frame)
|
||||
except ValidationError:
|
||||
return None
|
||||
|
||||
|
||||
def _bare_error_frame(frame: str) -> bool:
|
||||
payload: Final = _payload(frame)
|
||||
return isinstance(payload, dict) and "error" in payload and payload.get("type") != "error"
|
||||
|
||||
|
||||
@pytest.mark.provider_edge_host
|
||||
@pytest.mark.provider_live
|
||||
class TestMessagesUpstreamStreamFailure:
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS]
|
||||
)
|
||||
def test_interrupted_upstream_stream_raises_in_the_anthropic_sdk(
|
||||
self,
|
||||
proxy: ProxyClient,
|
||||
resources: ResourceManager,
|
||||
sdk: SdkClients,
|
||||
register: _CutRegistration,
|
||||
cut: StreamCut,
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
client: Final = sdk.anthropic(key)
|
||||
|
||||
stream: Final = client.messages.create(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[_user_turn(_STREAM_FAILURE_PROMPT)],
|
||||
extra_body=NO_PROXY_CACHE,
|
||||
)
|
||||
first: Final = next(stream)
|
||||
assert first.type == "message_start", (
|
||||
f"the stream produced a first event that is not message_start, so this run proves a "
|
||||
f"startup failure, not an interrupted stream: {first!r}"
|
||||
)
|
||||
with pytest.raises(anthropic.APIStatusError) as raised:
|
||||
for _ in stream:
|
||||
pass
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate(raised.value.body)
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f"the SDK raised on the interrupted stream but without the Anthropic error envelope a "
|
||||
f"client reads the failure from: body={raised.value.body!r} message={raised.value}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_event")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"), [case[1:] for case in _DROPPED_UPSTREAMS], ids=[case[0] for case in _DROPPED_UPSTREAMS]
|
||||
)
|
||||
def test_interrupted_upstream_stream_is_an_anthropic_error_event(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
|
||||
outcome: Final = proxy.messages_stream(
|
||||
key,
|
||||
AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[ChatMessage(role="user", content=_STREAM_FAILURE_PROMPT)],
|
||||
),
|
||||
)
|
||||
frames: Final = outcome.stream_events
|
||||
assert outcome.is_streaming, (
|
||||
f"/v1/messages did not answer with an SSE stream: status={outcome.status_code} body={outcome.body}"
|
||||
)
|
||||
assert frames, (
|
||||
f"the proxy sent no SSE data frames although the upstream hung up; stream_error={outcome.stream_error!r}"
|
||||
)
|
||||
assert outcome.stream_error == "event: error", (
|
||||
f"the interrupted stream was not announced by an 'event: error' line Anthropic clients read; "
|
||||
f"stream_error={outcome.stream_error!r} frames={frames}"
|
||||
)
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate_json(frames[-1])
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f'the last SSE frame was not an Anthropic {{"type": "error", "error": ...}} envelope; frames={frames}'
|
||||
)
|
||||
torn: Final = tuple(index for index, frame in enumerate(frames) if _payload(frame) is None)
|
||||
expected_torn: Final = 1 if cut.mid_chunk else 0
|
||||
assert len(torn) == expected_torn, (
|
||||
f"expected {expected_torn} data line(s) that are not JSON, since the edge tears one only when it "
|
||||
f"cuts mid-frame, but the proxy relayed {[frames[index] for index in torn]}; all frames={frames}"
|
||||
)
|
||||
for index in torn:
|
||||
assert _payload(frames[index + 1]) == {"type": "ping"}, (
|
||||
f"the frame the upstream tore was not closed as a ping event before the error, so an "
|
||||
f"Anthropic client parses the error inside it: after {frames[index]!r} came "
|
||||
f"{frames[index + 1]!r}; all frames={frames}"
|
||||
)
|
||||
bare: Final = tuple(frame for frame in frames if _bare_error_frame(frame))
|
||||
assert not bare, (
|
||||
f"the proxy emitted error frames without the Anthropic envelope, which Anthropic clients drop: "
|
||||
f"{bare}; all frames={frames}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"),
|
||||
[case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
)
|
||||
def test_upstream_that_hangs_up_before_the_first_byte_raises_with_its_status_in_the_anthropic_sdk(
|
||||
self,
|
||||
proxy: ProxyClient,
|
||||
resources: ResourceManager,
|
||||
sdk: SdkClients,
|
||||
register: _CutRegistration,
|
||||
cut: StreamCut,
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
client: Final = sdk.anthropic(key)
|
||||
|
||||
with pytest.raises(anthropic.APIStatusError) as raised:
|
||||
client.messages.create(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[_user_turn(_STREAM_FAILURE_PROMPT)],
|
||||
extra_body=NO_PROXY_CACHE,
|
||||
)
|
||||
assert 500 <= raised.value.status_code < 600, (
|
||||
f"an upstream that hung up before sending anything must answer with a server error status the SDK "
|
||||
f"retries on, not {raised.value.status_code}: {raised.value}"
|
||||
)
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate(raised.value.body)
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f"the SDK raised with the right status but without the Anthropic error envelope a client reads "
|
||||
f"the failure from: body={raised.value.body!r} message={raised.value}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.messages.anthropic.upstream_stream_failure.stream.error_status")
|
||||
@pytest.mark.parametrize(
|
||||
("register", "cut"),
|
||||
[case[1:] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
ids=[case[0] for case in _DROPPED_BEFORE_FIRST_BYTE],
|
||||
)
|
||||
def test_upstream_that_hangs_up_before_the_first_byte_is_a_json_error_with_its_status(
|
||||
self, proxy: ProxyClient, resources: ResourceManager, register: _CutRegistration, cut: StreamCut
|
||||
) -> None:
|
||||
model, key = register(proxy, resources, cut)
|
||||
|
||||
outcome: Final = proxy.messages_stream(
|
||||
key,
|
||||
AnthropicMessagesBody(
|
||||
model=model,
|
||||
max_tokens=300,
|
||||
stream=True,
|
||||
messages=[ChatMessage(role="user", content=_STREAM_FAILURE_PROMPT)],
|
||||
),
|
||||
)
|
||||
assert not outcome.is_streaming, (
|
||||
f"nothing had been streamed when the upstream hung up, yet /v1/messages opened a 200 SSE stream "
|
||||
f"instead of answering with the failure's status: stream_error={outcome.stream_error!r} "
|
||||
f"frames={outcome.stream_events}"
|
||||
)
|
||||
assert 500 <= outcome.status_code < 600, (
|
||||
f"/v1/messages answered {outcome.status_code} for an upstream that hung up before its first byte; "
|
||||
f"body={outcome.body}"
|
||||
)
|
||||
try:
|
||||
AnthropicErrorEvent.model_validate_json(outcome.body)
|
||||
except ValidationError:
|
||||
pytest.fail(
|
||||
f'the error body is not an Anthropic {{"type": "error", "error": ...}} envelope; body={outcome.body}'
|
||||
)
|
||||
|
|
|
|||
|
|
@ -598,6 +598,16 @@ class CountTokensResponse(BaseModel):
|
|||
input_tokens: int
|
||||
|
||||
|
||||
class AnthropicErrorBody(BaseModel):
|
||||
type: str
|
||||
message: str
|
||||
|
||||
|
||||
class AnthropicErrorEvent(BaseModel):
|
||||
type: Literal["error"]
|
||||
error: AnthropicErrorBody
|
||||
|
||||
|
||||
# ---------- mcp servers ----------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -45,6 +45,7 @@ import hashlib
|
|||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from collections import deque
|
||||
from collections.abc import Callable, Generator, Mapping, Sequence
|
||||
from contextlib import closing, contextmanager
|
||||
|
|
@ -56,6 +57,7 @@ from types import MappingProxyType
|
|||
from typing import Final, Literal, assert_never
|
||||
from urllib.parse import parse_qsl, urlsplit
|
||||
|
||||
from botocore.eventstream import EventStreamBuffer
|
||||
from e2e_http import (
|
||||
NetworkError,
|
||||
StreamChunk,
|
||||
|
|
@ -96,16 +98,18 @@ from fixture_mode import (
|
|||
)
|
||||
from fixture_profile import IneligibleRequest, MatchProfile, match_profile, strict_identity
|
||||
from provider_cache import (
|
||||
JSON_VALUE,
|
||||
SIGNATURE_HEADERS,
|
||||
CacheEdge,
|
||||
MountPolicy,
|
||||
RequestSigner,
|
||||
invoke_chunk_value,
|
||||
is_bedrock,
|
||||
scoped_edge_base,
|
||||
split_test_segment,
|
||||
)
|
||||
from provider_cache_routing import LIVE_PROVIDER_REQUIRED
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
BEDROCK_REGIONS: Final[tuple[str, ...]] = ("us-east-1",)
|
||||
|
||||
|
|
@ -537,10 +541,33 @@ class ReplayEdge:
|
|||
source: ReplaySource
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class StreamCut:
|
||||
"""Where a live edge hangs up on a streamed upstream body: before its first byte, or with
|
||||
``after_content`` set, right after the first transfer chunk carrying assistant output (a
|
||||
``content_block_delta``). That frame is what commits the proxy's mid-stream fallback
|
||||
wrapper to the client: it holds the lifecycle frames before it back and drops them when
|
||||
the transport fails first, so a cut after a fixed number of chunks landed on either side
|
||||
of that commit depending on how the provider batched its frames. With ``mid_chunk`` set
|
||||
the hang-up comes part way through the next ``data:`` line the provider sends after that,
|
||||
so the client is left inside an SSE frame the way a dropped transport leaves it.
|
||||
|
||||
Whatever was relayed sits on the wire for ``_CUT_SETTLE_SECONDS`` before the hang-up, so
|
||||
the client has read it by then instead of receiving the data and the close in one burst,
|
||||
where its reader can surface the close before what it buffered."""
|
||||
|
||||
after_content: bool
|
||||
mid_chunk: bool = False
|
||||
|
||||
|
||||
_CUT_SETTLE_SECONDS: Final = 1.0
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class LiveEdge:
|
||||
observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None
|
||||
sign: RequestSigner | None = None
|
||||
cut: StreamCut | None = None
|
||||
|
||||
|
||||
type EdgeBackend = RecordEdge | ReplayEdge | LiveEdge | CacheEdge
|
||||
|
|
@ -786,11 +813,128 @@ def _handle_record(
|
|||
assert_never(head)
|
||||
|
||||
|
||||
def _data_line_start(data: bytes) -> int:
|
||||
if data.startswith(b"data:"):
|
||||
return 0
|
||||
at_line_start: Final = data.find(b"\ndata:")
|
||||
return -1 if at_line_start < 0 else at_line_start + 1
|
||||
|
||||
|
||||
def _torn_prefix(data: bytes) -> bytes:
|
||||
start: Final = _data_line_start(data)
|
||||
line_end: Final = data.find(b"\n", start)
|
||||
end: Final = len(data) if line_end < 0 else line_end
|
||||
return data[: start + (end - start) // 2]
|
||||
|
||||
|
||||
class _DataLineTearer:
|
||||
__slots__ = ("_unfinished_line",)
|
||||
|
||||
_unfinished_line: bytes
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._unfinished_line = b""
|
||||
|
||||
def observe(self, data: bytes) -> None:
|
||||
self._unfinished_line = (self._unfinished_line + data).rsplit(b"\n", 1)[-1]
|
||||
|
||||
def tear(self, data: bytes) -> bytes | None:
|
||||
buffered: Final = self._unfinished_line + data
|
||||
if _data_line_start(buffered) < 0:
|
||||
self.observe(data)
|
||||
return None
|
||||
return _torn_prefix(buffered)[len(self._unfinished_line):]
|
||||
|
||||
|
||||
def _is_content_delta(value: JsonValue | None) -> bool:
|
||||
return isinstance(value, dict) and value.get("type") == "content_block_delta"
|
||||
|
||||
|
||||
def _sse_data_carries_content(line: bytes) -> bool:
|
||||
if not line.startswith(b"data:"):
|
||||
return False
|
||||
try:
|
||||
return _is_content_delta(JSON_VALUE.validate_json(line[len(b"data:"):].strip()))
|
||||
except ValidationError:
|
||||
return False
|
||||
|
||||
|
||||
class _AnthropicContentDetector:
|
||||
__slots__ = ("_unfinished_line",)
|
||||
|
||||
_unfinished_line: bytes
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._unfinished_line = b""
|
||||
|
||||
def __call__(self, data: bytes) -> bool:
|
||||
lines: Final = (self._unfinished_line + data).split(b"\n")
|
||||
self._unfinished_line = lines[-1]
|
||||
return any(_sse_data_carries_content(line.rstrip(b"\r")) for line in lines[:-1])
|
||||
|
||||
|
||||
def _invoke_frame_carries_content(payload: bytes) -> bool:
|
||||
try:
|
||||
return _is_content_delta(invoke_chunk_value(JSON_VALUE.validate_json(payload)))
|
||||
except ValidationError:
|
||||
return False
|
||||
|
||||
|
||||
def _bedrock_content_detector() -> Callable[[bytes], bool]:
|
||||
"""Bedrock's invoke stream wraps each Anthropic event in an eventstream frame that a
|
||||
transfer chunk can split, so the frames are reassembled across chunks before being read."""
|
||||
frames: Final = EventStreamBuffer()
|
||||
|
||||
def carries_content(data: bytes) -> bool:
|
||||
frames.add_data(data)
|
||||
return any(_invoke_frame_carries_content(frame.payload) for frame in frames)
|
||||
|
||||
return carries_content
|
||||
|
||||
|
||||
def _content_detector(mount: str) -> Callable[[bytes], bool]:
|
||||
return _bedrock_content_detector() if is_bedrock(mount) else _AnthropicContentDetector()
|
||||
|
||||
|
||||
def _cut_steps(
|
||||
steps: Generator[StreamStep, None, None], cut: StreamCut, carries_content: Callable[[bytes], bool]
|
||||
) -> Generator[StreamStep, None, None]:
|
||||
with closing(steps) as source:
|
||||
tearer: Final = _DataLineTearer()
|
||||
if cut.after_content:
|
||||
for step in source:
|
||||
yield step
|
||||
if isinstance(step, StreamTruncation):
|
||||
return
|
||||
tearer.observe(step.data)
|
||||
if carries_content(step.data):
|
||||
break
|
||||
else:
|
||||
return
|
||||
if cut.mid_chunk:
|
||||
for step in source:
|
||||
if isinstance(step, StreamTruncation):
|
||||
yield step
|
||||
return
|
||||
if (torn := tearer.tear(step.data)) is None:
|
||||
yield step
|
||||
continue
|
||||
if torn:
|
||||
yield StreamChunk(data=torn)
|
||||
break
|
||||
else:
|
||||
return
|
||||
if cut.after_content or cut.mid_chunk:
|
||||
time.sleep(_CUT_SETTLE_SECONDS)
|
||||
yield StreamTruncation(reason=f"edge cut the upstream stream: {cut!r}")
|
||||
|
||||
|
||||
def _handle_live(
|
||||
method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float,
|
||||
cache: CacheEdge | None = None, mount: str = "", test_key: str | None = None,
|
||||
observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None,
|
||||
sign: RequestSigner | None = None,
|
||||
cut: StreamCut | None = None,
|
||||
) -> EdgeOutcome:
|
||||
forwarded: Final = {
|
||||
name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS
|
||||
|
|
@ -805,6 +949,8 @@ def _handle_live(
|
|||
match head:
|
||||
case NetworkError(message=message):
|
||||
return _recorded_outcome(_network_error_response(message))
|
||||
case StreamHead() if cut is not None:
|
||||
return EdgeStream(head.status_code, _filtered_response_headers(head.headers), _cut_steps(head.steps, cut, _content_detector(mount)))
|
||||
case StreamHead() if _is_streamed(head.headers):
|
||||
return EdgeStream(head.status_code, _filtered_response_headers(head.headers), head.steps)
|
||||
case StreamHead():
|
||||
|
|
@ -875,10 +1021,10 @@ def handle_edge_request(
|
|||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
backend, mount, test_key,
|
||||
)
|
||||
case LiveEdge(observe_request=observe_request, sign=sign):
|
||||
case LiveEdge(observe_request=observe_request, sign=sign, cut=cut):
|
||||
return _handle_live(
|
||||
method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout,
|
||||
observe_request=observe_request, sign=sign,
|
||||
mount=mount, observe_request=observe_request, sign=sign, cut=cut,
|
||||
)
|
||||
case RecordEdge():
|
||||
return _handle_record(
|
||||
|
|
|
|||
|
|
@ -54,11 +54,13 @@ from provider_edge import (
|
|||
EdgeBackend,
|
||||
EdgeReply,
|
||||
EdgeStream,
|
||||
LiveEdge,
|
||||
ProviderEdge,
|
||||
ProviderRequestObservation,
|
||||
RecordEdge,
|
||||
ReplayEdge,
|
||||
ReplaySource,
|
||||
StreamCut,
|
||||
edge_request,
|
||||
handle_edge_request,
|
||||
observed_provider_edge,
|
||||
|
|
@ -1000,6 +1002,36 @@ def stream_chunks(response: RecordedStreamedResponse) -> list[bytes]:
|
|||
return [base64.b64decode(chunk) for chunk in response.chunks_b64]
|
||||
|
||||
|
||||
SECOND_DATA_LINE: Final = b'data: {"type":"content_block_delta","delta":{"text":" two"}}'
|
||||
SPLIT_MARKER_CHUNKS: tuple[bytes, ...] = (
|
||||
b'data: {"type":"content_block_delta","delta":{"text":"one"}}\n\nda',
|
||||
b"ta" + SECOND_DATA_LINE[4:] + b"\n\nda",
|
||||
b'ta: {"type":"message_delta","usage":{"output_tokens":7}}\n\nda',
|
||||
b"ta: [DONE]\n\n",
|
||||
)
|
||||
|
||||
|
||||
class TestStreamCut:
|
||||
def test_a_mid_frame_cut_tears_a_data_line_whose_marker_is_split_across_chunks(self) -> None:
|
||||
"""Every ``data:`` marker after the first content delta straddles a transfer
|
||||
chunk boundary, so a tearer that inspects each chunk on its own never finds
|
||||
one and lets the stream finish cleanly instead of cutting it."""
|
||||
backend: Final = LiveEdge(cut=StreamCut(after_content=True, mid_chunk=True))
|
||||
with chunked_provider(chunks=SPLIT_MARKER_CHUNKS) as provider:
|
||||
with running_edge(backend, {"openai": provider_url(provider)}) as edge:
|
||||
head, chunks, ending = raw_stream_post(edge.port, STREAM_PATH, STREAM_BODY)
|
||||
|
||||
assert head.startswith("HTTP/1.1 200 OK")
|
||||
assert ending == "truncated"
|
||||
relayed: Final = b"".join(chunks)
|
||||
whole: Final = b"".join(SPLIT_MARKER_CHUNKS)
|
||||
assert whole.startswith(relayed) and relayed != whole
|
||||
assert relayed.startswith(SPLIT_MARKER_CHUNKS[0])
|
||||
torn_line: Final = relayed.rsplit(b"\n", 1)[-1]
|
||||
assert torn_line and SECOND_DATA_LINE.startswith(torn_line) and torn_line != SECOND_DATA_LINE
|
||||
assert b"[DONE]" not in relayed
|
||||
|
||||
|
||||
class TestStreamingFidelity:
|
||||
"""LIT-5742: a streamed response records and replays as the chunk sequence the
|
||||
provider actually sent, not as one coalesced body. The unit of fidelity is the
|
||||
|
|
|
|||
|
|
@ -84,3 +84,42 @@ def test_jsonrpc_error_and_malformed_tool_result_remain_errors(gateway: Gateway)
|
|||
control: Final = call_tool(gateway, key, identity, names["add"], {"a": 3, "b": 5})
|
||||
assert control.status_code == 200 and control.json()["isError"] is False, control.text
|
||||
assert control.json()["content"][0]["text"] == "8"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("ingress", ("http", "sse"))
|
||||
def test_configured_revision_blocks_unadvertised_handshake_and_keeps_allowed_control(
|
||||
gateway: Gateway, tmp_path, ingress: str
|
||||
) -> None:
|
||||
import asyncio
|
||||
from pathlib import Path
|
||||
|
||||
import yaml
|
||||
from integration._support.mcp import mcp_peer
|
||||
from integration._support.process import owned_proxy
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from mcp import MCPError
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
with mcp_peer() as upstream, gateway.scenario() as scenario:
|
||||
alias: Final = "restricted" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, upstream, alias)
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["general_settings"]["mcp_advertised_versions"] = ["2024-11-05"]
|
||||
config_path: Final = tmp_path / "restricted.yaml"
|
||||
config_path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {"DISABLE_SCHEMA_UPDATE": "true"}, config=config_path) as restricted:
|
||||
endpoint: Final = str(restricted.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
|
||||
headers: Final = {"Authorization": f"Bearer {key}", "x-mcp-servers": identity}
|
||||
denied: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2025-11-25", extra_headers=headers)
|
||||
allowed: Final = MCPClient(server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version="2024-11-05", extra_headers=headers)
|
||||
|
||||
async def exercise() -> None:
|
||||
with pytest.raises(MCPError, match="Unsupported MCP protocol version"):
|
||||
await denied.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in await allowed.list_tools(raise_on_error=True))
|
||||
result: Final = await allowed.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 2, "b": 5}))
|
||||
assert result.is_error is False and result.content[0].text == "7"
|
||||
|
||||
asyncio.run(exercise())
|
||||
|
|
|
|||
|
|
@ -155,3 +155,43 @@ def test_server_initiated_sampling_and_elicitation_surface_as_errors_not_success
|
|||
assert outcome.error is not None, outcome.raw
|
||||
assert outcome.text is None or not outcome.text.startswith(("sampled:", "elicited:")), outcome.raw
|
||||
assert len(tool_calls(peer.drain())) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("downstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("upstream", ("2024-11-05", "2025-03-26", "2025-06-18", "2025-11-25"))
|
||||
@pytest.mark.parametrize("peer_kind", ("http", "sse", "stdio"))
|
||||
@pytest.mark.parametrize("ingress", ("http", "sse"))
|
||||
def test_pinned_revision_pairs_list_and_call_through_gateway(
|
||||
gateway: Gateway, downstream: str, upstream: str, peer_kind: PeerKind, ingress: str
|
||||
) -> None:
|
||||
import asyncio
|
||||
|
||||
from mcp.types import CallToolRequestParams
|
||||
|
||||
from litellm.experimental_mcp_client.client import MCPClient
|
||||
from litellm.types.mcp import MCPTransport
|
||||
|
||||
with peer_of(peer_kind) as peer, gateway.scenario() as scenario:
|
||||
alias: Final = "versions" + uuid.uuid4().hex[:8]
|
||||
identity: Final = register_mcp(scenario, peer, alias, mcp_info={"protocol_version": upstream})
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [identity]})
|
||||
endpoint: Final = str(gateway.client.base_url).rstrip("/") + ("/mcp/sse" if ingress == "sse" else "/mcp")
|
||||
client: Final = MCPClient(
|
||||
server_url=endpoint, transport_type=MCPTransport(ingress), protocol_version=downstream,
|
||||
extra_headers={"Authorization": f"Bearer {key}", "x-mcp-servers": identity}, timeout=15,
|
||||
)
|
||||
|
||||
async def exercise() -> None:
|
||||
tools: Final = await client.list_tools(raise_on_error=True)
|
||||
assert f"{alias}-add" in tuple(tool.name for tool in tools)
|
||||
result: Final = await client.call_tool(CallToolRequestParams(name=f"{alias}-add", arguments={"a": 3, "b": 4}))
|
||||
assert result.is_error is False
|
||||
assert result.content[0].text == "7"
|
||||
|
||||
peer.drain()
|
||||
asyncio.run(exercise())
|
||||
observed: Final = peer.drain()
|
||||
negotiations: Final = tuple(item["body"] for item in observed if item["body"].get("method") == "initialize")
|
||||
assert negotiations, "The operation must reach the upstream negotiation"
|
||||
assert all(request["params"]["protocolVersion"] == upstream for request in negotiations), negotiations
|
||||
assert len(tool_calls(observed)) == 1
|
||||
|
|
|
|||
|
|
@ -381,6 +381,36 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
)
|
||||
raise
|
||||
|
||||
@staticmethod
|
||||
def _config_completing_after_one_delta() -> Mock:
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
completed_response = ResponsesAPIResponse(
|
||||
id="resp_123",
|
||||
created_at=0,
|
||||
status="completed",
|
||||
model="gpt-5.5",
|
||||
object="response",
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=1, output_tokens=1, total_tokens=2),
|
||||
)
|
||||
|
||||
def _transform(model, parsed_chunk, logging_obj):
|
||||
if parsed_chunk.get("type") == "response.completed":
|
||||
return ResponseCompletedEvent(
|
||||
type=ResponsesAPIStreamEvents.RESPONSE_COMPLETED,
|
||||
response=completed_response,
|
||||
)
|
||||
return OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id="msg_123",
|
||||
output_index=0,
|
||||
content_index=0,
|
||||
delta=parsed_chunk["delta"],
|
||||
)
|
||||
|
||||
mock_config.transform_streaming_response.side_effect = _transform
|
||||
return mock_config
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stop_async_iteration_not_logged_as_failure(self):
|
||||
"""
|
||||
|
|
@ -399,6 +429,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
|
||||
async def mock_aiter_bytes():
|
||||
yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n'
|
||||
yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n'
|
||||
|
||||
mock_response.aiter_bytes = mock_aiter_bytes
|
||||
|
||||
|
|
@ -408,11 +439,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
mock_logging_obj.async_failure_handler = Mock()
|
||||
mock_logging_obj.failure_handler = Mock()
|
||||
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
mock_delta_event = Mock()
|
||||
mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
|
||||
mock_delta_event.delta = "test"
|
||||
mock_config.transform_streaming_response.return_value = mock_delta_event
|
||||
mock_config = self._config_completing_after_one_delta()
|
||||
|
||||
# Create the iterator instance
|
||||
iterator = ResponsesAPIStreamingIterator(
|
||||
|
|
@ -432,8 +459,9 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
except StopAsyncIteration:
|
||||
pass # This is expected
|
||||
|
||||
# Verify we got the chunk
|
||||
assert len(chunks_received) == 1
|
||||
# Verify we got the delta and the terminal event
|
||||
assert len(chunks_received) == 2
|
||||
assert iterator.completed_response is not None
|
||||
|
||||
# CRITICAL: Verify that failure handlers were NOT called
|
||||
# StopAsyncIteration is a normal end of stream, not a failure
|
||||
|
|
@ -460,6 +488,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
|
||||
def mock_iter_bytes():
|
||||
yield b'data: {"type": "response.output_text.delta", "delta": "test"}\n\n'
|
||||
yield b'data: {"type": "response.completed", "response": {"id": "resp_123"}}\n\n'
|
||||
|
||||
mock_response.iter_bytes = mock_iter_bytes
|
||||
|
||||
|
|
@ -469,11 +498,7 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
mock_logging_obj.async_failure_handler = Mock()
|
||||
mock_logging_obj.failure_handler = Mock()
|
||||
|
||||
mock_config = Mock(spec=BaseResponsesAPIConfig)
|
||||
mock_delta_event = Mock()
|
||||
mock_delta_event.type = ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA
|
||||
mock_delta_event.delta = "test"
|
||||
mock_config.transform_streaming_response.return_value = mock_delta_event
|
||||
mock_config = self._config_completing_after_one_delta()
|
||||
|
||||
# Create the iterator instance
|
||||
iterator = SyncResponsesAPIStreamingIterator(
|
||||
|
|
@ -493,8 +518,9 @@ class TestBaseResponsesAPIStreamingIterator:
|
|||
except StopIteration:
|
||||
pass # This is expected
|
||||
|
||||
# Verify we got the chunk
|
||||
assert len(chunks_received) == 1
|
||||
# Verify we got the delta and the terminal event
|
||||
assert len(chunks_received) == 2
|
||||
assert iterator.completed_response is not None
|
||||
|
||||
# CRITICAL: Verify that failure handlers were NOT called
|
||||
# StopIteration is a normal end of stream, not a failure
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ class TestBedrockGPTOSS(BaseLLMChatTest):
|
|||
"""Bedrock GPT-OSS intermittently emits truncated toolUse.input deltas on
|
||||
the live endpoint, which makes the inherited live integration test flaky.
|
||||
The accumulation side is covered deterministically by
|
||||
tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index;
|
||||
tests/unit/llms/bedrock/chat/test_invoke_handler.py::test_transform_tool_calls_index;
|
||||
the GPT-OSS-specific request-body transformation is covered by
|
||||
test_function_calling_request_body_gpt_oss below.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -324,7 +324,7 @@ def test_parallel_function_call_anthropic_error_msg(model, messages):
|
|||
Anthropic (and Bedrock Invoke via ``AnthropicConfig.transform_request``)
|
||||
inject a dummy tool so CLIs work with ``modify_params`` left off. Bedrock
|
||||
Converse's no-raise behavior is covered offline in
|
||||
``tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py``
|
||||
``tests/unit/llms/bedrock/chat/test_converse_transformation.py``
|
||||
(see #24158, #27138), which needs no live credentials.
|
||||
"""
|
||||
# Force modify_params off as a clean baseline: it exercises the Anthropic
|
||||
|
|
|
|||
|
|
@ -17,7 +17,7 @@ body can still arrive, released once the caller is done with the response.
|
|||
|
||||
Nothing here re-tests the shapes ``_handler_may_close_client`` covers -- a
|
||||
borrowed ``handler.client``, a caller-supplied client, an evicted-but-held
|
||||
client. Those are pinned in ``tests/test_litellm/llms/custom_httpx/
|
||||
client. Those are pinned in ``tests/unit/llms/custom_httpx/
|
||||
test_http_handler.py``. What is uncovered there is the in-flight response, so no
|
||||
test here may keep the client in a local: that inflates the very refcount under
|
||||
test, and the test then passes on a broken handler. They hold weak references
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Integration tests for SageMaker Nova provider.
|
|||
These tests require a live SageMaker Nova endpoint and AWS credentials.
|
||||
They are skipped by default — run manually with:
|
||||
|
||||
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py -v --no-header -rN
|
||||
pytest tests/local_testing/test_sagemaker_nova_integration.py -v --no-header -rN
|
||||
|
||||
Prerequisites:
|
||||
export AWS_PROFILE=<your-profile> # or set AWS_ACCESS_KEY_ID / AWS_SECRET_ACCESS_KEY
|
||||
|
|
@ -251,7 +251,7 @@ class TestSagemakerNova2LiteIntegration:
|
|||
|
||||
Run with:
|
||||
export SAGEMAKER_NOVA2_LITE_ENDPOINT=<your-nova-2-lite-endpoint>
|
||||
pytest tests/test_litellm/llms/sagemaker/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
|
||||
pytest tests/local_testing/test_sagemaker_nova_integration.py::TestSagemakerNova2LiteIntegration -v
|
||||
"""
|
||||
|
||||
def test_should_accept_reasoning_effort_low(self):
|
||||
|
|
|
|||
|
|
@ -85,7 +85,7 @@ class TestBingGroundingSearch(BaseSearchTest):
|
|||
class TestBingGroundingSearchTransformation:
|
||||
"""
|
||||
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
|
||||
Transformation details are unit-tested in tests/test_litellm/llms/azure/search/.
|
||||
Transformation details are unit-tested in tests/unit/llms/azure/search/.
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
|
|||
|
|
@ -58,7 +58,7 @@ class TestNimbleSearch(BaseSearchTest):
|
|||
class TestNimbleSearchTransformation:
|
||||
"""
|
||||
Full-stack tests through `litellm.search` / `litellm.asearch` with the HTTP layer mocked.
|
||||
Transformation details are unit-tested in tests/test_litellm/llms/nimble/search/.
|
||||
Transformation details are unit-tested in tests/unit/llms/nimble/search/.
|
||||
"""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
# Levo integration tests
|
||||
258
tests/test_litellm/integrations/zerobus/test_zerobus_client.py
Normal file
258
tests/test_litellm/integrations/zerobus/test_zerobus_client.py
Normal file
|
|
@ -0,0 +1,258 @@
|
|||
import base64
|
||||
import json
|
||||
from collections.abc import Iterator, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from itertools import chain, repeat
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.integrations.zerobus.client import ZerobusIngestClient
|
||||
from litellm.types.integrations.zerobus import ZerobusAccessToken, ZerobusConnection, ZerobusIngestFailure
|
||||
|
||||
CONNECTION = ZerobusConnection(
|
||||
workspace_url="https://dbc-a1b2c3d4-e5f6.cloud.databricks.com/",
|
||||
workspace_id="1234567890123456",
|
||||
server_endpoint="https://1234567890123456.zerobus.us-west-2.cloud.databricks.com",
|
||||
client_id="sp-client-id",
|
||||
client_secret="sp-client-secret",
|
||||
table_name="main.litellm.traces",
|
||||
)
|
||||
ROWS = ({"id": "a", "model": "gpt-4o"}, {"id": "b", "model": "gpt-4o"})
|
||||
|
||||
|
||||
def _token(value: str = "tok-1", expires_in: float = 3600) -> httpx.Response:
|
||||
return httpx.Response(200, text=json.dumps({"access_token": value, "expires_in": expires_in}))
|
||||
|
||||
|
||||
def _accepted() -> httpx.Response:
|
||||
return httpx.Response(200, text="{}")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TokenCall:
|
||||
url: str
|
||||
data: Mapping[str, str]
|
||||
headers: Mapping[str, str]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InsertCall:
|
||||
url: str
|
||||
content: bytes
|
||||
headers: Mapping[str, str]
|
||||
|
||||
|
||||
def _results(results: Sequence[httpx.Response | Exception]) -> Iterator[httpx.Response | Exception]:
|
||||
"""Results are served in order, and the last one repeats."""
|
||||
return chain(results[:-1], repeat(results[-1]))
|
||||
|
||||
|
||||
class FakeHTTPClient:
|
||||
"""Stands in for AsyncHTTPHandler, including its habit of raising on error statuses."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
token: Sequence[httpx.Response | Exception] = (),
|
||||
insert: Sequence[httpx.Response | Exception] = (),
|
||||
) -> None:
|
||||
self.token_results = _results(token or (_token(),))
|
||||
self.insert_results = _results(insert or (_accepted(),))
|
||||
self.token_calls: tuple[TokenCall, ...] = ()
|
||||
self.insert_calls: tuple[InsertCall, ...] = ()
|
||||
|
||||
async def post(
|
||||
self,
|
||||
url: str,
|
||||
data: Mapping[str, str] | None = None,
|
||||
content: bytes | None = None,
|
||||
headers: Mapping[str, str] | None = None,
|
||||
) -> httpx.Response:
|
||||
if url.endswith("/oidc/v1/token"):
|
||||
self.token_calls = (*self.token_calls, TokenCall(url, data or {}, headers or {}))
|
||||
return _raise_like_the_handler(next(self.token_results), url)
|
||||
self.insert_calls = (*self.insert_calls, InsertCall(url, content or b"", headers or {}))
|
||||
return _raise_like_the_handler(next(self.insert_results), url)
|
||||
|
||||
|
||||
def _raise_like_the_handler(result: httpx.Response | Exception, url: str) -> httpx.Response:
|
||||
if isinstance(result, Exception):
|
||||
raise result
|
||||
if result.status_code >= 300:
|
||||
raise httpx.HTTPStatusError(
|
||||
"boom",
|
||||
request=httpx.Request("POST", url),
|
||||
response=httpx.Response(result.status_code, text=result.text),
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
class FakeClock:
|
||||
def __init__(self, now: float = 1_000.0) -> None:
|
||||
self.now = now
|
||||
|
||||
def __call__(self) -> float:
|
||||
return self.now
|
||||
|
||||
|
||||
def _client(http_client: FakeHTTPClient, clock: FakeClock | None = None) -> ZerobusIngestClient:
|
||||
return ZerobusIngestClient(connection=CONNECTION, http_client=http_client, clock=clock or FakeClock())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rows_are_posted_as_one_json_list_to_the_table_insert_endpoint():
|
||||
http_client = FakeHTTPClient()
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert outcome is None
|
||||
(call,) = http_client.insert_calls
|
||||
# Insert endpoint per the Zerobus Ingest docs, read 2026-09-19:
|
||||
# https://docs.databricks.com/aws/en/ingestion/lakeflow-connect/zerobus-ingest
|
||||
assert call.url == (
|
||||
"https://1234567890123456.zerobus.us-west-2.cloud.databricks.com/zerobus/v1/tables/main.litellm.traces/insert"
|
||||
)
|
||||
assert json.loads(call.content) == [{"id": "a", "model": "gpt-4o"}, {"id": "b", "model": "gpt-4o"}]
|
||||
assert call.headers["Content-Type"] == "application/json"
|
||||
assert call.headers["Authorization"] == "Bearer tok-1"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_token_is_minted_for_the_zerobus_resource_with_the_table_privileges():
|
||||
"""Zerobus refuses a plain workspace token: it must name its own resource and the table's UC privileges."""
|
||||
http_client = FakeHTTPClient()
|
||||
|
||||
await _client(http_client).insert(ROWS)
|
||||
|
||||
(call,) = http_client.token_calls
|
||||
# Token form per the Zerobus Ingest docs (REST API authentication), read 2026-09-19:
|
||||
# https://docs.databricks.com/aws/en/ingestion/lakeflow-connect/zerobus-ingest
|
||||
assert call.url == "https://dbc-a1b2c3d4-e5f6.cloud.databricks.com/oidc/v1/token"
|
||||
assert call.data["grant_type"] == "client_credentials"
|
||||
assert call.data["scope"] == "all-apis"
|
||||
assert call.data["resource"] == "api://databricks/workspaces/1234567890123456/zerobusDirectWriteApi"
|
||||
details = json.loads(call.data["authorization_details"])
|
||||
assert [(d["object_type"], d["object_full_path"], d["privileges"]) for d in details] == [
|
||||
("CATALOG", "main", ["USE CATALOG"]),
|
||||
("SCHEMA", "main.litellm", ["USE SCHEMA"]),
|
||||
("TABLE", "main.litellm.traces", ["SELECT", "MODIFY"]),
|
||||
]
|
||||
assert all(d["type"] == "unity_catalog_privileges" for d in details)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_service_principal_authenticates_with_http_basic():
|
||||
http_client = FakeHTTPClient()
|
||||
|
||||
await _client(http_client).insert(ROWS)
|
||||
|
||||
scheme, credentials = http_client.token_calls[0].headers["Authorization"].split(" ")
|
||||
assert scheme == "Basic"
|
||||
assert base64.b64decode(credentials).decode() == "sp-client-id:sp-client-secret"
|
||||
|
||||
|
||||
def test_the_client_secret_and_minted_token_stay_out_of_reprs_and_tracebacks():
|
||||
token = ZerobusAccessToken(value="tok-secret", expires_at=1.0)
|
||||
|
||||
assert "sp-client-secret" not in repr(CONNECTION)
|
||||
assert "sp-client-id" in repr(CONNECTION)
|
||||
assert "tok-secret" not in repr(token)
|
||||
assert "expires_at=1.0" in repr(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_token_is_reused_across_inserts_until_it_nears_expiry():
|
||||
clock = FakeClock(now=1_000.0)
|
||||
http_client = FakeHTTPClient(token=[_token("tok-1", expires_in=600), _token("tok-2")])
|
||||
client = _client(http_client, clock)
|
||||
|
||||
await client.insert(ROWS)
|
||||
clock.now = 1_000.0 + 600 - 61
|
||||
await client.insert(ROWS)
|
||||
clock.now = 1_000.0 + 600 - 59
|
||||
await client.insert(ROWS)
|
||||
|
||||
assert len(http_client.token_calls) == 2
|
||||
assert [call.headers["Authorization"] for call in http_client.insert_calls] == [
|
||||
"Bearer tok-1",
|
||||
"Bearer tok-1",
|
||||
"Bearer tok-2",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_401_discards_the_token_so_the_next_insert_mints_a_fresh_one():
|
||||
http_client = FakeHTTPClient(
|
||||
token=[_token("tok-1"), _token("tok-2")],
|
||||
insert=[httpx.Response(401, text="expired"), _accepted()],
|
||||
)
|
||||
client = _client(http_client)
|
||||
|
||||
first = await client.insert(ROWS)
|
||||
second = await client.insert(ROWS)
|
||||
|
||||
assert first == ZerobusIngestFailure(detail="insert returned 401, token discarded", retryable=True)
|
||||
assert second is None
|
||||
assert http_client.insert_calls[1].headers["Authorization"] == "Bearer tok-2"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("status", [429, 500, 503])
|
||||
async def test_a_transient_insert_status_is_retryable(status: int):
|
||||
http_client = FakeHTTPClient(insert=[httpx.Response(status, text="later")])
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert isinstance(outcome, ZerobusIngestFailure)
|
||||
assert outcome.retryable is True
|
||||
assert str(status) in outcome.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_schema_rejection_is_not_retryable_and_says_why():
|
||||
http_client = FakeHTTPClient(insert=[httpx.Response(400, text="unknown column foo")])
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert outcome == ZerobusIngestFailure(detail="insert returned 400: unknown column foo", retryable=False)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_network_failure_on_insert_is_retryable():
|
||||
http_client = FakeHTTPClient(insert=[httpx.ConnectError("connection refused")])
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert isinstance(outcome, ZerobusIngestFailure)
|
||||
assert outcome.retryable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bad_credentials_fail_the_insert_without_posting_rows():
|
||||
http_client = FakeHTTPClient(token=[httpx.Response(401, text="invalid_client")])
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert outcome == ZerobusIngestFailure(detail="token request returned 401: invalid_client", retryable=False)
|
||||
assert http_client.insert_calls == ()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_token_endpoint_outage_is_retryable():
|
||||
http_client = FakeHTTPClient(token=[httpx.Response(503, text="try later")])
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert isinstance(outcome, ZerobusIngestFailure)
|
||||
assert outcome.retryable is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_token_response_without_a_token_is_reported_not_raised():
|
||||
http_client = FakeHTTPClient(token=[httpx.Response(200, text='{"token_type": "Bearer"}')])
|
||||
|
||||
outcome = await _client(http_client).insert(ROWS)
|
||||
|
||||
assert isinstance(outcome, ZerobusIngestFailure)
|
||||
assert outcome.retryable is False
|
||||
assert "token response" in outcome.detail
|
||||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue