mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
commit
95e3caefb0
126 changed files with 12440 additions and 6850 deletions
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
10
litellm-rust/AGENTS.md
Normal 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
|
||||
|
|
@ -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
|
|
@ -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
|
||||
",
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
",
|
||||
);
|
||||
});
|
||||
}
|
||||
|
|
@ -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")));
|
||||
});
|
||||
}
|
||||
|
|
@ -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())]);
|
||||
}
|
||||
}
|
||||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
|
|
@ -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",
|
||||
);
|
||||
});
|
||||
}
|
||||
|
|
@ -4,7 +4,6 @@ version = "0.1.0"
|
|||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
autotests = false
|
||||
|
||||
[dependencies]
|
||||
litellm-secrets.workspace = true
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -51,6 +51,3 @@ pub fn chat_completions_decline_reason(
|
|||
.unsupported_reason(&messages, optional_params)
|
||||
.map(|reason| reason.0)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)
|
||||
}
|
||||
|
||||
#[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, .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)
|
||||
}
|
||||
|
||||
#[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, .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -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",
|
||||
})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -46,6 +46,3 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesR
|
|||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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");
|
||||
|
|
@ -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");
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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()),
|
||||
]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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(), ¶ms, &[])
|
||||
.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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
|
|
@ -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())
|
||||
})
|
||||
}
|
||||
|
|
@ -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,
|
||||
¶ms,
|
||||
&[],
|
||||
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"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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"})
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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()
|
||||
)]));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -218,7 +218,3 @@ fn anthropic_body(
|
|||
);
|
||||
Value::Object(body)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "tests.rs"]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -302,7 +302,3 @@ fn has_blank_text(message: &ChatMessage) -> bool {
|
|||
}),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "tests.rs"]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -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!({})
|
||||
|
|
@ -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")
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ##########################
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 ""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -73,6 +73,7 @@ LlmCapability = Literal[
|
|||
"long_context_1m",
|
||||
"mid_conversation_system",
|
||||
"multi_turn",
|
||||
"native_passthrough",
|
||||
"pdf_input",
|
||||
"prompt_cache_1h",
|
||||
"prompt_cache_5m",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
522
tests/integration/spend/test_team_daily_activity_export.py
Normal file
522
tests/integration/spend/test_team_daily_activity_export.py
Normal 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
|
||||
128
tests/integration/spend/test_team_daily_activity_key_search.py
Normal file
128
tests/integration/spend/test_team_daily_activity_key_search.py
Normal 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
|
||||
106
tests/integration/spend/test_team_member_spend_flush.py
Normal file
106
tests/integration/spend/test_team_member_spend_flush.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
387
tests/test_litellm/batches/test_batch_utils.py
Normal file
387
tests/test_litellm/batches/test_batch_utils.py
Normal 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)
|
||||
0
tests/test_litellm/files/__init__.py
Normal file
0
tests/test_litellm/files/__init__.py
Normal file
71
tests/test_litellm/files/test_main.py
Normal file
71
tests/test_litellm/files/test_main.py
Normal 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}"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
0
tests/test_litellm/llms/vertex_ai/files/__init__.py
Normal file
0
tests/test_litellm/llms/vertex_ai/files/__init__.py
Normal file
310
tests/test_litellm/llms/vertex_ai/files/test_transformation.py
Normal file
310
tests/test_litellm/llms/vertex_ai/files/test_transformation.py
Normal 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)
|
||||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue