Merge remote-tracking branch 'origin/main' into litellm_stream_served_service_tier

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-25 19:45:20 +00:00
commit d07a7c5a01
499 changed files with 5782 additions and 5474 deletions

View file

@ -7,6 +7,8 @@ legacy_flags=(
caching-local
enterprise-package
enterprise-routing
llm-other-providers
llm-vertex-ai
mcp-integration
misc
proxy-db-auth-checks
@ -50,6 +52,8 @@ 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 ;;
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

View file

@ -354,6 +354,21 @@ 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-misc
flag: misc

View file

@ -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",

View file

@ -89,6 +89,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 +98,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

View file

@ -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

View file

@ -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"

View file

@ -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()
},

View file

@ -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

View file

@ -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)?)
}
}

View file

@ -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),
}

View 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()
}

View file

@ -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;

View file

@ -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(())
}
}

View file

@ -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());
}

View file

@ -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());
}

View file

@ -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());
}

View file

@ -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()
}

View file

@ -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,
})

View file

@ -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):

View file

@ -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

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -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.
"""

View file

@ -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

View file

@ -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

View file

@ -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):

View file

@ -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)

View file

@ -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)

View file

@ -13,7 +13,7 @@ def _claude_mapping(messages, response_obj):
def test_claude_mapping_serializes_custom_tool_calls(monkeypatch):
"""
Stub the anthropic module unconditionally: the SDK may be absent (it lives in the
proxy-runtime extra), and the tests/test_litellm/llms/anthropic test package can
proxy-runtime extra), and the tests/unit/llms/anthropic test package can
shadow it on sys.path, so an import probe proves nothing about the real SDK.
"""
stub = types.ModuleType("anthropic")

View file

@ -9,171 +9,6 @@ import os
import pytest
from litellm.llms.cometapi.chat.transformation import (
CometAPIChatCompletionStreamingHandler,
CometAPIConfig,
)
from litellm.llms.cometapi.common_utils import CometAPIException
class TestCometAPIChatCompletionStreamingHandler:
def test_chunk_parser_successful(self):
handler = CometAPIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Test input chunk
chunk = {
"id": "test_id",
"created": 1234567890,
"model": "gpt-3.5-turbo",
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"choices": [
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
],
}
# Parse chunk
result = handler.chunk_parser(chunk)
# Verify response
assert result.id == "test_id"
assert result.object == "chat.completion.chunk"
assert result.created == 1234567890
assert result.model == "gpt-3.5-turbo"
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
assert len(result.choices) == 1
assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning"
def test_chunk_parser_error_response(self):
handler = CometAPIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Test error chunk
error_chunk = {
"error": {
"message": "test error",
"code": 400,
}
}
# Verify error handling
with pytest.raises(CometAPIException) as exc_info:
handler.chunk_parser(error_chunk)
assert "CometAPI Error: test error" in str(exc_info.value)
assert exc_info.value.status_code == 400
def test_chunk_parser_key_error(self):
handler = CometAPIChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
# Test invalid chunk missing required fields
invalid_chunk = {"incomplete": "data"}
# Verify KeyError handling
with pytest.raises(CometAPIException) as exc_info:
handler.chunk_parser(invalid_chunk)
assert "KeyError" in str(exc_info.value)
assert exc_info.value.status_code == 400
class TestCometAPIConfig:
def test_transform_request_basic(self):
"""Test basic request transformation"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-3.5-turbo",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={},
litellm_params={},
headers={},
)
assert transformed_request["model"] == "cometapi/gpt-3.5-turbo"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_transform_request_with_extra_body(self):
"""Test request transformation with extra_body parameters"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-4",
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={"extra_body": {"custom_param": "custom_value"}},
litellm_params={},
headers={},
)
# Validate that extra_body parameters are merged into the request
assert transformed_request["custom_param"] == "custom_value"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_cache_control_flag_removal(self):
"""Test cache control flag removal from messages"""
config = CometAPIConfig()
transformed_request = config.transform_request(
model="cometapi/gpt-3.5-turbo",
messages=[
{
"role": "user",
"content": "Hello, world!",
"cache_control": {"type": "ephemeral"},
}
],
optional_params={},
litellm_params={},
headers={},
)
# CometAPI should remove cache_control flags by default
assert transformed_request["messages"][0].get("cache_control") is None
def test_map_openai_params(self):
"""Test OpenAI parameter mapping"""
config = CometAPIConfig()
non_default_params = {
"temperature": 0.7,
"max_tokens": 100,
"top_p": 0.9,
}
mapped_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params={},
model="cometapi/gpt-3.5-turbo",
drop_params=False,
)
assert mapped_params["temperature"] == 0.7
assert mapped_params["max_tokens"] == 100
assert mapped_params["top_p"] == 0.9
def test_get_error_class(self):
"""Test error class creation"""
config = CometAPIConfig()
error = config.get_error_class(
error_message="Test error",
status_code=400,
headers={"Content-Type": "application/json"},
)
assert isinstance(error, CometAPIException)
assert error.message == "Test error"
assert error.status_code == 400
# Integration test example (requires real API key)

View file

@ -1,79 +0,0 @@
import json
from typing import Final
import httpx
import respx
import litellm
def test_completion_merges_leading_system_and_developer_messages_for_chat_template_models(
respx_mock: respx.MockRouter,
):
upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock(
return_value=httpx.Response(
status_code=200,
json={
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "my-custom-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10},
},
)
)
response: Final = litellm.completion(
model="databricks/my-custom-model",
messages=[
{"role": "system", "content": "You are terse."},
{"role": "developer", "content": "Skills: none."},
{"role": "user", "content": "Hello"},
],
api_base="https://example.databricks.test/serving-endpoints",
api_key="fake-databricks-api-key",
num_retries=0,
)
assert upstream.call_count == 1
request_body: Final = json.loads(upstream.calls[0].request.read())
assert request_body["messages"] == [
{"role": "system", "content": "You are terse.\n\nSkills: none."},
{"role": "user", "content": "Hello"},
]
assert response.choices[0].message.content == "Answer"
def test_completion_merges_system_messages_when_one_has_empty_content(respx_mock: respx.MockRouter):
upstream: Final = respx_mock.post("https://example.databricks.test/serving-endpoints/chat/completions").mock(
return_value=httpx.Response(
status_code=200,
json={
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "my-custom-model",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "Answer"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 9, "completion_tokens": 1, "total_tokens": 10},
},
)
)
litellm.completion(
model="databricks/my-custom-model",
messages=[
{"role": "system", "content": "You are terse."},
{"role": "system", "content": ""},
{"role": "user", "content": "Hello"},
],
api_base="https://example.databricks.test/serving-endpoints",
api_key="fake-databricks-api-key",
num_retries=0,
)
request_body: Final = json.loads(upstream.calls[0].request.read())
assert request_body["messages"] == [
{"role": "system", "content": "You are terse."},
{"role": "user", "content": "Hello"},
]

View file

@ -1,433 +0,0 @@
"""
Integration tests for DeepInfra rerank functionality.
Tests the full rerank flow following the repository patterns.
"""
import asyncio
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
def assert_response_shape(response, custom_llm_provider):
"""Helper function to validate response structure specific to DeepInfra."""
assert hasattr(response, "id")
assert hasattr(response, "results")
assert hasattr(response, "meta")
assert isinstance(response.results, list)
for result in response.results:
assert "index" in result
assert "relevance_score" in result
assert isinstance(result["index"], int)
assert isinstance(result["relevance_score"], (int, float))
# Check meta structure
assert "tokens" in response.meta
assert "billed_units" in response.meta
assert "input_tokens" in response.meta["tokens"]
assert "total_tokens" in response.meta["billed_units"]
@pytest.mark.parametrize("sync_mode", [True, False])
@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post")
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_basic_rerank_deepinfra(mock_sync_post, mock_async_post, sync_mode):
"""Test basic DeepInfra rerank functionality."""
# Mock response data that matches DeepInfra API format
mock_response_data = {
"scores": [0.9, 0.1],
"input_tokens": 25,
"request_id": "deepinfra-request-123",
"inference_status": {
"status": "success",
"runtime_ms": 150,
"cost": 0.0001,
"tokens_generated": 0,
"tokens_input": 25,
},
}
def return_val():
return mock_response_data
api_key = "test_deepinfra_api_key"
api_base = "https://api.deepinfra.com"
if sync_mode:
# Create mock response object for sync
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_sync_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
top_n=2,
custom_llm_provider="deepinfra",
api_key=api_key,
api_base=api_base,
)
mock_sync_post.assert_called_once()
else:
# Create mock response object for async
mock_response = AsyncMock()
def return_val():
return mock_response_data
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_async_post.return_value = mock_response
response = asyncio.run(
litellm.arerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
top_n=2,
custom_llm_provider="deepinfra",
api_key=api_key,
api_base=api_base,
)
)
mock_async_post.assert_called_once()
# Verify response structure
assert response.id == "deepinfra-request-123"
assert response.results is not None
assert len(response.results) == 2
assert response.results[0]["index"] == 0
assert response.results[0]["relevance_score"] == 0.9
assert response.results[1]["index"] == 1
assert response.results[1]["relevance_score"] == 0.1
# Verify metadata
assert response.meta["tokens"]["input_tokens"] == 25
assert response.meta["billed_units"]["total_tokens"] == 25
# Verify hidden params specific to DeepInfra
assert response._hidden_params["status"] == "success"
assert response._hidden_params["runtime_ms"] == 150
assert response._hidden_params["cost"] == 0.0001
# Note: The model name is processed and the 'deepinfra/' prefix is removed
assert response._hidden_params["model"] == "Qwen/Qwen3-Reranker-0.6B"
assert_response_shape(response, custom_llm_provider="deepinfra")
@pytest.mark.parametrize("sync_mode", [True, False])
@patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post")
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_with_queries_param(
mock_sync_post, mock_async_post, sync_mode
):
"""Test DeepInfra rerank with multiple queries parameter."""
mock_response_data = {
"scores": [0.8, 0.6, 0.2],
"input_tokens": 35,
"request_id": "deepinfra-multi-query-123",
"inference_status": {"status": "success", "runtime_ms": 200},
}
def return_val():
return mock_response_data
if sync_mode:
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_sync_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-4B",
query="hello",
documents=["hello", "world", "test"],
queries=["hello", "hi there"], # DeepInfra specific param
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
mock_sync_post.assert_called_once()
# Verify that queries parameter was passed in request
call_data = json.loads(mock_sync_post.call_args.kwargs["data"])
assert "queries" in call_data
assert call_data["queries"] == ["hello", "hi there"]
else:
mock_response = AsyncMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_async_post.return_value = mock_response
response = asyncio.run(
litellm.arerank(
model="deepinfra/Qwen/Qwen3-Reranker-4B",
query="hello",
documents=["hello", "world", "test"],
queries=["hello", "hi there"],
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
)
mock_async_post.assert_called_once()
call_data = json.loads(mock_async_post.call_args.kwargs["data"])
assert "queries" in call_data
assert call_data["queries"] == ["hello", "hi there"]
assert response.results is not None
assert len(response.results) == 3
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_with_service_tier(mock_post):
"""Test DeepInfra rerank with service_tier parameter."""
mock_response_data = {
"scores": [0.95, 0.75],
"input_tokens": 30,
"request_id": "deepinfra-premium-123",
}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-8B",
query="premium search",
documents=["doc1", "doc2"],
service_tier="premium", # DeepInfra specific param
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
mock_post.assert_called_once()
# Verify URL
call_url = mock_post.call_args.kwargs["url"]
assert "api.deepinfra.com/inference/Qwen/Qwen3-Reranker-8B" in call_url
# Verify request contains service_tier
call_data = json.loads(mock_post.call_args.kwargs["data"])
assert call_data["service_tier"] == "premium"
assert response.results is not None
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_with_env_vars(mock_post, monkeypatch):
"""Test DeepInfra rerank with environment variable configuration."""
monkeypatch.setenv("DEEPINFRA_API_KEY", "env_test_key")
monkeypatch.setenv("DEEPINFRA_API_BASE", "https://custom-deepinfra.com")
mock_response_data = {
"scores": [0.88, 0.22],
"input_tokens": 28,
"request_id": "env-test-123",
}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
)
mock_post.assert_called_once()
# Verify headers contain env API key
headers = mock_post.call_args.kwargs.get("headers", {})
assert "Bearer env_test_key" in headers.get("Authorization", "")
assert response.results is not None
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_error_handling(mock_post):
"""Test DeepInfra rerank error handling."""
error_response = {"detail": {"error": "Invalid API key"}}
def return_val():
return error_response
mock_response = MagicMock()
mock_response.status_code = 401
mock_response.json = return_val
mock_response.text = json.dumps(error_response)
mock_response.headers = {"content-type": "application/json"}
mock_post.return_value = mock_response
# The current implementation handles errors gracefully, so we expect a successful response
# with the error information in the hidden params
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
api_key="invalid_key",
api_base="https://api.deepinfra.com",
)
# Verify that the response contains error information
assert (
response._hidden_params["status"] == "unknown"
) # Default status when error occurs
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_defaults_api_base_when_missing(mock_post, monkeypatch):
"""With no api_base anywhere, the call still goes out against DeepInfra's own base."""
monkeypatch.delenv("DEEPINFRA_API_BASE", raising=False)
mock_response = MagicMock()
mock_response.json = lambda: {"scores": [0.9, 0.1], "input_tokens": 20}
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
api_key="test_key",
# api_base is intentionally missing
)
assert "api.deepinfra.com" in mock_post.call_args.kwargs["url"]
assert [result["relevance_score"] for result in response.results] == [0.9, 0.1]
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_request_format(mock_post):
"""Test that the request is properly formatted for DeepInfra API."""
mock_response_data = {"scores": [0.9, 0.1], "input_tokens": 20}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="test query",
documents=["doc1", "doc2"],
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
instruction="custom instruction",
webhook="https://webhook.example.com",
)
mock_post.assert_called_once()
# Verify URL format
call_url = mock_post.call_args.kwargs["url"]
assert call_url == "https://api.deepinfra.com/inference/Qwen/Qwen3-Reranker-0.6B"
# Verify headers
headers = mock_post.call_args.kwargs["headers"]
assert headers["Authorization"] == "Bearer test_key"
assert headers["accept"] == "application/json"
assert headers["content-type"] == "application/json"
# Verify request body format
request_data = json.loads(mock_post.call_args.kwargs["data"])
assert request_data["queries"] == [
"test query",
"test query",
] # DeepInfra requires queries to match documents length
assert request_data["documents"] == ["doc1", "doc2"]
assert request_data["instruction"] == "custom instruction"
assert request_data["webhook"] == "https://webhook.example.com"
assert response.results is not None
def test_deepinfra_rerank_models():
"""Test that DeepInfra Qwen rerank models are recognized."""
# These should not raise errors during model validation
models = [
"deepinfra/Qwen/Qwen3-Reranker-0.6B",
"deepinfra/Qwen/Qwen3-Reranker-4B",
"deepinfra/Qwen/Qwen3-Reranker-8B",
]
for model in models:
resolved_model, provider, _, api_base = litellm.get_llm_provider(model=model)
assert provider == "deepinfra"
assert resolved_model == model.removeprefix("deepinfra/")
assert api_base == "https://api.deepinfra.com/v1/openai"
@patch("litellm.llms.custom_httpx.http_handler.HTTPHandler.post")
def test_deepinfra_rerank_minimal_response(mock_post):
"""Test handling of minimal DeepInfra response."""
# Minimal response with just scores
mock_response_data = {"scores": [0.7, 0.3]}
def return_val():
return mock_response_data
mock_response = MagicMock()
mock_response.json = return_val
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response.text = json.dumps(mock_response_data)
mock_post.return_value = mock_response
response = litellm.rerank(
model="deepinfra/Qwen/Qwen3-Reranker-0.6B",
query="hello",
documents=["hello", "world"],
custom_llm_provider="deepinfra",
api_key="test_key",
api_base="https://api.deepinfra.com",
)
# Should handle minimal response gracefully
assert response.results is not None
assert len(response.results) == 2
assert response.results[0]["relevance_score"] == 0.7
assert response.results[1]["relevance_score"] == 0.3
# Should have default values for missing fields
assert response.meta["tokens"]["input_tokens"] == 0 # Default when missing
assert response._hidden_params["status"] == "unknown" # Default when missing

View file

@ -1 +0,0 @@
"""Tests for Gemini files functionality"""

View file

@ -1 +0,0 @@
# Gemini Video Generation Tests

View file

@ -1 +0,0 @@
# Manus provider tests

View file

@ -1 +0,0 @@
# Manus Responses API tests

View file

@ -1 +0,0 @@
# MiniMax tests

View file

@ -1 +0,0 @@
# MiniMax chat tests

View file

@ -1 +0,0 @@
# MiniMax messages tests

View file

@ -1,19 +1,9 @@
import os
from typing import Dict
from unittest.mock import MagicMock
import httpx
import litellm
import pytest
from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
)
from litellm.llms.mistral.audio_transcription.transformation import (
MistralAudioTranscriptionConfig,
)
from litellm.types.utils import TranscriptionResponse
from litellm.utils import ProviderConfigManager
from tests.llm_translation.base_audio_transcription_unit_tests import (
BaseLLMAudioTranscriptionTest,
)
@ -37,184 +27,3 @@ class TestMistralAudioTranscription(BaseLLMAudioTranscriptionTest):
"Async audio transcription test for Mistral is skipped in this suite; "
"async test plugins (e.g. pytest-asyncio/anyio) are not configured here."
)
def test_mistral_audio_transcription_config_installed():
"""Ensure Mistral audio transcription config is registered with ProviderConfigManager."""
config = ProviderConfigManager.get_provider_audio_transcription_config(
model="mistral/voxtral-mini-latest",
provider=litellm.LlmProviders.MISTRAL,
)
assert config is not None
assert isinstance(config, BaseAudioTranscriptionConfig)
assert isinstance(config, MistralAudioTranscriptionConfig)
def test_mistral_audio_transcription_get_complete_url():
config = MistralAudioTranscriptionConfig()
url = config.get_complete_url(
api_base=None,
api_key="fake-key",
model="voxtral-mini-latest",
optional_params={},
litellm_params={},
)
assert url == "https://api.mistral.ai/v1/audio/transcriptions"
def test_mistral_audio_transcription_get_complete_url_custom_base():
config = MistralAudioTranscriptionConfig()
url = config.get_complete_url(
api_base="https://custom.api.example.com/v1/",
api_key="fake-key",
model="voxtral-mini-latest",
optional_params={},
litellm_params={},
)
assert url == "https://custom.api.example.com/v1/audio/transcriptions"
def test_mistral_audio_transcription_validate_environment():
config = MistralAudioTranscriptionConfig()
headers = config.validate_environment(
headers={},
model="voxtral-mini-latest",
messages=[],
optional_params={},
litellm_params={},
api_key="test-key-123",
)
assert headers["Authorization"] == "Bearer test-key-123"
assert headers["accept"] == "application/json"
def test_mistral_audio_transcription_supported_params():
config = MistralAudioTranscriptionConfig()
params = config.get_supported_openai_params("voxtral-mini-latest")
assert "language" in params
assert "temperature" in params
assert "response_format" in params
assert "timestamp_granularities" in params
def test_mistral_audio_transcription_request_transform():
config = MistralAudioTranscriptionConfig()
wav_path = os.path.join(
os.path.dirname(__file__),
"../../../../..",
"tests",
"llm_translation",
"gettysburg.wav",
)
audio_file = open(wav_path, "rb")
result = config.transform_audio_transcription_request(
model="voxtral-mini-latest",
audio_file=audio_file,
optional_params={"language": "en", "temperature": 0.0},
litellm_params={},
)
audio_file.close()
assert isinstance(result.data, dict)
assert result.data["model"] == "voxtral-mini-latest"
assert result.data["language"] == "en"
assert result.data["temperature"] == 0.0
assert result.files is not None
assert "file" in result.files
def test_mistral_audio_transcription_request_with_diarize():
"""Test that Mistral-specific params like diarize are passed through."""
config = MistralAudioTranscriptionConfig()
wav_path = os.path.join(
os.path.dirname(__file__),
"../../../../..",
"tests",
"llm_translation",
"gettysburg.wav",
)
audio_file = open(wav_path, "rb")
result = config.transform_audio_transcription_request(
model="voxtral-mini-latest",
audio_file=audio_file,
optional_params={"diarize": True},
litellm_params={},
)
audio_file.close()
assert isinstance(result.data, dict)
assert result.data["diarize"] == "true"
def test_mistral_audio_transcription_response_transform():
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {"text": "Four score and seven years ago..."}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == "Four score and seven years ago..."
def test_mistral_audio_transcription_response_transform_diarized():
"""Test that diarized responses preserve segments and language."""
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {
"model": "voxtral-mini-latest",
"text": "Hello, how are you? I am fine.",
"language": None,
"segments": [
{
"text": "Hello, how are you?",
"start": 0.3,
"end": 2.1,
"speaker_id": "speaker_1",
"type": "transcription_segment",
},
{
"text": "I am fine.",
"start": 2.5,
"end": 3.8,
"speaker_id": "speaker_2",
"type": "transcription_segment",
},
],
"usage": {
"prompt_audio_seconds": 4,
"prompt_tokens": 5,
"total_tokens": 50,
"completion_tokens": 20,
},
}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == "Hello, how are you? I am fine."
assert response["segments"] is not None
assert len(response["segments"]) == 2
assert response["segments"][0]["speaker_id"] == "speaker_1"
assert response["segments"][1]["speaker_id"] == "speaker_2"
assert response["language"] is None
def test_mistral_audio_transcription_response_transform_empty():
config = MistralAudioTranscriptionConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.json.return_value = {}
response = config.transform_audio_transcription_response(mock_response)
assert isinstance(response, TranscriptionResponse)
assert response.text == ""

View file

@ -3,321 +3,12 @@ Tests for JSON-based provider configuration system.
"""
import os
import sys
from unittest.mock import patch
try:
import pytest
except ImportError:
# pytest not available, will run as standalone script
pytest = None
# Add workspace to path
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
sys.path.insert(0, workspace_path)
import pytest
import litellm
class TestJSONProviderLoader:
"""Test JSON provider loading and configuration"""
def test_load_json_providers(self):
"""Test that JSON providers load correctly"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
# Verify publicai is loaded
assert JSONProviderRegistry.exists("publicai")
# Get publicai config
publicai = JSONProviderRegistry.get("publicai")
assert publicai is not None
assert publicai.base_url == "https://api.publicai.co/v1"
assert publicai.api_key_env == "PUBLICAI_API_KEY"
assert publicai.api_base_env == "PUBLICAI_API_BASE"
assert publicai.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_dynamic_config_generation(self):
"""Test dynamic config class creation"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Test API info resolution
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://api.publicai.co/v1"
# Test with custom base
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.api.com", "test-key"
)
assert api_base == "https://custom.api.com"
assert api_key == "test-key"
def test_parameter_mapping(self):
"""Test parameter mapping works"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Test parameter mapping
optional_params = {}
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
result = config.map_openai_params(
non_default_params, optional_params, "gpt-4", False
)
# max_completion_tokens should be mapped to max_tokens
assert "max_tokens" in result
assert result["max_tokens"] == 100
assert "max_completion_tokens" not in result
# temperature should be passed through
assert result["temperature"] == 0.7
def test_supported_params(self):
"""Test that config returns supported params"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Get supported params
supported = config.get_supported_openai_params("gpt-4")
# Should have standard OpenAI params
assert isinstance(supported, list)
assert len(supported) > 0
def test_tool_params_excluded_when_function_calling_not_supported(self):
"""Test that tool-related params are excluded for models that don't support
function calling. Regression test for https://github.com/BerriAI/litellm/issues/21125
"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Mock supports_function_calling to return False
with patch("litellm.utils.supports_function_calling", return_value=False):
supported = config.get_supported_openai_params("some-model-without-fc")
tool_params = [
"tools",
"tool_choice",
"function_call",
"functions",
"parallel_tool_calls",
]
for param in tool_params:
assert (
param not in supported
), f"'{param}' should not be in supported params when function calling is not supported"
# Non-tool params should still be present
assert "temperature" in supported
assert "max_tokens" in supported
assert "stop" in supported
def test_tool_params_included_when_function_calling_supported(self):
"""Test that tool-related params are included for models that support function calling."""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("publicai")
config_class = create_config_class(provider)
config = config_class()
# Mock supports_function_calling to return True
with patch("litellm.utils.supports_function_calling", return_value=True):
supported = config.get_supported_openai_params("some-model-with-fc")
assert "tools" in supported
assert "tool_choice" in supported
def test_provider_resolution(self):
"""Test that provider resolution finds JSON providers"""
from litellm.litellm_core_utils.get_llm_provider_logic import (
get_llm_provider,
)
model, provider, api_key, api_base = get_llm_provider(
model="publicai/gpt-4",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "gpt-4"
assert provider == "publicai"
assert api_base == "https://api.publicai.co/v1"
def test_provider_config_manager(self):
"""Test that ProviderConfigManager returns JSON-based configs"""
from litellm import LlmProviders
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model="gpt-4", provider=LlmProviders.PUBLICAI
)
assert config is not None
assert config.custom_llm_provider == "publicai"
class TestPinstripes:
"""Tests for Pinstripes JSON-configured provider"""
def test_pinstripes_json_config_exists(self):
"""Test that pinstripes is configured in providers.json"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
assert JSONProviderRegistry.exists("pinstripes")
pinstripes = JSONProviderRegistry.get("pinstripes")
assert pinstripes is not None
assert pinstripes.base_url == "https://pinstripes.io/v1"
assert pinstripes.api_key_env == "PINSTRIPES_API_KEY"
assert pinstripes.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_pinstripes_provider_resolution(self):
"""Test that provider resolution finds pinstripes and returns the default base URL"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="pinstripes/ps/glm-4.5-air",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "ps/glm-4.5-air"
assert provider == "pinstripes"
assert api_base == "https://pinstripes.io/v1"
def test_pinstripes_dynamic_config(self):
"""Test dynamic config class creation for pinstripes"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("pinstripes")
config_class = create_config_class(provider)
config = config_class()
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://pinstripes.io/v1"
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.pinstripes.io/v1", "test-key"
)
assert api_base == "https://custom.pinstripes.io/v1"
assert api_key == "test-key"
def test_pinstripes_parameter_mapping(self):
"""Test that max_completion_tokens is mapped to max_tokens for pinstripes"""
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("pinstripes")
config_class = create_config_class(provider)
config = config_class()
optional_params = {}
non_default_params = {"max_completion_tokens": 100, "temperature": 0.7}
result = config.map_openai_params(
non_default_params, optional_params, "ps/glm-4.5-air", False
)
assert "max_tokens" in result
assert result["max_tokens"] == 100
assert "max_completion_tokens" not in result
assert result["temperature"] == 0.7
class TestDarkbloom:
def test_darkbloom_json_config_exists(self):
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
darkbloom = JSONProviderRegistry.get("darkbloom")
assert darkbloom is not None
assert darkbloom.base_url == "https://api.darkbloom.dev/v1"
assert darkbloom.api_key_env == "DARKBLOOM_API_KEY"
assert darkbloom.api_base_env == "DARKBLOOM_API_BASE"
assert darkbloom.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_darkbloom_provider_resolution(self):
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="darkbloom/gemma-4-26b",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "gemma-4-26b"
assert provider == "darkbloom"
assert api_key is None
assert api_base == "https://api.darkbloom.dev/v1"
def test_darkbloom_dynamic_config(self):
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("darkbloom")
config_class = create_config_class(provider)
config = config_class()
api_base, api_key = config._get_openai_compatible_provider_info(None, None)
assert api_base == "https://api.darkbloom.dev/v1"
api_base, api_key = config._get_openai_compatible_provider_info(
"https://custom.darkbloom.dev/v1", "test-key"
)
assert api_base == "https://custom.darkbloom.dev/v1"
assert api_key == "test-key"
def test_darkbloom_complete_url_appends_endpoint(self):
from litellm.llms.openai_like.dynamic_config import create_config_class
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
provider = JSONProviderRegistry.get("darkbloom")
config_class = create_config_class(provider)
config = config_class()
url = config.get_complete_url(
api_base="https://api.darkbloom.dev/v1",
api_key="test-key",
model="darkbloom/gemma-4-26b",
optional_params={},
litellm_params={},
stream=True,
)
assert url == "https://api.darkbloom.dev/v1/chat/completions"
def test_darkbloom_provider_config_manager(self):
from litellm import LlmProviders
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model="gemma-4-26b", provider=LlmProviders.DARKBLOOM
)
assert config is not None
assert config.custom_llm_provider == "darkbloom"
class TestPublicAIIntegration:
"""Integration tests for PublicAI provider"""
@ -457,55 +148,3 @@ class TestPublicAIIntegration:
pytest.fail(f"Content list conversion test failed: {str(e)}")
else:
raise
if __name__ == "__main__":
# Run basic tests
print("Testing JSON Provider System...")
test_loader = TestJSONProviderLoader()
print("\n1. Testing JSON provider loading...")
test_loader.test_load_json_providers()
print(" ✓ JSON providers loaded")
print("\n2. Testing dynamic config generation...")
test_loader.test_dynamic_config_generation()
print(" ✓ Dynamic config works")
print("\n3. Testing parameter mapping...")
test_loader.test_parameter_mapping()
print(" ✓ Parameter mapping works")
print("\n4. Testing excluded params...")
test_loader.test_excluded_params()
print(" ✓ Excluded params work")
print("\n5. Testing provider resolution...")
test_loader.test_provider_resolution()
print(" ✓ Provider resolution works")
print("\n6. Testing provider config manager...")
test_loader.test_provider_config_manager()
print(" ✓ Config manager works")
print("\n" + "=" * 50)
print("PublicAI Integration Tests...")
print("=" * 50)
test_integration = TestPublicAIIntegration()
print("\n7. Testing basic completion...")
test_integration.test_publicai_completion_basic()
print("\n8. Testing streaming...")
test_integration.test_publicai_completion_with_streaming()
print("\n9. Testing parameter mapping...")
test_integration.test_publicai_parameter_mapping()
print("\n10. Testing content list conversion...")
test_integration.test_publicai_content_list_conversion()
print("\n" + "=" * 50)
print("✓ All tests passed!")
print("=" * 50)

View file

@ -4,86 +4,12 @@ Related to issue #18794
"""
import os
import sys
from unittest.mock import MagicMock, patch
try:
import pytest
except ImportError:
pytest = None
# Add workspace to path
workspace_path = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))
sys.path.insert(0, workspace_path)
import pytest
import litellm
class TestXiaomiMiMoProviderConfig:
"""Test Xiaomi MiMo provider configuration"""
def test_xiaomi_mimo_in_provider_list(self):
"""Test that xiaomi_mimo is in the provider list (fixes #18794)"""
from litellm import LlmProviders
# Verify xiaomi_mimo is in the enum
assert hasattr(LlmProviders, "XIAOMI_MIMO")
assert LlmProviders.XIAOMI_MIMO.value == "xiaomi_mimo"
# Verify it's in the provider list
assert "xiaomi_mimo" in litellm.provider_list
def test_xiaomi_mimo_json_config_exists(self):
"""Test that xiaomi_mimo is configured in providers.json"""
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
# Verify xiaomi_mimo is loaded
assert JSONProviderRegistry.exists("xiaomi_mimo")
# Get xiaomi_mimo config
xiaomi_mimo = JSONProviderRegistry.get("xiaomi_mimo")
assert xiaomi_mimo is not None
assert xiaomi_mimo.base_url == "https://api.xiaomimimo.com/v1"
assert xiaomi_mimo.api_key_env == "XIAOMI_MIMO_API_KEY"
assert xiaomi_mimo.param_mappings.get("max_completion_tokens") == "max_tokens"
def test_xiaomi_mimo_provider_resolution(self):
"""Test that provider resolution finds xiaomi_mimo"""
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
model, provider, api_key, api_base = get_llm_provider(
model="xiaomi_mimo/mimo-v2-flash",
custom_llm_provider=None,
api_base=None,
api_key=None,
)
assert model == "mimo-v2-flash"
assert provider == "xiaomi_mimo"
assert api_base == "https://api.xiaomimimo.com/v1"
def test_xiaomi_mimo_router_config(self):
"""Test that xiaomi_mimo can be used in Router configuration (fixes #18794)"""
from litellm import Router
# This should not raise "Unsupported provider - xiaomi_mimo"
router = Router(
model_list=[
{
"model_name": "mimo-v2-flash",
"litellm_params": {
"model": "xiaomi_mimo/mimo-v2-flash",
"api_key": "test-key",
},
}
]
)
# Verify the deployment was created successfully
assert len(router.model_list) == 1
assert router.model_list[0]["model_name"] == "mimo-v2-flash"
class TestXiaomiMiMoIntegration:
"""Integration tests for Xiaomi MiMo provider"""
@ -128,30 +54,3 @@ class TestXiaomiMiMoIntegration:
pytest.fail(f"Xiaomi MiMo completion failed: {str(e)}")
else:
raise
if __name__ == "__main__":
# Run basic tests
print("Testing Xiaomi MiMo Provider...")
test_config = TestXiaomiMiMoProviderConfig()
print("\n1. Testing provider in list...")
test_config.test_xiaomi_mimo_in_provider_list()
print(" ✓ xiaomi_mimo in provider list")
print("\n2. Testing JSON config...")
test_config.test_xiaomi_mimo_json_config_exists()
print(" ✓ xiaomi_mimo JSON config loaded")
print("\n3. Testing provider resolution...")
test_config.test_xiaomi_mimo_provider_resolution()
print(" ✓ Provider resolution works")
print("\n4. Testing router configuration...")
test_config.test_xiaomi_mimo_router_config()
print(" ✓ Router configuration works (issue #18794 fixed)")
print("\n" + "=" * 50)
print("✓ All configuration tests passed!")
print("=" * 50)

View file

@ -54,61 +54,3 @@ def test_ovhcloud_audio_transcription_config_installed():
assert config is not None
assert isinstance(config, BaseAudioTranscriptionConfig)
class TestOVHCloudDurationFieldMigration:
"""Tests for OVHCloud duration -> seconds field migration."""
def test_seconds_field_mapped_to_duration(self):
"""New `seconds` field should be normalized to `duration`."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello world",
"seconds": 3.14,
}
result = config.transform_audio_transcription_response(mock_response)
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 3.14
def test_legacy_duration_field_still_works(self):
"""Legacy `duration` field should still be accepted."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {
"text": "Hello world",
"duration": 2.71,
}
result = config.transform_audio_transcription_response(mock_response)
assert result.text == "Hello world"
assert result._hidden_params["duration"] == 2.71
def test_seconds_zero_mapped_to_duration(self):
"""seconds=0.0 must not be treated as falsy and lost."""
from litellm.llms.ovhcloud.audio_transcription.transformation import (
OVHCloudAudioTranscriptionConfig,
)
from unittest.mock import MagicMock
config = OVHCloudAudioTranscriptionConfig()
mock_response = MagicMock()
mock_response.json.return_value = {"text": "silence", "seconds": 0.0}
result = config.transform_audio_transcription_response(mock_response)
assert result._hidden_params["duration"] == 0.0

View file

@ -6,174 +6,12 @@ import os
import pytest
from litellm.llms.ovhcloud.utils import OVHCloudException
from litellm.utils import get_optional_params
from litellm.llms.ovhcloud.chat.transformation import (
OVHCloudChatCompletionStreamingHandler,
OVHCloudChatConfig,
)
config = OVHCloudChatConfig()
model = "ovhcloud/Mistral-7B-Instruct-v0.3"
class TestOvhCloudChatCompletionStreamingHandler:
def test_chunk_parser_successful(self):
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
chunk = {
"id": "test_id",
"created": 1234567890,
"model": "gpt-oss-20b",
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
"choices": [
{"delta": {"content": "test content", "reasoning": "test reasoning"}}
],
}
result = handler.chunk_parser(chunk)
assert result.id == "test_id"
assert result.object == "chat.completion.chunk"
assert result.created == 1234567890
assert result.model == "gpt-oss-20b"
assert result.usage.prompt_tokens == chunk["usage"]["prompt_tokens"]
assert result.usage.completion_tokens == chunk["usage"]["completion_tokens"]
assert result.usage.total_tokens == chunk["usage"]["total_tokens"]
assert len(result.choices) == 1
assert result.choices[0]["delta"]["reasoning_content"] == "test reasoning"
def test_chunk_parser_error_response(self):
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
error_chunk = {
"error": {
"message": "test error",
"code": 400,
}
}
with pytest.raises(OVHCloudException) as exc_info:
handler.chunk_parser(error_chunk)
assert "OVHCloud Error: test error" in str(exc_info.value)
assert exc_info.value.status_code == 400
def test_chunk_parser_key_error(self):
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=None, sync_stream=True
)
invalid_chunk = {"incomplete": "data"}
with pytest.raises(OVHCloudException) as exc_info:
handler.chunk_parser(invalid_chunk)
assert "KeyError" in str(exc_info.value)
assert exc_info.value.status_code == 400
class TestOVHCloudConfig:
def test_transform_request_basic(self):
"""Test basic request transformation"""
transformed_request = config.transform_request(
model,
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={},
litellm_params={},
headers={},
)
assert transformed_request["model"] == model
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_transform_request_with_extra_body(self):
"""Test request transformation with extra_body parameters"""
transformed_request = config.transform_request(
model,
messages=[{"role": "user", "content": "Hello, world!"}],
optional_params={"extra_body": {"custom_param": "custom_value"}},
litellm_params={},
headers={},
)
assert transformed_request["custom_param"] == "custom_value"
assert transformed_request["messages"] == [
{"role": "user", "content": "Hello, world!"}
]
def test_map_openai_params(self):
"""Test OpenAI parameter mapping"""
non_default_params = {
"temperature": 0.7,
"max_tokens": 100,
"top_p": 0.9,
}
mapped_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params={},
model=model,
drop_params=False,
)
assert mapped_params["temperature"] == 0.7
assert mapped_params["max_tokens"] == 100
assert mapped_params["top_p"] == 0.9
def test_get_error_class(self):
"""Test error class creation"""
error = config.get_error_class(
error_message="Test error",
status_code=400,
headers={"Content-Type": "application/json"},
)
assert isinstance(error, OVHCloudException)
assert error.message == "Test error"
assert error.status_code == 400
@pytest.mark.parametrize(
"model",
[
"Meta-Llama-3_3-70B-Instruct",
"Meta-Llama-3_1-70B-Instruct",
"Mixtral-8x7B-Instruct-v0.1",
"gpt-oss-120b",
"some-model-not-in-the-cost-map",
],
)
def test_tools_not_filtered_by_static_model_map(self, model):
"""
OVHCloud AI Endpoints are OpenAI-compatible; tools/tool_choice must pass
through for any model. The server is responsible for rejecting unsupported
tool calls — LiteLLM must not strip them based on a stale static catalog.
"""
params = get_optional_params(
model=model,
custom_llm_provider="ovhcloud",
tools=[
{
"type": "function",
"function": {"name": "x", "parameters": {}},
}
],
tool_choice="auto",
)
assert "tools" in params
assert "tool_choice" in params
def test_ovhcloud_integration():
from litellm import completion
@ -285,78 +123,3 @@ def test_ovhcloud_with_custom_base_url():
if __name__ == "__main__":
pytest.main([__file__, "-v"])
class TestOVHCloudReasoningFieldMigration:
"""Tests for OVHCloud reasoning_content -> reasoning field migration."""
def test_streaming_new_reasoning_field(self):
"""New `reasoning` field should be mapped to `reasoning_content`."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"role": "assistant",
"reasoning": "Let me think...",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "Let me think..."
def test_streaming_legacy_reasoning_content_unchanged(self):
"""Legacy `reasoning_content` field should pass through untouched."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"role": "assistant",
"reasoning_content": "Already correct field.",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "Already correct field."
def test_streaming_both_fields_legacy_wins(self):
"""When both fields present, existing `reasoning_content` is not overwritten."""
handler = OVHCloudChatCompletionStreamingHandler(
streaming_response=iter([]),
sync_stream=True,
)
chunk = {
"id": "test-id",
"created": 1234567890,
"model": "test-model",
"choices": [
{
"delta": {
"reasoning": "new field",
"reasoning_content": "legacy field",
},
"index": 0,
}
],
}
result = handler.chunk_parser(chunk)
assert result.choices[0]["delta"]["reasoning_content"] == "legacy field"

View file

@ -1 +0,0 @@
# S3 Vectors tests

View file

@ -1 +0,0 @@
# S3 Vectors vector store tests

View file

@ -1 +0,0 @@
"""Soniox provider tests."""

View file

@ -1 +0,0 @@
# Vertex AI Image Edit Tests

View file

@ -1,13 +1,9 @@
import os
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import httpx
import pytest
from litellm.llms.vertex_ai.image_generation import (
get_vertex_ai_image_generation_config,
)
from litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation import (
VertexAIGeminiImageGenerationConfig,
)
@ -16,588 +12,6 @@ from litellm.llms.vertex_ai.image_generation.vertex_imagen_transformation import
)
class TestVertexAIGeminiImageGenerationConfig:
def setup_method(self):
"""Set up test fixtures"""
self.config = VertexAIGeminiImageGenerationConfig()
def test_get_supported_openai_params(self):
"""Test get_supported_openai_params returns correct params"""
supported = self.config.get_supported_openai_params("gemini-2.5-flash-image")
assert "n" in supported
assert "size" in supported
def test_map_openai_params_n(self):
"""Test mapping n parameter to candidate_count"""
non_default_params = {"n": 3}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False)
assert result.get("candidate_count") == 3
def test_map_openai_params_size(self):
"""Test mapping size parameter to aspectRatio"""
non_default_params = {"size": "1024x1024"}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False)
assert result.get("aspectRatio") == "1:1"
def test_map_openai_params_size_16_9(self):
"""Test mapping 16:9 size"""
non_default_params = {"size": "1792x1024"}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "gemini-2.5-flash-image", False)
assert result.get("aspectRatio") == "16:9"
def test_map_size_to_aspect_ratio(self):
"""Test size to aspect ratio mapping"""
assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1"
assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9"
assert self.config._map_size_to_aspect_ratio("1024x1792") == "9:16"
assert self.config._map_size_to_aspect_ratio("1280x896") == "4:3"
assert self.config._map_size_to_aspect_ratio("896x1280") == "3:4"
assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default
def test_get_supported_openai_params_includes_native_gemini_params(self):
"""Test that native Gemini imageConfig params are supported"""
supported = self.config.get_supported_openai_params("gemini-3-pro-image-preview")
assert "aspectRatio" in supported
assert "aspect_ratio" in supported
assert "imageSize" in supported
assert "image_size" in supported
assert "imageConfig" in supported
def test_map_openai_params_aspect_ratio_camel_case(self):
"""Test mapping native aspectRatio parameter"""
result = self.config.map_openai_params({"aspectRatio": "9:16"}, {}, "gemini-3-pro-image-preview", False)
assert result["aspectRatio"] == "9:16"
def test_map_openai_params_aspect_ratio_snake_case(self):
"""Test mapping native aspect_ratio parameter"""
result = self.config.map_openai_params({"aspect_ratio": "16:9"}, {}, "gemini-3-pro-image-preview", False)
assert result["aspectRatio"] == "16:9"
def test_map_openai_params_image_size_camel_case(self):
"""Test mapping native imageSize parameter"""
result = self.config.map_openai_params({"imageSize": "4K"}, {}, "gemini-3-pro-image-preview", False)
assert result["imageSize"] == "4K"
def test_map_openai_params_image_size_snake_case(self):
"""Test mapping native image_size parameter"""
result = self.config.map_openai_params({"image_size": "2K"}, {}, "gemini-3-pro-image-preview", False)
assert result["imageSize"] == "2K"
def test_map_openai_params_image_config_dict_stored_whole(self):
"""imageConfig dict is stored as-is so all fields survive"""
result = self.config.map_openai_params(
{"imageConfig": {"aspectRatio": "16:9", "imageSize": "2K"}},
{},
"gemini-3.1-flash-image",
False,
)
assert result["imageConfig"] == {"aspectRatio": "16:9", "imageSize": "2K"}
def test_map_openai_params_image_config_all_fields(self):
"""All ImageConfig fields (personGeneration, imageOutputOptions) pass through"""
payload = {
"imageConfig": {
"aspectRatio": "9:16",
"imageSize": "4K",
"personGeneration": "DONT_ALLOW",
"imageOutputOptions": {
"mimeType": "image/jpeg",
"compressionQuality": 80,
},
}
}
result = self.config.map_openai_params(payload, {}, "gemini-3.1-flash-image", False)
assert result["imageConfig"] == payload["imageConfig"]
def test_map_openai_params_image_config_non_dict_warns_and_drops(self):
"""Non-dict imageConfig is dropped with a warning, not silently discarded"""
with patch("litellm.llms.vertex_ai.image_generation.vertex_gemini_transformation.verbose_logger") as mock_log:
result = self.config.map_openai_params(
{"imageConfig": "bad-string-value"}, {}, "gemini-3.1-flash-image", False
)
assert "imageConfig" not in result
mock_log.warning.assert_called_once()
def test_transform_image_generation_request_from_image_config(self):
"""Full imageConfig dict is forwarded verbatim into generationConfig"""
full_config = {
"aspectRatio": "16:9",
"imageSize": "2K",
"personGeneration": "DONT_ALLOW",
"imageOutputOptions": {"mimeType": "image/jpeg", "compressionQuality": 85},
}
mapped = self.config.map_openai_params(
{"imageConfig": full_config},
{},
"gemini-3.1-flash-image",
False,
)
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image",
prompt="A nano banana on a desk",
optional_params=mapped,
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"] == full_config
def test_transform_image_generation_flat_params_override_image_config(self):
"""Explicit flat params win over the same key inside imageConfig"""
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image",
prompt="A nano banana",
optional_params={
"imageConfig": {"aspectRatio": "1:1", "personGeneration": "DONT_ALLOW"},
"aspectRatio": "16:9", # should win
},
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9"
assert request["generationConfig"]["imageConfig"]["personGeneration"] == "DONT_ALLOW"
def test_transform_image_generation_request_basic(self):
"""Test basic request transformation"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={},
litellm_params={},
headers={},
)
assert "contents" in request
assert "generationConfig" in request
assert request["generationConfig"]["responseModalities"] == ["IMAGE"]
assert request["contents"][0]["parts"][0]["text"] == "A nano banana"
def test_transform_image_generation_request_with_aspect_ratio(self):
"""Test request transformation with aspectRatio"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={"aspectRatio": "16:9"},
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"]["aspectRatio"] == "16:9"
def test_transform_image_generation_request_with_image_size(self):
"""Test request transformation with imageSize (Gemini 3 Pro)"""
request = self.config.transform_image_generation_request(
model="gemini-3-pro-image-preview",
prompt="A nano banana",
optional_params={"imageSize": "4K"},
litellm_params={},
headers={},
)
assert request["generationConfig"]["imageConfig"]["imageSize"] == "4K"
def test_map_openai_params_web_search_options(self):
"""Test web_search_options maps to googleSearch tool"""
result = self.config.map_openai_params({"web_search_options": {}}, {}, "gemini-3.1-flash-image-preview", False)
assert result["tools"] == [{"googleSearch": {}}]
def test_transform_image_generation_request_with_web_search_tools(self):
"""Test request transformation includes googleSearch tools"""
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image-preview",
prompt="Generate an image of the latest iPhone",
optional_params={"tools": [{"googleSearch": {}}]},
litellm_params={},
headers={},
)
assert request["tools"] == [{"googleSearch": {}}]
def test_transform_image_generation_request_forwards_tool_config(self):
"""Test request transformation forwards toolConfig side-effects from tool mapping"""
mapped = self.config.map_openai_params(
{"tools": [{"googleMaps": {"latitude": 37.7, "longitude": -122.4}}]},
{},
"gemini-3.1-flash-image-preview",
False,
)
request = self.config.transform_image_generation_request(
model="gemini-3.1-flash-image-preview",
prompt="Generate an image of a coffee shop nearby",
optional_params=mapped,
litellm_params={},
headers={},
)
assert request["tools"] == [{"googleMaps": {}}]
assert request["toolConfig"] == {"retrievalConfig": {"latLng": {"latitude": 37.7, "longitude": -122.4}}}
def test_transform_image_generation_request_with_candidate_count(self):
"""Test request transformation with candidate_count"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={"candidate_count": 2},
litellm_params={},
headers={},
)
assert request["generationConfig"]["candidateCount"] == 2
def test_transform_image_generation_request_with_n(self):
"""Test request transformation with n parameter"""
request = self.config.transform_image_generation_request(
model="gemini-2.5-flash-image",
prompt="A nano banana",
optional_params={"n": 2},
litellm_params={},
headers={},
)
assert request["generationConfig"]["candidateCount"] == 2
def test_transform_image_generation_response(self):
"""Test response transformation"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "base64_encoded_image_data",
}
}
]
}
}
],
"usageMetadata": {
"promptTokenCount": 93,
"promptTokensDetails": [
{
"modality": "TEXT",
"tokenCount": 54,
},
{
"modality": "IMAGE",
"tokenCount": 39,
},
],
"candidatesTokenCount": 17,
"totalTokenCount": 110,
},
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="gemini-2.5-flash-image",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 1
assert result.data[0].b64_json == "base64_encoded_image_data"
assert result.data[0].url is None
assert result.usage.input_tokens == 93
assert result.usage.input_tokens_details.text_tokens == 54
assert result.usage.input_tokens_details.image_tokens == 39
assert result.usage.output_tokens == 17
assert result.usage.total_tokens == 110
def test_transform_image_generation_response_multiple_images(self):
"""Test response transformation with multiple images"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "image1",
}
},
{
"inlineData": {
"mimeType": "image/png",
"data": "image2",
}
},
]
}
}
]
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="gemini-2.5-flash-image",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 2
assert result.data[0].b64_json == "image1"
assert result.data[1].b64_json == "image2"
def test_transform_image_generation_response_signature(self):
"""Test response transformation includes thoughtSignature for Gemini 3 Pro"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "base64_encoded_image_data",
},
"thoughtSignature": "test_signature_abc123",
}
]
}
}
]
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="gemini-3-pro-image-preview",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 1
assert result.data[0].b64_json == "base64_encoded_image_data"
assert result.data[0].provider_specific_fields["thought_signature"] == "test_signature_abc123"
def test_transform_image_generation_response_tracks_web_search_requests(self):
"""Grounding queries are carried onto usage so search spend can be billed"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"candidates": [
{
"content": {
"parts": [
{
"inlineData": {
"mimeType": "image/png",
"data": "base64_encoded_image_data",
}
}
]
},
"groundingMetadata": {"webSearchQueries": ["eiffel tower", "paris skyline"]},
}
],
"usageMetadata": {
"promptTokenCount": 93,
"candidatesTokenCount": 17,
"totalTokenCount": 110,
},
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
result = self.config.transform_image_generation_response(
model="gemini-2.5-flash-image",
raw_response=mock_response,
model_response=ImageResponse(),
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert result.usage.web_search_requests == 2
class TestVertexAIImagenImageGenerationConfig:
def setup_method(self):
"""Set up test fixtures"""
self.config = VertexAIImagenImageGenerationConfig()
def test_get_supported_openai_params(self):
"""Test get_supported_openai_params returns correct params"""
supported = self.config.get_supported_openai_params("imagegeneration@006")
assert "n" in supported
assert "size" in supported
def test_map_openai_params_n(self):
"""Test mapping n parameter to sampleCount"""
non_default_params = {"n": 3}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False)
assert result.get("sampleCount") == 3
def test_map_openai_params_size(self):
"""Test mapping size parameter to aspectRatio"""
non_default_params = {"size": "1024x1024"}
optional_params = {}
result = self.config.map_openai_params(non_default_params, optional_params, "imagegeneration@006", False)
assert result.get("aspectRatio") == "1:1"
def test_map_size_to_aspect_ratio(self):
"""Test size to aspect ratio mapping"""
assert self.config._map_size_to_aspect_ratio("1024x1024") == "1:1"
assert self.config._map_size_to_aspect_ratio("1792x1024") == "16:9"
assert self.config._map_size_to_aspect_ratio("unknown") == "1:1" # default
def test_transform_image_generation_request_basic(self):
"""Test basic request transformation"""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={},
litellm_params={},
headers={},
)
assert "instances" in request
assert "parameters" in request
assert request["instances"][0]["prompt"] == "A cat"
assert request["parameters"]["sampleCount"] == 1
def test_transform_image_generation_request_with_params(self):
"""Test request transformation with parameters"""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={"sampleCount": 2, "aspectRatio": "16:9"},
litellm_params={},
headers={},
)
assert request["parameters"]["sampleCount"] == 2
assert request["parameters"]["aspectRatio"] == "16:9"
def test_transform_image_generation_request_labels_from_metadata(self):
"""Billing labels from litellm_params.metadata.requester_metadata on predict body."""
request = self.config.transform_image_generation_request(
model="imagegeneration@006",
prompt="A cat",
optional_params={},
litellm_params={"metadata": {"requester_metadata": {"team": "platform", "env": "prod"}}},
headers={},
)
assert request["labels"] == {"team": "platform", "env": "prod"}
assert "labels" not in request["parameters"]
def test_transform_image_generation_response(self):
"""Test response transformation"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {"predictions": [{"bytesBase64Encoded": "base64_encoded_image_data"}]}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="imagegeneration@006",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 1
assert result.data[0].b64_json == "base64_encoded_image_data"
assert result.data[0].url is None
def test_transform_image_generation_response_multiple_images(self):
"""Test response transformation with multiple images"""
mock_response = MagicMock(spec=httpx.Response)
mock_response.status_code = 200
mock_response.json.return_value = {
"predictions": [
{"bytesBase64Encoded": "image1"},
{"bytesBase64Encoded": "image2"},
]
}
mock_response.headers = {}
from litellm.types.utils import ImageResponse
model_response = ImageResponse()
result = self.config.transform_image_generation_response(
model="imagegeneration@006",
raw_response=mock_response,
model_response=model_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
)
assert len(result.data) == 2
assert result.data[0].b64_json == "image1"
assert result.data[1].b64_json == "image2"
class TestGetVertexAIImageGenerationConfig:
"""Test the router function that selects the correct config"""
def test_get_gemini_model_config(self):
"""Test that Gemini models return Gemini config"""
config = get_vertex_ai_image_generation_config("gemini-2.5-flash-image")
assert isinstance(config, VertexAIGeminiImageGenerationConfig)
config = get_vertex_ai_image_generation_config("gemini-3-pro-image-preview")
assert isinstance(config, VertexAIGeminiImageGenerationConfig)
config = get_vertex_ai_image_generation_config("vertex_ai/gemini-2.5-flash-image")
assert isinstance(config, VertexAIGeminiImageGenerationConfig)
def test_get_imagen_model_config(self):
"""Test that Imagen models return Imagen config"""
config = get_vertex_ai_image_generation_config("imagegeneration@006")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
config = get_vertex_ai_image_generation_config("imagen-4.0-generate-001")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
config = get_vertex_ai_image_generation_config("vertex_ai/imagegeneration@006")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
def test_get_non_gemini_model_config(self):
"""Test that non-Gemini models default to Imagen config"""
config = get_vertex_ai_image_generation_config("some-other-model")
assert isinstance(config, VertexAIImagenImageGenerationConfig)
class TestVertexAIImageGenerationIntegration:
"""Integration tests for Vertex AI image generation"""
@ -642,39 +56,3 @@ class TestVertexAIImageGenerationIntegration:
litellm_params={},
)
assert "Authorization" in headers
def test_gemini_get_complete_url(self):
"""Test Gemini config URL generation"""
config = VertexAIGeminiImageGenerationConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="gemini-2.5-flash-image",
optional_params={},
litellm_params={
"vertex_project": "test-project",
"vertex_location": "us-central1",
},
)
assert "test-project" in url
assert "us-central1" in url
assert "gemini-2.5-flash-image" in url
assert "generateContent" in url
def test_imagen_get_complete_url(self):
"""Test Imagen config URL generation"""
config = VertexAIImagenImageGenerationConfig()
url = config.get_complete_url(
api_base=None,
api_key=None,
model="imagegeneration@006",
optional_params={},
litellm_params={
"vertex_project": "test-project",
"vertex_location": "us-central1",
},
)
assert "test-project" in url
assert "us-central1" in url
assert "imagegeneration@006" in url
assert "predict" in url

View file

@ -1 +0,0 @@
"""Tests for Vertex AI Gemma-AI models"""

View file

@ -1,3 +0,0 @@
"""
Tests for Vertex AI video generation.
"""

View file

@ -30,7 +30,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockTextContent,
)
from litellm.types.utils import CallTypes, ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
@pytest.mark.asyncio

View file

@ -23,7 +23,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockGuardrailResponse,
)
from litellm.types.utils import Choices, Message, ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
CONTENT_FILTER_CHECKS = {"contentFilter": {"categories": [{"category": "VIOLENCE"}]}}

View file

@ -8029,6 +8029,37 @@ class TestPKCEStateCookieBinding:
assert result is not None
@pytest.mark.asyncio
@pytest.mark.parametrize("enable_sso_debug_value", [None, "false", "0"])
async def test_sso_debug_routes_return_404_unless_explicitly_enabled(enable_sso_debug_value):
"""
/sso/debug/login and /sso/debug/callback must 404 unless ENABLE_SSO_DEBUG is
explicitly set to a truthy value.
"""
from litellm.proxy.management_endpoints.ui_sso import debug_sso_callback, debug_sso_login
mock_request = MagicMock(spec=Request)
mock_request.base_url = "http://proxy.example.com/"
mock_request.cookies = {}
mock_request.query_params = {}
env = {"GENERIC_CLIENT_ID": "test_client_id"}
if enable_sso_debug_value is not None:
env["ENABLE_SSO_DEBUG"] = enable_sso_debug_value
with patch.dict(os.environ, env, clear=False):
if enable_sso_debug_value is None:
os.environ.pop("ENABLE_SSO_DEBUG", None)
with pytest.raises(HTTPException) as login_exc:
await debug_sso_login(mock_request)
with pytest.raises(HTTPException) as callback_exc:
await debug_sso_callback(mock_request)
assert login_exc.value.status_code == 404
assert callback_exc.value.status_code == 404
@pytest.mark.asyncio
async def test_debug_sso_callback_renders_full_jwt_claims():
"""
@ -8080,7 +8111,7 @@ async def test_debug_sso_callback_renders_full_jwt_claims():
with (
patch.dict(
os.environ,
{"GENERIC_CLIENT_ID": "test_client_id"},
{"GENERIC_CLIENT_ID": "test_client_id", "ENABLE_SSO_DEBUG": "true"},
clear=False,
),
patch(
@ -8165,7 +8196,7 @@ async def test_debug_sso_callback_handles_missing_raw_response():
with (
patch.dict(
os.environ,
{"MICROSOFT_CLIENT_ID": "test_microsoft_id"},
{"MICROSOFT_CLIENT_ID": "test_microsoft_id", "ENABLE_SSO_DEBUG": "true"},
clear=False,
),
patch.object(
@ -8213,7 +8244,7 @@ async def _render_debug_page(provider_env, id_jag_registered, force_inert=False)
return parsed
stack = [
patch.dict(os.environ, provider_env, clear=False),
patch.dict(os.environ, {**provider_env, "ENABLE_SSO_DEBUG": "true"}, clear=False),
patch( # test-quality-ok: endpoint test stubs the upstream generic IdP boundary
"litellm.proxy.management_endpoints.ui_sso.get_generic_sso_response", side_effect=fake_generic
),

View file

@ -24,7 +24,7 @@ from starlette.datastructures import FormData
import litellm
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,

View file

@ -1,11 +1,14 @@
import asyncio
import base64
import importlib
import os
from collections.abc import Coroutine, Iterator
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
import boto3
import httpx
import pytest
from pytest_socket import enable_socket, socket_allow_hosts
@ -15,9 +18,12 @@ import litellm # noqa: E402 # litellm reads LITELLM_LOCAL_MODEL_COST_MAP at im
import litellm.router as litellm_router_module # noqa: E402 # same import-time dependency
import litellm.utils as litellm_utils_module # noqa: E402 # same import-time dependency
from litellm._logging import ALL_LOGGERS # noqa: E402 # same import-time dependency
from litellm.anthropic_beta_headers_manager import reload_beta_headers_config # noqa: E402 # same import-time dependency
from litellm.litellm_core_utils.prompt_templates import factory as prompt_factory_module # noqa: E402 # same import-time dependency
from litellm.litellm_core_utils.prompt_templates import ( # noqa: E402 # same import-time dependency
image_handling as image_handling_module,
)
from litellm.llms.gemini.chat import transformation as gemini_chat_transformation_module # noqa: E402 # same import-time dependency
from litellm.llms.custom_httpx.async_client_cleanup import ( # noqa: E402 # same import-time dependency
close_litellm_async_clients,
)
@ -89,6 +95,9 @@ RESTORED_GLOBALS: Final = (
)
MODULE_LEVEL_CLIENTS: Final = ("module_level_client", "module_level_aclient")
SESSION_CLIENTS: Final = ("base_llm_aiohttp_handler", "httpx_client", "aclient", "client")
ONE_PIXEL_PNG: Final = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
)
def _allow_loopback_only() -> None:
@ -236,6 +245,47 @@ def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
litellm.get_model_info.cache_clear()
@pytest.fixture
def local_beta_headers_config(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]:
monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True")
reload_beta_headers_config()
yield
monkeypatch.delenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", raising=False)
reload_beta_headers_config()
@dataclass(slots=True)
class AsyncOnlyImageFetch:
fetched: list[str] = field(default_factory=list) # mutable-ok: tests assert on the URLs fetched, in order
base64_png: str = base64.b64encode(ONE_PIXEL_PNG).decode()
data_url: str = "data:image/png;base64," + base64.b64encode(ONE_PIXEL_PNG).decode()
@pytest.fixture
def async_only_image_fetch(monkeypatch: pytest.MonkeyPatch) -> AsyncOnlyImageFetch:
fetch: Final = AsyncOnlyImageFetch()
def forbid_sync_fetch(client: object, url: str, **kwargs: object) -> httpx.Response:
raise litellm.ImageFetchError(f"sync image fetch ran on the event loop: {url}")
async def serve_png(client: object, url: str, **kwargs: object) -> httpx.Response:
fetch.fetched.append(url)
return httpx.Response(
200, content=ONE_PIXEL_PNG, headers={"content-type": "image/png"}, request=httpx.Request("GET", url)
)
def forbid_sync_convert(url: str, *args: object, **kwargs: object) -> str:
if url.startswith(("http://", "https://")):
raise litellm.ImageFetchError(f"sync convert_url_to_base64 ran on the request path: {url}")
return url
monkeypatch.setattr(image_handling_module, "safe_get", forbid_sync_fetch)
monkeypatch.setattr(image_handling_module, "async_safe_get", serve_png)
for module in (image_handling_module, prompt_factory_module, gemini_chat_transformation_module):
monkeypatch.setattr(module, "convert_url_to_base64", forbid_sync_convert)
return fetch
@pytest.fixture
def no_ambient_azure_credentials(monkeypatch: pytest.MonkeyPatch) -> None:
for name in AMBIENT_AZURE_CREDENTIAL_ENV_VARS:

View file

@ -616,7 +616,7 @@ def test_transform_response_reraises_unexpected_error(config):
# automatically. See base_batches_config_test.py.
# --------------------------------------------------------------------------- #
from tests.test_litellm.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402
from tests.unit.llms.base_llm.batches.base_batches_config_test import ( # noqa: E402
BatchesConfigContractTests,
)

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