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

# Conflicts:
#	tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py
This commit is contained in:
mateo-berri 2026-09-24 14:24:05 -07:00
commit 95e3caefb0
126 changed files with 12440 additions and 6850 deletions

View file

@ -108,6 +108,7 @@ jobs:
tests/test_litellm/completion_extras
tests/test_litellm/containers
tests/test_litellm/endpoints
tests/test_litellm/files
tests/test_litellm/images
tests/test_litellm/interactions
tests/test_litellm/messages

10
litellm-rust/AGENTS.md Normal file
View file

@ -0,0 +1,10 @@
# Rust workspace rules
## Test placement
- Never create a `tests.rs` (or `test.rs`) file under `src/`, and never `#[path = "tests.rs"] mod tests;`
- A test that reaches private items lives inline, in a `#[cfg(test)] mod tests { ... }` at the bottom of the file that owns those items
- A test that only uses the crate's public API lives in `crates/<crate>/tests/<subject>.rs`, next to `src/`
- Split a mixed test file along that line instead of widening visibility to move it
- A test for another crate's item belongs in that crate, not in a downstream one
- Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests/<name>/mod.rs` or `tests/<subject>/support.rs` so it is not picked up as a test crate of its own

View file

@ -4,7 +4,6 @@ version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
autotests = false
[dependencies]
litellm-host.workspace = true

File diff suppressed because it is too large Load diff

View file

@ -63,5 +63,151 @@ impl PendingLogging {
}
#[cfg(test)]
#[path = "../tests/deferred.rs"]
mod tests;
mod tests {
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::{PendingLogging, PendingSuccess};
use crate::PythonLogger;
use crate::test_support::{local, namespace, run};
/// A deferred success for the namespace's `logger` and `response`, bound as `pending`.
fn defer<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = namespace(py, c"response = object()");
run(py, &locals, script);
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: PythonLogger::new(local(&locals, "logger").unbind()),
response: Some(local(&locals, "response").unbind()),
start: py.None(),
end: Some(py.None()),
}),
},
)
.unwrap();
locals.set_item("pending", pending).unwrap();
locals
}
#[test]
fn release_enqueues_the_success_once_in_the_releasing_context() {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
from contextvars import ContextVar
marker = ContextVar('marker', default='unset')
observed = []
def on_enqueue(coroutine):
observed.append(marker.get())
pending.release(True)
logger.on_enqueue = on_enqueue
",
);
run(
py,
&locals,
c"
marker.set('release')
pending.release(True)
pending.release(True)
assert observed == ['release'], observed
assert logger.names() == ['async_success_handler', 'enqueued'], logger.calls
assert logger.calls[0][1] is response
",
);
});
}
#[test]
fn a_blocked_release_drops_the_success_for_good() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
pending.release(False)
pending.release(True)
assert logger.calls == [], logger.calls
",
);
});
}
#[rstest]
#[case::ordinary_error(c"RuntimeError('queue full')", false)]
#[case::cancellation(c"asyncio.CancelledError()", true)]
fn a_failed_enqueue_closes_the_coroutine_and_is_never_replayed(
#[case] failure: &CStr,
#[case] propagates: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
import asyncio
def on_enqueue(coroutine):
raise failure
logger.on_enqueue = on_enqueue
",
);
locals
.set_item("failure", py.eval(failure, None, Some(&locals)).unwrap())
.unwrap();
let released = local(&locals, "pending").call_method1("release", (true,));
match released {
Ok(_) => assert!(!propagates),
Err(error) => {
assert!(propagates);
assert!(error.value(py).is(local(&locals, "failure")));
}
}
locals.set_item("propagates", propagates).unwrap();
run(
py,
&locals,
c"
pending.release(True)
assert logger.names() == ['async_success_handler', 'enqueued', 'closed'], logger.calls
assert unraisable_from(logger) == ([] if propagates else [failure])
",
);
});
}
#[test]
fn an_unreleased_success_does_not_keep_its_logger_alive() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
import gc
import weakref
logger.pending = pending
reference = weakref.ref(logger)
del logger, pending
gc.collect()
assert reference() is None
",
);
});
}
}

View file

@ -16,13 +16,218 @@ mod deferred;
mod logger;
mod preparation;
mod python;
#[cfg(test)]
#[path = "../tests/support.rs"]
mod test_support;
pub(crate) use adapter::LegacyLogging;
pub use adapter::{LegacySurface, PassThroughStream};
pub use call::{PublicCall, run_legacy_call};
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
pub(crate) use preparation::prepare;
#[cfg(test)]
mod test_support {
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use crate::{LegacyLogging, LegacySurface, PublicCall};
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
/// signatures, and [`namespace`] binds every fake call against it.
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests
/// share one interpreter and run concurrently, so each fake is installed idempotently and
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
/// Every fake is bound against the contract first, so a call the real module would reject
/// fails here too.
const STUBS: &CStr = c"
import contextvars
import inspect
import json
import sys
import traceback
import types
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'):
sys.modules.setdefault(name, types.ModuleType(name))
legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python']
CONTRACT = json.loads(python_contract)
def contracted(name, fake):
signature = inspect.Signature(
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
)
def checked(*args, **kwargs):
signature.bind(*args, **kwargs)
return fake(*args, **kwargs)
return checked
if not hasattr(legacy, 'is_internal'):
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
FAKES = {
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
),
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
custom_llm_provider=provider,
),
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
original_response, api_key, additional_args
),
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
response, start, end
),
'failure_handler': lambda logger, error, start, end, asynchronous: (
logger.async_failure_handler if asynchronous else logger.failure_handler
)(error, ''.join(traceback.format_exception(error)), start, end),
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
'restore_context': lambda logger: logger.record('restore', None),
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
'is_internal_call': lambda: legacy.is_internal.get(),
'credential_list': lambda: [],
'warn_unknown_credential': lambda name, loaded: None,
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
'success', response, call_type
),
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
'stream_opened': lambda logger: logger.record('stream_opened', None),
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success', list(chunks)
),
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
}
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
for name, fake in FAKES.items():
setattr(legacy, name, contracted(name, fake))
unraisable = sys.modules.setdefault(
'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable')
)
if not hasattr(unraisable, 'events'):
unraisable.events = []
sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value))
def unraisable_from(owner):
return [error for source, error in unraisable.events if source is owner]
class StubCoroutine:
def __init__(self, logger):
self.logger = logger
def enqueue(self):
self.logger.record('enqueued', None)
self.logger.on_enqueue(self)
def close(self):
self.logger.record('closed', None)
class StubLogger:
def __init__(self):
self.calls = []
self.hooks = {}
self.on_enqueue = lambda coroutine: None
def record(self, name, value):
self.calls.append((name, value))
def names(self):
return [name for name, _ in self.calls]
def hook(self, phase, value, call_type):
self.record(phase + '_hook', call_type)
return self.hooks.get(phase, lambda value: 'awaitable')(value)
def check_limits(self, arguments):
self.record('check_limits', arguments)
def failure_handler(self, error, trace, start, end):
self.record('failure_handler', error)
def async_failure_handler(self, error, trace, start, end):
self.record('async_failure_handler', error)
return 'awaitable'
def success_handler(self, response, start, end):
self.record('success_handler', response)
def async_success_handler(self, response, start, end):
self.record('async_success_handler', response)
return StubCoroutine(self)
def handle_sync_success_callbacks_for_async_calls(self, response, start, end):
self.record('sync_success_for_async_call', response)
logger = StubLogger()
";
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
py.run(script, Some(&locals), Some(&locals)).unwrap();
locals
}
pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) {
py.run(code, Some(locals), Some(locals)).unwrap();
}
pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
locals.get_item(name).unwrap().unwrap()
}
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
pub(crate) fn legacy_call(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
asynchronous: bool,
) -> LegacyLogging {
let request = locals
.get_item("request")
.unwrap()
.unwrap_or_else(|| py.None().into_bound(py));
let kwargs = locals
.get_item("kwargs")
.unwrap()
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
.unwrap_or_else(|| PyDict::new(py));
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
LegacyLogging::new(
py,
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,
)
}
}

View file

@ -1,146 +0,0 @@
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::{PendingLogging, PendingSuccess};
use crate::PythonLogger;
use crate::test_support::{local, namespace, run};
/// A deferred success for the namespace's `logger` and `response`, bound as `pending`.
fn defer<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = namespace(py, c"response = object()");
run(py, &locals, script);
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: PythonLogger::new(local(&locals, "logger").unbind()),
response: Some(local(&locals, "response").unbind()),
start: py.None(),
end: Some(py.None()),
}),
},
)
.unwrap();
locals.set_item("pending", pending).unwrap();
locals
}
#[test]
fn release_enqueues_the_success_once_in_the_releasing_context() {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
from contextvars import ContextVar
marker = ContextVar('marker', default='unset')
observed = []
def on_enqueue(coroutine):
observed.append(marker.get())
pending.release(True)
logger.on_enqueue = on_enqueue
",
);
run(
py,
&locals,
c"
marker.set('release')
pending.release(True)
pending.release(True)
assert observed == ['release'], observed
assert logger.names() == ['async_success_handler', 'enqueued'], logger.calls
assert logger.calls[0][1] is response
",
);
});
}
#[test]
fn a_blocked_release_drops_the_success_for_good() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
pending.release(False)
pending.release(True)
assert logger.calls == [], logger.calls
",
);
});
}
#[rstest]
#[case::ordinary_error(c"RuntimeError('queue full')", false)]
#[case::cancellation(c"asyncio.CancelledError()", true)]
fn a_failed_enqueue_closes_the_coroutine_and_is_never_replayed(
#[case] failure: &CStr,
#[case] propagates: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = defer(
py,
c"
import asyncio
def on_enqueue(coroutine):
raise failure
logger.on_enqueue = on_enqueue
",
);
locals
.set_item("failure", py.eval(failure, None, Some(&locals)).unwrap())
.unwrap();
let released = local(&locals, "pending").call_method1("release", (true,));
match released {
Ok(_) => assert!(!propagates),
Err(error) => {
assert!(propagates);
assert!(error.value(py).is(local(&locals, "failure")));
}
}
locals.set_item("propagates", propagates).unwrap();
run(
py,
&locals,
c"
pending.release(True)
assert logger.names() == ['async_success_handler', 'enqueued', 'closed'], logger.calls
assert unraisable_from(logger) == ([] if propagates else [failure])
",
);
});
}
#[test]
fn an_unreleased_success_does_not_keep_its_logger_alive() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"");
run(
py,
&locals,
c"
import gc
import weakref
logger.pending = pending
reference = weakref.ref(logger)
del logger, pending
gc.collect()
assert reference() is None
",
);
});
}

View file

@ -1,282 +0,0 @@
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::LegacyLogging;
use crate::test_support::{legacy_call, local, namespace, run};
const CALL: &CStr = c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'logger': logger, 'document': document}
";
const TIMING: Timing = Timing {
start_time: 0.0,
end_time: 1.0,
};
fn begin<'py>(
py: Python<'py>,
locals: &Bound<'py, PyDict>,
asynchronous: bool,
) -> (LegacyLogging, LifecycleStep) {
let mut logging = legacy_call(py, locals, asynchronous);
let kwargs = local(locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let step = logging.begin(py, kwargs, 0.0).unwrap();
(logging, step)
}
fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> {
let LifecycleStep::Arguments(arguments) = step else {
panic!("expected the prepared arguments");
};
arguments.into_bound(py)
}
fn awaits_deployment_hook(step: &LifecycleStep) -> bool {
matches!(step, LifecycleStep::Await(_))
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn deployment_pre_call_hook_runs_only_for_asynchronous_calls(#[case] asynchronous: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, CALL);
let (_, step) = begin(py, &locals, asynchronous);
assert_eq!(awaits_deployment_hook(&step), asynchronous);
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names.contains(&"pre_hook".to_string()), asynchronous);
});
}
#[test]
fn kwargs_returned_by_the_pre_call_hook_are_what_the_call_prepares() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
replacement = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,ZWRpdGVk'}
kwargs = {'logger': logger, 'document': document}
replaced_kwargs = {'logger': logger, 'document': replacement, 'pages': [0]}
",
);
let (mut logging, step) = begin(py, &locals, true);
assert!(awaits_deployment_hook(&step));
let step = logging
.resume(py, Ok(local(&locals, "replaced_kwargs").unbind()))
.unwrap();
locals.set_item("prepared", arguments(py, step)).unwrap();
run(
py,
&locals,
c"
assert prepared['document'] is replacement
assert prepared['pages'] is replaced_kwargs['pages']
assert prepared['litellm_logging_obj'] is logger
assert 'litellm_logging_obj' not in replaced_kwargs
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked is prepared
",
);
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_keyword_the_bridge_never_reads_reaches_every_reader_as_the_callers_object(
#[case] asynchronous: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
opaque = object()
hooked = []
logger.hooks = {'pre': lambda kwargs: hooked.append(kwargs['vendor_extension']) or kwargs}
kwargs = {'logger': logger, 'vendor_extension': opaque}
",
);
let (mut logging, step) = begin(py, &locals, asynchronous);
let step = match step {
LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(),
step => step,
};
locals.set_item("prepared", arguments(py, step)).unwrap();
locals.set_item("asynchronous", asynchronous).unwrap();
run(
py,
&locals,
c"
assert prepared['vendor_extension'] is opaque
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked['vendor_extension'] is opaque
assert hooked == ([opaque] if asynchronous else []), hooked
",
);
});
}
#[test]
fn response_returned_by_the_post_call_hook_is_finalized_and_returned() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
kwargs = {'logger': logger}
response = object()
replacement = object()
logger.hooks = {'pre': lambda kwargs: kwargs}
",
);
let (mut logging, _) = begin(py, &locals, true);
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let step = logging
.after_success(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
assert!(awaits_deployment_hook(&step));
let step = logging
.resume(py, Ok(local(&locals, "replacement").unbind()))
.unwrap();
let LifecycleStep::Response(returned) = step else {
panic!("expected the finalized response");
};
assert!(returned.bind(py).is(local(&locals, "replacement")));
run(
py,
&locals,
c"
[finalized] = [value for name, value in logger.calls if name == 'finalize']
assert finalized is replacement
",
);
});
}
#[rstest]
#[case::pre_call(false)]
#[case::post_call(true)]
fn cancelling_a_deployment_hook_ends_the_call_with_that_cancellation(#[case] post_call: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"kwargs = {'logger': logger}\nresponse = object()");
let (mut logging, _) = begin(py, &locals, true);
if post_call {
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
logging
.after_success(py, local(&locals, "response").unbind(), TIMING)
.unwrap();
}
let cancellation = CancelledError::new_err("cancelled");
let cancelled = cancellation.value(py).clone();
let error = logging.resume(py, Err(cancellation)).err().unwrap();
assert!(error.value(py).is(&cancelled));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert!(!names.iter().any(|name| name.contains("handler")));
});
}
#[rstest]
#[case::hook_completed(false)]
#[case::hook_cancelled(true)]
fn failure_callbacks_run_after_the_failure_hook_however_it_ends(#[case] cancelled: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"kwargs = {'logger': logger}\nfailure = ValueError('provider')",
);
let (mut logging, _) = begin(py, &locals, true);
logging
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let failure = PyErr::from_value(local(&locals, "failure"));
let failed = LifecycleEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &failure,
};
let step = logging.emit(py, failed).unwrap();
assert!(awaits_deployment_hook(&step));
let hook_result = if cancelled {
Err(CancelledError::new_err("cancelled"))
} else {
Ok(py.None())
};
assert!(matches!(
logging.resume(py, hook_result).unwrap(),
LifecycleStep::Await(_)
));
run(
py,
&locals,
c"
assert logger.names()[-3:] == ['failure_hook', 'failure_handler', 'async_failure_handler'], logger.calls
assert all(value is failure for name, value in logger.calls if name.endswith('_handler'))
",
);
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_limit_rejected_before_the_call_surfaces_as_the_callers_error(#[case] asynchronous: bool) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
class BudgetExceeded(Exception):
pass
rejection = BudgetExceeded('over budget')
class LimitedLogger(StubLogger):
def check_limits(self, arguments):
raise rejection
logger = LimitedLogger()
logger.hooks = {'pre': lambda kwargs: kwargs}
kwargs = {'logger': logger}
",
);
let mut logging = legacy_call(py, &locals, asynchronous);
let kwargs = local(&locals, "kwargs")
.cast_into::<PyDict>()
.unwrap()
.unbind();
let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step {
LifecycleStep::Await(_) => logging.resume(py, Ok(local(&locals, "kwargs").unbind())),
step => Ok(step),
});
let error = result.err().unwrap();
assert!(error.value(py).is(local(&locals, "rejection")));
});
}

View file

@ -1,523 +0,0 @@
use std::ffi::CStr;
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
use proptest::prelude::*;
use pyo3::prelude::*;
use rstest::rstest;
use serde_json::{Map, Value, json};
use super::LegacyLogging;
use crate::PythonLogger;
use crate::test_support::{legacy_call, local, namespace, run};
/// The payload phases of `Logging` on top of `StubLogger`, with `pre_call` handing the
/// payload to the case's `on_pre_call`.
const PAYLOAD_LOGGER: &CStr = c"
class Request:
pass
class PayloadLogger(StubLogger):
def update_from_kwargs(self, **update):
self.update = update
def pre_call(self, input, api_key, additional_args):
self.record('pre_call', None)
self.pre = additional_args
self.pre_api_key = api_key
on_pre_call(additional_args)
def post_call(self, original_response, api_key, additional_args):
self.record('post_call', None)
self.post = (original_response, api_key, additional_args)
request = Request()
kwargs = {}
logger = PayloadLogger()
on_pre_call = lambda additional_args: None
check = lambda: None
";
const DOCUMENT: &str = "data:application/pdf;base64,YWJj";
const EDITED: &str = "data:application/pdf;base64,ZWRpdGVk";
fn document(source: &str) -> Value {
json!({"type": "document_url", "document_url": source})
}
fn before_send(script: &CStr, body: Value) -> WireRequest {
before_send_with_secrets(script, json!({}), body, &[])
}
/// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with
/// the Python objects `script` binds, then delivers the provider's raw response the way the
/// driver does and runs the script's `check()`.
fn before_send_with_secrets(
script: &CStr,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
before_send_bound(&[], script, optional_params, body, secret_fields)
}
/// [`before_send_with_secrets`] with `bindings` placed in the namespace before `script` runs.
fn before_send_bound(
bindings: &[(&str, &Value)],
script: &CStr,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
for &(name, value) in bindings {
locals.set_item(name, to_py(py, value).unwrap()).unwrap();
}
run(py, &locals, script);
let mut logging = LegacyLogging {
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
..legacy_call(py, &locals, false)
};
let context = RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params,
secret_fields: secret_fields.iter().map(|name| name.to_string()).collect(),
api_key: Some(SecretValue::new("route-key")),
};
let wire = WireRequest {
url: "https://provider.invalid/ocr".into(),
headers: vec![("x-route".into(), "route".into())],
body,
};
let step = logging.before_send(py, Box::new(wire), &context).unwrap();
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
};
assert!(matches!(
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
LifecycleStep::Done
));
run(py, &locals, c"check()");
let LifecycleStep::Wire(wire) = step else {
panic!("before_send did not hand back the wire request");
};
*wire
})
}
#[rstest]
#[case::caller_keyword(c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
pages = [0]
kwargs = {'document': document, 'pages': pages}
observed = []
on_pre_call = lambda args: observed.append(
(args['complete_input_dict']['document'] is document, args['complete_input_dict']['pages'] is pages)
)
def check():
assert observed == [(True, True)], observed
")]
#[case::request_attribute_behind_an_omitted_keyword(c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
pages = [0]
request.document = document
kwargs = {'pages': pages}
observed = []
on_pre_call = lambda args: observed.append(
(args['complete_input_dict']['document'] is document, args['complete_input_dict']['pages'] is pages)
)
def check():
assert observed == [(True, True)], observed
")]
fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) {
let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]});
let wire = before_send(script, body.clone());
assert_eq!(wire.body, body);
}
#[test]
fn pre_call_edit_of_a_passthrough_object_reaches_the_caller_and_the_wire() {
let wire = before_send(
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'document': document}
def on_pre_call(args):
args['complete_input_dict']['document']['document_url'] = 'data:application/pdf;base64,ZWRpdGVk'
def check():
assert document['document_url'] == 'data:application/pdf;base64,ZWRpdGVk'
",
json!({"document": document(DOCUMENT)}),
);
assert_eq!(wire.body["document"], document(EDITED));
}
#[test]
fn a_body_key_the_route_rewrote_is_not_the_callers_object() {
let wire = before_send(
c"
document = {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
kwargs = {'document': document}
observed = []
def on_pre_call(args):
observed.append(args['complete_input_dict']['document'] is document)
args['complete_input_dict']['document']['document_name'] = 'edited.pdf'
def check():
assert observed == [False], observed
assert document == {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
",
json!({"document": document(DOCUMENT)}),
);
assert_eq!(
wire.body["document"],
json!({"type": "document_url", "document_url": DOCUMENT, "document_name": "edited.pdf"})
);
}
#[test]
fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() {
let body = json!({"pages": [0]});
let wire = before_send(
c"
opaque = object()
kwargs = {'pages': opaque}
observed = []
on_pre_call = lambda args: observed.append(args['complete_input_dict']['pages'])
def check():
assert observed == [[0]], observed
",
body.clone(),
);
assert_eq!(wire.body, body);
}
#[rstest]
#[case::body(
c"
def on_pre_call(args):
args['complete_input_dict'] = {'replacement': True}
"
)]
#[case::headers(
c"
def on_pre_call(args):
args['headers'] = {'x-replacement': 'yes'}
"
)]
fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, body.clone());
assert_eq!(wire.body, body);
assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
#[test]
fn pre_call_header_edit_reaches_the_wire() {
let wire = before_send(
c"
def on_pre_call(args):
args['headers']['x-callback'] = 'edited'
",
json!({}),
);
assert_eq!(
wire.headers,
[
("x-route".to_string(), "route".to_string()),
("x-callback".to_string(), "edited".to_string()),
]
);
}
#[test]
fn pre_call_receives_the_wire_request_and_the_logger_its_redacted_request() {
let body = json!({"model": "model", "document": document(DOCUMENT)});
before_send_with_secrets(
c"
logger_fn = lambda *args: None
kwargs = {
'litellm_call_id': 'call-1',
'client_secret': 'shh',
'proxy_server_request': {'body': {}},
'logger_fn': logger_fn,
'litellm_request_debug': True,
'ocr_cost_per_page': 0.05,
}
observed = []
on_pre_call = observed.append
def check():
[args] = observed
assert args['api_base'] == 'https://provider.invalid/ocr', args
assert args['complete_input_dict'] == {
'model': 'model',
'document': {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'},
}, args
update = logger.update
assert update['model'] == 'model' and update['custom_llm_provider'] == 'provider', update
assert update['litellm_params']['litellm_call_id'] == 'call-1', update
assert update['litellm_params']['api_base'] == 'https://provider.invalid/ocr', update
assert update['litellm_params']['logger_fn'] is logger_fn, update
assert update['litellm_params']['litellm_request_debug'] is True, update
assert update['litellm_params']['ocr_cost_per_page'] == 0.05, update
assert update['kwargs']['client_secret'] == '****', update
assert 'proxy_server_request' not in update['kwargs'], update
assert update['optional_params']['client_secret'] == '****', update
",
json!({"client_secret": "shh"}),
body,
&["client_secret"],
);
}
#[rstest]
#[case::added_key(
c"
def on_pre_call(args):
args['complete_input_dict']['include_image_base64'] = True
",
json!({"document": document(DOCUMENT), "include_image_base64": true})
)]
#[case::replaced_document(
c"
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
kwargs = {'document': document}
def on_pre_call(args):
args['complete_input_dict']['document'] = {
'type': 'document_url', 'document_url': 'data:application/pdf;base64,ZWRpdGVk'
}
def check():
assert document['document_url'] == 'data:application/pdf;base64,YWJj', document
",
json!({"document": document(EDITED)})
)]
#[case::retained_body_edited_after_rebinding(
c"
def on_pre_call(args):
retained = args['complete_input_dict']
args['complete_input_dict'] = {'rebound': True}
retained['include_image_base64'] = True
",
json!({"document": document(DOCUMENT), "include_image_base64": true})
)]
fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, body);
assert_eq!(wire.body, expected);
}
#[test]
fn retained_headers_edited_after_rebinding_reach_the_wire() {
let wire = before_send(
c"
def on_pre_call(args):
retained = args['headers']
args['headers'] = {'x-rebound': 'rebound'}
retained['x-retained'] = 'sent'
",
json!({}),
);
assert_eq!(
wire.headers,
[
("x-route".to_string(), "route".to_string()),
("x-retained".to_string(), "sent".to_string()),
]
);
}
#[test]
fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() {
before_send(
c"
def check():
original_response, api_key, additional_args = logger.post
assert original_response == 'raw response', original_response
assert api_key == logger.pre_api_key == 'route-key', (api_key, logger.pre_api_key)
assert additional_args == {
'complete_input_dict': logger.pre['complete_input_dict'],
'headers': logger.pre['headers'],
}, additional_args
assert additional_args['complete_input_dict'] is logger.pre['complete_input_dict']
assert additional_args['headers'] is logger.pre['headers']
",
json!({"document": document(DOCUMENT)}),
);
}
#[test]
fn every_request_runs_the_full_pre_call_and_post_call() {
let wire = before_send(
c"
def on_pre_call(args):
args['complete_input_dict']['include_image_base64'] = True
def check():
assert logger.names() == ['pre_call', 'post_call'], logger.calls
",
json!({"document": document(DOCUMENT)}),
);
assert_eq!(
wire.body,
json!({"document": document(DOCUMENT), "include_image_base64": true})
);
}
/// What one pre-call callback does to the payload it is handed.
#[derive(Clone, Debug)]
enum Edit {
Nothing,
Set(String, Value),
Remove(String),
Rebind(Value),
RebindThenSetRetained(String, Value),
}
impl Edit {
fn script(&self) -> Value {
match self {
Self::Nothing => json!({"kind": "nothing"}),
Self::Set(key, value) => json!({"kind": "set", "key": key, "value": value}),
Self::Remove(key) => json!({"kind": "remove", "key": key}),
Self::Rebind(value) => json!({"kind": "rebind", "value": value}),
Self::RebindThenSetRetained(key, value) => {
json!({"kind": "rebind_then_set_retained", "key": key, "value": value})
}
}
}
/// The legacy contract: the provider is sent the body object `pre_call` received, as
/// the callback left it. Rebinding the envelope's key points the envelope elsewhere and
/// leaves that object alone.
fn sent(&self, body: &Map<String, Value>) -> Value {
let mut sent = body.clone();
match self {
Self::Nothing | Self::Rebind(_) => {}
Self::Set(key, value) | Self::RebindThenSetRetained(key, value) => {
sent.insert(key.clone(), value.clone());
}
Self::Remove(key) => {
sent.remove(key);
}
}
Value::Object(sent)
}
}
/// How the caller's keyword for a body key relates to what the route sends under it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Caller {
PassedUnchanged,
RewrittenByTheRoute,
NotPassed,
}
const MODEL: &CStr = c"
aliased = {}
def on_pre_call(args):
body = args['complete_input_dict']
aliased.update({name: body[name] is kwargs[name] for name in unchanged})
kind = edit['kind']
if kind == 'set':
body[edit['key']] = edit['value']
elif kind == 'remove':
body.pop(edit['key'], None)
elif kind == 'rebind':
args['complete_input_dict'] = edit['value']
elif kind == 'rebind_then_set_retained':
args['complete_input_dict'] = {}
body[edit['key']] = edit['value']
def check():
assert aliased == {name: True for name in unchanged}, aliased
assert logger.names() == ['pre_call', 'post_call'], logger.calls
";
fn json_value() -> impl Strategy<Value = Value> {
let leaf = prop_oneof![
Just(Value::Null),
any::<bool>().prop_map(Value::from),
any::<i64>().prop_map(Value::from),
any::<f64>()
.prop_filter("JSON has no NaN or infinity", |number| number.is_finite())
.prop_map(Value::from),
".{0,8}".prop_map(Value::from),
];
leaf.prop_recursive(3, 24, 4, |inner| {
prop_oneof![
prop::collection::vec(inner.clone(), 0..4).prop_map(Value::from),
prop::collection::btree_map(key(), inner, 0..4)
.prop_map(|fields| Value::Object(fields.into_iter().collect())),
]
})
}
fn key() -> impl Strategy<Value = String> {
"[a-z]{1,6}"
}
fn caller() -> impl Strategy<Value = Caller> {
prop_oneof![
Just(Caller::PassedUnchanged),
Just(Caller::RewrittenByTheRoute),
Just(Caller::NotPassed),
]
}
fn edit() -> impl Strategy<Value = Edit> {
prop_oneof![
Just(Edit::Nothing),
(key(), json_value()).prop_map(|(key, value)| Edit::Set(key, value)),
key().prop_map(Edit::Remove),
json_value().prop_map(Edit::Rebind),
(key(), json_value()).prop_map(|(key, value)| Edit::RebindThenSetRetained(key, value)),
]
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(128))]
/// For any body, any caller keywords and any callback edit: every keyword the route
/// sends unchanged reaches `pre_call` as the caller's own object, and the provider is
/// sent exactly what the model says, so a callback that edits nothing changes nothing.
#[test]
fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it(
fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5),
edit in edit(),
) {
let body: Map<String, Value> = fields
.iter()
.map(|(name, (value, _))| (name.clone(), value.clone()))
.collect();
let kwargs: Map<String, Value> = fields
.iter()
.filter_map(|(name, (value, caller))| match caller {
Caller::PassedUnchanged => Some((name.clone(), value.clone())),
Caller::RewrittenByTheRoute => Some((name.clone(), json!([value]))),
Caller::NotPassed => None,
})
.collect();
let unchanged: Value = fields
.iter()
.filter(|(_, (_, caller))| *caller == Caller::PassedUnchanged)
.map(|(name, _)| Value::from(name.clone()))
.collect();
let wire = before_send_bound(
&[
("kwargs", &Value::Object(kwargs)),
("unchanged", &unchanged),
("edit", &edit.script()),
],
MODEL,
json!({}),
Value::Object(body.clone()),
&[],
);
prop_assert_eq!(wire.body, edit.sent(&body));
prop_assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
}

View file

@ -1,205 +0,0 @@
use std::ffi::CStr;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use crate::{LegacyLogging, LegacySurface, PublicCall};
/// The parameters of every `callbacks_legacy_python` function, as the real module declares them.
/// `tests/test_litellm/rust_bridge/test_callbacks_legacy_python.py` pins this file to the Python
/// signatures, and [`namespace`] binds every fake call against it.
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
/// Stand-ins for `callbacks_legacy_python`, the only Python module the crate calls. Tests
/// share one interpreter and run concurrently, so each fake is installed idempotently and
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
/// Every fake is bound against the contract first, so a call the real module would reject
/// fails here too.
const STUBS: &CStr = c"
import contextvars
import inspect
import json
import sys
import traceback
import types
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.callbacks_legacy_python'):
sys.modules.setdefault(name, types.ModuleType(name))
legacy = sys.modules['litellm.rust_bridge.callbacks_legacy_python']
CONTRACT = json.loads(python_contract)
def contracted(name, fake):
signature = inspect.Signature(
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
)
def checked(*args, **kwargs):
signature.bind(*args, **kwargs)
return fake(*args, **kwargs)
return checked
if not hasattr(legacy, 'is_internal'):
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
FAKES = {
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
),
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
custom_llm_provider=provider,
),
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
original_response, api_key, additional_args
),
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
response, start, end
),
'failure_handler': lambda logger, error, start, end, asynchronous: (
logger.async_failure_handler if asynchronous else logger.failure_handler
)(error, ''.join(traceback.format_exception(error)), start, end),
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
'restore_context': lambda logger: logger.record('restore', None),
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
'is_internal_call': lambda: legacy.is_internal.get(),
'credential_list': lambda: [],
'warn_unknown_credential': lambda name, loaded: None,
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
'success', response, call_type
),
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
'stream_opened': lambda logger: logger.record('stream_opened', None),
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success', list(chunks)
),
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
}
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
for name, fake in FAKES.items():
setattr(legacy, name, contracted(name, fake))
unraisable = sys.modules.setdefault(
'litellm_test_unraisable', types.ModuleType('litellm_test_unraisable')
)
if not hasattr(unraisable, 'events'):
unraisable.events = []
sys.unraisablehook = lambda event: unraisable.events.append((event.object, event.exc_value))
def unraisable_from(owner):
return [error for source, error in unraisable.events if source is owner]
class StubCoroutine:
def __init__(self, logger):
self.logger = logger
def enqueue(self):
self.logger.record('enqueued', None)
self.logger.on_enqueue(self)
def close(self):
self.logger.record('closed', None)
class StubLogger:
def __init__(self):
self.calls = []
self.hooks = {}
self.on_enqueue = lambda coroutine: None
def record(self, name, value):
self.calls.append((name, value))
def names(self):
return [name for name, _ in self.calls]
def hook(self, phase, value, call_type):
self.record(phase + '_hook', call_type)
return self.hooks.get(phase, lambda value: 'awaitable')(value)
def check_limits(self, arguments):
self.record('check_limits', arguments)
def failure_handler(self, error, trace, start, end):
self.record('failure_handler', error)
def async_failure_handler(self, error, trace, start, end):
self.record('async_failure_handler', error)
return 'awaitable'
def success_handler(self, response, start, end):
self.record('success_handler', response)
def async_success_handler(self, response, start, end):
self.record('async_success_handler', response)
return StubCoroutine(self)
def handle_sync_success_callbacks_for_async_calls(self, response, start, end):
self.record('sync_success_for_async_call', response)
logger = StubLogger()
";
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
py.run(script, Some(&locals), Some(&locals)).unwrap();
locals
}
pub(crate) fn run(py: Python<'_>, locals: &Bound<'_, PyDict>, code: &CStr) {
py.run(code, Some(locals), Some(locals)).unwrap();
}
pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyAny> {
locals.get_item(name).unwrap().unwrap()
}
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
pub(crate) fn legacy_call(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
asynchronous: bool,
) -> LegacyLogging {
let request = locals
.get_item("request")
.unwrap()
.unwrap_or_else(|| py.None().into_bound(py));
let kwargs = locals
.get_item("kwargs")
.unwrap()
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
.unwrap_or_else(|| PyDict::new(py));
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
LegacyLogging::new(
py,
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,
)
}

View file

@ -1,291 +0,0 @@
use std::ffi::CStr;
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use pyo3::exceptions::PyRuntimeError;
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use rstest::rstest;
use super::LegacyLogging;
use crate::PythonLogger;
use crate::test_support::{legacy_call, local, namespace, run};
const TIMING: Timing = Timing {
start_time: 0.0,
end_time: 1.0,
};
fn logged(py: Python<'_>, locals: &Bound<'_, PyDict>, asynchronous: bool) -> LegacyLogging {
LegacyLogging {
logger: Some(PythonLogger::new(local(locals, "logger").unbind())),
..legacy_call(py, locals, asynchronous)
}
}
fn succeed(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
logging: &mut LegacyLogging,
) -> LifecycleStep {
let response = local(locals, "response").unbind();
logging
.emit(
py,
LifecycleEvent::Succeeded {
timing: TIMING,
response: &response,
},
)
.unwrap()
}
fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> LifecycleStep {
let failure = PyErr::from_value(local(locals, "failure"));
logging
.emit(
py,
LifecycleEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Host,
error: &failure,
},
)
.unwrap()
}
#[rstest]
#[case::sync_listened(false, c"", &["submit"])]
#[case::async_listened(
true,
c"",
&["async_success_handler", "enqueued", "sync_success_for_async_call"]
)]
#[case::async_deferred(true, c"logger._defer_async_logging = True", &["sync_success_for_async_call"])]
#[case::async_with_fallbacks(true, c"kwargs = {'fallbacks': ['other']}", &["sync_success_for_async_call"])]
fn success_reaches_the_logging_handlers(
#[case] asynchronous: bool,
#[case] script: &CStr,
#[case] expected: &[&str],
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"response = object()");
run(py, &locals, script);
let mut logging = logged(py, &locals, asynchronous);
assert!(matches!(
succeed(py, &locals, &mut logging),
LifecycleStep::Done
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names, expected);
run(
py,
&locals,
c"
assert all(value is response for name, value in logger.calls if name.endswith('_handler'))
assert hasattr(logger, '_native_pending_logging') == getattr(logger, '_defer_async_logging', False)
",
);
});
}
#[rstest]
#[case::synchronous(false, &["failure_handler"])]
#[case::asynchronous(true, &[])]
fn internal_calls_skip_failure_callbacks_only_when_asynchronous(
#[case] asynchronous: bool,
#[case] expected: &[&str],
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
let mut logging = LegacyLogging {
internal: true,
..logged(py, &locals, asynchronous)
};
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Done
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names, expected);
});
}
#[test]
fn internal_async_calls_skip_the_async_success_fan_out() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"response = object()");
let mut logging = LegacyLogging {
internal: true,
..logged(py, &locals, true)
};
succeed(py, &locals, &mut logging);
run(
py,
&locals,
c"assert logger.names() == ['sync_success_for_async_call'], logger.calls",
);
});
}
#[test]
fn a_failing_success_callback_is_reported_without_replacing_the_response() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
response = object()
failure = ValueError('terminal diagnostic')
class FailingLogger(StubLogger):
def handle_sync_success_callbacks_for_async_calls(self, *args):
raise failure
logger = FailingLogger()
",
);
let mut logging = logged(py, &locals, true);
assert!(matches!(
succeed(py, &locals, &mut logging),
LifecycleStep::Done
));
assert!(
logging
.response
.as_ref()
.unwrap()
.bind(py)
.is(local(&locals, "response"))
);
run(py, &locals, c"assert unraisable_from(logger) == [failure]");
});
}
#[rstest]
#[case::sync_listened(false, c"", &["failure_handler"])]
#[case::async_listened(true, c"", &["failure_handler", "async_failure_handler"])]
fn failure_reaches_the_logging_handlers(
#[case] asynchronous: bool,
#[case] script: &CStr,
#[case] expected: &[&str],
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
run(py, &locals, script);
let mut logging = logged(py, &locals, asynchronous);
let step = fail(py, &locals, &mut logging);
let awaits_async_handler = expected.contains(&"async_failure_handler");
assert_eq!(
matches!(step, LifecycleStep::Await(_)),
awaits_async_handler
);
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
.extract()
.unwrap();
assert_eq!(names, expected);
run(
py,
&locals,
c"assert all(value is failure for name, value in logger.calls if name.endswith('_handler'))",
);
});
}
#[test]
fn a_failing_sync_failure_callback_keeps_the_error_and_still_runs_the_async_family() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
failure = ValueError('selected')
class FailingLogger(StubLogger):
def failure_handler(self, error, trace, start, end):
self.record('failure_handler', error)
raise RuntimeError('handler failed')
logger = FailingLogger()
",
);
let mut logging = logged(py, &locals, true);
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Await(_)
));
assert!(
logging
.error
.as_ref()
.unwrap()
.bind(py)
.is(local(&locals, "failure"))
);
run(
py,
&locals,
c"assert logger.names() == ['failure_handler', 'async_failure_handler'], logger.calls",
);
});
}
#[rstest]
#[case::completed(None, true)]
#[case::handler_error(Some(false), true)]
#[case::cancelled(Some(true), false)]
fn the_async_failure_handler_ends_the_call_unless_it_was_cancelled(
#[case] error: Option<bool>,
#[case] done: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"failure = ValueError('provider')");
let mut logging = logged(py, &locals, true);
fail(py, &locals, &mut logging);
let result = match error {
None => Ok(py.None()),
Some(false) => Err(PyRuntimeError::new_err("handler failed")),
Some(true) => Err(CancelledError::new_err("cancelled")),
};
let expected = result.as_ref().err().map(|error| error.value(py).clone());
match logging.resume(py, result) {
Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)),
Err(propagated) => {
assert!(!done);
assert!(propagated.value(py).is(expected.unwrap()));
}
}
});
}
#[test]
fn closing_restores_the_correlation_context_once() {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, c"");
let mut logging = logged(py, &locals, true);
logging.close(py);
logging.close(py);
run(
py,
&locals,
c"assert logger.names() == ['restore'], logger.calls",
);
});
}

View file

@ -4,7 +4,6 @@ version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
autotests = false
[dependencies]
litellm-secrets.workspace = true

View file

@ -14,6 +14,3 @@ pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Resu
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await
}
#[cfg(test)]
mod tests;

View file

@ -51,6 +51,3 @@ pub fn chat_completions_decline_reason(
.unsupported_reason(&messages, optional_params)
.map(|reason| reason.0)
}
#[cfg(test)]
mod tests;

View file

@ -143,3 +143,841 @@ pub(super) fn prepare_provider_request(
timeout: request.timeout,
})
}
#[cfg(test)]
mod tests {
use litellm_llms::base_llm::chat::transformation::RequestAuth;
use serde_json::{Map, Value, json};
use super::{prepare_provider_request, resolve_request};
use crate::chat_completions::{
Error,
types::{ChatCompletionsRequest, ProviderChatCompletionsRequest},
};
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
}
fn request<'a>(
model: &'a str,
provider: Option<&'a str>,
messages: Value,
optional_params: Value,
) -> ChatCompletionsRequest<'a> {
ChatCompletionsRequest {
model,
messages,
optional_params: match optional_params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: provider,
extra_headers: None,
timeout: None,
}
}
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
/// carry resolved credentials), so unwrap the failure case by hand.
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
match prepare_chat_completions_call(request) {
Err(error) => error,
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
}
}
#[test]
fn resolves_the_provider_from_the_model_prefix() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
assert_eq!(prepared.url, "https://api.anthropic.com/v1/messages");
assert_eq!(prepared.body["model"], json!("claude-sonnet-4-5"));
}
#[test]
fn strips_an_explicit_provider_prefix_from_the_model() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
}
#[test]
fn adds_the_auth_and_default_headers() {
let prepared = prepare_chat_completions_call(request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert!(
prepared
.upstream_headers
.contains(&("x-api-key".to_string(), "sk-test".to_string()))
);
assert!(
prepared
.upstream_headers
.contains(&("anthropic-version".to_string(), "2023-06-01".to_string()))
);
assert!(matches!(
prepared.auth,
RequestAuth::Header {
name: "x-api-key",
..
}
));
}
#[test]
fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
// Python builds `{**headers, **anthropic_headers}`, so the deployment's key
// overwrites a forwarded one. Honouring the caller's would let whoever sends
// the request choose the Anthropic principal it bills to.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"X-Api-Key".to_string(),
json!("sk-caller"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
}
#[test]
fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
// Anthropic's `validate_environment` pops `x-api-key` and sets `authorization`
// for an OAuth token, so re-adding the key here would put the credential into
// a header the host removed on purpose.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
(
"Authorization".to_string(),
json!("Bearer sk-ant-oat01-token"),
),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"),
"the resolved key must not be applied over an OAuth bearer, got {:?}",
prepared.upstream_headers
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-token")
);
}
#[test]
fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
// Only an OAuth bearer replaces the credential. Python sends the deployment's
// `x-api-key` alongside any other forwarded `authorization`, so deferring on
// the mere presence of that header would drop the deployment's auth.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer unrelated")),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer unrelated"),
"the unrelated authorization must survive, got {:?}",
prepared.upstream_headers
);
}
#[test]
fn declines_an_unsupported_request_before_resolving_credentials() {
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
);
call.api_key = None;
// No api_key is set and no env is consulted: the gate must run first, so the
// error is the decline rather than a missing-credential error.
assert_eq!(decline(call), Error::Unsupported("streaming"));
}
#[test]
fn rejects_an_unknown_provider() {
assert_eq!(
decline(request(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider("openai".to_string())
);
}
#[test]
fn rejects_a_model_with_no_resolvable_provider() {
assert!(matches!(
decline(request(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider(_)
));
}
#[test]
fn rejects_an_empty_or_malformed_message_list() {
assert_eq!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!([]),
json!({}),
)),
Error::InvalidRequest("chat completions requires at least one message".to_string())
);
assert!(matches!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!("not a list"),
json!({}),
)),
Error::InvalidRequest(_)
));
}
#[test]
fn rejects_non_string_extra_headers() {
let mut call = request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
assert_eq!(
decline(call),
Error::Headers(litellm_http::request::HeaderError {
context: "chat completions",
name: "x-trace".to_string(),
actual: "number",
})
);
}
#[test]
fn prepares_a_bedrock_call_without_resolving_credentials() {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.api_key = None;
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert_eq!(
prepared.url,
"https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse"
);
assert_eq!(
prepared.auth,
RequestAuth::AwsSigV4 {
region: "us-east-1".to_string(),
service: "bedrock",
}
);
// SigV4 signs the serialized body, so prepare must not have added an
// Authorization header; the handler does it.
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
);
assert_eq!(prepared.body["inferenceConfig"], json!({"maxTokens": 16}));
}
#[tokio::test]
async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
// Python signs only the AWS header set and reattaches the rest, so a header
// the caller forwarded rides along without joining the canonical request.
// Signing it makes Converse 403 on a deployment that works on Python.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
// A key would resolve to a bearer token and never reach the signer.
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(
"x-request-id".to_string(),
json!("abc-123"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let signed = crate::chat_completions::handler::outbound_request(&prepared)
.await
.expect("signs");
let authorization = signed
.header("authorization")
.expect("carries an authorization header")
.to_string();
assert!(
authorization.starts_with("AWS4-HMAC-SHA256"),
"expected a SigV4 signature, got {authorization}"
);
assert!(
!authorization.contains("x-request-id"),
"forwarded header reached SignedHeaders: {authorization}"
);
// It still goes on the wire, it is just not part of the signature.
assert!(
signed
.headers()
.iter()
.any(|(name, value)| name == "x-request-id" && value == "abc-123"),
"forwarded header was dropped instead of reattached"
);
}
#[tokio::test]
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
// Reattaching the caller's copy next to the computed one puts the name on
// the wire twice and Bedrock rejects the pair, so a request carrying one
// has to go to Python instead of being signed here.
for forwarded in [
"Authorization",
"x-amz-date",
"x-amz-security-token",
"Date",
] {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let error = crate::chat_completions::handler::outbound_request(&prepared)
.await
.expect_err("{forwarded} should decline instead of being signed");
assert!(
matches!(error, Error::Unsupported(_)),
"{forwarded} declined as {error:?}, which the host would not fall back on"
);
}
}
#[test]
fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
// `get_request_headers` assigns `headers["Authorization"]` unconditionally
// once a bearer token resolves, so the deployment's identity wins on
// Python. Keeping the caller's would authorize and bill the call as a
// different principal, and only when the deployment carries `rust: true`.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.extra_headers = Some(Map::from_iter([(
"Authorization".to_string(),
json!("Bearer caller-supplied"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let authorizations: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("authorization"))
.map(|(_, value)| value.as_str())
.collect();
assert_eq!(
authorizations,
vec!["Bearer sk-test"],
"the deployment token must be the only authorization on the wire"
);
}
#[test]
fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
// The opposite precedence, and deliberate: Anthropic's own transform
// honours a forwarded OAuth bearer, so the Bedrock fix above must not be
// generalized into a rule that the configured key always wins.
//
// An OAuth bearer is the whole of that exception. This forwarded a plain
// `x-api-key` until round 17, which read as the same claim and was not:
// Python overwrites a forwarded `x-api-key` with the deployment's.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"authorization".to_string(),
json!("Bearer sk-ant-oat01-forwarded"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.map(|(_, value)| value.as_str())
.collect();
assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-forwarded")
);
}
#[test]
fn a_bedrock_api_key_is_sent_as_a_bearer_token_instead_of_being_signed() {
// The configured bearer identity has its own account and quota boundary,
// so a request carrying one must not be signed as whatever principal the
// host's AWS credentials resolve to.
let prepared = prepare_chat_completions_call(request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
))
.expect("prepares");
assert_eq!(
prepared.auth,
RequestAuth::Bearer {
token: "sk-test".to_string()
}
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-test"),
"prepare did not carry the bearer token"
);
}
fn decline_reason(
model: &str,
provider: Option<&str>,
messages: Value,
params: Value,
) -> Option<&'static str> {
let params = match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
};
crate::chat_completions::chat_completions_decline_reason(model, provider, messages, &params)
}
#[test]
fn the_gate_accepts_what_prepare_accepts() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
None
);
}
#[test]
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
),
Some("streaming")
);
assert_eq!(
decline_reason(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!("nope"),
json!({})
),
Some("unreadable message list")
);
assert_eq!(
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
Some("empty message list")
);
}
#[test]
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
// A gate that accepts what prepare then declines would make the host emit
// its pre-call logging on a path that falls back, so pin the agreement.
for (messages, params) in [
(
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 8}),
),
(
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
json!({"temperature": 0.1}),
),
(
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
json!({}),
),
] {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params.clone()
),
None,
"gate declined {messages}"
);
prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params,
))
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
}
}
mod round_trip {
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::*;
use crate::chat_completions::chat_completions;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n")
{
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
fn http_response(status: &str, body: &str) -> String {
format!(
"HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
)
}
/// Serve one request from a stub upstream and hand back what it received.
async fn serve_once(
status: &'static str,
body: &'static str,
) -> (String, tokio::task::JoinHandle<String>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let port = listener.local_addr().expect("addr").port();
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts");
let received = read_http_request(&mut socket).await;
socket
.write_all(http_response(status, body).as_bytes())
.await
.expect("writes response");
socket.flush().await.expect("flushes");
received
});
(format!("http://127.0.0.1:{port}/v1/messages"), handle)
}
fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> {
ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages,
optional_params: match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: Some(api_base),
custom_llm_provider: None,
extra_headers: None,
timeout: Some(std::time::Duration::from_secs(10)),
}
}
const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
#[tokio::test]
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
let response = chat_completions(call(
&api_base,
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
]),
json!({"max_tokens": 16}),
))
.await
.expect("call succeeds");
let received = handle.await.expect("server task");
let sent: Value = serde_json::from_str(
received
.split_once("\r\n\r\n")
.expect("request has a body")
.1,
)
.expect("body is json");
assert_eq!(
sent["messages"],
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
);
assert_eq!(
sent["system"],
json!([{"type": "text", "text": "be terse"}])
);
assert_eq!(sent["max_tokens"], json!(16));
assert!(received.to_lowercase().contains("x-api-key: sk-test"));
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
assert_eq!(response.usage.total_tokens, 15);
}
#[tokio::test]
async fn a_response_it_cannot_normalize_is_reported_as_already_sent() {
// The provider was called and billed, so the host must not retry this
// on its own path. `MissingField` here would read as a pre-send
// decline and be retried; `InvalidResponse` cannot.
const NO_USAGE: &str =
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn an_upstream_error_status_keeps_its_code() {
let (api_base, handle) =
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("upstream rejects");
handle.await.expect("server task");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Http { status: 429, .. })
),
"expected a 429, got {err:?}"
);
}
#[tokio::test]
async fn a_connection_that_is_never_established_declines_instead_of_failing() {
// Nothing was sent, so nothing was billed and the host can still serve
// the request. Classing this with the post-send failures would turn a
// recoverable fallback into a user-facing error on exactly the
// deployments whose transport is configured only on the Python client.
let port = {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
listener.local_addr().expect("has an address").port()
// Dropped here, so the port is closed and the connect is refused.
};
let err = chat_completions(call(
&format!("http://127.0.0.1:{port}/v1/messages"),
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("nothing is listening");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Connect(_))
),
"expected a pre-send connect failure, got {err:?}"
);
}
#[test]
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
use crate::chat_completions::handler::as_response_error;
for original in [
Error::MissingField("usage"),
Error::Unsupported("non-text response content block"),
Error::InvalidRequest("whatever".to_string()),
Error::Auth(litellm_auth::Error::InvalidHeader),
] {
let label = format!("{original:?}");
assert!(
matches!(as_response_error(original), Error::InvalidResponse(_)),
"{label} must not stay retryable once the provider has answered"
);
}
// An upstream status is already unambiguous, so it survives intact.
assert!(matches!(
as_response_error(Error::Transport(litellm_http::transport::Error::Http {
status: 500,
body: "boom".to_string()
})),
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
));
}
}
}

View file

@ -1,833 +0,0 @@
use litellm_llms::base_llm::chat::transformation::RequestAuth;
use serde_json::{Map, Value, json};
use super::{
Error,
prepare::{prepare_provider_request, resolve_request},
};
use crate::chat_completions::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
}
fn request<'a>(
model: &'a str,
provider: Option<&'a str>,
messages: Value,
optional_params: Value,
) -> ChatCompletionsRequest<'a> {
ChatCompletionsRequest {
model,
messages,
optional_params: match optional_params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: provider,
extra_headers: None,
timeout: None,
}
}
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
/// carry resolved credentials), so unwrap the failure case by hand.
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
match prepare_chat_completions_call(request) {
Err(error) => error,
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
}
}
#[test]
fn resolves_the_provider_from_the_model_prefix() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
assert_eq!(prepared.url, "https://api.anthropic.com/v1/messages");
assert_eq!(prepared.body["model"], json!("claude-sonnet-4-5"));
}
#[test]
fn strips_an_explicit_provider_prefix_from_the_model() {
let prepared = prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert_eq!(prepared.model, "claude-sonnet-4-5");
}
#[test]
fn adds_the_auth_and_default_headers() {
let prepared = prepare_chat_completions_call(request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
))
.expect("prepares");
assert!(
prepared
.upstream_headers
.contains(&("x-api-key".to_string(), "sk-test".to_string()))
);
assert!(
prepared
.upstream_headers
.contains(&("anthropic-version".to_string(), "2023-06-01".to_string()))
);
assert!(matches!(
prepared.auth,
RequestAuth::Header {
name: "x-api-key",
..
}
));
}
#[test]
fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
// Python builds `{**headers, **anthropic_headers}`, so the deployment's key
// overwrites a forwarded one. Honouring the caller's would let whoever sends
// the request choose the Anthropic principal it bills to.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"X-Api-Key".to_string(),
json!("sk-caller"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
}
#[test]
fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
// Anthropic's `validate_environment` pops `x-api-key` and sets `authorization`
// for an OAuth token, so re-adding the key here would put the credential into
// a header the host removed on purpose.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
(
"Authorization".to_string(),
json!("Bearer sk-ant-oat01-token"),
),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"),
"the resolved key must not be applied over an OAuth bearer, got {:?}",
prepared.upstream_headers
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-token")
);
}
#[test]
fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
// Only an OAuth bearer replaces the credential. Python sends the deployment's
// `x-api-key` alongside any other forwarded `authorization`, so deferring on
// the mere presence of that header would drop the deployment's auth.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer unrelated")),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.collect();
assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers);
assert_eq!(keys[0].1, "sk-test");
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer unrelated"),
"the unrelated authorization must survive, got {:?}",
prepared.upstream_headers
);
}
#[test]
fn declines_an_unsupported_request_before_resolving_credentials() {
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
);
call.api_key = None;
// No api_key is set and no env is consulted: the gate must run first, so the
// error is the decline rather than a missing-credential error.
assert_eq!(decline(call), Error::Unsupported("streaming"));
}
#[test]
fn rejects_an_unknown_provider() {
assert_eq!(
decline(request(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider("openai".to_string())
);
}
#[test]
fn rejects_a_model_with_no_resolvable_provider() {
assert!(matches!(
decline(request(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
)),
Error::InvalidProvider(_)
));
}
#[test]
fn rejects_an_empty_or_malformed_message_list() {
assert_eq!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!([]),
json!({}),
)),
Error::InvalidRequest("chat completions requires at least one message".to_string())
);
assert!(matches!(
decline(request(
"anthropic/claude-sonnet-4-5",
None,
json!("not a list"),
json!({}),
)),
Error::InvalidRequest(_)
));
}
#[test]
fn rejects_non_string_extra_headers() {
let mut call = request(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
assert_eq!(
decline(call),
Error::Headers(litellm_http::request::HeaderError {
context: "chat completions",
name: "x-trace".to_string(),
actual: "number",
})
);
}
#[test]
fn prepares_a_bedrock_call_without_resolving_credentials() {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.api_key = None;
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert_eq!(
prepared.url,
"https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse"
);
assert_eq!(
prepared.auth,
RequestAuth::AwsSigV4 {
region: "us-east-1".to_string(),
service: "bedrock",
}
);
// SigV4 signs the serialized body, so prepare must not have added an
// Authorization header; the handler does it.
assert!(
!prepared
.upstream_headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization"))
);
assert_eq!(prepared.body["inferenceConfig"], json!({"maxTokens": 16}));
}
#[tokio::test]
async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
// Python signs only the AWS header set and reattaches the rest, so a header
// the caller forwarded rides along without joining the canonical request.
// Signing it makes Converse 403 on a deployment that works on Python.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
// A key would resolve to a bearer token and never reach the signer.
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(
"x-request-id".to_string(),
json!("abc-123"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let signed = super::handler::outbound_request(&prepared)
.await
.expect("signs");
let authorization = signed
.header("authorization")
.expect("carries an authorization header")
.to_string();
assert!(
authorization.starts_with("AWS4-HMAC-SHA256"),
"expected a SigV4 signature, got {authorization}"
);
assert!(
!authorization.contains("x-request-id"),
"forwarded header reached SignedHeaders: {authorization}"
);
// It still goes on the wire, it is just not part of the signature.
assert!(
signed
.headers()
.iter()
.any(|(name, value)| name == "x-request-id" && value == "abc-123"),
"forwarded header was dropped instead of reattached"
);
}
#[tokio::test]
async fn a_forwarded_header_the_signer_computes_declines_to_python() {
// Reattaching the caller's copy next to the computed one puts the name on
// the wire twice and Bedrock rejects the pair, so a request carrying one
// has to go to Python instead of being signed here.
for forwarded in [
"Authorization",
"x-amz-date",
"x-amz-security-token",
"Date",
] {
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
);
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let error = super::handler::outbound_request(&prepared)
.await
.expect_err("{forwarded} should decline instead of being signed");
assert!(
matches!(error, Error::Unsupported(_)),
"{forwarded} declined as {error:?}, which the host would not fall back on"
);
}
}
#[test]
fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
// `get_request_headers` assigns `headers["Authorization"]` unconditionally
// once a bearer token resolves, so the deployment's identity wins on
// Python. Keeping the caller's would authorize and bill the call as a
// different principal, and only when the deployment carries `rust: true`.
let mut call = request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.extra_headers = Some(Map::from_iter([(
"Authorization".to_string(),
json!("Bearer caller-supplied"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let authorizations: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("authorization"))
.map(|(_, value)| value.as_str())
.collect();
assert_eq!(
authorizations,
vec!["Bearer sk-test"],
"the deployment token must be the only authorization on the wire"
);
}
#[test]
fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
// The opposite precedence, and deliberate: Anthropic's own transform
// honours a forwarded OAuth bearer, so the Bedrock fix above must not be
// generalized into a rule that the configured key always wins.
//
// An OAuth bearer is the whole of that exception. This forwarded a plain
// `x-api-key` until round 17, which read as the same claim and was not:
// Python overwrites a forwarded `x-api-key` with the deployment's.
let mut call = request(
"claude-sonnet-4-5",
Some("anthropic"),
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
"authorization".to_string(),
json!("Bearer sk-ant-oat01-forwarded"),
)]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let keys: Vec<_> = prepared
.upstream_headers
.iter()
.filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key"))
.map(|(_, value)| value.as_str())
.collect();
assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-ant-oat01-forwarded")
);
}
#[test]
fn a_bedrock_api_key_is_sent_as_a_bearer_token_instead_of_being_signed() {
// The configured bearer identity has its own account and quota boundary,
// so a request carrying one must not be signed as whatever principal the
// host's AWS credentials resolve to.
let prepared = prepare_chat_completions_call(request(
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
))
.expect("prepares");
assert_eq!(
prepared.auth,
RequestAuth::Bearer {
token: "sk-test".to_string()
}
);
assert!(
prepared
.upstream_headers
.iter()
.any(|(name, value)| name.eq_ignore_ascii_case("authorization")
&& value == "Bearer sk-test"),
"prepare did not carry the bearer token"
);
}
fn decline_reason(
model: &str,
provider: Option<&str>,
messages: Value,
params: Value,
) -> Option<&'static str> {
let params = match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
};
super::chat_completions_decline_reason(model, provider, messages, &params)
}
#[test]
fn the_gate_accepts_what_prepare_accepts() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
None
);
}
#[test]
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
),
Some("streaming")
);
assert_eq!(
decline_reason(
"openai/gpt-4o",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"claude-sonnet-4-5",
None,
json!([{"role": "user", "content": "hi"}]),
json!({}),
),
Some("provider is not on the rust chat completions path")
);
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
json!("nope"),
json!({})
),
Some("unreadable message list")
);
assert_eq!(
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
Some("empty message list")
);
}
#[test]
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
// A gate that accepts what prepare then declines would make the host emit
// its pre-call logging on a path that falls back, so pin the agreement.
for (messages, params) in [
(
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 8}),
),
(
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
json!({"temperature": 0.1}),
),
(
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
json!({}),
),
] {
assert_eq!(
decline_reason(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params.clone()
),
None,
"gate declined {messages}"
);
prepare_chat_completions_call(request(
"anthropic/claude-sonnet-4-5",
None,
messages.clone(),
params,
))
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
}
}
mod round_trip {
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::*;
use crate::chat_completions::chat_completions;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
fn http_response(status: &str, body: &str) -> String {
format!(
"HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
)
}
/// Serve one request from a stub upstream and hand back what it received.
async fn serve_once(
status: &'static str,
body: &'static str,
) -> (String, tokio::task::JoinHandle<String>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let port = listener.local_addr().expect("addr").port();
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts");
let received = read_http_request(&mut socket).await;
socket
.write_all(http_response(status, body).as_bytes())
.await
.expect("writes response");
socket.flush().await.expect("flushes");
received
});
(format!("http://127.0.0.1:{port}/v1/messages"), handle)
}
fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> {
ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages,
optional_params: match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: Some(api_base),
custom_llm_provider: None,
extra_headers: None,
timeout: Some(std::time::Duration::from_secs(10)),
}
}
const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
#[tokio::test]
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
let response = chat_completions(call(
&api_base,
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
]),
json!({"max_tokens": 16}),
))
.await
.expect("call succeeds");
let received = handle.await.expect("server task");
let sent: Value = serde_json::from_str(
received
.split_once("\r\n\r\n")
.expect("request has a body")
.1,
)
.expect("body is json");
assert_eq!(
sent["messages"],
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
);
assert_eq!(
sent["system"],
json!([{"type": "text", "text": "be terse"}])
);
assert_eq!(sent["max_tokens"], json!(16));
assert!(received.to_lowercase().contains("x-api-key: sk-test"));
assert_eq!(
response.choices[0].message.content.as_deref(),
Some("hello")
);
assert_eq!(response.usage.total_tokens, 15);
}
#[tokio::test]
async fn a_response_it_cannot_normalize_is_reported_as_already_sent() {
// The provider was called and billed, so the host must not retry this
// on its own path. `MissingField` here would read as a pre-send
// decline and be retried; `InvalidResponse` cannot.
const NO_USAGE: &str =
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
assert!(
matches!(err, Error::InvalidResponse(_)),
"expected a post-send error, got {err:?}"
);
}
#[tokio::test]
async fn an_upstream_error_status_keeps_its_code() {
let (api_base, handle) =
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("upstream rejects");
handle.await.expect("server task");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Http { status: 429, .. })
),
"expected a 429, got {err:?}"
);
}
#[tokio::test]
async fn a_connection_that_is_never_established_declines_instead_of_failing() {
// Nothing was sent, so nothing was billed and the host can still serve
// the request. Classing this with the post-send failures would turn a
// recoverable fallback into a user-facing error on exactly the
// deployments whose transport is configured only on the Python client.
let port = {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
listener.local_addr().expect("has an address").port()
// Dropped here, so the port is closed and the connect is refused.
};
let err = chat_completions(call(
&format!("http://127.0.0.1:{port}/v1/messages"),
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
.await
.expect_err("nothing is listening");
assert!(
matches!(
err,
Error::Transport(litellm_http::transport::Error::Connect(_))
),
"expected a pre-send connect failure, got {err:?}"
);
}
#[test]
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
use crate::chat_completions::handler::as_response_error;
for original in [
Error::MissingField("usage"),
Error::Unsupported("non-text response content block"),
Error::InvalidRequest("whatever".to_string()),
Error::Auth(litellm_auth::Error::InvalidHeader),
] {
let label = format!("{original:?}");
assert!(
matches!(as_response_error(original), Error::InvalidResponse(_)),
"{label} must not stay retryable once the provider has answered"
);
}
// An upstream status is already unambiguous, so it survives intact.
assert!(matches!(
as_response_error(Error::Transport(litellm_http::transport::Error::Http {
status: 500,
body: "boom".to_string()
})),
Error::Transport(litellm_http::transport::Error::Http { status: 500, .. })
));
}
}

View file

@ -26,3 +26,186 @@ pub(super) fn string_headers(
) -> Result<Vec<(String, String)>, Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from)
}
#[cfg(test)]
mod tests {
use std::{sync::Arc, time::Duration};
use futures_util::future::BoxFuture;
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::{TcpListener, TcpStream},
};
use super::{messages_provider_config, string_headers, truncate_error_body};
use crate::messages::{
Error,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
types::MessagesShaping,
};
struct RecordingSecrets {
values: Vec<(&'static str, String)>,
requested: std::sync::Mutex<Vec<String>>,
}
impl SecretSource for RecordingSecrets {
fn get_secret_str<'a>(
&'a self,
name: &'a str,
) -> BoxFuture<'a, Result<Option<SecretValue>, litellm_secrets::Error>> {
Box::pin(async move {
self.requested.lock().unwrap().push(name.to_string());
Ok(self
.values
.iter()
.find(|(key, _)| *key == name)
.map(|(_, value)| SecretValue::new(value.clone())))
})
}
}
fn secrets_call() -> MessagesCall {
let Value::Object(body) = json!({
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hi"}]
}) else {
unreachable!("literal object")
};
MessagesCall {
model: "claude-sonnet-4-5".into(),
body,
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic".into()),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),
shaping: MessagesShaping::default(),
}
}
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
#[tokio::test]
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let request = read_http_request(&mut socket).await;
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let secrets = Arc::new(RecordingSecrets {
values: vec![
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
],
requested: std::sync::Mutex::new(Vec::new()),
});
let output = litellm_host::run::run(
messages_machine(secrets.clone()),
&LocalMessagesHost::new(secrets_call()),
)
.await
.expect("messages request succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
let request = server.await.expect("server task completes");
assert!(
request
.to_ascii_lowercase()
.contains("x-api-key: sk-from-manager"),
"{request}"
);
let requested = secrets.requested.lock().unwrap().clone();
assert_eq!(
requested,
messages_provider_config("anthropic")
.unwrap()
.secret_names()
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
);
}
#[test]
fn provider_config_resolves_anthropic_and_azure_ai() {
assert!(messages_provider_config("anthropic").is_some());
assert!(messages_provider_config("azure_ai").is_some());
assert!(messages_provider_config("openai").is_none());
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(400);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
Error::Headers(litellm_http::request::HeaderError {
context: "messages",
name: "x-count".to_string(),
actual: "number",
})
);
}
}

View file

@ -46,6 +46,3 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
)),
}
}
#[cfg(test)]
mod tests;

View file

@ -248,3 +248,160 @@ mod tests {
}
}
}
#[cfg(test)]
mod document_tests {
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::route::LocalOcrHost;
use crate::ocr::test_support::{
MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with,
request_body, wire_request_with_document,
};
#[derive(Clone, Copy, Debug)]
enum Route {
Mistral,
AzureAi,
VertexMistral,
AzureCohereParse,
Cohere,
}
impl Route {
fn model(self) -> &'static str {
match self {
Self::Mistral => "mistral/model",
Self::AzureAi => "azure_ai/model",
Self::VertexMistral => "vertex_ai/mistral-ocr-maas",
Self::AzureCohereParse => "azure_ai/cohere-parse",
Self::Cohere => "cohere/model",
}
}
fn document_type(self) -> &'static str {
match self {
Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url",
Self::AzureCohereParse | Self::Cohere => "image_url",
}
}
fn options(self) -> Value {
match self {
Self::Mistral | Self::AzureAi => json!({"pages": [0]}),
Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}),
Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}),
}
}
}
/// What the host does to the wire request in `before_send`.
#[derive(Clone, Copy, Debug)]
enum Host {
Detached,
ReplacesDocument,
}
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
impl Host {
fn before_send(self, wire: WireRequest) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
let body = fields
.into_iter()
.map(|(name, value)| match self {
Self::Detached => (name, value),
Self::ReplacesDocument if name == "document" => {
let document_type = value["type"].clone();
let key = document_type.as_str().unwrap_or_default().to_string();
(name, json!({"type": document_type, key: REPLACED_DOCUMENT}))
}
Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
struct Sent {
result: Result<(), Error>,
provider_body: Option<Value>,
}
async fn send(route: Route, host: Host, document_base: &str) -> Sent {
let (base, seen, provider) =
mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let document_type = route.document_type();
let document =
json!({"type": document_type, document_type: format!("{document_base}/scan.png")});
let request = wire_request_with_document(route.model(), &base, document, route.options());
let local =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire)));
let result = perform_ocr_with(local).await.map(|_| ());
match result {
Ok(()) => provider.await.unwrap(),
Err(_) => provider.abort(),
}
let provider_body = seen
.lock()
.unwrap()
.first()
.map(|request| request_body(request));
Sent {
result,
provider_body,
}
}
fn served_document_uri() -> String {
use base64::Engine;
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
}
#[rstest]
#[case::azure_ai(Route::AzureAi)]
#[case::vertex_mistral(Route::VertexMistral)]
#[case::azure_cohere_parse(Route::AzureCohereParse)]
#[tokio::test]
async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::Detached, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(served_document_uri())
);
}
#[rstest]
#[tokio::test]
async fn document_replaced_by_the_host_reaches_the_provider(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::ReplacesDocument, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(REPLACED_DOCUMENT)
);
}
}

View file

@ -9,36 +9,210 @@ pub mod types;
pub mod wire;
#[cfg(test)]
#[path = "../../tests/aws_textract_ocr.rs"]
mod aws_textract_tests;
pub(crate) mod test_support {
use std::sync::{Arc, Mutex};
#[cfg(test)]
#[path = "../../tests/azure_ai_ocr.rs"]
mod azure_ai_tests;
#[cfg(test)]
#[path = "../../tests/azure_document_intelligence_ocr.rs"]
mod azure_document_intelligence_tests;
#[cfg(test)]
#[path = "../../tests/cohere_ocr.rs"]
mod cohere_tests;
#[cfg(test)]
#[path = "../../tests/deepseek_ocr.rs"]
mod deepseek_tests;
#[cfg(test)]
#[path = "../../tests/ocr/document.rs"]
mod document_tests;
#[cfg(test)]
#[path = "../../tests/reducto_ocr.rs"]
mod reducto_tests;
#[cfg(test)]
#[path = "../../tests/ocr/support.rs"]
pub(crate) mod test_support;
#[cfg(test)]
#[path = "../../tests/ocr.rs"]
pub(crate) mod tests;
#[cfg(test)]
#[path = "../../tests/vertex_ai_deepseek_ocr.rs"]
mod vertex_ai_deepseek_tests;
#[cfg(test)]
#[path = "../../tests/vertex_ai_ocr.rs"]
mod vertex_ai_tests;
use futures_util::future::BoxFuture;
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient},
transformation::LiteLLMOcrResponse,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
use crate::ocr::{
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
/// Stands in for a host with no hooks registered: the wire request goes out unchanged
/// and response events go nowhere.
pub(crate) struct NoHooks;
impl CallHooks<Error> for NoHooks {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(async move { Ok(wire) })
}
fn response_received<'a>(&'a self, _body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(async { Ok(()) })
}
}
pub(crate) fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test document client builds");
OcrClient::for_test(reqwest::Client::new(), document_http)
}
pub(crate) async fn perform_ocr(
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
crate::ocr::client::perform(&ocr_client(), request).await
}
pub(crate) async fn perform_ocr_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
wire_request_with_document(
model,
base,
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
options,
)
}
pub(crate) fn wire_request_with_document(
model: &str,
base: &str,
document: Value,
options: Value,
) -> LiteLLMOcrRequest {
decode_request(OcrWireRequest {
model: model.into(),
document,
api_key: Some(litellm_auth::SecretValue::new("test-key")),
api_base: Some(base.into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
})
.unwrap()
}
pub(crate) fn resolved_request(
request: LiteLLMOcrRequest,
) -> crate::ocr::types::ResolvedOcrRequest {
request
.map_document(crate::ocr::document::prepare_document)
.unwrap()
}
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let request = resolved_request(request);
let document = request.document.clone().with_source(source.into());
request.with_document(document.into())
}
pub(crate) fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
pub(crate) const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document";
/// Serves [`SERVED_DOCUMENT`] as `image/png` to every connection until aborted.
pub(crate) async fn document_server() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let task = tokio::spawn(async move {
loop {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buffer = [0u8; 4096];
let _ = socket.read(&mut buffer).await.unwrap();
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
SERVED_DOCUMENT.len()
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(SERVED_DOCUMENT).await.unwrap();
}
});
(base, task)
}
pub(crate) struct MockResponse {
pub status: u16,
pub headers: Vec<(&'static str, String)>,
pub body: Value,
}
impl MockResponse {
pub fn json(body: Value) -> Self {
Self {
status: 200,
headers: vec![],
body,
}
}
}
pub(crate) async fn mock_server(
responses: Vec<MockResponse>,
) -> (String, Arc<Mutex<Vec<String>>>, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let requests = Arc::new(Mutex::new(Vec::new()));
let seen = requests.clone();
let server_base = base.clone();
let task = tokio::spawn(async move {
for response in responses {
let (mut socket, _) = listener.accept().await.unwrap();
let mut bytes = Vec::new();
let mut buffer = [0u8; 4096];
let header_end = loop {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") {
break index + 4;
}
};
let length = String::from_utf8_lossy(&bytes[..header_end])
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().unwrap())
})
.unwrap_or(0);
while bytes.len() < header_end + length {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
}
seen.lock()
.unwrap()
.push(String::from_utf8_lossy(&bytes).into_owned());
let body = serde_json::to_vec(&response.body).unwrap();
let headers = response
.headers
.into_iter()
.map(|(name, value)| {
format!("{name}: {}\r\n", value.replace("{base}", &server_base))
})
.collect::<String>();
let head = format!(
"HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n",
response.status,
body.len(),
headers
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(&body).await.unwrap();
}
});
(base, requests, task)
}
pub(crate) fn header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
request
.lines()
.take_while(|line| !line.is_empty())
.find_map(|line| {
let (key, value) = line.split_once(':')?;
key.eq_ignore_ascii_case(name).then(|| value.trim())
})
}
}

File diff suppressed because it is too large Load diff

View file

@ -4,11 +4,9 @@ use std::{
thread,
};
use litellm_core::audio_transcription::{audio_transcription, types::AudioTranscriptionRequest};
use serde_json::{Map, json};
use super::audio_transcription;
use crate::audio_transcription::types::AudioTranscriptionRequest;
#[tokio::test]
async fn bedrock_request_is_signed_and_contains_audio() {
let listener = TcpListener::bind("127.0.0.1:0").expect("listener");

View file

@ -1,193 +0,0 @@
use std::{collections::BTreeMap, time::SystemTime};
use litellm_auth_aws::{Credentials, aws_signature_headers, sign_post};
use litellm_llms::base_llm::ocr::error::Error;
use serde_json::{Value, json};
use time::{PrimitiveDateTime, format_description};
use crate::ocr::{
route::LocalOcrHost,
test_support::{
MockResponse, header, mock_server, perform_ocr_with, request_body,
wire_request_with_document,
},
types::LiteLLMOcrRequest,
};
const ACCESS_KEY_ID: &str = "AKIDEXAMPLE";
const SECRET_ACCESS_KEY: &str = "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY";
fn textract_request(base: &str) -> LiteLLMOcrRequest {
textract_request_for("aws_textract/detect-document-text", base)
}
fn textract_request_for(model: &str, base: &str) -> LiteLLMOcrRequest {
wire_request_with_document(
model,
&format!("{base}/"),
json!({"type": "image_url", "image_url": "data:image/png;base64,b3JpZ2luYWw="}),
json!({
"aws_access_key_id": ACCESS_KEY_ID,
"aws_secret_access_key": SECRET_ACCESS_KEY,
"aws_region_name": "eu-west-1"
}),
)
}
fn textract_response() -> MockResponse {
MockResponse::json(json!({
"DocumentMetadata": {"Pages": 1},
"Blocks": [{"BlockType": "PAGE"}, {"BlockType": "LINE", "Text": "Invoice 12345"}]
}))
}
/// Recomputes SigV4 over the bytes the server received, at the time the client claimed.
fn expected_authorization(url: &str, raw_request: &str) -> String {
let format =
format_description::parse_borrowed::<2>("[year][month][day]T[hour][minute][second]Z")
.unwrap();
let signed_at: SystemTime =
PrimitiveDateTime::parse(header(raw_request, "x-amz-date").unwrap(), &format)
.unwrap()
.assume_utc()
.into();
let headers: BTreeMap<String, String> = ["content-type", "x-amz-target"]
.into_iter()
.map(|name| {
(
name.to_string(),
header(raw_request, name).unwrap().to_string(),
)
})
.collect();
let body = raw_request.split_once("\r\n\r\n").unwrap().1;
sign_post(
url,
body.as_bytes(),
&aws_signature_headers(&headers),
"eu-west-1",
"textract",
&Credentials::new(ACCESS_KEY_ID, SECRET_ACCESS_KEY, None, None, "test"),
signed_at,
)
.unwrap()["Authorization"]
.clone()
}
#[tokio::test]
async fn the_request_is_signed_for_textract_and_lines_become_the_page() {
let (base, seen, server) = mock_server(vec![textract_response()]).await;
let response = perform_ocr_with(LocalOcrHost::new(textract_request(&base)))
.await
.unwrap();
server.await.unwrap();
let raw = seen.lock().unwrap()[0].clone();
assert_eq!(
header(&raw, "x-amz-target"),
Some("Textract.DetectDocumentText")
);
assert_eq!(
header(&raw, "content-type"),
Some("application/x-amz-json-1.1")
);
assert_eq!(
request_body(&raw),
json!({"Document": {"Bytes": "b3JpZ2luYWw="}})
);
assert_eq!(
header(&raw, "authorization"),
Some(expected_authorization(&format!("{base}/"), &raw).as_str())
);
assert_eq!(response.pages[0].markdown, "Invoice 12345");
assert_eq!(response.usage_info.unwrap().pages_processed, Some(1));
}
#[tokio::test]
async fn a_body_rewritten_by_before_send_is_what_gets_signed_and_sent() {
let (base, seen, server) = mock_server(vec![textract_response()]).await;
let host = LocalOcrHost::new(textract_request(&base)).with_before_send(|mut wire, _| {
assert!(
!wire
.headers
.iter()
.any(|(name, _)| name.eq_ignore_ascii_case("authorization")),
"the hook ran after signing"
);
wire.body["Document"]["Bytes"] = Value::from("cmVkYWN0ZWQ=");
Ok(wire)
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let raw = seen.lock().unwrap()[0].clone();
assert_eq!(
request_body(&raw),
json!({"Document": {"Bytes": "cmVkYWN0ZWQ="}})
);
assert_eq!(
header(&raw, "authorization"),
Some(expected_authorization(&format!("{base}/"), &raw).as_str())
);
}
#[tokio::test]
async fn a_multi_page_rejection_reaches_the_caller_with_the_single_page_limit() {
let (base, _, server) = mock_server(vec![MockResponse {
status: 400,
headers: vec![],
body: json!({
"__type": "UnsupportedDocumentException",
"Message": "Request has unsupported document format"
}),
}])
.await;
let error = perform_ocr_with(LocalOcrHost::new(textract_request(&base)))
.await
.unwrap_err();
server.await.unwrap();
let Error::Provider { status, body, .. } = error else {
panic!("expected a provider error, got {error:?}");
};
assert_eq!(status, 400);
assert!(
body.contains("multi-page documents are not supported"),
"{body}"
);
}
#[tokio::test]
async fn analyze_document_asks_for_layout_and_tables_and_returns_markdown() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"DocumentMetadata": {"Pages": 1},
"Blocks": [
{"Id": "l1", "BlockType": "LINE", "Text": "Quarterly Report"},
{"Id": "t", "BlockType": "LAYOUT_TITLE",
"Relationships": [{"Type": "CHILD", "Ids": ["l1"]}]}
]
}))])
.await;
let request = textract_request_for("aws_textract/analyze-document", &base);
let response = perform_ocr_with(LocalOcrHost::new(request)).await.unwrap();
server.await.unwrap();
let raw = seen.lock().unwrap()[0].clone();
assert_eq!(
header(&raw, "x-amz-target"),
Some("Textract.AnalyzeDocument")
);
assert_eq!(
request_body(&raw)["FeatureTypes"],
json!(["LAYOUT", "TABLES"])
);
assert_eq!(
header(&raw, "authorization"),
Some(expected_authorization(&format!("{base}/"), &raw).as_str())
);
assert_eq!(response.pages[0].markdown, "# Quarterly Report");
}

View file

@ -1,293 +0,0 @@
use litellm_llms::base_llm::ocr::error::Error;
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request};
use crate::ocr::route::LocalOcrHost;
#[tokio::test]
async fn facade_executes_azure_mistral_with_prepared_auth() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"include_image_base64":true}),
);
request.credentials.api_key = None;
request.transport.extra_headers = vec![(
"Authorization".into(),
"Bearer python-prepared-token".into(),
)];
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer python-prepared-token\r\n")
);
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"include_image_base64":true
})
);
}
#[tokio::test]
async fn facade_acquires_supplied_entra_token_for_final_request() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"azure_ad_token":"rust-owned-token"}),
);
request.credentials.api_key = None;
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer rust-owned-token\r\n")
);
}
#[tokio::test]
async fn rejects_non_inline_body_after_guardrails() {
let request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({}));
let host = LocalOcrHost::new(request).with_before_send(|mut wire, _| {
wire.body["document"] = json!({
"type":"document_url",
"document_url":"https://example.com/not-inline.pdf"
});
Ok(wire)
});
let error = perform_ocr_with(host).await.unwrap_err();
assert!(error.to_string().contains("data URI"));
}
mod transformation {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};
use litellm_auth::{
ResolvedCredential, SecretValue, TokenFuture, TokenProvider, TokenProviderHandle,
};
use rstest::rstest;
use serde_json::json;
use super::*;
use crate::ocr::{
test_support::{MockResponse, header, mock_server, perform_ocr},
types::LiteLLMOcrRequest,
wire::decode_request,
};
#[derive(Debug)]
struct CountingToken {
token: fn(usize) -> String,
calls: AtomicUsize,
}
impl CountingToken {
fn new(token: fn(usize) -> String) -> Arc<Self> {
Arc::new(Self {
token,
calls: AtomicUsize::new(0),
})
}
fn calls(&self) -> usize {
self.calls.load(Ordering::SeqCst)
}
}
impl TokenProvider for CountingToken {
fn acquire(&self) -> TokenFuture<'_> {
let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
let token = SecretValue::new((self.token)(call));
Box::pin(async move {
Ok(ResolvedCredential::AccessToken {
token,
expires_on: None,
})
})
}
}
fn numbered_token(call: usize) -> String {
format!("callback-{call}")
}
fn azure_request(
provider: &Arc<CountingToken>,
api_base: Option<&str>,
api_key: Option<&str>,
extra_headers: Value,
optional_params: Value,
) -> LiteLLMOcrRequest {
let wire = serde_json::from_value(json!({
"model": "azure_ai/mistral-ocr-latest",
"document": {"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": null,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": 2.0
}))
.unwrap();
LiteLLMOcrRequest {
azure_ad_token_provider: Some(TokenProviderHandle::new(provider.clone())),
..decode_request(wire).unwrap()
}
}
fn ocr_page() -> MockResponse {
MockResponse::json(json!({"pages":[{"index":0,"markdown":"hello"}]}))
}
#[tokio::test]
async fn token_provider_result_is_the_bearer_and_is_acquired_for_each_request() {
let provider = CountingToken::new(numbered_token);
let (base, seen, server) = mock_server(vec![ocr_page(), ocr_page()]).await;
for _ in 0..2 {
perform_ocr(azure_request(
&provider,
Some(&base),
None,
Value::Null,
json!({}),
))
.await
.unwrap();
}
server.await.unwrap();
assert_eq!(provider.calls(), 2);
let requests = seen.lock().unwrap();
assert_eq!(
requests
.iter()
.map(|request| header(request, "authorization"))
.collect::<Vec<_>>(),
[Some("Bearer callback-1"), Some("Bearer callback-2")]
);
}
#[rstest]
#[case::api_key_skips_provider(Some("resource-key"), Value::Null, json!({}), "Bearer resource-key", 0)]
#[case::provider_beats_static_token(
None,
Value::Null,
json!({"azure_ad_token":"static-token"}),
"Bearer callback-1",
1
)]
#[case::header_wins_on_the_wire_but_provider_still_runs(
None,
json!({"Authorization":"Bearer override"}),
json!({}),
"Bearer override",
1
)]
#[tokio::test]
async fn credential_precedence(
#[case] api_key: Option<&str>,
#[case] extra_headers: Value,
#[case] optional_params: Value,
#[case] expected_authorization: &str,
#[case] expected_calls: usize,
) {
let provider = CountingToken::new(numbered_token);
let (base, seen, server) = mock_server(vec![ocr_page()]).await;
perform_ocr(azure_request(
&provider,
Some(&base),
api_key,
extra_headers,
optional_params,
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(provider.calls(), expected_calls);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert_eq!(
header(&requests[0], "authorization"),
Some(expected_authorization)
);
}
#[rstest]
#[case::missing_api_base(
false,
json!({}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: "AZURE_AI_API_BASE",
})),
0
)]
#[case::unsupported_oidc_reference(
true,
json!({"azure_ad_token":"oidc/assertion","client_id":"client","tenant_id":"tenant"}),
numbered_token,
|error: &Error| matches!(error, Error::Auth(litellm_auth::Error::UnsupportedOidcReference)),
0
)]
#[case::empty_provider_token_ignores_static_token(
true,
json!({"azure_ad_token":"static-token"}),
|_| String::new(),
|error: &Error| matches!(error, Error::MissingAzureAiCredentials),
1
)]
#[tokio::test]
async fn credential_failures_send_no_provider_request(
#[case] with_api_base: bool,
#[case] optional_params: Value,
#[case] token: fn(usize) -> String,
#[case] expected: fn(&Error) -> bool,
#[case] expected_calls: usize,
) {
let provider = CountingToken::new(token);
let (base, seen, server) = mock_server(vec![ocr_page()]).await;
let error = perform_ocr(azure_request(
&provider,
with_api_base.then_some(base.as_str()),
None,
Value::Null,
optional_params,
))
.await
.unwrap_err();
server.abort();
assert!(expected(&error), "unexpected error: {error:?}");
assert_eq!(provider.calls(), expected_calls);
assert!(seen.lock().unwrap().is_empty());
}
}

View file

@ -1,712 +0,0 @@
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::{error::Error, settings::OcrSettings};
use rstest::rstest;
use serde_json::{Value, json};
use super::{
test_support::{
MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request,
},
wire::{OcrWireRequest, decode_request},
};
use crate::ocr::route::LocalOcrHost;
fn query_value(url: &str, key: &str) -> Option<String> {
url::Url::parse(url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == key).then(|| value.into_owned()))
}
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}),
);
request.document =
serde_json::from_value::<litellm_llms::base_llm::ocr::transformation::OcrDocument>(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap()
.into();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf"})
);
}
#[rstest]
#[case(json!({"pages":[true]}), Error::Pages("expected only integers or only strings".into()))]
#[case(json!({"pages":[1,"2"]}), Error::Pages("expected only integers or only strings".into()))]
#[case(json!({"pages":[-1]}), Error::Pages("negative page index".into()))]
#[case(json!({"pages":"1&&features=bad"}), Error::Pages("invalid native page range".into()))]
#[case(json!({"features":"languages&pages=1"}), Error::Features)]
#[case(json!({"req_format":"azure"}), Error::RequestFormat)]
#[tokio::test]
async fn rejects_invalid_pages_features_and_format(
#[case] options: Value,
#[case] expected: Error,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({}))]).await;
let result = decode_request(OcrWireRequest {
model: "azure_ai/doc-intelligence/prebuilt-read".into(),
document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
api_key: Some(litellm_auth::SecretValue::new("key")),
api_base: Some(base),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
});
let result = match result {
Ok(request) => perform_ocr(request).await,
Err(error) => Err(error),
};
server.abort();
let _ = server.await;
assert!(
seen.lock().unwrap().is_empty(),
"sent invalid options: {options}"
);
let error = result.unwrap_err();
assert_eq!(
std::mem::discriminant(&error),
std::mem::discriminant(&expected)
);
assert_eq!(error.http_status_code(), Some(400));
assert_eq!(error.to_string(), expected.to_string());
}
#[rstest]
#[case(json!({}))]
#[case(json!({"req_format":"litellm"}))]
#[tokio::test]
async fn missing_native_fields_keep_page_text_without_retaining_raw_response(
#[case] options: Value,
) {
let operation = json!({
"status":"succeeded",
"analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"hello"}]}]}
});
let (base, seen, server) = mock_server(vec![MockResponse::json(operation)]).await;
let response = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
options,
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(response.pages.len(), 1);
assert_eq!(response.pages[0].index, 0);
assert_eq!(response.pages[0].markdown, "hello");
assert_eq!(response.provider_native_response, None);
let serialized = response.into_json();
assert_eq!(serialized.get("content"), Some(&Value::Null));
assert_eq!(serialized.get("tables"), Some(&Value::Null));
assert_eq!(serialized.get("keyValuePairs"), Some(&Value::Null));
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
let target = requests[0].split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
for field in ["pages", "features", "req_format"] {
assert_eq!(query_value(&url, field), None);
}
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn inline_document_decodes_to_base64_source() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
}
#[tokio::test]
async fn client_settings_choose_the_api_version_and_the_inch_to_pixel_dpi() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[{"pageNumber":1,"width":8.5,"height":11,"unit":"inch"}]}
}))])
.await;
let client = ocr_client().with_settings(OcrSettings {
document_intelligence_api_version: "2099-01-01".into(),
document_intelligence_dpi: 72,
..OcrSettings::default()
});
let result = crate::ocr::client::perform(
&client,
wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})),
)
.await
.unwrap();
server.await.unwrap();
let target = seen.lock().unwrap()[0]
.split_whitespace()
.nth(1)
.unwrap()
.to_string();
assert_eq!(
query_value(&format!("{base}{target}"), "api-version").as_deref(),
Some("2099-01-01")
);
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":612,"height":792,"dpi":72})
);
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.transport
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
#[tokio::test]
async fn accepted_response_emits_response_received_before_polling() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.with_observer(move |event| {
let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event else {
return;
};
match request_count.lock().unwrap().len() {
1 => assert_eq!(raw.body, r#"{"submitted":true}"#),
2 => assert!(raw.body.contains("succeeded")),
count => panic!("unexpected callback after {count} requests"),
}
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.credentials.api_key = None;
request.transport.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert!(
requests[1]
.to_ascii_lowercase()
.contains("authorization: bearer token")
);
}
#[tokio::test]
async fn polling_does_not_follow_redirects() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 302,
headers: vec![("Location", "{base}/redirected".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
assert!(error.to_string().contains("status 302"), "{error}");
assert_eq!(seen.lock().unwrap().len(), 2);
server.abort();
}
#[tokio::test]
async fn polling_rejects_terminal_failure() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"failed"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("status failed"));
}
#[tokio::test]
async fn malformed_provider_pages_report_response_paths() {
for (analysis, path) in [
(json!({"pages":null}), "pages"),
(json!({"pages":[null]}), "pages[0]"),
(json!({"pages":[{"lines":null}]}), "lines"),
(json!({"pages":[{"width":"bad"}]}), "width"),
] {
let (base, _, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":analysis
}))])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains(path), "{error}");
}
}
#[tokio::test]
async fn rejects_missing_invalid_and_cross_origin_operation_locations() {
for headers in [
Vec::new(),
vec![("Operation-Location", "/relative".into())],
vec![("Operation-Location", "http://example.com/operation".into())],
vec![(
"Operation-Location",
"http://user:password@127.0.0.1/operation".into(),
)],
] {
let (base, _, server) = mock_server(vec![MockResponse {
status: 202,
headers,
body: json!({}),
}])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("operation-location"));
}
}
#[tokio::test]
async fn polling_deadline_bounds_retry_delay() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "9999".into())],
body: json!({"status":"notStarted"}),
},
])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
let client = ocr_client().with_settings(OcrSettings {
poll_timeout: std::time::Duration::from_millis(100),
..OcrSettings::default()
});
let error = tokio::time::timeout(
std::time::Duration::from_secs(1),
crate::ocr::client::perform(&client, request),
)
.await
.unwrap()
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("timed out"));
}
#[tokio::test]
async fn model_id_is_encoded_and_dot_segments_are_rejected() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
perform_ocr(wire_request(
"azure_ai/doc-intelligence/a ?#é",
&base,
json!({}),
))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze"));
for model in [
"azure_ai/doc-intelligence/.",
"azure_ai/doc-intelligence/..",
] {
let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("dot segment"));
}
}
mod transformation {
use std::sync::{Arc, Mutex};
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::transformation::OcrDocument;
use serde_json::{Value, json};
use super::*;
use crate::ocr::{
route::LocalOcrHost,
test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request},
};
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"], "future_option": {"nested":null}, "extra_body":{"provider_option":false}}),
);
request.document = serde_json::from_value::<OcrDocument>(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap()
.into();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf", "future_option":{"nested":null}, "provider_option":false})
);
}
#[tokio::test]
async fn rejects_invalid_pages_features_and_format() {
for options in [
json!({"pages":[true]}),
json!({"pages":[1,"2"]}),
json!({"pages":[-1]}),
json!({"pages":"1&&features=bad"}),
json!({"features":"languages&pages=1"}),
json!({"req_format":"azure"}),
] {
let request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
"http://127.0.0.1:1",
options.clone(),
);
let rejected = perform_ocr(request).await.is_err();
assert!(rejected, "accepted {options}");
}
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.transport
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
#[tokio::test]
async fn accepted_response_emits_response_received_for_submission_and_completed_poll() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let responses_received = Arc::new(Mutex::new(Vec::new()));
let request_count = seen.clone();
let observed = responses_received.clone();
let host = LocalOcrHost::new(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
observed
.lock()
.unwrap()
.push((request_count.lock().unwrap().len(), raw.body.clone()));
}
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
assert_eq!(
*responses_received.lock().unwrap(),
[
(1, r#"{"submitted":true}"#.to_string()),
(2, r#"{"status":"succeeded"}"#.to_string()),
]
);
}
}

View file

@ -1,136 +0,0 @@
mod transformation {
use litellm_llms::{
base_llm::ocr::{
error::Error,
transformation::{BaseOcrConfig, OcrDocument, OcrResponseFormat},
},
cohere::ocr::transformation::*,
};
use rstest::rstest;
use serde_json::{Value, json};
#[tokio::test]
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({
"output_format":"markdown", "timeout":30,
"extra_body":{
"output_format": {"future":true},
"document":{"type":"image_url","image_url":"https://example.com/a.png",
"provider_options":{"nested":[false,0,null]}}
}
}),
);
let request = request.with_document(
serde_json::from_value(json!({
"type":"image_url","image_url":"https://example.com/original.png"
}))
.unwrap(),
);
let request = crate::ocr::prepare::prepare_request_for_test(request);
let http = CohereParseConfig
.prepare_request(
&request,
&crate::ocr::test_support::ocr_client(),
&crate::ocr::test_support::NoHooks,
)
.await
.unwrap();
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(
body,
json!({
"model":"parse", "output_format":{"future":true},
"document":{"type":"image_url","image_url":"https://example.com/a.png",
"provider_options":{"nested":[false,0,null]}}
})
);
}
#[tokio::test]
async fn explicit_null_options_use_defaults_before_http() {
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({"output_format":null,"req_format":null}),
);
let request = request.with_document(
serde_json::from_value(
json!({"type":"image_url","image_url":"https://example.com/a.png"}),
)
.unwrap(),
);
assert_eq!(
request.response_format().unwrap(),
OcrResponseFormat::Litellm
);
let request = crate::ocr::prepare::prepare_request_for_test(request);
let http = CohereParseConfig
.prepare_request(
&request,
&crate::ocr::test_support::ocr_client(),
&crate::ocr::test_support::NoHooks,
)
.await
.unwrap();
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(body["output_format"], "markdown");
assert!(body.get("req_format").is_none());
}
#[rstest]
#[case::cohere("cohere/parse-v5.0", "POST /v2/parse ")]
#[case::azure_ai("azure_ai/Cohere-parse-v5.0", "POST /providers/cohere/v2/parse ")]
#[tokio::test]
async fn route_sends_image_to_its_parse_endpoint_with_the_bearer_key(
#[case] model: &str,
#[case] request_line: &str,
) {
use crate::ocr::test_support::{MockResponse, header, mock_server, perform_ocr};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = crate::ocr::test_support::wire_request(model, &base, json!({}))
.with_document(
serde_json::from_value::<OcrDocument>(
json!({"type":"image_url","image_url":"data:image/png;base64,YWJj"}),
)
.unwrap()
.into(),
);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(request_line), "{}", requests[0]);
assert_eq!(
header(&requests[0], "authorization"),
Some("Bearer test-key")
);
}
#[rstest]
#[tokio::test]
async fn route_rejects_non_image_document_without_a_request(
#[values("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")] model: &str,
) {
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let error = perform_ocr(crate::ocr::test_support::wire_request(
model,
&base,
json!({}),
))
.await
.unwrap_err();
server.abort();
assert!(matches!(error, Error::CohereImageOnly), "{error:?}");
assert!(seen.lock().unwrap().is_empty());
}
}

View file

@ -1,133 +0,0 @@
use litellm_llms::{
base_llm::ocr::transformation::{BaseOcrConfig, OcrDocument},
vertex_ai::ocr::deepseek_transformation::{
DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig,
normalize_response as transform_ocr_response,
},
};
use rstest::rstest;
use serde_json::{Value, json};
fn document() -> OcrDocument {
serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap()
}
#[rstest]
#[case("stream", json!(true))]
#[case("temperature", json!(0.1))]
#[case("max_tokens", json!(1024))]
#[case("top_p", json!(0.9))]
#[case("n", json!(2))]
#[case("stop", json!("done"))]
#[case("stop", json!(["done", "stop"]))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
VertexAIDeepSeekOCRConfig
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/a.png"})
);
assert_eq!(result[name], value);
assert!(result.get("ignored").is_none());
}
#[rstest]
#[case(json!({"type":"image_url","image_url":"data:image/png;base64,AA=="}))]
#[case(json!({"type":"document_url","document_url":"data:application/pdf;base64,AA=="}))]
fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
let source = document
.get("image_url")
.or_else(|| document.get("document_url"))
.unwrap()
.clone();
let request = VertexAIDeepSeekOCRConfig
.transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
&[],
)
.unwrap();
let result = serde_json::to_value(request).unwrap();
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":source})
);
}
#[rstest]
#[case(json!("# hello"), "# hello")]
#[case(json!("{broken"), "{broken")]
#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")]
#[case(json!({"pages":[]}), "")]
#[case(json!("[]"), "[]")]
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case] expected: &str) {
let structured = content
.as_object()
.is_some_and(|object| object.contains_key("pages"))
|| content
.as_str()
.is_some_and(|text| text.contains("\"pages\""));
let response: DeepSeekOcrResponse = serde_json::from_value(
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
)
.unwrap();
let result = transform_ocr_response("model", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["markdown"], expected);
assert_eq!(result["pages"][0]["index"], 0);
if structured {
assert!(result["usage_info"].is_null());
} else {
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
}
}
#[test]
fn structured_result_maps_pages_usage_model_and_annotation() {
let response: DeepSeekOcrResponse = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}],
"model":"provider-model",
"usage_info":{"pages_processed":1},
"document_annotation":{"language":"en"},
"future":"kept"
}}}]
}))
.unwrap();
let result = transform_ocr_response("requested", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["index"], 2);
assert_eq!(result["pages"][0]["images"][0]["id"], "one");
assert_eq!(result["model"], "provider-model");
assert_eq!(result["usage_info"]["pages_processed"], 1);
assert_eq!(result["document_annotation"]["language"], "en");
assert_eq!(result["future"], "kept");
}
#[test]
fn response_codec_rejects_missing_empty_and_malformed_content() {
for value in [
json!({"choices":[{"message":{"content":{}}}]}),
json!({"choices":[]}),
json!({"choices":[{"message":{"content":""}}]}),
json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}),
json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}),
] {
let result = serde_json::from_value::<DeepSeekOcrResponse>(value)
.map_err(|_| ())
.and_then(|response| transform_ocr_response("model", response).map_err(|_| ()));
assert!(result.is_err());
}
}

View file

@ -1,7 +1,11 @@
use std::{sync::Arc, time::Duration};
use futures_util::future::BoxFuture;
use litellm_http::request::{has_bearer_auth, has_header};
use litellm_core::messages::{
Error, messages,
route::{LocalMessagesHost, MessagesCall, messages_machine},
types::{MessagesRequest, MessagesShaping},
};
use litellm_secrets::{SecretValue, source::SecretSource};
use serde_json::{Map, Value, json};
use tokio::{
@ -9,14 +13,6 @@ use tokio::{
net::{TcpListener, TcpStream},
};
use super::{
Error,
common_utils::{messages_provider_config, string_headers, truncate_error_body},
messages,
route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine},
};
use crate::messages::types::{MessagesRequest, MessagesShaping};
struct RecordingSecrets {
values: Vec<(&'static str, String)>,
fails: bool,
@ -73,55 +69,6 @@ fn secrets_call() -> MessagesCall {
}
}
#[tokio::test]
async fn route_reads_the_provider_credential_and_base_from_the_secret_source() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let addr = listener.local_addr().expect("addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let request = read_http_request(&mut socket).await;
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#;
socket
.write_all(write_response(response_body).as_bytes())
.await
.expect("writes response");
request
});
let secrets = Arc::new(RecordingSecrets::new(
vec![
("ANTHROPIC_API_KEY", "sk-from-manager".to_string()),
("ANTHROPIC_BASE_URL", format!("http://{addr}")),
],
false,
));
let output = litellm_host::run::run(
messages_machine(secrets.clone()),
&LocalMessagesHost::new(secrets_call()),
)
.await
.expect("messages request succeeds");
assert!(matches!(output, MessagesOutput::Message(_)));
let request = server.await.expect("server task completes");
assert!(
request
.to_ascii_lowercase()
.contains("x-api-key: sk-from-manager"),
"{request}"
);
let requested = secrets.requested.lock().unwrap().clone();
assert_eq!(
requested,
messages_provider_config("anthropic")
.unwrap()
.secret_names()
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
);
}
#[tokio::test]
async fn route_surfaces_a_secret_manager_failure_before_the_call() {
let Err(error) = litellm_host::run::run(
@ -179,75 +126,6 @@ fn write_response(body: &str) -> String {
)
}
#[test]
fn provider_config_resolves_anthropic_and_azure_ai() {
assert!(messages_provider_config("anthropic").is_some());
assert!(messages_provider_config("azure_ai").is_some());
assert!(messages_provider_config("openai").is_none());
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(400);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({"x-count": 3}).as_object().unwrap().clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
Error::Headers(litellm_http::request::HeaderError {
context: "messages",
name: "x-count".to_string(),
actual: "number",
})
);
}
#[test]
fn has_header_is_case_insensitive() {
let headers = vec![("X-Api-Key".to_string(), "secret".to_string())];
assert!(has_header(&headers, "x-api-key"));
assert!(!has_header(&headers, "authorization"));
}
#[test]
fn has_bearer_auth_requires_a_nonempty_bearer_token() {
assert!(has_bearer_auth(&[(
"Authorization".to_string(),
"Bearer tok".to_string()
)]));
assert!(has_bearer_auth(&[(
"authorization".to_string(),
"bearer tok".to_string()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
"Bearer ".to_string()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
String::new()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
"Basic abc".to_string()
)]));
assert!(!has_bearer_auth(&[(
"x-api-key".to_string(),
"sk".to_string()
)]));
}
#[tokio::test]
async fn messages_round_trip_builds_azure_request_and_passes_response_through() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");

File diff suppressed because it is too large Load diff

View file

@ -1,152 +0,0 @@
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use serde_json::{Value, json};
use super::test_support::{
MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, request_body,
wire_request_with_document,
};
use crate::ocr::route::LocalOcrHost;
#[derive(Clone, Copy, Debug)]
enum Route {
Mistral,
AzureAi,
VertexMistral,
AzureCohereParse,
Cohere,
}
impl Route {
fn model(self) -> &'static str {
match self {
Self::Mistral => "mistral/model",
Self::AzureAi => "azure_ai/model",
Self::VertexMistral => "vertex_ai/mistral-ocr-maas",
Self::AzureCohereParse => "azure_ai/cohere-parse",
Self::Cohere => "cohere/model",
}
}
fn document_type(self) -> &'static str {
match self {
Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url",
Self::AzureCohereParse | Self::Cohere => "image_url",
}
}
fn options(self) -> Value {
match self {
Self::Mistral | Self::AzureAi => json!({"pages": [0]}),
Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}),
Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}),
}
}
}
/// What the host does to the wire request in `before_send`.
#[derive(Clone, Copy, Debug)]
enum Host {
Detached,
ReplacesDocument,
}
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
impl Host {
fn before_send(self, wire: WireRequest) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
let body = fields
.into_iter()
.map(|(name, value)| match self {
Self::Detached => (name, value),
Self::ReplacesDocument if name == "document" => {
let document_type = value["type"].clone();
let key = document_type.as_str().unwrap_or_default().to_string();
(name, json!({"type": document_type, key: REPLACED_DOCUMENT}))
}
Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
struct Sent {
result: Result<(), Error>,
provider_body: Option<Value>,
}
async fn send(route: Route, host: Host, document_base: &str) -> Sent {
let (base, seen, provider) = mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let document_type = route.document_type();
let document =
json!({"type": document_type, document_type: format!("{document_base}/scan.png")});
let request = wire_request_with_document(route.model(), &base, document, route.options());
let local =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire)));
let result = perform_ocr_with(local).await.map(|_| ());
match result {
Ok(()) => provider.await.unwrap(),
Err(_) => provider.abort(),
}
let provider_body = seen
.lock()
.unwrap()
.first()
.map(|request| request_body(request));
Sent {
result,
provider_body,
}
}
fn served_document_uri() -> String {
use base64::Engine;
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
}
#[rstest]
#[case::azure_ai(Route::AzureAi)]
#[case::vertex_mistral(Route::VertexMistral)]
#[case::azure_cohere_parse(Route::AzureCohereParse)]
#[tokio::test]
async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::Detached, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(served_document_uri())
);
}
#[rstest]
#[tokio::test]
async fn document_replaced_by_the_host_reaches_the_provider(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::ReplacesDocument, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(REPLACED_DOCUMENT)
);
}

View file

@ -1,203 +0,0 @@
use std::sync::{Arc, Mutex};
use futures_util::future::BoxFuture;
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::{
error::Error,
handler::{CallHooks, OcrClient},
transformation::LiteLLMOcrResponse,
};
use serde_json::{Value, json};
use tokio::{
io::{AsyncReadExt, AsyncWriteExt},
net::TcpListener,
};
use crate::ocr::{
route::{LocalOcrHost, ocr_machine},
types::LiteLLMOcrRequest,
wire::{OcrWireRequest, decode_request},
};
/// Stands in for a host with no hooks registered: the wire request goes out unchanged
/// and response events go nowhere.
pub(crate) struct NoHooks;
impl CallHooks<Error> for NoHooks {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(async move { Ok(wire) })
}
fn response_received<'a>(&'a self, _body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(async { Ok(()) })
}
}
pub(crate) fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.build()
.expect("test document client builds");
OcrClient::for_test(reqwest::Client::new(), document_http)
}
pub(crate) async fn perform_ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
crate::ocr::client::perform(&ocr_client(), request).await
}
pub(crate) async fn perform_ocr_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
wire_request_with_document(
model,
base,
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
options,
)
}
pub(crate) fn wire_request_with_document(
model: &str,
base: &str,
document: Value,
options: Value,
) -> LiteLLMOcrRequest {
decode_request(OcrWireRequest {
model: model.into(),
document,
api_key: Some(litellm_auth::SecretValue::new("test-key")),
api_base: Some(base.into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
})
.unwrap()
}
pub(crate) fn resolved_request(
request: LiteLLMOcrRequest,
) -> crate::ocr::types::ResolvedOcrRequest {
request
.map_document(crate::ocr::document::prepare_document)
.unwrap()
}
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let request = resolved_request(request);
let document = request.document.clone().with_source(source.into());
request.with_document(document.into())
}
pub(crate) fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
pub(crate) const SERVED_DOCUMENT: &[u8] = b"\x89PNG served document";
/// Serves [`SERVED_DOCUMENT`] as `image/png` to every connection until aborted.
pub(crate) async fn document_server() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let task = tokio::spawn(async move {
loop {
let (mut socket, _) = listener.accept().await.unwrap();
let mut buffer = [0u8; 4096];
let _ = socket.read(&mut buffer).await.unwrap();
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Type: image/png\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
SERVED_DOCUMENT.len()
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(SERVED_DOCUMENT).await.unwrap();
}
});
(base, task)
}
pub(crate) struct MockResponse {
pub status: u16,
pub headers: Vec<(&'static str, String)>,
pub body: Value,
}
impl MockResponse {
pub fn json(body: Value) -> Self {
Self {
status: 200,
headers: vec![],
body,
}
}
}
pub(crate) async fn mock_server(
responses: Vec<MockResponse>,
) -> (String, Arc<Mutex<Vec<String>>>, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let base = format!("http://{}", listener.local_addr().unwrap());
let requests = Arc::new(Mutex::new(Vec::new()));
let seen = requests.clone();
let server_base = base.clone();
let task = tokio::spawn(async move {
for response in responses {
let (mut socket, _) = listener.accept().await.unwrap();
let mut bytes = Vec::new();
let mut buffer = [0u8; 4096];
let header_end = loop {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") {
break index + 4;
}
};
let length = String::from_utf8_lossy(&bytes[..header_end])
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().unwrap())
})
.unwrap_or(0);
while bytes.len() < header_end + length {
let n = socket.read(&mut buffer).await.unwrap();
assert!(n > 0);
bytes.extend_from_slice(&buffer[..n]);
}
seen.lock()
.unwrap()
.push(String::from_utf8_lossy(&bytes).into_owned());
let body = serde_json::to_vec(&response.body).unwrap();
let headers = response
.headers
.into_iter()
.map(|(name, value)| {
format!("{name}: {}\r\n", value.replace("{base}", &server_base))
})
.collect::<String>();
let head = format!(
"HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n",
response.status,
body.len(),
headers
);
socket.write_all(head.as_bytes()).await.unwrap();
socket.write_all(&body).await.unwrap();
}
});
(base, requests, task)
}
pub(crate) fn header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
request
.lines()
.take_while(|line| !line.is_empty())
.find_map(|line| {
let (key, value) = line.split_once(':')?;
key.eq_ignore_ascii_case(name).then(|| value.trim())
})
}

View file

@ -1,584 +0,0 @@
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_llms::base_llm::ocr::{error::Error, transformation::OcrDocument};
use rstest::rstest;
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request};
use crate::ocr::route::LocalOcrHost;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[rstest]
#[case(
"reducto/parse-v3",
json!({
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://already.pdf",
json!({
"input":"reducto://already.pdf",
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[case(
"reducto/parse-legacy",
json!({
"enhance":{"agentic":[{"type":"table"}]},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://legacy.pdf",
json!({
"document_url":"reducto://legacy.pdf",
"options":{"enhance":{"agentic":[{"type":"table"}]}},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[tokio::test]
async fn request_mapping_matches_python(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[]}
}))])
.await;
let request = super::test_support::with_source(wire_request(model, &base, options), source);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(request_body(&requests[0]), expected);
}
#[rstest]
#[case("parse-v3")]
#[case("parse-legacy")]
#[tokio::test]
async fn data_uri_upload_preserves_multipart_headers(
#[case] model: &str,
#[values("application/pdf", "image/png")] mime_type: &str,
) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})),
])
.await;
let document = if mime_type.starts_with("image/") {
json!({"type":"image_url","image_url":format!("data:{mime_type};base64,YWJj")})
} else {
json!({"type":"document_url","document_url":format!("data:{mime_type};base64,YWJj")})
};
let mut request = crate::ocr::types::LiteLLMOcrRequest {
document: serde_json::from_value::<OcrDocument>(document)
.unwrap()
.into(),
..wire_request(&format!("reducto/{model}"), &base, json!({}))
};
request.transport.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("content-type: multipart/form-data; boundary=")
);
assert!(requests[0].contains("x-trace: upload-test"));
let multipart = requests[0].split_once("\r\n\r\n").unwrap().1;
assert!(multipart.contains(&format!("Content-Type: {mime_type}\r\n")));
assert!(multipart.contains("\r\n\r\nabc\r\n--"));
assert!(requests[1].starts_with("POST /parse "));
let source_field = if model == "parse-legacy" {
"document_url"
} else {
"input"
};
assert_eq!(
request_body(&requests[1]),
json!({source_field:"reducto://uploaded.pdf"})
);
for request in requests.iter() {
assert!(
request
.to_ascii_lowercase()
.contains("authorization: bearer test-key\r\n")
);
}
}
#[tokio::test]
async fn response_received_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))).with_observer(
move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
assert_eq!(request_count.lock().unwrap().len(), 2);
assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#);
}
},
);
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case(json!({"file_id":""}))]
#[case(json!({}))]
#[case(json!({"file_id":null}))]
#[tokio::test]
async fn invalid_upload_ids_stop_before_parse(#[case] response: Value) {
let (base, seen, server) = mock_server(vec![MockResponse::json(response)]).await;
let error = perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("file_id"));
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn upload_failure_stops_before_parse() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 503,
headers: vec![],
body: json!({"error":"unavailable"}),
}])
.await;
assert!(
perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.is_err()
);
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[rstest]
#[case("https://example.com/a.pdf", Error::ReductoSource)]
#[case("reducto://", Error::RequestField { path: "document file id".into() })]
#[case("data:application/pdf;base64", Error::InvalidDataUri)]
#[case("data:application/pdf;base64,INVALID!", Error::InvalidDataUri)]
#[tokio::test]
async fn rejects_invalid_document_sources_before_network(
#[case] source: &str,
#[case] expected: Error,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({}))]).await;
let request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
source,
);
let result = perform_ocr(request).await;
server.abort();
let _ = server.await;
assert!(
seen.lock().unwrap().is_empty(),
"sent invalid source: {source}"
);
let error = result.unwrap_err();
assert_eq!(
std::mem::discriminant(&error),
std::mem::discriminant(&expected)
);
assert_eq!(error.http_status_code(), Some(400));
assert_eq!(error.to_string(), expected.to_string());
}
#[test]
fn response_normalization_groups_blocks_and_distinguishes_null_result() {
use litellm_llms::reducto::ocr::transformation::{
ReductoResponse, normalize_response as transform_ocr_response,
};
let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[
{"blocks":[{
"type":"Table",
"content":"B",
"bbox":{"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4},
"confidence":"high",
"granular_confidence":{"parse_confidence":0.95,"extract_confidence":null},
"image_url":null
}]},
{"blocks":[{"content":"A","bbox":{"page":1},"type":"Text"},{"content":"C","bbox":{"page":1}}]}
]}});
let response: ReductoResponse = serde_json::from_value(raw).unwrap();
let normalized = transform_ocr_response("parse-v3", response)
.unwrap()
.into_json();
assert_eq!(normalized["pages"][0]["markdown"], "A\n\nC");
assert_eq!(normalized["pages"][1]["markdown"], "B");
assert_eq!(normalized["pages"][1]["blocks"][0]["type"], "Table");
assert_eq!(
normalized["pages"][1]["blocks"][0]["bbox"],
json!({"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4})
);
assert_eq!(normalized["pages"][1]["blocks"][0]["confidence"], "high");
assert_eq!(
normalized["pages"][1]["blocks"][0]["granular_confidence"]["parse_confidence"],
0.95
);
assert!(normalized["pages"][1]["blocks"][0]["image_url"].is_null());
assert_eq!(normalized["usage_info"]["pages_processed"], 2);
assert_eq!(normalized["usage_info"]["credits"], 3.0);
let missing: ReductoResponse =
serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap();
let missing = transform_ocr_response("parse-v3", missing).unwrap();
assert_eq!(missing.pages[0].markdown, "text");
let null: ReductoResponse = serde_json::from_value(
json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}),
)
.unwrap();
let null = transform_ocr_response("parse-v3", null).unwrap();
assert!(null.pages.is_empty());
}
#[tokio::test]
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
let mut request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
"reducto://ready.pdf",
);
request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, None);
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer existing")
);
}
#[tokio::test]
async fn native_format_retains_the_provider_response() {
let raw = json!({
"result":{"chunks":[{"content":"native OCR response"}]},
"usage":{"num_pages":1}
});
let (base, _, server) = mock_server(vec![MockResponse::json(raw.clone())]).await;
let request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({"req_format":"native"})),
"reducto://ready.pdf",
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "native OCR response");
assert_eq!(response.provider_native_response.as_ref(), raw.as_object());
}
#[tokio::test]
async fn unknown_model_reaches_parse_and_keeps_its_name() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[{"content":"future model response"}]}
}))])
.await;
let request = super::test_support::with_source(
wire_request("reducto/future-parse-model", &base, json!({})),
"reducto://ready.pdf",
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.model, "future-parse-model");
assert_eq!(response.pages[0].markdown, "future model response");
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(
request_body(&requests[0]),
json!({"input":"reducto://ready.pdf"})
);
}
#[tokio::test]
async fn guardrail_rewrites_document_before_upload() {
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await;
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({})))
.with_before_send(|wire, _| {
assert_eq!(
wire.body["document_url"],
"data:application/pdf;base64,YWJj"
);
Ok(WireRequest {
body: json!({"type":"document_url","document_url":"reducto://guarded.pdf"}),
..wire
})
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert!(requests[0].contains("reducto://guarded.pdf"));
}
mod transformation {
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_llms::{
base_llm::ocr::transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext},
reducto::ocr::transformation::*,
};
use rstest::rstest;
use super::*;
use crate::ocr::{
route::LocalOcrHost,
test_support::{MockResponse, mock_server, perform_ocr, perform_ocr_with, wire_request},
};
#[tokio::test]
async fn v3_options_preserve_explicit_null() {
let overrides =
serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true}))
.unwrap();
let params = ReductoParseV3Config
.map_ocr_params(&overrides, "parse-v3")
.unwrap();
let client = crate::ocr::test_support::ocr_client();
let connection = OcrConnection::default();
let document = serde_json::from_value(
json!({"type":"document_url","document_url":"reducto://ready.pdf"}),
)
.unwrap();
let body = ReductoParseV3Config
.async_transform_ocr_request(
"parse-v3",
document,
&params,
&[],
OcrRequestContext {
client: &client,
connection: &connection,
},
)
.await
.unwrap();
assert_eq!(
serde_json::to_value(body).unwrap(),
json!({
"input":"reducto://ready.pdf", "formatting":null, "settings":{}
})
);
let absent = ReductoParseV3Config
.map_ocr_params(
&litellm_core_utils::call_arguments::CallArguments::default(),
"parse-v3",
)
.unwrap();
assert_eq!(serde_json::to_value(absent).unwrap(), json!({}));
}
#[rstest]
#[case(
"reducto/parse-v3",
json!({
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://already.pdf",
json!({
"input":"reducto://already.pdf",
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[case(
"reducto/parse-legacy",
json!({
"enhance":{"agentic":[{"type":"table"}]},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://legacy.pdf",
json!({
"document_url":"reducto://legacy.pdf",
"options":{"enhance":{"agentic":[{"type":"table"}]}},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[tokio::test]
async fn request_mapping_matches_python(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[]}
}))])
.await;
let request =
crate::ocr::test_support::with_source(wire_request(model, &base, options), source);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(request_body(&requests[0]), expected);
}
#[rstest]
#[case("parse-v3")]
#[case("parse-legacy")]
#[tokio::test]
async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})),
])
.await;
let mut request = wire_request(&format!("reducto/{model}"), &base, json!({}));
request.transport.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("content-type: multipart/form-data; boundary=")
);
assert!(requests[0].contains("x-trace: upload-test"));
assert!(requests[0].contains("application/pdf"));
assert!(requests[0].contains("abc"));
assert!(requests[1].starts_with("POST /parse "));
}
#[tokio::test]
async fn response_received_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({})))
.with_observer(move |event| {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
assert_eq!(request_count.lock().unwrap().len(), 2);
assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#);
}
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case("https://example.com/a.pdf")]
#[case("reducto://")]
#[case("data:application/pdf;base64")]
#[case("data:application/pdf;base64,INVALID!")]
#[tokio::test]
async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
let request = crate::ocr::test_support::with_source(
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})),
source,
);
assert!(perform_ocr(request).await.is_err());
}
#[tokio::test]
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
let mut request = crate::ocr::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
"reducto://ready.pdf",
);
request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, None);
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer existing")
);
}
#[rstest]
#[case("reducto/parse-v3")]
#[case("reducto/parse-legacy")]
#[tokio::test]
async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let mut request = wire_request(model, &base, json!({}));
request.transport.extra_headers = vec![("authorization".into(), "Bearer original".into())];
let host = LocalOcrHost::new(request).with_before_send(|wire, _| {
Ok(WireRequest {
headers: vec![("authorization".into(), "Bearer guarded".into())],
..wire
})
});
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(requests[1].starts_with("POST /parse "));
for request in requests.iter() {
assert!(request.contains("authorization: Bearer guarded"));
assert!(!request.contains("Bearer original"));
}
}
}

View file

@ -1,143 +0,0 @@
use litellm_auth::InputSource;
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
let request = super::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
#[test]
fn host_registration_selects_deepseek_without_affecting_mistral() {
assert!(crate::ocr::arguments::is_supported_request(
"deepseek-ocr-maas",
Some("vertex_ai")
));
assert!(crate::ocr::arguments::is_supported_request(
"mistral-ocr-maas",
Some("vertex_ai")
));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
mod deepseek_transformation {
use serde_json::json;
use super::*;
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
let request = crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert_eq!(body["provider_option"], "value");
assert!(body.get("vertex_project").is_none());
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
}

View file

@ -1,293 +0,0 @@
use litellm_auth::InputSource;
use litellm_llms::base_llm::ocr::{settings::OcrSettings, transformation::OcrResponseFormat};
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, ocr_client, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let request = wire_request(
"vertex_ai/mistral-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"extract_footer":true
}),
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
assert_eq!(
request_body(&requests[0]),
json!({
"model":"mistral-ocr-maas",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"extract_footer":true
})
);
}
#[tokio::test]
async fn configured_project_and_location_apply_when_the_call_sets_neither() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let client = ocr_client().with_settings(OcrSettings {
vertex_project: Some("configured-project".into()),
vertex_location: Some("europe-west4".into()),
..OcrSettings::default()
});
crate::ocr::client::perform(
&client,
wire_request("vertex_ai/mistral-ocr-maas", &base, json!({})),
)
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].starts_with(
"POST /v1/projects/configured-project/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
}
#[tokio::test]
async fn supplied_authorization_is_forwarded_without_a_static_token() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"vertex_ai/model",
&base,
json!({"vertex_project":"project-1"}),
);
request.credentials.api_key = None;
request.transport.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer supplied")
);
}
#[tokio::test]
async fn invalid_credentials_fail_before_provider_http() {
let request = wire_request(
"vertex_ai/model",
"http://127.0.0.1:1",
json!({"vertex_credentials": true}),
);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("vertex_credentials"));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/mistral-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
#[tokio::test]
async fn adapters_build_complete_requests_and_share_mistral_normalization() {
use std::time::Duration;
use litellm_llms::{
base_llm::ocr::transformation::BaseOcrConfig,
mistral::ocr::transformation::MistralOcrConfig,
vertex_ai::ocr::transformation::VertexAiOcrConfig,
};
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "ignored"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct = crate::ocr::prepare::prepare_request_for_test(
super::test_support::resolved_request(direct),
);
let vertex = crate::ocr::prepare::prepare_request_for_test(
super::test_support::resolved_request(vertex),
);
let direct_http = MistralOcrConfig
.prepare_request(&direct, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
let vertex_http = VertexAiOcrConfig
.prepare_request(&vertex, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
for http in [&direct_http, &vertex_http] {
assert_eq!(http.header("authorization").unwrap(), "Bearer test-key");
assert_eq!(http.header("content-type").unwrap(), "application/json");
assert_eq!(http.timeout(), Some(Duration::from_secs(2)));
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "ignored"
})
);
}
let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"});
let raw = serde_json::to_vec(&payload).unwrap();
let direct_response = MistralOcrConfig
.transform_ocr_response(&direct.model, &raw, OcrResponseFormat::Litellm)
.unwrap()
.into_json();
let vertex_response = VertexAiOcrConfig
.transform_ocr_response(&vertex.model, &raw, OcrResponseFormat::Litellm)
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert_eq!(direct_response["extra"], "preserved");
}
mod transformation {
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::test_support::wire_request;
#[rstest]
#[case::mistral(false)]
#[case::vertex(true)]
#[tokio::test]
async fn configs_build_complete_requests_and_share_mistral_normalization(
#[case] use_vertex: bool,
) {
use std::time::Duration;
use litellm_llms::{
base_llm::ocr::transformation::BaseOcrConfig,
mistral::ocr::transformation::MistralOcrConfig,
vertex_ai::ocr::transformation::VertexAiOcrConfig,
};
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "preserved"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct = crate::ocr::prepare::prepare_request_for_test(
crate::ocr::test_support::resolved_request(direct),
);
let vertex = crate::ocr::prepare::prepare_request_for_test(
crate::ocr::test_support::resolved_request(vertex),
);
let direct_http = MistralOcrConfig
.prepare_request(&direct, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
let vertex_http = VertexAiOcrConfig
.prepare_request(&vertex, &client, &crate::ocr::test_support::NoHooks)
.await
.unwrap();
assert_eq!(direct_http.url(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
let http = if use_vertex {
&vertex_http
} else {
&direct_http
};
assert_eq!(http.header("authorization").unwrap(), "Bearer test-key");
assert_eq!(http.header("content-type").unwrap(), "application/json");
assert_eq!(http.timeout(), Some(Duration::from_secs(2)));
let body: Value = serde_json::from_slice(http.body()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "preserved"
})
);
let payload = serde_json::to_vec(
&json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}),
)
.unwrap();
let direct_response = MistralOcrConfig
.transform_ocr_response(&direct.model, &payload, Default::default())
.unwrap()
.into_json();
let vertex_response = VertexAiOcrConfig
.transform_ocr_response(&vertex.model, &payload, Default::default())
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert_eq!(direct_response["extra"], "preserved");
}
}

View file

@ -217,13 +217,25 @@ mod tests {
"Authorization".to_string(),
"Bearer abc".to_string()
)]));
assert!(has_bearer_auth(&[(
"authorization".to_string(),
"bearer abc".to_string()
)]));
assert!(!has_bearer_auth(&[(
"Authorization".to_string(),
"Bearer ".to_string()
)]));
assert!(!has_bearer_auth(&[(
"authorization".to_string(),
String::new()
)]));
assert!(!has_bearer_auth(&[(
"Authorization".to_string(),
"Basic abc".to_string()
)]));
assert!(!has_bearer_auth(&[(
"x-api-key".to_string(),
"abc".to_string()
)]));
}
}

View file

@ -218,7 +218,3 @@ fn anthropic_body(
);
Value::Object(body)
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;

View file

@ -302,7 +302,3 @@ fn has_blank_text(message: &ChatMessage) -> bool {
}),
}
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;

View file

@ -1,7 +1,11 @@
use serde_json::json;
use super::*;
use crate::base_llm::chat::transformation::Error;
use litellm_llms::{
anthropic::chat::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG,
base_llm::chat::transformation::{
BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported,
},
};
use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse};
use serde_json::{Map, Value, json};
fn messages(value: Value) -> Vec<ChatMessage> {
serde_json::from_value(value).expect("valid messages")
@ -205,7 +209,8 @@ fn declines_tool_calls_tool_results_and_multimodal_content() {
);
assert_eq!(
reason(
json!([{"role": "user", "content": [
json!([
{"role": "user", "content": [
{"type": "image_url", "image_url": {"url": "https://x/y.png"}}
]}]),
json!({})
@ -214,7 +219,8 @@ fn declines_tool_calls_tool_results_and_multimodal_content() {
);
assert_eq!(
reason(
json!([{"role": "user", "content": [
json!([
{"role": "user", "content": [
{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}
]}]),
json!({})

View file

@ -1,7 +1,11 @@
use serde_json::json;
use super::*;
use crate::base_llm::chat::transformation::Error;
use litellm_llms::{
base_llm::chat::transformation::{
BaseConfig, Error, ProviderChatResponseData, RequestAuth, Unsupported,
},
bedrock::chat::converse_transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG,
};
use litellm_types::{llms::openai::ChatMessage, utils::ChatCompletionsResponse};
use serde_json::{Map, Value, json};
fn messages(value: Value) -> Vec<ChatMessage> {
serde_json::from_value(value).expect("valid messages")

View file

@ -56,7 +56,7 @@ pub struct ChatCompletionsChoice {
///
/// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the
/// `ModelResponse` it already created, and echoing the provider's own id here
/// would change it. Pinned by `response_carries_no_id` in `tests.rs`.
/// would change it. Pinned by `response_carries_no_id` in the Anthropic chat transformation tests.
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatCompletionsResponse {
pub created: u64,

View file

@ -11,7 +11,10 @@ from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_
from litellm.litellm_core_utils.llm_cost_calc.utils import parse_prompt_tokens_details
from litellm.llms.base_llm.ocr.transformation import OCRUsageInfo
from litellm.llms.bedrock.batches.transformation import titan_embedding_usage_from_batch_output
from litellm.llms.vertex_ai.batches.transformation import vertex_prompt_tokens_details
from litellm.llms.vertex_ai.batches.transformation import (
is_native_vertex_batch_output_row,
native_vertex_batch_row_stats,
)
from litellm.types.llms.openai import Batch
from litellm.types.utils import ModelInfo, Usage
from litellm.utils import token_counter
@ -31,6 +34,20 @@ class BatchCostUsageResult:
_COMPLETED_BATCH_STATUSES: Final = frozenset({"completed", "complete"})
def _uses_native_vertex_output(
custom_llm_provider: str,
model_name: str | None,
first_row: Mapping[str, object] | None,
) -> bool:
if custom_llm_provider != "vertex_ai":
return False
if model_name and getattr(litellm, "disable_vertex_batch_output_transformation", False):
return True
return first_row is not None and is_native_vertex_batch_output_row(first_row)
_TERMINAL_BATCH_STATUSES: Final = _COMPLETED_BATCH_STATUSES | frozenset({"failed", "cancelled", "expired"})
@ -66,12 +83,9 @@ async def calculate_batch_cost_and_usage(
deployment-specific pricing (e.g. input_cost_per_token_batches)
is used instead of the global cost map.
"""
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
):
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name)
first_row: Final = file_content_dictionary[0] if file_content_dictionary else None
if _uses_native_vertex_output(custom_llm_provider, model_name, first_row):
return calculate_vertex_ai_batch_cost_and_usage(file_content_dictionary, model_name, model_info=model_info)
return _aggregate_batch_cost_usage_models(
entries=file_content_dictionary,
@ -126,11 +140,11 @@ async def _handle_completed_batch(
)
output_file_result: Final = (
calculate_vertex_ai_batch_cost_and_usage(_get_file_content_as_dictionary(file_content), model_name)
if (
custom_llm_provider == "vertex_ai"
and model_name
and getattr(litellm, "disable_vertex_batch_output_transformation", False)
calculate_vertex_ai_batch_cost_and_usage(
_iter_batch_output_entries(file_content), model_name, model_info=model_info
)
if _uses_native_vertex_output(
custom_llm_provider, model_name, next(_iter_batch_output_entries(file_content), None)
)
else _aggregate_batch_cost_usage_models(
entries=_iter_batch_output_entries(file_content),
@ -332,69 +346,36 @@ def _aggregate_batch_cost_usage_models(
def calculate_vertex_ai_batch_cost_and_usage(
vertex_ai_batch_responses: list[dict],
vertex_ai_batch_responses: Iterable[dict],
model_name: str | None = None,
model_info: ModelInfo | None = None,
) -> BatchCostUsageResult:
"""
Calculate both cost and usage from raw Vertex AI batch responses.
Used only when ``litellm.disable_vertex_batch_output_transformation = True``.
In that case the GCS predictions.jsonl is returned as-is, with each line in
the native Vertex format:
{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}}}
usageMetadata contains promptTokenCount, candidatesTokenCount, totalTokenCount.
A row with no ``response`` is counted as failed - the same signal already
used to skip it from cost/usage aggregation, since Vertex batch prediction
output doesn't establish a distinct error shape in this (non-default) path.
Cost and usage of a native Vertex predictions.jsonl, one
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
generateContent row or `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
embedding row per line. `model_name` (the deployment model) prices every row, else each row's own
`modelVersion` does; a row without a usable response counts as failed.
"""
from litellm.cost_calculator import batch_cost_calculator
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig
total_prompt_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
total_completion_cost = 0.0 # rebind-ok: loop accumulator, matches total_tokens below
total_tokens = 0
prompt_tokens = 0
completion_tokens = 0
successful_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
failed_requests = 0 # rebind-ok: loop accumulator, matches total_cost/total_tokens above
actual_model_name: Final = model_name or "gemini-2.0-flash-001"
for response in vertex_ai_batch_responses:
response_body = response.get("response")
if response_body is None:
failed_requests += 1
continue
successful_requests += 1
usage_metadata = response_body.get("usageMetadata", {})
_prompt = usage_metadata.get("promptTokenCount", 0) or 0
_completion = usage_metadata.get("candidatesTokenCount", 0) or 0
_total = usage_metadata.get("totalTokenCount", 0) or (_prompt + _completion)
line_usage = Usage(
prompt_tokens=_prompt,
completion_tokens=_completion,
total_tokens=_total,
prompt_tokens_details=vertex_prompt_tokens_details(usage_metadata),
row_stats: Final = tuple(
native_vertex_batch_row_stats(
row,
model_name,
model_info=model_info,
calculate_usage=VertexGeminiConfig._calculate_usage,
cost_calculator=batch_cost_calculator,
)
try:
p_cost, c_cost = batch_cost_calculator(
usage=line_usage,
model=actual_model_name,
custom_llm_provider="vertex_ai",
)
total_prompt_cost += p_cost
total_completion_cost += c_cost
except Exception as e:
verbose_logger.debug("vertex_ai batch cost calculation error for line: %s", str(e))
prompt_tokens += _prompt
completion_tokens += _completion
total_tokens += _total
for row in vertex_ai_batch_responses
)
priced: Final = tuple(stats for stats in row_stats if stats is not None)
total_prompt_cost: Final = sum(stats.prompt_cost for stats in priced)
total_completion_cost: Final = sum(stats.completion_cost for stats in priced)
prompt_tokens: Final = sum(stats.usage.prompt_tokens for stats in priced)
completion_tokens: Final = sum(stats.usage.completion_tokens for stats in priced)
total_tokens: Final = sum(stats.total_tokens for stats in priced)
total_cost: Final = total_prompt_cost + total_completion_cost
verbose_logger.info(
"vertex_ai batch cost: cost=%s, prompt=%d, completion=%d, total=%d, successful=%d, failed=%d",
@ -402,8 +383,8 @@ def calculate_vertex_ai_batch_cost_and_usage(
prompt_tokens,
completion_tokens,
total_tokens,
successful_requests,
failed_requests,
len(priced),
len(row_stats) - len(priced),
)
return BatchCostUsageResult(
@ -413,9 +394,13 @@ def calculate_vertex_ai_batch_cost_and_usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
),
models=[actual_model_name],
successful_requests=successful_requests,
failed_requests=failed_requests,
models=(
[model_name]
if model_name
else list(dict.fromkeys(stats.model for stats in priced if stats.model is not None))
),
successful_requests=len(priced),
failed_requests=len(row_stats) - len(priced),
prompt_cost=total_prompt_cost,
completion_cost=total_completion_cost,
)

View file

@ -176,6 +176,22 @@ def create_file(
if logging_obj is None:
raise ValueError("logging_obj is required")
client: Final = kwargs.get("client")
if litellm_params_dict.get("passthrough") is True and (
custom_llm_provider != "vertex_ai" or purpose != "batch"
):
raise litellm.exceptions.BadRequestError(
message=(
"`passthrough=True` uploads the file bytes unchanged for a native Vertex AI batch, so it needs "
f"custom_llm_provider='vertex_ai' and purpose='batch', got '{custom_llm_provider}' and '{purpose}'."
),
model="n/a",
llm_provider=custom_llm_provider or "n/a",
response=httpx.Response(
status_code=400,
content="passthrough needs a vertex_ai batch",
request=httpx.Request(method="create_file", url="https://github.com/BerriAI/litellm"),
),
)
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600

View file

@ -118,6 +118,7 @@ class SlackAlerting(CustomBatchLogger):
self.default_webhook_url = default_webhook_url
self.flush_lock = asyncio.Lock()
self.periodic_started = False
self._periodic_flush_task: asyncio.Task[None] | None = None
self.hanging_request_check = AlertingHangingRequestCheck(
slack_alerting_object=self,
)
@ -129,6 +130,12 @@ class SlackAlerting(CustomBatchLogger):
self.digest_lock = asyncio.Lock()
super().__init__(**kwargs, flush_lock=self.flush_lock)
def _ensure_periodic_flush_task(self) -> None:
if self.periodic_started and (self._periodic_flush_task is None or not self._periodic_flush_task.done()):
return
self._periodic_flush_task = asyncio.create_task(self.periodic_flush())
self.periodic_started = True
def update_values(
self,
alerting: list | None = None,
@ -141,17 +148,14 @@ class SlackAlerting(CustomBatchLogger):
):
if alerting is not None:
self.alerting = alerting
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
self._ensure_periodic_flush_task()
if alerting_threshold is not None:
self.alerting_threshold = alerting_threshold
if alert_types is not None:
self.alert_types = alert_types
if alerting_args is not None:
self.alerting_args = SlackAlertingArgs(**alerting_args)
if not self.periodic_started:
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
self._ensure_periodic_flush_task()
if alert_type_config is not None:
for key, val in alert_type_config.items():
self.alert_type_config[key] = AlertTypeConfig(**val) if isinstance(val, dict) else val
@ -1446,9 +1450,8 @@ Model Info:
return
# Start periodic flush if not already started
if not self.periodic_started and self.alerting is not None and len(self.alerting) > 0:
asyncio.create_task(self.periodic_flush())
self.periodic_started = True
if self.alerting is not None and len(self.alerting) > 0:
self._ensure_periodic_flush_task()
if "webhook" in self.alerting and alert_type == "budget_alerts" and user_info is not None:
await self.send_webhook_alert(webhook_event=user_info)

View file

@ -765,3 +765,30 @@ def set_response_cost_in_hidden_params(response: _CarriesHiddenParams, cost: flo
RESPONSE_COST_HEADER: cost,
}
hidden_params["additional_headers"] = merged
_HIDDEN_PARAMS_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_PROVIDER_HEADERS_ADAPTER: Final = TypeAdapter(Mapping[str, str])
def set_provider_response_headers_in_hidden_params(
response: _CarriesHiddenParams, headers: httpx.Headers | Mapping[str, str]
) -> None:
hidden_params: Final = response._hidden_params # pyright: ignore[reportPrivateUsage] # no public accessor
existing_additional_headers: Final[object] = hidden_params.get("additional_headers")
raw_headers: Final[dict[str, str]] = dict(headers) # mutable-ok: stored as the plain-dict hidden param
additional_headers: Final[dict[str, object]] = { # mutable-ok: assigned into the plain-dict hidden params
**process_response_headers(raw_headers),
**(existing_additional_headers if isinstance(existing_additional_headers, Mapping) else _NO_HEADERS),
}
hidden_params["headers"] = raw_headers
hidden_params["additional_headers"] = additional_headers
def get_provider_response_headers_from_hidden_params(response: object) -> Mapping[str, str] | None:
hidden_params: Final[object] = getattr(response, "_hidden_params", None)
try:
validated: Final = _HIDDEN_PARAMS_ADAPTER.validate_python(hidden_params)
return _PROVIDER_HEADERS_ADAPTER.validate_python(validated.get("headers"))
except ValidationError:
return None

View file

@ -72,6 +72,7 @@ from litellm.litellm_core_utils.classifier_logging import (
is_classifier_call,
)
from litellm.litellm_core_utils.core_helpers import (
get_provider_response_headers_from_hidden_params,
is_expected_client_error,
reconstruct_model_name,
set_response_cost_in_hidden_params,
@ -2353,6 +2354,15 @@ class Logging(LiteLLMLoggingBaseClass):
)
return logging_result
def _surface_response_headers_from_result(self, logging_result: object) -> None:
existing: Final[object] = self.model_call_details.get("response_headers")
if existing is not None:
return
headers: Final = get_provider_response_headers_from_hidden_params(logging_result)
if headers is None:
return
self.model_call_details["response_headers"] = headers
def _merge_hidden_params_from_response_into_metadata(self, logging_result: object) -> None:
"""
Copy response._hidden_params into litellm_params.metadata['hidden_params'].
@ -2386,6 +2396,7 @@ class Logging(LiteLLMLoggingBaseClass):
build_logging_payload: bool = True,
):
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
self._surface_response_headers_from_result(logging_result)
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
if hidden_params:
if self.model_call_details.get("litellm_params") is not None:
@ -2788,6 +2799,7 @@ class Logging(LiteLLMLoggingBaseClass):
if complete_streaming_response is not None:
verbose_logger.debug("Logging Details LiteLLM-Success Call streaming complete")
self.model_call_details["complete_streaming_response"] = complete_streaming_response
self._surface_response_headers_from_result(complete_streaming_response)
self.model_call_details["response_cost"] = self._response_cost_calculator(
result=complete_streaming_response
)
@ -3302,6 +3314,7 @@ class Logging(LiteLLMLoggingBaseClass):
print_verbose("Async success callbacks: Got a complete streaming response")
self.model_call_details["async_complete_streaming_response"] = complete_streaming_response
self._surface_response_headers_from_result(complete_streaming_response)
try:
if self.model_call_details.get("cache_hit", False) is True:
@ -6362,12 +6375,15 @@ def _extract_response_obj_and_hidden_params(
original_exception: Exception | None,
) -> tuple[dict, dict | None]:
"""Extract response_obj and hidden_params from init_response_obj."""
hidden_params: dict | None = None
hidden_params: dict | None = (
getattr(init_response_obj, "_hidden_params", None)
if isinstance(init_response_obj, BaseModel | HttpxBinaryResponseContent)
else None
)
if init_response_obj is None:
response_obj = {}
elif isinstance(init_response_obj, BaseModel):
response_obj = init_response_obj.model_dump()
hidden_params = getattr(init_response_obj, "_hidden_params", None)
elif isinstance(init_response_obj, dict):
response_obj = init_response_obj
elif isinstance(init_response_obj, HttpxBinaryResponseContent):

View file

@ -1218,6 +1218,8 @@ class BedrockModelInfo(BaseLLMModelInfo):
alt_model: Final = BedrockModelInfo.get_non_litellm_routing_model_name(model=model)
if base_model in litellm.bedrock_converse_models or alt_model in litellm.bedrock_converse_models:
return "converse"
if _OPENAI_FAMILY_MODEL_RE.search(base_model):
return "converse"
return "invoke"
@staticmethod

View file

@ -45,6 +45,7 @@ from litellm.litellm_core_utils.audio_utils.subtitle_utils import (
SUBTITLE_RESPONSE_FORMATS,
synthesize_subtitle_document,
)
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.get_litellm_params import AWS_CREDENTIAL_KWARGS_KEYS
from litellm.litellm_core_utils.llm_request_utils import serialize_multipart_form_fields
from litellm.litellm_core_utils.realtime_errors import (
@ -1461,6 +1462,7 @@ class BaseLLMHTTPHandler:
transformed: Final = provider_config.transform_audio_transcription_response(
raw_response=response,
)
set_provider_response_headers_in_hidden_params(transformed, response.headers)
if not provider_config.supports_subtitle_synthesis:
return transformed
requested_format: Final = optional_params.get("response_format")
@ -6960,11 +6962,13 @@ class BaseLLMHTTPHandler:
provider_config=image_edit_provider_config,
)
return image_edit_provider_config.transform_image_edit_response(
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
return image_edit_response
async def async_image_edit_handler(
self,
@ -7059,11 +7063,13 @@ class BaseLLMHTTPHandler:
provider_config=image_edit_provider_config,
)
return image_edit_provider_config.transform_image_edit_response(
image_edit_response: Final = image_edit_provider_config.transform_image_edit_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(image_edit_response, response.headers)
return image_edit_response
def image_generation_handler(
self,
@ -7186,6 +7192,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
encoding=None,
)
set_provider_response_headers_in_hidden_params(model_response, response.headers)
return model_response
@ -7293,6 +7300,7 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
encoding=None,
)
set_provider_response_headers_in_hidden_params(model_response, response.headers)
return model_response
@ -12077,11 +12085,13 @@ class BaseLLMHTTPHandler:
provider_config=text_to_speech_provider_config,
)
return text_to_speech_provider_config.transform_text_to_speech_response(
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
return speech_response
async def async_text_to_speech_handler(
self,
@ -12176,11 +12186,13 @@ class BaseLLMHTTPHandler:
provider_config=text_to_speech_provider_config,
)
return text_to_speech_provider_config.transform_text_to_speech_response(
speech_response: Final = text_to_speech_provider_config.transform_text_to_speech_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
set_provider_response_headers_in_hidden_params(speech_response, response.headers)
return speech_response
#########################################################
########## SKILLS API HANDLERS ##########################

View file

@ -27,6 +27,7 @@ from litellm import LlmProviders
from litellm._logging import verbose_logger
from litellm.constants import DEFAULT_MAX_RETRIES
from litellm.files.types import FileContentStreamingResult
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.logging_utils import speech_request_body, track_llm_api_timing
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
@ -1404,7 +1405,6 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization: str | None = None,
headers: dict | None = None,
):
response = None
try:
openai_aclient: Final = self._get_openai_client(
is_async=True,
@ -1428,8 +1428,10 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
request_data: Final = {**data, "extra_headers": headers} if headers else data
response = await openai_aclient.images.generate(**request_data, timeout=timeout)
stringified_response: Final = response.model_dump()
raw_response: Final = await openai_aclient.images.with_raw_response.generate(
**request_data, timeout=timeout
)
stringified_response: Final = raw_response.parse().model_dump()
## LOGGING
logging_obj.post_call(
input=prompt,
@ -1437,11 +1439,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
additional_args={"complete_input_dict": data},
original_response=stringified_response,
)
return convert_to_model_response_object(
image_response: Final[ImageResponse] = convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
response_type="image_generation",
)
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
return image_response
except Exception as e:
## LOGGING
logging_obj.post_call(
@ -1512,9 +1516,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
## COMPLETION CALL
request_data: Final = {**data, "extra_headers": headers} if headers else data
_response: Final = openai_client.images.generate(**request_data, timeout=timeout)
raw_response: Final = openai_client.images.with_raw_response.generate(**request_data, timeout=timeout)
response: Final = _response.model_dump()
response: Final = raw_response.parse().model_dump()
## LOGGING
logging_obj.post_call(
input=prompt,
@ -1522,11 +1526,13 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
additional_args={"complete_input_dict": data},
original_response=response,
)
return convert_to_model_response_object(
image_response: Final[ImageResponse] = convert_to_model_response_object(
response_object=response,
model_response_object=model_response,
response_type="image_generation",
)
set_provider_response_headers_in_hidden_params(image_response, raw_response.headers)
return image_response
except OpenAIError as e:
## LOGGING
logging_obj.post_call(
@ -1609,7 +1615,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
input=input,
**optional_params,
)
return HttpxBinaryResponseContent(response=response.response)
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
return speech_response
async def async_audio_speech(
self,
@ -1655,8 +1663,9 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
input=input,
**optional_params,
)
return HttpxBinaryResponseContent(response=response.response)
speech_response: Final = HttpxBinaryResponseContent(response=response.response)
set_provider_response_headers_in_hidden_params(speech_response, response.response.headers)
return speech_response
class OpenAIFilesAPI(BaseLLM):

View file

@ -4,11 +4,10 @@ import httpx
from openai import AsyncOpenAI, OpenAI
from pydantic import BaseModel
import litellm
if TYPE_CHECKING:
from aiohttp import ClientSession
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_name
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
@ -31,11 +30,6 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
data: dict,
timeout: float | httpx.Timeout,
):
"""
Helper to:
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
- call openai_aclient.audio.transcriptions.create by default
"""
try:
raw_response = await openai_aclient.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
@ -51,20 +45,11 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
data: dict,
timeout: float | httpx.Timeout,
):
"""
Helper to:
- call openai_aclient.audio.transcriptions.with_raw_response when litellm.return_response_headers is True
- call openai_aclient.audio.transcriptions.create by default
"""
try:
if litellm.return_response_headers is True:
raw_response = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response = raw_response.parse()
return headers, response
else:
response = openai_client.audio.transcriptions.create(**data, timeout=timeout)
return None, response
raw_response: Final = openai_client.audio.transcriptions.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response: Final = raw_response.parse()
return headers, response
except Exception as e:
raise e
@ -133,11 +118,12 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
"complete_input_dict": data,
},
)
_, response = self.make_sync_openai_audio_transcriptions_request(
headers, response = self.make_sync_openai_audio_transcriptions_request(
openai_client=openai_client,
data=data,
timeout=timeout,
)
logging_obj.model_call_details["response_headers"] = headers
if isinstance(response, BaseModel):
stringified_response = response.model_dump()
@ -158,6 +144,7 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
hidden_params=hidden_params,
response_type="audio_transcription",
)
set_provider_response_headers_in_hidden_params(final_response, headers)
return final_response
async def async_audio_transcriptions(
@ -217,12 +204,14 @@ class OpenAIAudioTranscription(OpenAIChatCompletion):
actual_model: Final = data.get("model", "whisper-1")
hidden_params: Final = {"model": actual_model, "custom_llm_provider": "openai"}
return convert_to_model_response_object(
final_response: Final[TranscriptionResponse] = convert_to_model_response_object(
response_object=stringified_response,
model_response_object=model_response,
hidden_params=hidden_params,
response_type="audio_transcription",
)
set_provider_response_headers_in_hidden_params(final_response, headers)
return final_response
except Exception as e:
## LOGGING
logging_obj.post_call(

View file

@ -1,7 +1,11 @@
from collections.abc import Mapping
from typing import Any, Final
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Any, Final, Protocol
from urllib.parse import unquote
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm._uuid import uuid
from litellm.llms.vertex_ai.common_utils import (
VertexAIError,
@ -9,35 +13,128 @@ from litellm.llms.vertex_ai.common_utils import (
)
from litellm.types.llms.openai import BatchJobStatus, CreateBatchRequest
from litellm.types.llms.vertex_ai import *
from litellm.types.utils import LiteLLMBatch, PromptTokensDetailsWrapper
from litellm.types.llms.vertex_ai import GenerateContentResponseBody
from litellm.types.utils import LiteLLMBatch, ModelInfo, Usage
_NATIVE_VERTEX_RESPONSE: Final = TypeAdapter(GenerateContentResponseBody)
def vertex_prompt_tokens_details(
usage_metadata: Mapping[str, object],
) -> PromptTokensDetailsWrapper | None:
raw_details: Final = usage_metadata.get("promptTokensDetails")
if not isinstance(raw_details, list):
return None
def _int_field(mapping: Mapping[str, object], key: str) -> int:
value: Final = mapping.get(key)
if isinstance(value, int):
return value
return int(value) if isinstance(value, str) and value.isdigit() else 0
def _normalize(detail: object) -> tuple[str, int] | None:
if not isinstance(detail, Mapping):
def vertex_embedding_prompt_token_count(vertex_response: Mapping[str, object]) -> int:
"""
Prompt tokens billed for one Vertex Gemini Embedding batch row.
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
a fallback.
"""
usage_metadata: Final = vertex_response.get("usageMetadata")
if isinstance(usage_metadata, Mapping):
return _int_field(usage_metadata, "promptTokenCount")
return _int_field(vertex_response, "tokenCount")
def is_vertex_embedding_batch_output_response(response_body: Mapping[str, object]) -> bool:
return isinstance(response_body.get("embedding"), dict)
def is_native_vertex_batch_output_row(row: Mapping[str, object]) -> bool:
return isinstance(row.get("request"), dict)
class NativeVertexBatchCostCalculator(Protocol):
def __call__(
self,
usage: Usage,
model: str,
custom_llm_provider: str | None = None,
model_info: ModelInfo | None = None,
) -> tuple[float, float]: ...
@dataclass(frozen=True, slots=True)
class NativeVertexBatchRowStats:
usage: Usage
total_tokens: int
model: str | None
prompt_cost: float
completion_cost: float
def _native_vertex_row_usage(
response_body: Mapping[str, object],
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
) -> Usage | None:
if "usageMetadata" not in response_body:
if not is_vertex_embedding_batch_output_response(response_body):
return None
modality: Final = detail.get("modality")
token_count: Final = detail.get("tokenCount")
if not isinstance(modality, str) or not isinstance(token_count, int):
return None
return modality.upper(), token_count
parsed_details: Final = tuple(_normalize(detail) for detail in raw_details)
normalized: Final = tuple(detail for detail in parsed_details if detail is not None)
if len(normalized) != len(parsed_details):
prompt_tokens: Final = vertex_embedding_prompt_token_count(response_body)
return Usage(prompt_tokens=prompt_tokens, completion_tokens=0, total_tokens=prompt_tokens)
try:
completion_response: Final = _NATIVE_VERTEX_RESPONSE.validate_python(response_body)
except ValidationError as e:
verbose_logger.debug("vertex_ai batch row response is not a GenerateContentResponse: %s", str(e))
return None
return calculate_usage(completion_response)
return PromptTokensDetailsWrapper(
text_tokens=sum(token_count for modality, token_count in normalized if modality in ("TEXT", "DOCUMENT")),
audio_tokens=sum(token_count for modality, token_count in normalized if modality == "AUDIO"),
image_tokens=sum(token_count for modality, token_count in normalized if modality == "IMAGE"),
video_tokens=sum(token_count for modality, token_count in normalized if modality == "VIDEO"),
def native_vertex_batch_row_stats(
row: Mapping[str, object],
model_name: str | None,
*,
model_info: ModelInfo | None,
calculate_usage: Callable[[GenerateContentResponseBody], Usage],
cost_calculator: NativeVertexBatchCostCalculator,
) -> NativeVertexBatchRowStats | None:
"""
Usage and cost of one native Vertex predictions.jsonl row, a
`{"request": ..., "response": {"candidates": [...], "usageMetadata": {...}, "modelVersion": ...}}`
generateContent object or a `{"request": ..., "response": {"embedding": {...}, "usageMetadata": {...}}}`
embedding object (an embedding row without `usageMetadata` is billed from its documented `tokenCount`).
`model_name` (the deployment model) prices the row unless it is a wildcard, else its own `modelVersion`
does, else the wildcard name so explicit deployment prices still apply; a row without a response, a
generateContent row without `response.usageMetadata`, and a row whose response fails validation are
None (failed).
"""
response_body: Final = row.get("response")
if not isinstance(response_body, dict):
return None
usage: Final = _native_vertex_row_usage(response_body, calculate_usage)
if usage is None:
return None
total_tokens: Final = usage.total_tokens or (usage.prompt_tokens + usage.completion_tokens)
model_version: Final = response_body.get("modelVersion")
deployment_model: Final = model_name if model_name and "*" not in model_name else None
model: Final = deployment_model or (model_version if isinstance(model_version, str) else model_name)
if model is None:
verbose_logger.warning(
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
"is still billed: the row has no modelVersion and the batch has no deployment model"
)
return NativeVertexBatchRowStats(
usage=usage, total_tokens=total_tokens, model=None, prompt_cost=0.0, completion_cost=0.0
)
try:
prompt_cost, completion_cost = cost_calculator(
usage=usage, model=model, custom_llm_provider="vertex_ai", model_info=model_info
)
except Exception as e: # noqa: BLE001 # one unpriceable row must not abort the batch's cost accounting
verbose_logger.warning(
"vertex_ai batch output row could not be costed, so it is billed at $0 and the rest of the batch "
"is still billed. model=%s error=%s",
model,
str(e),
)
return NativeVertexBatchRowStats(
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=0.0, completion_cost=0.0
)
return NativeVertexBatchRowStats(
usage=usage, total_tokens=total_tokens, model=model, prompt_cost=prompt_cost, completion_cost=completion_cost
)
@ -156,30 +253,15 @@ class VertexAIBatchTransformation:
return uris[0]
@classmethod
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str:
def _get_output_file_id_from_vertex_ai_batch_response(cls, response: VertexBatchPredictionResponse) -> str | None:
"""
Gets the output file id from the Vertex AI Batch response
Gets the output file id from the Vertex AI Batch response, None until Vertex reports outputInfo
"""
output_info: Final = response.get("outputInfo") or OutputInfo()
output_file_id: str = output_info.get("gcsOutputDirectory", "")
if output_file_id:
output_file_id = output_file_id.rstrip("/") + "/predictions.jsonl"
if output_file_id and output_file_id != "/predictions.jsonl":
return output_file_id
output_config: Final = response.get("outputConfig")
if output_config is None:
return output_file_id
gcs_destination: Final = output_config.get("gcsDestination")
if gcs_destination is None:
return output_file_id
output_uri_prefix: Final = gcs_destination.get("outputUriPrefix", "")
if output_uri_prefix.endswith("/predictions.jsonl"):
return output_uri_prefix
return output_uri_prefix.rstrip("/") + "/predictions.jsonl"
gcs_output_directory: Final = (output_info.get("gcsOutputDirectory") or "").rstrip("/")
if not gcs_output_directory:
return None
return f"{gcs_output_directory}/predictions.jsonl"
@classmethod
def _get_batch_job_status_from_vertex_ai_batch_response(

View file

@ -9,7 +9,7 @@ from collections.abc import AsyncGenerator, Callable, Iterable, Iterator, Mappin
from contextlib import aclosing
from dataclasses import dataclass
from types import MappingProxyType
from typing import Any, Final, TypedDict
from typing import IO, Any, Final, TypedDict
from urllib.parse import quote, unquote
import httpx
@ -41,6 +41,7 @@ from litellm.llms.base_llm.files.transformation import (
BaseFileUploadStream,
LiteLLMLoggingObj,
)
from litellm.llms.vertex_ai.batches.transformation import vertex_embedding_prompt_token_count
from litellm.llms.vertex_ai.common_utils import (
_convert_vertex_datetime_to_openai_datetime,
get_vertex_ai_fine_tuned_endpoint_id,
@ -56,6 +57,7 @@ from litellm.types.files import StreamingMediaUploadConfig
from litellm.types.llms.openai import (
AllMessageValues,
CreateFileRequest,
FileContent,
FileTypes,
HttpxBinaryResponseContent,
OpenAICreateFileRequestOptionalParams,
@ -87,6 +89,8 @@ _EMBED_REQUEST_FIELD_BY_GEMINI_PARAM: Final = (
_VERTEX_BATCH_FANNED_OUT_KEY_PATTERN: Final = re.compile(r"(?P<custom_id>[^#]*)#(?P<index>\d+)/(?P<total>\d+)")
_JSONL_NEWLINE: Final = b"\n"
_BATCH_OUTPUT_FIRST_ROW_PEEK_LIMIT_BYTES: Final = 32 * 1024 * 1024
_PASSTHROUGH_MANAGED_GCS_PREFIX: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}passthrough/"
_RAW_UPLOAD_CHUNK_BYTES: Final = 1024 * 1024
class _GcsObjectMetadataJson(TypedDict, total=False):
@ -418,19 +422,6 @@ def _split_vertex_batch_key(vertex_output_row: Mapping[str, object]) -> tuple[st
return unquote(match["custom_id"]), int(match["index"]), int(match["total"])
def _embedding_prompt_token_count(vertex_response: _VertexEmbeddingResponse) -> int:
"""
Prompt tokens billed for one Vertex Gemini Embedding batch row.
Live rows report usage under `usageMetadata`; the documented `tokenCount` is kept as
a fallback.
"""
usage_metadata = vertex_response.get("usageMetadata")
if isinstance(usage_metadata, Mapping):
return int(usage_metadata.get("promptTokenCount") or 0)
return int(vertex_response.get("tokenCount") or 0)
def _vertex_embeddings_rows_to_openai_batch_output_row(
custom_id: str,
vertex_output_rows: tuple[_VertexEmbeddingBatchRow, ...],
@ -471,7 +462,7 @@ def _vertex_embeddings_rows_to_openai_batch_output_row(
)
responses = tuple(row["response"] for row in vertex_output_rows)
token_count = sum(_embedding_prompt_token_count(response) for response in responses)
token_count = sum(vertex_embedding_prompt_token_count(response) for response in responses)
body = EmbeddingResponse(
model=model or "",
data=[
@ -528,6 +519,16 @@ def _model_from_managed_gcs_url(url: str) -> str | None:
return match.group(1) if match else None
def is_passthrough_managed_gcs_url(url: str) -> bool:
decoded_url: Final = unquote(url)
managed_prefix_start: Final = decoded_url.find(VERTEX_AI_MANAGED_GCS_PREFIX)
return managed_prefix_start >= 0 and decoded_url.startswith(_PASSTHROUGH_MANAGED_GCS_PREFIX, managed_prefix_start)
def is_passthrough_batch_upload(create_file_data: Mapping[str, object], litellm_params: Mapping[str, object]) -> bool:
return create_file_data.get("purpose") == "batch" and litellm_params.get("passthrough") is True
def _is_embeddings_batch_entry(openai_entry: Mapping[str, object]) -> bool:
"""
Whether an OpenAI batch JSONL line targets the embeddings endpoint.
@ -791,6 +792,58 @@ class _OpenAIToVertexBatchUploadStream(BaseFileUploadStream):
return self._iter_vertex_jsonl_chunks()
def _read_chunk_as_bytes(handle: IO[bytes]) -> bytes:
chunk: Final[bytes | str] = handle.read(_RAW_UPLOAD_CHUNK_BYTES)
return chunk.encode("utf-8") if isinstance(chunk, str) else bytes(chunk)
def _iter_raw_file_chunks(file_content: FileTypes) -> Iterator[bytes]:
content: Final[FileContent | str] = file_content[1] if isinstance(file_content, tuple) else file_content
if isinstance(content, (bytes, bytearray)):
yield from (
bytes(content[offset : offset + _RAW_UPLOAD_CHUNK_BYTES])
for offset in range(0, len(content), _RAW_UPLOAD_CHUNK_BYTES)
)
return
if isinstance(content, str):
yield content.encode("utf-8")
return
if isinstance(content, PathLike):
with open(str(content), "rb") as handle:
yield from iter(lambda: handle.read(_RAW_UPLOAD_CHUNK_BYTES), b"")
return
if not hasattr(content, "read"):
raise ValueError("Unsupported file content type")
seek: Final = getattr(content, "seek", None)
if seek is None:
raise ValueError(
"Batch upload file handle must be seekable; got a non-seekable "
"stream. Pass bytes, a path, or a seekable handle."
)
seek(0)
yield from iter(lambda: _read_chunk_as_bytes(content), b"")
class _RawFileUploadStream(BaseFileUploadStream):
def __init__(self, file_content: FileTypes) -> None:
self._file_content = file_content
def iter_bytes(self) -> Iterator[bytes]:
return _iter_raw_file_chunks(self._file_content)
def _managed_batch_object_name(raw_model: str, *, passthrough: bool) -> str:
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
model_path: Final = (
f"endpoints/{endpoint_id}"
if endpoint_id is not None
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
)
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
prefix: Final = _PASSTHROUGH_MANAGED_GCS_PREFIX if passthrough else VERTEX_AI_MANAGED_GCS_PREFIX
return f"{prefix}{safe_model_path}/{uuid.uuid4()}"
class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
"""
Config for VertexAI Files
@ -848,23 +901,34 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
if deployment_model
else openai_jsonl_content[0].get("body", {}).get("model", "")
)
endpoint_id: Final = get_vertex_ai_fine_tuned_endpoint_id(raw_model)
model_path: Final = (
f"endpoints/{endpoint_id}"
if endpoint_id is not None
else (raw_model if "publishers/google/models" in raw_model else f"publishers/google/models/{raw_model}")
)
safe_model_path: Final = sanitize_cloud_object_path(model_path, fallback="model")
object_name: Final = f"{VERTEX_AI_MANAGED_GCS_PREFIX}{safe_model_path}/{uuid.uuid4()}"
return object_name
return _managed_batch_object_name(raw_model, passthrough=False)
def get_object_name(self, file_data: FileTypes, purpose: str, deployment_model: str | None = None) -> str:
def _get_passthrough_gcs_object_name(self, deployment_model: str | None) -> str:
if not deployment_model:
raise VertexAIError(
status_code=400,
message=(
"Native Vertex batch passthrough uploads need the deployment model to name the GCS object, "
"since native rows carry no model: pass `target_model_names` (proxy) or `model` (SDK)."
),
)
return _managed_batch_object_name(deployment_model.removeprefix("vertex_ai/"), passthrough=True)
def get_object_name(
self,
file_data: FileTypes,
purpose: str,
deployment_model: str | None = None,
passthrough: bool = False,
) -> str:
"""
Get the object name for the request.
Reads only the first JSONL entry (streamed) for batch files, so a large
upload is never materialized just to derive the GCS object name.
"""
if purpose == "batch" and passthrough:
return self._get_passthrough_gcs_object_name(deployment_model)
if purpose == "batch":
## 1. If jsonl, derive the object name from the deployment model (or the first entry's)
first_entry: Final = next(_iter_openai_jsonl_entries(file_data), None)
@ -922,6 +986,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
file_data,
purpose,
deployment_model=configured_model if isinstance(configured_model, str) else None,
passthrough=is_passthrough_batch_upload(data, litellm_params),
)
if object_prefix:
object_name = f"{object_prefix}/{object_name}"
@ -984,6 +1049,14 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
if file_data is None:
raise ValueError("file is required")
if is_passthrough_batch_upload(create_file_data, litellm_params):
return {
"streaming_media_upload": StreamingMediaUploadConfig(
body_stream=_RawFileUploadStream(file_data),
content_type="application/json",
)
}
_, content_type = extract_file_metadata(file_data)
if FilesAPIUtils.is_batch_jsonl_request(
create_file_data=create_file_data,
@ -1164,6 +1237,8 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
# transformation, e.g. if they consume raw `predictions.jsonl` directly.
if getattr(litellm, "disable_vertex_batch_output_transformation", False):
return HttpxBinaryResponseContent(response=raw_response)
if is_passthrough_managed_gcs_url(str(raw_response.request.url)):
return HttpxBinaryResponseContent(response=raw_response)
# Try to transform batch output if it's a JSONL file
content: Final = raw_response.content
@ -1209,7 +1284,7 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig):
Everything else is passed through unchanged, including a row that fails to
transform mid-stream.
"""
if litellm.disable_vertex_batch_output_transformation:
if litellm.disable_vertex_batch_output_transformation or is_passthrough_managed_gcs_url(request_url):
return FileContentStreamingResult(stream_iterator=stream_iterator, headers=headers)
first_line, buffered = await _peek_first_jsonl_line(

View file

@ -68565,6 +68565,7 @@
"source": "https://api.together.ai/v1/models"
},
"vertex_ai/gemini-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
@ -68611,7 +68612,9 @@
"input_cost_per_token": 1.5e-06,
"litellm_provider": "vertex_ai",
"mode": "chat",
"output_cost_per_reasoning_token": 9e-06,
"output_cost_per_token": 9e-06,
"output_cost_per_video_token": 1.75e-05,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-omni-1.1-flash-preview": {
@ -69773,6 +69776,7 @@
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"deprecation_date": "2026-10-31",
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image_token": 5e-06,
"input_cost_per_token": 4e-06,
@ -69804,6 +69808,7 @@
"supports_tool_choice": true
},
"azure/gpt-live-1": {
"deprecation_date": "2027-09-10",
"input_cost_per_second": 0.000833333333333,
"litellm_provider": "azure",
"mode": "realtime",
@ -69821,6 +69826,7 @@
"supports_function_calling": true
},
"azure/gpt-live-transcribe": {
"deprecation_date": "2028-02-01",
"input_cost_per_second": 0.000283333333333,
"litellm_provider": "azure",
"max_input_tokens": 32000,
@ -69842,6 +69848,7 @@
"supports_audio_input": true
},
"azure/gpt-transcribe": {
"deprecation_date": "2028-02-01",
"input_cost_per_second": 7.5e-05,
"litellm_provider": "azure",
"mode": "audio_transcription",
@ -69860,6 +69867,7 @@
"supports_audio_input": true
},
"azure/gpt-realtime-translate": {
"deprecation_date": "2027-05-06",
"input_cost_per_second": 0.000566666666667,
"litellm_provider": "azure",
"max_input_tokens": 32000,

View file

@ -307,6 +307,8 @@ class KeyManagementRoutes(str, enum.Enum):
# team usage routes
TEAM_DAILY_ACTIVITY = "/team/daily/activity"
TEAM_DAILY_ACTIVITY_AGGREGATED = "/team/daily/activity/aggregated"
TEAM_DAILY_ACTIVITY_EXPORT = "/team/daily/activity/export"
TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH = "/team/daily/activity/aggregated/search"
# team spend-log viewing
SPEND_LOGS = "/spend/logs"
@ -673,6 +675,7 @@ class LiteLLMRoutes(enum.Enum):
KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED.value,
KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED_SEARCH.value,
KeyManagementRoutes.SPEND_LOGS.value,
KeyManagementRoutes.SPEND_LOGS_V2.value,
KeyManagementRoutes.KEY_RESET_SPEND.value,
@ -699,6 +702,7 @@ class LiteLLMRoutes(enum.Enum):
"/user/list",
"/user/daily/activity",
"/user/daily/activity/aggregated",
"/user/daily/activity/aggregated/search",
# team
"/team/new",
"/team/update",
@ -716,6 +720,8 @@ class LiteLLMRoutes(enum.Enum):
"/team/permissions_bulk_update",
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/export",
"/team/daily/activity/aggregated/search",
"/team/spend/by_user",
# gateway request counts (SGR); deployment-wide, admin-only
"/gateway/daily/activity",
@ -886,6 +892,8 @@ class LiteLLMRoutes(enum.Enum):
"/team/permissions_update",
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/export",
"/team/daily/activity/aggregated/search",
"/team/spend/by_user",
"/team/{team_id}/members/me",
# POST/GET the team's logging callbacks, and DELETE one of them. Every
@ -901,6 +909,7 @@ class LiteLLMRoutes(enum.Enum):
"/model/delete",
"/user/daily/activity",
"/user/daily/activity/aggregated",
"/user/daily/activity/aggregated/search",
# Endpoint restricts results to organizations the caller is ORG_ADMIN
# of; a caller who administers none gets an empty result set.
"/organization/daily/activity",
@ -984,6 +993,8 @@ class LiteLLMRoutes(enum.Enum):
"/user/daily/activity",
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/export",
"/team/daily/activity/aggregated/search",
"/tag/daily/activity",
"/tag/list",
"/audit",

View file

@ -124,6 +124,12 @@ class _SpendIncrement(TypedDict):
increment: ReadOnly[float]
class _MemberSpendRow(TypedDict):
user_id: ReadOnly[str]
team_id: ReadOnly[str]
cost: ReadOnly[float]
class _SpendBatch(Protocol):
litellm_usertable: BatchTable
litellm_verificationtoken: BatchTable
@ -351,17 +357,22 @@ _TEAM_ADVISORY_LOCK_SQL: Final = "SELECT pg_advisory_xact_lock(hashtext($1)) IS
# One statement adds every member's cost to their membership row. A missing row is created only
# while the user is still on the team's roster, so a spend flush landing after a removal never
# recreates the member.
# recreates the member. The rows travel as one JSON document, not as a numeric array: Prisma
# types a raw array parameter from the first batch a connection sees, so after an all-$0 batch
# (integers) every later fractional batch on that connection failed with "improper binary format".
_TEAM_MEMBER_SPEND_SQL: Final = """
INSERT INTO "LiteLLM_TeamMembership" (user_id, team_id, spend, total_spend)
SELECT p.user_id, p.team_id, p.cost, p.cost
FROM unnest($1::text[], $2::text[], $3::float8[]) AS p(user_id, team_id, cost)
SELECT member.user_id, member.team_id, member.cost, member.cost
FROM jsonb_to_recordset($1::jsonb) AS member(user_id text, team_id text, cost float8)
WHERE EXISTS (
SELECT 1 FROM "LiteLLM_TeamTable" t
WHERE t.team_id = p.team_id
AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))
WHERE t.team_id = member.team_id
AND t.members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', member.user_id))
)
OR EXISTS (
SELECT 1 FROM "LiteLLM_TeamMembership" m
WHERE m.user_id = member.user_id AND m.team_id = member.team_id
)
OR EXISTS (SELECT 1 FROM "LiteLLM_TeamMembership" m WHERE m.user_id = p.user_id AND m.team_id = p.team_id)
ON CONFLICT (user_id, team_id) DO UPDATE
SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend,
total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend
@ -371,15 +382,12 @@ SET spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend,
async def _write_team_member_spend(transaction: _SpendTransaction, spend_by_member_key: Mapping[str, float]) -> None:
# key is "team_id::<value>::user_id::<value>"; locks are taken in sorted team_id order like the team endpoints
rows: Final = sorted((key.split("::")[1], key.split("::")[3], cost) for key, cost in spend_by_member_key.items())
team_ids: Final = tuple(team_id for team_id, _user_id, _cost in rows)
for team_id in dict.fromkeys(team_ids):
for team_id in dict.fromkeys(team_id for team_id, _user_id, _cost in rows):
_ = await transaction.execute_raw(_TEAM_ADVISORY_LOCK_SQL, team_id)
_ = await transaction.execute_raw(
_TEAM_MEMBER_SPEND_SQL,
tuple(user_id for _team_id, user_id, _cost in rows),
team_ids,
tuple(cost for _team_id, _user_id, cost in rows),
members: Final = tuple(
_MemberSpendRow(user_id=user_id, team_id=team_id, cost=cost) for team_id, user_id, cost in rows
)
_ = await transaction.execute_raw(_TEAM_MEMBER_SPEND_SQL, json.dumps(members))
def get_llm_router():

View file

@ -1,4 +1,6 @@
import asyncio
import dataclasses
import itertools
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from datetime import datetime, timedelta, timezone
@ -35,6 +37,10 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import (
SpendAnalyticsPaginatedResponse,
SpendMetrics,
)
from litellm.types.proxy.management_endpoints.team_endpoints import (
TeamDailyActivityExportRow,
TeamDailyActivityExportType,
)
if TYPE_CHECKING:
from prisma.models import (
@ -198,7 +204,7 @@ class _AggregatedQueryKwargs(TypedDict):
include_current_utc_day: ReadOnly[bool]
_SqlQuery = tuple[str, list[str]]
_SqlQuery = tuple[str, Sequence[str]]
async def _query_raw_optional(
@ -974,6 +980,291 @@ def _build_entity_rollup_sql_query(
return sql_query, sql_params
def _build_export_sql_query(
*,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
start_date: str,
end_date: str,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None,
export_type: TeamDailyActivityExportType,
) -> tuple[str, tuple[str, ...]]:
"""One unbounded rollup for the export route, on the aggregated path's WHERE clause.
No LIMIT anywhere: the export exists so a caller can reach keys past
USAGE_TOP_API_KEYS_LIMIT. PTU sentinel rows stay in `daily` so per-team
totals match breakdown.entities, and are excluded from the key, user and
model exports where the flat-cost row has no meaning.
"""
pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name)
if pg_table is None:
raise ValueError(f"Unknown table name: {table_name}")
adjusted_start, adjusted_end = _adjust_dates_for_timezone(start_date, end_date, timezone_offset_minutes)
where_clause, where_params = _build_aggregated_where_clause(
entity_id_field=entity_id_field,
entity_id=entity_id,
adjusted_start=adjusted_start,
adjusted_end=adjusted_end,
model=None,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
)
keyed: Final = export_type in ("daily_with_keys", "daily_with_users")
by_model: Final = export_type == "daily_with_models"
group_extras: Final = tuple(field for field in ("api_key" if keyed else "", "model" if by_model else "") if field)
group_by: Final = f'date, "{entity_id_field}"' + "".join(f", {field}" for field in group_extras)
sentinel_clause: Final = f" AND api_key <> ${len(where_params) + 1}" if (keyed or by_model) else ""
sentinel_params: Final = (PTU_SENTINEL_API_KEY,) if (keyed or by_model) else ()
sql_query: Final = f"""
SELECT
date,
"{entity_id_field}" AS entity_id,
{"api_key" if keyed else "NULL::text AS api_key"},
{"model" if by_model else "NULL::text AS model"},{_rollup_metric_select(table_name)}
FROM "{pg_table}"
WHERE {where_clause}{sentinel_clause}
GROUP BY {group_by}
ORDER BY {group_by}
"""
return sql_query, (*where_params, *sentinel_params)
class _ExportRow(_RollupMetricsRow):
entity_id: str | None
model: str | None
def _export_team_alias(entity_metadata_field: Mapping[str, dict[str, object]] | None, entity_id: str) -> str | None:
alias: Final = _entity_metadata(entity_metadata_field, entity_id).get("team_alias")
return alias if isinstance(alias, str) else None
@dataclasses.dataclass(frozen=True, slots=True)
class _ExportMetrics:
spend: float
api_requests: int
successful_requests: int
failed_requests: int
total_tokens: int
prompt_tokens: int
completion_tokens: int
cache_read_input_tokens: int
cache_creation_input_tokens: int
@classmethod
def from_record(cls, record: _RollupMetricsRow) -> "_ExportMetrics":
prompt_tokens: Final = record.prompt_tokens or 0
completion_tokens: Final = record.completion_tokens or 0
return cls(
spend=record.spend or 0.0,
api_requests=record.api_requests or 0,
successful_requests=record.successful_requests or 0,
failed_requests=record.failed_requests or 0,
total_tokens=prompt_tokens + completion_tokens,
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
cache_read_input_tokens=record.cache_read_input_tokens or 0,
cache_creation_input_tokens=record.cache_creation_input_tokens or 0,
)
@classmethod
def zero(cls) -> "_ExportMetrics":
return cls(
spend=0.0,
api_requests=0,
successful_requests=0,
failed_requests=0,
total_tokens=0,
prompt_tokens=0,
completion_tokens=0,
cache_read_input_tokens=0,
cache_creation_input_tokens=0,
)
def __add__(self, other: "_ExportMetrics") -> "_ExportMetrics":
return _ExportMetrics(
spend=self.spend + other.spend,
api_requests=self.api_requests + other.api_requests,
successful_requests=self.successful_requests + other.successful_requests,
failed_requests=self.failed_requests + other.failed_requests,
total_tokens=self.total_tokens + other.total_tokens,
prompt_tokens=self.prompt_tokens + other.prompt_tokens,
completion_tokens=self.completion_tokens + other.completion_tokens,
cache_read_input_tokens=self.cache_read_input_tokens + other.cache_read_input_tokens,
cache_creation_input_tokens=self.cache_creation_input_tokens + other.cache_creation_input_tokens,
)
def _export_base_row(
record: _ExportRow,
entity_metadata_field: Mapping[str, dict[str, object]] | None,
) -> TeamDailyActivityExportRow:
entity_id: Final = record.entity_id or "Unassigned"
metrics: Final = _ExportMetrics.from_record(record)
return TeamDailyActivityExportRow(
date=record.date,
team_id=entity_id,
team_alias=_export_team_alias(entity_metadata_field, entity_id),
model=record.model,
spend=metrics.spend,
flat_cost=_reported_flat_cost(record),
api_requests=metrics.api_requests,
successful_requests=metrics.successful_requests,
failed_requests=metrics.failed_requests,
total_tokens=metrics.total_tokens,
prompt_tokens=metrics.prompt_tokens,
completion_tokens=metrics.completion_tokens,
cache_read_input_tokens=metrics.cache_read_input_tokens,
cache_creation_input_tokens=metrics.cache_creation_input_tokens,
)
def _export_key_row(
record: _ExportRow,
entity_metadata_field: Mapping[str, dict[str, object]] | None,
api_key_metadata: Mapping[str, _KeyMetadataDict],
) -> TeamDailyActivityExportRow:
entity_id: Final = record.entity_id or "Unassigned"
metadata: Final = _key_metadata(api_key_metadata, record.api_key or "")
metrics: Final = _ExportMetrics.from_record(record)
return TeamDailyActivityExportRow(
date=record.date,
team_id=entity_id,
team_alias=_export_team_alias(entity_metadata_field, entity_id),
api_key=record.api_key,
key_alias=metadata.key_alias,
user_id=metadata.user_id,
user_email=metadata.user_email,
spend=metrics.spend,
api_requests=metrics.api_requests,
successful_requests=metrics.successful_requests,
failed_requests=metrics.failed_requests,
total_tokens=metrics.total_tokens,
prompt_tokens=metrics.prompt_tokens,
completion_tokens=metrics.completion_tokens,
cache_read_input_tokens=metrics.cache_read_input_tokens,
cache_creation_input_tokens=metrics.cache_creation_input_tokens,
)
def _fold_export_users(
records: Sequence[_ExportRow],
entity_metadata_field: Mapping[str, dict[str, object]] | None,
api_key_metadata: Mapping[str, _KeyMetadataDict],
) -> tuple[TeamDailyActivityExportRow, ...]:
"""Fold (date, team, api_key) rows into (date, team, user) rows."""
def bucket_of(record: _ExportRow) -> tuple[str, str, str]:
return (
record.date,
record.entity_id or "Unassigned",
_key_metadata(api_key_metadata, record.api_key or "").user_id or "Unassigned",
)
key_sets: Final = MappingProxyType(
{
bucket: frozenset(record.api_key or "" for record in group)
for bucket, group in itertools.groupby(sorted(records, key=bucket_of), key=bucket_of)
}
)
sums: Final[dict[tuple[str, str, str], _ExportMetrics]] = {} # mutable-ok: local fold accumulator
emails: Final[dict[tuple[str, str, str], str | None]] = {} # mutable-ok: local fold accumulator
for record in records:
metadata = _key_metadata(api_key_metadata, record.api_key or "")
bucket_key = bucket_of(record)
sums[bucket_key] = sums.get(bucket_key, _ExportMetrics.zero()) + _ExportMetrics.from_record(record)
emails.setdefault(bucket_key, metadata.user_email)
if emails[bucket_key] is None and metadata.user_email is not None:
emails[bucket_key] = metadata.user_email
return tuple(
_export_folded_user_row(
bucket_key, sums[bucket_key], emails[bucket_key], len(key_sets[bucket_key]), entity_metadata_field
)
for bucket_key in sorted(sums)
)
def _export_folded_user_row(
bucket_key: tuple[str, str, str],
metrics: _ExportMetrics,
user_email: str | None,
keys: int,
entity_metadata_field: Mapping[str, dict[str, object]] | None,
) -> TeamDailyActivityExportRow:
date, entity_id, user_id = bucket_key
return TeamDailyActivityExportRow(
date=date,
team_id=entity_id,
team_alias=_export_team_alias(entity_metadata_field, entity_id),
user_id=user_id if user_id != "Unassigned" else None,
user_email=user_email,
keys=keys,
spend=metrics.spend,
api_requests=metrics.api_requests,
successful_requests=metrics.successful_requests,
failed_requests=metrics.failed_requests,
total_tokens=metrics.total_tokens,
prompt_tokens=metrics.prompt_tokens,
completion_tokens=metrics.completion_tokens,
cache_read_input_tokens=metrics.cache_read_input_tokens,
cache_creation_input_tokens=metrics.cache_creation_input_tokens,
)
async def get_daily_activity_export_rows(
*,
prisma_client: PrismaClient,
table_name: str,
entity_id_field: str,
entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
entity_metadata_field: Mapping[str, dict[str, object]] | None,
start_date: str,
end_date: str,
api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path
exclude_entity_ids: list[str] | None, # mutable-ok: filter union shared with the paginated path
timezone_offset_minutes: int | None,
export_type: TeamDailyActivityExportType,
) -> tuple[TeamDailyActivityExportRow, ...]:
"""Every (date, entity[, api_key|model]) rollup row in the range, uncapped."""
sql_query, sql_params = _build_export_sql_query(
table_name=table_name,
entity_id_field=entity_id_field,
entity_id=entity_id,
start_date=start_date,
end_date=end_date,
api_key=api_key,
exclude_entity_ids=exclude_entity_ids,
timezone_offset_minutes=timezone_offset_minutes,
export_type=export_type,
)
raw_rows: Final = await _query_raw_optional(prisma_client, (sql_query, sql_params))
records: Final = tuple(_ExportRow(**row) for row in (raw_rows or ()))
if export_type in ("daily", "daily_with_models"):
return await asyncio.to_thread(
lambda: tuple(_export_base_row(record, entity_metadata_field) for record in records)
)
api_keys: Final = frozenset(record.api_key for record in records if record.api_key)
api_key_metadata: Final = (
await get_api_key_metadata(prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records)))
if api_keys
else _EMPTY_KEY_METADATA
)
if export_type == "daily_with_keys":
return await asyncio.to_thread(
lambda: tuple(_export_key_row(record, entity_metadata_field, api_key_metadata) for record in records)
)
return await asyncio.to_thread(_fold_export_users, records, entity_metadata_field, api_key_metadata)
def _aggregate_spend_records_sync(
*,
records: Sequence[DailySpendRecord],

View file

@ -28,6 +28,7 @@ from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import (
@ -87,11 +88,13 @@ from litellm.repositories.verification_token_repository import (
VerificationTokenRepository,
)
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
BulkUpdateUserResponse,
KeyActivitySearchWhere,
UserListResponse,
UserSearchWhere,
UserUpdateResult,
@ -2991,6 +2994,27 @@ async def get_user_daily_activity(
)
def _resolve_user_daily_activity_entity_id(
user_api_key_dict: UserAPIKeyAuth,
user_id: str | None,
) -> str | None:
is_admin: Final = _user_has_admin_view(user_api_key_dict)
if is_admin:
return user_id
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
effective_user_id: Final = user_id if user_id is not None else caller_user_id
if effective_user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={ # mutable-ok: FastAPI detail payload shape
"error": "Non-admin users can only view their own spend data."
},
)
return effective_user_id
@router.get(
"/user/daily/activity/aggregated",
tags=["Budget & Spend Tracking", "Internal User management"],
@ -3057,20 +3081,7 @@ async def get_user_daily_activity_aggregated(
)
try:
is_admin: Final = _user_has_admin_view(user_api_key_dict)
if is_admin:
entity_id = user_id # None means global view, otherwise filter by user
else:
caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict)
if user_id is None:
user_id = caller_user_id
if user_id != caller_user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"error": "Non-admin users can only view their own spend data."},
)
entity_id = user_id
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
@ -3094,3 +3105,117 @@ async def get_user_daily_activity_aggregated(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to fetch analytics: {e}"},
)
@router.get(
"/user/daily/activity/aggregated/search",
tags=["Budget & Spend Tracking", "Internal User management"], # mutable-ok: FastAPI route tags shape
dependencies=[Depends(user_api_key_auth)], # mutable-ok: FastAPI route dependencies shape
response_model=SpendAnalyticsPaginatedResponse,
)
@management_endpoint_wrapper
async def search_user_daily_activity_keys(
search: str = fastapi.Query(
...,
min_length=1,
description="Matches keys whose hash equals the value, or whose key alias or user ID contains it (case-insensitive)",
),
start_date: str | None = fastapi.Query(
default=None,
description="Start date in YYYY-MM-DD format",
),
end_date: str | None = fastapi.Query(
default=None,
description="End date in YYYY-MM-DD format",
),
user_id: str | None = fastapi.Query(
default=None,
description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.",
),
timezone: int | None = fastapi.Query(
default=None,
description="Timezone offset in minutes from UTC (e.g., 480 for PST). "
"Matches JavaScript's Date.getTimezoneOffset() convention.",
),
include_current_utc_day: bool = fastapi.Query(
default=False,
description="When the range ends on the caller's current local day, extend it to "
"today's UTC bucket so spend written after the caller's local midnight (in UTC "
"terms) is included. Requires the timezone parameter. Historical ranges are "
"never extended.",
),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI dependency injection
) -> SpendAnalyticsPaginatedResponse:
"""
Search verification tokens by exact token hash or by a case-insensitive substring of
the key alias or owning user ID, then return the aggregated daily activity for the
matches. Lets the Usage page surface keys that fell outside the top-spend subset
the aggregated endpoint loads.
"""
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={ # mutable-ok: FastAPI detail payload shape
"error": CommonProxyErrors.db_not_connected_error.value
},
)
if start_date is None or end_date is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": "Please provide start_date and end_date"}, # mutable-ok: FastAPI detail payload shape
)
try:
entity_id: Final = _resolve_user_daily_activity_entity_id(user_api_key_dict, user_id)
search_or: Final = (
{"token": search}, # mutable-ok: prisma serializes where clauses, keep plain dicts
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
)
where: Final[KeyActivitySearchWhere] = (
{"OR": search_or} # mutable-ok: prisma where clause root
if entity_id is None
else {"user_id": entity_id, "OR": search_or} # mutable-ok: prisma where clause root
)
matched_keys: Final = await VerificationTokenRepository(prisma_client).table.find_many(
where=where,
take=USAGE_TOP_API_KEYS_LIMIT,
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
)
tokens: Final = [key.token for key in matched_keys] # mutable-ok: api_key filter union expects a list
if not tokens:
return SpendAnalyticsPaginatedResponse(
results=[], # mutable-ok: response model field shape
metadata=DailySpendMetadata(
api_key_limit=USAGE_TOP_API_KEYS_LIMIT,
total_api_keys=0,
),
)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=entity_id,
entity_metadata_field=None,
start_date=start_date,
end_date=end_date,
model=None,
api_key=tokens,
timezone_offset_minutes=timezone,
include_current_utc_day=include_current_utc_day,
)
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.exception("/user/daily/activity/aggregated/search: Exception occured - %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail={"error": f"Failed to fetch analytics: {e}"}, # mutable-ok: FastAPI detail payload shape
)

View file

@ -11,6 +11,8 @@ All /team management endpoints
import asyncio
import copy
import csv
import io
import json
import math
import traceback
@ -33,13 +35,15 @@ from typing import (
)
import fastapi
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, Response, status
from fastapi.responses import JSONResponse
from pydantic import BaseModel, JsonValue, TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict, assert_never
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
from litellm.integrations.prometheus import PrometheusLogger
from litellm.litellm_core_utils.ptu_pricing import is_ptu_cost_attribution_enabled
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
@ -126,6 +130,7 @@ from litellm.proxy.hooks.model_max_budget_limiter import (
)
from litellm.proxy.management_endpoints.common_daily_activity import (
get_daily_activity_aggregated,
get_daily_activity_export_rows,
)
from litellm.proxy.management_endpoints.common_utils import (
_check_disable_global_guardrails_caller_permission,
@ -199,6 +204,7 @@ from litellm.router import Router
from litellm.router_utils.ptu_shares import model_group_ptu_capacity
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendAnalyticsPaginatedResponse,
)
from litellm.types.proxy.management_endpoints.team_endpoints import (
@ -207,7 +213,14 @@ from litellm.types.proxy.management_endpoints.team_endpoints import (
BulkUpdateTeamMemberPermissionsRequest,
BulkUpdateTeamMemberPermissionsResponse,
GetTeamMemberPermissionsResponse,
TeamDailyActivityExportFormat,
TeamDailyActivityExportMetadata,
TeamDailyActivityExportResponse,
TeamDailyActivityExportRow,
TeamDailyActivityExportType,
TeamIdSearchFilter,
TeamIdSearchMatch,
TeamKeyActivitySearchWhere,
TeamListItem,
TeamListResponse,
TeamMemberAddResult,
@ -6823,6 +6836,283 @@ async def get_team_daily_activity_aggregated(
return _with_ptu_consumption(activity, llm_router)
_EXPORT_CSV_METRIC_HEADERS: Final = (
"Spend ($)",
"Requests",
"Successful Requests",
"Failed Requests",
"Total Tokens",
"Prompt Tokens",
"Completion Tokens",
"Cache Read Input Tokens",
"Cache Creation Input Tokens",
)
def _export_csv_headers(export_type: TeamDailyActivityExportType) -> tuple[str, ...]:
base: Final = ("Date", "Team", "Team ID")
if export_type == "daily_with_keys":
return (*base, "Key Alias", "Key ID", "User ID", "User Email", *_EXPORT_CSV_METRIC_HEADERS)
if export_type == "daily_with_users":
return (*base, "User ID", "User Email", "Keys", *_EXPORT_CSV_METRIC_HEADERS)
if export_type == "daily_with_models":
return (
*base,
"Model",
"Spend ($)",
"Requests",
"Successful",
"Failed",
"Total Tokens",
"Prompt Tokens",
"Completion Tokens",
"Cache Read Input Tokens",
"Cache Creation Input Tokens",
)
return (*base, *_EXPORT_CSV_METRIC_HEADERS)
def _csv_safe(value: str) -> str:
return "'" + value if value[:1] in ("=", "+", "-", "@", "\t", "\r") else value
def _export_csv_record(row: TeamDailyActivityExportRow) -> dict[str, object]:
return { # mutable-ok: csv.DictWriter consumes a plain mapping per row
"Date": row.date,
"Team": _csv_safe(row.team_alias) if row.team_alias else "-",
"Team ID": row.team_id,
"Key Alias": _csv_safe(row.key_alias) if row.key_alias else "-",
"Key ID": row.api_key or "-",
"User ID": _csv_safe(row.user_id) if row.user_id else "-",
"User Email": _csv_safe(row.user_email) if row.user_email else "-",
"Keys": row.keys,
"Model": _csv_safe(row.model) if row.model else "-",
"Spend ($)": f"{row.spend:.4f}",
"Flat Cost ($)": f"{row.flat_cost:.4f}",
"Total Cost ($)": f"{row.spend + row.flat_cost:.4f}",
"Requests": row.api_requests,
"Successful Requests": row.successful_requests,
"Failed Requests": row.failed_requests,
"Successful": row.successful_requests,
"Failed": row.failed_requests,
"Total Tokens": row.total_tokens,
"Prompt Tokens": row.prompt_tokens,
"Completion Tokens": row.completion_tokens,
"Cache Read Input Tokens": row.cache_read_input_tokens,
"Cache Creation Input Tokens": row.cache_creation_input_tokens,
}
def _team_export_csv(export_type: TeamDailyActivityExportType, rows: Sequence[TeamDailyActivityExportRow]) -> str:
base_headers: Final = _export_csv_headers(export_type)
spend_index: Final = base_headers.index("Spend ($)") + 1
headers: Final = (
(*base_headers[:spend_index], "Flat Cost ($)", "Total Cost ($)", *base_headers[spend_index:])
if sum(row.flat_cost for row in rows) > 0
else base_headers
)
buffer: Final = io.StringIO()
writer: Final = csv.DictWriter(buffer, fieldnames=headers, extrasaction="ignore")
writer.writeheader()
writer.writerows(_export_csv_record(row) for row in rows)
return buffer.getvalue()
@router.get(
"/team/daily/activity/export",
response_model=TeamDailyActivityExportResponse,
responses={200: {"content": {"text/csv": {}, "application/json": {}}}}, # mutable-ok: OpenAPI content map
tags=["team management"], # mutable-ok: fastapi's decorator signature types tags as a list
)
async def get_team_daily_activity_export(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
start_date: str | None = None,
end_date: str | None = None,
export_type: TeamDailyActivityExportType = "daily",
format: TeamDailyActivityExportFormat = "csv",
team_id: str | None = None,
exclude_team_ids: str | None = None,
timezone_offset: Annotated[int | None, Query(alias="timezone")] = None,
) -> Response:
"""
Server-side Team Usage export, not subject to USAGE_TOP_API_KEYS_LIMIT.
Same scoping as /team/daily/activity/aggregated, answered by one unbounded
rollup query, returned as CSV or JSON. For daily_with_keys,
daily_with_users and daily_with_models the PTU sentinel flat-cost rows are
excluded, so metadata totals under those export types cover request spend
only; the plain daily export includes them.
"""
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value)
range_error: Final = _aggregated_date_range_error(start_date, end_date)
if range_error is not None or start_date is None or end_date is None:
raise _daily_activity_error(status_code=400, message=range_error or "Please provide start_date and end_date")
scope: Final = await _resolve_team_daily_activity_scope(
team_ids=team_id,
exclude_team_ids=exclude_team_ids,
api_key=None,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
rows: Final = await get_daily_activity_export_rows(
prisma_client=prisma_client,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=scope.team_ids,
entity_metadata_field=scope.team_alias_metadata,
start_date=start_date,
end_date=end_date,
api_key=scope.api_key_filter,
exclude_entity_ids=scope.exclude_team_ids,
timezone_offset_minutes=timezone_offset,
export_type=export_type,
)
now: Final = datetime.now(timezone.utc)
metadata: Final = TeamDailyActivityExportMetadata(
export_date=now.isoformat(),
export_type=export_type,
start_date=start_date,
end_date=end_date,
team_ids=list(scope.team_ids) if scope.team_ids else None, # mutable-ok: response model field type
total_spend=sum(row.spend for row in rows),
total_flat_cost=sum(row.flat_cost for row in rows),
total_api_requests=sum(row.api_requests for row in rows),
total_successful_requests=sum(row.successful_requests for row in rows),
total_failed_requests=sum(row.failed_requests for row in rows),
total_tokens=sum(row.total_tokens for row in rows),
)
if format == "json":
return JSONResponse(
content=TeamDailyActivityExportResponse(metadata=metadata, data=rows).model_dump(mode="json")
)
return Response(
content=_team_export_csv(export_type, rows),
media_type="text/csv; charset=utf-8",
headers={ # mutable-ok: starlette Response headers is a dict
"Content-Disposition": f'attachment; filename="team_usage_{export_type}_{now.date().isoformat()}.csv"'
},
)
def _team_key_search_where(*, search: str, scope: _TeamDailyActivityScope) -> TeamKeyActivitySearchWhere:
"""Caller scoping lives inside the same Prisma where as the search term so `take`
never trims visible matches in favour of keys the caller is not allowed to see."""
search_or: Final = (
{"token": search}, # mutable-ok: prisma where clause leaf
{"key_alias": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
{"user_id": {"contains": search, "mode": "insensitive"}}, # mutable-ok: prisma where clause leaf
)
own_keys: Final = tuple(scope.api_key_filter) if isinstance(scope.api_key_filter, list) else None
team_filter: Final[TeamIdSearchFilter | None] = (
{ # mutable-ok: prisma where clause leaf
"in": tuple(scope.team_ids),
"notIn": tuple(scope.exclude_team_ids),
}
if scope.team_ids is not None and scope.exclude_team_ids is not None
else {"in": tuple(scope.team_ids)} # mutable-ok: prisma where clause leaf
if scope.team_ids is not None
else {"notIn": tuple(scope.exclude_team_ids)} # mutable-ok: prisma where clause leaf
if scope.exclude_team_ids is not None
else None
)
if team_filter is None and own_keys is None:
return {"OR": search_or} # mutable-ok: prisma where clause root
if team_filter is None and own_keys is not None:
return {"token": {"in": own_keys}, "OR": search_or} # mutable-ok: prisma where clause root
if team_filter is not None and own_keys is None:
return {"team_id": team_filter, "OR": search_or} # mutable-ok: prisma where clause root
assert team_filter is not None and own_keys is not None
return { # mutable-ok: prisma where clause root
"team_id": team_filter,
"token": {"in": own_keys}, # mutable-ok: prisma where clause leaf
"OR": search_or,
}
@router.get(
"/team/daily/activity/aggregated/search",
response_model=SpendAnalyticsPaginatedResponse,
tags=["team management"], # mutable-ok: FastAPI route tags shape
)
async def search_team_daily_activity_keys(
user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)],
search: str = fastapi.Query(
...,
min_length=1,
description="Exact token hash, or a case-insensitive substring of the key alias or owning user id",
),
team_ids: str | None = None,
start_date: str | None = None,
end_date: str | None = None,
exclude_team_ids: str | None = None,
timezone: int | None = None,
) -> SpendAnalyticsPaginatedResponse:
"""Aggregated daily team activity for the keys matching `search`, across every key the caller may
see rather than only the top USAGE_TOP_API_KEYS_LIMIT keys by spend."""
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if prisma_client is None:
raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value)
range_error: Final = _aggregated_date_range_error(start_date, end_date)
if range_error is not None:
raise _daily_activity_error(status_code=400, message=range_error)
scope: Final = await _resolve_team_daily_activity_scope(
team_ids=team_ids,
exclude_team_ids=exclude_team_ids,
api_key=None,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
matched_keys: Final = await _tokens_db(prisma_client).find_many(
where=_team_key_search_where(search=search, scope=scope),
take=USAGE_TOP_API_KEYS_LIMIT,
order={"spend": "desc"}, # mutable-ok: prisma serializes order, keep it a plain dict
)
tokens: Final = [key.token for key in matched_keys] # mutable-ok: get_daily_activity_aggregated takes list[str]
if not tokens:
return SpendAnalyticsPaginatedResponse(
results=[], # mutable-ok: response model field shape
metadata=DailySpendMetadata(api_key_limit=USAGE_TOP_API_KEYS_LIMIT, total_api_keys=0),
)
return await get_daily_activity_aggregated(
prisma_client=prisma_client,
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id=scope.team_ids,
entity_metadata_field=scope.team_alias_metadata,
start_date=start_date,
end_date=end_date,
model=None,
api_key=tokens,
exclude_entity_ids=scope.exclude_team_ids,
timezone_offset_minutes=timezone,
include_entity_breakdown=True,
)
def _team_user_spend_sql(*, team_count: int, restrict_to_user: bool) -> str:
team_placeholders: Final = ", ".join(f"${i}" for i in range(3, 3 + team_count))
user_clause: Final = f' AND sl."user" = ${3 + team_count}' if restrict_to_user else ""

View file

@ -8,10 +8,25 @@ from typing_extensions import assert_never
from litellm.proxy._types import ProxyException
BATCH_LINE_REQUIRED_KEYS: Final = ("custom_id", "method", "url", "body")
_MB: Final = 1024 * 1024
@dataclass(frozen=True, slots=True)
class BatchLineShape:
required_keys: tuple[str, ...]
hint: str
BATCH_LINE_SHAPE: Final = BatchLineShape(
required_keys=("custom_id", "method", "url", "body"),
hint="Each line must be a JSON object with keys custom_id, method, url, body",
)
PASSTHROUGH_BATCH_LINE_SHAPE: Final = BatchLineShape(
required_keys=("request",),
hint="A passthrough upload takes native Vertex batch rows, so each line must be a JSON object with a request key",
)
@dataclass(frozen=True, slots=True)
class BatchFileTooLarge:
size_bytes: int
@ -42,6 +57,7 @@ class BatchFileLineNotObject:
class BatchFileMissingLineKey:
line_number: int
key: str
line_shape: BatchLineShape = BATCH_LINE_SHAPE
BatchFileValidationFailure = (
@ -70,20 +86,20 @@ def _iter_lines(file_source: bytes | BinaryIO) -> Iterator[bytes]:
return iter(file_source)
def _check_line(line_number: int, raw_line: bytes) -> BatchFileValidationFailure | None:
def _check_line(line_number: int, raw_line: bytes, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
try:
parsed: Final = json.loads(raw_line)
except (json.JSONDecodeError, UnicodeDecodeError):
return BatchFileInvalidJsonLine(line_number=line_number)
if not isinstance(parsed, dict):
return BatchFileLineNotObject(line_number=line_number)
missing: Final = next((key for key in BATCH_LINE_REQUIRED_KEYS if key not in parsed), None)
missing: Final = next((key for key in line_shape.required_keys if key not in parsed), None)
if missing is None:
return None
return BatchFileMissingLineKey(line_number=line_number, key=missing)
return BatchFileMissingLineKey(line_number=line_number, key=missing, line_shape=line_shape)
def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | None:
def _scan_lines(file_source: bytes | BinaryIO, line_shape: BatchLineShape) -> BatchFileValidationFailure | None:
content_lines: Final = (
(line_number, raw_line)
for line_number, raw_line in enumerate(_iter_lines(file_source), start=1)
@ -96,7 +112,7 @@ def _scan_lines(file_source: bytes | BinaryIO) -> BatchFileValidationFailure | N
(
failure
for line_number, raw_line in chain((first_line,), content_lines)
for failure in (_check_line(line_number, raw_line),)
for failure in (_check_line(line_number, raw_line, line_shape),)
if failure is not None
),
None,
@ -107,6 +123,7 @@ def check_batch_file_upload(
filename: str | None,
file_source: bytes | BinaryIO,
max_batch_file_size_mb: int | None,
line_shape: BatchLineShape = BATCH_LINE_SHAPE,
) -> BatchFileValidationFailure | None:
if filename is None or not filename.lower().endswith(".jsonl"):
return BatchFileWrongExtension(filename=filename or "")
@ -114,7 +131,7 @@ def check_batch_file_upload(
size_bytes: Final = _file_size_bytes(file_source)
if size_bytes > max_batch_file_size_mb * _MB:
return BatchFileTooLarge(size_bytes=size_bytes, limit_mb=max_batch_file_size_mb)
scan_failure: Final = _scan_lines(file_source)
scan_failure: Final = _scan_lines(file_source, line_shape)
if not isinstance(file_source, bytes):
file_source.seek(0)
return scan_failure
@ -169,11 +186,11 @@ def raise_batch_file_validation_failure(failure: BatchFileValidationFailure) ->
param="file",
code=400,
)
case BatchFileMissingLineKey(line_number=line_number, key=key):
case BatchFileMissingLineKey(line_number=line_number, key=key, line_shape=line_shape):
raise ProxyException(
message=(
f"Missing required parameter: '{key}' (batch input file line {line_number}). "
f"Each line must be a JSON object with keys {', '.join(BATCH_LINE_REQUIRED_KEYS)}. "
f"{line_shape.hint}. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",

View file

@ -57,6 +57,8 @@ from litellm.proxy.common_utils.openai_error_payload import (
openai_error_type,
)
from litellm.proxy.openai_files_endpoints.batch_file_validation import (
BATCH_LINE_SHAPE,
PASSTHROUGH_BATCH_LINE_SHAPE,
check_batch_file_upload,
raise_batch_file_validation_failure,
)
@ -207,10 +209,91 @@ def get_files_provider_config(
return None
def _deployment_provider(llm_router: Router, model_id: str, team_id: str | None) -> str | None:
credentials: Final = llm_router.get_deployment_credentials_with_provider(model_id=model_id, team_id=team_id)
return None if credentials is None else credentials.get("custom_llm_provider")
def _resolves_to_vertex_deployments_only(llm_router: Router | None, model_name: str, team_id: str | None) -> bool:
if llm_router is None or _deployment_provider(llm_router, model_name, team_id) != "vertex_ai":
return False
return all(
_deployment_provider(llm_router, str(deployment["model_info"]["id"]), team_id) == "vertex_ai"
for deployment in llm_router.get_model_list(model_name=model_name, team_id=team_id) or ()
if "id" in deployment.get("model_info", {})
)
def _validate_passthrough_upload(
*,
purpose: str,
target_model_names: Sequence[str],
model: str | None,
target_storage: str | None,
llm_router: Router | None,
team_id: str | None,
) -> None:
if purpose != "batch":
raise ProxyException(
message=(
"`passthrough` uploads the file bytes unchanged for a native Vertex batch, "
f"so purpose must be 'batch', got '{purpose}'."
),
type="invalid_request_error",
param="passthrough",
code=400,
)
if target_storage and target_storage != "default":
raise ProxyException(
message=(
"`passthrough` writes the native batch file to the Vertex AI deployment's GCS bucket, "
f"so it cannot be combined with target_storage='{target_storage}'."
),
type="invalid_request_error",
param="target_storage",
code=400,
)
named_deployments: Final = (
*(("target_model_names", name) for name in target_model_names),
*((("model", model),) if model else ()),
)
if not named_deployments:
raise ProxyException(
message=(
"`passthrough` needs the Vertex AI deployment that will run the batch, "
"since native rows carry no model: pass `target_model_names` or `model`."
),
type="invalid_request_error",
param="target_model_names",
code=400,
)
offending: Final = next(
(
(param, name)
for param, name in named_deployments
if not _resolves_to_vertex_deployments_only(llm_router, name, team_id)
),
None,
)
if offending is None:
return
param, name = offending
raise ProxyException(
message=(
f"`passthrough` is only supported for Vertex AI deployments; '{name}' does not resolve "
"to vertex_ai deployments only."
),
type="invalid_request_error",
param=param,
code=400,
)
async def _scan_batch_upload(
*,
file_source: bytes | BinaryIO,
purpose: str,
passthrough: bool,
request_metadata: Mapping[str, object],
user_api_key_dict: UserAPIKeyAuth,
proxy_logging_obj: ProxyLogging,
@ -222,6 +305,17 @@ async def _scan_batch_upload(
or not proxy_logging_obj.has_pre_call_guardrails(request_metadata)
):
return None
if passthrough:
raise ProxyException(
message=(
"Batch guardrails cannot scan native Vertex batch rows, so a `passthrough` upload is refused "
"when the key, team, or request has pre-call guardrails configured. "
"The file was not forwarded to the provider."
),
type="invalid_request_error",
param="passthrough",
code=400,
)
outcome: Final = await scan_batch_input_file(
file_source=file_source,
request_metadata=request_metadata,
@ -458,6 +552,7 @@ async def create_file(
custom_llm_provider: str = Form(default="openai"),
file: UploadFile = File(...),
litellm_metadata: str | None = Form(default=None),
passthrough: bool = Form(default=False),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
@ -560,17 +655,28 @@ async def create_file(
if blocked_extension_failure is not None:
raise_upload_validation_failure(blocked_extension_failure)
if passthrough:
_validate_passthrough_upload(
purpose=purpose,
target_model_names=target_model_names_list,
model=model_param,
target_storage=target_storage,
llm_router=llm_router,
team_id=user_api_key_dict.team_id,
)
if purpose == "batch":
batch_file_failure: Final = await asyncio.to_thread(
check_batch_file_upload,
file.filename,
file_source,
_MAX_BATCH_FILE_SIZE_MB_ADAPTER.validate_python(general_settings.get("max_batch_file_size_mb")),
PASSTHROUGH_BATCH_LINE_SHAPE if passthrough else BATCH_LINE_SHAPE,
)
if batch_file_failure is not None:
raise_batch_file_validation_failure(batch_file_failure)
data = {}
data = {"passthrough": True} if passthrough else {}
# Parse expires_after if provided
expires_after: FileExpiresAfter | None = None
@ -673,6 +779,7 @@ async def create_file(
scan_result: Final = await _scan_batch_upload(
file_source=file_source,
purpose=purpose,
passthrough=passthrough,
request_metadata=request_metadata,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,

View file

@ -118,6 +118,7 @@ from litellm.llms.openai_like.model_info import (
MODEL_INFO_REFRESH_SECONDS,
get_openai_compatible_model_info,
)
from litellm.router_strategy.base_routing_strategy import BaseRoutingStrategy
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
from litellm.router_strategy.complexity_router.context_compaction import (
arm_compaction,
@ -1378,6 +1379,9 @@ class Router:
`_init_routing_groups`) so repeated `update_settings` calls don't
accumulate dead selectors that keep receiving callback events.
"""
for selector in selectors:
if isinstance(selector, BaseRoutingStrategy):
selector.retire()
selector_ids: Final = {id(s) for s in selectors if s is not None}
if not selector_ids:
return
@ -5981,6 +5985,7 @@ class Router:
replace_model_in_jsonl_bool: Final = should_replace_model_in_jsonl(
purpose=purpose,
passthrough=kwargs.get("passthrough") is True,
)
if replace_model_in_jsonl_bool:
file = replace_model_in_jsonl(
@ -12126,7 +12131,7 @@ class Router:
)
rebuild_routing_groups = True
elif var == "routing_strategy_args":
routing_args_updated = True
routing_args_updated = value != self.routing_strategy_args
setattr(self, var, value)
else:
verbose_router_logger.debug("Setting %s is not allowed", var)

View file

@ -40,10 +40,24 @@ class BaseRoutingStrategy(ABC):
self.periodic_sync_in_memory_spend_with_redis(default_sync_interval=default_sync_interval)
)
def cancel_sync_task(self) -> None:
if self._sync_task is not None:
self._sync_task.cancel()
def retire(self) -> None:
self.cancel_sync_task()
if not self.redis_increment_operation_queue:
return
try:
loop: Final = asyncio.get_running_loop()
except RuntimeError:
return
loop.create_task(self._push_in_memory_increments_to_redis())
async def cleanup(self):
"""Cleanup method to be called when shutting down"""
if self._sync_task is not None:
self._sync_task.cancel()
self.cancel_sync_task()
try:
await self._sync_task
except asyncio.CancelledError:

View file

@ -62,15 +62,15 @@ def parse_jsonl_with_embedded_newlines(content: str) -> list[dict]:
def should_replace_model_in_jsonl(
purpose: OpenAIFilesPurpose,
passthrough: bool = False,
) -> bool:
"""
Check if the model name should be replaced in the JSONL file for the deployment model name.
Azure raises an error on create batch if the model name for deployment is not in the .jsonl.
A passthrough upload keeps the caller's bytes untouched, so its rows are never rewritten.
"""
if purpose == "batch":
return True
return False
return purpose == "batch" and not passthrough
def replace_model_in_jsonl(file_content: FileTypes, new_model_name: str) -> FileTypes:

View file

@ -2,7 +2,7 @@ from collections.abc import Mapping, Sequence
from typing import Any, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator
from typing_extensions import ReadOnly, TypedDict
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.proxy._types import (
LiteLLM_UserTableWithKeyCount,
@ -28,6 +28,16 @@ class UserSearchWhere(TypedDict):
OR: ReadOnly[tuple[Mapping[Literal["user_id", "user_email"], InsensitiveContains], ...]]
class KeyActivitySearchWhere(TypedDict):
"""Prisma filter behind `/user/daily/activity/aggregated/search`: exact token hash, or key alias
or user id containing the term, case-insensitive."""
user_id: NotRequired[ReadOnly[str]]
OR: ReadOnly[
tuple[Mapping[Literal["token"], str] | Mapping[Literal["key_alias", "user_id"], InsensitiveContains], ...]
]
class UserListResponse(BaseModel):
"""
Response model for the user list endpoint

View file

@ -1,6 +1,8 @@
from collections.abc import Mapping, Sequence
from typing import Any, Final, Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from typing_extensions import NotRequired, ReadOnly, TypedDict
from litellm.proxy._types import (
KeyManagementRoutes,
@ -11,10 +13,32 @@ from litellm.proxy._types import (
MemberDeleteRequest,
)
from litellm.proxy.common_utils.timezone_utils import budget_duration_error
from litellm.types.proxy.management_endpoints.internal_user_endpoints import InsensitiveContains
from litellm.types.proxy.management_endpoints.management_v1 import ResourceResponse
TeamIdSearchMatch = Literal["exact", "prefix"]
TeamIdSearchFilter = TypedDict(
"TeamIdSearchFilter",
{ # mutable-ok: functional TypedDict field map
"in": NotRequired[ReadOnly[Sequence[str]]],
"notIn": NotRequired[ReadOnly[Sequence[str]]],
},
)
class TeamKeyActivitySearchWhere(TypedDict):
"""Prisma filter behind `/team/daily/activity/aggregated/search`: exact token hash, or key alias
or user id containing the term, case-insensitive, narrowed to the teams and keys the caller may see."""
team_id: NotRequired[ReadOnly[TeamIdSearchFilter]]
token: NotRequired[ReadOnly[Mapping[Literal["in"], Sequence[str]]]]
OR: ReadOnly[
tuple[Mapping[Literal["token"], str] | Mapping[Literal["key_alias", "user_id"], InsensitiveContains], ...]
]
MAX_BULK_TEAM_MEMBER_DELETES: Final = 500
MAX_BULK_TEAM_MEMBER_BUDGET_UPDATES: Final = 500
@ -253,3 +277,48 @@ class TeamUserSpendResponse(BaseModel):
start_date: str
end_date: str
results: tuple[TeamUserSpendRow, ...]
TeamDailyActivityExportType = Literal["daily", "daily_with_keys", "daily_with_users", "daily_with_models"]
TeamDailyActivityExportFormat = Literal["csv", "json"]
class TeamDailyActivityExportRow(BaseModel):
date: str
team_id: str
team_alias: str | None = None
api_key: str | None = None
key_alias: str | None = None
user_id: str | None = None
user_email: str | None = None
keys: int | None = None
model: str | None = None
spend: float
flat_cost: float = 0.0
api_requests: int
successful_requests: int
failed_requests: int
total_tokens: int
prompt_tokens: int
completion_tokens: int
cache_read_input_tokens: int
cache_creation_input_tokens: int
class TeamDailyActivityExportMetadata(BaseModel):
export_date: str
export_type: TeamDailyActivityExportType
start_date: str
end_date: str
team_ids: list[str] | None
total_spend: float
total_flat_cost: float = 0.0
total_api_requests: int
total_successful_requests: int
total_failed_requests: int
total_tokens: int
class TeamDailyActivityExportResponse(BaseModel):
metadata: TeamDailyActivityExportMetadata
data: list[TeamDailyActivityExportRow]

View file

@ -68565,6 +68565,7 @@
"source": "https://api.together.ai/v1/models"
},
"vertex_ai/gemini-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai",
@ -68611,7 +68612,9 @@
"input_cost_per_token": 1.5e-06,
"litellm_provider": "vertex_ai",
"mode": "chat",
"output_cost_per_reasoning_token": 9e-06,
"output_cost_per_token": 9e-06,
"output_cost_per_video_token": 1.75e-05,
"source": "https://cloud.google.com/gemini-enterprise-agent-platform/generative-ai/pricing"
},
"vertex_ai/gemini-omni-1.1-flash-preview": {
@ -69773,6 +69776,7 @@
"cache_creation_input_audio_token_cost": 4e-07,
"cache_read_input_audio_token_cost": 4e-07,
"cache_read_input_token_cost": 4e-07,
"deprecation_date": "2026-10-31",
"input_cost_per_audio_token": 3.2e-05,
"input_cost_per_image_token": 5e-06,
"input_cost_per_token": 4e-06,
@ -69804,6 +69808,7 @@
"supports_tool_choice": true
},
"azure/gpt-live-1": {
"deprecation_date": "2027-09-10",
"input_cost_per_second": 0.000833333333333,
"litellm_provider": "azure",
"mode": "realtime",
@ -69821,6 +69826,7 @@
"supports_function_calling": true
},
"azure/gpt-live-transcribe": {
"deprecation_date": "2028-02-01",
"input_cost_per_second": 0.000283333333333,
"litellm_provider": "azure",
"max_input_tokens": 32000,
@ -69842,6 +69848,7 @@
"supports_audio_input": true
},
"azure/gpt-transcribe": {
"deprecation_date": "2028-02-01",
"input_cost_per_second": 7.5e-05,
"litellm_provider": "azure",
"mode": "audio_transcription",
@ -69860,6 +69867,7 @@
"supports_audio_input": true
},
"azure/gpt-realtime-translate": {
"deprecation_date": "2027-05-06",
"input_cost_per_second": 0.000566666666667,
"litellm_provider": "azure",
"max_input_tokens": 32000,

View file

@ -28,10 +28,13 @@ GET /tag/user-agent/per-user-analytics
GET /tag/wau
GET /team/daily/activity
GET /team/daily/activity/aggregated
GET /team/daily/activity/export
GET /team/daily/activity/aggregated/search
GET /team/spend/by_user
GET /team/spend/report
GET /user/daily/activity
GET /user/daily/activity/aggregated
GET /user/daily/activity/aggregated/search
GET /user/spend/report
# Admin UI helper endpoints; serve UI forms and caller-scoped views, not desired state

View file

@ -61,6 +61,7 @@ from e2e_http import (
StreamingResponse,
Success,
UnknownApiError,
proxy_error,
require_successful_call,
unwrap,
)
@ -1752,3 +1753,114 @@ class TestBatchTerminalState:
assert (cost_row.total_tokens or 0) > 0, (
f"batch cost row has no token usage: {cost_row.total_tokens!r}"
)
NATIVE_VERTEX_BATCH_ROWS: Final = b"".join(
json.dumps(
{
"request": {
"contents": [{"role": "user", "parts": [{"text": text}]}],
"tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}],
}
}
).encode()
+ b"\n"
for text in ("What is the tallest building in the world?", "Who won the last FIFA World Cup?")
)
VERTEX_BATCH_PROVIDER: Final = next(p for p in PROVIDERS if p.name == "vertex_ai")
class TestVertexNativePassthrough:
"""`passthrough=true` on POST /v1/files uploads native Vertex batch JSONL byte for
byte (no OpenAI-to-Vertex translation, so `googleSearch` tools and the grounding
metadata they produce survive), and a batch created from that file is accepted.
Terminal-state assertions (native output rows with groundingMetadata, the spend
row) are deliberately not here: retrieving a non-terminal batch books a $0 spend
row that blocks the real-cost row, the same reason TestBatchTerminalState polls
the list endpoint only. Those are proven by the PR's live curl proof instead.
"""
@pytest.mark.covers(
"llm.files.vertex.native_passthrough.nonstream.works",
"llm.batches.vertex.native_passthrough.nonstream.works",
exercised_on=["files", "batches"],
)
def test_native_jsonl_round_trips_untouched_and_starts_a_batch(
self, client: BatchClient, resources: ResourceManager, batch_deployments: None
) -> None:
key = resources.key()
file = unwrap(
client.upload_file(
content=NATIVE_VERTEX_BATCH_ROWS,
form=FileUploadForm(
purpose="batch", target_model_names=VERTEX_BATCH_PROVIDER.model, passthrough=True
),
key=key,
)
)
resources.defer(lambda: cleanup_file(client, file.id, key=key))
assert_file_object(file, provider="vertex_ai")
assert is_managed_id(file.id), f"passthrough upload must return a managed file id, got {file.id!r}"
assert file.bytes == len(NATIVE_VERTEX_BATCH_ROWS), (
f"passthrough upload must report the caller's byte count, got {file.bytes}"
)
downloaded = client.proxy.transport.download(
f"/v1/files/{file.id}/content", headers=client.proxy.transport.bearer(key)
)
assert downloaded.status_code == 200, (
f"file content must be 200, got {downloaded.status_code}: {downloaded.body[:300]}"
)
assert downloaded.body.encode() == NATIVE_VERTEX_BATCH_ROWS, (
"passthrough file content must be the uploaded native rows byte for byte"
)
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(lambda: cleanup_batch(client, batch.id, key=key, delete_output_files=True))
assert is_managed_id(batch.id), f"passthrough batch must be LiteLLM-managed, got {batch.id!r}"
assert batch.status in CREATED_BATCH_STATUSES, f"passthrough batch has non-transitional status {batch.status!r}"
assert batch.input_file_id == file.id
@pytest.mark.covers("llm.files.vertex.native_passthrough_validation.nonstream.works", exercised_on=["files"])
@pytest.mark.parametrize(
"content, form, expected_param",
[
pytest.param(
NATIVE_VERTEX_BATCH_ROWS,
FileUploadForm(purpose="batch", passthrough=True),
"target_model_names",
id="no-target-model",
),
pytest.param(
NATIVE_VERTEX_BATCH_ROWS,
FileUploadForm(purpose="batch", target_model_names=OPENAI_BATCH_MODEL, passthrough=True),
"target_model_names",
id="non-vertex-target-model",
),
pytest.param(
render_jsonl(VERTEX_BATCH_PROVIDER.raw_model),
FileUploadForm(purpose="batch", target_model_names=VERTEX_BATCH_PROVIDER.model, passthrough=True),
"request",
id="openai-shaped-rows",
),
],
)
def test_passthrough_upload_is_rejected_outside_a_native_vertex_batch(
self,
content: bytes,
form: FileUploadForm,
expected_param: str,
client: BatchClient,
resources: ResourceManager,
batch_deployments: None,
) -> None:
key = resources.key()
result = client.upload_file(content=content, form=form, key=key)
assert isinstance(result, UnknownApiError), f"expected a 400, got {result!r}"
assert result.status_code == 400, f"expected 400, got {result.status_code}: {result.body[:300]}"
error = proxy_error(result.body)
assert error.param == expected_param, f"unexpected error param in {error!r}"
assert "passthrough" in error.message

View file

@ -21,6 +21,7 @@
- {id: llm.batches.openai_provider_fallback.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py", rationale: "Provider-fallback raw-id scenario"}
- {id: llm.batches.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Azure batches all scenarios"}
- {id: llm.batches.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Vertex batches"}
- {id: llm.batches.vertex.native_passthrough.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: vertex, capability: native_passthrough, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-4790", rationale: "A batch created from a passthrough-uploaded native Vertex JSONL file is accepted and starts on the deployment named at upload"}
- {id: llm.batches.bedrock.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Bedrock batches (encoded/unified only)"}
- {id: llm.batches.bedrock.assume_role.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: assume_role, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create under STS assume-role credentials"}
- {id: llm.batches.bedrock.govcloud_partition.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: govcloud_partition, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create in the us-gov-west-1 partition"}
@ -46,6 +47,8 @@
- {id: llm.files.openai.passthrough.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_passthrough_e2e.py", rationale: "POST/DELETE /openai_passthrough/v1/files relay OpenAI's own file object; the dedicated prefix must not bind as a provider name on the /{provider}/v1/files route (GitHub issue #36086)"}
- {id: llm.files.azure_openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:45", rationale: "Azure file upload managed backend"}
- {id: llm.files.vertex.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:52", rationale: "Vertex file upload to GCS"}
- {id: llm.files.vertex.native_passthrough.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: vertex, capability: native_passthrough, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-4790", rationale: "POST /v1/files with passthrough=true ships native Vertex batch JSONL (googleSearch tools and all) to GCS untouched and GET /v1/files/{id}/content returns the same bytes"}
- {id: llm.files.vertex.native_passthrough_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: vertex, capability: input_validation, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-4790", rationale: "passthrough=true without a Vertex target_model_names, or with OpenAI-shaped rows, is a 400 naming the offending field and nothing is uploaded"}
- {id: llm.files.bedrock.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:59", rationale: "Bedrock file upload to S3"}
- {id: llm.files.bedrock.govcloud_partition.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: govcloud_partition, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock file upload to an S3 bucket in the us-gov-west-1 partition"}
- {id: llm.files.bedrock.split_s3_credentials.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: split_s3_credentials, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py / LIT-8297", rationale: "Bedrock file upload, content and delete sign S3 with s3_access_key_id / s3_secret_access_key when they differ from the aws_* identity"}

View file

@ -73,6 +73,7 @@ LlmCapability = Literal[
"long_context_1m",
"mid_conversation_system",
"multi_turn",
"native_passthrough",
"pdf_input",
"prompt_cache_1h",
"prompt_cache_5m",

View file

@ -69,6 +69,7 @@ class FileUploadForm(BaseModel):
purpose: str = "batch"
target_model_names: str | None = None
custom_llm_provider: str | None = None
passthrough: bool | None = None
# ---------- Result types ----------
@ -376,12 +377,18 @@ class ProxyErrorDetail(BaseModel):
message: str
type: str
code: str
param: str | None = None
class _ProxyErrorBody(BaseModel):
error: ProxyErrorDetail
def proxy_error(body: str) -> ProxyErrorDetail:
"""The proxy's own error envelope (`{"error": {message, type, param, code}}`) parsed off a rejected call."""
return _ProxyErrorBody.model_validate_json(body).error
def relayed_provider_rate_limit(outcome: RateLimitedError) -> ProxyErrorDetail | None:
"""The provider's own 429 as the proxy relayed it, or None when the 429 is the proxy's own."""
if PROVIDER_RATE_LIMIT_MARKER not in outcome.body:

View file

@ -8,7 +8,8 @@ caller can hand AWS support the request id behind a completion. Regional
inference-profile ids are the deployment shape most Bedrock customers run; a
v1.90.0 regression timed them out, and the Converse route keeps them covered in
test_chat_completions_regression_e2e.py, so the invoke route carries its own
rows here.
rows here. The file also covers Bedrock-native OpenAI model ids taking the
default (Converse) route with max_tokens.
"""
from __future__ import annotations
@ -26,6 +27,7 @@ pytestmark = pytest.mark.e2e
CONVERSE_REGIONAL_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
INVOKE_REGIONAL_BACKEND = "bedrock/invoke/us.anthropic.claude-haiku-4-5-20251001-v1:0"
OPENAI_FAMILY_BACKEND = "bedrock/global.openai.gpt-6-sol"
PROVIDER_HEADER_PREFIX = "llm_provider-"
BEDROCK_REQUEST_ID_HEADER = "llm_provider-x-amzn-requestid"
@ -199,3 +201,18 @@ class TestBedrockInvokeRegionalModelIds:
)
_assert_streamed_completion(result)
class TestBedrockOpenAIFamilyDefaultRoute:
@pytest.mark.covers("llm.chat_completions.bedrock_converse.basic.nonstream.works", exercised_on=[])
def test_openai_family_model_id_completes_with_max_tokens(
self, client: PassthroughClient, resources: ResourceManager
) -> None:
model = _register_bedrock_model(
client, resources, "e2e-bedrock-openai-family", OPENAI_FAMILY_BACKEND
)
key = resources.key()
response = unwrap(client.proxy.chat(key, ChatBody(model=model, messages=_prompt(), max_tokens=64)))
_assert_completion(response)

View file

@ -250,6 +250,7 @@ async def test_azure_image_edit_litellm_sdk():
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data
@ -370,6 +371,7 @@ async def test_openai_image_edit_cost_tracking():
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data
@ -460,6 +462,7 @@ async def test_azure_image_edit_cost_tracking():
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data
@ -737,6 +740,7 @@ async def test_image_edit_array_handling():
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data

View file

@ -24,9 +24,14 @@ async def test_xinference_image_generation():
def model_dump(self):
return mock_openai_response
# Create a mock client with the images.generate method
class MockRawResponse:
headers = {}
def parse(self):
return MockResponse()
mock_client = AsyncMock()
mock_client.images.generate = AsyncMock(return_value=MockResponse())
mock_client.images.with_raw_response.generate = AsyncMock(return_value=MockRawResponse())
# Capture the actual arguments sent to OpenAI client
captured_args = None
@ -36,9 +41,9 @@ async def test_xinference_image_generation():
nonlocal captured_args, captured_kwargs
captured_args = args
captured_kwargs = kwargs
return MockResponse()
return MockRawResponse()
mock_client.images.generate.side_effect = capture_generate_call
mock_client.images.with_raw_response.generate.side_effect = capture_generate_call
# Mock the _get_openai_client method to return our mock client
with patch.object(
@ -65,7 +70,7 @@ async def test_xinference_image_generation():
assert response.data[0].url == "https://example.com/image.png"
# Validate that the OpenAI client was called with correct parameters
mock_client.images.generate.assert_called_once()
mock_client.images.with_raw_response.generate.assert_called_once()
assert captured_kwargs is not None
assert (
captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large"
@ -97,9 +102,14 @@ async def test_xinference_image_generation_with_response_format():
def model_dump(self):
return mock_openai_response
# Create a mock client with the images.generate method
class MockRawResponse:
headers = {}
def parse(self):
return MockResponse()
mock_client = AsyncMock()
mock_client.images.generate = AsyncMock(return_value=MockResponse())
mock_client.images.with_raw_response.generate = AsyncMock(return_value=MockRawResponse())
# Capture the actual arguments sent to OpenAI client
captured_args = None
@ -109,9 +119,9 @@ async def test_xinference_image_generation_with_response_format():
nonlocal captured_args, captured_kwargs
captured_args = args
captured_kwargs = kwargs
return MockResponse()
return MockRawResponse()
mock_client.images.generate.side_effect = capture_generate_call
mock_client.images.with_raw_response.generate.side_effect = capture_generate_call
# Mock the _get_openai_client method to return our mock client
with patch.object(
@ -141,7 +151,7 @@ async def test_xinference_image_generation_with_response_format():
assert response.data[0].b64_json is not None
# Validate that the OpenAI client was called with correct parameters
mock_client.images.generate.assert_called_once()
mock_client.images.with_raw_response.generate.assert_called_once()
assert captured_kwargs is not None
assert (
captured_kwargs["model"] == "stabilityai/stable-diffusion-3.5-large"

View file

@ -116,7 +116,7 @@ def test_vertex_batch_create_survives_explicit_null_output_info(gateway: Gateway
"batch",
"validating",
_encoded(INPUT_FILE_ID, model, "file-"),
_encoded(f"{OUTPUT_PREFIX}/predictions.jsonl", model, "file-"),
None,
None,
"24h",
), response.text

View file

@ -0,0 +1,522 @@
import csv
import io
import os
import signal
import uuid
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime, timedelta, timezone
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import openai
import pytest
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import group_members, owned_proxy, owned_proxy_process
def _export_range() -> dict[str, str]:
today: Final = datetime.now(timezone.utc)
return {
"start_date": (today - timedelta(days=1)).strftime("%Y-%m-%d"),
"end_date": (today + timedelta(days=1)).strftime("%Y-%m-%d"),
"timezone": "0",
}
def _team_with_three_keys(
gateway: Gateway, scenario: Scenario, model: str
) -> tuple[str, tuple[str, ...], tuple[str, ...], dict[str, float]]:
team: Final = scenario.team(models=[model])
keys: Final = tuple(scenario.key(team_id=team, models=[model]) for _ in range(3))
digests: Final = tuple(sha256(key.encode()).hexdigest() for key in keys)
for key in keys:
reply: Final = gateway.chat(model, key=key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: len({row["api_key"] for row in values}) == 3,
seconds=70,
)
spend_by_key: Final = {row["api_key"]: float(row["spend"]) for row in daily}
return team, keys, digests, spend_by_key
def _export_json(gateway: Gateway, **params: str) -> httpx.Response:
return gateway.request("GET", "/team/daily/activity/export", params={**_export_range(), **params})
def test_team_activity_export_returns_every_key_beyond_the_top_n_cap(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
keys: Final = tuple(scenario.key(team_id=team, models=[model]) for _ in range(3))
digests: Final = tuple(sha256(key.encode()).hexdigest() for key in keys)
for key in keys:
reply: Final = gateway.chat(model, key=key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: len({row["api_key"] for row in values}) == 3,
seconds=70,
)
spend_by_key: Final = {row["api_key"]: float(row["spend"]) for row in daily}
response: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={
**_export_range(),
"team_id": team,
"export_type": "daily_with_keys",
"format": "json",
},
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
rows: Final = tuple(object_value(row) for row in body["data"])
assert sorted(string_value(row["api_key"]) for row in rows) == sorted(digests), response.text
for row in rows:
assert row["team_id"] == team, response.text
assert float(row["spend"]) == pytest.approx(spend_by_key[string_value(row["api_key"])]), response.text
metadata: Final = object_value(body["metadata"])
assert (
metadata["export_type"],
metadata["team_ids"],
metadata["total_api_requests"],
metadata["total_successful_requests"],
metadata["total_failed_requests"],
) == ("daily_with_keys", [team], 3, 3, 0), response.text
assert float(metadata["total_spend"]) == pytest.approx(sum(spend_by_key.values())), response.text
def test_team_activity_export_csv_downloads_every_key(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
keys: Final = tuple(scenario.key(team_id=team, models=[model]) for _ in range(3))
digests: Final = tuple(sha256(key.encode()).hexdigest() for key in keys)
for key in keys:
reply: Final = gateway.chat(model, key=key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: len({row["api_key"] for row in values}) == 3,
seconds=70,
)
spend_by_key: Final = {row["api_key"]: float(row["spend"]) for row in daily}
response: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={
**_export_range(),
"team_id": team,
"export_type": "daily_with_keys",
"format": "csv",
},
)
assert response.status_code == 200, response.text
assert response.headers["content-type"].startswith("text/csv"), response.headers
assert "attachment" in response.headers["content-disposition"], response.headers
records: Final = tuple(csv.DictReader(io.StringIO(response.text)))
assert len(records) == 3, response.text
assert sorted(record["Key ID"] for record in records) == sorted(digests), response.text
assert sorted(record["Team ID"] for record in records) == [team, team, team], response.text
for record in records:
assert record["Spend ($)"] == f"{spend_by_key[record['Key ID']]:.4f}", response.text
def test_team_activity_export_denies_a_member_another_team(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team_a: Final = scenario.team(models=[model])
team_b: Final = scenario.team(models=[model])
member: Final = scenario.user(user_role="internal_user", teams=[team_a])
member_key: Final = scenario.key(user_id=member, team_id=team_a, models=[model])
reply: Final = gateway.chat(model, key=member_key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team_a,)),
lambda values: len(values) == 1,
seconds=70,
)
denied: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={**_export_range(), "team_id": team_b, "export_type": "daily", "format": "json"},
key=member_key,
)
assert denied.status_code == 404, denied.text
assert f"User does not belong to Team= {team_b}" in denied.text, denied.text
allowed: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={**_export_range(), "team_id": team_a, "export_type": "daily", "format": "json"},
key=member_key,
)
assert allowed.status_code == 200, allowed.text
rows: Final = tuple(object_value(row) for row in object_value(allowed.json())["data"])
assert len(rows) == 1, allowed.text
assert rows[0]["team_id"] == team_a, allowed.text
assert float(rows[0]["spend"]) == pytest.approx(float(daily[0]["spend"])), allowed.text
def test_export_daily_total_matches_the_capped_aggregated_team_spend(gateway: Gateway, tmp_path: Path) -> None:
with owned_proxy(gateway, tmp_path, {"USAGE_TOP_API_KEYS_LIMIT": "2"}, workers=2) as candidate:
with candidate.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team, keys, digests, spend_by_key = _team_with_three_keys(candidate, scenario, model)
aggregated: Final = candidate.request(
"GET",
"/team/daily/activity/aggregated",
params={**_export_range(), "team_ids": team},
)
assert aggregated.status_code == 200, aggregated.text
body: Final = object_value(aggregated.json())
metadata: Final = object_value(body["metadata"])
assert metadata["api_key_limit"] == 2, aggregated.text
assert metadata["total_api_keys"] == 3, aggregated.text
day: Final = object_value(body["results"][0])
breakdown: Final = object_value(day["breakdown"])
assert len(object_value(breakdown["api_keys"])) == 2, aggregated.text
team_spend: Final = float(
object_value(object_value(object_value(breakdown["entities"])[team])["metrics"])["spend"]
)
response: Final = _export_json(candidate, team_id=team, export_type="daily", format="json")
assert response.status_code == 200, response.text
rows: Final = tuple(object_value(row) for row in object_value(response.json())["data"])
assert len(rows) == 1, response.text
assert rows[0]["team_id"] == team, response.text
assert float(rows[0]["spend"]) == pytest.approx(team_spend), response.text
assert float(rows[0]["spend"]) == pytest.approx(sum(spend_by_key.values())), response.text
def test_export_users_folds_spend_per_user_and_leaves_keyless_keys_unassigned(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
user_a: Final = scenario.user(user_role="internal_user", teams=[team])
user_b: Final = scenario.user(user_role="internal_user", teams=[team])
key_a: Final = scenario.key(team_id=team, user_id=user_a, models=[model])
key_b: Final = scenario.key(team_id=team, user_id=user_b, models=[model])
key_none: Final = scenario.key(team_id=team, models=[model])
for key in (key_a, key_b, key_none):
reply: Final = gateway.chat(model, key=key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: len({row["api_key"] for row in values}) == 3,
seconds=70,
)
spend_by_key: Final = {row["api_key"]: float(row["spend"]) for row in daily}
response: Final = _export_json(gateway, team_id=team, export_type="daily_with_users", format="json")
assert response.status_code == 200, response.text
rows: Final = tuple(object_value(row) for row in object_value(response.json())["data"])
by_user: Final = {row["user_id"]: row for row in rows}
assert by_user[user_a]["spend"] == pytest.approx(spend_by_key[sha256(key_a.encode()).hexdigest()]), (
response.text
)
assert by_user[user_b]["spend"] == pytest.approx(spend_by_key[sha256(key_b.encode()).hexdigest()]), (
response.text
)
assert None in by_user, response.text
assert by_user[None]["spend"] == pytest.approx(spend_by_key[sha256(key_none.encode()).hexdigest()]), (
response.text
)
metadata: Final = object_value(object_value(response.json())["metadata"])
assert float(metadata["total_spend"]) == pytest.approx(sum(spend_by_key.values())), response.text
def test_export_models_reports_one_row_per_model_with_matching_spend(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
upstream_a: Final = f"openai/export-{uuid.uuid4().hex}"
upstream_b: Final = f"openai/export-{uuid.uuid4().hex}"
model_a: Final = scenario.model(model=upstream_a, input_cost_per_token=0.001, output_cost_per_token=0.002)
model_b: Final = scenario.model(model=upstream_b, input_cost_per_token=0.0005, output_cost_per_token=0.001)
upstream_models: Final = (upstream_a, upstream_b)
team: Final = scenario.team(models=[model_a, model_b])
key: Final = scenario.key(team_id=team, models=[model_a, model_b])
for model in (model_a, model_b):
reply: Final = gateway.chat(model, key=key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT model, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: len({row["model"] for row in values}) == 2,
seconds=70,
)
spend_by_model: Final = {row["model"]: float(row["spend"]) for row in daily}
response: Final = _export_json(gateway, team_id=team, export_type="daily_with_models", format="json")
assert response.status_code == 200, response.text
rows: Final = tuple(object_value(row) for row in object_value(response.json())["data"])
assert {row["model"] for row in rows} == set(upstream_models), response.text
for row in rows:
assert float(row["spend"]) == pytest.approx(spend_by_model[row["model"]]), response.text
csv_response: Final = _export_json(gateway, team_id=team, export_type="daily_with_models", format="csv")
assert csv_response.status_code == 200, csv_response.text
records: Final = tuple(csv.DictReader(io.StringIO(csv_response.text)))
assert sorted(record["Model"] for record in records) == sorted(upstream_models), csv_response.text
def test_export_without_team_id_returns_only_the_callers_teams(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team_a: Final = scenario.team(models=[model])
team_b: Final = scenario.team(models=[model])
member: Final = scenario.user(user_role="internal_user", teams=[team_a])
member_key: Final = scenario.key(user_id=member, team_id=team_a, models=[model])
other_key: Final = scenario.key(team_id=team_b, models=[model])
reply: Final = gateway.chat(model, key=member_key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
reply_b: Final = gateway.chat(model, key=other_key, text=f"team export {uuid.uuid4().hex}")
assert reply_b["usage"]["total_tokens"] == 40, reply_b
eventually(
lambda: read_rows('SELECT team_id FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team_b,)),
lambda values: len(values) == 1,
seconds=70,
)
eventually(
lambda: read_rows('SELECT team_id FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team_a,)),
lambda values: len(values) == 1,
seconds=70,
)
response: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={**_export_range(), "export_type": "daily", "format": "json"},
key=member_key,
)
assert response.status_code == 200, response.text
rows: Final = tuple(object_value(row) for row in object_value(response.json())["data"])
assert len(rows) == 1, response.text
assert rows[0]["team_id"] == team_a, response.text
def test_export_rejects_requests_without_a_valid_key(gateway: Gateway) -> None:
params: Final = {**_export_range(), "export_type": "daily", "format": "json"}
anonymous: Final = gateway.client.get("/team/daily/activity/export", params=params)
assert anonymous.status_code == 401, anonymous.text
garbage: Final = gateway.request("GET", "/team/daily/activity/export", params=params, key="sk-nope")
assert garbage.status_code == 401, garbage.text
def test_export_rejects_bad_parameters(gateway: Gateway) -> None:
weekly: Final = _export_json(gateway, export_type="weekly", format="json")
assert weekly.status_code == 422, weekly.text
xml: Final = _export_json(gateway, export_type="daily", format="xml")
assert xml.status_code == 422, xml.text
no_dates: Final = gateway.request(
"GET", "/team/daily/activity/export", params={"export_type": "daily", "format": "json"}
)
assert no_dates.status_code == 400, no_dates.text
assert "start_date and end_date" in no_dates.text, no_dates.text
reversed_range: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={"start_date": "2026-09-25", "end_date": "2026-09-23", "export_type": "daily", "format": "json"},
)
assert reversed_range.status_code == 400, reversed_range.text
assert "end_date must be on or after start_date" in reversed_range.text, reversed_range.text
bad_date: Final = gateway.request(
"GET",
"/team/daily/activity/export",
params={"start_date": "2026-13-40", "end_date": "2026-12-31", "export_type": "daily", "format": "json"},
)
assert bad_date.status_code == 400, bad_date.text
assert "valid YYYY-MM-DD" in bad_date.text, bad_date.text
def test_export_of_a_team_without_spend_returns_empty(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
fresh: Final = _export_json(gateway, team_id=team, export_type="daily", format="json")
assert fresh.status_code == 200, fresh.text
body: Final = object_value(fresh.json())
assert body["data"] == [], fresh.text
assert float(object_value(body["metadata"])["total_spend"]) == 0, fresh.text
unknown: Final = _export_json(gateway, team_id=str(uuid.uuid4()), export_type="daily", format="json")
assert unknown.status_code == 200, unknown.text
unknown_body: Final = object_value(unknown.json())
assert unknown_body["data"] == [], unknown.text
assert float(object_value(unknown_body["metadata"])["total_spend"]) == 0, unknown.text
def test_export_csv_is_deterministic_and_omits_flat_cost_without_ptu(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
_team_with_three_keys(gateway, scenario, model)
params: Final = {**_export_range(), "export_type": "daily_with_keys", "format": "csv"}
first: Final = gateway.request("GET", "/team/daily/activity/export", params=params)
second: Final = gateway.request("GET", "/team/daily/activity/export", params=params)
assert first.status_code == 200 and second.status_code == 200, first.text
assert first.text == second.text, "daily_with_keys csv is not byte-identical across calls"
header: Final = first.text.splitlines()[0]
assert "Flat Cost" not in header and "Total Cost" not in header, header
def test_export_csv_escapes_formula_like_key_aliases(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
alias: Final = f'=HYPERLINK("http://x.{uuid.uuid4().hex}","x")'
keys: Final = (
scenario.key(team_id=team, models=[model], key_alias=alias),
scenario.key(team_id=team, models=[model]),
)
for key in keys:
reply: Final = gateway.chat(model, key=key, text=f"team export {uuid.uuid4().hex}")
assert reply["usage"]["total_tokens"] == 40, reply
digests: Final = tuple(sha256(key.encode()).hexdigest() for key in keys)
eventually(
lambda: read_rows('SELECT api_key FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: len({row["api_key"] for row in values}) == 2,
seconds=70,
)
response: Final = _export_json(gateway, team_id=team, export_type="daily_with_keys", format="csv")
assert response.status_code == 200, response.text
records: Final = {record["Key ID"]: record for record in csv.DictReader(io.StringIO(response.text))}
assert records[digests[0]]["Key Alias"] == "'" + alias, response.text
assert records[digests[1]]["Key Alias"] == "-", response.text
def test_aggregated_route_keeps_the_top_n_key_cap(gateway: Gateway, tmp_path: Path) -> None:
with owned_proxy(gateway, tmp_path, {"USAGE_TOP_API_KEYS_LIMIT": "2"}, workers=2) as candidate:
with candidate.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team, keys, digests, spend_by_key = _team_with_three_keys(candidate, scenario, model)
response: Final = candidate.request(
"GET",
"/team/daily/activity/aggregated",
params={**_export_range(), "team_ids": team},
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
metadata: Final = object_value(body["metadata"])
assert metadata["api_key_limit"] == 2, response.text
assert metadata["total_api_keys"] == 3, response.text
breakdown: Final = object_value(object_value(body["results"][0])["breakdown"])
assert len(object_value(breakdown["api_keys"])) == 2, response.text
def test_paginated_team_daily_activity_still_lists_the_team(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team, keys, digests, spend_by_key = _team_with_three_keys(gateway, scenario, model)
response: Final = gateway.request(
"GET",
"/team/daily/activity",
params={
"team_ids": team,
"start_date": _export_range()["start_date"],
"end_date": _export_range()["end_date"],
},
)
assert response.status_code == 200, response.text
results: Final = object_value(response.json())["results"]
assert isinstance(results, list), response.text
days: Final = tuple(
object_value(day)
for day in results
if team in object_value(object_value(object_value(day)["breakdown"])["entities"])
)
assert len(days) == 1, response.text
entity: Final = object_value(object_value(object_value(days[0]["breakdown"])["entities"])[team])
assert float(object_value(entity["metrics"])["spend"]) == pytest.approx(sum(spend_by_key.values())), (
response.text
)
def test_openai_sdk_chat_still_lands_one_spend_log(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
key: Final = scenario.key(team_id=team, models=[model])
client: Final = openai.OpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=key, max_retries=0)
reply: Final = client.chat.completions.create(
model=model, messages=[{"role": "user", "content": f"sdk {uuid.uuid4().hex}"}], stream=False
)
rows: Final = eventually(
lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (reply.id,)),
lambda values: len(values) == 1,
seconds=70,
)
assert len(rows) == 1 and rows[0]["request_id"] == reply.id, rows
def test_export_and_chat_burst_survives_worker_kill(gateway: Gateway, tmp_path: Path) -> None:
with owned_proxy_process(gateway, tmp_path, {"USAGE_TOP_API_KEYS_LIMIT": "2"}, workers=2) as owned:
candidate: Final = owned.gateway
with candidate.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team, keys, digests, spend_by_key = _team_with_three_keys(candidate, scenario, model)
workers: Final = eventually(
lambda: tuple(member for member in group_members(owned.process.pid) if member.pid != owned.process.pid),
lambda members: len(members) >= 2,
seconds=30,
)
assert len(workers) >= 2, workers
params: Final = {
**_export_range(),
"team_id": team,
"export_type": "daily_with_keys",
"format": "json",
}
def burst(tag: str) -> tuple[tuple[httpx.Response, ...], tuple[httpx.Response, ...]]:
with ThreadPoolExecutor(max_workers=30) as pool:
futures: Final = tuple(
(
pool.submit(
candidate.request,
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"{tag}-{index}-{uuid.uuid4().hex}"}],
},
key=keys[index % 3],
)
if index % 2 == 0
else pool.submit(candidate.request, "GET", "/team/daily/activity/export", params=params)
)
for index in range(30)
)
results: Final = tuple(future.result() for future in futures)
return results[0::2], results[1::2]
chat_a, export_a = burst("bursta")
assert all(response.status_code == 200 for response in chat_a), [r.text for r in chat_a]
assert all(response.status_code == 200 for response in export_a), [r.text for r in export_a]
victim: Final = workers[0]
os.kill(victim.pid, signal.SIGKILL)
chat_b, export_b = burst("burstb")
all_chats: Final = chat_a + chat_b
all_exports: Final = export_a + export_b
assert all(response.status_code == 200 for response in all_chats), [
(r.status_code, r.text) for r in all_chats
]
for response in all_exports:
assert response.status_code == 200, response.text
returned: Final = {string_value(row["api_key"]) for row in object_value(response.json())["data"]}
assert returned == set(digests), response.text
chat_ids: Final = tuple(string_value(object_value(r.json())["id"]) for r in all_chats)
assert len(set(chat_ids)) == 30
id_slots: Final = ", ".join("%s" for _ in chat_ids)
rows: Final = eventually(
lambda: read_rows(
f'SELECT request_id, COUNT(*)::int AS n FROM "LiteLLM_SpendLogs" WHERE request_id IN ({id_slots}) GROUP BY request_id',
chat_ids,
),
lambda values: len(values) == 30,
seconds=70,
)
assert all(row["n"] == 1 for row in rows), rows

View file

@ -0,0 +1,128 @@
import uuid
from datetime import datetime, timedelta, timezone
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from pydantic import JsonValue
_SEARCH_PATH: Final = "/team/daily/activity/aggregated/search"
def _range_around_today() -> dict[str, str]:
today: Final = datetime.now(timezone.utc)
return {
"start_date": (today - timedelta(days=1)).strftime("%Y-%m-%d"),
"end_date": (today + timedelta(days=1)).strftime("%Y-%m-%d"),
"timezone": "0",
}
def _team_key_breakdown(body: dict[str, JsonValue], team: str) -> dict[str, JsonValue]:
results: Final = body["results"]
assert isinstance(results, list) and len(results) == 1, body
entities: Final = object_value(object_value(object_value(results[0])["breakdown"])["entities"])
return object_value(object_value(entities[team])["api_key_breakdown"])
def test_team_key_search_returns_only_the_matching_key_spend_by_alias_and_by_hash(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
needle_alias: Final = f"needle-{uuid.uuid4().hex}"
needle: Final = scenario.key(team_id=team, models=[model], key_alias=needle_alias)
other: Final = scenario.key(team_id=team, models=[model], key_alias=f"other-{uuid.uuid4().hex}")
needle_digest: Final = sha256(needle.encode()).hexdigest()
other_digest: Final = sha256(other.encode()).hexdigest()
for key in (needle, other):
reply: Final = gateway.chat(model, key=key, text=f"key search {uuid.uuid4().hex}")
assert object_value(reply["usage"])["total_tokens"] == 40, reply
daily: Final = eventually(
lambda: read_rows('SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: sorted(row["api_key"] for row in values) == sorted((needle_digest, other_digest)),
seconds=70,
)
assert all(float(row["spend"]) == pytest.approx(0.06) for row in daily), daily
for search in (needle_alias.upper(), needle_digest):
response: Final = gateway.request(
"GET", _SEARCH_PATH, params={"team_ids": team, "search": search, **_range_around_today()}
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
assert object_value(body["metadata"])["total_spend"] == pytest.approx(0.06), response.text
per_key: Final = _team_key_breakdown(body, team)
assert set(per_key) == {needle_digest}, response.text
assert object_value(object_value(per_key[needle_digest])["metrics"])["spend"] == pytest.approx(0.06)
def test_team_key_search_is_scoped_to_the_teams_the_caller_belongs_to(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model])
needle_alias: Final = f"needle-{uuid.uuid4().hex}"
needle: Final = scenario.key(team_id=team, models=[model], key_alias=needle_alias)
needle_digest: Final = sha256(needle.encode()).hexdigest()
reply: Final = gateway.chat(model, key=needle, text=f"key search {uuid.uuid4().hex}")
assert object_value(reply["usage"])["total_tokens"] == 40, reply
eventually(
lambda: read_rows('SELECT api_key FROM "LiteLLM_DailyTeamSpend" WHERE team_id=%s', (team,)),
lambda values: [row["api_key"] for row in values] == [needle_digest],
seconds=70,
)
outsider: Final = scenario.user(user_role="internal_user")
outsider_team: Final = scenario.team(models=[model], members_with_roles=[{"user_id": outsider, "role": "user"}])
outsider_key: Final = scenario.key(user_id=outsider, team_id=outsider_team, models=[model])
params: Final = {"search": needle_alias, **_range_around_today()}
admin_view: Final = gateway.request("GET", _SEARCH_PATH, params={"team_ids": team, **params})
assert admin_view.status_code == 200, admin_view.text
assert set(_team_key_breakdown(object_value(admin_view.json()), team)) == {needle_digest}, admin_view.text
own_teams_view: Final = gateway.request("GET", _SEARCH_PATH, params=params, key=outsider_key)
assert own_teams_view.status_code == 200, own_teams_view.text
own_teams_body: Final = object_value(own_teams_view.json())
assert own_teams_body["results"] == [], own_teams_view.text
assert object_value(own_teams_body["metadata"])["total_api_keys"] == 0, own_teams_view.text
foreign_team_view: Final = gateway.request(
"GET", _SEARCH_PATH, params={"team_ids": team, **params}, key=outsider_key
)
assert foreign_team_view.status_code == 404, foreign_team_view.text
def test_team_key_search_excludes_teams_inside_the_where(gateway: Gateway) -> None:
"""The dashboard always sends exclude_team_ids; a matching key in an excluded
team with higher spend must not consume a take slot nor appear in the result."""
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team_keep: Final = scenario.team(models=[model])
team_drop: Final = scenario.team(models=[model])
shared_alias: Final = f"needle-{uuid.uuid4().hex}"
keep: Final = scenario.key(team_id=team_keep, models=[model], key_alias=f"{shared_alias}-keep")
drop: Final = scenario.key(team_id=team_drop, models=[model], key_alias=f"{shared_alias}-drop")
keep_digest: Final = sha256(keep.encode()).hexdigest()
drop_digest: Final = sha256(drop.encode()).hexdigest()
for _ in range(2):
reply: Final = gateway.chat(model, key=drop, text=f"key search {uuid.uuid4().hex}")
assert object_value(reply["usage"])["total_tokens"] == 40, reply
reply = gateway.chat(model, key=keep, text=f"key search {uuid.uuid4().hex}")
assert object_value(reply["usage"])["total_tokens"] == 40, reply
eventually(
lambda: read_rows(
'SELECT api_key, spend FROM "LiteLLM_DailyTeamSpend" WHERE team_id IN (%s, %s)',
(team_keep, team_drop),
),
lambda values: sorted(row["api_key"] for row in values) == sorted((keep_digest, drop_digest)),
seconds=70,
)
response: Final = gateway.request(
"GET",
_SEARCH_PATH,
params={"search": shared_alias, "exclude_team_ids": team_drop, **_range_around_today()},
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
results: Final = body["results"]
assert isinstance(results, list) and len(results) == 1, body
entities: Final = object_value(object_value(object_value(results[0])["breakdown"])["entities"])
assert set(entities) == {team_keep}, response.text
assert set(_team_key_breakdown(body, team_keep)) == {keep_digest}, response.text

View file

@ -0,0 +1,106 @@
"""Team member spend keeps landing after a flush in which every cost was a whole number.
The proxy runs on a one-connection pool so every spend flush reuses the same database
connection. A batch of $0 requests (a free model here) is the whole-number batch, and the
fractional batches that follow it must still land on that connection.
The $0 batch has to be flushed on its own before the paid request is sent. The spend log
row cannot prove that, since a separate monitor writes spend logs whenever they queue up,
but the daily user spend row is written by the flush cycle right after the member spend
statement, so its arrival means the whole-number batch has already been sent.
"""
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, eventually, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy
from pydantic import JsonValue
SINGLE_CONNECTION_CONFIG: Final = """
model_list: []
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
database_url: os.environ/DATABASE_URL
store_model_in_db: true
proxy_batch_write_at: 1
proxy_batch_polling_interval: 1
database_connection_pool_limit: 1
router_settings:
disable_cooldowns: true
"""
def _member_row(team_id: str, user_id: str) -> list[dict[str, JsonValue]]:
return read_rows(
'SELECT spend, total_spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s',
(team_id, user_id),
)
def _member_spend_is(rows: list[dict[str, JsonValue]], amount: float) -> bool:
return len(rows) == 1 and all(
float(str(rows[0][column])) == pytest.approx(amount) for column in ("spend", "total_spend")
)
def _daily_user_spend_rows(user_id: str) -> list[dict[str, JsonValue]]:
return read_rows('SELECT spend FROM "LiteLLM_DailyUserSpend" WHERE user_id=%s', (user_id,))
def _logged_spend(request_id: str) -> float:
rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (request_id,)),
lambda found: len(found) == 1,
seconds=30,
)
return float(str(rows[0]["spend"]))
def _chat(gateway: Gateway, key: str, model: str) -> str:
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "member spend control"}]},
key=key,
)
assert response.status_code == 200, response.text
return string_value(response.json()["id"])
def test_fractional_member_spend_lands_after_a_whole_number_flush_on_the_same_connection(
gateway: Gateway, tmp_path: Path
) -> None:
config: Final = tmp_path / "single_connection_proxy.yaml"
config.write_text(SINGLE_CONNECTION_CONFIG)
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate, candidate.scenario() as scenario:
free: Final = scenario.model(input_cost_per_token=0, output_cost_per_token=0, num_retries=0)
paid: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002, num_retries=0)
team: Final = scenario.team(models=[free, paid])
first: Final = scenario.user()
second: Final = scenario.user()
candidate.post(
"/team/member_add",
{"team_id": team, "member": [{"role": "user", "user_id": first}, {"role": "user", "user_id": second}]},
)
first_key: Final = scenario.key(team_id=team, user_id=first)
second_key: Final = scenario.key(team_id=team, user_id=second)
_chat(candidate, first_key, free)
flushed: Final = eventually(lambda: _daily_user_spend_rows(first), lambda rows: len(rows) == 1, seconds=30)
assert float(str(flushed[0]["spend"])) == 0
assert _member_spend_is(_member_row(team, first), 0)
paid_spend: Final = _logged_spend(_chat(candidate, second_key, paid))
assert paid_spend > 0
eventually(lambda: _member_row(team, second), lambda rows: _member_spend_is(rows, paid_spend), seconds=30)
repeat_spend: Final = _logged_spend(_chat(candidate, first_key, paid))
eventually(lambda: _member_row(team, first), lambda rows: _member_spend_is(rows, repeat_spend), seconds=30)
eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id=%s', (team,)),
lambda rows: float(str(rows[0]["spend"])) == pytest.approx(paid_spend + repeat_spend),
seconds=30,
)

View file

@ -210,11 +210,14 @@ async def test_litellm_gateway_image_generation_direct(is_async):
"created": 1,
"data": [{"url": "https://example.com/image.png"}],
}
mock_raw_response = MagicMock()
mock_raw_response.parse.return_value = mock_openai_response
mock_raw_response.headers = {}
if is_async:
# Mock the AsyncOpenAI client that gets created inside _get_openai_client
mock_async_client = AsyncMock()
mock_async_client.images.generate = AsyncMock(return_value=mock_openai_response)
mock_async_client.images.with_raw_response.generate = AsyncMock(return_value=mock_raw_response)
with patch(
"litellm.llms.openai.openai.AsyncOpenAI", return_value=mock_async_client
@ -234,14 +237,14 @@ async def test_litellm_gateway_image_generation_direct(is_async):
assert constructor_kwargs["base_url"] == "http://my-proxy"
# Verify the AsyncOpenAI client was called correctly
mock_async_client.images.generate.assert_awaited_once()
call_kwargs = mock_async_client.images.generate.call_args.kwargs
mock_async_client.images.with_raw_response.generate.assert_awaited_once()
call_kwargs = mock_async_client.images.with_raw_response.generate.call_args.kwargs
assert call_kwargs["model"] == "dall-e-3"
assert call_kwargs["prompt"] == "A beautiful sunset over mountains"
else:
# Mock the sync OpenAI client that gets created inside _get_openai_client
mock_sync_client = MagicMock()
mock_sync_client.images.generate.return_value = mock_openai_response
mock_sync_client.images.with_raw_response.generate.return_value = mock_raw_response
with patch(
"litellm.llms.openai.openai.OpenAI", return_value=mock_sync_client
@ -260,8 +263,8 @@ async def test_litellm_gateway_image_generation_direct(is_async):
assert constructor_kwargs["base_url"] == "http://my-proxy"
# Verify the OpenAI client was called correctly
mock_sync_client.images.generate.assert_called_once()
call_kwargs = mock_sync_client.images.generate.call_args.kwargs
mock_sync_client.images.with_raw_response.generate.assert_called_once()
call_kwargs = mock_sync_client.images.with_raw_response.generate.call_args.kwargs
assert call_kwargs["model"] == "dall-e-3"
assert call_kwargs["prompt"] == "A beautiful sunset over mountains"
@ -285,6 +288,7 @@ async def test_litellm_gateway_from_sdk_image_edit(is_async):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data

View file

@ -313,7 +313,7 @@ def test_openai_max_retries_0(mock_get_openai_client):
def test_openai_image_generation_forwards_organization(mock_get_openai_client):
"""Ensure organization flows to OpenAI client for image generation."""
class _DummyImages:
class _DummyRawImages:
def generate(self, **kwargs): # type: ignore
class _Resp:
def model_dump(self_inner): # minimal OpenAI ImagesResponse shape
@ -327,7 +327,16 @@ def test_openai_image_generation_forwards_organization(mock_get_openai_client):
},
}
return _Resp()
class _RawResp:
headers = {}
def parse(self_inner):
return _Resp()
return _RawResp()
class _DummyImages:
with_raw_response = _DummyRawImages()
class _DummyClient:
def __init__(self):

View file

@ -5,8 +5,9 @@ from .actors import Actor
pytestmark = pytest.mark.asyncio(loop_scope="session")
# GET /team/daily/activity and its /aggregated variant (same shared scope
# resolver, so the matrix must hold for both). A proxy admin (admin view) sees
# GET /team/daily/activity, its /aggregated variant, and the key-search
# variant (same shared scope resolver, so the matrix must hold for all
# three). A proxy admin (admin view) sees
# activity for any team. A non-admin is scoped to user_info.teams: a bare query
# defaults to its own teams (200), and an explicit team_ids filter naming a
# team it does not belong to is 404 (the VERIA-43 fix). Org admins have no
@ -43,8 +44,13 @@ _DATES = "start_date=2024-01-01&end_date=2024-12-31"
@pytest.mark.parametrize(
"endpoint",
("/team/daily/activity", "/team/daily/activity/aggregated"),
ids=("paginated", "aggregated"),
(
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/aggregated/search",
"/team/daily/activity/export",
),
ids=("paginated", "aggregated", "search", "export"),
)
@pytest.mark.parametrize(
"actor,team,expected_status",
@ -54,16 +60,15 @@ _DATES = "start_date=2024-01-01&end_date=2024-12-31"
async def test_team_daily_activity_matrix(
actor: Actor, team: str, expected_status: int, endpoint: str, proxy_client, world
):
query = _DATES
filter_param = "team_id" if endpoint.endswith("/export") else "team_ids"
query = _DATES + ("&search=x" if endpoint.endswith("/search") else "")
if team == "alpha":
query += f"&team_ids={world.team_alpha_id}"
query += f"&{filter_param}={world.team_alpha_id}"
elif team == "beta":
query += f"&team_ids={world.team_beta_id}"
query += f"&{filter_param}={world.team_beta_id}"
resp = await proxy_client.get(
f"{endpoint}?{query}",
headers={"Authorization": f"Bearer {world.keys[actor].cleartext}"},
)
assert (
resp.status_code == expected_status
), f"{actor.value} -> {team}: {resp.status_code} {resp.text}"
assert resp.status_code == expected_status, f"{actor.value} -> {team}: {resp.status_code} {resp.text}"

View file

@ -141,6 +141,7 @@ def test_should_replace_model_in_jsonl():
from litellm.router_utils.batch_utils import should_replace_model_in_jsonl
assert should_replace_model_in_jsonl(purpose="batch") is True
assert should_replace_model_in_jsonl(purpose="batch", passthrough=True) is False
assert should_replace_model_in_jsonl(purpose="test") is False
assert should_replace_model_in_jsonl(purpose="user_data") is False

View file

@ -0,0 +1,387 @@
import json
import pytest
import litellm
import litellm.batches.batch_utils as bu
from litellm.types.llms.openai import Batch
GROUNDED_USAGE_METADATA = {
"promptTokenCount": 19,
"candidatesTokenCount": 59,
"thoughtsTokenCount": 406,
"toolUsePromptTokenCount": 73,
"totalTokenCount": 557,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 19}],
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 59}],
"toolUsePromptTokensDetails": [{"modality": "TEXT", "tokenCount": 73}],
"trafficType": "ON_DEMAND",
}
PASSTHROUGH_OUTPUT_URI = (
"gs://litellm-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/u/"
"predictions.jsonl"
)
UNGROUNDED_USAGE_METADATA = {
"promptTokenCount": 20,
"candidatesTokenCount": 48,
"thoughtsTokenCount": 195,
"toolUsePromptTokenCount": 73,
"totalTokenCount": 336,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 20}],
"trafficType": "ON_DEMAND",
}
def _batch(output_file_id: str) -> Batch:
return Batch(
id="b",
completion_window="24h",
created_at=1,
endpoint="/v1/chat/completions",
input_file_id="f",
object="batch",
status="completed",
output_file_id=output_file_id,
)
def _vertex_jsonl(rows: list[dict]) -> bytes:
return "\n".join(json.dumps(row) for row in rows).encode()
def _vertex_openai_row(custom_id: str, model: str, prompt_tokens: int, completion_tokens: int) -> dict:
return {
"id": f"batch_req_{custom_id}",
"custom_id": custom_id,
"response": {
"status_code": 200,
"request_id": custom_id,
"body": {
"id": f"chatcmpl-{custom_id}",
"object": "chat.completion",
"model": model,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
},
},
"error": None,
}
def _native_vertex_row(usage_metadata: dict, *, grounded: bool, model_version: str | None = "gemini-2.5-flash"):
candidate = {"content": {"role": "model", "parts": [{"text": "ok"}]}, "finishReason": "STOP"}
grounding = {"groundingMetadata": {"webSearchQueries": ["q"]}} if grounded else {}
response = {"candidates": [{**candidate, **grounding}], "usageMetadata": usage_metadata}
return {
"request": {"contents": [{"role": "user", "parts": [{"text": "q"}]}], "tools": [{"googleSearch": {}}]},
"status": "",
"response": {**response, **({"modelVersion": model_version} if model_version else {})},
"processed_time": "2026-09-23T19:02:00.000+00:00",
}
def _capture_cost_calls(monkeypatch, prompt_cost=0.5, completion_cost=0.25) -> list:
import litellm.cost_calculator as cc
calls: list = []
def _calc(**kw):
calls.append(kw)
return (prompt_cost, completion_cost)
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
return calls
def test_vertex_native_cost_bills_embedding_rows(monkeypatch):
monkeypatch.setitem(litellm.model_cost, "vertex_ai/gemini-embedding-2", {"input_cost_per_token_batches": 1e-7})
rows = [
{
"key": "id_1",
"status": "",
"request": {"content": {"parts": [{"text": "hello world"}]}},
"response": {"embedding": {"values": [0.1, 0.2]}, "usageMetadata": {"promptTokenCount": 2}},
},
{
"key": "id_2",
"status": "",
"request": {"content": {"parts": [{"text": "hello"}]}},
"response": {"embedding": {"values": [0.3]}, "tokenCount": "3"},
},
{"key": "id_3", "status": "INVALID_ARGUMENT", "request": {"content": {"parts": [{"text": ""}]}}},
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-embedding-2")
assert (result.successful_requests, result.failed_requests) == (2, 1)
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (5, 0, 5)
assert result.cost == pytest.approx(5 * 1e-7)
assert result.models == ["gemini-embedding-2"]
@pytest.mark.asyncio
async def test_native_vertex_rows_route_to_vertex_cost_path_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False),
]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert result.cost == pytest.approx(1.5)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.models == ["gemini-2.5-flash"]
assert {(call["model"], call["custom_llm_provider"]) for call in calls} == {("gemini-2.5-flash", "vertex_ai")}
@pytest.mark.asyncio
async def test_openai_shaped_vertex_rows_keep_the_generic_path_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
monkeypatch.setattr(
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
)
_capture_cost_calls(monkeypatch)
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert result.successful_requests == 1
@pytest.mark.asyncio
async def test_native_vertex_rows_on_another_provider_keep_the_generic_path(monkeypatch):
monkeypatch.setattr(
bu, "calculate_vertex_ai_batch_cost_and_usage", lambda *a, **kw: pytest.fail("native path should not run")
)
_capture_cost_calls(monkeypatch)
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
custom_llm_provider="openai",
)
assert result.successful_requests == 0
@pytest.mark.asyncio
async def test_handle_completed_batch_routes_native_rows_without_flag(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", False, raising=False)
raw_rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)]
async def fake_fetch(batch, custom_llm_provider, litellm_params=None):
return _vertex_jsonl(raw_rows)
monkeypatch.setattr(bu, "_fetch_batch_output_file_content", fake_fetch)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch, prompt_cost=0.7, completion_cost=0.3)
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
result = await bu._handle_completed_batch(
_batch(PASSTHROUGH_OUTPUT_URI),
custom_llm_provider="vertex_ai",
model_name="gemini-2.5-flash",
model_info=deployment_model_info,
)
assert result.cost == pytest.approx(1.0)
assert result.usage.total_tokens == 557
assert [call["model_info"] for call in calls] == [deployment_model_info]
def test_native_vertex_usage_is_billed_like_the_online_path(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
grounded = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)
ungrounded = _native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False)
result = bu.calculate_vertex_ai_batch_cost_and_usage([grounded, ungrounded], "gemini-2.5-flash")
grounded_usage, ungrounded_usage = (call["usage"] for call in calls)
assert grounded_usage.prompt_tokens == 19
assert grounded_usage.completion_tokens == 59 + 406
assert grounded_usage.completion_tokens_details.reasoning_tokens == 406
assert ungrounded_usage.prompt_tokens == 20 + 73
assert ungrounded_usage.completion_tokens == 48 + 195
assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (
19 + 93,
465 + 243,
557 + 336,
)
def test_native_vertex_rows_are_priced_by_model_version_without_a_model_name(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-pro"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
assert [call["model"] for call in calls] == ["gemini-2.5-flash", "gemini-2.5-pro"]
assert result.models == ["gemini-2.5-flash", "gemini-2.5-pro"]
assert result.cost == pytest.approx(1.5)
assert result.successful_requests == 3
assert result.usage.total_tokens == 557 + 336 + 336
def test_native_vertex_rows_without_usage_metadata_count_as_failed(monkeypatch):
_capture_cost_calls(monkeypatch)
rows = [
{"request": {"contents": []}, "status": "Error: bad request", "processed_time": "t"},
{"request": {"contents": []}, "response": {"candidates": []}},
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert (result.successful_requests, result.failed_requests) == (1, 2)
assert result.usage.total_tokens == 557
def test_native_vertex_batch_whose_rows_all_failed_still_names_the_deployment_model(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [{"request": {"contents": []}, "status": "Error: quota exceeded", "processed_time": "t"}] * 2
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert result.models == ["gemini-2.5-flash"]
assert (result.successful_requests, result.failed_requests, result.cost) == (0, 2, 0.0)
assert calls == []
def test_native_vertex_rows_are_priced_with_the_deployment_model_info(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
bu.calculate_vertex_ai_batch_cost_and_usage(
[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
"gemini-2.5-flash",
model_info=deployment_model_info,
)
assert [call["model_info"] for call in calls] == [deployment_model_info]
@pytest.mark.asyncio
async def test_native_vertex_rows_keep_the_deployment_model_info_through_the_batch_entrypoint(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
deployment_model_info = {"input_cost_per_token_batches": 1e-6}
await bu.calculate_batch_cost_and_usage(
file_content_dictionary=[_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True)],
custom_llm_provider="vertex_ai",
model_name="gemini-2.5-flash",
model_info=deployment_model_info,
)
assert [call["model_info"] for call in calls] == [deployment_model_info]
def test_native_vertex_rows_are_priced_by_the_deployment_model_over_model_version(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-pro")]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert [call["model"] for call in calls] == ["gemini-2.5-flash"]
assert result.models == ["gemini-2.5-flash"]
def test_native_vertex_rows_that_fail_response_validation_count_as_failed(monkeypatch):
calls = _capture_cost_calls(monkeypatch)
rows = [
{"request": {"contents": []}, "response": {"candidates": "nope", "usageMetadata": GROUNDED_USAGE_METADATA}},
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, "gemini-2.5-flash")
assert (result.successful_requests, result.failed_requests) == (1, 1)
assert result.usage.total_tokens == 557
assert len(calls) == 1
@pytest.mark.parametrize("wildcard_model", ["*", "vertex_ai/*"])
def test_native_vertex_rows_under_a_wildcard_deployment_are_priced_by_model_version(monkeypatch, wildcard_model):
calls = _capture_cost_calls(monkeypatch)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version=None),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows, wildcard_model)
assert [call["model"] for call in calls] == ["gemini-2.5-flash", wildcard_model]
assert result.cost == pytest.approx(1.5)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.usage.total_tokens == 557 + 336
def test_native_vertex_row_without_model_version_under_a_wildcard_deployment_bills_its_explicit_prices():
deployment_model_info = {"input_cost_per_token_batches": 1e-6, "output_cost_per_token_batches": 2e-6}
with_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-2.5-flash")
without_version = _native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version=None)
twin = bu.calculate_vertex_ai_batch_cost_and_usage([with_version], "vertex_ai/*", model_info=deployment_model_info)
both = bu.calculate_vertex_ai_batch_cost_and_usage(
[with_version, without_version], "vertex_ai/*", model_info=deployment_model_info
)
assert twin.cost > 0
assert both.cost == pytest.approx(2 * twin.cost)
assert (both.successful_requests, both.failed_requests) == (2, 0)
def test_native_vertex_row_the_cost_map_cannot_price_is_billed_at_zero_and_the_rest_still_bills(monkeypatch):
import litellm.cost_calculator as cc
def _calc(**kw):
if kw["model"] == "gemini-unpriced":
raise ValueError("no pricing")
return (0.5, 0.25)
monkeypatch.setattr(cc, "batch_cost_calculator", _calc)
rows = [
_native_vertex_row(GROUNDED_USAGE_METADATA, grounded=True, model_version="gemini-unpriced"),
_native_vertex_row(UNGROUNDED_USAGE_METADATA, grounded=False, model_version="gemini-2.5-flash"),
]
result = bu.calculate_vertex_ai_batch_cost_and_usage(rows)
assert result.cost == pytest.approx(0.75)
assert (result.successful_requests, result.failed_requests) == (2, 0)
assert result.usage.total_tokens == 557 + 336
assert result.models == ["gemini-unpriced", "gemini-2.5-flash"]
@pytest.mark.asyncio
async def test_flag_sends_every_vertex_row_down_the_native_path_when_a_model_is_known(monkeypatch):
monkeypatch.setattr(litellm, "disable_vertex_batch_output_transformation", True, raising=False)
monkeypatch.setattr(
bu, "_aggregate_batch_cost_usage_models", lambda **kw: pytest.fail("generic path should not run")
)
calls = _capture_cost_calls(monkeypatch)
rows = [_vertex_openai_row("request-1", "gemini-2.5-flash", 10, 5)]
result = await bu.calculate_batch_cost_and_usage(
file_content_dictionary=rows, custom_llm_provider="vertex_ai", model_name="gemini-2.5-flash"
)
assert calls == []
assert (result.successful_requests, result.failed_requests) == (0, 1)

View file

View file

@ -0,0 +1,71 @@
from typing import Final
from urllib.parse import parse_qs, urlparse
import httpx
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import HTTPHandler
NATIVE_VERTEX_ROWS: Final = (
b'{"request": {"contents": [{"role": "user", "parts": [{"text": "Who won the 2024 Tour de France?"}]}],'
b' "tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}]}}\n'
b'{"request": {"contents": [{"role": "user", "parts": [{"text": "What is the tallest building in Tokyo?"}]}],'
b' "tools": [{"googleSearch": {}}]}}\n'
)
@pytest.mark.parametrize(
"custom_llm_provider, purpose",
[("openai", "batch"), ("vertex_ai", "assistants")],
ids=["non-vertex-provider", "non-batch-purpose"],
)
def test_create_file_passthrough_is_rejected_outside_a_vertex_batch(custom_llm_provider, purpose):
with pytest.raises(litellm.BadRequestError) as exc_info:
litellm.create_file(
file=("batch.jsonl", b'{"request": {"contents": []}}\n', "application/jsonl"),
purpose=purpose,
custom_llm_provider=custom_llm_provider,
passthrough=True,
api_key="sk-test",
api_base="http://127.0.0.1:9",
)
assert "vertex_ai" in str(exc_info.value)
assert "batch" in str(exc_info.value)
def _gcs_upload_transport(uploads: list[httpx.Request]) -> httpx.MockTransport:
def respond(request: httpx.Request) -> httpx.Response:
uploads.append(request)
object_name: Final = parse_qs(urlparse(str(request.url)).query)["name"][0]
return httpx.Response(
200,
json={
"id": f"my-bucket/{object_name}/1758585600000000",
"name": object_name,
"size": str(len(request.read())),
"timeCreated": "2026-09-23T00:00:00.000Z",
},
)
return httpx.MockTransport(respond)
def test_create_file_passthrough_kwarg_ships_native_rows_byte_for_byte_under_the_passthrough_prefix():
uploads: Final[list[httpx.Request]] = []
file_object = litellm.create_file(
file=("batch.jsonl", NATIVE_VERTEX_ROWS, "application/jsonl"),
purpose="batch",
custom_llm_provider="vertex_ai",
passthrough=True,
model="vertex_ai/gemini-2.5-flash",
gcs_bucket_name="my-bucket",
api_key="test-token",
client=HTTPHandler(client=httpx.Client(transport=_gcs_upload_transport(uploads))),
)
(upload,) = uploads
object_name: Final = parse_qs(urlparse(str(upload.url)).query)["name"][0]
assert upload.read() == NATIVE_VERTEX_ROWS
assert object_name.startswith("litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/")
assert file_object.id == f"gs://my-bucket/{object_name}"

View file

@ -526,3 +526,29 @@ async def test_async_send_batch_collapses_only_identical_alerts() -> None:
{"text": f"[Num Alerts: 2]\n\n{THRESHOLD_ALERT}"},
{"text": CROSSED_ALERT},
)
def _periodic_flush_tasks() -> list[asyncio.Task[object]]:
return [
t
for t in asyncio.all_tasks()
if t.get_coro() is not None and t.get_coro().__qualname__ == "SlackAlerting.periodic_flush"
]
@pytest.mark.asyncio
async def test_update_values_repeated_alerting_reload_keeps_single_periodic_flush_task() -> None:
slack_alerting: Final = SlackAlerting(alerting=["slack"])
try:
for _ in range(5):
slack_alerting.update_values(alerting=["slack"])
await asyncio.sleep(0)
flush_tasks: Final = _periodic_flush_tasks()
assert len(flush_tasks) == 1, f"expected 1 periodic_flush task, found {len(flush_tasks)}"
finally:
for t in _periodic_flush_tasks():
t.cancel()
try:
await t
except asyncio.CancelledError:
pass

View file

@ -1207,14 +1207,18 @@ def test_speech_response_without_a_byte_count_produces_no_output() -> None:
def test_speech_binary_response_is_logged_as_its_summary_not_dropped() -> None:
import httpx
from litellm.litellm_core_utils.core_helpers import set_provider_response_headers_in_hidden_params
from litellm.litellm_core_utils.litellm_logging import _extract_response_obj_and_hidden_params
from litellm.types.llms.openai import HttpxBinaryResponseContent
raw: Final = httpx.Response(200, headers={"content-type": "audio/mpeg"}, content=b"\x00" * 1234)
response_obj, hidden_params = _extract_response_obj_and_hidden_params(HttpxBinaryResponseContent(raw), None)
speech: Final = HttpxBinaryResponseContent(raw)
set_provider_response_headers_in_hidden_params(speech, raw.headers)
response_obj, hidden_params = _extract_response_obj_and_hidden_params(speech, None)
assert response_obj == {"object": "binary", "content_type": "audio/mpeg", "num_bytes": 1234}
assert hidden_params is None
assert hidden_params is not None
assert hidden_params["headers"]["content-type"] == "audio/mpeg"
def test_speech_binary_response_still_streaming_reports_the_bytes_downloaded_so_far() -> None:

View file

@ -2,22 +2,27 @@
import logging
import httpx
import pytest
from litellm.litellm_core_utils.core_helpers import (
_FINISH_REASON_MAP,
RESPONSE_COST_HEADER,
bind_budget_reservation_to_callbacks,
budget_reservation_from_metadata,
drop_params_env_flag,
drop_params_flag,
get_or_create_metadata_bucket,
get_provider_response_headers_from_hidden_params,
map_finish_reason,
normalize_drop_params,
reconstruct_model_name,
redact_nested_match_and_regex_keys,
set_provider_response_headers_in_hidden_params,
unbind_budget_reservation_from_callbacks,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import ImageResponse, TranscriptionResponse
class TestBudgetReservationBinding:
@ -489,3 +494,66 @@ class TestIsExpectedClientError:
category=RateLimitErrorCategory.VENDOR_RATE_LIMIT,
)
assert is_expected_client_error(vendor_limit) is False
class TestProviderResponseHeadersInHiddenParams:
def test_records_raw_headers_and_the_processed_additional_headers(self):
response = ImageResponse()
response._hidden_params = {"additional_headers": {RESPONSE_COST_HEADER: 0.04}}
set_provider_response_headers_in_hidden_params(
response, httpx.Headers({"X-Request-Id": "req_img", "x-ratelimit-remaining-requests": "41"})
)
assert response._hidden_params["headers"] == {
"x-request-id": "req_img",
"x-ratelimit-remaining-requests": "41",
}
additional_headers = response._hidden_params["additional_headers"]
assert additional_headers["llm_provider-x-request-id"] == "req_img"
assert additional_headers["x-ratelimit-remaining-requests"] == "41"
assert additional_headers[RESPONSE_COST_HEADER] == 0.04
def test_litellm_owned_additional_headers_win_over_provider_headers(self):
response = TranscriptionResponse(text="hi")
response._hidden_params = {"additional_headers": {"llm_provider-x-request-id": "kept"}}
set_provider_response_headers_in_hidden_params(response, {"x-request-id": "provider"})
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "kept"
assert response._hidden_params["headers"] == {"x-request-id": "provider"}
def test_getter_returns_the_recorded_headers(self):
response = ImageResponse()
set_provider_response_headers_in_hidden_params(response, {"x-request-id": "req_img"})
assert get_provider_response_headers_from_hidden_params(response) == {"x-request-id": "req_img"}
@pytest.mark.parametrize(
"hidden_params",
[
None,
"headers",
{"additional_headers": {}},
{"headers": "x-request-id: req_img"},
{"headers": {"x-request-id": 7}},
],
)
def test_getter_returns_none_without_a_string_header_mapping(self, hidden_params):
response = ImageResponse()
response._hidden_params = hidden_params
assert get_provider_response_headers_from_hidden_params(response) is None
def test_getter_returns_none_for_an_object_without_hidden_params(self):
assert get_provider_response_headers_from_hidden_params(object()) is None
def test_headers_never_leak_into_a_sibling_response(self):
recorded = TranscriptionResponse()
sibling = TranscriptionResponse()
set_provider_response_headers_in_hidden_params(recorded, {"x-request-id": "req_stt"})
assert get_provider_response_headers_from_hidden_params(sibling) is None
assert "additional_headers" not in sibling._hidden_params

View file

@ -24,6 +24,7 @@ from litellm.cost_calculator import ocr_batch_cost
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LitellmLogging
from litellm.litellm_core_utils.litellm_logging import (
_extract_response_obj_and_hidden_params,
_get_status_fields,
set_callbacks,
)
@ -32,6 +33,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse
from litellm.types.utils import (
CallTypes,
ImageResponse,
LiteLLMRealtimeStreamLoggingObject,
ModelResponse,
TextCompletionResponse,
@ -8694,3 +8696,88 @@ async def test_async_failure_handler_delivers_failure_payload_to_custom_logger()
assert "smoke-failure" in payload["error_str"]
assert payload["model"] == "openai/gpt-5.6"
assert events.empty()
def _image_logging_obj() -> LitellmLogging:
logging_obj = LitellmLogging(
model="gpt-image-2",
messages="a cat",
stream=False,
call_type="aimage_generation",
start_time=time.time(),
litellm_call_id="response-headers-test",
function_id="response-headers-test",
)
logging_obj.model_call_details["litellm_params"] = {"metadata": {}}
logging_obj.optional_params = {}
return logging_obj
def _image_result_with_headers(request_id: str) -> ImageResponse:
result = ImageResponse(created=1, data=[])
result._hidden_params = {"headers": {"x-request-id": request_id}}
return result
def test_process_hidden_params_surfaces_response_headers_from_the_result():
logging_obj = _image_logging_obj()
logging_obj._process_hidden_params_and_response_cost(
_image_result_with_headers("req_img"), datetime.datetime.now(), datetime.datetime.now()
)
assert logging_obj.model_call_details["response_headers"] == {"x-request-id": "req_img"}
def test_process_hidden_params_keeps_handler_set_response_headers():
logging_obj = _image_logging_obj()
logging_obj.model_call_details["response_headers"] = {"x-request-id": "from-handler"}
logging_obj._process_hidden_params_and_response_cost(
_image_result_with_headers("from-result"), datetime.datetime.now(), datetime.datetime.now()
)
assert logging_obj.model_call_details["response_headers"] == {"x-request-id": "from-handler"}
def _assembled_stream_result_with_headers() -> ModelResponse:
result = _assembled_stream_result()
result._hidden_params = {"headers": {"x-request-id": "req_stream"}}
return result
@pytest.mark.asyncio
async def test_async_streaming_success_passes_result_headers_to_callback_kwargs():
releasing = CustomLogger()
releasing.async_log_success_event = AsyncMock()
patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing])
with patcher:
await logging_obj.async_success_handler(result=_assembled_stream_result_with_headers())
kwargs = releasing.async_log_success_event.await_args.kwargs["kwargs"]
assert kwargs["response_headers"] == {"x-request-id": "req_stream"}
def test_sync_streaming_success_passes_result_headers_to_callback_kwargs():
releasing = CustomLogger()
releasing.log_success_event = MagicMock()
patcher, logging_obj = _streaming_logging_obj_with_callbacks([releasing])
with patcher:
logging_obj.success_handler(result=_assembled_stream_result_with_headers())
kwargs = releasing.log_success_event.call_args.kwargs["kwargs"]
assert kwargs["response_headers"] == {"x-request-id": "req_stream"}
def test_extract_response_obj_and_hidden_params_reads_binary_content_hidden_params():
from litellm.types.llms.openai import HttpxBinaryResponseContent as LiteLLMBinaryResponseContent
result = LiteLLMBinaryResponseContent(response=httpx.Response(status_code=200, content=b"audio bytes"))
result._hidden_params = {"headers": {"x-request-id": "req_tts"}}
response_obj, hidden_params = _extract_response_obj_and_hidden_params(result, None)
assert hidden_params == {"headers": {"x-request-id": "req_tts"}}
assert response_obj["object"] == "binary"

View file

@ -346,7 +346,7 @@ def test_route_prefix_matched_as_path_segment_not_substring():
BedrockModelInfo.get_bedrock_route("bedrock_mantle/openai.gpt-5.5") != "mantle"
)
assert (
BedrockModelInfo.get_bedrock_route("bedrock_mantle/openai.gpt-5.4") == "invoke"
BedrockModelInfo.get_bedrock_route("bedrock_mantle/openai.gpt-5.4") == "converse"
)
assert (
BedrockModelInfo._explicit_mantle_route("bedrock_mantle/openai.gpt-5.5")
@ -964,3 +964,20 @@ def test_s3_static_key_pair_is_none_without_a_full_pair(partial_s3_pair):
from litellm.llms.bedrock.common_utils import s3_static_key_pair
assert s3_static_key_pair({"aws_access_key_id": "bedrock-key", **partial_s3_pair}) is None
def test_unmapped_openai_family_model_routes_to_converse():
"""A Bedrock-native OpenAI model that is not in the cost map yet must not fall to the invoke route.
The invoke ``openai`` provider is the imported-model path and sends ``max_tokens``, which Bedrock
rejects for these models; Converse maps it to ``inferenceConfig.maxTokens``.
"""
from typing import Final
import litellm
unmapped: Final = "bedrock/global.openai.gpt-99-unmapped"
assert unmapped.removeprefix("bedrock/") not in litellm.bedrock_converse_models
assert BedrockModelInfo.get_bedrock_route(unmapped) == "converse"
imported: Final = "bedrock/openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/abc123"
assert BedrockModelInfo.get_bedrock_route(imported) == "openai"

View file

@ -26,6 +26,8 @@ from litellm.llms.base_llm.search.transformation import BaseSearchConfig, Search
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS
from litellm.llms.brave.search.transformation import BraveSearchConfig
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.base_llm.image_generation.transformation import BaseImageGenerationConfig
from litellm.llms.base_llm.text_to_speech.transformation import BaseTextToSpeechConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
@ -41,7 +43,7 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.llms.openai import HttpxBinaryResponseContent, ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
@ -4166,3 +4168,202 @@ async def test_chat_completion_agentic_followup_does_not_repeat_request_params_f
assert followup_calls[0]["temperature"] == 0.2
assert followup_calls[0]["api_base"] == "https://a"
assert followup_calls[0]["model"] == "openai/gpt-5"
_UPSTREAM_HEADERS: Final = {"x-request-id": "req_upstream", "x-ratelimit-remaining-requests": "41"}
def _assert_upstream_headers_recorded(response) -> None:
assert response._hidden_params["headers"]["x-request-id"] == "req_upstream"
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_upstream"
assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41"
def _json_with_upstream_headers(payload: dict) -> httpx.MockTransport:
return httpx.MockTransport(lambda request: httpx.Response(200, json=payload, headers=_UPSTREAM_HEADERS))
def _binary_with_upstream_headers() -> httpx.MockTransport:
return httpx.MockTransport(
lambda request: httpx.Response(
200, content=b"audio-bytes", headers={**_UPSTREAM_HEADERS, "content-type": "audio/mpeg"}
)
)
def test_audio_transcriptions_records_upstream_response_headers():
client = HTTPHandler(client=httpx.Client(transport=_json_with_upstream_headers({"text": "transcribed"})))
response = BaseLLMHTTPHandler().audio_transcriptions(
client=client,
atranscription=False,
**_json_transcription_call_kwargs(_JSONBodyAudioTranscriptionConfig()),
)
_assert_upstream_headers_recorded(response)
@pytest.mark.asyncio
async def test_async_audio_transcriptions_records_upstream_response_headers():
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"text": "transcribed"}))
response = await BaseLLMHTTPHandler().async_audio_transcriptions(
client=client,
**_json_transcription_call_kwargs(_JSONBodyAudioTranscriptionConfig()),
)
_assert_upstream_headers_recorded(response)
def _image_edit_call_kwargs() -> dict:
return {
"model": "edit-model",
"image": b"raw-image",
"prompt": "add a hat",
"image_edit_provider_config": _ImageEditRecordingConfig(),
"image_edit_optional_request_params": {},
"custom_llm_provider": "openai",
"litellm_params": GenericLiteLLMParams(),
"logging_obj": Mock(),
"timeout": 10.0,
}
def test_image_edit_handler_records_upstream_response_headers():
client = HTTPHandler()
client.client = httpx.Client(transport=_json_with_upstream_headers({"transformed_by": "sync"}))
response = BaseLLMHTTPHandler().image_edit_handler(client=client, **_image_edit_call_kwargs())
_assert_upstream_headers_recorded(response)
@pytest.mark.asyncio
async def test_async_image_edit_handler_records_upstream_response_headers():
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"transformed_by": "async"}))
response = await BaseLLMHTTPHandler().async_image_edit_handler(client=client, **_image_edit_call_kwargs())
_assert_upstream_headers_recorded(response)
class _HeaderImageGenerationConfig(BaseImageGenerationConfig):
def get_supported_openai_params(self, model):
return []
def map_openai_params(self, non_default_params, optional_params, model, drop_params):
return optional_params
def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None):
return "https://images.example/v1/generations"
def transform_image_generation_request(self, model, prompt, optional_params, litellm_params, headers):
return {"prompt": prompt}
def transform_image_generation_response(
self,
model,
raw_response,
model_response,
logging_obj,
request_data,
optional_params,
litellm_params,
encoding,
api_key=None,
json_mode=None,
):
return ImageResponse(data=[ImageObject(b64_json=raw_response.json()["b64_json"])])
def _image_generation_call_kwargs() -> dict:
return {
"model": "image-model",
"prompt": "a cat",
"image_generation_provider_config": _HeaderImageGenerationConfig(),
"image_generation_optional_request_params": {},
"custom_llm_provider": "openai",
"litellm_params": {},
"logging_obj": Mock(),
"timeout": 10.0,
}
def test_image_generation_handler_records_upstream_response_headers():
client = HTTPHandler()
client.client = httpx.Client(transport=_json_with_upstream_headers({"b64_json": "abc"}))
response = BaseLLMHTTPHandler().image_generation_handler(client=client, **_image_generation_call_kwargs())
assert response.data[0].b64_json == "abc"
_assert_upstream_headers_recorded(response)
@pytest.mark.asyncio
async def test_async_image_generation_handler_records_upstream_response_headers():
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=_json_with_upstream_headers({"b64_json": "abc"}))
response = await BaseLLMHTTPHandler().async_image_generation_handler(
client=client, **_image_generation_call_kwargs()
)
assert response.data[0].b64_json == "abc"
_assert_upstream_headers_recorded(response)
class _HeaderTextToSpeechConfig(BaseTextToSpeechConfig):
def get_supported_openai_params(self, model):
return []
def map_openai_params(self, model, optional_params, voice=None, drop_params=False, kwargs=None):
return voice, optional_params
def validate_environment(self, headers, model, api_key=None, api_base=None):
return {}
def get_complete_url(self, model, api_base, litellm_params):
return "https://tts.example/v1/speech"
def transform_text_to_speech_request(self, model, input, voice, optional_params, litellm_params, headers):
return {"dict_body": {"input": input}}
def transform_text_to_speech_response(self, model, raw_response, logging_obj):
return HttpxBinaryResponseContent(response=raw_response)
def _text_to_speech_call_kwargs() -> dict:
return {
"model": "tts-model",
"input": "hello",
"voice": "alloy",
"text_to_speech_provider_config": _HeaderTextToSpeechConfig(),
"text_to_speech_optional_params": {},
"custom_llm_provider": "openai",
"litellm_params": {},
"logging_obj": Mock(),
"timeout": 10.0,
}
def test_text_to_speech_handler_records_upstream_response_headers():
client = HTTPHandler()
client.client = httpx.Client(transport=_binary_with_upstream_headers())
response = BaseLLMHTTPHandler().text_to_speech_handler(client=client, **_text_to_speech_call_kwargs())
assert response.content == b"audio-bytes"
_assert_upstream_headers_recorded(response)
@pytest.mark.asyncio
async def test_async_text_to_speech_handler_records_upstream_response_headers():
client = AsyncHTTPHandler()
client.client = httpx.AsyncClient(transport=_binary_with_upstream_headers())
response = await BaseLLMHTTPHandler().async_text_to_speech_handler(client=client, **_text_to_speech_call_kwargs())
assert response.content == b"audio-bytes"
_assert_upstream_headers_recorded(response)

View file

@ -1,13 +1,15 @@
import asyncio
import json
from typing import Final
from unittest.mock import Mock
import httpx
import pytest
from openai import AsyncOpenAI
from openai import AsyncOpenAI, OpenAI
import litellm
from litellm.llms.openai.openai import OpenAIChatCompletion
from litellm.types.utils import ImageResponse
@pytest.mark.parametrize(
@ -253,3 +255,98 @@ async def test_acompletion_streams_tool_call_arguments_over_injected_transport()
assert tool_call.function.name == "get_weather"
assert json.loads(tool_call.function.arguments) == {"city": "Paris"}
assert rebuilt.choices[0].finish_reason == "tool_calls"
_PROVIDER_HEADERS: Final = {"x-request-id": "req_openai", "x-ratelimit-remaining-requests": "41"}
def _image_generation_transport() -> httpx.MockTransport:
return httpx.MockTransport(
lambda request: httpx.Response(
200, json={"created": 1, "data": [{"b64_json": "abc"}]}, headers=_PROVIDER_HEADERS
)
)
def _speech_transport() -> httpx.MockTransport:
return httpx.MockTransport(
lambda request: httpx.Response(
200, content=b"audio-bytes", headers={**_PROVIDER_HEADERS, "content-type": "audio/mpeg"}
)
)
def _assert_provider_headers_recorded(response) -> None:
assert response._hidden_params["headers"]["x-request-id"] == "req_openai"
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_openai"
assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41"
def _image_generation_kwargs() -> dict:
return {
"model": "gpt-image-2",
"prompt": "a cat",
"timeout": 10,
"optional_params": {},
"logging_obj": Mock(),
"api_key": "transport-only",
"model_response": ImageResponse(),
}
def test_image_generation_records_provider_response_headers():
with httpx.Client(transport=_image_generation_transport()) as http_client:
response = OpenAIChatCompletion().image_generation(
client=OpenAI(api_key="transport-only", http_client=http_client), **_image_generation_kwargs()
)
_assert_provider_headers_recorded(response)
@pytest.mark.asyncio
async def test_aimage_generation_records_provider_response_headers():
async with httpx.AsyncClient(transport=_image_generation_transport()) as http_client:
response = await OpenAIChatCompletion().image_generation(
client=AsyncOpenAI(api_key="transport-only", http_client=http_client),
aimg_generation=True,
**_image_generation_kwargs(),
)
_assert_provider_headers_recorded(response)
def _audio_speech_kwargs() -> dict:
return {
"model": "gpt-4o-mini-tts",
"input": "hello",
"voice": "alloy",
"optional_params": {},
"api_key": "transport-only",
"api_base": None,
"organization": None,
"project": None,
"max_retries": 0,
"timeout": 10,
"logging_obj": Mock(),
}
def test_audio_speech_records_provider_response_headers():
with httpx.Client(transport=_speech_transport()) as http_client:
response = OpenAIChatCompletion().audio_speech(
client=OpenAI(api_key="transport-only", http_client=http_client), **_audio_speech_kwargs()
)
_assert_provider_headers_recorded(response)
@pytest.mark.asyncio
async def test_async_audio_speech_records_provider_response_headers():
async with httpx.AsyncClient(transport=_speech_transport()) as http_client:
response = await OpenAIChatCompletion().audio_speech(
client=AsyncOpenAI(api_key="transport-only", http_client=http_client),
aspeech=True,
**_audio_speech_kwargs(),
)
_assert_provider_headers_recorded(response)

View file

@ -0,0 +1,71 @@
from typing import Final
from unittest.mock import Mock
import httpx
import pytest
from openai import AsyncOpenAI, OpenAI
from litellm.llms.openai.transcriptions.handler import OpenAIAudioTranscription
from litellm.types.utils import TranscriptionResponse
_PROVIDER_HEADERS: Final = {"x-request-id": "req_stt", "x-ratelimit-remaining-requests": "41"}
def _transcription_transport() -> httpx.MockTransport:
return httpx.MockTransport(lambda request: httpx.Response(200, json={"text": "hello"}, headers=_PROVIDER_HEADERS))
def _logging_obj() -> Mock:
logging_obj = Mock()
logging_obj.model_call_details = {}
return logging_obj
def _call_kwargs(logging_obj: Mock) -> dict:
return {
"model": "gpt-4o-mini-transcribe",
"audio_file": ("audio.wav", b"riff-bytes", "audio/wav"),
"optional_params": {},
"litellm_params": {},
"model_response": TranscriptionResponse(),
"timeout": 10.0,
"max_retries": 0,
"logging_obj": logging_obj,
"api_key": "transport-only",
"api_base": None,
}
def _assert_headers_recorded(response: TranscriptionResponse, logging_obj: Mock) -> None:
assert response.text == "hello"
assert response._hidden_params["headers"]["x-request-id"] == "req_stt"
assert response._hidden_params["additional_headers"]["llm_provider-x-request-id"] == "req_stt"
assert response._hidden_params["additional_headers"]["x-ratelimit-remaining-requests"] == "41"
assert logging_obj.model_call_details["response_headers"]["x-request-id"] == "req_stt"
def test_audio_transcriptions_records_provider_response_headers():
logging_obj = _logging_obj()
with httpx.Client(transport=_transcription_transport()) as http_client:
response = OpenAIAudioTranscription().audio_transcriptions(
client=OpenAI(api_key="transport-only", http_client=http_client),
atranscription=False,
**_call_kwargs(logging_obj),
)
_assert_headers_recorded(response, logging_obj)
@pytest.mark.asyncio
async def test_async_audio_transcriptions_records_provider_response_headers():
logging_obj = _logging_obj()
async with httpx.AsyncClient(transport=_transcription_transport()) as http_client:
response = await OpenAIAudioTranscription().audio_transcriptions(
client=AsyncOpenAI(api_key="transport-only", http_client=http_client),
atranscription=True,
**_call_kwargs(logging_obj),
)
_assert_headers_recorded(response, logging_obj)

View file

@ -13,13 +13,14 @@ There are no real I/O seams here; ``uuid.uuid4`` is the only nondeterministic
dependency and is patched where the displayName is asserted.
"""
from collections.abc import Mapping
from typing import Final
from unittest.mock import patch
import pytest
from litellm.llms.vertex_ai.batches.transformation import ( # noqa: E402
VertexAIBatchTransformation,
vertex_prompt_tokens_details,
)
from litellm.llms.vertex_ai.common_utils import ( # noqa: E402
VertexAIError,
@ -36,32 +37,10 @@ INPUT_FILE = (
ENDPOINT_ID = "7768560373388541952"
ENDPOINT_INPUT_FILE = (
f"gs://litellm-testing-bucket/litellm-vertex-files/endpoints/{ENDPOINT_ID}/"
"e9412502-2c91-42a6-8e61-f5c294cc0fc8"
f"gs://litellm-testing-bucket/litellm-vertex-files/endpoints/{ENDPOINT_ID}/e9412502-2c91-42a6-8e61-f5c294cc0fc8"
)
def test_vertex_prompt_tokens_details_rejects_malformed_details():
assert vertex_prompt_tokens_details({"promptTokensDetails": [1]}) is None
assert vertex_prompt_tokens_details({"promptTokensDetails": [{"modality": "AUDIO"}]}) is None
assert (
vertex_prompt_tokens_details(
{
"promptTokensDetails": [
{"modality": "AUDIO", "tokenCount": 1},
"malformed",
]
}
)
is None
)
# =========================================================================== #
# transform_openai_batch_request_to_vertex_ai_batch_request
# =========================================================================== #
def test_transform_openai_request_builds_full_vertex_job():
with patch(
"litellm.llms.vertex_ai.batches.transformation.uuid.uuid4",
@ -270,9 +249,78 @@ def test_get_input_file_id_empty_uris():
# =========================================================================== #
# _get_output_file_id_from_vertex_ai_batch_response
# _get_output_file_id_from_vertex_ai_batch_response: None until Vertex reports outputInfo
# =========================================================================== #
SHARED_OUTPUT_PREFIX: Final = "gs://bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-flash"
SUCCEEDED_OUTPUT_DIRECTORY: Final = f"{SHARED_OUTPUT_PREFIX}/prediction-model-2026-09-24T19:41:00.000000Z"
def _vertex_job(state: str) -> dict[str, object]:
return {
"name": "projects/510528649030/locations/us-central1/batchPredictionJobs/3814889423749775360",
"state": state,
"createTime": "2026-09-24T19:37:25.775603Z",
"inputConfig": {
"instancesFormat": "jsonl",
"gcsSource": {"uris": [f"{SHARED_OUTPUT_PREFIX}/0586ba52-4f8b-4988-aa8d-3573550a4b0f"]},
},
"outputConfig": {
"predictionsFormat": "jsonl",
"gcsDestination": {"outputUriPrefix": SHARED_OUTPUT_PREFIX},
},
}
@pytest.mark.parametrize(
"vertex_state,output_info_field,expected_status,expected_output_file_id",
[
("JOB_STATE_PENDING", {}, "validating", None),
("JOB_STATE_RUNNING", {"outputInfo": {}}, "in_progress", None),
("JOB_STATE_CANCELLED", {"outputInfo": None}, "cancelled", None),
(
"JOB_STATE_SUCCEEDED",
{"outputInfo": {"gcsOutputDirectory": SUCCEEDED_OUTPUT_DIRECTORY}},
"completed",
f"{SUCCEEDED_OUTPUT_DIRECTORY}/predictions.jsonl",
),
],
ids=["create_or_pending", "running", "cancelled", "succeeded"],
)
def test_transform_vertex_response_output_file_id_is_none_until_output_info(
vertex_state: str,
output_info_field: Mapping[str, object],
expected_status: str,
expected_output_file_id: str | None,
) -> None:
batch: Final = T.transform_vertex_ai_batch_response_to_openai_batch_response(
{**_vertex_job(vertex_state), **output_info_field}
)
assert batch.status == expected_status
assert batch.output_file_id == expected_output_file_id
@pytest.mark.parametrize(
"response",
[
{},
{"outputConfig": {}},
{"outputInfo": None},
{"outputInfo": {"gcsOutputDirectory": ""}},
{"outputInfo": {"gcsOutputDirectory": None}},
],
ids=[
"no_fields",
"output_config_without_destination",
"null_output_info",
"empty_output_directory",
"null_output_directory",
],
)
def test_get_output_file_id_is_none_without_output_directory(response: Mapping[str, object]) -> None:
assert T._get_output_file_id_from_vertex_ai_batch_response(response) is None
def test_get_output_file_id_from_output_info():
# outputInfo branch: rstrip trailing slash, append predictions.jsonl
@ -289,49 +337,7 @@ def test_get_output_file_id_output_info_no_trailing_slash():
)
def test_get_output_file_id_empty_output_info_falls_through_to_output_config():
# gcsOutputDirectory missing -> "" -> the "/predictions.jsonl" guard skips
# the outputInfo branch, falls through to outputConfig
resp = {
"outputInfo": {},
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}},
}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_explicit_none_falls_through_to_output_config():
resp = {
"outputInfo": None,
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg"}},
}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_explicit_none_and_no_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response({"outputInfo": None}) == ""
def test_get_output_file_id_no_output_info_and_no_output_config():
assert T._get_output_file_id_from_vertex_ai_batch_response({}) == ""
def test_get_output_file_id_output_config_missing_gcs_destination():
# outputConfig present but no gcsDestination -> returns the running "" value
assert T._get_output_file_id_from_vertex_ai_batch_response({"outputConfig": {}}) == ""
def test_get_output_file_id_output_config_already_has_suffix():
# outputUriPrefix already ends in /predictions.jsonl -> returned as-is (no double append)
resp = {"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg/predictions.jsonl"}}}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_config_strips_trailing_slash():
resp = {"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://b/cfg/"}}}
assert T._get_output_file_id_from_vertex_ai_batch_response(resp) == "gs://b/cfg/predictions.jsonl"
def test_get_output_file_id_output_info_takes_precedence_over_output_config():
def test_get_output_file_id_output_info_ignores_output_uri_prefix():
resp = {
"outputInfo": {"gcsOutputDirectory": "gs://from-info"},
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://from-config"}},
@ -477,3 +483,19 @@ def test_list_response_none_jobs_treated_as_empty():
out = T.transform_vertex_ai_batch_list_response_to_openai_list_response({"batchPredictionJobs": None})
assert out["data"] == []
assert out["first_id"] is None
PASSTHROUGH_INPUT_FILE = (
"gs://litellm-testing-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/uuid-1"
)
def test_get_model_from_passthrough_gcs_file():
assert T._get_model_from_gcs_file(PASSTHROUGH_INPUT_FILE) == "publishers/google/models/gemini-2.5-flash"
def test_get_gcs_uri_prefix_keeps_passthrough_segment_so_output_lands_beside_input():
assert (
T._get_gcs_uri_prefix_from_file(PASSTHROUGH_INPUT_FILE)
== "gs://litellm-testing-bucket/litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash"
)

View file

@ -0,0 +1,310 @@
import io
import json
import urllib.parse
from pathlib import Path
from unittest.mock import MagicMock
from urllib.parse import parse_qs, urlparse
import httpx
import pytest
from litellm.llms.vertex_ai.common_utils import VertexAIError
from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig, is_passthrough_managed_gcs_url
NATIVE_VERTEX_ROW = json.dumps(
{
"request": {
"contents": [{"role": "user", "parts": [{"text": "What is the tallest building in the world?"}]}],
"tools": [{"googleSearch": {"excludeDomains": ["example.com"]}}],
}
}
).encode()
NATIVE_VERTEX_JSONL = NATIVE_VERTEX_ROW + b"\n" + NATIVE_VERTEX_ROW + b"\n"
OPENAI_BATCH_JSONL = (
b'{"custom_id": "r1", "method": "POST", "url": "/v1/chat/completions",'
b' "body": {"model": "gemini-2.5-flash", "messages": [{"role": "user", "content": "hi"}]}}\n'
)
PASSTHROUGH_OBJECT = (
"litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/uuid-1/predictions.jsonl"
)
TRANSFORMED_OBJECT = "litellm-vertex-files/publishers/google/models/gemini-2.5-flash/uuid-1/predictions.jsonl"
UPLOAD_CHUNK_BYTES = 1024 * 1024
@pytest.fixture
def config() -> VertexAIFilesConfig:
return VertexAIFilesConfig()
def _gcs_media_url(object_name: str) -> str:
return (
f"https://storage.googleapis.com/storage/v1/b/my-bucket/o/{urllib.parse.quote(object_name, safe='')}?alt=media"
)
def _native_output_jsonl() -> bytes:
return (
json.dumps(
{
"request": json.loads(NATIVE_VERTEX_ROW)["request"],
"status": "",
"response": {
"candidates": [
{
"content": {"role": "model", "parts": [{"text": "The Burj Khalifa."}]},
"finishReason": "STOP",
"groundingMetadata": {"webSearchQueries": ["tallest building in the world"]},
}
],
"modelVersion": "gemini-2.5-flash",
"usageMetadata": {"promptTokenCount": 20, "candidatesTokenCount": 48, "totalTokenCount": 68},
},
"processed_time": "2026-09-23T19:02:00.000+00:00",
}
).encode()
+ b"\n"
)
def _upload_chunks(config: VertexAIFilesConfig, file: object, litellm_params: dict) -> list[bytes]:
body = config.transform_create_file_request(
model="",
create_file_data={"file": file, "purpose": "batch"},
optional_params={},
litellm_params=litellm_params,
)
return list(body["streaming_media_upload"]["body_stream"].iter_bytes())
def _upload_body_bytes(config: VertexAIFilesConfig, file: object, litellm_params: dict) -> bytes:
return b"".join(_upload_chunks(config, file, litellm_params))
class TestPassthroughBatchUpload:
"""`passthrough=True` on a batch upload ships the caller's native Vertex JSONL
to GCS byte for byte, filed under a `passthrough/` object path so the batch
output that lands beside it is recognized and returned untouched as well."""
def _upload_url(self, config, litellm_params, file, purpose="batch") -> str:
return config.get_complete_file_url(
api_base=None,
api_key=None,
model="",
optional_params={},
litellm_params=litellm_params,
data={"file": file, "purpose": purpose},
)
def test_passthrough_object_is_filed_under_passthrough_prefix_named_by_deployment_model(self, config):
url = self._upload_url(
config,
{"gcs_bucket_name": "my-bucket", "model": "vertex_ai/gemini-2.5-flash", "passthrough": True},
("batch.jsonl", NATIVE_VERTEX_JSONL, "application/jsonl"),
)
object_name = parse_qs(urlparse(url).query)["name"][0]
assert object_name.startswith("litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash/")
def test_passthrough_upload_without_deployment_model_is_rejected(self, config):
with pytest.raises(VertexAIError) as exc_info:
self._upload_url(
config,
{"gcs_bucket_name": "my-bucket", "passthrough": True},
("batch.jsonl", NATIVE_VERTEX_JSONL, "application/jsonl"),
)
assert exc_info.value.status_code == 400
assert "target_model_names" in exc_info.value.message
def test_passthrough_flag_does_not_ship_a_non_batch_upload_raw(self, config):
result = config.transform_create_file_request(
model="",
create_file_data={"file": ("notes.txt", b"plain text", "text/plain"), "purpose": "user_data"},
optional_params={},
litellm_params={"gcs_bucket_name": "my-bucket", "passthrough": True},
)
assert result == b"plain text"
def test_passthrough_flag_is_ignored_for_non_batch_purposes(self, config):
url = self._upload_url(
config,
{"gcs_bucket_name": "my-bucket", "model": "vertex_ai/gemini-2.5-flash", "passthrough": True},
("notes.txt", b"plain text", "text/plain"),
purpose="user_data",
)
object_name = parse_qs(urlparse(url).query)["name"][0]
assert object_name.startswith("litellm-vertex-files/uploads/")
assert "passthrough" not in object_name
@pytest.mark.parametrize(
"file",
[
("batch.jsonl", NATIVE_VERTEX_JSONL, "application/jsonl"),
NATIVE_VERTEX_JSONL,
("batch.jsonl", io.BytesIO(NATIVE_VERTEX_JSONL), "application/jsonl"),
("batch.jsonl", NATIVE_VERTEX_JSONL.decode(), "application/jsonl"),
],
ids=["bytes-tuple", "bare-bytes", "handle-tuple", "text-tuple"],
)
def test_passthrough_upload_body_is_the_callers_bytes(self, config, file):
body = config.transform_create_file_request(
model="",
create_file_data={"file": file, "purpose": "batch"},
optional_params={},
litellm_params={"passthrough": True},
)
stream = body["streaming_media_upload"]["body_stream"]
assert b"".join(stream.iter_bytes()) == NATIVE_VERTEX_JSONL
assert b"".join(stream.iter_bytes()) == NATIVE_VERTEX_JSONL
assert body["streaming_media_upload"]["content_type"] == "application/json"
def test_passthrough_upload_streams_a_large_handle_in_bounded_chunks(self, config):
content = NATIVE_VERTEX_ROW * (3 * UPLOAD_CHUNK_BYTES // len(NATIVE_VERTEX_ROW) + 1)
chunks = _upload_chunks(
config, ("batch.jsonl", io.BytesIO(content), "application/jsonl"), {"passthrough": True}
)
assert len(chunks) >= 3
assert max(len(chunk) for chunk in chunks) <= UPLOAD_CHUNK_BYTES
assert b"".join(chunks) == content
def test_passthrough_upload_streams_a_path_in_bounded_chunks(self, config, tmp_path: Path):
content = NATIVE_VERTEX_ROW * (2 * UPLOAD_CHUNK_BYTES // len(NATIVE_VERTEX_ROW) + 1)
batch_path = tmp_path / "batch.jsonl"
batch_path.write_bytes(content)
chunks = _upload_chunks(config, ("batch.jsonl", batch_path, "application/jsonl"), {"passthrough": True})
assert len(chunks) >= 2
assert max(len(chunk) for chunk in chunks) <= UPLOAD_CHUNK_BYTES
assert b"".join(chunks) == content
def test_passthrough_upload_rejects_a_non_seekable_handle(self, config):
class _Pipe:
def read(self, size=-1):
return b""
with pytest.raises(ValueError, match="seekable"):
_upload_body_bytes(config, ("batch.jsonl", _Pipe(), "application/jsonl"), {"passthrough": True})
def test_passthrough_upload_rejects_content_that_is_neither_bytes_path_nor_handle(self, config):
with pytest.raises(ValueError, match="Unsupported file content type"):
_upload_body_bytes(config, ("batch.jsonl", 42, "application/jsonl"), {"passthrough": True})
def test_openai_rows_are_translated_unless_passthrough_is_set(self, config):
file = ("batch.jsonl", OPENAI_BATCH_JSONL, "application/jsonl")
translated = _upload_body_bytes(config, file, {})
untouched = _upload_body_bytes(config, file, {"passthrough": True})
assert untouched == OPENAI_BATCH_JSONL
assert translated != OPENAI_BATCH_JSONL
assert b'"contents"' in translated
def test_passthrough_output_content_is_returned_untouched(self, config):
raw_jsonl = _native_output_jsonl()
def _download(object_name: str) -> bytes:
raw_response = httpx.Response(
status_code=200,
content=raw_jsonl,
headers={"content-type": "application/octet-stream"},
request=httpx.Request("GET", _gcs_media_url(object_name)),
)
result = config.transform_file_content_response(
raw_response=raw_response, logging_obj=MagicMock(), litellm_params={}
)
return result.response.content
assert _download(PASSTHROUGH_OBJECT) == raw_jsonl
assert _download(f"team-a/{PASSTHROUGH_OBJECT}") == raw_jsonl
transformed = _download(TRANSFORMED_OBJECT)
assert transformed != raw_jsonl
assert json.loads(transformed.splitlines()[0])["response"]["body"]["choices"]
nested = _download(f"litellm-vertex-files/{PASSTHROUGH_OBJECT}")
assert nested != raw_jsonl
assert json.loads(nested.splitlines()[0])["response"]["body"]["choices"]
def test_output_of_an_upload_whose_model_smuggles_the_passthrough_segment_is_still_transformed(self, config):
smuggled_model = b"litellm-vertex-files/passthrough/publishers/google/models/gemini-2.5-flash"
upload_url = self._upload_url(
config,
{"gcs_bucket_name": "my-bucket"},
("batch.jsonl", OPENAI_BATCH_JSONL.replace(b"gemini-2.5-flash", smuggled_model), "application/jsonl"),
)
object_name = parse_qs(urlparse(upload_url).query)["name"][0]
raw_jsonl = _native_output_jsonl()
raw_response = httpx.Response(
status_code=200,
content=raw_jsonl,
headers={"content-type": "application/octet-stream"},
request=httpx.Request("GET", _gcs_media_url(f"{object_name}/predictions.jsonl")),
)
result = config.transform_file_content_response(
raw_response=raw_response, logging_obj=MagicMock(), litellm_params={}
)
assert object_name.startswith("litellm-vertex-files/litellm-vertex-files/passthrough/")
assert json.loads(result.response.content.splitlines()[0])["response"]["body"]["choices"]
@pytest.mark.parametrize(
"url, expected",
[
(f"gs://my-bucket/{PASSTHROUGH_OBJECT}", True),
(f"gs://my-bucket/team-a/{PASSTHROUGH_OBJECT}", True),
(f"gs://my-bucket/litellm-vertex-files/{PASSTHROUGH_OBJECT}", False),
(_gcs_media_url(f"team-a/sub/{PASSTHROUGH_OBJECT}"), True),
(_gcs_media_url(f"litellm-vertex-files/publishers/google/models/x/{PASSTHROUGH_OBJECT}"), False),
(_gcs_media_url(TRANSFORMED_OBJECT), False),
],
ids=["gs", "gs-prefixed", "gs-smuggled", "https-prefixed", "https-model-path-smuggled", "https-transformed"],
)
def test_passthrough_detection_anchors_on_the_first_managed_segment(self, url, expected):
assert is_passthrough_managed_gcs_url(url) is expected
@pytest.mark.asyncio
async def test_passthrough_output_stream_is_returned_untouched(self, config):
stream_iterator = object()
headers = {"content-type": "application/octet-stream"}
result = await config.transform_file_content_stream(
stream_iterator=stream_iterator,
headers=headers,
request_url=_gcs_media_url(f"team-a/{PASSTHROUGH_OBJECT}"),
logging_obj=MagicMock(),
litellm_params={},
)
assert result.stream_iterator is stream_iterator
assert result.headers == headers
class TestEmbeddingOutputTranslation:
EMBEDDING_OBJECT = (
"litellm-vertex-files/publishers/google/models/gemini-embedding-2/prediction-model-1/predictions.jsonl"
)
def _transform(self, config: VertexAIFilesConfig, rows: list[dict]) -> list[dict]:
raw_response = httpx.Response(
status_code=200,
content="\n".join(json.dumps(row) for row in rows).encode(),
headers={"content-type": "application/octet-stream"},
request=httpx.Request("GET", _gcs_media_url(self.EMBEDDING_OBJECT)),
)
result = config.transform_file_content_response(
raw_response=raw_response, logging_obj=MagicMock(), litellm_params={}
)
return [json.loads(line) for line in result.response.content.decode().splitlines()]
def test_embedding_rows_become_openai_batch_rows_billed_by_their_prompt_tokens(self, config):
live_row = {
"key": "request-1",
"request": {"content": {"parts": [{"text": "hello world"}]}},
"response": {"embedding": {"values": [-0.015, 0.024]}, "usageMetadata": {"promptTokenCount": 2}},
}
documented_row = {
"key": "request-2",
"request": {"content": {"parts": [{"text": "hello"}]}},
"response": {"embedding": {"values": [0.5]}, "tokenCount": "3"},
}
live, documented = self._transform(config, [live_row, documented_row])
assert (live["custom_id"], live["error"], live["response"]["status_code"]) == ("request-1", None, 200)
assert live["response"]["body"]["model"] == "gemini-embedding-2"
assert live["response"]["body"]["data"] == [{"embedding": [-0.015, 0.024], "index": 0, "object": "embedding"}]
live_usage, documented_usage = (row["response"]["body"]["usage"] for row in (live, documented))
assert (live_usage["prompt_tokens"], live_usage["total_tokens"]) == (2, 2)
assert (documented_usage["prompt_tokens"], documented_usage["total_tokens"]) == (3, 3)

View file

@ -26,12 +26,12 @@ def test_output_file_id_uses_predictions_jsonl_with_output_info():
)
def test_output_file_id_falls_back_to_output_uri_prefix_with_predictions_jsonl():
def test_output_file_id_is_none_until_output_info():
response = {
"outputInfo": {},
"outputConfig": {
"gcsDestination": {
"outputUriPrefix": "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456"
"outputUriPrefix": "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro"
}
},
}
@ -42,10 +42,7 @@ def test_output_file_id_falls_back_to_output_uri_prefix_with_predictions_jsonl()
)
)
assert (
output_file_id
== "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456/predictions.jsonl"
)
assert output_file_id is None
def test_vertex_ai_cancel_batch():

View file

@ -3517,6 +3517,7 @@ def test_internal_user_still_blocked_from_another_users_info():
[
"/user/daily/activity",
"/user/daily/activity/aggregated",
"/user/daily/activity/aggregated/search",
],
)
@pytest.mark.parametrize(
@ -3599,6 +3600,55 @@ def test_user_daily_activity_aggregated_not_covered_by_prefix_match():
)
@pytest.mark.parametrize(
"route",
[
"/team/daily/activity",
"/team/daily/activity/aggregated",
"/team/daily/activity/aggregated/search",
],
)
@pytest.mark.parametrize(
"user_role",
[
LitellmUserRoles.INTERNAL_USER.value,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value,
],
)
def test_team_daily_activity_routes_reachable_by_non_admin(route, user_role):
"""The Team Usage dashboard calls all three team daily-activity routes, and
each handler self-scopes to the caller's teams and own keys
(_resolve_team_daily_activity_scope). self_managed_routes is the only list
granting them to a non-admin, and check_route_access is exact-match, so each
sub-path needs its own entry: dropping one 401s the dashboard before the
handler ever runs.
"""
user_obj = LiteLLM_UserTable(
user_id="test_user",
user_email="test@example.com",
user_role=user_role,
)
valid_token = UserAPIKeyAuth(user_id="test_user", user_role=user_role)
request = MagicMock(spec=Request)
request.query_params = {}
def outcome() -> str:
try:
RouteChecks.non_proxy_admin_allowed_routes_check(
user_obj=user_obj,
_user_role=user_role,
route=route,
request=request,
valid_token=valid_token,
request_data={},
)
except Exception as exc:
return f"denied: {exc}"
return "allowed"
assert outcome() == "allowed"
@pytest.mark.parametrize(
"user_role",
[

View file

@ -994,11 +994,11 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster
assert lock_statement is _TEAM_ADVISORY_LOCK_SQL
assert locked_team_id == team_id
assert "pg_advisory_xact_lock(hashtext($1))" in lock_statement
statement, user_ids, team_ids, costs = spend_call.args
statement, members = spend_call.args
assert statement is _TEAM_MEMBER_SPEND_SQL
assert (list(user_ids), list(team_ids), list(costs)) == ([user_id], [team_id], [response_cost])
assert json.loads(members) == [{"user_id": user_id, "team_id": team_id, "cost": response_cost}]
assert 'INSERT INTO "LiteLLM_TeamMembership"' in statement
assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', p.user_id))" in statement
assert "members_with_roles @> jsonb_build_array(jsonb_build_object('user_id', member.user_id))" in statement
assert "ON CONFLICT (user_id, team_id) DO UPDATE" in statement
assert 'spend = "LiteLLM_TeamMembership".spend + EXCLUDED.spend' in statement
assert 'total_spend = "LiteLLM_TeamMembership".total_spend + EXCLUDED.total_spend' in statement
@ -1007,7 +1007,7 @@ async def test_commit_spend_updates_to_db_writes_team_member_spend_in_one_roster
@pytest.mark.asyncio
async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_user():
"""
The member spend statement touches rows in the order of its input arrays, so the batch
The member spend statement touches rows in the order of its input rows, so the batch
is handed over sorted by (team_id, user_id), with each cost kept next to its member, and
each distinct team is locked once, in `sorted(team_ids)` order, the order /team/delete
locks in, so a concurrent flush and delete cannot deadlock. `eng` and `eng2` pin that:
@ -1034,13 +1034,13 @@ async def test_commit_spend_updates_to_db_orders_team_member_rows_by_team_then_u
)
*lock_calls, spend_call = mock_transaction.execute_raw.await_args_list
_statement, user_ids, team_ids, costs = spend_call.args
_statement, members = spend_call.args
assert [lock_call.args for lock_call in lock_calls] == [
(_TEAM_ADVISORY_LOCK_SQL, "eng"),
(_TEAM_ADVISORY_LOCK_SQL, "eng-b"),
(_TEAM_ADVISORY_LOCK_SQL, "eng2"),
]
assert list(zip(team_ids, user_ids, costs)) == [
assert [(row["team_id"], row["user_id"], row["cost"]) for row in json.loads(members)] == [
("eng", "user_x", 0.3),
("eng", "user_y", 0.2),
("eng-b", "user_x", 0.4),

View file

@ -25,6 +25,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
get_api_key_metadata,
get_daily_activity,
get_daily_activity_aggregated,
get_daily_activity_export_rows,
global_rollup_reconciled_through,
update_metrics,
)
@ -2868,3 +2869,307 @@ def test_spend_logs_window_is_none_when_no_date_parses():
from litellm.proxy.management_endpoints.common_daily_activity import _spend_logs_window
assert _spend_logs_window({"garbage", ""}) is None
_DAILY_TEAM_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyTeamSpend" (
id TEXT PRIMARY KEY,
team_id TEXT,
date TEXT NOT NULL,
api_key TEXT NOT NULL,
model TEXT,
model_group TEXT,
custom_llm_provider TEXT,
mcp_namespaced_tool_name TEXT,
endpoint TEXT,
prompt_tokens BIGINT DEFAULT 0,
completion_tokens BIGINT DEFAULT 0,
cache_read_input_tokens BIGINT DEFAULT 0,
cache_creation_input_tokens BIGINT DEFAULT 0,
compression_saved_tokens BIGINT DEFAULT 0,
compression_savings_spend DOUBLE PRECISION DEFAULT 0,
prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
autorouter_savings_spend DOUBLE PRECISION DEFAULT 0,
spend DOUBLE PRECISION DEFAULT 0,
ptu_flat_cost DOUBLE PRECISION DEFAULT 0,
api_requests BIGINT DEFAULT 0,
successful_requests BIGINT DEFAULT 0,
failed_requests BIGINT DEFAULT 0,
total_response_time_ms BIGINT DEFAULT 0,
timed_requests BIGINT DEFAULT 0
)
"""
def _seed_daily_team_spend(conn: psycopg.Connection, rows: Sequence[tuple[object, ...]]) -> None:
with conn.cursor() as cur:
cur.execute(_DAILY_TEAM_SPEND_DDL)
cur.executemany(
"""
INSERT INTO "LiteLLM_DailyTeamSpend"
(id, team_id, date, api_key, model, model_group, custom_llm_provider,
endpoint, prompt_tokens, spend, ptu_flat_cost, api_requests, successful_requests)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""",
rows,
)
conn.commit()
def _team_spend_row(
row_id: str,
team_id: str,
api_key: str,
spend: float,
*,
date: str = "2026-06-01",
model: str = "gpt-5",
ptu_flat_cost: float = 0.0,
) -> tuple[object, ...]:
return (
row_id,
team_id,
date,
api_key,
model,
"",
"openai",
"/v1/chat/completions",
10,
spend,
ptu_flat_cost,
1,
1,
)
def _export_prisma(conn: psycopg.Connection, token_rows: Sequence[SimpleNamespace] = ()) -> MagicMock:
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(conn, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=list(token_rows))
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
return mock_prisma
@pytest.mark.asyncio
async def test_export_keys_returns_every_key_beyond_the_top_n_cap(
_aggregated_postgresql: psycopg.Connection,
):
"""The export route exists because the aggregated route caps the per-key arm at
USAGE_TOP_API_KEYS_LIMIT. With more keys than the cap every one of them must
land in the export, while the PTU sentinel stays out of the key view."""
n_keys: Final = USAGE_TOP_API_KEYS_LIMIT + 7
_seed_daily_team_spend(
_aggregated_postgresql,
[
*[_team_spend_row(f"row-{i:03d}", "team-1", f"key-{i:03d}", float(i + 1)) for i in range(n_keys)],
_team_spend_row("row-ptu", "team-1", PTU_SENTINEL_API_KEY, 0.0, ptu_flat_cost=1000.0),
],
)
rows = await get_daily_activity_export_rows(
prisma_client=_export_prisma(_aggregated_postgresql),
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
api_key=None,
exclude_entity_ids=None,
timezone_offset_minutes=None,
export_type="daily_with_keys",
)
assert {row.api_key for row in rows} == {f"key-{i:03d}" for i in range(n_keys)}
assert len(rows) == n_keys
assert all(row.team_id == "team-1" for row in rows)
by_key: Final = {row.api_key: row for row in rows}
assert by_key["key-000"].spend == pytest.approx(1.0)
assert sum(row.spend for row in rows) == pytest.approx(n_keys * (n_keys + 1) / 2)
assert all(row.total_tokens == 10 and row.api_requests == 1 for row in rows)
@pytest.mark.asyncio
async def test_export_daily_keeps_ptu_sentinel_in_the_team_rollup(
_aggregated_postgresql: psycopg.Connection,
):
"""The plain daily export groups by (date, team), so the sentinel's flat cost
must land in the team row exactly like breakdown.entities on the aggregated
route; dropping it would silently under-report team spend."""
_seed_daily_team_spend(
_aggregated_postgresql,
[
_team_spend_row("row-1", "team-1", "key-1", 2.0),
_team_spend_row("row-ptu", "team-1", PTU_SENTINEL_API_KEY, 0.0, ptu_flat_cost=0.0),
],
)
with _aggregated_postgresql.cursor() as cur:
cur.execute("UPDATE \"LiteLLM_DailyTeamSpend\" SET spend = 1000.0 WHERE id = 'row-ptu'")
_aggregated_postgresql.commit()
rows = await get_daily_activity_export_rows(
prisma_client=_export_prisma(_aggregated_postgresql),
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field={"team-1": {"team_alias": "Alpha"}},
start_date="2026-06-01",
end_date="2026-06-01",
api_key=None,
exclude_entity_ids=None,
timezone_offset_minutes=None,
export_type="daily",
)
assert len(rows) == 1
assert rows[0].team_id == "team-1"
assert rows[0].team_alias == "Alpha"
assert rows[0].api_key is None
assert rows[0].spend == pytest.approx(1002.0)
@pytest.mark.asyncio
async def test_export_users_folds_keys_into_one_row_per_user(
_aggregated_postgresql: psycopg.Connection,
):
"""daily_with_users runs the per-key rollup then folds in Python: two keys of
user-1 merge into one row with keys=2 and summed metrics, and the distinct
user keeps its own row."""
_seed_daily_team_spend(
_aggregated_postgresql,
[
_team_spend_row("row-1", "team-1", "key-1", 2.0),
_team_spend_row("row-2", "team-1", "key-2", 3.0),
_team_spend_row("row-3", "team-1", "key-3", 5.0),
],
)
tokens: Final = tuple(
SimpleNamespace(token=token, key_alias=None, team_id="team-1", user_id=user_id)
for token, user_id in (("key-1", "user-1"), ("key-2", "user-1"), ("key-3", "user-2"))
)
rows = await get_daily_activity_export_rows(
prisma_client=_export_prisma(_aggregated_postgresql, tokens),
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
api_key=None,
exclude_entity_ids=None,
timezone_offset_minutes=None,
export_type="daily_with_users",
)
assert [(row.user_id, row.keys, row.spend, row.api_requests, row.total_tokens) for row in rows] == [
("user-1", 2, 5.0, 2, 20),
("user-2", 1, 5.0, 1, 10),
]
@pytest.mark.asyncio
async def test_export_models_rolls_up_per_team_and_model(
_aggregated_postgresql: psycopg.Connection,
):
_seed_daily_team_spend(
_aggregated_postgresql,
[
_team_spend_row("row-1", "team-1", "key-1", 2.0, model="gpt-5"),
_team_spend_row("row-2", "team-1", "key-2", 3.0, model="gpt-5"),
_team_spend_row("row-3", "team-1", "key-1", 5.0, model="claude"),
],
)
rows = await get_daily_activity_export_rows(
prisma_client=_export_prisma(_aggregated_postgresql),
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
api_key=None,
exclude_entity_ids=None,
timezone_offset_minutes=None,
export_type="daily_with_models",
)
assert [(row.model, row.spend, row.api_requests) for row in rows] == [
("claude", 5.0, 1),
("gpt-5", 5.0, 2),
]
@pytest.mark.asyncio
async def test_export_daily_reports_ptu_flat_cost_on_the_team_row(
_aggregated_postgresql: psycopg.Connection, ptu_cost_attribution_enabled
):
"""The CSV the dashboard hands to finance must match the client-side export,
which shows flat cost columns once any PTU spend exists for the day."""
from litellm.proxy.management_endpoints.team_endpoints import _team_export_csv
_seed_daily_team_spend(
_aggregated_postgresql,
[
_team_spend_row("row-1", "team-1", "key-1", 2.0),
_team_spend_row("row-ptu", "team-1", PTU_SENTINEL_API_KEY, 0.0, ptu_flat_cost=240.0),
],
)
rows = await get_daily_activity_export_rows(
prisma_client=_export_prisma(_aggregated_postgresql),
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
api_key=None,
exclude_entity_ids=None,
timezone_offset_minutes=None,
export_type="daily",
)
assert len(rows) == 1
assert rows[0].flat_cost == pytest.approx(240.0)
header: Final = _team_export_csv("daily", rows).splitlines()[0]
assert "Spend ($),Flat Cost ($),Total Cost ($)" in header
record: Final = _team_export_csv("daily", rows).splitlines()[1].split(",")
spend_index: Final = header.split(",").index("Spend ($)")
assert record[spend_index : spend_index + 3] == ["2.0000", "240.0000", "242.0000"]
@pytest.mark.asyncio
async def test_export_csv_omits_flat_cost_columns_when_no_ptu_spend_exists(
_aggregated_postgresql: psycopg.Connection,
):
from litellm.proxy.management_endpoints.team_endpoints import _team_export_csv
_seed_daily_team_spend(
_aggregated_postgresql,
[_team_spend_row("row-1", "team-1", "key-1", 2.0)],
)
rows = await get_daily_activity_export_rows(
prisma_client=_export_prisma(_aggregated_postgresql),
table_name="litellm_dailyteamspend",
entity_id_field="team_id",
entity_id="team-1",
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
api_key=None,
exclude_entity_ids=None,
timezone_offset_minutes=None,
export_type="daily",
)
assert rows[0].flat_cost == 0.0
header: Final = _team_export_csv("daily", rows).splitlines()[0]
assert "Flat Cost" not in header
assert "Total Cost" not in header

View file

@ -2659,6 +2659,172 @@ async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_us
assert mock_get_daily_agg.call_args.kwargs["entity_id"] == "regular-user-123"
@pytest.mark.asyncio
async def test_search_user_daily_activity_keys_passes_matched_tokens_to_aggregation(monkeypatch):
"""The search endpoint resolves matching verification tokens by hash, alias, or
user id, then aggregates daily spend for exactly those tokens. This is what lets
the Usage page find keys outside the top-spend subset the aggregated endpoint caps."""
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy.management_endpoints.internal_user_endpoints import (
search_user_daily_activity_keys,
)
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[SimpleNamespace(token="tok-a"), SimpleNamespace(token="tok-b")]
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_response = MagicMock()
mock_get_daily_agg = AsyncMock(return_value=mock_response)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
mock_get_daily_agg,
)
admin_key_dict = UserAPIKeyAuth(
user_id="admin-user-001",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
result = await search_user_daily_activity_keys(
search="gamma",
start_date="2025-02-01",
end_date="2025-02-28",
user_id=None,
timezone=480,
include_current_utc_day=False,
user_api_key_dict=admin_key_dict,
)
assert result is mock_response
find_many_kwargs = mock_prisma_client.db.litellm_verificationtoken.find_many.call_args.kwargs
assert find_many_kwargs["take"] == USAGE_TOP_API_KEYS_LIMIT
assert find_many_kwargs["where"]["OR"] == (
{"token": "gamma"},
{"key_alias": {"contains": "gamma", "mode": "insensitive"}},
{"user_id": {"contains": "gamma", "mode": "insensitive"}},
)
assert "user_id" not in find_many_kwargs["where"]
mock_get_daily_agg.assert_called_once_with(
prisma_client=mock_prisma_client,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2025-02-01",
end_date="2025-02-28",
model=None,
api_key=["tok-a", "tok-b"],
timezone_offset_minutes=480,
include_current_utc_day=False,
)
@pytest.mark.asyncio
async def test_search_user_daily_activity_keys_no_match_returns_empty_without_aggregating(monkeypatch):
from unittest.mock import AsyncMock, MagicMock
from litellm.constants import USAGE_TOP_API_KEYS_LIMIT
from litellm.proxy.management_endpoints.internal_user_endpoints import (
search_user_daily_activity_keys,
)
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_get_daily_agg = AsyncMock()
monkeypatch.setattr(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
mock_get_daily_agg,
)
admin_key_dict = UserAPIKeyAuth(
user_id="admin-user-001",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
result = await search_user_daily_activity_keys(
search="nothing-matches",
start_date="2025-02-01",
end_date="2025-02-28",
user_id=None,
timezone=None,
include_current_utc_day=False,
user_api_key_dict=admin_key_dict,
)
assert result.results == []
assert result.metadata.api_key_limit == USAGE_TOP_API_KEYS_LIMIT
assert result.metadata.total_api_keys == 0
mock_get_daily_agg.assert_not_called()
@pytest.mark.asyncio
async def test_search_user_daily_activity_keys_non_admin_scoped_to_caller(monkeypatch):
"""Same scoping contract as the aggregated route: a non-admin with no user_id
is scoped to their own rows, and any other user_id is a 403."""
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
from fastapi import HTTPException
from litellm.proxy.management_endpoints.internal_user_endpoints import (
search_user_daily_activity_keys,
)
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[SimpleNamespace(token="tok-a")])
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
non_admin_key_dict = UserAPIKeyAuth(
user_id="user-1",
user_role=LitellmUserRoles.INTERNAL_USER,
)
mock_response = MagicMock()
mock_get_daily_agg = AsyncMock(return_value=mock_response)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated",
mock_get_daily_agg,
)
result = await search_user_daily_activity_keys(
search="gamma",
start_date="2025-02-01",
end_date="2025-02-28",
user_id=None,
timezone=None,
include_current_utc_day=False,
user_api_key_dict=non_admin_key_dict,
)
assert result is mock_response
assert mock_get_daily_agg.call_args.kwargs["entity_id"] == "user-1"
find_many_kwargs = mock_prisma_client.db.litellm_verificationtoken.find_many.call_args.kwargs
assert find_many_kwargs["where"]["user_id"] == "user-1"
with pytest.raises(HTTPException) as exc_info:
await search_user_daily_activity_keys(
search="gamma",
start_date="2025-02-01",
end_date="2025-02-28",
user_id="user-2",
timezone=None,
include_current_utc_day=False,
user_api_key_dict=non_admin_key_dict,
)
assert exc_info.value.status_code == 403
assert "Non-admin users can only view their own spend data" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_delete_user_cleans_up_created_by_invitation_links(mocker):
"""

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