From e9491d31b5881ae991ffa537fe74f3474379f8ce Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 01:09:36 +0000 Subject: [PATCH 01/10] refactor(rust): move credential inheritance and the SDK limits out of the legacy callback crate into a driver preflight (#43259) The legacy callback crate carried two rewrites that have nothing to do with the Logging contract: litellm_credential_name inheritance and the max_budget and num_retries_per_request checks. Any later callback host would need them unchanged, which is the smell the crate's AGENTS.md now names. They are now a Preflight the driver in litellm-host-python runs on the keyword view begin returned, before the host projects from it, supplied by python-bridge and passed through run_legacy_call. The call order is unchanged (setup, deployment hook, credentials, limits) and a rejection still fails the call as a host failure, so the failure callbacks run as before. The preflight rewrites the adapter's own copy in place, so no extra dict copy and no new lifecycle method Co-authored-by: Yujong Lee Co-authored-by: Claude Fable 5.1 --- litellm-rust/Cargo.lock | 1 + .../crates/callbacks-legacy-python/AGENTS.md | 24 +-- .../python_contract.json | 8 - .../callbacks-legacy-python/src/adapter.rs | 53 +---- .../callbacks-legacy-python/src/call.rs | 10 +- .../crates/callbacks-legacy-python/src/lib.rs | 17 +- .../callbacks-legacy-python/src/python.rs | 10 +- litellm-rust/crates/host-python/AGENTS.md | 2 +- .../crates/host-python/src/adapter.rs | 6 + litellm-rust/crates/host-python/src/driver.rs | 111 +++++++++- litellm-rust/crates/host-python/src/lib.rs | 3 +- litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../python-bridge/preflight_contract.json | 10 + litellm-rust/crates/python-bridge/src/lib.rs | 1 + .../src/preflight.rs} | 203 ++++++++++++++++-- .../python-bridge/src/routes/messages/mod.rs | 1 + .../python-bridge/src/routes/ocr/mod.rs | 1 + .../rust_bridge/callbacks_legacy_python.py | 32 --- litellm/rust_bridge/preflight.py | 45 ++++ .../test_callbacks_legacy_python.py | 30 +-- tests/unit/rust_bridge/test_preflight.py | 63 ++++++ 21 files changed, 460 insertions(+), 172 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/preflight_contract.json rename litellm-rust/crates/{callbacks-legacy-python/src/preparation.rs => python-bridge/src/preflight.rs} (57%) create mode 100644 litellm/rust_bridge/preflight.py create mode 100644 tests/unit/rust_bridge/test_preflight.py diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 3677d1d654f..ffcc6a5496b 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3311,6 +3311,7 @@ dependencies = [ "serde_json", "serde_with", "sha2 0.10.9", + "strum", "thiserror 2.0.19", "tokio", "tokio-tungstenite", diff --git a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md index 8b2e1c15f6e..de6e0c1b225 100644 --- a/litellm-rust/crates/callbacks-legacy-python/AGENTS.md +++ b/litellm-rust/crates/callbacks-legacy-python/AGENTS.md @@ -1,19 +1,19 @@ - Target invariants, not completion claims -- Keep this crate the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits) +- This crate is the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) + - Smell test: if a future callback host (`callbacks-v1-python`, WASM, in-process Rust) could share a piece of this crate, it does not belong here + - SDK request policy (credential inheritance, the budget and retry-count limits) is the driver's preflight, supplied by `python-bridge`; this crate only adopts the keyword view it produces - The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call -- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`) +- Every litellm Python internal Rust still borrows is a variant of `LegacyPython`, grouped by subsystem, with its signature pinned in `python_contract.json` - The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it - Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython` - - `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy -- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does - - Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's +- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the call rewrites it (setup, deployment hook, preflight) and the bound request object backing omitted keywords; routes hand it over through `run_legacy_call` and keep no copy +- `setup` reuses a `Logging` passed as `litellm_logging_obj` (the proxy and Router) and otherwise builds one through `function_setup`; which callbacks run is `Logging`'s decision, never this crate's - Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation - - Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view - - Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup` - - Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only - - A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view -- Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts - - Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch - - Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once + - Retain complete boundary arguments, opaque values, aliases, omitted/default distinctions and deliberate copies; preserve the deployment-hook kwargs view + - Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object, resolved through `litellm_host_python::lookup` + - Retain body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only +- Success and failure handlers receive the exact selected public response or exception + - A failure-handler error cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch + - Dispatch errors never replay provider work or trigger the opposite outcome; the proxy releases deferred success at most once - Delivery follows the registry, not the callable's type: direct, awaited, executor-submitted, logging-worker and deferred paths stay distinct - Traverse every retained Python edge; `close` is idempotent and restores the correlation context once diff --git a/litellm-rust/crates/callbacks-legacy-python/python_contract.json b/litellm-rust/crates/callbacks-legacy-python/python_contract.json index 8a7f3b98f47..9ed13ae5ed5 100644 --- a/litellm-rust/crates/callbacks-legacy-python/python_contract.json +++ b/litellm-rust/crates/callbacks-legacy-python/python_contract.json @@ -6,9 +6,6 @@ "start_time", "asynchronous" ], - "check_limits": [ - "kwargs" - ], "finalize": [ "response", "logger", @@ -76,11 +73,6 @@ ], "custom_pricing_fields": [], "is_internal_call": [], - "credential_list": [], - "warn_unknown_credential": [ - "name", - "loaded" - ], "before_deployment_call": [ "kwargs", "call_type" diff --git a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs index 75a635e9c63..718cc615f30 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/adapter.rs @@ -19,7 +19,7 @@ use serde_json::Value; use crate::{ DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger, deferred::{PendingLogging, PendingSuccess}, - finalize, is_internal_call, prepare, + finalize, is_internal_call, python::Streaming, setup, }; @@ -117,9 +117,13 @@ impl LegacyLogging { }) } + /// The keyword view the rest of the call reads: a copy, so the deployment hook's own + /// dict is left as the hook returned it, carrying the logger as `@client` injects it. + /// The driver's preflight rewrites this same dict before the host projects from it. fn prepare(&mut self, py: Python<'_>) -> PyResult { - let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind(); - self.call.set_kwargs(prepared); + let prepared = self.call.kwargs().bind(py).copy()?; + prepared.set_item("litellm_logging_obj", self.logger()?.object(py))?; + self.call.set_kwargs(prepared.unbind()); Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py))) } @@ -580,8 +584,6 @@ 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 ", ); }); @@ -616,8 +618,6 @@ kwargs = {'logger': logger, 'vendor_extension': opaque} &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 ", ); @@ -733,45 +733,6 @@ assert all(value is failure for name, value in logger.calls if name.endswith('_h ); }); } - - #[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::() - .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"))); - }); - } } #[cfg(test)] diff --git a/litellm-rust/crates/callbacks-legacy-python/src/call.rs b/litellm-rust/crates/callbacks-legacy-python/src/call.rs index 9b921070839..3fa638ac6d3 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/call.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/call.rs @@ -4,7 +4,7 @@ //! this crate holds them. use litellm_host::{machine::Machine, protocol::Protocol}; -use litellm_host_python::{ProtocolHost, lookup, run_call}; +use litellm_host_python::{Preflight, ProtocolHost, lookup, run_call}; use pyo3::{ gc::{PyTraverseError, PyVisit}, prelude::*, @@ -39,7 +39,8 @@ impl PublicCall { } /// The keyword view the legacy path currently reads: the caller's copy until - /// `function_setup`, then each rewrite (setup, deployment hook, prepare) in turn. + /// `function_setup`, then each rewrite (setup, deployment hook, the driver's preflight) + /// in turn. pub(crate) fn kwargs(&self) -> &Py { &self.kwargs } @@ -64,13 +65,15 @@ impl PublicCall { } /// Runs one native call under the legacy `Logging` contract: the protocol host projects from -/// the keyword view the contract prepares, and the contract observes the call. +/// the keyword view the contract prepares and `preflight` rewrites, and the contract +/// observes the call. pub fn run_legacy_call( py: Python<'_>, surface: LegacySurface, call: PublicCall, machine: M, host: H, + preflight: Preflight, asynchronous: bool, ) -> PyResult> where @@ -83,6 +86,7 @@ where machine, host, Box::new(LegacyLogging::new(py, surface, call, asynchronous)), + preflight, arguments, asynchronous, ) diff --git a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs index 69f72fbc177..869c534acf4 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/lib.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/lib.rs @@ -1,9 +1,10 @@ //! The legacy `@client` wrapper as the native call sees it: litellm's `Logging` object, the -//! sync and async callback registries it fans out to, the deployment hooks, the deferred -//! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name -//! inheritance, budget and retry-count limits). All of it sits behind one +//! sync and async callback registries it fans out to, the deployment hooks and the deferred +//! proxy release. All of it sits behind one //! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and -//! core never learn which Python object is on the other end. +//! core never learn which Python object is on the other end. The SDK's own request policy +//! (credential inheritance, the budget and retry limits) is the driver's preflight, not this +//! crate's. //! //! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`] //! is where those objects live, and [`run_legacy_call`] is how a route hands them over @@ -14,14 +15,12 @@ mod call; mod callbacks; mod deferred; mod logger; -mod preparation; mod python; 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 { @@ -77,7 +76,6 @@ FAKES = { 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, @@ -104,8 +102,6 @@ FAKES = { '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 @@ -162,9 +158,6 @@ class StubLogger: 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) diff --git a/litellm-rust/crates/callbacks-legacy-python/src/python.rs b/litellm-rust/crates/callbacks-legacy-python/src/python.rs index cb609d52878..47331f369f5 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/python.rs +++ b/litellm-rust/crates/callbacks-legacy-python/src/python.rs @@ -19,18 +19,12 @@ pub(crate) enum LegacyPython { Streaming(Streaming), } -/// The `@client` wrapper around the call: `function_setup`, limits, credentials, -/// response metadata and the correlation context. +/// The `@client` wrapper around the call: `function_setup`, response metadata and the +/// correlation context. #[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] pub(crate) enum Wrapper { #[strum(serialize = "setup")] Setup, - #[strum(serialize = "check_limits")] - CheckLimits, - #[strum(serialize = "credential_list")] - CredentialList, - #[strum(serialize = "warn_unknown_credential")] - WarnUnknownCredential, #[strum(serialize = "is_internal_call")] IsInternalCall, #[strum(serialize = "finalize")] diff --git a/litellm-rust/crates/host-python/AGENTS.md b/litellm-rust/crates/host-python/AGENTS.md index 7c1919f9f39..cadc55a35a7 100644 --- a/litellm-rust/crates/host-python/AGENTS.md +++ b/litellm-rust/crates/host-python/AGENTS.md @@ -2,7 +2,7 @@ - Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits - No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features - The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business - - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance) + - `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, rewritten in place by the route's `Preflight`, not the caller's dict; a protocol host that projects from it inherits the adapter's rewrites (for the legacy adapter: setup, deployment hooks) and the preflight's (credential inheritance) - A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is - A failing `classify` is raised with the native error's text as its `__context__`, never swallowed - Use standard PyO3 ownership and conversion APIs diff --git a/litellm-rust/crates/host-python/src/adapter.rs b/litellm-rust/crates/host-python/src/adapter.rs index 7f07475bc4c..83ed6416d6e 100644 --- a/litellm-rust/crates/host-python/src/adapter.rs +++ b/litellm-rust/crates/host-python/src/adapter.rs @@ -9,6 +9,12 @@ pub fn missing_state() -> PyErr { PyRuntimeError::new_err("missing native call state") } +/// The SDK's request policy, run by the driver on the keyword view `begin` returned and +/// before the protocol host projects from it. It rewrites that view in place, so the +/// lifecycle that returned it sees the rewrite too; a rejection fails the call as a host +/// failure, so the lifecycle still observes it. +pub type Preflight = fn(Python<'_>, &Bound<'_, PyDict>) -> PyResult<()>; + /// What an adapter step produced: either the value the driver asked for, or a Python /// awaitable the driver hands back to the caller's task before asking again. pub enum LifecycleStep { diff --git a/litellm-rust/crates/host-python/src/driver.rs b/litellm-rust/crates/host-python/src/driver.rs index 372af2843bd..50eae1e0225 100644 --- a/litellm-rust/crates/host-python/src/driver.rs +++ b/litellm-rust/crates/host-python/src/driver.rs @@ -14,7 +14,8 @@ use pyo3::types::PyDict; use tokio::sync::Mutex; use crate::adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle, + missing_state, }; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; use crate::handle::{Execution, ExecutionBody, ExecutionStep}; @@ -83,6 +84,7 @@ where { host: H, adapter: Box, + preflight: Preflight, machine: Option>>>, arguments: Option>, started_at: f64, @@ -95,12 +97,14 @@ where } /// Runs one native call for Python: synchronously, or as a coroutine that awaits every -/// host suspension inline in the caller's task. +/// host suspension inline in the caller's task. `preflight` runs once, on the keyword view +/// the adapter's `begin` returned, before the host projects from it. pub fn run_call( py: Python<'_>, machine: M, host: H, adapter: Box, + preflight: Preflight, arguments: Py, asynchronous: bool, ) -> PyResult> @@ -111,6 +115,7 @@ where let mut driver = PythonDriver { host, adapter, + preflight, machine: Some(Arc::new(Mutex::new(MachineState { machine, result: None, @@ -213,6 +218,9 @@ where match (expect, step) { (Expect::Started, LifecycleStep::Done) => self.begin(py), (Expect::Arguments, LifecycleStep::Arguments(arguments)) => { + if let Err(error) = (self.preflight)(py, arguments.bind(py)) { + return self.adapter_failed(py, error); + } self.arguments = Some(arguments); self.stage = Stage::Call; self.resume_machine(py, None) @@ -869,6 +877,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri host: SyntheticHost, script: AdapterScript, asynchronous: bool, + ) -> (PyResult>, Vec) { + run_preflighted(py, machine, host, script, no_preflight, asynchronous) + } + + fn no_preflight(_: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> { + Ok(()) + } + + fn run_preflighted( + py: Python<'_>, + machine: CallMachine, + host: SyntheticHost, + script: AdapterScript, + preflight: Preflight, + asynchronous: bool, ) -> (PyResult>, Vec) { let log = Log(host.log.0.clone()); let adapter = SyntheticAdapter { @@ -882,6 +905,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri machine, host, Box::new(adapter), + preflight, arguments.unbind(), asynchronous, ); @@ -1088,6 +1112,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri streaming_machine(), StreamingHost, Box::new(adapter), + no_preflight, PyDict::new(py).unbind(), asynchronous, ) @@ -1291,6 +1316,87 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri }); } + /// The rejection a preflight raised, kept so a test can check the caller receives that + /// exact object. A `Preflight` is a plain `fn`, so it cannot capture one itself. + static REJECTION: Mutex>> = Mutex::new(None); + + fn rejecting_preflight(py: Python<'_>, _: &Bound<'_, PyDict>) -> PyResult<()> { + let error = PyValueError::new_err("over budget"); + *REJECTION.lock().unwrap() = Some(error.value(py).clone().unbind()); + Err(error) + } + + fn inheriting_preflight(_: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> { + arguments.set_item("api_key", "inherited") + } + + #[test] + fn a_preflight_rejection_is_the_callers_error_and_the_machine_never_starts() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + for asynchronous in [false, true] { + let (result, log) = run_preflighted( + py, + success_machine(), + SyntheticHost { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + AdapterScript::Plain, + rejecting_preflight, + asynchronous, + ); + let error = result.unwrap_err(); + let raised = REJECTION.lock().unwrap().take().unwrap(); + assert!(error.value(py).is(&raised)); + assert_eq!( + log, + [ + "started", + "begin", + "failed:Host:over budget", + "adapter.close", + "host.close" + ] + ); + } + }); + } + + #[test] + fn the_host_projects_from_the_keyword_view_the_preflight_rewrote() { + let _guard = PYTHON_GLOBALS + .lock() + .unwrap_or_else(|error| error.into_inner()); + crate::initialize_python(); + Python::attach(|py| { + install_lifecycle_module(py); + for asynchronous in [false, true] { + let (result, _) = run_preflighted( + py, + success_machine(), + SyntheticHost { + log: Log::default(), + op: OpScript::Answer, + classifier_fails: false, + }, + AdapterScript::Plain, + inheriting_preflight, + asynchronous, + ); + assert_eq!( + result.unwrap().extract::(py).unwrap(), + "project:2|sign|rewritten" + ); + } + }); + } + #[test] fn the_adapters_finalized_response_is_what_the_call_returns_and_reports() { let _guard = PYTHON_GLOBALS @@ -1412,6 +1518,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri success_machine(), host, Box::new(adapter), + no_preflight, PyDict::new(py).unbind(), false, ) diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 7e17c4da51e..2f9e37fe968 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -15,7 +15,8 @@ mod handle; mod marshal; pub use adapter::{ - InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state, + InvokeError, LifecycleEvent, LifecycleStep, Preflight, ProtocolHost, PythonLifecycle, + missing_state, }; pub use argument::lookup; pub use callable::wrap_failure; diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index a02adfaa064..057cad2f42e 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -55,6 +55,7 @@ pyo3-async-runtimes.workspace = true reqwest.workspace = true redis = { version = "1.7.0", features = ["tls-rustls"] } serde_json.workspace = true +strum.workspace = true veil.workspace = true thiserror.workspace = true tokio = { workspace = true, features = ["rt", "sync"] } diff --git a/litellm-rust/crates/python-bridge/preflight_contract.json b/litellm-rust/crates/python-bridge/preflight_contract.json new file mode 100644 index 00000000000..343dea268cd --- /dev/null +++ b/litellm-rust/crates/python-bridge/preflight_contract.json @@ -0,0 +1,10 @@ +{ + "credential_list": [], + "warn_unknown_credential": [ + "name", + "loaded" + ], + "check_limits": [ + "kwargs" + ] +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 7c814f540a8..51e112fa1be 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -6,6 +6,7 @@ mod errors; mod http; mod logger; mod marshal; +mod preflight; mod python_settings; mod routes; mod secrets; diff --git a/litellm-rust/crates/callbacks-legacy-python/src/preparation.rs b/litellm-rust/crates/python-bridge/src/preflight.rs similarity index 57% rename from litellm-rust/crates/callbacks-legacy-python/src/preparation.rs rename to litellm-rust/crates/python-bridge/src/preflight.rs index aab654c9893..34813672c09 100644 --- a/litellm-rust/crates/callbacks-legacy-python/src/preparation.rs +++ b/litellm-rust/crates/python-bridge/src/preflight.rs @@ -1,9 +1,51 @@ +//! The SDK's request policy the driver runs on every route's keyword view before the host +//! projects from it: credential-name inheritance from `litellm.credential_list`, then the +//! budget and retry-count limits. It is the `@client` prologue after `function_setup` and the +//! deployment hook, and belongs to no callback contract. + use pyo3::{ prelude::*, types::{PyDict, PyList}, }; +use strum::{IntoStaticStr, VariantArray}; -use crate::python::Wrapper; +const MODULE: &str = "litellm.rust_bridge.preflight"; + +/// The litellm globals the preflight still reads through Python. `preflight_contract.json` +/// pins each function's parameters on both sides. +#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)] +pub(crate) enum PythonPreflight { + #[strum(serialize = "credential_list")] + CredentialList, + #[strum(serialize = "warn_unknown_credential")] + WarnUnknownCredential, + #[strum(serialize = "check_limits")] + CheckLimits, +} + +impl PythonPreflight { + fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult> + where + A: pyo3::call::PyCallArgs<'py>, + { + py.import(MODULE)?.getattr(<&str>::from(self))?.call1(args) + } +} + +#[cfg(test)] +pub(crate) const PYTHON_CONTRACT: &str = include_str!("../preflight_contract.json"); + +/// Rewrites `arguments` in place, in the order the Python wrapper runs: credentials first, +/// so the limits see the same view the provider request is built from. +pub(crate) fn sdk_preflight(py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<()> { + inherit_credentials(py, arguments, || { + Ok(PythonPreflight::CredentialList + .call(py, ())? + .cast_into::()?) + })?; + PythonPreflight::CheckLimits.call(py, (arguments,))?; + Ok(()) +} struct CredentialEntry<'py>(Bound<'py, PyAny>); @@ -17,22 +59,6 @@ impl<'py> CredentialEntry<'py> { } } -pub fn prepare<'py>( - py: Python<'py>, - kwargs: &Bound<'py, PyDict>, - logger: &crate::PythonLogger, -) -> PyResult> { - let arguments = kwargs.copy()?; - arguments.set_item("litellm_logging_obj", logger.object(py))?; - inherit_credentials(py, &arguments, || { - Ok(Wrapper::CredentialList - .call(py, ())? - .cast_into::()?) - })?; - Wrapper::CheckLimits.call(py, (&arguments,))?; - Ok(arguments) -} - fn inherit_credentials<'py>( py: Python<'py>, arguments: &Bound<'py, PyDict>, @@ -54,7 +80,7 @@ fn inherit_credentials<'py>( .map(|credential| CredentialEntry(credential).name()) .collect::>>()?; let Some(index) = names.iter().position(|name| *name == requested) else { - Wrapper::WarnUnknownCredential.call(py, (requested, names.len()))?; + PythonPreflight::WarnUnknownCredential.call(py, (requested, names.len()))?; return Ok(()); }; let selected = CredentialEntry(credentials.get_item(index)?); @@ -71,7 +97,42 @@ fn inherit_credentials<'py>( #[cfg(test)] mod tests { + use std::collections::BTreeSet; + use std::sync::Mutex; + use super::*; + use strum::VariantArray; + + /// Tests share one interpreter, and the stub module below is global state, so the + /// tests that install it run one at a time. + static PREFLIGHT_MODULE: Mutex<()> = Mutex::new(()); + + /// A fresh stand-in for `litellm.rust_bridge.preflight` that records every call, then + /// `script` run against it with the module bound as `preflight`. + fn preflight_stubs<'py>(py: Python<'py>, script: &std::ffi::CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + py.run( + c" +import sys +import types + +for name in ('litellm', 'litellm.rust_bridge'): + sys.modules.setdefault(name, types.ModuleType(name)) +preflight = types.ModuleType('litellm.rust_bridge.preflight') +preflight.warnings = [] +preflight.checked = [] +preflight.credential_list = lambda: [] +preflight.warn_unknown_credential = lambda name, loaded: preflight.warnings.append((name, loaded)) +preflight.check_limits = lambda kwargs: preflight.checked.append(kwargs) +sys.modules['litellm.rust_bridge.preflight'] = preflight +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + py.run(script, Some(&locals), Some(&locals)).unwrap(); + locals + } fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { let locals = PyDict::new(py); @@ -312,4 +373,110 @@ arguments = {'litellm_credential_name': 'ocr-test'} } }); } + + #[test] + fn every_borrowed_function_is_in_the_python_contract() { + Python::initialize(); + Python::attach(|py| { + let contract = litellm_host_python::json_loads(py, PYTHON_CONTRACT.as_bytes()).unwrap(); + let declared: BTreeSet = contract + .bind(py) + .cast::() + .unwrap() + .keys() + .extract() + .map(|names: Vec| names.into_iter().collect()) + .unwrap(); + let called: BTreeSet = PythonPreflight::VARIANTS + .iter() + .map(|&function| <&str>::from(function).to_owned()) + .collect(); + assert_eq!( + called.len(), + PythonPreflight::VARIANTS.len(), + "a function is borrowed twice" + ); + assert_eq!(called, declared); + }); + } + + #[test] + fn an_unknown_name_is_reported_with_the_loaded_count_and_leaves_the_arguments_alone() { + let _guard = PREFLIGHT_MODULE + .lock() + .unwrap_or_else(|error| error.into_inner()); + Python::initialize(); + Python::attach(|py| { + let locals = preflight_stubs( + py, + c" +class Credential: + credential_name = 'listed' + credential_values = {'api_key': 'listed-key'} +preflight.credential_list = lambda: [Credential(), Credential()] +arguments = {'litellm_credential_name': 'missing'} +", + ); + sdk_preflight(py, &argument_dict(&locals)).unwrap(); + py.run( + c" +assert arguments == {'litellm_credential_name': 'missing'}, arguments +assert preflight.warnings == [('missing', 2)], preflight.warnings +assert preflight.checked == [arguments] +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn limits_are_checked_on_the_arguments_after_credentials_are_inherited() { + let _guard = PREFLIGHT_MODULE + .lock() + .unwrap_or_else(|error| error.into_inner()); + Python::initialize(); + Python::attach(|py| { + let locals = preflight_stubs( + py, + c" +class Credential: + credential_name = 'ocr-test' + credential_values = {'api_key': 'inherited'} +preflight.credential_list = lambda: [Credential()] +rejection = RuntimeError('Max retries per request hit!') +def check_limits(arguments): + preflight.checked.append(dict(arguments)) + raise rejection +preflight.check_limits = check_limits +arguments = {'litellm_credential_name': 'ocr-test'} +", + ); + let error = sdk_preflight(py, &argument_dict(&locals)).unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("rejection").unwrap().unwrap()) + ); + py.run( + c" +assert preflight.checked == [{'litellm_credential_name': 'ocr-test', 'api_key': 'inherited'}], preflight.checked +assert arguments['api_key'] == 'inherited' +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + fn argument_dict<'py>(locals: &Bound<'py, PyDict>) -> Bound<'py, PyDict> { + locals + .get_item("arguments") + .unwrap() + .unwrap() + .cast_into::() + .unwrap() + } } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index dae8623979a..52cebb7c903 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -33,6 +33,7 @@ fn run_messages( PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(messages_machine(secrets)), MessagesPythonHost::new(request.unbind()), + crate::preflight::sdk_preflight, asynchronous, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index a4f2bf851d7..e00c57fad64 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -70,6 +70,7 @@ fn run_ocr( PublicCall::capture(&request, &args, &kwargs)?, crate::logger::LoggedMachine::new(ocr_machine(client)), OcrPythonHost::new(request.unbind()), + crate::preflight::sdk_preflight, asynchronous, ) } diff --git a/litellm/rust_bridge/callbacks_legacy_python.py b/litellm/rust_bridge/callbacks_legacy_python.py index 6bbf2ffed6b..e39324d3348 100644 --- a/litellm/rust_bridge/callbacks_legacy_python.py +++ b/litellm/rust_bridge/callbacks_legacy_python.py @@ -22,7 +22,6 @@ from typing import ( if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.types.utils import CredentialItem class MetadataUpdater(Protocol): @@ -72,21 +71,6 @@ def _claim_budget_reservation(call_setup: CallSetup, asynchronous: bool) -> Call return call_setup -def check_limits(kwargs: Mapping[str, object]) -> None: - from litellm import ( - BudgetExceededError, - _current_cost, # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor - max_budget, - num_retries_per_request, - ) - from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit - - if max_budget and _current_cost > max_budget: - raise BudgetExceededError(current_cost=_current_cost, max_budget=max_budget) - if max_retries_per_request_hit(kwargs, num_retries_per_request): - raise RuntimeError("Max retries per request hit!") - - def finalize( response: object, logger: Logging, @@ -299,22 +283,6 @@ def is_internal_call() -> bool: return internal.get() -def credential_list() -> list[CredentialItem]: - from litellm import credential_list as credentials - - return credentials - - -def warn_unknown_credential(name: str, loaded: int) -> None: - from litellm._logging import verbose_logger - - verbose_logger.warning( - "litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", - name, - loaded, - ) - - def before_deployment_call(kwargs: dict[str, object], call_type: str) -> Awaitable[object]: from litellm import utils diff --git a/litellm/rust_bridge/preflight.py b/litellm/rust_bridge/preflight.py new file mode 100644 index 00000000000..e030382bfdc --- /dev/null +++ b/litellm/rust_bridge/preflight.py @@ -0,0 +1,45 @@ +"""The SDK request policy the native driver runs before a route's host projects. + +These are the `@client` prologue steps after `function_setup` and the deployment hook: +credential-name inheritance and the budget and retry-count limits. Rust owns the +inheritance itself; it borrows only the globals below. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from litellm.types.utils import CredentialItem + + +def credential_list() -> list[CredentialItem]: + from litellm import credential_list as credentials + + return credentials + + +def warn_unknown_credential(name: str, loaded: int) -> None: + from litellm._logging import verbose_logger + + verbose_logger.warning( + "litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", + name, + loaded, + ) + + +def check_limits(kwargs: Mapping[str, object]) -> None: + from litellm import ( + BudgetExceededError, + _current_cost, # pyright: ignore[reportPrivateUsage] # shared SDK budget counter has no public accessor + max_budget, + num_retries_per_request, + ) + from litellm.litellm_core_utils.core_helpers import max_retries_per_request_hit + + if max_budget and _current_cost > max_budget: + raise BudgetExceededError(current_cost=_current_cost, max_budget=max_budget) + if max_retries_per_request_hit(kwargs, num_retries_per_request): + raise RuntimeError("Max retries per request hit!") diff --git a/tests/unit/rust_bridge/test_callbacks_legacy_python.py b/tests/unit/rust_bridge/test_callbacks_legacy_python.py index 05f2d13a079..7365679a28c 100644 --- a/tests/unit/rust_bridge/test_callbacks_legacy_python.py +++ b/tests/unit/rust_bridge/test_callbacks_legacy_python.py @@ -8,11 +8,10 @@ from typing import Final import pytest from pydantic import TypeAdapter -import litellm from litellm._internal_context import is_internal_call from litellm.litellm_core_utils.litellm_logging import Logging from litellm.rust_bridge import callbacks_legacy_python as legacy -from litellm.rust_bridge.callbacks_legacy_python import check_limits, failure_handler, setup +from litellm.rust_bridge.callbacks_legacy_python import failure_handler, setup _OCR_KWARGS: Final = MappingProxyType( { @@ -22,33 +21,6 @@ _OCR_KWARGS: Final = MappingProxyType( ) -@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) -@pytest.mark.parametrize( - "cap, request_retry_count, refused", - [(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)], - ids=[ - "cap-above-four-reached", - "cap-above-four-not-reached", - "first-attempt-passes-cap-of-zero", - "cap-of-zero-refuses-first-retry", - ], -) -def test_check_limits_reads_request_retry_count( - monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool -) -> None: - monkeypatch.setattr(litellm, "num_retries_per_request", cap) - monkeypatch.setattr(litellm, "max_budget", None) - kwargs: Final = { - "model": "mistral/mistral-ocr-latest", - metadata_key: {"request_retry_count": request_retry_count}, - } - if refused: - with pytest.raises(RuntimeError, match="Max retries per request hit!"): - check_limits(kwargs) - else: - check_limits(kwargs) - - def _supplied_logger() -> Logging: return Logging( model="mistral/mistral-ocr-latest", diff --git a/tests/unit/rust_bridge/test_preflight.py b/tests/unit/rust_bridge/test_preflight.py new file mode 100644 index 00000000000..a8b1a40a00b --- /dev/null +++ b/tests/unit/rust_bridge/test_preflight.py @@ -0,0 +1,63 @@ +import inspect +from collections.abc import Callable +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import pytest +from pydantic import TypeAdapter + +import litellm +from litellm.rust_bridge import preflight +from litellm.rust_bridge.preflight import check_limits + + +@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"]) +@pytest.mark.parametrize( + "cap, request_retry_count, refused", + [(5, 5, True), (5, 4, False), (0, 0, False), (0, 1, True)], + ids=[ + "cap-above-four-reached", + "cap-above-four-not-reached", + "first-attempt-passes-cap-of-zero", + "cap-of-zero-refuses-first-retry", + ], +) +def test_check_limits_reads_request_retry_count( + monkeypatch: pytest.MonkeyPatch, metadata_key: str, cap: int, request_retry_count: int, refused: bool +) -> None: + monkeypatch.setattr(litellm, "num_retries_per_request", cap) + monkeypatch.setattr(litellm, "max_budget", None) + kwargs: Final = { + "model": "mistral/mistral-ocr-latest", + metadata_key: {"request_retry_count": request_retry_count}, + } + if refused: + with pytest.raises(RuntimeError, match="Max retries per request hit!"): + check_limits(kwargs) + else: + check_limits(kwargs) + + +def test_check_limits_refuses_a_call_over_the_budget(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "num_retries_per_request", None) + monkeypatch.setattr(litellm, "max_budget", 1.0) + monkeypatch.setattr(litellm, "_current_cost", 1.5) + with pytest.raises(litellm.BudgetExceededError): + check_limits({"model": "mistral/mistral-ocr-latest"}) + + +CONTRACT_PATH: Final = Path(__file__).parents[3] / "litellm-rust/crates/python-bridge/preflight_contract.json" +_SHIMS: Final[MappingProxyType[str, Callable[..., object]]] = MappingProxyType( + { + "credential_list": preflight.credential_list, + "warn_unknown_credential": preflight.warn_unknown_credential, + "check_limits": preflight.check_limits, + } +) + + +def test_the_rust_contract_matches_the_shim_signatures() -> None: + contract: Final = TypeAdapter(dict[str, list[str]]).validate_json(CONTRACT_PATH.read_text()) + + assert contract == {name: list(inspect.signature(_SHIMS[name]).parameters) for name in contract} From 4adbc13d7996fe93fa8df8b5de3aa1f82a29e1dd Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 18:16:05 -0700 Subject: [PATCH 02/10] fix(router): hold Responses lifecycle events until output so a pre-output fallback announces one response (#43238) * fix(router): hold Responses lifecycle events until output so a pre-output fallback announces one response * fix(router): narrow the responses wrapper close guards to Exception and test the hold helpers directly * test(router): type the responses fallback test helpers * fix(router): replay the held lifecycle events when the fallback stream fails before its first event --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/responses/streaming_iterator.py | 4 +- litellm/router.py | 203 ++++++++++++++-------- tests/unit/test_router/test_router.py | 219 ++++++++++++++++++++++-- 3 files changed, 338 insertions(+), 88 deletions(-) diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 1ef39775bd3..9f537d24eaa 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -265,7 +265,7 @@ def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool: return not isinstance(status_code, int) or status_code >= 500 or status_code == 429 -_PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"}) +PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: Final = frozenset({"response.created", "response.in_progress", "response.queued"}) class BaseResponsesAPIStreamingIterator: @@ -885,7 +885,7 @@ class BaseResponsesAPIStreamingIterator: def _note_yielded_event(self, event: ResponsesAPIStreamingResponse) -> None: self._yielded_first_chunk = True - if event.type not in _PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: + if event.type not in PRE_OUTPUT_LIFECYCLE_EVENT_TYPES: self._output_started = True def _fallback_error(self, original: Exception) -> MidStreamFallbackError: diff --git a/litellm/router.py b/litellm/router.py index 023b99cd64e..6f416c416c0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -614,6 +614,17 @@ def _anthropic_stream_commits_now(chunk: object, has_generated_content: bool, bu return is_anthropic_content_delta_chunk(chunk) or buffered_chunk_count >= MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS +MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: Final = 200 + + +def _responses_stream_holds_event(item: object, held_event_count: int) -> bool: + from litellm.responses.streaming_iterator import PRE_OUTPUT_LIFECYCLE_EVENT_TYPES + + if held_event_count >= MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS: + return False + return getattr(item, "type", None) in PRE_OUTPUT_LIFECYCLE_EVENT_TYPES + + class FallbackAwareAnthropicMessagesStream: """ Bare async generators can't carry the `_hidden_params` attribute the @@ -3332,100 +3343,140 @@ class Router: await self._async_generator.aclose() async def stream_with_fallbacks(): - fallback_response = None + held_lifecycle_events: tuple[object, ...] = () # rebind-ok: flushed at first output, dropped on fallback try: async for item in source_iterator: + if _responses_stream_holds_event(item, len(held_lifecycle_events)): + held_lifecycle_events = (*held_lifecycle_events, item) + continue + for held_event in held_lifecycle_events: + yield held_event + held_lifecycle_events = () yield item + for held_event in held_lifecycle_events: + yield held_event except MidStreamFallbackError as e: - partial_usage: Final = Router._extract_partial_responses_usage(source_iterator) - try: - model_group: Final = cast(str, initial_kwargs.get("model")) - fallbacks: Final[list | None] = initial_kwargs.get("fallbacks", self.fallbacks) - context_window_fallbacks: Final[list | None] = initial_kwargs.get( - "context_window_fallbacks", self.context_window_fallbacks + async with contextlib.aclosing( + self._aresponses_fallback_attempt( + e, source_iterator, initial_kwargs, wrapper.adopt_fallback_headers, held_lifecycle_events ) - content_policy_fallbacks: Final[list | None] = initial_kwargs.get( - "content_policy_fallbacks", self.content_policy_fallbacks - ) - initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_responses_attempt - if e.is_pre_first_chunk or not e.generated_content: - # No content generated before the error — retry with the - # original input. Adding a continuation prompt would - # waste tokens and confuse the model. - pass - else: - initial_kwargs["input"] = Router._build_responses_continuation_input( - initial_kwargs.get("input"), - e.generated_content, - ) - # The Responses-API path stores observability metadata - # under "litellm_metadata" (not the default "metadata") — - # see _ageneric_api_call_with_fallbacks. Mirroring that - # here ensures model_group, model_group_alias, and trace - # ids land in the same key litellm.aresponses reads from. - self._update_kwargs_before_fallbacks( - model=model_group, - kwargs=initial_kwargs, - metadata_variable_name="litellm_metadata", - ) - # The content-policy dispatch branch matches on the trigger's own type, so a refusal's - # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. - fallback_trigger: Final[Exception] = ( - e.original_exception - if isinstance(e.original_exception, litellm.ContentPolicyViolationError) - else e - ) - fallback_response = await self.async_function_with_fallbacks_common_utils( - e=fallback_trigger, - disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), - fallbacks=fallbacks, - context_window_fallbacks=context_window_fallbacks, - content_policy_fallbacks=content_policy_fallbacks, - model_group=model_group, - args=(), - kwargs=initial_kwargs, - include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, - ) - - prepared_fallback_hidden_params = wrapper.adopt_fallback_headers(fallback_response) - if hasattr(fallback_response, "__aiter__"): - async for fallback_item in fallback_response: - Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) - if partial_usage is not None: - Router._combine_responses_fallback_usage(fallback_item, partial_usage) - yield fallback_item - else: - yield fallback_response - except Exception as fallback_error: - verbose_router_logger.error("Responses streaming fallback also failed: %s", fallback_error) - if ( - isinstance(fallback_error, MidStreamFallbackError) - and fallback_error.original_exception is not None - ): - raise fallback_error.original_exception from fallback_error - raise fallback_error + ) as fallback_stream: + async for fallback_item in fallback_stream: + yield fallback_item + except Exception: + for held_event in held_lifecycle_events: + yield held_event + raise finally: with anyio.CancelScope(shield=True): if hasattr(source_iterator, "aclose"): try: await source_iterator.aclose() - except BaseException as exc: + except Exception as exc: verbose_router_logger.debug( "stream_with_fallbacks(aresponses): error closing source: %s", exc, ) - if fallback_response is not None and hasattr(fallback_response, "aclose"): - try: - await fallback_response.aclose() - except BaseException as exc: - verbose_router_logger.debug( - "stream_with_fallbacks(aresponses): error closing fallback: %s", - exc, - ) wrapper: Final = FallbackResponsesStreamWrapper(stream_with_fallbacks()) return wrapper + async def _aresponses_fallback_attempt( + self, + e: "MidStreamFallbackError", + source_iterator: "BaseResponsesAPIStreamingIterator", + initial_kwargs: dict[str, Any], # mutable-ok: mutated in-place before re-entering the fallback chain + adopt_headers: Callable[[object], tuple[dict[str, object], dict[str, object]]], # mutable-ok: hidden params + held_lifecycle_events: tuple[object, ...], + ) -> AsyncGenerator[object, None]: + """ + Re-enters the Router's fallback chain for a mid-stream Responses API error and yields + whatever the fallback attempt produces. The lifecycle events the primary stream held + back reach the client only when no fallback lands, so the client sees exactly one + response announced, the one whose id completes. Split out of + _aresponses_streaming_iterator to keep each function's cyclomatic complexity within + the repo's C901 budget. + """ + from litellm.exceptions import MidStreamFallbackError + + partial_usage: Final = Router._extract_partial_responses_usage(source_iterator) + fallback_response = None # rebind-ok: pre-init so finally can close it if a fallback was actually attempted + fallback_yielded = False # rebind-ok: flipped on the first fallback item so a fallback that dies before its first event still replays the primary's held announcement + try: + model_group: Final = cast(str, initial_kwargs.get("model")) # cast-ok: model group + fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the common_utils list|None param + "fallbacks", self.fallbacks + ) + context_window_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below + "context_window_fallbacks", self.context_window_fallbacks + ) + content_policy_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below + "content_policy_fallbacks", self.content_policy_fallbacks + ) + initial_kwargs["original_function"] = ( # rebind-ok: the fallback chain re-enters on the same kwargs + self._ageneric_api_call_with_fallbacks_responses_attempt + ) + if e.generated_content and not e.is_pre_first_chunk: + initial_kwargs["input"] = Router._build_responses_continuation_input( # rebind-ok: fallback hop input + initial_kwargs.get("input"), + e.generated_content, + ) + # The Responses-API path stores observability metadata + # under "litellm_metadata" (not the default "metadata") — + # see _ageneric_api_call_with_fallbacks. Mirroring that + # here ensures model_group, model_group_alias, and trace + # ids land in the same key litellm.aresponses reads from. + self._update_kwargs_before_fallbacks( + model=model_group, + kwargs=initial_kwargs, + metadata_variable_name="litellm_metadata", + ) + # The content-policy dispatch branch matches on the trigger's own type, so a refusal's + # MidStreamFallbackError envelope is unwrapped here or the wrong fallback list is consulted. + fallback_trigger: Final[Exception] = ( + e.original_exception if isinstance(e.original_exception, litellm.ContentPolicyViolationError) else e + ) + fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success + e=fallback_trigger, + disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs), + fallbacks=fallbacks, + context_window_fallbacks=context_window_fallbacks, + content_policy_fallbacks=content_policy_fallbacks, + model_group=model_group, + args=(), + kwargs=initial_kwargs, + include_fallback_errors=initial_kwargs.get("include_fallback_errors", False) is True, + ) + prepared_fallback_hidden_params: Final = adopt_headers(fallback_response) + if hasattr(fallback_response, "__aiter__"): + async for fallback_item in fallback_response: + Router._apply_fallback_hidden_params_to_item(fallback_item, prepared_fallback_hidden_params) + if partial_usage is not None: + Router._combine_responses_fallback_usage(fallback_item, partial_usage) + fallback_yielded = True + yield fallback_item + else: + fallback_yielded = True # rebind-ok: see the pre-init above + yield fallback_response + except Exception as fallback_error: + verbose_router_logger.error("Responses streaming fallback also failed: %s", fallback_error) + if not fallback_yielded: + for held_event in held_lifecycle_events: + yield held_event + if isinstance(fallback_error, MidStreamFallbackError) and fallback_error.original_exception is not None: + raise fallback_error.original_exception from fallback_error + raise + finally: + if fallback_response is not None and hasattr(fallback_response, "aclose"): + with anyio.CancelScope(shield=True): + try: + await fallback_response.aclose() + except Exception as exc: + verbose_router_logger.debug( + "stream_with_fallbacks(aresponses): error closing fallback: %s", + exc, + ) + def _completion_streaming_iterator( self, model_response: CustomStreamWrapper, diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 3393c2f0d3c..e99cefb35dd 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -8,7 +8,7 @@ import os import sys import threading import warnings -from collections.abc import Awaitable, Callable, Mapping +from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable, Mapping from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final, Literal @@ -37,6 +37,7 @@ from litellm.models.access_group import LiteLLM_AccessGroupTable from litellm.proxy._types import LiteLLM_TeamTable, LitellmUserRoles, Member, ProxyException, UserAPIKeyAuth from litellm.router import ( MAX_BUFFERED_PRE_CONTENT_ANTHROPIC_CHUNKS, + MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS, FallbackAwareAnthropicMessagesStream, _anthropic_stream_commits_now, _anthropic_stream_error_is_gateway_verdict, @@ -46,6 +47,7 @@ from litellm.router import ( _anthropic_stream_should_decline_fallback, _anthropic_stream_should_drop_pre_content_ping, _is_retriable_anthropic_status, + _responses_stream_holds_event, ) from litellm.router_strategy import simple_shuffle from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit @@ -4213,6 +4215,7 @@ async def test_aresponses_streaming_iterator_fallback(): hidden_params={"model_id": "src-deployment-1"}, ) fallback_chunks = [ + MagicMock(type="response.created"), MagicMock(type="response.output_text.delta"), MagicMock(type="response.completed"), ] @@ -4235,7 +4238,7 @@ async def test_aresponses_streaming_iterator_fallback(): assert wrapped._hidden_params.get("model_id") == "src-deployment-1" collected = [c async for c in wrapped] - assert len(collected) == 3 # 1 primary chunk + 2 fallback chunks + assert collected == fallback_chunks call_kwargs = mock_fallback_utils.call_args.kwargs fbk = call_kwargs["kwargs"] # Bound methods compare equal when they share the same instance + __func__. @@ -4522,20 +4525,35 @@ def _make_native_responses_iterator(*, sse_payloads: tuple[dict[str, str], ...], _RESPONSES_LIFECYCLE_PAYLOADS: Final = ({"type": "response.created"}, {"type": "response.in_progress"}) +async def _events_until_error(stream: AsyncIterable[object]) -> AsyncIterator[object]: + try: + async for chunk in stream: + yield chunk + except Exception as error: + yield error + + @pytest.mark.asyncio async def test_aresponses_streaming_iterator_falls_back_on_transport_drop_before_output(): """A connection lost after response.created but before any output item is re-routed to the - fallback with the original input, the same as a provider error event would be.""" + fallback with the original input, the same as a provider error event would be, and the client + sees one response lifecycle: the fallback's, whose id the completed event carries.""" router: Final = _make_router_with_fallback() src: Final = _make_native_responses_iterator( sse_payloads=_RESPONSES_LIFECYCLE_PAYLOADS, trailing_error=httpx.ReadError("Response payload is not completed"), ) + fallback_chunks: Final = [ + MagicMock(type="response.created", response=MagicMock(id="resp_fallback")), + MagicMock(type="response.in_progress", response=MagicMock(id="resp_fallback")), + MagicMock(type="response.output_text.delta"), + MagicMock(type="response.completed", response=MagicMock(id="resp_fallback")), + ] with patch.object( router, "async_function_with_fallbacks_common_utils", - return_value=_AsyncList([MagicMock(type="response.completed")]), + return_value=_AsyncList(fallback_chunks), ) as mock_fallback_utils: wrapped: Final = await router._aresponses_streaming_iterator( response=src, @@ -4546,9 +4564,10 @@ async def test_aresponses_streaming_iterator_falls_back_on_transport_drop_before "original_generic_function": litellm.aresponses, }, ) - seen: Final = [chunk.type async for chunk in wrapped] + collected: Final = [chunk async for chunk in wrapped] - assert seen == ["response.created", "response.in_progress", "response.completed"] + assert collected == fallback_chunks + assert [chunk.response.id for chunk in collected if chunk.type == "response.created"] == ["resp_fallback"] assert isinstance(mock_fallback_utils.call_args.kwargs["e"], MidStreamFallbackError) assert mock_fallback_utils.call_args.kwargs["kwargs"]["input"] == "Hello" @@ -4576,17 +4595,197 @@ async def test_aresponses_streaming_iterator_surfaces_transport_drop_when_no_fal "original_generic_function": litellm.aresponses, }, ) - with pytest.raises(httpx.ReadError) as exc_info: - async for _ in wrapped: - pass + outcome: Final = [item async for item in _events_until_error(wrapped)] - assert exc_info.value is transport_error + assert [item.type for item in outcome[:-1]] == ["response.created", "response.in_progress"] + assert outcome[-1] is transport_error assert mock_fallback_utils.await_count == 1 trigger: Final = mock_fallback_utils.await_args.kwargs["e"] assert isinstance(trigger, MidStreamFallbackError) assert trigger.original_exception is transport_error +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_forwards_lifecycle_events_in_order_once_output_starts(): + router: Final = _make_router_with_fallback() + chunks: Final = [ + MagicMock(type="response.created"), + MagicMock(type="response.in_progress"), + MagicMock(type="response.output_text.delta"), + MagicMock(type="response.completed"), + ] + wrapped: Final = await router._aresponses_streaming_iterator( + response=_make_responses_iterator(chunks=chunks), + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + + assert [chunk async for chunk in wrapped] == chunks + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_flushes_held_lifecycle_events_when_the_stream_ends_without_output(): + router: Final = _make_router_with_fallback() + chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.in_progress")] + wrapped: Final = await router._aresponses_streaming_iterator( + response=_make_responses_iterator(chunks=chunks), + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + + assert [chunk async for chunk in wrapped] == chunks + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_forwards_held_lifecycle_events_before_a_non_fallback_error(): + router: Final = _make_router_with_fallback() + chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.in_progress")] + client_error: Final = litellm.BadRequestError(message="bad input", model="gpt-4", llm_provider="openai") + wrapped: Final = await router._aresponses_streaming_iterator( + response=_make_responses_iterator(chunks=chunks, error=client_error), + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + outcome: Final = [item async for item in _events_until_error(wrapped)] + + assert outcome[:-1] == chunks + assert outcome[-1] is client_error + + +@pytest.mark.asyncio +async def test_aresponses_streaming_iterator_commits_held_lifecycle_events_at_the_hold_cap(): + router: Final = _make_router_with_fallback() + chunks: Final = [MagicMock(type="response.in_progress") for _ in range(MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS + 1)] + src: Final = _make_responses_iterator( + chunks=chunks, + error=MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ), + ) + fallback_chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.completed")] + + with patch.object(router, "async_function_with_fallbacks_common_utils", return_value=_AsyncList(fallback_chunks)): + wrapped: Final = await router._aresponses_streaming_iterator( + response=src, + initial_kwargs={"model": "gpt-4", "stream": True, "input": "Hello"}, + ) + collected: Final = [chunk async for chunk in wrapped] + + assert collected == [*chunks, *fallback_chunks] + + +@pytest.mark.parametrize( + ("event_type", "held_event_count", "expected"), + [ + ("response.created", 0, True), + ("response.in_progress", 1, True), + ("response.queued", MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS - 1, True), + ("response.in_progress", MAX_HELD_PRE_OUTPUT_RESPONSES_EVENTS, False), + ("response.output_item.added", 0, False), + ("response.output_text.delta", 0, False), + ("response.completed", 0, False), + ], +) +def test_responses_stream_holds_event_holds_only_pre_output_lifecycle_events_under_the_cap( + event_type: str, held_event_count: int, expected: bool +): + assert _responses_stream_holds_event(MagicMock(type=event_type), held_event_count) is expected + + +@pytest.mark.asyncio +async def test_aresponses_fallback_attempt_drops_held_lifecycle_events_when_a_fallback_lands(): + router: Final = _make_router_with_fallback() + trigger: Final = MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ) + held: Final = (MagicMock(type="response.created"), MagicMock(type="response.in_progress")) + fallback_chunks: Final = [MagicMock(type="response.created"), MagicMock(type="response.completed")] + adopt_headers: Final = MagicMock(return_value=({}, {})) + + with patch.object( + router, "async_function_with_fallbacks_common_utils", return_value=_AsyncList(fallback_chunks) + ) as mock_fallback_utils: + collected: Final = [ + chunk + async for chunk in router._aresponses_fallback_attempt( + trigger, + _make_responses_iterator(), + {"model": "gpt-4", "stream": True, "input": "Hello"}, + adopt_headers, + held, + ) + ] + + assert collected == fallback_chunks + adopt_headers.assert_called_once() + assert mock_fallback_utils.await_args.kwargs["e"] is trigger + + +@pytest.mark.asyncio +async def test_aresponses_fallback_attempt_replays_held_lifecycle_events_when_the_fallback_dies_before_its_first_event(): + """A fallback stream that raises before yielding anything announced no response of its own, so the + primary's held created/in_progress pair is replayed ahead of the error and the client sees the + announcement the failure belongs to, the same as when no fallback was attempted at all.""" + router: Final = _make_router_with_fallback() + trigger: Final = MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ) + held: Final = (MagicMock(type="response.created"), MagicMock(type="response.in_progress")) + fallback_error: Final = RuntimeError("fallback closed before its first event") + adopt_headers: Final = MagicMock(return_value=({}, {})) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=_make_responses_iterator(error=fallback_error), + ): + outcome: Final = [ + item + async for item in _events_until_error( + router._aresponses_fallback_attempt( + trigger, + _make_responses_iterator(), + {"model": "gpt-4", "stream": True, "input": "Hello"}, + adopt_headers, + held, + ) + ) + ] + + assert outcome == [*held, fallback_error] + + +@pytest.mark.asyncio +async def test_aresponses_fallback_attempt_does_not_replay_held_lifecycle_events_once_the_fallback_announced_itself(): + """Once the fallback has yielded its own created event, a later failure must not replay the + primary's held pair on top of it, or the client would again see two announced response ids.""" + router: Final = _make_router_with_fallback() + trigger: Final = MidStreamFallbackError( + message="dropped before output", model="gpt-4", llm_provider="openai", is_pre_first_chunk=True + ) + held: Final = (MagicMock(type="response.created"), MagicMock(type="response.in_progress")) + fallback_created: Final = MagicMock(type="response.created") + fallback_error: Final = RuntimeError("fallback dropped after announcing itself") + adopt_headers: Final = MagicMock(return_value=({}, {})) + + with patch.object( + router, + "async_function_with_fallbacks_common_utils", + return_value=_make_responses_iterator(chunks=(fallback_created,), error=fallback_error), + ): + outcome: Final = [ + item + async for item in _events_until_error( + router._aresponses_fallback_attempt( + trigger, + _make_responses_iterator(), + {"model": "gpt-4", "stream": True, "input": "Hello"}, + adopt_headers, + held, + ) + ) + ] + + assert outcome == [fallback_created, fallback_error] + + @pytest.mark.asyncio async def test_aresponses_streaming_iterator_partial_content_injects_continuation(): """Mid-stream error: input is rewritten to include user prompt + From 1a581626301c8a09a2cd29578c5175802b1ebb5a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 18:31:15 -0700 Subject: [PATCH 03/10] refactor(http): hand out an owned Client and route all providers through the pool (#43245) * refactor(messages): take the provider client from the injected HTTP pool The messages route kept its own process-wide reqwest client, so it ignored ssl_verify, CA bundles, client certs, proxies and every other setting that litellm-http resolves. The machine now takes the HttpClientPool and the call's HttpClientConfig, as OCR does, and the bridge passes its shared pool. Co-Authored-By: Claude Opus 5.5 * refactor(http): hand out an owned Client and move chat, audio and OIDC onto the pool HttpClientPool now returns litellm_http::Client, a newtype only crates/http can build, so every provider client carries the resolved TLS, proxy and timeout settings. Chat completions and audio transcription drop their process-wide reqwest clients and take the pool and call config like messages; their 600s ceiling moves to the request. OidcResolver takes its client instead of building one, and the bridge hands it the pooled one. Co-Authored-By: Claude Opus 5.5 * refactor(secrets): build Google, Azure and CyberArk manager clients from the pool The native secret managers built bare reqwest clients, so they ignored the host's TLS and proxy settings. load_native_manager now takes the pool and the host config and hands each manager a pooled client. CyberArk's CYBERARK_SSL_VERIFY and CYBERARK_CLIENT_CERT/KEY become an override on the host config instead of a hand-built client. To express a certificate and key in separate files, HttpClientConfig::client_certificate is now a ClientIdentity that is either one PEM or a split pair. Co-Authored-By: Claude Opus 5.5 * chore(clippy): only crates/http may build a reqwest client Fence reqwest::Client, ClientBuilder and the TLS builder methods with disallowed-types and disallowed-methods so new code takes a litellm_http::Client from the pool. crates/http is exempt as the one place clients are built, and testkit as a dev-only installer. Tests move to litellm_http::Client::plain_for_test or a pooled client. Co-Authored-By: Claude Opus 5.5 * fix(secrets-cyberark): keep verifying certificates when the host disables it Python hands CyberArk its own ssl_verify, which wins over the global setting, so CYBERARK_SSL_VERIFY unset or true still verifies even when the host sets ssl_verify false. The pooled client copied the host's Disabled and would send the API key unverified; fall back to the built-in roots instead. Co-Authored-By: Claude Opus 5.5 * fix(python-bridge): treat a missing litellm package as no host HTTP settings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Claude Opus 5.5 Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 10 +++ litellm-rust/clippy.toml | 12 ++++ litellm-rust/crates/auth-aws/Cargo.toml | 1 + litellm-rust/crates/auth-aws/src/aws.rs | 2 +- .../crates/cache-azure-blob/Cargo.toml | 2 + .../crates/cache-azure-blob/src/cache.rs | 2 +- .../crates/cache-azure-blob/src/transport.rs | 2 +- .../cache-azure-blob/tests/transport.rs | 2 +- litellm-rust/crates/cache-gcs/Cargo.toml | 2 + litellm-rust/crates/cache-gcs/src/cache.rs | 2 +- litellm-rust/crates/cache-gcs/tests/cache.rs | 2 +- .../crates/cache-gcs/tests/support/mod.rs | 2 +- .../crates/cache-qdrant-semantic/Cargo.toml | 2 + .../cache-qdrant-semantic/src/embedder.rs | 2 +- .../cache-qdrant-semantic/tests/embedder.rs | 22 +++++-- litellm-rust/crates/cache-s3/Cargo.toml | 2 + litellm-rust/crates/cache-s3/src/cache.rs | 7 ++- litellm-rust/crates/cache-s3/src/transport.rs | 2 +- .../crates/cache-s3/tests/support/mod.rs | 7 ++- litellm-rust/crates/core/Cargo.toml | 1 + .../core/src/audio_transcription/client.rs | 13 ---- .../core/src/audio_transcription/handler.rs | 20 ++++-- .../core/src/audio_transcription/mod.rs | 13 ++-- .../core/src/chat_completions/client.rs | 14 ----- .../core/src/chat_completions/handler.rs | 20 ++++-- .../crates/core/src/chat_completions/mod.rs | 8 ++- litellm-rust/crates/core/src/constants.rs | 6 -- .../crates/core/src/messages/client.rs | 14 ----- .../crates/core/src/messages/error.rs | 2 + .../crates/core/src/messages/handler.rs | 12 ++-- litellm-rust/crates/core/src/messages/mod.rs | 19 ++++-- .../crates/core/src/messages/route.rs | 15 ++++- litellm-rust/crates/core/src/ocr/prepare.rs | 5 +- .../crates/core/tests/audio_transcription.rs | 20 +++--- .../crates/core/tests/chat_completions.rs | 21 ++++--- .../crates/core/tests/messages/host.rs | 2 +- .../crates/core/tests/messages/main.rs | 11 +++- .../crates/core/tests/messages/response.rs | 35 +++++++---- .../crates/core/tests/messages/stream.rs | 2 +- litellm-rust/crates/core/tests/ocr/main.rs | 7 +-- litellm-rust/crates/core/tests/ocr/mistral.rs | 13 ++-- litellm-rust/crates/core/tests/support/mod.rs | 13 +++- litellm-rust/crates/http/Cargo.toml | 2 + litellm-rust/crates/http/src/client.rs | 38 +++++++++++ litellm-rust/crates/http/src/config.rs | 12 +++- litellm-rust/crates/http/src/lib.rs | 10 ++- litellm-rust/crates/http/src/media.rs | 25 +++----- litellm-rust/crates/http/src/outbound.rs | 2 +- litellm-rust/crates/http/src/pool.rs | 14 +++-- litellm-rust/crates/http/src/tls.rs | 58 ++++++++++++----- litellm-rust/crates/llms/Cargo.toml | 1 + .../document_intelligence/transformation.rs | 4 +- .../crates/llms/src/base_llm/ocr/document.rs | 25 +++++--- .../crates/llms/src/base_llm/ocr/handler.rs | 27 ++++---- .../llms/src/reducto/ocr/transformation.rs | 5 +- litellm-rust/crates/llms/tests/ocr_handler.rs | 2 +- litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../python-bridge/src/cache/activation.rs | 2 +- .../crates/python-bridge/src/cache/config.rs | 2 +- .../crates/python-bridge/src/cache/handle.rs | 2 +- .../crates/python-bridge/src/cache/mod.rs | 10 --- .../crates/python-bridge/src/cache/native.rs | 8 +-- litellm-rust/crates/python-bridge/src/http.rs | 19 ++++-- .../python-bridge/src/python_settings.rs | 12 +++- .../src/routes/audio_transcription.rs | 34 ++++++---- .../src/routes/chat_completions.rs | 42 +++++++++---- .../python-bridge/src/routes/messages/mod.rs | 5 +- .../python-bridge/src/secrets/callback.rs | 2 +- .../crates/python-bridge/src/secrets/mod.rs | 7 ++- .../python-bridge/src/secrets/resolved.rs | 29 ++++++--- .../python-bridge/src/secrets/runtime.rs | 15 ++++- litellm-rust/crates/secrets-azure/Cargo.toml | 2 + .../crates/secrets-azure/src/key_vault.rs | 11 ++-- .../crates/secrets-azure/tests/key_vault.rs | 19 +++--- .../crates/secrets-azure/tests/live.rs | 2 +- .../crates/secrets-cyberark/Cargo.toml | 2 + .../crates/secrets-cyberark/src/error.rs | 2 + .../secrets-cyberark/src/secret_manager.rs | 7 ++- .../src/secret_manager/client.rs | 63 +++++++++++++++---- .../secrets-cyberark/tests/secret_manager.rs | 2 + .../tests/secret_manager/configuration.rs | 26 ++++---- .../tests/secret_manager/support.rs | 14 ++++- .../tests/secret_manager/writes.rs | 6 +- litellm-rust/crates/secrets-google/Cargo.toml | 2 + .../secrets-google/src/secret_manager.rs | 7 ++- .../secrets-google/tests/secret_manager.rs | 16 +++-- litellm-rust/crates/secrets/Cargo.toml | 2 + litellm-rust/crates/secrets/src/error.rs | 2 + litellm-rust/crates/secrets/src/native.rs | 33 ++++++---- litellm-rust/crates/secrets/src/oidc.rs | 34 +++++----- litellm-rust/crates/secrets/src/resolver.rs | 15 +---- litellm-rust/crates/secrets/src/source.rs | 5 +- litellm-rust/crates/secrets/tests/aws.rs | 6 +- litellm-rust/crates/secrets/tests/azure.rs | 6 +- .../secrets/tests/common_read_contract.rs | 8 +-- litellm-rust/crates/secrets/tests/cyberark.rs | 2 +- litellm-rust/crates/secrets/tests/google.rs | 4 +- .../crates/secrets/tests/hashicorp.rs | 4 +- litellm-rust/crates/secrets/tests/oidc.rs | 36 ++++++----- .../crates/secrets/tests/resolution.rs | 14 ++--- litellm-rust/crates/secrets/tests/source.rs | 4 +- litellm-rust/crates/testkit/src/lib.rs | 6 ++ 102 files changed, 745 insertions(+), 403 deletions(-) delete mode 100644 litellm-rust/crates/core/src/audio_transcription/client.rs delete mode 100644 litellm-rust/crates/core/src/chat_completions/client.rs delete mode 100644 litellm-rust/crates/core/src/messages/client.rs create mode 100644 litellm-rust/crates/http/src/client.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ffcc6a5496b..5bb65b6b05f 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2911,6 +2911,7 @@ dependencies = [ "litellm-cache", "litellm-cache-response", "litellm-cache-testing", + "litellm-http", "reqwest 0.12.28", "rstest", "serde_json", @@ -2944,6 +2945,7 @@ dependencies = [ "litellm-auth-types", "litellm-cache", "litellm-cache-testing", + "litellm-http", "percent-encoding", "reqwest 0.12.28", "rstest", @@ -2970,6 +2972,7 @@ dependencies = [ "futures-util", "litellm-cache", "litellm-cache-testing", + "litellm-http", "qdrant-client", "reqwest 0.12.28", "rstest", @@ -3043,6 +3046,7 @@ dependencies = [ "litellm-auth-aws", "litellm-cache", "litellm-cache-testing", + "litellm-http", "reqwest 0.12.28", "rstest", "serde_json", @@ -3207,11 +3211,13 @@ dependencies = [ "http 1.4.2", "hyper-util", "litellm-core-utils", + "rcgen", "reqwest 0.12.28", "rstest", "rustls 0.23.42", "serde", "serde_json", + "tempfile", "thiserror 2.0.19", "tokio", "veil", @@ -3345,6 +3351,7 @@ dependencies = [ "google-cloud-auth", "google-cloud-kms-v1", "litellm-core-utils", + "litellm-http", "litellm-python-compat", "litellm-secrets-aws", "litellm-secrets-azure", @@ -3392,6 +3399,7 @@ dependencies = [ "litellm-auth-azure", "litellm-auth-types", "litellm-core-utils", + "litellm-http", "litellm-secrets-types", "percent-encoding", "reqwest 0.12.28", @@ -3411,6 +3419,7 @@ version = "0.1.0" dependencies = [ "base64 0.22.1", "litellm-core-utils", + "litellm-http", "litellm-secrets-types", "litellm-tracing", "moka", @@ -3439,6 +3448,7 @@ dependencies = [ "litellm-auth-gcp", "litellm-auth-types", "litellm-core-utils", + "litellm-http", "litellm-secrets-types", "moka", "percent-encoding", diff --git a/litellm-rust/clippy.toml b/litellm-rust/clippy.toml index f7e3293069b..0e2ff770d27 100644 --- a/litellm-rust/clippy.toml +++ b/litellm-rust/clippy.toml @@ -7,4 +7,16 @@ disallowed-methods = [ { path = "pyo3_async_runtimes::tokio::local_future_into_py", reason = "use litellm_host_python::run_async / run_async_value" }, { path = "pyo3_async_runtimes::tokio::run", reason = "use litellm_host_python::run_sync / run_sync_value" }, { path = "pyo3_async_runtimes::tokio::run_until_complete", reason = "use litellm_host_python::run_sync / run_sync_value" }, + { path = "reqwest::Client::new", reason = "take litellm_http::Client from HttpClientPool" }, + { path = "reqwest::Client::builder", reason = "HttpClientConfig owns client construction" }, + { path = "reqwest::ClientBuilder::danger_accept_invalid_certs", reason = "set HttpClientConfig::verify instead" }, + { path = "reqwest::ClientBuilder::identity", reason = "set HttpClientConfig::client_certificate instead" }, + { path = "reqwest::ClientBuilder::use_preconfigured_tls", reason = "HttpClientConfig owns the TLS configuration" }, +] + +# Every outbound client comes from litellm_http::HttpClientPool so it honors the host's TLS, +# proxy and timeout settings. Only crates/http builds one. +disallowed-types = [ + { path = "reqwest::Client", reason = "take litellm_http::Client from HttpClientPool; only crates/http builds one" }, + { path = "reqwest::ClientBuilder", reason = "HttpClientConfig owns client construction" }, ] diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml index 1a35af48574..9592f278d94 100644 --- a/litellm-rust/crates/auth-aws/Cargo.toml +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -22,5 +22,6 @@ aws-types = "1.4.0" aws-smithy-runtime-api = "1.13.0" [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } reqwest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/auth-aws/src/aws.rs b/litellm-rust/crates/auth-aws/src/aws.rs index cb9195ffeb6..409ff78867f 100644 --- a/litellm-rust/crates/auth-aws/src/aws.rs +++ b/litellm-rust/crates/auth-aws/src/aws.rs @@ -962,7 +962,7 @@ mod tests { &no_env, ) .await?; - let client = reqwest::Client::new(); + let client = litellm_http::Client::plain_for_test(); let mut failures = Vec::new(); for region in ["us-west-2", "us-east-1"] { diff --git a/litellm-rust/crates/cache-azure-blob/Cargo.toml b/litellm-rust/crates/cache-azure-blob/Cargo.toml index baa1b0f5482..5bdfa16ef53 100644 --- a/litellm-rust/crates/cache-azure-blob/Cargo.toml +++ b/litellm-rust/crates/cache-azure-blob/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true litellm-auth-azure.workspace = true litellm-auth-types.workspace = true litellm-cache.workspace = true @@ -19,6 +20,7 @@ tokio.workspace = true url.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-response.workspace = true litellm-cache-testing.workspace = true rstest.workspace = true diff --git a/litellm-rust/crates/cache-azure-blob/src/cache.rs b/litellm-rust/crates/cache-azure-blob/src/cache.rs index 489b08d485e..c5c1fdd8ab9 100644 --- a/litellm-rust/crates/cache-azure-blob/src/cache.rs +++ b/litellm-rust/crates/cache-azure-blob/src/cache.rs @@ -31,7 +31,7 @@ impl AzureBlobCache { pub async fn connect( account_url: &str, container: &str, - http: reqwest::Client, + http: litellm_http::Client, codec: C, runtime: Handle, ) -> Result { diff --git a/litellm-rust/crates/cache-azure-blob/src/transport.rs b/litellm-rust/crates/cache-azure-blob/src/transport.rs index ed038b8d69d..3914b92365c 100644 --- a/litellm-rust/crates/cache-azure-blob/src/transport.rs +++ b/litellm-rust/crates/cache-azure-blob/src/transport.rs @@ -8,7 +8,7 @@ use azure_core::{ use futures_util::TryStreamExt; #[derive(Debug)] -pub struct ReqwestTransport(pub reqwest::Client); +pub struct ReqwestTransport(pub litellm_http::Client); #[async_trait::async_trait] impl HttpClient for ReqwestTransport { diff --git a/litellm-rust/crates/cache-azure-blob/tests/transport.rs b/litellm-rust/crates/cache-azure-blob/tests/transport.rs index cd1e10aa3d8..c8b14cd8543 100644 --- a/litellm-rust/crates/cache-azure-blob/tests/transport.rs +++ b/litellm-rust/crates/cache-azure-blob/tests/transport.rs @@ -31,7 +31,7 @@ async fn connect(server: &MockServer) -> AzureBlobCache> { None, ClientOptions { transport: Some(Transport::new(Arc::new(ReqwestTransport( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), )))), ..ClientOptions::default() }, diff --git a/litellm-rust/crates/cache-gcs/Cargo.toml b/litellm-rust/crates/cache-gcs/Cargo.toml index da0acf554f9..1a06683e615 100644 --- a/litellm-rust/crates/cache-gcs/Cargo.toml +++ b/litellm-rust/crates/cache-gcs/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true futures-util.workspace = true litellm-auth-gcp.workspace = true litellm-auth-types.workspace = true @@ -15,6 +16,7 @@ reqwest.workspace = true tokio.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-testing.workspace = true rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/cache-gcs/src/cache.rs b/litellm-rust/crates/cache-gcs/src/cache.rs index a8a7fbc9a7b..bad81573cb4 100644 --- a/litellm-rust/crates/cache-gcs/src/cache.rs +++ b/litellm-rust/crates/cache-gcs/src/cache.rs @@ -5,8 +5,8 @@ use litellm_cache::{ BaseCache, BatchCache, BatchEntry, CacheCodec, DisconnectCache, Error, ExactCacheContext, FlushCache, }; +use litellm_http::Client; use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, percent_encode}; -use reqwest::Client; use crate::{GcpTokenSource, TokenSource}; diff --git a/litellm-rust/crates/cache-gcs/tests/cache.rs b/litellm-rust/crates/cache-gcs/tests/cache.rs index cdce6a00bdd..12bb5344570 100644 --- a/litellm-rust/crates/cache-gcs/tests/cache.rs +++ b/litellm-rust/crates/cache-gcs/tests/cache.rs @@ -96,7 +96,7 @@ async fn cache_exposes_its_configuration(#[future(awt)] server: MockServer) { path_service_account: Some("/secrets/sa.json".into()), ..support::config(&server, Some("folder")) }, - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), litellm_cache::JsonCodec::::new(), ); assert_eq!(cache.bucket_name(), "bucket"); diff --git a/litellm-rust/crates/cache-gcs/tests/support/mod.rs b/litellm-rust/crates/cache-gcs/tests/support/mod.rs index 6097f0ee1bd..beb9aa39d9c 100644 --- a/litellm-rust/crates/cache-gcs/tests/support/mod.rs +++ b/litellm-rust/crates/cache-gcs/tests/support/mod.rs @@ -29,7 +29,7 @@ pub fn cache_with_token( ) -> JsonGcsCache { GcsCache::with_token_source( config(server, gcs_path), - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), JsonCodec::new(), token, ) diff --git a/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml b/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml index 950c2db7491..a44bef0a731 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml +++ b/litellm-rust/crates/cache-qdrant-semantic/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true futures-util.workspace = true litellm-cache.workspace = true qdrant-client = { workspace = true, features = ["serde"] } @@ -17,6 +18,7 @@ tokio.workspace = true uuid.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } futures-executor = "0.3" litellm-cache-testing.workspace = true rstest.workspace = true diff --git a/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs b/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs index 340393600f2..b7fbcd9b02d 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/src/embedder.rs @@ -1,7 +1,7 @@ use std::time::Duration; use litellm_cache::{Error, semantic::Embedder}; -use reqwest::Client; +use litellm_http::Client; use serde_json::Value; pub struct OpenAiEmbedder { diff --git a/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs b/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs index de0fab0a66f..24e6e5eba3e 100644 --- a/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs +++ b/litellm-rust/crates/cache-qdrant-semantic/tests/embedder.rs @@ -5,6 +5,10 @@ use std::{ use litellm_cache::{Error, semantic::Embedder}; use litellm_cache_qdrant_semantic::{OpenAiEmbedder, OpenAiEmbedderConfig}; +use litellm_http::{ + ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, Resolution, + media::PublicDnsResolver, +}; use rstest::rstest; use serde_json::{Value, json}; use tokio::{ @@ -104,7 +108,7 @@ fn config(base: String, timeout: Option) -> OpenAiEmbedderConfig { async fn posts_embeddings_request_and_parses_vector() { let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await; let embedder = OpenAiEmbedder::new( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), config( format!("{}/", server.base_url()), Some(Duration::from_secs(1)), @@ -156,14 +160,17 @@ async fn status_timeout_and_body_errors_are_unavailable( ) { let server = TestHttpServer::response_after(status, body, Duration::from_millis(delay_ms)).await; - let embedder = OpenAiEmbedder::new(reqwest::Client::new(), config(server.base_url(), timeout)); + let embedder = OpenAiEmbedder::new( + litellm_http::Client::plain_for_test(), + config(server.base_url(), timeout), + ); assert_eq!(embedder.async_embed("hello", None).await, expected); } #[rstest] fn sync_embedding_is_unsupported() { let embedder = OpenAiEmbedder::new( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), config("http://127.0.0.1:9".to_owned(), None), ); assert_eq!( @@ -176,9 +183,12 @@ fn sync_embedding_is_unsupported() { #[tokio::test] async fn uses_the_injected_client() { let server = TestHttpServer::response("200 OK", r#"{"data":[{"embedding":[0.1,0.2]}]}"#).await; - let client = reqwest::Client::builder() - .user_agent("litellm-embedder-test") - .build() + let config_with_agent = HttpClientConfig { + user_agent: Some("litellm-embedder-test".into()), + ..Resolution::from(&HttpSettings::default()).config + }; + let client = HttpClientPool::new(Arc::new(PublicDnsResolver)) + .client(&config_with_agent, ClientVariant::Provider) .unwrap(); let embedder = OpenAiEmbedder::new(client, config(server.base_url(), None)); assert_eq!( diff --git a/litellm-rust/crates/cache-s3/Cargo.toml b/litellm-rust/crates/cache-s3/Cargo.toml index c8150180e7c..680f2da8215 100644 --- a/litellm-rust/crates/cache-s3/Cargo.toml +++ b/litellm-rust/crates/cache-s3/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true litellm-cache.workspace = true litellm-auth-aws.workspace = true aws-sdk-s3 = { version = "1.146.1", default-features = false, features = ["rustls", "rt-tokio"] } @@ -19,6 +20,7 @@ reqwest.workspace = true tokio.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-cache-testing.workspace = true rstest.workspace = true wiremock = "0.6.5" diff --git a/litellm-rust/crates/cache-s3/src/cache.rs b/litellm-rust/crates/cache-s3/src/cache.rs index 91cd5e8ef54..ced948e80c6 100644 --- a/litellm-rust/crates/cache-s3/src/cache.rs +++ b/litellm-rust/crates/cache-s3/src/cache.rs @@ -42,7 +42,12 @@ pub struct S3Cache { } impl S3Cache { - pub fn new(config: S3CacheConfig, http: reqwest::Client, codec: C, runtime: Handle) -> Self { + pub fn new( + config: S3CacheConfig, + http: litellm_http::Client, + codec: C, + runtime: Handle, + ) -> Self { let endpoint_url: Option = config.endpoint.map(|endpoint| endpoint.url); let base = aws_sdk_s3::Config::builder() .behavior_version(BehaviorVersion::latest()) diff --git a/litellm-rust/crates/cache-s3/src/transport.rs b/litellm-rust/crates/cache-s3/src/transport.rs index 3e5ce578c31..eabc54ac9e2 100644 --- a/litellm-rust/crates/cache-s3/src/transport.rs +++ b/litellm-rust/crates/cache-s3/src/transport.rs @@ -9,7 +9,7 @@ use aws_smithy_runtime_api::client::{ use aws_smithy_types::body::SdkBody; #[derive(Clone, Debug)] -pub(crate) struct ReqwestHttpClient(pub(crate) reqwest::Client); +pub(crate) struct ReqwestHttpClient(pub(crate) litellm_http::Client); impl HttpClient for ReqwestHttpClient { fn http_connector( diff --git a/litellm-rust/crates/cache-s3/tests/support/mod.rs b/litellm-rust/crates/cache-s3/tests/support/mod.rs index 046b042c66a..b628c1df431 100644 --- a/litellm-rust/crates/cache-s3/tests/support/mod.rs +++ b/litellm-rust/crates/cache-s3/tests/support/mod.rs @@ -32,7 +32,12 @@ pub fn config(endpoint: &str) -> S3CacheConfig { } pub fn cache_with(config: S3CacheConfig, runtime: Handle) -> JsonS3Cache { - S3Cache::new(config, reqwest::Client::new(), JsonCodec::new(), runtime) + S3Cache::new( + config, + litellm_http::Client::plain_for_test(), + JsonCodec::new(), + runtime, + ) } pub fn cache(endpoint: &str) -> JsonS3Cache { diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 12410c187e2..8700c8df308 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -36,6 +36,7 @@ url.workspace = true veil.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-auth-gcp.workspace = true litellm-llms = { workspace = true, features = ["test-support"] } rstest.workspace = true diff --git a/litellm-rust/crates/core/src/audio_transcription/client.rs b/litellm-rust/crates/core/src/audio_transcription/client.rs deleted file mode 100644 index 3cf131839b8..00000000000 --- a/litellm-rust/crates/core/src/audio_transcription/client.rs +++ /dev/null @@ -1,13 +0,0 @@ -use std::{sync::OnceLock, time::Duration}; - -use crate::constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index a1862f341a5..30900bc14c6 100644 --- a/litellm-rust/crates/core/src/audio_transcription/handler.rs +++ b/litellm-rust/crates/core/src/audio_transcription/handler.rs @@ -1,10 +1,16 @@ -use litellm_http::request::truncate_error_body; +use std::time::Duration; + +use litellm_http::{Client, request::truncate_error_body}; use serde_json::Value; -use super::{Error, client::http_client}; -use crate::audio_transcription::types::ProviderAudioTranscriptionRequest; +use super::Error; +use crate::{ + audio_transcription::types::ProviderAudioTranscriptionRequest, + constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS, +}; pub async fn execute_audio_transcription_provider_call( + http: &Client, request: ProviderAudioTranscriptionRequest, ) -> Result { let response = crate::outbound::outbound_request::( @@ -12,11 +18,15 @@ pub async fn execute_audio_transcription_provider_call( request.url.clone(), request.upstream_headers.clone(), &request.body, - request.timeout, + Some( + request + .timeout + .unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)), + ), &request.optional_params, ) .await? - .send(http_client()) + .send(http) .await .map_err(|error| { Error::Transport(litellm_http::transport::Error::Network(error.to_string())) diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 801fd5e9673..dc75326d5c3 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -1,16 +1,21 @@ mod error; pub mod types; pub use error::Error; -mod client; mod handler; mod prepare; pub use handler::execute_audio_transcription_provider_call; +use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool}; pub use prepare::prepare_audio_transcription_provider_call; use serde_json::Value; use crate::audio_transcription::types::AudioTranscriptionRequest; -pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result { - execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?) - .await +pub async fn audio_transcription( + pool: &HttpClientPool, + config: &HttpClientConfig, + request: AudioTranscriptionRequest<'_>, +) -> Result { + let request = prepare_audio_transcription_provider_call(request)?; + let http = pool.client(config, ClientVariant::Provider)?; + execute_audio_transcription_provider_call(&http, request).await } diff --git a/litellm-rust/crates/core/src/chat_completions/client.rs b/litellm-rust/crates/core/src/chat_completions/client.rs deleted file mode 100644 index d8ad6c49b7b..00000000000 --- a/litellm-rust/crates/core/src/chat_completions/client.rs +++ /dev/null @@ -1,14 +0,0 @@ -use std::{sync::OnceLock, time::Duration}; - -use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS}; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)) - .connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 2391ab83a60..f3404fcaa8a 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -1,20 +1,24 @@ -use litellm_http::{outbound::OutboundRequest, request::truncate_error_body}; +use std::time::Duration; + +use litellm_http::{Client, outbound::OutboundRequest, request::truncate_error_body}; use litellm_llms::base_llm::chat::transformation::ProviderChatResponseData; use litellm_types::utils::ChatCompletionsResponse; use serde_json::Value; -use super::{Error, client::http_client, prepare::prepare_provider_request}; -use crate::chat_completions::types::{ - ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest, +use super::{Error, prepare::prepare_provider_request}; +use crate::{ + chat_completions::types::{ProviderChatCompletionsRequest, ResolvedChatCompletionsRequest}, + constants::CHAT_COMPLETIONS_TIMEOUT_SECS, }; pub(super) async fn execute_chat_completions_provider_call( + http: &Client, request: ResolvedChatCompletionsRequest<'_>, ) -> Result { let request = prepare_provider_request(request)?; let outbound = outbound_request(&request).await?; - let response = outbound.send(http_client()).await.map_err(|err| { + let response = outbound.send(http).await.map_err(|err| { // Failing to establish the connection means the request never went out, // so the host can still serve it. Everything else here, a timeout // above all, may have reached the provider and been answered. @@ -72,7 +76,11 @@ pub(super) async fn outbound_request( request.url.clone(), request.upstream_headers.clone(), &request.body, - request.timeout, + Some( + request + .timeout + .unwrap_or(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)), + ), &request.optional_params, ) .await diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 224c9d8cfed..be22aea5669 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -9,11 +9,11 @@ mod error; pub mod types; pub use error::Error; -mod client; mod common_utils; pub(crate) mod handler; mod prepare; use handler::execute_chat_completions_provider_call; +use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool}; use litellm_types::utils::ChatCompletionsResponse; use prepare::{parse_messages, resolve_provider_config, resolve_request}; use serde_json::{Map, Value}; @@ -21,9 +21,13 @@ use serde_json::{Map, Value}; use crate::chat_completions::types::ChatCompletionsRequest; pub async fn chat_completions( + pool: &HttpClientPool, + config: &HttpClientConfig, request: ChatCompletionsRequest<'_>, ) -> Result { - execute_chat_completions_provider_call(resolve_request(request)?).await + let request = resolve_request(request)?; + let http = pool.client(config, ClientVariant::Provider)?; + execute_chat_completions_provider_call(&http, request).await } /// Whether the core would accept this request, without resolving credentials or diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 3d740e39677..455c3258799 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -5,9 +5,6 @@ pub const OPENAI_DEFAULT_API_BASE: &str = "https://api.openai.com"; /// timeout from the caller still overrides this on the request builder. pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600; -/// Connect timeout for Anthropic Messages provider calls, in seconds. -pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; - /// Provider name used for Anthropic Messages when a deployment's provider model /// does not carry an explicit provider prefix. pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; @@ -16,9 +13,6 @@ pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; /// seconds. Mirrors the Python chat completions default. pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600; -/// Connect timeout for chat completions provider calls, in seconds. -pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10; - pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600; /// `object` field every non-streaming chat completion response carries. diff --git a/litellm-rust/crates/core/src/messages/client.rs b/litellm-rust/crates/core/src/messages/client.rs deleted file mode 100644 index ca70b1b03eb..00000000000 --- a/litellm-rust/crates/core/src/messages/client.rs +++ /dev/null @@ -1,14 +0,0 @@ -use std::{sync::OnceLock, time::Duration}; - -use crate::constants::{MESSAGES_CONNECT_TIMEOUT_SECS, MESSAGES_TIMEOUT_SECS}; - -pub(super) fn http_client() -> &'static reqwest::Client { - static CLIENT: OnceLock = OnceLock::new(); - CLIENT.get_or_init(|| { - reqwest::Client::builder() - .timeout(Duration::from_secs(MESSAGES_TIMEOUT_SECS)) - .connect_timeout(Duration::from_secs(MESSAGES_CONNECT_TIMEOUT_SECS)) - .build() - .unwrap_or_else(|_| reqwest::Client::new()) - }) -} diff --git a/litellm-rust/crates/core/src/messages/error.rs b/litellm-rust/crates/core/src/messages/error.rs index 2a9723beb38..76f8813e330 100644 --- a/litellm-rust/crates/core/src/messages/error.rs +++ b/litellm-rust/crates/core/src/messages/error.rs @@ -17,6 +17,8 @@ pub enum Error { #[error(transparent)] Auth(#[from] litellm_auth::Error), #[error(transparent)] + Client(#[from] litellm_http::Error), + #[error(transparent)] Transport(#[from] litellm_http::transport::Error), #[error(transparent)] Headers(#[from] litellm_http::request::HeaderError), diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index de1a5f476ed..f90cb8cb454 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -5,13 +5,15 @@ use litellm_llms::base_llm::anthropic_messages::transformation::BaseAnthropicMes use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use serde_json::Value; -use super::{Error, client::http_client, common_utils::truncate_error_body}; +use super::{Error, common_utils::truncate_error_body}; +use crate::constants::MESSAGES_TIMEOUT_SECS; pub(super) fn network(error: reqwest::Error) -> Error { Error::Transport(TransportError::Network(error.to_string())) } pub(super) async fn send( + http: &litellm_http::Client, url: &str, headers: &[(String, String)], body: &Value, @@ -20,13 +22,11 @@ pub(super) async fn send( let encoded = serde_json::to_vec(body) .map_err(|err| Error::InvalidRequest(format!("failed to encode messages body: {err}")))?; let builder = headers.iter().fold( - http_client().post(url).body(encoded), + http.post(url) + .body(encoded) + .timeout(timeout.unwrap_or(Duration::from_secs(MESSAGES_TIMEOUT_SECS))), |builder, (key, value)| builder.header(key, value), ); - let builder = match timeout { - Some(duration) => builder.timeout(duration), - None => builder, - }; http_request(builder).await.map_err(network) } diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 180eb08810e..5cb83b4e34d 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -7,13 +7,13 @@ mod error; pub mod types; pub use error::Error; -mod client; mod common_utils; mod handler; mod prepare; pub mod route; use std::sync::Arc; +use litellm_http::{ClientVariant, HttpClientConfig, HttpClientPool}; use litellm_secrets::source::EnvironmentSecrets; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}; @@ -21,7 +21,11 @@ use serde_json::Value; use crate::messages::types::MessagesRequest; -pub async fn messages(request: MessagesRequest<'_>) -> Result { +pub async fn messages( + pool: &HttpClientPool, + config: &HttpClientConfig, + request: MessagesRequest<'_>, +) -> Result { let Value::Object(body) = request.body else { return Err(Error::InvalidRequest( "messages body must be an object".into(), @@ -38,8 +42,15 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result Ok(*message), MessagesOutput::Streamed => Err(Error::Unsupported( "streamed responses need a streaming host", diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 40aff185e81..7f6589cdf3e 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -12,6 +12,7 @@ use litellm_host::{ machine::{CallMachine, HostChannel, MachineFault}, protocol::Protocol, }; +use litellm_http::{Client, ClientVariant, HttpClientConfig, HttpClientPool}; use litellm_secrets::source::SecretSource; use litellm_types::{ llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, @@ -108,12 +109,20 @@ impl Host for LocalMessagesHost { } } -pub fn messages_machine(secrets: Arc) -> MessagesMachine { - CallMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) +pub fn messages_machine( + pool: &HttpClientPool, + config: &HttpClientConfig, + secrets: Arc, +) -> Result { + let http = pool.client(config, ClientVariant::Provider)?; + Ok(CallMachine::new(move |host| { + Box::pin(execute(host, http.clone(), secrets.clone())) + })) } async fn execute( host: MessagesHost, + http: Client, secrets: Arc, ) -> Result { let call = host.project().await?; @@ -164,7 +173,7 @@ async fn execute( context, ) .await?; - let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?; + let response = send(&http, &wire.url, &wire.headers, &wire.body, request.timeout).await?; if !response.status().is_success() { return Err(provider_error(response).await); } diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 18961ec96fa..f13d6984763 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -110,7 +110,10 @@ mod tests { } fn client() -> OcrClient { - OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()) + OcrClient::for_test( + litellm_http::Client::plain_for_test(), + litellm_http::Client::no_redirect_for_test(), + ) } fn request(model: &str, base: &str, document: Value, options: Value) -> LiteLLMOcrRequest { diff --git a/litellm-rust/crates/core/tests/audio_transcription.rs b/litellm-rust/crates/core/tests/audio_transcription.rs index 196f085a6c3..612395fe63a 100644 --- a/litellm-rust/crates/core/tests/audio_transcription.rs +++ b/litellm-rust/crates/core/tests/audio_transcription.rs @@ -10,6 +10,10 @@ use support::*; const MODEL: &str = "mistral.voxtral-mini-3b-2507"; +async fn transcribe(request: AudioTranscriptionRequest<'_>) -> Result { + audio_transcription(&http_pool(), &http_config(), request).await +} + fn transcript_response(text: &str) -> ResponseTemplate { json_response(json!({"output": {"message": {"content": [{"text": text}]}}})) } @@ -47,7 +51,7 @@ async fn bedrock_converse_request_is_signed_for_the_requested_region( let upstream = upstream([transcript_response("hello")]).await; let base = upstream.uri(); - let response = audio_transcription(AudioTranscriptionRequest { + let response = transcribe(AudioTranscriptionRequest { api_base: Some(&base), optional_params: aws_params(region), ..request @@ -79,7 +83,7 @@ async fn the_provider_can_come_from_the_model_prefix(request: AudioTranscription let base = upstream.uri(); let model = format!("bedrock/{MODEL}"); - audio_transcription(AudioTranscriptionRequest { + transcribe(AudioTranscriptionRequest { model: &model, custom_llm_provider: None, api_base: Some(&base), @@ -110,7 +114,7 @@ async fn audio_and_transcription_params_reach_the_converse_body( ]) .collect(); - audio_transcription(AudioTranscriptionRequest { + transcribe(AudioTranscriptionRequest { audio: json!({"data": "AQI=", "format": format}), api_base: Some(&base), optional_params, @@ -142,7 +146,7 @@ async fn invalid_audio_is_rejected_before_sending( let upstream = upstream([transcript_response("hello")]).await; let base = upstream.uri(); - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { audio, api_base: Some(&base), ..request @@ -174,7 +178,7 @@ async fn unsupported_providers_are_rejected_before_sending( #[case] provider: Option<&'static str>, #[case] reported: &str, ) { - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { model, custom_llm_provider: provider, api_base: Some(UNREACHABLE_BASE), @@ -189,7 +193,7 @@ async fn unsupported_providers_are_rejected_before_sending( #[rstest] #[tokio::test] async fn a_non_string_extra_header_is_rejected(request: AudioTranscriptionRequest<'static>) { - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { extra_headers: Some(Map::from_iter([("x-count".to_string(), json!(3))])), api_base: Some(UNREACHABLE_BASE), ..request @@ -212,7 +216,7 @@ async fn an_upstream_error_keeps_its_status_and_body( upstream([ResponseTemplate::new(status).set_body_string("upstream said no")]).await; let base = upstream.uri(); - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { api_base: Some(&base), ..request }) @@ -239,7 +243,7 @@ async fn an_unreadable_success_body_is_an_invalid_response( let upstream = upstream([response]).await; let base = upstream.uri(); - let error = audio_transcription(AudioTranscriptionRequest { + let error = transcribe(AudioTranscriptionRequest { api_base: Some(&base), ..request }) diff --git a/litellm-rust/crates/core/tests/chat_completions.rs b/litellm-rust/crates/core/tests/chat_completions.rs index ae96509fe2e..d1f6cde19e8 100644 --- a/litellm-rust/crates/core/tests/chat_completions.rs +++ b/litellm-rust/crates/core/tests/chat_completions.rs @@ -4,6 +4,7 @@ use litellm_core::chat_completions::{ Error, chat_completions, chat_completions_decline_reason, types::ChatCompletionsRequest, }; use litellm_http::transport::Error as TransportError; +use litellm_types::utils::ChatCompletionsResponse; use rstest::{fixture, rstest}; use serde_json::{Map, Value, json}; use wiremock::ResponseTemplate; @@ -13,6 +14,10 @@ use support::*; const ANTHROPIC_MESSAGE: &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}}"#; +async fn complete(request: ChatCompletionsRequest<'_>) -> Result { + chat_completions(&http_pool(), &http_config(), request).await +} + fn object(value: Value) -> Map { let Value::Object(map) = value else { panic!("expected a json object, got {value}"); @@ -50,7 +55,7 @@ async fn anthropic_round_trip_translates_the_conversation_and_normalizes_the_res let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; let base = upstream.uri(); - let response = chat_completions(ChatCompletionsRequest { + let response = complete(ChatCompletionsRequest { messages: json!([ {"role": "system", "content": "be terse"}, {"role": "user", "content": "hi"} @@ -90,7 +95,7 @@ async fn the_deployment_key_replaces_a_caller_supplied_x_api_key( let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; let base = upstream.uri(); - chat_completions(ChatCompletionsRequest { + complete(ChatCompletionsRequest { api_base: Some(&base), extra_headers: Some(object( json!({"x-api-key": "caller-key", "x-trace": "kept"}), @@ -116,7 +121,7 @@ async fn bedrock_round_trip_is_signed_and_normalized(request: ChatCompletionsReq .await; let base = upstream.uri(); - let response = chat_completions(ChatCompletionsRequest { + let response = complete(ChatCompletionsRequest { model: "bedrock/anthropic.claude-sonnet-4-5", optional_params: object(json!({ "aws_access_key_id": "access-key", @@ -167,7 +172,7 @@ async fn a_response_it_cannot_normalize_is_reported_as_already_sent( let upstream = upstream([anthropic_response(body)]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(&base), ..request }) @@ -188,7 +193,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body( let upstream = upstream([ResponseTemplate::new(status).set_body_string("slow down")]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(&base), ..request }) @@ -210,7 +215,7 @@ async fn an_upstream_error_status_keeps_its_code_and_body( async fn a_connection_that_is_never_established_declines_instead_of_failing( request: ChatCompletionsRequest<'static>, ) { - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(UNREACHABLE_BASE), ..request }) @@ -232,7 +237,7 @@ async fn a_timeout_after_sending_is_not_a_pre_send_decline( upstream([anthropic_response(ANTHROPIC_MESSAGE).set_delay(Duration::from_secs(5))]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { api_base: Some(&base), timeout: Some(Duration::from_millis(100)), ..request @@ -307,7 +312,7 @@ async fn a_declined_request_fails_the_call_before_sending( let upstream = upstream([anthropic_response(ANTHROPIC_MESSAGE)]).await; let base = upstream.uri(); - let error = chat_completions(ChatCompletionsRequest { + let error = complete(ChatCompletionsRequest { optional_params: object(json!({"stream": true})), api_base: Some(&base), ..request diff --git a/litellm-rust/crates/core/tests/messages/host.rs b/litellm-rust/crates/core/tests/messages/host.rs index ca2aece5ebd..844ada3e1ad 100644 --- a/litellm-rust/crates/core/tests/messages/host.rs +++ b/litellm-rust/crates/core/tests/messages/host.rs @@ -78,7 +78,7 @@ impl Host for RecordingHost { } async fn run_through(host: &RecordingHost) -> Result { - litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await + litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await } fn authenticated(call: MessagesCall, api_base: String) -> MessagesCall { diff --git a/litellm-rust/crates/core/tests/messages/main.rs b/litellm-rust/crates/core/tests/messages/main.rs index 21ee678ced3..1ae822e5437 100644 --- a/litellm-rust/crates/core/tests/messages/main.rs +++ b/litellm-rust/crates/core/tests/messages/main.rs @@ -2,9 +2,11 @@ use std::{sync::Arc, time::Duration}; use litellm_core::messages::{ Error, - route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}, + route::{LocalMessagesHost, MessagesCall, MessagesMachine, MessagesOutput, messages_machine}, types::MessagesShaping, }; +use litellm_http::{HttpSettings, Resolution}; +use litellm_secrets::source::SecretSource; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use rstest::fixture; use serde_json::{Map, Value, json}; @@ -75,11 +77,16 @@ fn headers<'a>(pairs: impl IntoIterator) -> Option) -> MessagesMachine { + messages_machine(&http_pool(), &http_config(), secrets) + .expect("default HTTP settings build a client") +} + async fn run_with( secrets: Arc, call: MessagesCall, ) -> Result { - litellm_host::run::run(messages_machine(secrets), &LocalMessagesHost::new(call)).await + litellm_host::run::run(machine(secrets), &LocalMessagesHost::new(call)).await } /// Runs the route with a secret source that knows nothing, so no environment leaks in. diff --git a/litellm-rust/crates/core/tests/messages/response.rs b/litellm-rust/crates/core/tests/messages/response.rs index 133b7d2b162..431dd4f4b93 100644 --- a/litellm-rust/crates/core/tests/messages/response.rs +++ b/litellm-rust/crates/core/tests/messages/response.rs @@ -190,29 +190,40 @@ fn facade_request(body: Value, api_base: &str) -> MessagesRequest<'_> { } #[tokio::test] -async fn the_facade_runs_the_route_in_process() { +async fn the_facade_sends_through_the_injected_http_pool_configuration() { let upstream = upstream([message_response()]).await; let base = upstream.uri(); + let settings = HttpSettings { + user_agent: Some("host-owned/1".into()), + ..HttpSettings::default() + }; - let message = messages(facade_request( - json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}), - &base, - )) + let message = messages( + &http_pool(), + &Resolution::from(&settings).config, + facade_request( + json!({"model": MODEL, "max_tokens": 16, "messages": [{"role": "user", "content": "hi"}]}), + &base, + ), + ) .await .expect("messages request succeeds"); assert_eq!(message.id, "msg_1"); - assert_eq!( - only_request(&upstream).await.header("x-api-key"), - Some("sk-ant") - ); + let sent = only_request(&upstream).await; + assert_eq!(sent.header("x-api-key"), Some("sk-ant")); + assert_eq!(sent.header("user-agent"), Some("host-owned/1")); } #[tokio::test] async fn the_facade_rejects_a_body_that_is_not_an_object() { - let error = messages(facade_request(json!([]), UNREACHABLE_BASE)) - .await - .expect_err("a non-object body is rejected"); + let error = messages( + &http_pool(), + &http_config(), + facade_request(json!([]), UNREACHABLE_BASE), + ) + .await + .expect_err("a non-object body is rejected"); assert_eq!( error, diff --git a/litellm-rust/crates/core/tests/messages/stream.rs b/litellm-rust/crates/core/tests/messages/stream.rs index c4be3127d66..4ca6e609052 100644 --- a/litellm-rust/crates/core/tests/messages/stream.rs +++ b/litellm-rust/crates/core/tests/messages/stream.rs @@ -87,7 +87,7 @@ fn sse_response() -> ResponseTemplate { } async fn stream_through(host: &RecordingStreamHost) -> Result { - litellm_host::run::run(messages_machine(Arc::new(RecordingSecrets::empty())), host).await + litellm_host::run::run(machine(Arc::new(RecordingSecrets::empty())), host).await } #[rstest] diff --git a/litellm-rust/crates/core/tests/ocr/main.rs b/litellm-rust/crates/core/tests/ocr/main.rs index 1a915389b20..e1f6b8cb5c1 100644 --- a/litellm-rust/crates/core/tests/ocr/main.rs +++ b/litellm-rust/crates/core/tests/ocr/main.rs @@ -4,6 +4,7 @@ use litellm_core::ocr::{ types::LiteLLMOcrRequest, wire::{OcrWireRequest, decode_request}, }; +use litellm_http::Client; use litellm_llms::base_llm::ocr::{ error::Error, handler::OcrClient, @@ -37,11 +38,7 @@ fn object(value: Value) -> Map { } 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) + OcrClient::for_test(Client::plain_for_test(), Client::no_redirect_for_test()) } async fn perform(request: LiteLLMOcrRequest) -> Result { diff --git a/litellm-rust/crates/core/tests/ocr/mistral.rs b/litellm-rust/crates/core/tests/ocr/mistral.rs index f80e564b03f..4c3f1c5cc39 100644 --- a/litellm-rust/crates/core/tests/ocr/mistral.rs +++ b/litellm-rust/crates/core/tests/ocr/mistral.rs @@ -1,10 +1,7 @@ use std::sync::Arc; use litellm_auth_gcp::VertexAuth; -use litellm_http::{ - HttpClientPool, HttpSettings, Resolution, - media::{PublicDnsResolver, UrlPolicy}, -}; +use litellm_http::{HttpSettings, Resolution, media::UrlPolicy}; use litellm_llms::{ base_llm::ocr::{ settings::OcrSettings, @@ -192,12 +189,16 @@ async fn the_client_uses_the_injected_http_pool_configuration() { ..HttpSettings::default() }; let client = OcrClient::new( - &HttpClientPool::new(Arc::new(PublicDnsResolver)), + &http_pool(), &Resolution::from(&settings).config, UrlPolicy::default(), VertexAuth::default(), OcrSettings::default(), - Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), + Arc::new( + litellm_secrets::source::EnvironmentSecrets::python_compatible( + litellm_http::Client::plain_for_test(), + ), + ), ) .unwrap(); diff --git a/litellm-rust/crates/core/tests/support/mod.rs b/litellm-rust/crates/core/tests/support/mod.rs index 4d2fe0232d0..1d9af236811 100644 --- a/litellm-rust/crates/core/tests/support/mod.rs +++ b/litellm-rust/crates/core/tests/support/mod.rs @@ -3,9 +3,12 @@ #![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset -use std::sync::Mutex; +use std::sync::{Arc, Mutex}; use futures_util::future::BoxFuture; +use litellm_http::{ + HttpClientConfig, HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver, +}; use litellm_secrets::{SecretValue, source::SecretSource}; use serde_json::Value; use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; @@ -13,6 +16,14 @@ use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any}; /// A port nothing listens on, for calls that must fail before any request is sent. pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1"; +pub fn http_pool() -> HttpClientPool { + HttpClientPool::new(Arc::new(PublicDnsResolver)) +} + +pub fn http_config() -> HttpClientConfig { + Resolution::from(&HttpSettings::default()).config +} + /// Starts an upstream that answers its n-th request with the n-th response and 404s after. pub async fn upstream(responses: impl IntoIterator) -> MockServer { let server = MockServer::start().await; diff --git a/litellm-rust/crates/http/Cargo.toml b/litellm-rust/crates/http/Cargo.toml index cad5aa87e49..0cb2b15b768 100644 --- a/litellm-rust/crates/http/Cargo.toml +++ b/litellm-rust/crates/http/Cargo.toml @@ -22,5 +22,7 @@ veil.workspace = true webpki-roots.workspace = true [dev-dependencies] +rcgen = "0.14.10" +tempfile.workspace = true rstest.workspace = true tokio.workspace = true diff --git a/litellm-rust/crates/http/src/client.rs b/litellm-rust/crates/http/src/client.rs new file mode 100644 index 00000000000..1f7017d083b --- /dev/null +++ b/litellm-rust/crates/http/src/client.rs @@ -0,0 +1,38 @@ +use std::ops::Deref; + +#[derive(Clone, Debug)] +pub struct Client(reqwest::Client); + +impl Client { + pub(crate) fn new(client: reqwest::Client) -> Self { + Self(client) + } + + #[cfg(any(test, feature = "test-support"))] + pub fn for_test(client: reqwest::Client) -> Self { + Self(client) + } + + #[cfg(any(test, feature = "test-support"))] + pub fn plain_for_test() -> Self { + Self(reqwest::Client::new()) + } + + #[cfg(any(test, feature = "test-support"))] + pub fn no_redirect_for_test() -> Self { + Self( + reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("a client without TLS or proxy settings builds"), + ) + } +} + +impl Deref for Client { + type Target = reqwest::Client; + + fn deref(&self) -> &reqwest::Client { + &self.0 + } +} diff --git a/litellm-rust/crates/http/src/config.rs b/litellm-rust/crates/http/src/config.rs index cb0173369d5..2f36784bc70 100644 --- a/litellm-rust/crates/http/src/config.rs +++ b/litellm-rust/crates/http/src/config.rs @@ -18,10 +18,16 @@ pub enum Verify { BuiltInRoots, } +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +pub enum ClientIdentity { + Pem(PathBuf), + Split { certificate: PathBuf, key: PathBuf }, +} + #[derive(Clone, Debug, PartialEq, Eq, Hash)] pub struct HttpClientConfig { pub verify: Verify, - pub client_certificate: Option, + pub client_certificate: Option, pub key_exchange_group: Option, pub tls12_cipher_suites: Option>, pub force_ipv4: bool, @@ -67,7 +73,7 @@ impl From<&HttpSettings> for Resolution { Self { config: HttpClientConfig { verify: Verify::from(settings), - client_certificate: settings.ssl_certificate.clone(), + client_certificate: settings.ssl_certificate.clone().map(ClientIdentity::Pem), key_exchange_group: curve.clone().ok().flatten(), tls12_cipher_suites: ciphers.tls12_cipher_suites, force_ipv4: settings.force_ipv4, @@ -276,7 +282,7 @@ mod tests { config, HttpClientConfig { verify: Verify::BuiltInRoots, - client_certificate: Some("/client.pem".into()), + client_certificate: Some(ClientIdentity::Pem("/client.pem".into())), key_exchange_group: None, tls12_cipher_suites: None, force_ipv4: true, diff --git a/litellm-rust/crates/http/src/lib.rs b/litellm-rust/crates/http/src/lib.rs index a1456208bb3..3e55a1843c8 100644 --- a/litellm-rust/crates/http/src/lib.rs +++ b/litellm-rust/crates/http/src/lib.rs @@ -1,3 +1,10 @@ +#![allow( + clippy::disallowed_types, + clippy::disallowed_methods, + reason = "this crate is the one place reqwest clients are built" +)] + +mod client; mod config; mod error; pub mod media; @@ -9,7 +16,8 @@ mod settings; mod tls; pub mod transport; -pub use config::{HttpClientConfig, Resolution, Verify}; +pub use client::Client; +pub use config::{ClientIdentity, HttpClientConfig, Resolution, Verify}; pub use error::{Error, TlsSource}; pub use pool::{ClientVariant, HttpClientPool}; pub use proxy::EnvironmentProxies; diff --git a/litellm-rust/crates/http/src/media.rs b/litellm-rust/crates/http/src/media.rs index 1b9159973ef..1dac68305b0 100644 --- a/litellm-rust/crates/http/src/media.rs +++ b/litellm-rust/crates/http/src/media.rs @@ -12,7 +12,7 @@ use reqwest::{ dns::{Addrs, Name, Resolve, Resolving}, }; -use crate::{ClientVariant, HttpClientConfig, HttpClientPool}; +use crate::{Client, ClientVariant, HttpClientConfig, HttpClientPool}; #[derive(Debug, thiserror::Error)] pub enum Error { @@ -93,8 +93,8 @@ type ProxyMatch = Arc bool + Send + Sync>; #[derive(Clone)] pub struct MediaFetcher { - pinned: reqwest::Client, - unpinned: reqwest::Client, + pinned: Client, + unpinned: Client, uses_proxy: ProxyMatch, address_resolver: Arc, url_policy: UrlPolicy, @@ -154,7 +154,7 @@ impl MediaFetcher { } #[cfg(any(test, feature = "test-support"))] - pub fn for_test(client: reqwest::Client) -> Self { + pub fn for_test(client: Client) -> Self { Self { pinned: client.clone(), unpinned: client, @@ -230,7 +230,7 @@ impl MediaFetcher { } } - async fn client_for(&self, url: &Url) -> Result<&reqwest::Client, Error> { + async fn client_for(&self, url: &Url) -> Result<&Client, Error> { if !self.url_policy.validate { return Ok(&self.unpinned); } @@ -520,10 +520,7 @@ mod tests { b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf; charset=binary\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc", ) .await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test client builds"); + let client = Client::no_redirect_for_test(); let media = MediaFetcher::for_test(client) .fetch(url, policy(3, 0)) .await @@ -539,10 +536,7 @@ mod tests { b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc", ) .await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test client builds"); + let client = Client::no_redirect_for_test(); let error = MediaFetcher::for_test(client) .fetch(url, policy(2, 0)) .await @@ -557,10 +551,7 @@ mod tests { b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n2\r\nab\r\n2\r\ncd\r\n0\r\n\r\n", ) .await; - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test client builds"); + let client = Client::no_redirect_for_test(); let error = MediaFetcher::for_test(client) .fetch(url, policy(3, 0)) .await diff --git a/litellm-rust/crates/http/src/outbound.rs b/litellm-rust/crates/http/src/outbound.rs index d100bdf624b..c2cfb00d79b 100644 --- a/litellm-rust/crates/http/src/outbound.rs +++ b/litellm-rust/crates/http/src/outbound.rs @@ -107,7 +107,7 @@ impl OutboundRequest { self.timeout } - pub async fn send(self, client: &reqwest::Client) -> Result { + pub async fn send(self, client: &crate::Client) -> Result { let builder = with_headers( client.post(&self.url).body(self.body), &self.headers, diff --git a/litellm-rust/crates/http/src/pool.rs b/litellm-rust/crates/http/src/pool.rs index ee47e5dc52a..1187c34f2d7 100644 --- a/litellm-rust/crates/http/src/pool.rs +++ b/litellm-rust/crates/http/src/pool.rs @@ -6,7 +6,7 @@ use std::{ use reqwest::dns::Resolve; -use crate::{config::HttpClientConfig, error::Error, proxy::EnvironmentProxies}; +use crate::{client::Client, config::HttpClientConfig, error::Error, proxy::EnvironmentProxies}; #[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] pub enum ClientVariant { @@ -48,7 +48,7 @@ impl HttpClientPool { &self, config: &HttpClientConfig, variant: ClientVariant, - ) -> Result { + ) -> Result { let effective = match variant { ClientVariant::Media => HttpClientConfig { client_certificate: None, @@ -65,7 +65,7 @@ impl HttpClientPool { if let Some(pooled) = self.lock().get(&key) && pooled.built_at.elapsed() < self.ttl { - return Ok(pooled.client.clone()); + return Ok(Client::new(pooled.client.clone())); } let client = self .apply(variant, reqwest::ClientBuilder::try_from(&key.0)?) @@ -77,7 +77,7 @@ impl HttpClientPool { built_at: Instant::now(), }, ); - Ok(client) + Ok(Client::new(client)) } fn lock(&self) -> MutexGuard<'_, Clients> { @@ -116,7 +116,7 @@ mod tests { }; use super::*; - use crate::{HttpSettings, Resolution, Verify}; + use crate::{ClientIdentity, HttpSettings, Resolution, Verify}; struct FixedResolver(SocketAddr); @@ -288,7 +288,9 @@ mod tests { fn media_variant_never_loads_the_client_certificate() { let pool = pool(); let with_identity = HttpClientConfig { - client_certificate: Some(std::env::temp_dir().join("litellm-http-absent-client.pem")), + client_certificate: Some(ClientIdentity::Pem( + std::env::temp_dir().join("litellm-http-absent-client.pem"), + )), ..config("a") }; assert!( diff --git a/litellm-rust/crates/http/src/tls.rs b/litellm-rust/crates/http/src/tls.rs index e2e6d27cd54..c58076607e4 100644 --- a/litellm-rust/crates/http/src/tls.rs +++ b/litellm-rust/crates/http/src/tls.rs @@ -8,7 +8,7 @@ use rustls::{ }; use crate::{ - config::{HttpClientConfig, Verify}, + config::{ClientIdentity, HttpClientConfig, Verify}, error::{Error, TlsSource}, }; @@ -203,11 +203,15 @@ impl TryFrom<&HttpClientConfig> for ClientConfig { }; let mut tls = match &config.client_certificate { None => verified.with_no_client_auth(), - Some(path) => { - let (chain, key) = identity(path, TlsSource::ClientIdentity)?; + Some(identity) => { + let (certificate, key) = match identity { + ClientIdentity::Pem(path) => (path, path), + ClientIdentity::Split { certificate, key } => (certificate, key), + }; + let (chain, private_key) = client_identity(certificate, key)?; verified - .with_client_auth_cert(chain, key) - .map_err(|error| invalid_pem(path, TlsSource::ClientIdentity, error))? + .with_client_auth_cert(chain, private_key) + .map_err(|error| invalid_pem(key, TlsSource::ClientIdentity, error))? } }; tls.alpn_protocols = if config.http2 { @@ -233,17 +237,18 @@ fn bundle_roots(path: &Path, source: TlsSource) -> Result Ok(store) } -fn identity( - path: &Path, - source: TlsSource, +fn client_identity( + certificate: &Path, + key: &Path, ) -> Result<(Vec>, PrivateKeyDer<'static>), Error> { - let chain = certificates(path, source)?; + let source = TlsSource::ClientIdentity; + let chain = certificates(certificate, source)?; if chain.is_empty() { - return Err(invalid_pem(path, source, "no certificates found")); + return Err(invalid_pem(certificate, source, "no certificates found")); } - let key = PrivateKeyDer::from_pem_slice(&read(path, source)?) - .map_err(|error| invalid_pem(path, source, error))?; - Ok((chain, key)) + let private_key = PrivateKeyDer::from_pem_slice(&read(key, source)?) + .map_err(|error| invalid_pem(key, source, error))?; + Ok((chain, private_key)) } fn certificates(path: &Path, source: TlsSource) -> Result>, Error> { @@ -405,7 +410,7 @@ mod tests { ) .unwrap(); let result = ClientConfig::try_from(&HttpClientConfig { - client_certificate: Some(path.clone()), + client_certificate: Some(ClientIdentity::Pem(path.clone())), ..config(HttpSettings::default()) }) .map(drop); @@ -419,4 +424,29 @@ mod tests { }) if reported == path )); } + + #[test] + fn split_client_identity_reads_the_key_from_its_own_file() { + let identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let directory = tempfile::tempdir().unwrap(); + let certificate = directory.path().join("client.crt"); + let key = directory.path().join("client.key"); + std::fs::write(&certificate, identity.cert.pem()).unwrap(); + std::fs::write(&key, identity.signing_key.serialize_pem()).unwrap(); + + let split = ClientConfig::try_from(&HttpClientConfig { + client_certificate: Some(ClientIdentity::Split { + certificate: certificate.clone(), + key, + }), + ..config(HttpSettings::default()) + }); + let combined = ClientConfig::try_from(&HttpClientConfig { + client_certificate: Some(ClientIdentity::Pem(certificate)), + ..config(HttpSettings::default()) + }); + + assert!(split.unwrap().client_auth_cert_resolver.has_certs()); + assert!(combined.is_err()); + } } diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index ed15d9f7cdb..36ccd18f220 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -35,6 +35,7 @@ tokio = { workspace = true, features = ["sync"] } url.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } aws-smithy-eventstream = "=0.61.4" aws-smithy-types = "1.6.1" rstest.workspace = true diff --git a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs index 51a2668310e..b4e9d01f867 100644 --- a/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/ocr/document_intelligence/transformation.rs @@ -450,7 +450,7 @@ fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result: Send + Sync { #[derive(Clone)] pub struct OcrClient { - provider_http: reqwest::Client, - polling_http: reqwest::Client, + provider_http: Client, + polling_http: Client, document_fetcher: MediaFetcher, vertex_auth: VertexAuth, settings: OcrSettings, @@ -60,11 +60,11 @@ impl OcrClient { }) } - pub fn provider_http(&self) -> &reqwest::Client { + pub fn provider_http(&self) -> &Client { &self.provider_http } - pub fn polling_http(&self) -> &reqwest::Client { + pub fn polling_http(&self) -> &Client { &self.polling_http } @@ -85,17 +85,18 @@ impl OcrClient { } #[cfg(any(test, feature = "test-support"))] - pub fn for_test(provider_http: reqwest::Client, document_http: reqwest::Client) -> Self { + pub fn for_test(provider_http: Client, no_redirect_http: Client) -> Self { Self { + secrets: Arc::new( + litellm_secrets::source::EnvironmentSecrets::python_compatible( + provider_http.clone(), + ), + ), provider_http, - polling_http: reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .expect("test polling client builds"), - document_fetcher: MediaFetcher::for_test(document_http), + polling_http: no_redirect_http.clone(), + document_fetcher: MediaFetcher::for_test(no_redirect_http), vertex_auth: VertexAuth::default(), settings: OcrSettings::default(), - secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), } } @@ -311,7 +312,7 @@ mod tests { let _connection = listener.accept().await.unwrap(); tokio::time::sleep(Duration::from_secs(1)).await; }); - let error = reqwest::Client::new() + let error = litellm_http::Client::plain_for_test() .get(format!("http://{address}")) .timeout(Duration::from_millis(10)) .send() diff --git a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs index c3377536545..147056dab8d 100644 --- a/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/reducto/ocr/transformation.rs @@ -564,7 +564,10 @@ mod tests { let params = ReductoParseV3Config .map_ocr_params(&overrides, "parse-v3") .unwrap(); - let client = OcrClient::for_test(reqwest::Client::new(), reqwest::Client::new()); + let client = OcrClient::for_test( + litellm_http::Client::plain_for_test(), + litellm_http::Client::no_redirect_for_test(), + ); let connection = OcrConnection::default(); let document = serde_json::from_value( json!({"type":"document_url","document_url":"reducto://ready.pdf"}), diff --git a/litellm-rust/crates/llms/tests/ocr_handler.rs b/litellm-rust/crates/llms/tests/ocr_handler.rs index 6e46e6f76d4..5ed3244087d 100644 --- a/litellm-rust/crates/llms/tests/ocr_handler.rs +++ b/litellm-rust/crates/llms/tests/ocr_handler.rs @@ -19,7 +19,7 @@ async fn read_bounded(response: String, limit: usize) -> Result().await; }); - let response = reqwest::Client::new() + let response = litellm_http::Client::plain_for_test() .get(format!("http://{address}")) .send() .await diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 057cad2f42e..d965223bd59 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -62,6 +62,7 @@ tokio = { workspace = true, features = ["rt", "sync"] } url.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } litellm-secrets-aws.workspace = true serde.workspace = true serde_with.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/cache/activation.rs b/litellm-rust/crates/python-bridge/src/cache/activation.rs index f77032c579d..58735679554 100644 --- a/litellm-rust/crates/python-bridge/src/cache/activation.rs +++ b/litellm-rust/crates/python-bridge/src/cache/activation.rs @@ -9,10 +9,10 @@ use super::{ cache_error, config::{CacheBackendConfig, NativeCacheConfig, UnsupportedCacheConfig}, embedder::PythonEmbedder, - host_client, native::NativeResponseCache, }; use crate::errors::RustBridgeDeclined; +use crate::http::host_client; fn declined(reason: UnsupportedCacheConfig) -> PyErr { RustBridgeDeclined::new_err(reason.message()) diff --git a/litellm-rust/crates/python-bridge/src/cache/config.rs b/litellm-rust/crates/python-bridge/src/cache/config.rs index 6e25f07efa1..e58902b07ee 100644 --- a/litellm-rust/crates/python-bridge/src/cache/config.rs +++ b/litellm-rust/crates/python-bridge/src/cache/config.rs @@ -1511,7 +1511,7 @@ mod tests { path_service_account: Some("credentials.json".into()), endpoint: litellm_cache_gcs::DEFAULT_ENDPOINT.into(), }, - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), Some("token".into()), ); let matching_config = NativeCacheConfig { diff --git a/litellm-rust/crates/python-bridge/src/cache/handle.rs b/litellm-rust/crates/python-bridge/src/cache/handle.rs index e14916b25c6..fcc8aa6218a 100644 --- a/litellm-rust/crates/python-bridge/src/cache/handle.rs +++ b/litellm-rust/crates/python-bridge/src/cache/handle.rs @@ -1,3 +1,4 @@ +use crate::http::host_client; use crate::logger::run_sync_value; use litellm_auth_aws::AwsAuthConfig; use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig}; @@ -19,7 +20,6 @@ use super::{ config::{QdrantSemanticCacheConfig, project_redis_semantic}, embedder::PythonEmbedder, facade::FacadeGuard, - host_client, native::NativeResponseCache, request::duration, }; diff --git a/litellm-rust/crates/python-bridge/src/cache/mod.rs b/litellm-rust/crates/python-bridge/src/cache/mod.rs index ac1e00d5273..00b0c71684a 100644 --- a/litellm-rust/crates/python-bridge/src/cache/mod.rs +++ b/litellm-rust/crates/python-bridge/src/cache/mod.rs @@ -13,11 +13,9 @@ mod resolver; mod semantic; use litellm_cache::Error; -use litellm_http::ClientVariant; use pyo3::{ exceptions::{PyNotImplementedError, PyRuntimeError, PyValueError}, prelude::*, - types::PyDict, }; pub(crate) use self::{binding::ResolvedCache, handle::CacheTestHandle, resolver::CacheResolver}; @@ -29,11 +27,3 @@ fn cache_error(error: Error) -> PyErr { _ => PyRuntimeError::new_err(error.to_string()), } } - -/// The host's pooled HTTP client, configured from the proxy's HTTP settings. -fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult { - let http_config = crate::http::call_config(py, &PyDict::new(py), true)?; - crate::http::pool() - .client(&http_config, variant) - .map_err(crate::http::client_error) -} diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native.rs index 0e279046812..460136baa1f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native.rs @@ -92,7 +92,7 @@ impl NativeResponseCache { )) } - pub async fn s3(config: S3CacheConfig, http: reqwest::Client) -> Self { + pub async fn s3(config: S3CacheConfig, http: litellm_http::Client) -> Self { let runtime = tokio::runtime::Handle::current(); let backend = S3Cache::new(config, http, ResponseCacheCodec, runtime); let identity = BackendIdentity::S3 { @@ -112,7 +112,7 @@ impl NativeResponseCache { Ok(Self::exact(ResponseCache::new(Arc::new(backend)), identity)) } - pub fn gcs(config: GcsConfig, client: reqwest::Client, token: Option) -> Self { + pub fn gcs(config: GcsConfig, client: litellm_http::Client, token: Option) -> Self { let backend = match token { Some(token) => GcsCache::with_token_source( config, @@ -133,7 +133,7 @@ impl NativeResponseCache { pub async fn azure_blob( account_url: &str, container: &str, - http: reqwest::Client, + http: litellm_http::Client, ) -> Result { let backend = AzureBlobCache::connect( account_url, @@ -242,7 +242,7 @@ impl NativeResponseCache { pub async fn qdrant_semantic( config: QdrantSemanticCacheConfig, - client: reqwest::Client, + client: litellm_http::Client, runtime: tokio::runtime::Handle, ) -> Result { let qdrant = qdrant_client::Qdrant::from_url(&config.grpc_url) diff --git a/litellm-rust/crates/python-bridge/src/http.rs b/litellm-rust/crates/python-bridge/src/http.rs index 3dad3447f45..4d8f0fd7147 100644 --- a/litellm-rust/crates/python-bridge/src/http.rs +++ b/litellm-rust/crates/python-bridge/src/http.rs @@ -6,8 +6,8 @@ use std::{ use litellm_core_utils::settings::ProcessEnvironment; use litellm_http::{ - HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, Resolution, SslVerify, - TlsSource, Unsupported, + Client, ClientVariant, HttpClientConfig, HttpClientPool, HttpSettings, HttpSettingsLayer, + Resolution, SslVerify, TlsSource, Unsupported, media::{PublicDnsResolver, UrlPolicy}, }; use pyo3::{ @@ -97,7 +97,10 @@ pub(crate) fn call_config( let settings = HttpSettings::from_layers([ for_call(call_ssl_verify(kwargs)?, asynchronous), HttpSettingsLayer::from_environment(&ProcessEnvironment), - configured(&PythonSettings::Http.read(py)?)?, + match PythonSettings::Http.read_or_unset(py)? { + Some(snapshot) => configured(&snapshot)?, + None => HttpSettingsLayer::default(), + }, ]) .without_missing_files(&|path: &Path| path.exists()); let resolution = Resolution::from(&settings); @@ -107,6 +110,11 @@ pub(crate) fn call_config( Ok(resolution.config) } +pub(crate) fn host_client(py: Python<'_>, variant: ClientVariant) -> PyResult { + let config = call_config(py, &PyDict::new(py), true)?; + pool().client(&config, variant).map_err(client_error) +} + pub(crate) fn client_error(error: litellm_http::Error) -> PyErr { match error { litellm_http::Error::Read { @@ -143,7 +151,10 @@ fn unreported( } pub(crate) fn url_policy(py: Python<'_>) -> PyResult { - project_url_policy(&PythonSettings::UrlPolicy.read(py)?) + match PythonSettings::UrlPolicy.read_or_unset(py)? { + Some(snapshot) => project_url_policy(&snapshot), + None => Ok(UrlPolicy::default()), + } } fn project_url_policy(snapshot: &Snapshot<'_>) -> PyResult { diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index abf664b795d..f03e5fdce7f 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -1,4 +1,4 @@ -use pyo3::prelude::*; +use pyo3::{exceptions::PyModuleNotFoundError, prelude::*}; use crate::coercion::{FieldSpec, ProjectionError}; @@ -40,6 +40,16 @@ impl PythonSettings { Ok(Snapshot { group: self, value }) } + /// Reads the accessor, or `None` when the litellm package is not installed + /// (a bare extension module), meaning there are no configured values. + pub(crate) fn read_or_unset(self, py: Python<'_>) -> PyResult>> { + match self.read(py) { + Ok(snapshot) => Ok(Some(snapshot)), + Err(error) if error.is_instance_of::(py) => Ok(None), + Err(error) => Err(error), + } + } + #[cfg(test)] pub(crate) fn snapshot(self, value: Bound<'_, PyAny>) -> Snapshot<'_> { Snapshot { group: self, value } diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index dec4dcea21c..93d0e11d323 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -3,7 +3,8 @@ use litellm_core::audio_transcription::{ Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest, }; use litellm_host_python::from_py_argument; -use pyo3::prelude::*; +use litellm_http::HttpClientConfig; +use pyo3::{prelude::*, types::PyDict}; use serde_json::{Map, Value}; use crate::{ @@ -12,6 +13,7 @@ use crate::{ }; async fn execute( + config: HttpClientConfig, audio: Value, optional_params: Map, options: RouteOptions, @@ -24,16 +26,20 @@ async fn execute( extra_headers, timeout, } = options; - run_audio_transcription(AudioTranscriptionRequest { - model: &model, - audio, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - optional_params, - timeout, - }) + run_audio_transcription( + crate::http::pool(), + &config, + AudioTranscriptionRequest { + model: &model, + audio, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + optional_params, + timeout, + }, + ) .await } @@ -62,9 +68,10 @@ pub(crate) fn transcription( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), false)?; run_sync( py, - execute(audio, optional_params.unwrap_or_default(), options), + execute(config, audio, optional_params.unwrap_or_default(), options), audio_transcription_error_to_pyerr, ) } @@ -94,9 +101,10 @@ pub(crate) fn atranscription<'py>( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), true)?; run_async( py, - execute(audio, optional_params.unwrap_or_default(), options), + execute(config, audio, optional_params.unwrap_or_default(), options), audio_transcription_error_to_pyerr, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index b96b12bfc43..6d7fad0d69c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -7,6 +7,7 @@ use litellm_core::chat_completions::{ types::ChatCompletionsRequest, }; use litellm_host_python::from_py_argument; +use litellm_http::HttpClientConfig; use litellm_types::utils::ChatCompletionsResponse; use pyo3::prelude::*; use serde_json::{Map, Value}; @@ -20,6 +21,7 @@ use crate::{ }; async fn execute( + config: HttpClientConfig, messages: Vec, optional_params: Map, options: RouteOptions, @@ -32,16 +34,20 @@ async fn execute( extra_headers, timeout, } = options; - run_chat_completions(ChatCompletionsRequest { - model: &model, - messages: Value::Array(messages), - optional_params, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), - extra_headers, - timeout, - }) + run_chat_completions( + crate::http::pool(), + &config, + ChatCompletionsRequest { + model: &model, + messages: Value::Array(messages), + optional_params, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + timeout, + }, + ) .await } @@ -87,9 +93,15 @@ pub(crate) fn chat_completions( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), false)?; run_sync( py, - execute(messages, optional_params.unwrap_or_default(), options), + execute( + config, + messages, + optional_params.unwrap_or_default(), + options, + ), chat_completions_error_to_pyerr, ) } @@ -119,9 +131,15 @@ pub(crate) fn achat_completions<'py>( extra_headers, timeout: optional_timeout(timeout_seconds), }; + let config = crate::http::call_config(py, &PyDict::new(py), true)?; run_async( py, - execute(messages, optional_params.unwrap_or_default(), options), + execute( + config, + messages, + optional_params.unwrap_or_default(), + options, + ), chat_completions_error_to_pyerr, ) } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index 52cebb7c903..a59c9360c36 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -27,11 +27,14 @@ fn run_messages( asynchronous: bool, ) -> PyResult> { let secrets = crate::secrets::source(py)?; + let config = crate::http::call_config(py, &kwargs, asynchronous)?; + let machine = messages_machine(crate::http::pool(), &config, secrets) + .map_err(crate::http::client_error)?; run_legacy_call( py, SURFACE, PublicCall::capture(&request, &args, &kwargs)?, - crate::logger::LoggedMachine::new(messages_machine(secrets)), + crate::logger::LoggedMachine::new(machine), MessagesPythonHost::new(request.unbind()), crate::preflight::sdk_preflight, asynchronous, diff --git a/litellm-rust/crates/python-bridge/src/secrets/callback.rs b/litellm-rust/crates/python-bridge/src/secrets/callback.rs index 82bb4443f98..6ba60630b3a 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/callback.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/callback.rs @@ -210,7 +210,7 @@ handler.get_secret_from_manager = get_secret_from_manager KeyManagementSettings::default(), )), Arc::new(move |_: &str| fallback.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::EnvironmentFallback); (resolver, locals, handler) diff --git a/litellm-rust/crates/python-bridge/src/secrets/mod.rs b/litellm-rust/crates/python-bridge/src/secrets/mod.rs index 439f9ddddd1..ed54c306397 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/mod.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/mod.rs @@ -27,7 +27,12 @@ const NATIVE: FieldSpec = FieldSpec::new("native", |field| field.schema_bo pub(crate) fn source(py: Python<'_>) -> PyResult> { if PythonSettings::SecretManager.read(py)?.read(&NATIVE)? { let context = litellm_host_python::PythonContext::capture(py)?; - return Ok(Arc::new(ResolvedSecrets::new(config::read(py)?, context))); + let client = crate::http::host_client(py, litellm_http::ClientVariant::Provider)?; + return Ok(Arc::new(ResolvedSecrets::new( + config::read(py)?, + context, + client, + ))); } Ok(Arc::new(PythonSecrets::new(py)?)) } diff --git a/litellm-rust/crates/python-bridge/src/secrets/resolved.rs b/litellm-rust/crates/python-bridge/src/secrets/resolved.rs index 40b187c99de..5a606ab1039 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/resolved.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/resolved.rs @@ -3,6 +3,7 @@ use std::sync::Arc; use futures_util::future::BoxFuture; use litellm_core_utils::settings::ProcessEnvironment; use litellm_host_python::PythonContext; +use litellm_http::Client; use litellm_secrets::source::SecretSource; use litellm_secrets::{ Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue, @@ -15,16 +16,20 @@ pub(crate) struct ResolvedSecrets { } impl ResolvedSecrets { - pub(crate) fn new(snapshot: SecretManagerSnapshot, context: PythonContext) -> Self { - Self::from_state(snapshot.into_state(context)) + pub(crate) fn new( + snapshot: SecretManagerSnapshot, + context: PythonContext, + client: Client, + ) -> Self { + Self::from_state(snapshot.into_state(context), client) } - fn from_state(state: Arc) -> Self { + fn from_state(state: Arc, client: Client) -> Self { Self { resolver: SecretResolver::new_python_compatible( state, Arc::new(ProcessEnvironment), - OidcResolver::default(), + OidcResolver::new(client), ) .with_failure_policy(FailurePolicy::EnvironmentFallback), } @@ -79,7 +84,7 @@ mod tests { } async fn resolve(state: Arc, name: &'static str) -> Option { - ResolvedSecrets::from_state(state) + ResolvedSecrets::from_state(state, litellm_http::Client::plain_for_test()) .resolve(&[name]) .await .unwrap() @@ -175,7 +180,10 @@ mod tests { .expect(1) .mount(&server) .await; - let source = ResolvedSecrets::from_state(state(&server, KeyManagementSettings::default())); + let source = ResolvedSecrets::from_state( + state(&server, KeyManagementSettings::default()), + litellm_http::Client::plain_for_test(), + ); let snapshot = source.resolve(&[declared]).await.unwrap(); assert_eq!(snapshot.get(undeclared), None); let result = source @@ -238,9 +246,12 @@ mod tests { #[tokio::test] async fn oidc_failures_are_not_converted_to_missing_secrets() { - let result = ResolvedSecrets::from_state(Arc::new(SecretManagerState::default())) - .resolve(&["oidc/"]) - .await; + let result = ResolvedSecrets::from_state( + Arc::new(SecretManagerState::default()), + litellm_http::Client::plain_for_test(), + ) + .resolve(&["oidc/"]) + .await; assert!(matches!(result, Err(litellm_secrets::Error::InvalidOidc))); } diff --git a/litellm-rust/crates/python-bridge/src/secrets/runtime.rs b/litellm-rust/crates/python-bridge/src/secrets/runtime.rs index 1a89130ee82..4d2e88115c8 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/runtime.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/runtime.rs @@ -10,6 +10,7 @@ use litellm_secrets_types::PythonSecretRead; use pyo3::{ exceptions::{PyAttributeError, PyRuntimeError, PyValueError}, prelude::*, + types::PyDict, }; #[derive(Clone, PartialEq)] @@ -44,10 +45,18 @@ impl NativeSecretManager { let system = configuration.system; let settings = configuration.settings.clone(); let enterprise_enabled = configuration.enterprise_enabled; + let http_config = crate::http::call_config(py, &PyDict::new(py), false)?; let backend = run_sync_value(py, async move { - load_native_manager(system, settings, environment, enterprise_enabled) - .await - .map_err(|error| PyValueError::new_err(error.to_string())) + load_native_manager( + crate::http::pool(), + &http_config, + system, + settings, + environment, + enterprise_enabled, + ) + .await + .map_err(|error| PyValueError::new_err(error.to_string())) })?; Ok(Self { backend, diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml index 7e8a79f89ef..efdf681e2bc 100644 --- a/litellm-rust/crates/secrets-azure/Cargo.toml +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true tokio.workspace = true litellm-auth-azure.workspace = true litellm-auth-types.workspace = true @@ -18,6 +19,7 @@ veil.workspace = true percent-encoding = "2.3" [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } wiremock = "0.6.5" rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/secrets-azure/src/key_vault.rs b/litellm-rust/crates/secrets-azure/src/key_vault.rs index 13e59f8e4ac..095c451927c 100644 --- a/litellm-rust/crates/secrets-azure/src/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/src/key_vault.rs @@ -19,7 +19,7 @@ const PATH_SEGMENT: &AsciiSet = &NON_ALPHANUMERIC #[derive(Clone)] pub struct AzureKeyVault { - client: reqwest::Client, + client: litellm_http::Client, vault: reqwest::Url, auth: Arc, inputs: Arc, @@ -33,7 +33,7 @@ struct SecretResponse { impl AzureKeyVault { pub fn with_client( - client: reqwest::Client, + client: litellm_http::Client, vault: reqwest::Url, environment: Arc, ) -> Result { @@ -57,7 +57,10 @@ impl AzureKeyVault { }) } - pub fn new(environment: Arc) -> Result { + pub fn new( + client: litellm_http::Client, + environment: Arc, + ) -> Result { let value = environment .get(AZURE_KEY_VAULT_URI) .ok_or(Error::MissingEnvironment(AZURE_KEY_VAULT_URI))?; @@ -65,7 +68,7 @@ impl AzureKeyVault { if vault.scheme() != "https" || vault.host_str().is_none() { return Err(Error::VaultUri); } - Self::with_client(reqwest::Client::new(), vault, environment) + Self::with_client(client, vault, environment) } pub fn scope(&self) -> &str { diff --git a/litellm-rust/crates/secrets-azure/tests/key_vault.rs b/litellm-rust/crates/secrets-azure/tests/key_vault.rs index a21149db345..fcbc46092e1 100644 --- a/litellm-rust/crates/secrets-azure/tests/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/tests/key_vault.rs @@ -11,7 +11,7 @@ use wiremock::{ fn manager(server: &MockServer) -> AzureKeyVault { AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), ) @@ -130,11 +130,14 @@ fn new_validates_vault_environment( #[case] uri: Option<&'static str>, #[case] missing_environment: bool, ) { - let result = AzureKeyVault::new(Arc::new(move |name: &str| { - (name == "AZURE_KEY_VAULT_URI") - .then(|| uri.map(str::to_owned)) - .flatten() - })); + let result = AzureKeyVault::new( + litellm_http::Client::plain_for_test(), + Arc::new(move |name: &str| { + (name == "AZURE_KEY_VAULT_URI") + .then(|| uri.map(str::to_owned)) + .flatten() + }), + ); if missing_environment { assert!(matches!( @@ -155,7 +158,7 @@ fn new_validates_vault_environment( #[case::local("http://localhost:8080", "https://localhost/.default")] fn derives_scope_from_vault_host(#[case] uri: &str, #[case] expected: &str) { let manager = AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), uri.parse().unwrap(), Arc::new(|_: &str| None), ) @@ -184,7 +187,7 @@ async fn missing_credentials_do_not_request_vault() { fn manager_without_credentials(server: &MockServer) -> AzureKeyVault { AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), Arc::new(|name: &str| { (name == "AZURE_CREDENTIAL").then(|| "ClientSecretCredential".to_owned()) diff --git a/litellm-rust/crates/secrets-azure/tests/live.rs b/litellm-rust/crates/secrets-azure/tests/live.rs index 18306382613..429cd3013f7 100644 --- a/litellm-rust/crates/secrets-azure/tests/live.rs +++ b/litellm-rust/crates/secrets-azure/tests/live.rs @@ -10,7 +10,7 @@ use rstest::rstest; #[ignore] async fn reads_a_real_secret() { let environment = Arc::new(ProcessEnvironment); - let manager = AzureKeyVault::new(environment).unwrap(); + let manager = AzureKeyVault::new(litellm_http::Client::plain_for_test(), environment).unwrap(); let name = std::env::var("AZURE_KEY_VAULT_LIVE_SECRET_NAME").unwrap(); let secret = manager.get_secret(&name).await.unwrap().unwrap(); assert!(matches!(&secret, Secret::String(_))); diff --git a/litellm-rust/crates/secrets-cyberark/Cargo.toml b/litellm-rust/crates/secrets-cyberark/Cargo.toml index 1c280171f4c..0a91c61ade9 100644 --- a/litellm-rust/crates/secrets-cyberark/Cargo.toml +++ b/litellm-rust/crates/secrets-cyberark/Cargo.toml @@ -8,6 +8,7 @@ repository.workspace = true [dependencies] litellm-secrets-types.workspace = true litellm-core-utils.workspace = true +litellm-http.workspace = true base64.workspace = true moka.workspace = true reqwest.workspace = true @@ -19,6 +20,7 @@ percent-encoding = "2.3" tokio = { workspace = true, features = ["sync"] } [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } rcgen = "0.14.10" rstest.workspace = true tempfile = "3.27.0" diff --git a/litellm-rust/crates/secrets-cyberark/src/error.rs b/litellm-rust/crates/secrets-cyberark/src/error.rs index 3dfeb95fe26..4f21647225b 100644 --- a/litellm-rust/crates/secrets-cyberark/src/error.rs +++ b/litellm-rust/crates/secrets-cyberark/src/error.rs @@ -18,6 +18,8 @@ pub enum Error { MissingCredentials, #[error("CyberArk client certificate could not be loaded")] ClientCertificate, + #[error("CyberArk Conjur HTTP client could not be built")] + Client(#[redact] Box), #[error("invalid refresh interval")] RefreshInterval, #[error("invalid CyberArk Conjur endpoint")] diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs index 252a99c917f..37d547ff6c1 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs @@ -2,10 +2,13 @@ mod client; mod read; mod write; -use std::{fs, sync::Arc, time::Duration}; +use std::{sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_core_utils::settings::Lookup; +use litellm_http::{ + Client, ClientIdentity, ClientVariant, HttpClientConfig, HttpClientPool, TlsSource, Verify, +}; use litellm_secrets_types::{ BaseSecretManager, CyberarkOperationContext, RotationError, SecretCache, SecretDeleter, SecretRotator, SecretValue, SecretWriteContext, SecretWriter, async_rotate_secret, @@ -37,7 +40,7 @@ const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC #[derive(Clone)] pub struct CyberArkSecretManager { - client: reqwest::Client, + client: Client, endpoint: reqwest::Url, account: String, username: String, diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs index 1d99fe474d5..052d5570896 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs @@ -2,7 +2,7 @@ use super::*; impl CyberArkSecretManager { pub fn with_client( - client: reqwest::Client, + client: Client, endpoint: reqwest::Url, account: String, username: String, @@ -30,6 +30,8 @@ impl CyberArkSecretManager { } pub fn new( + pool: &HttpClientPool, + config: &HttpClientConfig, environment: Arc, enterprise_enabled: bool, ) -> Result { @@ -46,21 +48,34 @@ impl CyberArkSecretManager { .get(CYBERARK_SSL_VERIFY) .map(|value| !value.trim().eq_ignore_ascii_case("false")) .unwrap_or(true); - let mut builder = reqwest::Client::builder(); if !verify { litellm_tracing::warn!( "CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates." ); - builder = builder.danger_accept_invalid_certs(true); } - if !cert.is_empty() && !key.is_empty() { - let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?; - let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?; - let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat()) - .map_err(|_| Error::ClientCertificate)?; - builder = builder.identity(identity); - } - let client = builder.build()?; + let config = HttpClientConfig { + verify: effective_verify(verify, &config.verify), + client_certificate: (!cert.is_empty() && !key.is_empty()).then(|| { + ClientIdentity::Split { + certificate: cert.into(), + key: key.into(), + } + }), + ..config.clone() + }; + let client = + pool.client(&config, ClientVariant::Provider) + .map_err(|error| match error { + litellm_http::Error::Read { + tls_source: TlsSource::ClientIdentity, + .. + } + | litellm_http::Error::InvalidPem { + tls_source: TlsSource::ClientIdentity, + .. + } => Error::ClientCertificate, + other => Error::Client(Box::new(other)), + })?; let endpoint = reqwest::Url::parse( &environment .get(CYBERARK_API_BASE) @@ -139,9 +154,35 @@ impl CyberArkSecretManager { } } +fn effective_verify(cyberark_verify: bool, host: &Verify) -> Verify { + match (cyberark_verify, host) { + (false, _) => Verify::Disabled, + (true, Verify::Disabled) => Verify::BuiltInRoots, + (true, host) => host.clone(), + } +} + fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url { if !endpoint.path().ends_with('/') { endpoint.set_path(&format!("{}/", endpoint.path())); } endpoint } + +#[cfg(test)] +mod tests { + use std::path::PathBuf; + + use super::*; + + #[test] + fn cyberark_verification_does_not_follow_a_host_that_disabled_it() { + let bundle = Verify::CaBundle(PathBuf::from("/ca.pem")); + assert_eq!( + effective_verify(true, &Verify::Disabled), + Verify::BuiltInRoots + ); + assert_eq!(effective_verify(true, &bundle), bundle); + assert_eq!(effective_verify(false, &bundle), Verify::Disabled); + } +} diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs index 2048e067b6e..783ba6ff67a 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs @@ -7,6 +7,8 @@ use std::{ }; use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_core_utils::settings::Lookup; +use litellm_http::{HttpClientPool, HttpSettings, Resolution, media::PublicDnsResolver}; use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error}; use litellm_secrets_types::{BaseSecretManager, CyberarkOperationContext, SecretValue}; use rstest::{fixture, rstest}; diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs index fbd4317f446..1bea77e49ae 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs @@ -77,7 +77,7 @@ async fn authentication_encodes_login(#[case] username: &str, #[case] expected_p .mount(&server) .await; let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), username.into(), @@ -211,25 +211,25 @@ fn new_validates_credentials_before_license_and_configuration() { let empty: Arc = Arc::new(|_: &str| None); assert!(matches!( - CyberArkSecretManager::new(empty, true), + from_environment(empty, true), Err(Error::MissingCredentials) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())), false ), Err(Error::EnterpriseRequired) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())), true ), Err(Error::MissingCredentials) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| match name { "CYBERARK_API_KEY" => Some("k3y".into()), "CYBERARK_REFRESH_INTERVAL" => Some("abc".into()), @@ -240,7 +240,7 @@ fn new_validates_credentials_before_license_and_configuration() { Err(Error::RefreshInterval) )); assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| match name { "CYBERARK_API_KEY" => Some("k3y".into()), "CYBERARK_API_BASE" => Some("not a url".into()), @@ -254,7 +254,7 @@ fn new_validates_credentials_before_license_and_configuration() { #[rstest] fn certificate_only_credentials_are_validated_as_a_client_identity() { - let result = CyberArkSecretManager::new( + let result = from_environment( Arc::new(|name: &str| match name { "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), "CYBERARK_CLIENT_KEY" => Some("/missing/key".into()), @@ -295,7 +295,7 @@ async fn configured_client_identity_preserves_auth_request_and_read_result( let endpoint = server.uri(); let certificate = client_identity_directory.path().join("client.crt"); let key = client_identity_directory.path().join("client.key"); - let manager = CyberArkSecretManager::new( + let manager = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_BASE" => Some(endpoint.clone()), "CYBERARK_API_KEY" => Some(api_key.into()), @@ -337,7 +337,7 @@ fn invalid_client_identity_is_not_ignored( let certificate = client_identity_directory.path().join("client.crt"); let key = client_identity_directory.path().join("client.key"); - let result = CyberArkSecretManager::new( + let result = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_KEY" => Some(api_key.into()), "CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()), @@ -354,7 +354,7 @@ fn invalid_client_identity_is_not_ignored( #[case::certificate_only("")] #[case::certificate_and_api_key("k3y")] fn client_identity_does_not_bypass_the_enterprise_requirement(#[case] api_key: &'static str) { - let result = CyberArkSecretManager::new( + let result = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_KEY" => Some(api_key.into()), "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), @@ -381,7 +381,7 @@ async fn new_reads_environment_defaults_end_to_end() { .mount(&server) .await; let endpoint = server.uri(); - let manager = CyberArkSecretManager::new( + let manager = from_environment( Arc::new(move |name: &str| match name { "CYBERARK_API_BASE" => Some(endpoint.clone()), "CYBERARK_API_KEY" => Some("k3y".into()), @@ -404,7 +404,7 @@ async fn new_reads_environment_defaults_end_to_end() { #[rstest] fn new_reports_missing_client_certificate_files() { assert!(matches!( - CyberArkSecretManager::new( + from_environment( Arc::new(|name: &str| match name { "CYBERARK_API_KEY" => Some("k3y".into()), "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), @@ -433,7 +433,7 @@ async fn trailing_slash_endpoint_preserves_base_path() { .await; let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap(); let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), endpoint, "acct".into(), "admin".into(), diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs index f5ba7a63273..6fa15ecba9b 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs @@ -48,9 +48,21 @@ pub(super) fn client_identity_directory() -> tempfile::TempDir { directory } +pub(super) fn from_environment( + environment: Arc, + enterprise_enabled: bool, +) -> Result { + CyberArkSecretManager::new( + &HttpClientPool::new(Arc::new(PublicDnsResolver)), + &Resolution::from(&HttpSettings::default()).config, + environment, + enterprise_enabled, + ) +} + pub(super) fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager { CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), "admin".into(), diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs index 331a26c6119..e9a027091fa 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs @@ -166,7 +166,7 @@ async fn writes_match_python_parity_fixture(parity_fixture: ParityFixture) { .mount(&server) .await; let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), parity_fixture.account, parity_fixture.username, @@ -230,7 +230,7 @@ async fn live_conjur_round_trip() { .as_nanos() ); let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), endpoint.clone(), account.clone(), username.clone(), @@ -245,7 +245,7 @@ async fn live_conjur_round_trip() { .await .unwrap(); let verifier = CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), endpoint.clone(), account.clone(), username.clone(), diff --git a/litellm-rust/crates/secrets-google/Cargo.toml b/litellm-rust/crates/secrets-google/Cargo.toml index 805eb80740d..208b5ddd03f 100644 --- a/litellm-rust/crates/secrets-google/Cargo.toml +++ b/litellm-rust/crates/secrets-google/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +litellm-http.workspace = true moka.workspace = true tokio.workspace = true litellm-auth-gcp = { workspace = true, features = ["google-sdk"] } @@ -24,6 +25,7 @@ serde.workspace = true reqwest.workspace = true [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } google-cloud-auth.workspace = true rstest.workspace = true wiremock = "0.6.5" diff --git a/litellm-rust/crates/secrets-google/src/secret_manager.rs b/litellm-rust/crates/secrets-google/src/secret_manager.rs index b8787999e12..08ba466b799 100644 --- a/litellm-rust/crates/secrets-google/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/src/secret_manager.rs @@ -23,7 +23,7 @@ const CACHE_CAPACITY: u64 = 200; #[derive(Clone)] pub struct GoogleSecretManager { - client: reqwest::Client, + client: litellm_http::Client, credentials: Arc, endpoint: reqwest::Url, project: String, @@ -46,7 +46,7 @@ struct Payload { impl GoogleSecretManager { pub fn with_client( - client: reqwest::Client, + client: litellm_http::Client, endpoint: reqwest::Url, project: String, environment: Arc, @@ -79,6 +79,7 @@ impl GoogleSecretManager { } pub fn new( + client: litellm_http::Client, environment: Arc, enterprise_enabled: bool, ) -> Result { @@ -104,7 +105,7 @@ impl GoogleSecretManager { .get(GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER) .is_some_and(|v| v.eq_ignore_ascii_case("true")); Self::with_client( - reqwest::Client::new(), + client, reqwest::Url::parse("https://secretmanager.googleapis.com").expect("static URL"), project, environment, diff --git a/litellm-rust/crates/secrets-google/tests/secret_manager.rs b/litellm-rust/crates/secrets-google/tests/secret_manager.rs index 0d7efc4b1b3..e9bee7633f5 100644 --- a/litellm-rust/crates/secrets-google/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/tests/secret_manager.rs @@ -11,7 +11,7 @@ use wiremock::{ fn manager(server: &MockServer, always_read: bool, ttl: Duration) -> GoogleSecretManager { GoogleSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "project".into(), Arc::new(|name: &str| (name == "VERTEX_AI_API_KEY").then(|| "token".into())), @@ -214,11 +214,19 @@ async fn always_read_and_expired_cache_fetch_again( #[rstest] fn google_manager_requires_host_license_and_project_configuration() { assert!(matches!( - GoogleSecretManager::new(Arc::new(|_: &str| None), false), + GoogleSecretManager::new( + litellm_http::Client::plain_for_test(), + Arc::new(|_: &str| None), + false + ), Err(Error::EnterpriseRequired) )); assert!(matches!( - GoogleSecretManager::new(Arc::new(|_: &str| None), true), + GoogleSecretManager::new( + litellm_http::Client::plain_for_test(), + Arc::new(|_: &str| None), + true + ), Err(Error::MissingEnvironment( "GOOGLE_SECRET_MANAGER_PROJECT_ID" )) @@ -236,7 +244,7 @@ fn google_manager_rejects_invalid_refresh_intervals(#[case] variable: &'static s }); assert!(matches!( - GoogleSecretManager::new(environment, true), + GoogleSecretManager::new(litellm_http::Client::plain_for_test(), environment, true), Err(Error::RefreshInterval) )); } diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index 17acc01682b..fe61d6cb5b4 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -23,6 +23,7 @@ litellm-secrets-hashicorp = { workspace = true, optional = true } litellm-secrets-azure = { workspace = true, optional = true } litellm-secrets-cyberark = { workspace = true, optional = true } litellm-core-utils.workspace = true +litellm-http.workspace = true base64.workspace = true serde.workspace = true strum.workspace = true @@ -33,6 +34,7 @@ moka.workspace = true tokio = { workspace = true, features = ["fs"] } [dev-dependencies] +litellm-http = { workspace = true, features = ["test-support"] } rstest.workspace = true wiremock = "0.6.5" tempfile = "3" diff --git a/litellm-rust/crates/secrets/src/error.rs b/litellm-rust/crates/secrets/src/error.rs index 07f2f205bec..7fc756c5c7d 100644 --- a/litellm-rust/crates/secrets/src/error.rs +++ b/litellm-rust/crates/secrets/src/error.rs @@ -10,6 +10,8 @@ pub enum Error { InvalidCiphertext, #[error("decrypted value is not UTF-8")] Utf8, + #[error(transparent)] + Client(#[from] litellm_http::Error), #[error("unsupported OIDC provider or missing build feature")] UnsupportedOidc, #[error("OIDC reference requires a provider and audience")] diff --git a/litellm-rust/crates/secrets/src/native.rs b/litellm-rust/crates/secrets/src/native.rs index 80f0e46245c..f1dc7ccb732 100644 --- a/litellm-rust/crates/secrets/src/native.rs +++ b/litellm-rust/crates/secrets/src/native.rs @@ -1,10 +1,13 @@ use std::sync::Arc; use litellm_core_utils::settings::Lookup; +use litellm_http::{HttpClientConfig, HttpClientPool}; use crate::{Error, KeyManagementSettings, KeyManagementSystem, SecretManager}; pub async fn load_native_manager( + pool: &HttpClientPool, + config: &HttpClientConfig, system: KeyManagementSystem, settings: KeyManagementSettings, environment: Arc, @@ -29,14 +32,19 @@ pub async fn load_native_manager( } #[cfg(feature = "azure")] (KeyManagementSystem::AzureKeyVault, _, environment, _) => Ok( - SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new(environment)?), + SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new( + pool.client(config, litellm_http::ClientVariant::Provider)?, + environment, + )?), ), #[cfg(feature = "google")] - (KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => { - Ok(SecretManager::GoogleSecretManager( - crate::google::GoogleSecretManager::new(environment, enterprise_enabled)?, - )) - } + (KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => Ok( + SecretManager::GoogleSecretManager(crate::google::GoogleSecretManager::new( + pool.client(config, litellm_http::ClientVariant::Provider)?, + environment, + enterprise_enabled, + )?), + ), #[cfg(feature = "google")] (KeyManagementSystem::GoogleKms, _, environment, _) => { crate::google::load_google_kms(Some(true), environment) @@ -51,11 +59,14 @@ pub async fn load_native_manager( )) } #[cfg(feature = "cyberark")] - (KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => { - Ok(SecretManager::Cyberark( - crate::cyberark::CyberArkSecretManager::new(environment, enterprise_enabled)?, - )) - } + (KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => Ok( + SecretManager::Cyberark(crate::cyberark::CyberArkSecretManager::new( + pool, + config, + environment, + enterprise_enabled, + )?), + ), _ => Err(Error::NativeBackendUnavailable), } } diff --git a/litellm-rust/crates/secrets/src/oidc.rs b/litellm-rust/crates/secrets/src/oidc.rs index f3c1e38ce7b..b6e8dbc123b 100644 --- a/litellm-rust/crates/secrets/src/oidc.rs +++ b/litellm-rust/crates/secrets/src/oidc.rs @@ -5,6 +5,7 @@ use std::{ use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; use litellm_core_utils::settings::Lookup; +use litellm_http::Client; use moka::future::Cache; use serde::Deserialize; @@ -82,7 +83,7 @@ impl NumericDate { } pub struct OidcResolver { - client: reqwest::Client, + client: Client, google_identity_endpoint: reqwest::Url, cache: Cache, clock: fn() -> SystemTime, @@ -90,25 +91,17 @@ pub struct OidcResolver { azure_token_provider: std::sync::Arc, } -impl Default for OidcResolver { - fn default() -> Self { - let client = reqwest::Client::builder() - .timeout(Duration::from_secs(600)) - .connect_timeout(Duration::from_secs(5)) - .build() - .expect("HTTP client configuration"); - Self::new( - client, - reqwest::Url::parse("http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity").expect("static URL"), - ) - } -} +const GOOGLE_IDENTITY_ENDPOINT: &str = + "http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity"; + +const REQUEST_TIMEOUT: Duration = Duration::from_secs(600); impl OidcResolver { - pub fn new(client: reqwest::Client, google_identity_endpoint: reqwest::Url) -> Self { + pub fn new(client: Client) -> Self { Self { client, - google_identity_endpoint, + google_identity_endpoint: reqwest::Url::parse(GOOGLE_IDENTITY_ENDPOINT) + .expect("static URL"), cache: Cache::builder() .max_capacity(200) .time_to_live(GOOGLE_TOKEN_MAX_TTL) @@ -121,6 +114,13 @@ impl OidcResolver { } } + pub fn with_google_identity_endpoint(self, google_identity_endpoint: reqwest::Url) -> Self { + Self { + google_identity_endpoint, + ..self + } + } + #[cfg(feature = "azure")] pub fn with_azure_token_provider( self, @@ -180,6 +180,7 @@ impl OidcResolver { let response = self .client .get(url) + .timeout(REQUEST_TIMEOUT) .query(&[("audience", audience)]) .bearer_auth(authorization) .header("Accept", "application/json; api-version=2.0") @@ -214,6 +215,7 @@ impl OidcResolver { let response = self .client .get(self.google_identity_endpoint.clone()) + .timeout(REQUEST_TIMEOUT) .query(&[("audience", audience)]) .header("Metadata-Flavor", "Google") .send() diff --git a/litellm-rust/crates/secrets/src/resolver.rs b/litellm-rust/crates/secrets/src/resolver.rs index 69445e5b410..830e7c22cd5 100644 --- a/litellm-rust/crates/secrets/src/resolver.rs +++ b/litellm-rust/crates/secrets/src/resolver.rs @@ -1,10 +1,7 @@ use std::sync::Arc; use crate::compatibility::python_manager_string; -use litellm_core_utils::{ - serde_compat::parse_str_bool, - settings::{Lookup, ProcessEnvironment}, -}; +use litellm_core_utils::{serde_compat::parse_str_bool, settings::Lookup}; use crate::state::{LookupTarget, normalize_secret_name}; use crate::{Error, OidcResolver, Secret, SecretManagerState, SecretValue}; @@ -24,16 +21,6 @@ pub struct SecretResolver { python_compatible: bool, } -impl Default for SecretResolver { - fn default() -> Self { - Self::new( - Arc::new(SecretManagerState::default()), - Arc::new(ProcessEnvironment), - OidcResolver::default(), - ) - } -} - impl SecretResolver { pub fn new( state: Arc, diff --git a/litellm-rust/crates/secrets/src/source.rs b/litellm-rust/crates/secrets/src/source.rs index a1615bad055..3c86fd25e9a 100644 --- a/litellm-rust/crates/secrets/src/source.rs +++ b/litellm-rust/crates/secrets/src/source.rs @@ -37,15 +37,14 @@ impl SecretSource for SecretResolver { } } -#[derive(Default)] pub struct EnvironmentSecrets(SecretResolver); impl EnvironmentSecrets { - pub fn python_compatible() -> Self { + pub fn python_compatible(client: litellm_http::Client) -> Self { Self(SecretResolver::new_python_compatible( Arc::new(crate::SecretManagerState::default()), Arc::new(litellm_core_utils::settings::ProcessEnvironment), - crate::OidcResolver::default(), + crate::OidcResolver::new(client), )) } } diff --git a/litellm-rust/crates/secrets/tests/aws.rs b/litellm-rust/crates/secrets/tests/aws.rs index 174d0881339..910b1855b77 100644 --- a/litellm-rust/crates/secrets/tests/aws.rs +++ b/litellm-rust/crates/secrets/tests/aws.rs @@ -53,7 +53,7 @@ async fn read_results_follow_the_selected_failure_policy( }, )), Arc::new(move |_: &str| environment.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(policy); let result = resolver @@ -97,7 +97,7 @@ async fn primary_secret_values_other_than_strings_resolve_to_none( let resolver = SecretResolver::new_python_compatible( Arc::new(state(&server, settings)), Arc::new(|_: &str| Some("fallback".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let text = value.as_str(); assert_eq!( @@ -157,7 +157,7 @@ async fn gating_prediction_matches_actual_lookup( let resolver = SecretResolver::new_python_compatible( Arc::new(state), Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver diff --git a/litellm-rust/crates/secrets/tests/azure.rs b/litellm-rust/crates/secrets/tests/azure.rs index b844b198cd4..60f165add54 100644 --- a/litellm-rust/crates/secrets/tests/azure.rs +++ b/litellm-rust/crates/secrets/tests/azure.rs @@ -22,7 +22,7 @@ async fn azure_handler_reads_missing_and_failed_secrets() { .await; let manager = SecretManager::AzureKeyVault( AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), std::sync::Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), ) @@ -81,7 +81,7 @@ async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_ .mount(&server) .await; let manager = AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "token".into())), ) @@ -92,7 +92,7 @@ async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_ Default::default(), )), Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver diff --git a/litellm-rust/crates/secrets/tests/common_read_contract.rs b/litellm-rust/crates/secrets/tests/common_read_contract.rs index 698e1ad8f63..2e8fbf46008 100644 --- a/litellm-rust/crates/secrets/tests/common_read_contract.rs +++ b/litellm-rust/crates/secrets/tests/common_read_contract.rs @@ -59,7 +59,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager { ), Provider::Azure => SecretManager::AzureKeyVault( AzureKeyVault::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), environment, ) @@ -67,7 +67,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager { ), Provider::Google => SecretManager::GoogleSecretManager( GoogleSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "project".into(), environment, @@ -84,7 +84,7 @@ fn manager(provider: Provider, server: &MockServer) -> SecretManager { .unwrap(), ), Provider::Cyberark => SecretManager::Cyberark(CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), "admin".into(), @@ -240,7 +240,7 @@ async fn python_read_failures_preserve_provider_fallback_rules( KeyManagementSettings::default(), )), Arc::new(move |_: &str| environment_value.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let expected = if matches!(provider, Provider::Aws) { None diff --git a/litellm-rust/crates/secrets/tests/cyberark.rs b/litellm-rust/crates/secrets/tests/cyberark.rs index 706c35752d7..b94bd9dc0ad 100644 --- a/litellm-rust/crates/secrets/tests/cyberark.rs +++ b/litellm-rust/crates/secrets/tests/cyberark.rs @@ -24,7 +24,7 @@ async fn cyberark_handler_reads_values_and_surfaces_errors() { .mount(&server) .await; let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "acct".into(), "admin".into(), diff --git a/litellm-rust/crates/secrets/tests/google.rs b/litellm-rust/crates/secrets/tests/google.rs index 67954fedb78..0b45a90c10b 100644 --- a/litellm-rust/crates/secrets/tests/google.rs +++ b/litellm-rust/crates/secrets/tests/google.rs @@ -25,7 +25,7 @@ async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) _ => None, }); let manager = GoogleSecretManager::with_client( - reqwest::Client::new(), + litellm_http::Client::plain_for_test(), server.uri().parse().unwrap(), "project".into(), environment.clone(), @@ -40,7 +40,7 @@ async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) let resolver = SecretResolver::new_python_compatible( Arc::new(state), environment, - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::Propagate); let result = resolver.get_secret_str("KEY", None).await; diff --git a/litellm-rust/crates/secrets/tests/hashicorp.rs b/litellm-rust/crates/secrets/tests/hashicorp.rs index bc35b88018e..e10e903c25a 100644 --- a/litellm-rust/crates/secrets/tests/hashicorp.rs +++ b/litellm-rust/crates/secrets/tests/hashicorp.rs @@ -56,7 +56,7 @@ async fn hashicorp_handler_resolves_found_missing_and_failed_values() { }, )), Arc::new(|_: &str| None), - litellm_secrets::OidcResolver::default(), + litellm_secrets::OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( found_resolver @@ -131,7 +131,7 @@ async fn hashicorp_handler_resolves_found_missing_and_failed_values() { let failed_resolver = SecretResolver::new_python_compatible( Arc::new(failed_state), Arc::new(|_: &str| None), - litellm_secrets::OidcResolver::default(), + litellm_secrets::OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::Propagate); assert!(matches!( diff --git a/litellm-rust/crates/secrets/tests/oidc.rs b/litellm-rust/crates/secrets/tests/oidc.rs index afc49e8231d..6f72bd9d645 100644 --- a/litellm-rust/crates/secrets/tests/oidc.rs +++ b/litellm-rust/crates/secrets/tests/oidc.rs @@ -30,7 +30,7 @@ async fn environment_sources_resolve_expected_value( ("CIRCLE_OIDC_TOKEN_V2", "circle-v2"), ]); assert_eq!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve(reference, env.as_ref()) .await .unwrap() @@ -43,7 +43,7 @@ async fn environment_sources_resolve_expected_value( #[tokio::test] async fn environment_sources_bypass_boolean_conversion_and_defaults() { let env = environment(&[("TOKEN", "true")]); - let oidc = OidcResolver::default(); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()); let resolver = SecretResolver::new(Arc::new(SecretManagerState::default()), env, oidc); assert_eq!( resolver @@ -94,7 +94,7 @@ async fn github_requests_are_authenticated_cached_and_revalidate_environment() { ), ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token"), ]); - let oidc = OidcResolver::default(); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()); for _ in 0..2 { assert_eq!( oidc.resolve("oidc/github/https://service/oidc/path", env.as_ref()) @@ -131,7 +131,7 @@ async fn file_allowlist_resolves_symlinks_while_environment_paths_remain_explici ("PATH_TOKEN", private.to_str().unwrap()), ("AZURE_FEDERATED_TOKEN_FILE", token.to_str().unwrap()), ]); - let oidc = OidcResolver::default(); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()); assert_eq!( oidc.resolve(&format!("oidc/file/{}", token.display()), env.as_ref()) .await @@ -213,8 +213,9 @@ async fn google_expiry_caps_cache_and_preserves_audience( .expect(calls) .mount(&server) .await; - let oidc = - OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()).with_clock(now); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()) + .with_clock(now); for _ in 0..2 { assert_eq!( oidc.resolve( @@ -234,7 +235,7 @@ async fn google_expiry_caps_cache_and_preserves_audience( #[tokio::test] async fn google_oidc_requires_its_build_feature() { assert!(matches!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve("oidc/google/audience", environment(&[]).as_ref()) .await, Err(Error::UnsupportedOidc) @@ -245,7 +246,7 @@ async fn google_oidc_requires_its_build_feature() { #[tokio::test] async fn azure_oidc_without_a_token_file_requires_its_build_feature() { assert!(matches!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve("oidc/azure/scope", environment(&[]).as_ref()) .await, Err(Error::UnsupportedOidc) @@ -261,7 +262,7 @@ async fn invalid_references_fail_before_environment_lookup( #[case] reference: &str, #[case] unsupported: bool, ) { - let error = OidcResolver::default() + let error = OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve(reference, &|_: &str| { panic!("invalid reference reached environment lookup") }) @@ -283,7 +284,8 @@ async fn unreadable_expiry_keeps_python_cache_fallback(#[case] token: &str) { .expect(1) .mount(&server) .await; - let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()); + let resolver = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()); for _ in 0..2 { assert_eq!( resolver @@ -334,7 +336,8 @@ async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case] }) } } - let oidc = OidcResolver::default().with_azure_token_provider(Arc::new(Provider(failed))); + let oidc = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_azure_token_provider(Arc::new(Provider(failed))); let resolver = SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), environment(&[("AZURE_CLIENT_ID", "client-id")]), @@ -361,7 +364,7 @@ async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case] #[tokio::test] async fn missing_oidc_environment_is_an_error(#[case] reference: &str) { assert!(matches!( - OidcResolver::default() + OidcResolver::new(litellm_http::Client::plain_for_test()) .resolve(reference, environment(&[]).as_ref()) .await, Err(Error::MissingEnvironment) @@ -380,7 +383,8 @@ async fn google_oidc_failures_are_not_cached_or_hidden_by_defaults() { let resolver = SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), environment(&[]), - OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()), + OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()), ); for _ in 0..2 { assert!(matches!( @@ -420,7 +424,8 @@ async fn google_tokens_expire_at_the_python_cache_deadline( .expect(2) .mount(&server) .await; - let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()) + let resolver = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()) .with_clock(|| UNIX_EPOCH + Duration::from_secs(1000)); assert_eq!( resolver @@ -472,7 +477,8 @@ async fn google_cache_uses_payload_expiry_without_requiring_a_jwt_header() { .expect(2) .mount(&server) .await; - let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()) + let resolver = OidcResolver::new(litellm_http::Client::plain_for_test()) + .with_google_identity_endpoint(server.uri().parse().unwrap()) .with_clock(|| UNIX_EPOCH + Duration::from_secs(1000)); for _ in 0..2 { assert_eq!( diff --git a/litellm-rust/crates/secrets/tests/resolution.rs b/litellm-rust/crates/secrets/tests/resolution.rs index bed762adc59..37557c36fb5 100644 --- a/litellm-rust/crates/secrets/tests/resolution.rs +++ b/litellm-rust/crates/secrets/tests/resolution.rs @@ -12,7 +12,7 @@ fn resolver(value: Option<&str>) -> SecretResolver { SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), Arc::new(move |_: &str| value.clone()), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) } @@ -35,7 +35,7 @@ async fn native_reads_preserve_strings_and_report_conversion_errors(#[case] mana let resolver = SecretResolver::new( Arc::new(state), Arc::new(move |_: &str| Some(raw.to_owned())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver @@ -75,7 +75,7 @@ async fn native_defaults_apply_to_absence_but_never_hide_provider_failures() { KeyManagementSettings::default(), )), Arc::new(|_: &str| None), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let result = resolver .get_secret_str("key", Some(SecretValue::new("default"))) @@ -189,7 +189,7 @@ fn managed(reply: Result, ()>, environment: Option<&'static str>) KeyManagementSettings::default(), )), Arc::new(move |_: &str| environment.map(str::to_owned)), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ) .with_failure_policy(FailurePolicy::EnvironmentFallback) } @@ -281,7 +281,7 @@ async fn prefix_is_removed_once_and_resolved_from_environment() { let resolver = SecretResolver::new_python_compatible( Arc::new(state), Arc::new(|name: &str| (name == "os.environ/KEY").then(|| "value".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver @@ -340,7 +340,7 @@ async fn excluded_hosted_keys_keep_the_python_manager_conversion_path( }, )), Arc::new(move |_: &str| Some(raw.to_owned())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver.get_secret("KEY", None).await.unwrap(), @@ -373,7 +373,7 @@ async fn azure_callback_absence_preserves_none_but_errors_fall_back( KeyManagementSettings::default(), )), Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); assert_eq!( resolver diff --git a/litellm-rust/crates/secrets/tests/source.rs b/litellm-rust/crates/secrets/tests/source.rs index b4782c6af86..17210940a67 100644 --- a/litellm-rust/crates/secrets/tests/source.rs +++ b/litellm-rust/crates/secrets/tests/source.rs @@ -15,7 +15,7 @@ mod tests { #[case] expected: Option<&str>, ) { unsafe { std::env::set_var(name, value) }; - let secret = EnvironmentSecrets::python_compatible() + let secret = EnvironmentSecrets::python_compatible(litellm_http::Client::plain_for_test()) .resolve(&[name]) .await .unwrap() @@ -42,7 +42,7 @@ async fn dynamic_names_use_the_same_resolver_and_snapshots_never_do_fresh_lookup reads.fetch_add(1, Ordering::SeqCst); (name != "missing").then(|| name.to_owned()) }), - OidcResolver::default(), + OidcResolver::new(litellm_http::Client::plain_for_test()), ); let snapshot = source.resolve(&["declared", "missing"]).await.unwrap(); let name = format!("runtime-{}", "key"); diff --git a/litellm-rust/crates/testkit/src/lib.rs b/litellm-rust/crates/testkit/src/lib.rs index 9ea6123a176..423adeec426 100644 --- a/litellm-rust/crates/testkit/src/lib.rs +++ b/litellm-rust/crates/testkit/src/lib.rs @@ -1,3 +1,9 @@ +#![allow( + clippy::disallowed_types, + clippy::disallowed_methods, + reason = "a dev-only installer tool that never talks to providers" +)] + mod agent; mod error; mod install; From 7fc22061715487220826135298231fc723e8993a Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 25 Sep 2026 19:10:20 -0700 Subject: [PATCH 04/10] test: fix stale and state-leaking tests red on scheduled CircleCI (#43266) test_update_config_success_callback_normalization replaced proxy_server.proxy_logging_obj with a MagicMock and never restored it. Since the proxy unit tests joined tests/unit (#42903), 14 JWT mapping, end-user and MCP tests on the same xdist worker awaited that mock and failed. The test now uses monkeypatch. test_prometheus_logging_callbacks set verbose_logger to DEBUG and litellm.set_verbose at import, so every worker in the unit job ran with DEBUG on. That broke caplog equality in the JEV classifier test, the vertex streaming memory ratio, and four event-loop lag checks. The module-level setup is removed; nothing in the file depended on it. #43081 removed the OCR harness modules but left them in the importability parametrize list. test_get_model_info_bedrock_region reassigned litellm.model_cost and set LITELLM_LOCAL_MODEL_COST_MAP without restoring either, and never cleared the get_model_info caches, so it failed whenever an earlier test had looked up the regional model. It now uses monkeypatch and invalidates the caches; the local_testing isolation fixture also invalidates them after restoring model_cost. The Windows job hit CircleCI's 10 minute no-output limit while cargo compiles the Rust crates inside uv sync and uv build. Those two steps now allow 30 minutes of silence. --- .circleci/config.yml | 2 ++ tests/local_testing/conftest.py | 2 ++ tests/local_testing/test_get_model_info.py | 17 +++++++++-------- tests/test_rust_python_harness.py | 2 -- .../test_prometheus_logging_callbacks.py | 6 ------ tests/unit/proxy/test_proxy_server.py | 8 ++++---- 6 files changed, 17 insertions(+), 20 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index 370424dca86..3ff8061fb48 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -323,6 +323,7 @@ jobs: CHOCOLATEY_CONFIRM_ALL: "true" - run: name: Install Dependencies + no_output_timeout: 30m environment: UV_HTTP_TIMEOUT: "300" command: | @@ -381,6 +382,7 @@ jobs: uv run --no-sync python -m pytest tests/windows_tests/ -v - run: name: Guard against MAX_PATH-busting packaged wheel paths + no_output_timeout: 30m environment: UV_HTTP_TIMEOUT: "300" command: | diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index d03f074f557..df3dacac3b2 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -19,6 +19,7 @@ import pytest import litellm from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER +from litellm.utils import _invalidate_model_cost_lowercase_map # ``litellm.model_cost`` is loaded at import time from the URL pinned to ``main`` # (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with this branch @@ -232,6 +233,7 @@ def isolate_litellm_state(): for attr, original_value in original_state.items(): if hasattr(litellm, attr): setattr(litellm, attr, original_value) + _invalidate_model_cost_lowercase_map() @pytest.fixture(scope="module", autouse=True) diff --git a/tests/local_testing/test_get_model_info.py b/tests/local_testing/test_get_model_info.py index 1e46a1bf853..dbe1fc3b69b 100644 --- a/tests/local_testing/test_get_model_info.py +++ b/tests/local_testing/test_get_model_info.py @@ -9,6 +9,7 @@ import pytest import litellm from litellm import get_model_info +from litellm.utils import _invalidate_model_cost_lowercase_map from unittest.mock import MagicMock, patch @@ -74,15 +75,15 @@ def test_get_model_info_ollama_chat(): assert mock_client.call_args.kwargs["json"]["name"] == "unknown-model" -def test_get_model_info_bedrock_region(): - os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" - litellm.model_cost = litellm.get_model_cost_map(url="") - args = { - "model": "us.anthropic.claude-haiku-4-5-20251001-v1:0", - "custom_llm_provider": "bedrock", +def test_get_model_info_bedrock_region(monkeypatch): + regional_model = "us.anthropic.claude-haiku-4-5-20251001-v1:0" + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + model_cost_without_regional_entry = { + key: value for key, value in litellm.get_model_cost_map(url="").items() if key != regional_model } - litellm.model_cost.pop("us.anthropic.claude-haiku-4-5-20251001-v1:0", None) - info = litellm.get_model_info(**args) + monkeypatch.setattr(litellm, "model_cost", model_cost_without_regional_entry) + _invalidate_model_cost_lowercase_map() + info = litellm.get_model_info(model=regional_model, custom_llm_provider="bedrock") print("info", info) assert info["key"] == "anthropic.claude-haiku-4-5-20251001-v1:0" assert info["litellm_provider"] == "bedrock_converse" diff --git a/tests/test_rust_python_harness.py b/tests/test_rust_python_harness.py index a1bb370a074..9af38941684 100644 --- a/tests/test_rust_python_harness.py +++ b/tests/test_rust_python_harness.py @@ -36,8 +36,6 @@ def _case(module: str = "tests.example") -> HarnessCase: @pytest.mark.parametrize( "module", [ - "tests.rust-python-harness.strategies.e2e_parity.sdk.ocr.test_sdk_parity", - "tests.rust-python-harness.strategies.trace_parity.sdk.ocr.case", "tests.rust-python-harness.strategies.trace_parity.sdk.messages.case", "tests.rust-python-harness.strategies.trace_parity.sdk.chat_completions.case", "tests.rust-python-harness.strategies.trace_parity.sdk.transcription.case", diff --git a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 92ff3d5813c..f1c80bb11ea 100644 --- a/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -1,7 +1,6 @@ import asyncio -import logging from datetime import datetime, timedelta, timezone from unittest.mock import MagicMock, call, patch @@ -9,7 +8,6 @@ import pytest from prometheus_client import REGISTRY import litellm -from litellm._logging import verbose_logger from litellm.types.utils import ( StandardLoggingHiddenParams, StandardLoggingMetadata, @@ -27,10 +25,6 @@ except Exception: PrometheusLogger = None from litellm.proxy._types import UserAPIKeyAuth -verbose_logger.setLevel(logging.DEBUG) - -litellm.set_verbose = True - @pytest.fixture def prometheus_logger() -> PrometheusLogger: diff --git a/tests/unit/proxy/test_proxy_server.py b/tests/unit/proxy/test_proxy_server.py index eae80f311d8..65b368ca9e3 100644 --- a/tests/unit/proxy/test_proxy_server.py +++ b/tests/unit/proxy/test_proxy_server.py @@ -2979,7 +2979,7 @@ async def test_get_config_callbacks_environment_variables(client_no_auth): @pytest.mark.asyncio -async def test_update_config_success_callback_normalization(): +async def test_update_config_success_callback_normalization(monkeypatch): """ Ensure success_callback values are normalized to lowercase when updating config. This prevents delete_callback (which searches lowercase) from failing on mixed case inputs like 'SQS'. @@ -2987,7 +2987,7 @@ async def test_update_config_success_callback_normalization(): import litellm.proxy.proxy_server as proxy_server from litellm.proxy._types import ConfigYAML - setattr(proxy_server, "proxy_logging_obj", MagicMock()) + monkeypatch.setattr(proxy_server, "proxy_logging_obj", MagicMock()) existing_litellm_settings = {"success_callback": ["langfuse"]} @@ -3013,7 +3013,7 @@ async def test_update_config_success_callback_normalization(): self.db.litellm_config.find_first = AsyncMock(side_effect=fake_find_first) self.db.litellm_config.upsert = AsyncMock(side_effect=fake_upsert) - setattr(proxy_server, "prisma_client", MockPrisma()) + monkeypatch.setattr(proxy_server, "prisma_client", MockPrisma()) class MockProxyConfig: async def add_deployment(self, prisma_client=None, proxy_logging_obj=None): # noqa: F811 # pytest fixture, not a redefinition @@ -3022,7 +3022,7 @@ async def test_update_config_success_callback_normalization(): def reject_config_owned_writes(self, *, section_name, changed_keys): return None - setattr(proxy_server, "proxy_config", MockProxyConfig()) + monkeypatch.setattr(proxy_server, "proxy_config", MockProxyConfig()) config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]}) from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth From 4179860a17ec1b054db1a1e6e1d1402c4693089e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 25 Sep 2026 19:12:36 -0700 Subject: [PATCH 05/10] fix(cost-map): retirement dates, chatgpt reasoning flags, bing pricing, bedrock mantle and mythos, azure gpt-5.6 alias, anthropic batch rates, new nebius, openrouter and xai rows (#42951) --- .../crates/model-catalog/src/model_info.rs | 44 ++ litellm/litellm_core_utils/litellm_logging.py | 1 + ...odel_prices_and_context_window_backup.json | 493 +++++++++++++++--- litellm/types/utils.py | 2 + litellm/utils.py | 3 + model_prices_and_context_window.json | 493 +++++++++++++++--- model_prices_and_context_window.schema.json | 5 + tests/local_testing/test_get_model_info.py | 27 + .../test_bing_grounding_search.py | 3 +- .../llm_cost_calc/test_llm_cost_calc_utils.py | 12 + .../test_litellm_logging.py | 23 + tests/unit/test_model_prices_schema.py | 55 ++ tests/unit/test_utils.py | 1 + ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 + 14 files changed, 1029 insertions(+), 137 deletions(-) diff --git a/litellm-rust/crates/model-catalog/src/model_info.rs b/litellm-rust/crates/model-catalog/src/model_info.rs index 361cb56e9b1..9e9a4220e50 100644 --- a/litellm-rust/crates/model-catalog/src/model_info.rs +++ b/litellm-rust/crates/model-catalog/src/model_info.rs @@ -42,6 +42,9 @@ pub struct ModelInfo { pub cache_creation_input_token_cost_above_200k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] + pub cache_creation_input_token_cost_above_200k_tokens_batches: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] pub cache_creation_input_token_cost_above_256k_tokens: Option, /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] @@ -78,6 +81,9 @@ pub struct ModelInfo { /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_200k_tokens: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] + pub cache_read_input_token_cost_above_200k_tokens_batches: Option, /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub cache_read_input_token_cost_above_200k_tokens_priority: Option, @@ -113,6 +119,10 @@ pub struct ModelInfo { pub code_interpreter_cost_per_session: Option, #[serde(skip_serializing_if = "Option::is_none")] pub comment: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub computer_use_input_cost_per_1k_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub computer_use_output_cost_per_1k_tokens: Option, /// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'. #[serde(skip_serializing_if = "Option::is_none")] pub default_reasoning_effort: Option, @@ -120,6 +130,10 @@ pub struct ModelInfo { #[serde(skip_serializing_if = "Option::is_none")] pub deprecation_date: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub file_search_cost_per_1k_calls: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub file_search_cost_per_gb_per_day: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub gemini_audio_only_live: Option, #[serde(skip_serializing_if = "Option::is_none")] pub gemini_native_audio: Option, @@ -174,6 +188,9 @@ pub struct ModelInfo { /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_200k_tokens: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] + pub input_cost_per_token_above_200k_tokens_batches: Option, /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub input_cost_per_token_above_200k_tokens_priority: Option, @@ -265,6 +282,26 @@ pub struct ModelInfo { pub output_cost_per_image_1536: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_512: Option, + #[serde( + rename = "output_cost_per_image_0.5K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_0_5k: Option, + #[serde( + rename = "output_cost_per_image_1K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_1k: Option, + #[serde( + rename = "output_cost_per_image_2K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_2k: Option, + #[serde( + rename = "output_cost_per_image_4K", + skip_serializing_if = "Option::is_none" + )] + pub output_cost_per_image_4k: Option, #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_image_token: Option, #[serde(skip_serializing_if = "Option::is_none")] @@ -297,6 +334,9 @@ pub struct ModelInfo { /// Rate applied once the prompt exceeds the token threshold in the field name. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_200k_tokens: Option, + /// Rate applied once the prompt exceeds the token threshold in the field name. + #[serde(skip_serializing_if = "Option::is_none")] + pub output_cost_per_token_above_200k_tokens_batches: Option, /// Priority service-tier rate for the same-named base field. #[serde(skip_serializing_if = "Option::is_none")] pub output_cost_per_token_above_200k_tokens_priority: Option, @@ -357,6 +397,8 @@ pub struct ModelInfo { /// Provider default requests-per-minute limit. #[serde(skip_serializing_if = "Option::is_none")] pub rpm: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub rules: Option>, /// USD cost per web search query, keyed by search context size. #[serde(skip_serializing_if = "Option::is_none")] pub search_context_cost_per_query: Option, @@ -475,6 +517,8 @@ pub struct ModelInfo { #[serde(skip_serializing_if = "Option::is_none")] pub uses_embed_content: Option, #[serde(skip_serializing_if = "Option::is_none")] + pub vector_store_cost_per_gb_per_day: Option, + #[serde(skip_serializing_if = "Option::is_none")] pub vertex_ai_audio_api: Option, /// Whether web search is billed per query or per prompt. #[serde(skip_serializing_if = "Option::is_none")] diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 28d72702f3e..e955c0157c6 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -384,6 +384,7 @@ _DEPLOYMENT_PRICING_KEYS: Final = ( "cache_read_input_token_cost_above_200k_tokens_batches", "cache_read_input_token_cost_above_272k_tokens_batches", "cache_creation_input_token_cost_batches", + "cache_creation_input_token_cost_above_200k_tokens_batches", "cache_creation_input_token_cost_above_272k_tokens_batches", "ocr_cost_per_page", "ocr_cost_per_page_batches", diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 56129fef135..ec31039fee0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1279,8 +1279,11 @@ "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "anthropic.claude-mythos-preview": { - "input_cost_per_token": 0, - "output_cost_per_token": 0, + "cache_creation_input_token_cost": 3.4375e-05, + "cache_creation_input_token_cost_above_1hr": 5.5e-05, + "cache_read_input_token_cost": 2.75e-06, + "input_cost_per_token": 2.75e-05, + "output_cost_per_token": 0.0001375, "litellm_provider": "bedrock", "max_input_tokens": 1000000, "max_output_tokens": 128000, @@ -1289,10 +1292,11 @@ "thinking_always_on": true, "supports_function_calling": true, "supports_vision": true, - "supports_prompt_caching": false, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, - "supports_output_config": true + "supports_output_config": true, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "global.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -7773,27 +7777,27 @@ "output_cost_per_token_above_272k_tokens_batches": 0.000135 }, "azure/gpt-5.6": { - "cache_creation_input_token_cost": 6.25e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.25e-05, - "cache_creation_input_token_cost_priority": 1.25e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.5e-05, - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, - "cache_read_input_token_cost_priority": 1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, - "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, - "input_cost_per_token_priority": 1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "cache_creation_input_token_cost": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1e-05, + "cache_creation_input_token_cost_priority": 1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2e-05, + "cache_read_input_token_cost": 4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, + "cache_read_input_token_cost_priority": 8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.6e-06, + "input_cost_per_token": 4e-06, + "input_cost_per_token_above_272k_tokens": 8e-06, + "input_cost_per_token_priority": 8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.6e-05, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "output_cost_per_token_priority": 6e-05, - "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "output_cost_per_token": 2e-05, + "output_cost_per_token_above_272k_tokens": 3e-05, + "output_cost_per_token_priority": 4e-05, + "output_cost_per_token_above_272k_tokens_priority": 6e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8579,27 +8583,27 @@ "supports_web_search": true }, "azure/us/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8985,27 +8989,27 @@ "supports_web_search": true }, "azure/eu/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14640,14 +14644,18 @@ "claude-haiku-4-5-20251001": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14663,14 +14671,18 @@ "claude-haiku-4-5": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14693,13 +14705,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14727,13 +14747,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14757,14 +14785,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, + "cache_creation_input_token_cost_batches": 1.25e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14797,14 +14829,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14865,14 +14901,18 @@ "claude-opus-4-5-20251101": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14895,14 +14935,18 @@ "claude-opus-4-5": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14927,14 +14971,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14966,14 +15014,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15004,14 +15056,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15044,14 +15100,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15083,14 +15143,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15123,14 +15187,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15164,14 +15232,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 5e-06, "cache_creation_input_token_cost_above_1hr": 8e-06, + "cache_creation_input_token_cost_batches": 2.5e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 4e-06, + "input_cost_per_token_batches": 2e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15207,14 +15279,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15250,14 +15326,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15756,6 +15836,16 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "c4ai-aya-expanse-32b": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.5e-06, + "litellm_provider": "cohere_chat", + "max_input_tokens": 128000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "chat", + "source": "https://docs.cohere.com/docs/models" + }, "command-a-plus-05-2026": { "input_cost_per_token": 0.0, "litellm_provider": "cohere_chat", @@ -22145,11 +22235,11 @@ } }, "bing_grounding/search": { - "input_cost_per_query": 0.035, + "input_cost_per_query": 0.014, "litellm_provider": "bing_grounding", "mode": "search", "metadata": { - "notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." + "notes": "Grounding with Bing Search (G1 SKU): $14 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." } }, "tinyfish/search": { @@ -28036,6 +28126,7 @@ "tpm": 10000000 }, "gemini/gemini-3-pro-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -28083,6 +28174,7 @@ "supports_reasoning": false }, "gemini/gemini-3.1-flash-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -28130,6 +28222,7 @@ "cache_read_input_token_cost_batches": 1.25e-08, "cache_read_input_token_cost_flex": 1.25e-08, "cache_read_input_token_cost_priority": 4.5e-08, + "deprecation_date": "2026-05-25", "input_cost_per_audio_token": 5e-07, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, @@ -30551,7 +30644,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.6-luna": { "litellm_provider": "chatgpt", @@ -30567,7 +30664,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-sol": { "litellm_provider": "chatgpt", @@ -30583,7 +30684,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-terra": { "litellm_provider": "chatgpt", @@ -30599,7 +30704,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.4": { "litellm_provider": "chatgpt", @@ -30614,7 +30723,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.4-pro": { "litellm_provider": "chatgpt", @@ -30628,7 +30742,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.3-codex": { "litellm_provider": "chatgpt", @@ -30642,7 +30760,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": false, + "supports_minimal_reasoning_effort": true }, "chatgpt/gpt-5.3-codex-spark": { "litellm_provider": "chatgpt", @@ -30686,7 +30808,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2-codex": { "litellm_provider": "chatgpt", @@ -30700,7 +30823,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2": { "litellm_provider": "chatgpt", @@ -30715,7 +30839,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.1-codex-max": { "litellm_provider": "chatgpt", @@ -30729,7 +30858,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.1-codex-mini": { "litellm_provider": "chatgpt", @@ -30743,7 +30873,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "gigachat/GigaChat-2": { "input_cost_per_token": 0.0, @@ -39214,6 +39345,17 @@ "supports_function_calling": true, "supports_reasoning": true }, + "nebius/deepseek-ai/DeepSeek-V4.1-Flash": { + "input_cost_per_token": 3e-07, + "litellm_provider": "nebius", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://tokenfactory.nebius.com/endpoints?modals=endpoint-details&model-id=deepseek-ai/DeepSeek-V4.1-Flash", + "supports_vision": true + }, "nebius/MiniMaxAI/MiniMax-M2.5": { "max_tokens": 196608, "max_input_tokens": 196608, @@ -56230,7 +56372,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -56417,7 +56559,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -59653,14 +59795,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "prompt_cache_min_tokens": 512, "search_context_cost_per_query": { "search_context_size_high": 0.01, @@ -59693,14 +59839,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -63000,6 +63150,22 @@ "image" ] }, + "xai/grok-imagine-image-pro": { + "input_cost_per_image": 0.05, + "litellm_provider": "xai", + "mode": "image_generation", + "source": "https://docs.x.ai/docs/models", + "supported_endpoints": [ + "/v1/images/generations" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "image" + ] + }, "xai/grok-imagine-image-2.0": { "input_cost_per_image": 0.06, "litellm_provider": "xai", @@ -64006,6 +64172,55 @@ "supports_response_schema": true, "supports_vision": true }, + "azure_ai/deepseek-r1": { + "input_cost_per_token": 1.35e-06, + "output_cost_per_token": 5.4e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3-0324": { + "input_cost_per_token": 1.14e-06, + "output_cost_per_token": 4.56e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3.1": { + "input_cost_per_token": 1.23e-06, + "output_cost_per_token": 4.94e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3": { + "input_cost_per_token": 3e-06, + "output_cost_per_token": 1.5e-05, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3-mini": { + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.27e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-non-reasoning": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-reasoning": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", @@ -64834,6 +65049,131 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "bedrock_mantle/deepseek.v3.1": { + "input_cost_per_token": 5.8e-07, + "output_cost_per_token": 1.68e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html" + }, + "bedrock_mantle/moonshotai.kimi-k2-thinking": { + "input_cost_per_token": 6e-07, + "output_cost_per_token": 2.5e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html" + }, + "bedrock_mantle/qwen.qwen3-235b-a22b-2507": { + "input_cost_per_token": 2.2e-07, + "output_cost_per_token": 8.8e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-235b-a22b-2507.html" + }, + "bedrock_mantle/qwen.qwen3-32b": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 32000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-32b.html" + }, + "bedrock_mantle/qwen.qwen3-coder-30b-a3b-instruct": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-30b-a3b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-coder-480b-a35b-instruct": { + "input_cost_per_token": 4.5e-07, + "output_cost_per_token": 1.8e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-480b-a35b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-next-80b-a3b-instruct": { + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-next-80b-a3b.html" + }, + "bedrock_mantle/qwen.qwen3-vl-235b-a22b-instruct": { + "input_cost_per_token": 5.3e-07, + "output_cost_per_token": 2.66e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-vl-235b-a22b.html" + }, "azure/us-gov/gpt-5.1": { "cache_read_input_token_cost": 1.71875e-07, "default_reasoning_effort": "none", @@ -76578,6 +76918,23 @@ "supports_vision": true, "supports_web_search": false }, + "openrouter/perceptron/perceptron-mk1.5": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 36864, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, "vertex_ai/gemini-2.0-flash": { "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 1e-06, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index cd336c9b989..2e518af4da4 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -302,6 +302,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): cache_read_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None] cache_read_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None] cache_creation_input_token_cost_batches: ReadOnly[float | None] + cache_creation_input_token_cost_above_200k_tokens_batches: ReadOnly[float | None] cache_creation_input_token_cost_above_272k_tokens_batches: ReadOnly[float | None] # Smallest prefix this model will actually cache, whatever caching mechanism its provider uses. # Absent means the provider-agnostic default applies; see MINIMUM_PROMPT_CACHE_TOKEN_COUNT. @@ -3735,6 +3736,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): cache_read_input_token_cost_above_200k_tokens_batches: float | None = None cache_read_input_token_cost_above_272k_tokens_batches: float | None = None cache_creation_input_token_cost_batches: float | None = None + cache_creation_input_token_cost_above_200k_tokens_batches: float | None = None cache_creation_input_token_cost_above_272k_tokens_batches: float | None = None cache_read_input_audio_token_cost: float | None = None cache_read_input_image_token_cost: float | None = None diff --git a/litellm/utils.py b/litellm/utils.py index be4388802f9..42da2e2a7b7 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6167,6 +6167,9 @@ def _get_model_info_helper( "cache_read_input_token_cost_above_272k_tokens_batches" ), cache_creation_input_token_cost_batches=_model_info.get("cache_creation_input_token_cost_batches"), + cache_creation_input_token_cost_above_200k_tokens_batches=_model_info.get( + "cache_creation_input_token_cost_above_200k_tokens_batches" + ), cache_creation_input_token_cost_above_272k_tokens_batches=_model_info.get( "cache_creation_input_token_cost_above_272k_tokens_batches" ), diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 56129fef135..ec31039fee0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1279,8 +1279,11 @@ "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "anthropic.claude-mythos-preview": { - "input_cost_per_token": 0, - "output_cost_per_token": 0, + "cache_creation_input_token_cost": 3.4375e-05, + "cache_creation_input_token_cost_above_1hr": 5.5e-05, + "cache_read_input_token_cost": 2.75e-06, + "input_cost_per_token": 2.75e-05, + "output_cost_per_token": 0.0001375, "litellm_provider": "bedrock", "max_input_tokens": 1000000, "max_output_tokens": 128000, @@ -1289,10 +1292,11 @@ "thinking_always_on": true, "supports_function_calling": true, "supports_vision": true, - "supports_prompt_caching": false, + "supports_prompt_caching": true, "supports_reasoning": true, "supports_tool_choice": true, - "supports_output_config": true + "supports_output_config": true, + "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" }, "global.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -7773,27 +7777,27 @@ "output_cost_per_token_above_272k_tokens_batches": 0.000135 }, "azure/gpt-5.6": { - "cache_creation_input_token_cost": 6.25e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.25e-05, - "cache_creation_input_token_cost_priority": 1.25e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.5e-05, - "cache_read_input_token_cost": 5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1e-06, - "cache_read_input_token_cost_priority": 1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, - "input_cost_per_token": 5e-06, - "input_cost_per_token_above_272k_tokens": 1e-05, - "input_cost_per_token_priority": 1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "cache_creation_input_token_cost": 5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1e-05, + "cache_creation_input_token_cost_priority": 1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2e-05, + "cache_read_input_token_cost": 4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8e-07, + "cache_read_input_token_cost_priority": 8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.6e-06, + "input_cost_per_token": 4e-06, + "input_cost_per_token_above_272k_tokens": 8e-06, + "input_cost_per_token_priority": 8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.6e-05, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-05, - "output_cost_per_token_above_272k_tokens": 4.5e-05, - "output_cost_per_token_priority": 6e-05, - "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "output_cost_per_token": 2e-05, + "output_cost_per_token_above_272k_tokens": 3e-05, + "output_cost_per_token_priority": 4e-05, + "output_cost_per_token_above_272k_tokens_priority": 6e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8579,27 +8583,27 @@ "supports_web_search": true }, "azure/us/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -8985,27 +8989,27 @@ "supports_web_search": true }, "azure/eu/gpt-5.6": { - "cache_creation_input_token_cost": 6.875e-06, - "cache_creation_input_token_cost_above_272k_tokens": 1.375e-05, - "cache_creation_input_token_cost_above_272k_tokens_priority": 2.75e-05, - "cache_creation_input_token_cost_priority": 1.375e-05, - "cache_read_input_token_cost": 5.5e-07, - "cache_read_input_token_cost_above_272k_tokens": 1.1e-06, - "cache_read_input_token_cost_above_272k_tokens_priority": 2.2e-06, - "cache_read_input_token_cost_priority": 1.1e-06, - "input_cost_per_token": 5.5e-06, - "input_cost_per_token_above_272k_tokens": 1.1e-05, - "input_cost_per_token_above_272k_tokens_priority": 2.2e-05, - "input_cost_per_token_priority": 1.1e-05, + "cache_creation_input_token_cost": 5.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 1.1e-05, + "cache_creation_input_token_cost_above_272k_tokens_priority": 2.2e-05, + "cache_creation_input_token_cost_priority": 1.1e-05, + "cache_read_input_token_cost": 4.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 8.8e-07, + "cache_read_input_token_cost_above_272k_tokens_priority": 1.76e-06, + "cache_read_input_token_cost_priority": 8.8e-07, + "input_cost_per_token": 4.4e-06, + "input_cost_per_token_above_272k_tokens": 8.8e-06, + "input_cost_per_token_above_272k_tokens_priority": 1.76e-05, + "input_cost_per_token_priority": 8.8e-06, "litellm_provider": "azure", "max_input_tokens": 922000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3.3e-05, - "output_cost_per_token_above_272k_tokens": 4.95e-05, - "output_cost_per_token_above_272k_tokens_priority": 9.9e-05, - "output_cost_per_token_priority": 6.6e-05, + "output_cost_per_token": 2.2e-05, + "output_cost_per_token_above_272k_tokens": 3.3e-05, + "output_cost_per_token_above_272k_tokens_priority": 6.6e-05, + "output_cost_per_token_priority": 4.4e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14640,14 +14644,18 @@ "claude-haiku-4-5-20251001": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14663,14 +14671,18 @@ "claude-haiku-4-5": { "cache_creation_input_token_cost": 1.25e-06, "cache_creation_input_token_cost_above_1hr": 2e-06, + "cache_creation_input_token_cost_batches": 6.25e-07, "cache_read_input_token_cost": 1e-07, + "cache_read_input_token_cost_batches": 5e-08, "input_cost_per_token": 1e-06, + "input_cost_per_token_batches": 5e-07, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 5e-06, + "output_cost_per_token_batches": 2.5e-06, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_native_structured_output": true, @@ -14693,13 +14705,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14727,13 +14747,21 @@ "input_cost_per_token_above_200k_tokens": 6e-06, "output_cost_per_token_above_200k_tokens": 2.25e-05, "cache_creation_input_token_cost_above_200k_tokens": 7.5e-06, + "cache_creation_input_token_cost_above_200k_tokens_batches": 3.75e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost_above_200k_tokens": 6e-07, + "cache_read_input_token_cost_above_200k_tokens_batches": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, + "input_cost_per_token_above_200k_tokens_batches": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_above_200k_tokens_batches": 1.125e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14757,14 +14785,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 2.5e-06, "cache_creation_input_token_cost_above_1hr": 4e-06, + "cache_creation_input_token_cost_batches": 1.25e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1e-05, + "output_cost_per_token_batches": 5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14797,14 +14829,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 3.75e-06, "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_batches": 1.875e-06, "cache_read_input_token_cost": 3e-07, + "cache_read_input_token_cost_batches": 1.5e-07, "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14865,14 +14901,18 @@ "claude-opus-4-5-20251101": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14895,14 +14935,18 @@ "claude-opus-4-5": { "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 200000, "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14927,14 +14971,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -14966,14 +15014,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15004,14 +15056,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15044,14 +15100,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15083,14 +15143,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15123,14 +15187,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15164,14 +15232,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 5e-06, "cache_creation_input_token_cost_above_1hr": 8e-06, + "cache_creation_input_token_cost_batches": 2.5e-06, "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_batches": 1e-07, "input_cost_per_token": 4e-06, + "input_cost_per_token_batches": 2e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2e-05, + "output_cost_per_token_batches": 1e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15207,14 +15279,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15250,14 +15326,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 6.25e-06, "cache_creation_input_token_cost_above_1hr": 1e-05, + "cache_creation_input_token_cost_batches": 3.125e-06, "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_batches": 2.5e-07, "input_cost_per_token": 5e-06, + "input_cost_per_token_batches": 2.5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 2.5e-05, + "output_cost_per_token_batches": 1.25e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -15756,6 +15836,16 @@ "supports_function_calling": true, "supports_tool_choice": true }, + "c4ai-aya-expanse-32b": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.5e-06, + "litellm_provider": "cohere_chat", + "max_input_tokens": 128000, + "max_output_tokens": 4000, + "max_tokens": 4000, + "mode": "chat", + "source": "https://docs.cohere.com/docs/models" + }, "command-a-plus-05-2026": { "input_cost_per_token": 0.0, "litellm_provider": "cohere_chat", @@ -22145,11 +22235,11 @@ } }, "bing_grounding/search": { - "input_cost_per_query": 0.035, + "input_cost_per_query": 0.014, "litellm_provider": "bing_grounding", "mode": "search", "metadata": { - "notes": "Grounding with Bing Search (G1 SKU): $35 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." + "notes": "Grounding with Bing Search (G1 SKU): $14 per 1,000 transactions. Tokens for the Foundry model deployment that runs the grounded search are billed separately on that deployment." } }, "tinyfish/search": { @@ -28036,6 +28126,7 @@ "tpm": 10000000 }, "gemini/gemini-3-pro-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_image": 0.0011, "input_cost_per_token": 2e-06, "input_cost_per_token_batches": 1e-06, @@ -28083,6 +28174,7 @@ "supports_reasoning": false }, "gemini/gemini-3.1-flash-image-preview": { + "deprecation_date": "2026-06-25", "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, "litellm_provider": "gemini", @@ -28130,6 +28222,7 @@ "cache_read_input_token_cost_batches": 1.25e-08, "cache_read_input_token_cost_flex": 1.25e-08, "cache_read_input_token_cost_priority": 4.5e-08, + "deprecation_date": "2026-05-25", "input_cost_per_audio_token": 5e-07, "input_cost_per_token": 2.5e-07, "input_cost_per_token_batches": 1.25e-07, @@ -30551,7 +30644,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.6-luna": { "litellm_provider": "chatgpt", @@ -30567,7 +30664,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-sol": { "litellm_provider": "chatgpt", @@ -30583,7 +30684,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.6-terra": { "litellm_provider": "chatgpt", @@ -30599,7 +30704,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true }, "chatgpt/gpt-5.4": { "litellm_provider": "chatgpt", @@ -30614,7 +30723,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.4-pro": { "litellm_provider": "chatgpt", @@ -30628,7 +30742,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.3-codex": { "litellm_provider": "chatgpt", @@ -30642,7 +30760,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": false, + "supports_xhigh_reasoning_effort": false, + "supports_minimal_reasoning_effort": true }, "chatgpt/gpt-5.3-codex-spark": { "litellm_provider": "chatgpt", @@ -30686,7 +30808,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2-codex": { "litellm_provider": "chatgpt", @@ -30700,7 +30823,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.2": { "litellm_provider": "chatgpt", @@ -30715,7 +30839,12 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true, + "supports_none_reasoning_effort": true, + "default_reasoning_effort": "none", + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false }, "chatgpt/gpt-5.1-codex-max": { "litellm_provider": "chatgpt", @@ -30729,7 +30858,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "chatgpt/gpt-5.1-codex-mini": { "litellm_provider": "chatgpt", @@ -30743,7 +30873,8 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_vision": true + "supports_vision": true, + "supports_reasoning": true }, "gigachat/GigaChat-2": { "input_cost_per_token": 0.0, @@ -39214,6 +39345,17 @@ "supports_function_calling": true, "supports_reasoning": true }, + "nebius/deepseek-ai/DeepSeek-V4.1-Flash": { + "input_cost_per_token": 3e-07, + "litellm_provider": "nebius", + "max_input_tokens": 1048576, + "max_output_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://tokenfactory.nebius.com/endpoints?modals=endpoint-details&model-id=deepseek-ai/DeepSeek-V4.1-Flash", + "supports_vision": true + }, "nebius/MiniMaxAI/MiniMax-M2.5": { "max_tokens": 196608, "max_input_tokens": 196608, @@ -56230,7 +56372,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -56417,7 +56559,7 @@ "input_cost_per_audio_token": 3e-06, "input_cost_per_token": 5e-07, "litellm_provider": "gemini", - "max_input_tokens": 1048576, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "realtime", @@ -59653,14 +59795,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 1e-06, + "cache_read_input_token_cost_batches": 5e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "prompt_cache_min_tokens": 512, "search_context_cost_per_query": { "search_context_size_high": 0.01, @@ -59693,14 +59839,18 @@ "supports_anthropic_compaction": true, "cache_creation_input_token_cost": 1.25e-05, "cache_creation_input_token_cost_above_1hr": 2e-05, + "cache_creation_input_token_cost_batches": 6.25e-06, "cache_read_input_token_cost": 2.5e-07, + "cache_read_input_token_cost_batches": 1.25e-07, "input_cost_per_token": 1e-05, + "input_cost_per_token_batches": 5e-06, "litellm_provider": "anthropic", "max_input_tokens": 1000000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", "output_cost_per_token": 5e-05, + "output_cost_per_token_batches": 2.5e-05, "search_context_cost_per_query": { "search_context_size_high": 0.01, "search_context_size_low": 0.01, @@ -63000,6 +63150,22 @@ "image" ] }, + "xai/grok-imagine-image-pro": { + "input_cost_per_image": 0.05, + "litellm_provider": "xai", + "mode": "image_generation", + "source": "https://docs.x.ai/docs/models", + "supported_endpoints": [ + "/v1/images/generations" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "image" + ] + }, "xai/grok-imagine-image-2.0": { "input_cost_per_image": 0.06, "litellm_provider": "xai", @@ -64006,6 +64172,55 @@ "supports_response_schema": true, "supports_vision": true }, + "azure_ai/deepseek-r1": { + "input_cost_per_token": 1.35e-06, + "output_cost_per_token": 5.4e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3-0324": { + "input_cost_per_token": 1.14e-06, + "output_cost_per_token": 4.56e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/deepseek-v3.1": { + "input_cost_per_token": 1.23e-06, + "output_cost_per_token": 4.94e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3": { + "input_cost_per_token": 3e-06, + "output_cost_per_token": 1.5e-05, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-3-mini": { + "input_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.27e-06, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-non-reasoning": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, + "azure_ai/grok-4-fast-reasoning": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 5e-07, + "litellm_provider": "azure_ai", + "mode": "chat", + "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + }, "bedrock/us-gov-west-1/nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 7.2e-08, "litellm_provider": "bedrock", @@ -64834,6 +65049,131 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "bedrock_mantle/deepseek.v3.1": { + "input_cost_per_token": 5.8e-07, + "output_cost_per_token": 1.68e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-deepseek-deepseek-v3-1.html" + }, + "bedrock_mantle/moonshotai.kimi-k2-thinking": { + "input_cost_per_token": 6e-07, + "output_cost_per_token": 2.5e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-moonshot-ai-kimi-k2-thinking.html" + }, + "bedrock_mantle/qwen.qwen3-235b-a22b-2507": { + "input_cost_per_token": 2.2e-07, + "output_cost_per_token": 8.8e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-235b-a22b-2507.html" + }, + "bedrock_mantle/qwen.qwen3-32b": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 32000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-32b.html" + }, + "bedrock_mantle/qwen.qwen3-coder-30b-a3b-instruct": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-30b-a3b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-coder-480b-a35b-instruct": { + "input_cost_per_token": 4.5e-07, + "output_cost_per_token": 1.8e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 128000, + "max_output_tokens": 16000, + "max_tokens": 16000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-coder-480b-a35b-instruct.html" + }, + "bedrock_mantle/qwen.qwen3-next-80b-a3b-instruct": { + "input_cost_per_token": 1.4e-07, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-next-80b-a3b.html" + }, + "bedrock_mantle/qwen.qwen3-vl-235b-a22b-instruct": { + "input_cost_per_token": 5.3e-07, + "output_cost_per_token": 2.66e-06, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 256000, + "max_output_tokens": 8000, + "max_tokens": 8000, + "mode": "chat", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-qwen-qwen3-vl-235b-a22b.html" + }, "azure/us-gov/gpt-5.1": { "cache_read_input_token_cost": 1.71875e-07, "default_reasoning_effort": "none", @@ -76578,6 +76918,23 @@ "supports_vision": true, "supports_web_search": false }, + "openrouter/perceptron/perceptron-mk1.5": { + "input_cost_per_token": 1.5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 36864, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "source": "https://openrouter.ai/api/v1/models", + "supports_audio_input": true, + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, "vertex_ai/gemini-2.0-flash": { "deprecation_date": "2026-06-01", "input_cost_per_audio_token": 1e-06, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index fa1828c780a..c4048cac905 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -103,6 +103,11 @@ "minimum": 0, "description": "Rate applied once the prompt exceeds the token threshold in the field name." }, + "cache_creation_input_token_cost_above_200k_tokens_batches": { + "type": "number", + "minimum": 0, + "description": "Rate applied once the prompt exceeds the token threshold in the field name." + }, "cache_creation_input_token_cost_above_256k_tokens": { "type": "number", "minimum": 0, diff --git a/tests/local_testing/test_get_model_info.py b/tests/local_testing/test_get_model_info.py index dbe1fc3b69b..79f6739a423 100644 --- a/tests/local_testing/test_get_model_info.py +++ b/tests/local_testing/test_get_model_info.py @@ -320,6 +320,33 @@ def test_get_model_info_bedrock_cross_region_capability_parity(): assert checked > 0, "no cross-region bedrock profiles found - the filter is inert" + +def test_get_model_info_bedrock_priced_cross_region_profile_has_priced_base(): + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + prefixes = ("us.", "eu.", "apac.", "us-gov.", "au.", "global.") + checked = 0 + + for k, v in litellm.model_cost.items(): + if not str(v.get("litellm_provider", "")).startswith("bedrock"): + continue + base_model_key = next( + (k[len(p) :] for p in prefixes if k.startswith(p)), + None, + ) + if base_model_key is None or base_model_key not in litellm.model_cost: + continue + checked += 1 + base = litellm.model_cost[base_model_key] + for cost_key in ("input_cost_per_token", "output_cost_per_token"): + if (v.get(cost_key) or 0) > 0: + assert ( + base.get(cost_key) or 0 + ) > 0, f"{k} charges {cost_key} but its base {base_model_key} is free" + + assert checked > 0, "no cross-region bedrock profiles found - the filter is inert" + def test_get_model_info_huggingface_models(monkeypatch): from litellm import Router from litellm.types.router import ModelGroupInfo diff --git a/tests/search_tests/test_bing_grounding_search.py b/tests/search_tests/test_bing_grounding_search.py index f532158e462..6ea79076370 100644 --- a/tests/search_tests/test_bing_grounding_search.py +++ b/tests/search_tests/test_bing_grounding_search.py @@ -196,4 +196,5 @@ class TestBingGroundingSearchTransformation: ): response = litellm.search(query="pricing check", search_provider="bing_grounding") - assert response._hidden_params["response_cost"] == pytest.approx(0.035) + # Grounding with Bing Search (G1 SKU): $14 per 1,000 transactions, https://www.microsoft.com/en-us/bing/apis, checked 2026-09-24 + assert response._hidden_params["response_cost"] == pytest.approx(0.014) diff --git a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 781a3a7c4ed..0afd989272e 100644 --- a/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/unit/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -3789,3 +3789,15 @@ def test_azure_gpt_6_foundry_price_sheet(_local_model_cost_map, model_base): assert azure_ai_info[field] == base assert azure_us_info[field] == pytest.approx(1.1 * base) assert azure_eu_info[field] == pytest.approx(1.2 * base) + + +@pytest.mark.parametrize("region_prefix", ["azure/", "azure/us/", "azure/eu/"]) +def test_azure_gpt_5_6_alias_matches_sol_pricing(_local_model_cost_map, region_prefix): + """The bare gpt-5.6 alias routes to GPT-5.6 Sol, so every Azure region must bill the + alias exactly like the Sol entry (including the Sept 2026 $4/$20 promo).""" + alias = litellm.model_cost[f"{region_prefix}gpt-5.6"] + sol = litellm.model_cost[f"{region_prefix}gpt-5.6-sol"] + shared_cost_fields = [f for f in alias if "cost" in f and f in sol and not isinstance(alias[f], dict)] + assert shared_cost_fields + for field in shared_cost_fields: + assert alias[field] == sol[field], field diff --git a/tests/unit/litellm_core_utils/test_litellm_logging.py b/tests/unit/litellm_core_utils/test_litellm_logging.py index d717718cba2..cb1e281e356 100644 --- a/tests/unit/litellm_core_utils/test_litellm_logging.py +++ b/tests/unit/litellm_core_utils/test_litellm_logging.py @@ -7969,6 +7969,29 @@ def test_deployment_pricing_model_info_honors_a_tier_only_batch_override_over_th assert {key: info[key] for key in carried_keys} == {key: _PUBLISHED_BATCH_RATES[key] for key in carried_keys} +@pytest.mark.parametrize( + "override_key", + ( + "output_cost_per_token_above_200k_tokens_batches", + "cache_read_input_token_cost_above_200k_tokens_batches", + "cache_creation_input_token_cost_above_200k_tokens_batches", + ), +) +def test_deployment_pricing_model_info_honors_a_200k_tier_batch_override( + _published_batch_model: None, override_key: str +) -> None: + from litellm.litellm_core_utils.litellm_logging import deployment_pricing_model_info + + info: Final = deployment_pricing_model_info(_batch_deployment_id({override_key: 1e-3}), _PUBLISHED_BATCH_DEPLOYMENT) + carried_keys: Final = tuple( + key for key in (*_PUBLISHED_INPUT_BATCH_KEYS, *_PUBLISHED_OUTPUT_BATCH_KEYS) if key != override_key + ) + + assert info is not None + assert info[override_key] == 1e-3 + assert {key: info[key] for key in carried_keys} == {key: _PUBLISHED_BATCH_RATES[key] for key in carried_keys} + + def test_get_status_fields_ranks_guardrail_flagged_between_success_and_intervened(): """LIT-6894: a non-blocking flagged verdict must outrank success in the request-level guardrail_status but never mask an intervention.""" diff --git a/tests/unit/test_model_prices_schema.py b/tests/unit/test_model_prices_schema.py index 052278631e2..a05c345b5cd 100644 --- a/tests/unit/test_model_prices_schema.py +++ b/tests/unit/test_model_prices_schema.py @@ -266,6 +266,61 @@ def test_openai_reasoning_family_entries_carry_supports_reasoning(prices: dict): ) +_ABSENT: Final = object() + +REASONING_ANNOTATION_KEYS: Final = ( + "supports_reasoning", + "supports_minimal_reasoning_effort", + "supports_none_reasoning_effort", + "supports_xhigh_reasoning_effort", + "default_reasoning_effort", +) + + +def chatgpt_openai_twins(prices: dict) -> list[tuple[str, str]]: + """`chatgpt/` rows paired with the bare `` row served by the openai provider. + + Scoped to openai twins on purpose. `ChatGPTConfig` and `ChatGPTResponsesAPIConfig` subclass + their openai counterparts, so a chatgpt row's reasoning behaviour is whatever the openai row + describes. The azure rows are a separate registry that already diverges from openai here, and + pinning them to each other would assert something this repository does not control. + """ + pairs = [] + for name, entry in prices.items(): + if not isinstance(entry, dict) or not name.startswith("chatgpt/"): + continue + bare = name.split("/", 1)[1] + twin = prices.get(bare) + if isinstance(twin, dict) and twin.get("litellm_provider") == "openai": + pairs.append((name, bare)) + return pairs + + +def test_chatgpt_rows_carry_their_openai_twin_reasoning_annotations(prices: dict): + """A chatgpt row must not silently drop the reasoning annotations of the model it proxies. + + `litellm.utils._get_model_info_from_generalization` refuses to fall back when an exact cost-map + key exists, so an unannotated `chatgpt/` row wins over its annotated twin and + `/model/info` reports the model as non-reasoning. + """ + twins = chatgpt_openai_twins(prices) + assert twins, "no chatgpt/* row has an openai twin any more; this guard has stopped guarding" + + mismatched = [] + for name, bare in twins: + for key in REASONING_ANNOTATION_KEYS: + if prices[name].get(key, _ABSENT) != prices[bare].get(key, _ABSENT): + mismatched.append( + f"{name}.{key} is {prices[name].get(key)!r}, {bare}.{key} is {prices[bare].get(key)!r}" + ) + + assert mismatched == [], ( + "chatgpt/* entries proxy their openai twin through ChatGPTConfig, so they must carry the " + "same reasoning annotations; an exact cost-map key blocks the generalization fallback, so " + "a missing flag here is reported to callers as 'not a reasoning model':\n" + "\n".join(mismatched) + ) + + def test_chat_latest_declares_the_one_effort_openai_accepts(prices: dict): """OpenAI rejects every reasoning.effort on chat-latest except medium, and a reasoning entry with no declared levels resolves to None, which lets /model_group/info and the dashboard offer diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 0cdc52c9a93..5b327305e31 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -766,6 +766,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_creation_input_token_cost_above_272k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_flex": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_priority": {"type": "number"}, + "cache_creation_input_token_cost_above_200k_tokens_batches": {"type": "number"}, "cache_creation_input_token_cost_above_272k_tokens_batches": {"type": "number"}, "cache_creation_input_token_cost_batches": {"type": "number"}, "cache_creation_input_token_cost_flex": {"type": "number"}, diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 1e0aed46923..6d46af04730 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -32957,6 +32957,8 @@ export interface components { cache_creation_input_token_cost_above_1hr?: number | null; /** Cache Creation Input Token Cost Above 200K Tokens */ cache_creation_input_token_cost_above_200k_tokens?: number | null; + /** Cache Creation Input Token Cost Above 200K Tokens Batches */ + cache_creation_input_token_cost_above_200k_tokens_batches?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens */ cache_creation_input_token_cost_above_272k_tokens?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens Batches */ @@ -46750,6 +46752,8 @@ export interface components { cache_creation_input_token_cost_above_1hr?: number | null; /** Cache Creation Input Token Cost Above 200K Tokens */ cache_creation_input_token_cost_above_200k_tokens?: number | null; + /** Cache Creation Input Token Cost Above 200K Tokens Batches */ + cache_creation_input_token_cost_above_200k_tokens_batches?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens */ cache_creation_input_token_cost_above_272k_tokens?: number | null; /** Cache Creation Input Token Cost Above 272K Tokens Batches */ From 2ef3250ec3bdd275dae3281d1e0f2347d995e5fe Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 02:16:32 +0000 Subject: [PATCH 06/10] refactor(rust): promote anthropic messages out of experimental_pass_through (#43269) Co-authored-by: Yujong Lee --- litellm-rust/crates/core/src/messages/common_utils.rs | 2 +- litellm-rust/crates/core/src/messages/prepare.rs | 2 +- .../crates/llms/src/anthropic/batches/transformation.rs | 2 +- litellm-rust/crates/llms/src/anthropic/chat/handler.rs | 2 +- .../crates/llms/src/anthropic/chat/transformation.rs | 4 +--- .../llms/src/anthropic/experimental_pass_through/mod.rs | 1 - .../{experimental_pass_through => }/messages/handler.rs | 0 .../{experimental_pass_through => }/messages/headers.rs | 0 .../{experimental_pass_through => }/messages/mod.rs | 0 .../messages/streaming_iterator.rs | 0 .../{experimental_pass_through => }/messages/thinking.rs | 0 .../messages/transformation.rs | 0 litellm-rust/crates/llms/src/anthropic/mod.rs | 5 +++-- .../llms/src/azure_ai/anthropic/messages_transformation.rs | 2 +- .../llms/src/base_llm/anthropic_messages/transformation.rs | 3 +-- 15 files changed, 10 insertions(+), 13 deletions(-) delete mode 100644 litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs rename litellm-rust/crates/llms/src/anthropic/{experimental_pass_through => }/messages/handler.rs (100%) rename litellm-rust/crates/llms/src/anthropic/{experimental_pass_through => }/messages/headers.rs (100%) rename litellm-rust/crates/llms/src/anthropic/{experimental_pass_through => }/messages/mod.rs (100%) rename litellm-rust/crates/llms/src/anthropic/{experimental_pass_through => }/messages/streaming_iterator.rs (100%) rename litellm-rust/crates/llms/src/anthropic/{experimental_pass_through => }/messages/thinking.rs (100%) rename litellm-rust/crates/llms/src/anthropic/{experimental_pass_through => }/messages/transformation.rs (100%) diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index d27b79bdc04..95142e87519 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,7 +1,7 @@ use litellm_http::request::string_headers as shared_string_headers; pub(super) use litellm_http::request::truncate_error_body; use litellm_llms::{ - anthropic::experimental_pass_through::messages::transformation::ANTHROPIC_MESSAGES_CONFIG, + anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG, azure_ai::anthropic::messages_transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG, base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, }; diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index dc4b3562e3f..84884ab279e 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -5,7 +5,7 @@ use litellm_core_utils::{ settings::Lookup, }; use litellm_llms::{ - anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request, + anthropic::messages::handler::shape_anthropic_messages_request, base_llm::anthropic_messages::transformation::{ BaseAnthropicMessagesConfig, MessagesTransformContext, }, diff --git a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs index 94e4dc7838a..1c26684901a 100644 --- a/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/batches/transformation.rs @@ -5,7 +5,7 @@ use time::OffsetDateTime; use url::Url; use crate::{ - anthropic::experimental_pass_through::messages::transformation::resolve_anthropic_api_base, + anthropic::messages::transformation::resolve_anthropic_api_base, base_llm::chat::transformation::Error, }; diff --git a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs index a80cfbf28bd..9160cdf28ee 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/handler.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/handler.rs @@ -7,7 +7,7 @@ use litellm_types::{ use serde_json::Value; use crate::{ - anthropic::experimental_pass_through::messages::streaming_iterator::{ + anthropic::messages::streaming_iterator::{ AnthropicContentBlock, AnthropicContentBlockDelta, AnthropicMessagesStreamEvent, AnthropicStreamUsage, }, diff --git a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs index fd86c5ca25a..07ed6ba6ed1 100644 --- a/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/chat/transformation.rs @@ -11,9 +11,7 @@ use serde_json::{Map, Value, json}; use crate::{ anthropic::{ ANTHROPIC_OAUTH_TOKEN_PREFIX, - experimental_pass_through::messages::transformation::{ - complete_anthropic_url, resolve_anthropic_api_key, - }, + messages::transformation::{complete_anthropic_url, resolve_anthropic_api_key}, }, base_llm::chat::transformation::{ BaseConfig, Error, ProviderChatRequestData, ProviderChatResponseData, RequestAuth, diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs deleted file mode 100644 index ba63992f3cb..00000000000 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod messages; diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs b/litellm-rust/crates/llms/src/anthropic/messages/handler.rs similarity index 100% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs rename to litellm-rust/crates/llms/src/anthropic/messages/handler.rs diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs b/litellm-rust/crates/llms/src/anthropic/messages/headers.rs similarity index 100% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs rename to litellm-rust/crates/llms/src/anthropic/messages/headers.rs diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs b/litellm-rust/crates/llms/src/anthropic/messages/mod.rs similarity index 100% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs rename to litellm-rust/crates/llms/src/anthropic/messages/mod.rs diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs b/litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs similarity index 100% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/streaming_iterator.rs rename to litellm-rust/crates/llms/src/anthropic/messages/streaming_iterator.rs diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/thinking.rs b/litellm-rust/crates/llms/src/anthropic/messages/thinking.rs similarity index 100% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/thinking.rs rename to litellm-rust/crates/llms/src/anthropic/messages/thinking.rs diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs similarity index 100% rename from litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs rename to litellm-rust/crates/llms/src/anthropic/messages/transformation.rs diff --git a/litellm-rust/crates/llms/src/anthropic/mod.rs b/litellm-rust/crates/llms/src/anthropic/mod.rs index 755bc7d1907..a884c146dca 100644 --- a/litellm-rust/crates/llms/src/anthropic/mod.rs +++ b/litellm-rust/crates/llms/src/anthropic/mod.rs @@ -1,7 +1,8 @@ +pub mod common_utils; + pub mod batches; pub mod chat; -pub mod common_utils; pub mod count_tokens; -pub mod experimental_pass_through; +pub mod messages; pub const ANTHROPIC_OAUTH_TOKEN_PREFIX: &str = "sk-ant-oat"; diff --git a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs index c409f7f687e..137239bbeaf 100644 --- a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs @@ -6,7 +6,7 @@ use litellm_types::llms::anthropic_messages::{ }; use crate::{ - anthropic::experimental_pass_through::messages::transformation::{ + anthropic::messages::transformation::{ ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty, }, base_llm::{ diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs index 8db14687214..eff1dd1cf0b 100644 --- a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs @@ -4,8 +4,7 @@ use litellm_types::llms::anthropic_messages::{ }; use crate::{ - anthropic::experimental_pass_through::messages::thinking::ThinkingContext, - base_llm::chat::transformation::Error, + anthropic::messages::thinking::ThinkingContext, base_llm::chat::transformation::Error, }; pub type Headers = Vec<(String, String)>; From 2530255624904e8e9868bc29dbf2f39f54fd9b70 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Fri, 25 Sep 2026 19:27:48 -0700 Subject: [PATCH 07/10] test: stop CI tests from downloading tokenizer files and images (#43257) * test: load the embedding base image from a committed 100x100 PNG instead of downloading it * test: move the volcengine embedding test into tests/unit * test: check gpt2 and r50k_base tokenizer parity against committed tiktoken reference files * test: check hub tokenizer selection against an in-memory Hugging Face hub * test: serve image URLs from respx in the gemini tool-result and format-param tests * ci: drop the emptied legacy core-utils test path * test: cover the cohere and anthropic tokenizer paths in the hub tokenizer test * test: fetch every format-param image through respx and check its bytes reach the request * test: drop the gpt2 and r50k_base parity tests, which no litellm path uses * test: drop comments that restate assertions in the format-param test --- .github/workflows/test-unit.yml | 2 +- .../base_embedding_unit_tests.py | 7 +- .../litellm_core_utils/__init__.py | 0 .../litellm_core_utils/test_token_counter.py | 51 ------ .../litellm_core_utils/test_tokenizer.py | 20 --- .../llms/vertex_ai/gemini/__init__.py | 0 .../test_vertex_ai_gemini_transformation.py | 54 ------ .../test_litellm/llms/volcengine/__init__.py | 1 - tests/test_litellm/test_main.py | 164 ------------------ .../litellm_core_utils/test_token_counter.py | 92 ++++++++++ .../unit/litellm_core_utils/test_tokenizer.py | 12 +- .../test_vertex_ai_gemini_transformation.py | 42 +++++ .../volcengine/test_volcengine_embedding.py | 0 tests/unit/test_main.py | 100 +++++++++++ tests/white_100x100.png | Bin 0 -> 214 bytes 15 files changed, 239 insertions(+), 306 deletions(-) delete mode 100644 tests/test_litellm/litellm_core_utils/__init__.py delete mode 100644 tests/test_litellm/litellm_core_utils/test_token_counter.py delete mode 100644 tests/test_litellm/litellm_core_utils/test_tokenizer.py delete mode 100644 tests/test_litellm/llms/vertex_ai/gemini/__init__.py delete mode 100644 tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py delete mode 100644 tests/test_litellm/llms/volcengine/__init__.py delete mode 100644 tests/test_litellm/test_main.py rename tests/{test_litellm => unit}/llms/volcengine/test_volcengine_embedding.py (100%) create mode 100644 tests/white_100x100.png diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index d75213d37ea..2212b276b0d 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -61,7 +61,7 @@ jobs: - shard: core-utils artifact-name: core-utils - test-path: "tests/test_litellm/litellm_core_utils" + test-path: "" unit-flag: core-utils workers: 2 reruns: 1 diff --git a/tests/llm_translation/base_embedding_unit_tests.py b/tests/llm_translation/base_embedding_unit_tests.py index 1a88f0e9d6b..469416fc0cf 100644 --- a/tests/llm_translation/base_embedding_unit_tests.py +++ b/tests/llm_translation/base_embedding_unit_tests.py @@ -16,15 +16,12 @@ from litellm.utils import ( get_optional_params, get_optional_params_embeddings, ) -import requests import base64 +from pathlib import Path -# test_example.py from abc import ABC, abstractmethod -url = "https://dummyimage.com/100/100/fff&text=Test+image" -response = requests.get(url) -file_data = response.content +file_data = (Path(__file__).parent.parent / "white_100x100.png").read_bytes() encoded_file = base64.b64encode(file_data).decode("utf-8") base64_image = f"data:image/png;base64,{encoded_file}" diff --git a/tests/test_litellm/litellm_core_utils/__init__.py b/tests/test_litellm/litellm_core_utils/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py deleted file mode 100644 index 1e10b7e82b1..00000000000 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ /dev/null @@ -1,51 +0,0 @@ -import pytest -from litellm import create_pretrained_tokenizer -from tests.unit.litellm_core_utils.test_token_counter import token_counter - - -def test_tokenizers(): - try: - ### test the openai, claude, cohere and llama2 tokenizers. - ### The tokenizer value should be different for all - sample_text = "Hellö World, this is my input string! My name is ishaan CTO" - - # openai tokenizer - openai_tokens = token_counter(model="gpt-3.5-turbo", text=sample_text) - - # claude tokenizer - claude_tokens = token_counter(model="claude-3-5-haiku-20241022", text=sample_text) - - # cohere tokenizer - cohere_tokens = token_counter(model="command-nightly", text=sample_text) - - # llama2 tokenizer - llama2_tokens = token_counter(model="meta-llama/Llama-2-7b-chat", text=sample_text) - - # llama3 tokenizer (also testing custom tokenizer) - llama3_tokens_1 = token_counter(model="meta-llama/llama-3-70b-instruct", text=sample_text) - - try: - llama3_tokenizer = create_pretrained_tokenizer("Xenova/llama-3-tokenizer") - except Exception as e: - pytest.skip(f"custom tokenizer download failed (HF hub unreachable): {e}") - llama3_tokens_2 = token_counter(custom_tokenizer=llama3_tokenizer, text=sample_text) - - print( - f"openai tokens: {openai_tokens}; claude tokens: {claude_tokens}; cohere tokens: {cohere_tokens}; llama2 tokens: {llama2_tokens}; llama3 tokens: {llama3_tokens_1}" - ) - - # assert that all token values are different - # llama2 may fall back to the tiktoken tokenizer when the HuggingFace - # model hub is unreachable (e.g. in CI). In that case the count will - # equal the openai count and the differentiation assertion is skipped. - if openai_tokens == llama2_tokens: - pytest.skip("llama2 fell back to tiktoken (HF hub unreachable); skipping differentiation assertion") - assert llama2_tokens != llama3_tokens_1, "Token values are not different." - - assert llama3_tokens_1 == llama3_tokens_2, ( - "Custom tokenizer is not being used! It has been configured to use the same tokenizer as the built in llama3 tokenizer and the results should be the same." - ) - - print("test tokenizer: It worked!") - except Exception as e: - pytest.fail(f"An exception occured: {e}") diff --git a/tests/test_litellm/litellm_core_utils/test_tokenizer.py b/tests/test_litellm/litellm_core_utils/test_tokenizer.py deleted file mode 100644 index 2171044970c..00000000000 --- a/tests/test_litellm/litellm_core_utils/test_tokenizer.py +++ /dev/null @@ -1,20 +0,0 @@ -import pytest - -from tests.unit.litellm_core_utils.test_tokenizer import ( - UNICODE_TEXTS, - assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface, - assert_openai_encoding_matches_python, -) - -NETWORK_ENCODINGS = ("r50k_base", "gpt2") - - -@pytest.mark.parametrize("name", NETWORK_ENCODINGS) -@pytest.mark.parametrize("text", UNICODE_TEXTS) -def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None: - assert_openai_encoding_matches_python(name, text) - - -@pytest.mark.parametrize("name", ("gpt2",)) -def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: - assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name) diff --git a/tests/test_litellm/llms/vertex_ai/gemini/__init__.py b/tests/test_litellm/llms/vertex_ai/gemini/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py deleted file mode 100644 index d3a7ba7a1bd..00000000000 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ /dev/null @@ -1,54 +0,0 @@ -import pytest - -from litellm.litellm_core_utils.prompt_templates.factory import ( - convert_to_gemini_tool_call_result, -) -from litellm.types.llms.vertex_ai import BlobType - - -def test_convert_tool_response_with_url_image(): - """Test tool response with HTTP URL image (will download and convert).""" - # Use a publicly accessible test image URL - test_image_url = "https://via.placeholder.com/1x1.png" - - tool_message = { - "role": "tool", - "tool_call_id": "call_test456", - "content": [ - {"type": "text", "text": '{"url": "https://example.com"}'}, - {"type": "input_image", "image_url": test_image_url}, - ], - } - - last_message_with_tool_calls = { - "tool_calls": [ - { - "id": "call_test456", - "function": { - "name": "type_text_at", - "arguments": '{"x": 300, "y": 400, "text": "hello"}', - }, - } - ] - } - - try: - result = convert_to_gemini_tool_call_result(tool_message, last_message_with_tool_calls) - - assert isinstance(result, list), "Should return a parts list when media is present" - assert len(result) == 1, "Should return one function_response part" - result_part = result[0] - assert "function_response" in result_part - assert "inline_data" not in result_part - function_response = result_part["function_response"] - assert function_response["name"] == "type_text_at" - - # Check inline_data is nested under functionResponse.parts. - assert "parts" in function_response - assert len(function_response["parts"]) == 1 - inline_data: BlobType = function_response["parts"][0]["inline_data"] - assert "data" in inline_data - assert "mime_type" in inline_data - except Exception as e: - # Skip test if URL download fails (no internet connection, etc.) - pytest.skip(f"Failed to download image from URL: {e}") diff --git a/tests/test_litellm/llms/volcengine/__init__.py b/tests/test_litellm/llms/volcengine/__init__.py deleted file mode 100644 index 825e259b1fc..00000000000 --- a/tests/test_litellm/llms/volcengine/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# Volcengine tests diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py deleted file mode 100644 index 78728d6fd58..00000000000 --- a/tests/test_litellm/test_main.py +++ /dev/null @@ -1,164 +0,0 @@ -import json -import os - -import pytest - - -from unittest.mock import MagicMock, patch - -import litellm - - -async def _async_fake_bedrock_image_details(image_url): - return "ZmFrZS1pbWFnZQ==", "image/png" - - -@pytest.fixture(autouse=True) -def clear_client_cache(): - """ - Clear the HTTP client cache before each test to ensure mocks are used. - This prevents cached real clients from being reused across tests. - """ - cache = getattr(litellm, "in_memory_llm_clients_cache", None) - if cache is not None: - cache.flush_cache() - yield - if cache is not None: - cache.flush_cache() - - -@pytest.fixture(autouse=True) -def add_api_keys_to_env(monkeypatch): - monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-api03-1234567890") - monkeypatch.setenv("OPENAI_API_KEY", "sk-openai-api03-1234567890") - monkeypatch.setenv("AWS_ACCESS_KEY_ID", "my-fake-aws-access-key-id") - monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "my-fake-aws-secret-access-key") - monkeypatch.setenv("AWS_REGION", "us-east-1") - # Keep these transformation tests on the simple access-key path. A leaked - # session token or role/web-identity env var pushes Bedrock auth down a - # different branch and fails before the mocked HTTP client is exercised. - monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False) - monkeypatch.delenv("AWS_ROLE_ARN", raising=False) - monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) - - -@pytest.mark.parametrize( - "model", - [ - "gemini/gemini-1.5-flash", - "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", - "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", - "anthropic/claude-3-5-sonnet", - ], -) -@pytest.mark.parametrize("sync_mode", [True, False]) -@pytest.mark.asyncio -async def test_url_with_format_param(model, sync_mode, monkeypatch): - from litellm import acompletion, completion - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler - from litellm.litellm_core_utils.prompt_templates import factory as prompt_factory - - if sync_mode: - client = HTTPHandler() - else: - client = AsyncHTTPHandler() - - # This test is about request shaping, not live image downloads. Stub the - # URL->image conversion helpers so suite-level network/client state from - # earlier tests cannot prevent the mocked provider client from being hit. - fake_base64_image = "data:image/png;base64,ZmFrZS1pbWFnZQ==" - monkeypatch.setattr( - prompt_factory, "convert_url_to_base64", lambda url: fake_base64_image - ) - monkeypatch.setattr( - prompt_factory.BedrockImageProcessor, - "get_image_details", - staticmethod(lambda image_url: ("ZmFrZS1pbWFnZQ==", "image/png")), - ) - monkeypatch.setattr( - prompt_factory.BedrockImageProcessor, - "get_image_details_async", - staticmethod(_async_fake_bedrock_image_details), - ) - - args = { - "model": model, - "messages": [ - { - "role": "user", - "content": [ - { - "type": "image_url", - "image_url": { - "url": "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png", - "format": "image/png", - }, - }, - {"type": "text", "text": "Describe this image"}, - ], - } - ], - } - if model.startswith("gemini/"): - args["api_key"] = "test-api-key" - with patch.object(client, "post", new=MagicMock()) as mock_client: - try: - if sync_mode: - response = completion(**args, client=client) - else: - response = await acompletion(**args, client=client) - print(response) - except Exception as e: - pass - - mock_client.assert_called() - - print(mock_client.call_args.kwargs) - - if "data" in mock_client.call_args.kwargs: - json_str = mock_client.call_args.kwargs["data"] - else: - json_str = json.dumps(mock_client.call_args.kwargs["json"]) - - if isinstance(json_str, bytes): - json_str = json_str.decode("utf-8") - - print(f"type of json_str: {type(json_str)}") - - # Bedrock models convert URLs to base64, while direct Anthropic models support URLs - # bedrock/invoke models use Anthropic messages API which supports URLs - if model.startswith("bedrock/invoke/"): - # bedrock/invoke should convert URLs to base64 (doesn't support URL references) - # URL should NOT be in the JSON (it should be converted to base64) - assert "https://awsmp-logos.s3.amazonaws.com" not in json_str - # Should have base64 data in the source (type="base64", not type="url") - assert '"type":"base64"' in json_str or '"type": "base64"' in json_str - # Should have "data" field containing base64 content - assert '"data"' in json_str - elif model.startswith("bedrock/"): - # Regular Bedrock models should convert URLs to base64 (uses "bytes" field) - # URL should NOT be in the JSON (it should be converted to base64) - assert "https://awsmp-logos.s3.amazonaws.com" not in json_str - # Should have "bytes" field (Bedrock uses "bytes" not "base64" in the field name) - assert '"bytes"' in json_str or '"bytes":' in json_str - elif model.startswith("anthropic/"): - # Direct Anthropic models should pass HTTPS URLs directly (HTTP URLs are converted to base64) - # Since we're using HTTPS URL, it should be passed as-is - assert "https://awsmp-logos.s3.amazonaws.com" in json_str - # For Anthropic, URL references use "url" type, not base64 - assert '"type":"url"' in json_str or '"type": "url"' in json_str - else: - # For other models, check format parameter is respected - assert "png" in json_str - assert "jpeg" not in json_str - - -@pytest.fixture(autouse=True) -def set_openrouter_api_key(): - original_api_key = os.environ.get("OPENROUTER_API_KEY") - os.environ["OPENROUTER_API_KEY"] = "fake-key-for-testing" - yield - if original_api_key is not None: - os.environ["OPENROUTER_API_KEY"] = original_api_key - else: - del os.environ["OPENROUTER_API_KEY"] diff --git a/tests/unit/litellm_core_utils/test_token_counter.py b/tests/unit/litellm_core_utils/test_token_counter.py index b1a14e61b96..ae9b30d862b 100644 --- a/tests/unit/litellm_core_utils/test_token_counter.py +++ b/tests/unit/litellm_core_utils/test_token_counter.py @@ -3,16 +3,22 @@ import asyncio import base64 import importlib +import json +import os +import subprocess +import sys import threading import time import traceback from concurrent.futures import Future, wait +from pathlib import Path from typing import Final from unittest.mock import MagicMock import anyio.to_thread import pytest import tiktoken +from tokenizers import Regex, Tokenizer, models, pre_tokenizers from unittest.mock import AsyncMock, patch @@ -1439,3 +1445,89 @@ def test_high_detail_image_token_upper_bound_covers_every_image_size(width: int, def test_high_detail_image_token_upper_bound_is_reached_by_the_largest_high_res_image() -> None: assert calculate_img_tokens(_png_data_url(2000, 768), mode="high") == high_detail_image_token_upper_bound() assert calculate_img_tokens(_png_data_url(1, 1), mode="high") < high_detail_image_token_upper_bound() + + +HUB_TOKENIZER_SCRIPT: Final = """ +import json +import sys +sys.path.insert(0, sys.argv[1]) +import httpx +import huggingface_hub +import litellm +served = json.loads(sys.argv[2]) +text = sys.argv[3] +requested = [] +def handle(request): + repo = request.url.path.lstrip("/").split("/resolve/")[0] + if repo not in served or not request.url.path.endswith("/tokenizer.json"): + return httpx.Response(404) + requested.append(repo) + payload = served[repo].encode() + headers = {"content-length": str(len(payload)), "etag": '"fixture"', "x-repo-commit": "a" * 40} + return httpx.Response(200, headers=headers, content=payload if request.method == "GET" else b"") +huggingface_hub.set_client_factory(lambda: httpx.Client(transport=httpx.MockTransport(handle))) +litellm.cohere_models = {"command-r-v1"} +litellm.anthropic_models = {"claude-2"} +custom = litellm.create_pretrained_tokenizer("Xenova/llama-3-tokenizer") +print(json.dumps({ + "llama2": litellm.token_counter(model="meta-llama/Llama-2-7b-chat", text=text), + "llama3": litellm.token_counter(model="meta-llama/llama-3-70b-instruct", text=text), + "cohere": litellm.token_counter(model="command-r-v1", text=text), + "anthropic": litellm.token_counter(model="claude-2", text=text), + "custom": litellm.token_counter(custom_tokenizer=custom, text=text), + "requested": sorted(set(requested)), +})) +""" + + +def _word_level_tokenizer_json(pre_tokenizer: pre_tokenizers.PreTokenizer) -> str: + tokenizer: Final = Tokenizer(models.WordLevel(vocab={"[UNK]": 0}, unk_token="[UNK]")) + tokenizer.pre_tokenizer = pre_tokenizer + return tokenizer.to_str() + + +def test_token_counter_uses_the_tokenizer_of_each_model_family_and_of_a_custom_tokenizer(tmp_path: Path) -> None: + sample: Final = "Tokenizers disagree: anthropic, tiktoken; llama-2 & llama-3!" + served: Final = { + "hf-internal-testing/llama-tokenizer": _word_level_tokenizer_json(pre_tokenizers.WhitespaceSplit()), + "Xenova/llama-3-tokenizer": _word_level_tokenizer_json(pre_tokenizers.Split(Regex("."), "isolated")), + "Xenova/c4ai-command-r-v01-tokenizer": _word_level_tokenizer_json(pre_tokenizers.Whitespace()), + } + expected: Final = {repo: len(Tokenizer.from_str(payload).encode(sample).ids) for repo, payload in served.items()} + anthropic_count: Final = len(Tokenizer.from_str(claude_json_str).encode(sample).ids) + tiktoken_count: Final = litellm.token_counter(model="gpt-3.5-turbo", text=sample) + assert len({*expected.values(), anthropic_count, tiktoken_count}) == len(expected) + 2 + + result: Final = subprocess.run( + [ + sys.executable, + "-I", + "-c", + HUB_TOKENIZER_SCRIPT, + str(Path(litellm.__file__).parent.parent), + json.dumps(served), + sample, + ], + capture_output=True, + text=True, + timeout=60, + env={ + **os.environ, + "HF_HOME": str(tmp_path / "home"), + "HF_HUB_CACHE": str(tmp_path / "cache"), + "HF_ENDPOINT": "http://127.0.0.1:9", + "HF_HUB_OFFLINE": "0", + "LITELLM_LOCAL_MODEL_COST_MAP": "True", + }, + ) + + assert result.returncode == 0, result.stdout + result.stderr + counts: Final = json.loads(result.stdout.strip().splitlines()[-1]) + assert counts == { + "llama2": expected["hf-internal-testing/llama-tokenizer"], + "llama3": expected["Xenova/llama-3-tokenizer"], + "cohere": expected["Xenova/c4ai-command-r-v01-tokenizer"], + "anthropic": anthropic_count, + "custom": expected["Xenova/llama-3-tokenizer"], + "requested": sorted(served), + } diff --git a/tests/unit/litellm_core_utils/test_tokenizer.py b/tests/unit/litellm_core_utils/test_tokenizer.py index a9005ff6a86..9d08442b164 100644 --- a/tests/unit/litellm_core_utils/test_tokenizer.py +++ b/tests/unit/litellm_core_utils/test_tokenizer.py @@ -17,17 +17,13 @@ from litellm.utils import claude_json_str from tests.unit.litellm_core_utils.test_decode_special_tokens import TOKENIZER_JSON -OFFLINE_ENCODINGS: Final = ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "o200k_harmony") +ENCODINGS: Final = ("cl100k_base", "o200k_base", "p50k_base", "p50k_edit", "o200k_harmony") UNICODE_TEXTS: Final = ("hello world", "café 漢字 🙂", "", "a\ud800b", "\ud83d\ude42", "🙂\ud83d\ude42\udfff", " " * 64) -@pytest.mark.parametrize("name", OFFLINE_ENCODINGS) +@pytest.mark.parametrize("name", ENCODINGS) @pytest.mark.parametrize("text", UNICODE_TEXTS) def test_openai_encoding_matches_python_unicode_and_batches(name: str, text: str) -> None: - assert_openai_encoding_matches_python(name, text) - - -def assert_openai_encoding_matches_python(name: str, text: str) -> None: reference: Final = tiktoken.get_encoding(name) encoding: Final = OpenAIEncoding.from_tiktoken(name) expected: Final = reference.encode(text) @@ -309,10 +305,6 @@ def test_huggingface_batch_sequence_containers_match_python(is_pretokenized: boo @pytest.mark.parametrize("name", ("cl100k_base", "o200k_base", "p50k_edit")) def test_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: - assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name) - - -def assert_openai_encoding_exposes_the_tiktoken_vocabulary_surface(name: str) -> None: reference: Final = tiktoken.get_encoding(name) encoding: Final = OpenAIEncoding.from_tiktoken(name) text: Final = "hello fanta" diff --git a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 4f23ac1773a..0b37e033023 100644 --- a/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/unit/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1,7 +1,12 @@ import base64 +from pathlib import Path +from typing import Final +import httpx import pytest +import respx +import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( convert_to_gemini_tool_call_result, ) @@ -2727,3 +2732,40 @@ def test_gemini_server_side_tool_signature_not_duplicated_on_text(): assert "thoughtSignature" not in text_part tool_call_part = next(p for p in parts if "toolCall" in p) assert tool_call_part["thoughtSignature"] == "server_side_signature" + + +WHITE_PNG: Final = (Path(__file__).parents[4] / "white_100x100.png").read_bytes() + + +@respx.mock +def test_convert_tool_response_with_url_image(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "user_url_validation", False) + image_url: Final = "https://tool-result-images.test/gemini-tool-response.png" + respx.get(image_url).mock(return_value=httpx.Response(200, content=WHITE_PNG, headers={"content-type": "image/png"})) + tool_message: Final = { + "role": "tool", + "tool_call_id": "call_test456", + "content": [ + {"type": "text", "text": '{"url": "https://example.com"}'}, + {"type": "input_image", "image_url": image_url}, + ], + } + last_message_with_tool_calls: Final = { + "tool_calls": [ + { + "id": "call_test456", + "function": {"name": "type_text_at", "arguments": '{"x": 300, "y": 400, "text": "hello"}'}, + } + ] + } + + result: Final = convert_to_gemini_tool_call_result(tool_message, last_message_with_tool_calls) + + assert isinstance(result, list) + assert len(result) == 1 + assert "inline_data" not in result[0] + function_response: Final = result[0]["function_response"] + assert function_response["name"] == "type_text_at" + assert len(function_response["parts"]) == 1 + inline_data: Final[BlobType] = function_response["parts"][0]["inline_data"] + assert inline_data == {"data": base64.b64encode(WHITE_PNG).decode(), "mime_type": "image/png"} diff --git a/tests/test_litellm/llms/volcengine/test_volcengine_embedding.py b/tests/unit/llms/volcengine/test_volcengine_embedding.py similarity index 100% rename from tests/test_litellm/llms/volcengine/test_volcengine_embedding.py rename to tests/unit/llms/volcengine/test_volcengine_embedding.py diff --git a/tests/unit/test_main.py b/tests/unit/test_main.py index c06216e4f4e..57200a79a8c 100644 --- a/tests/unit/test_main.py +++ b/tests/unit/test_main.py @@ -17,6 +17,7 @@ import respx import urllib.parse from importlib import import_module +from pathlib import Path from unittest.mock import MagicMock, patch import litellm @@ -56,6 +57,9 @@ def add_api_keys_to_env(monkeypatch): monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) +WHITE_PNG: Final = (Path(__file__).parents[1] / "white_100x100.png").read_bytes() + + @pytest.fixture def openai_api_response(): mock_response_data = { @@ -213,6 +217,102 @@ async def test_url_with_format_param_openai(model, sync_mode): assert "format" not in json_str +@pytest.mark.parametrize( + "model", + [ + "gemini/gemini-1.5-flash", + "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "bedrock/invoke/anthropic.claude-haiku-4-5-20251001-v1:0", + "anthropic/claude-3-5-sonnet", + ], +) +@pytest.mark.parametrize("sync_mode", [True, False]) +@pytest.mark.asyncio +async def test_url_with_format_param(model, sync_mode, monkeypatch): + from litellm import acompletion, completion + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + + if sync_mode: + client = HTTPHandler() + else: + client = AsyncHTTPHandler() + + image_url: Final = ( + "https://awsmp-logos.s3.amazonaws.com/seller-xw5kijmvmzasy/c233c9ade2ccb5491072ae232c814942.png" + f"?case={sync_mode}-{model}" + ) + args = { + "model": model, + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": { + "url": image_url, + "format": "image/png", + }, + }, + {"type": "text", "text": "Describe this image"}, + ], + } + ], + } + if model.startswith("gemini/"): + args["api_key"] = "test-api-key" + monkeypatch.setattr(litellm, "user_url_validation", False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "module_level_aclient", AsyncHTTPHandler(transport=httpx.AsyncHTTPTransport())) + with ( + respx.mock(assert_all_called=False) as image_host, + patch.object(client, "post", new=MagicMock()) as mock_client, + ): + image_route = image_host.get(image_url).mock( + return_value=httpx.Response(200, content=WHITE_PNG, headers={"content-type": "image/png"}) + ) + try: + if sync_mode: + response = completion(**args, client=client) + else: + response = await acompletion(**args, client=client) + print(response) + except Exception as e: + pass + + mock_client.assert_called() + + print(mock_client.call_args.kwargs) + + if "data" in mock_client.call_args.kwargs: + json_str = mock_client.call_args.kwargs["data"] + else: + json_str = json.dumps(mock_client.call_args.kwargs["json"]) + + if isinstance(json_str, bytes): + json_str = json_str.decode("utf-8") + + print(f"type of json_str: {type(json_str)}") + + if model.startswith("bedrock/invoke/"): + assert "https://awsmp-logos.s3.amazonaws.com" not in json_str + assert '"type":"base64"' in json_str or '"type": "base64"' in json_str + assert '"data"' in json_str + elif model.startswith("bedrock/"): + assert "https://awsmp-logos.s3.amazonaws.com" not in json_str + assert '"bytes"' in json_str or '"bytes":' in json_str + elif model.startswith("anthropic/"): + assert "https://awsmp-logos.s3.amazonaws.com" in json_str + assert '"type":"url"' in json_str or '"type": "url"' in json_str + else: + assert "png" in json_str + assert "jpeg" not in json_str + + fetches_image: Final = not model.startswith("anthropic/") + assert image_route.called is fetches_image + assert (base64.b64encode(WHITE_PNG).decode() in json_str) is fetches_image + + def test_bedrock_latency_optimized_inference(): from litellm.llms.custom_httpx.http_handler import HTTPHandler diff --git a/tests/white_100x100.png b/tests/white_100x100.png new file mode 100644 index 0000000000000000000000000000000000000000..fdd268ded88d11837418170a350a045d3395e9d7 GIT binary patch literal 214 zcmeAS@N?(olHy`uVBq!ia0vp^DIm~-MNWmM-(43$OW{8}a1ZHrhoCGsifoebuCXiwvqY Date: Fri, 25 Sep 2026 19:41:53 -0700 Subject: [PATCH 08/10] fix(rust): preserve nested optional import failures (#43265) * fix(rust): preserve nested optional import failures Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(rust): restore Python modules after settings tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../python-bridge/src/python_settings.rs | 175 +++++++++++++++++- 1 file changed, 172 insertions(+), 3 deletions(-) diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index f03e5fdce7f..6f8388471dc 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -45,8 +45,13 @@ impl PythonSettings { pub(crate) fn read_or_unset(self, py: Python<'_>) -> PyResult>> { match self.read(py) { Ok(snapshot) => Ok(Some(snapshot)), - Err(error) if error.is_instance_of::(py) => Ok(None), - Err(error) => Err(error), + Err(error) => { + if missing_module(py, &error, "litellm")? { + Ok(None) + } else { + Err(error) + } + } } } @@ -56,9 +61,24 @@ impl PythonSettings { } } +fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult { + if !error.is_instance_of::(py) { + return Ok(false); + } + Ok(error + .value(py) + .getattr("name")? + .extract::>()? + .is_some_and(|name| name == expected)) +} + #[cfg(test)] mod tests { - use pyo3::{exceptions::PyRuntimeError, prelude::*, types::PyDict}; + use pyo3::{ + exceptions::{PyImportError, PyModuleNotFoundError, PyRuntimeError}, + prelude::*, + types::PyDict, + }; use super::PythonSettings; use crate::coercion::FieldSpec; @@ -150,4 +170,153 @@ values = (Descriptor(), SimpleNamespace(flag=Truth())) ); }); } + + #[test] + fn read_or_unset_returns_none_when_litellm_is_missing() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +import sys +class MissingLitellm: + def find_spec(self, fullname, path=None, target=None): + if fullname == 'litellm': + raise ModuleNotFoundError('No module named litellm', name='litellm') +finder = MissingLitellm() +previous_litellm = sys.modules.get('litellm') +had_litellm = 'litellm' in sys.modules +sys.meta_path.insert(0, finder) +sys.modules.pop('litellm', None) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let result = PythonSettings::Http.read_or_unset(py); + assert!(result.unwrap().is_none()); + py.run( + c" +sys.meta_path.remove(finder) +if had_litellm: + sys.modules['litellm'] = previous_litellm +else: + sys.modules.pop('litellm', None) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn read_or_unset_propagates_nested_module_not_found_errors() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +import sys +import types +previous_modules = { + name: sys.modules[name] + for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.settings') + if name in sys.modules +} +litellm = types.ModuleType('litellm') +litellm.__path__ = [] +rust_bridge = types.ModuleType('litellm.rust_bridge') +rust_bridge.__path__ = [] +settings = types.ModuleType('litellm.rust_bridge.settings') +def http_settings(): + raise ModuleNotFoundError('No module named certifi', name='certifi') +settings.http_settings = http_settings +litellm.rust_bridge = rust_bridge +rust_bridge.settings = settings +sys.modules['litellm'] = litellm +sys.modules['litellm.rust_bridge'] = rust_bridge +sys.modules['litellm.rust_bridge.settings'] = settings +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let error = match PythonSettings::Http.read_or_unset(py) { + Ok(_) => panic!("nested module errors must propagate"), + Err(error) => error, + }; + assert!(error.is_instance_of::(py)); + assert_eq!( + error + .value(py) + .getattr("name") + .unwrap() + .extract::() + .unwrap(), + "certifi" + ); + py.run( + c" +for name in ('litellm.rust_bridge.settings', 'litellm.rust_bridge', 'litellm'): + sys.modules.pop(name, None) +sys.modules.update(previous_modules) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } + + #[test] + fn read_or_unset_propagates_import_errors() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c" +import sys +import types +previous_modules = { + name: sys.modules[name] + for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.settings') + if name in sys.modules +} +litellm = types.ModuleType('litellm') +litellm.__path__ = [] +rust_bridge = types.ModuleType('litellm.rust_bridge') +rust_bridge.__path__ = [] +settings = types.ModuleType('litellm.rust_bridge.settings') +def http_settings(): + raise ImportError('cannot import name setting') +settings.http_settings = http_settings +litellm.rust_bridge = rust_bridge +rust_bridge.settings = settings +sys.modules['litellm'] = litellm +sys.modules['litellm.rust_bridge'] = rust_bridge +sys.modules['litellm.rust_bridge.settings'] = settings +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let error = match PythonSettings::Http.read_or_unset(py) { + Ok(_) => panic!("import errors must propagate"), + Err(error) => error, + }; + assert!(error.is_instance_of::(py)); + assert_eq!(error.to_string(), "ImportError: cannot import name setting"); + py.run( + c" +for name in ('litellm.rust_bridge.settings', 'litellm.rust_bridge', 'litellm'): + sys.modules.pop(name, None) +sys.modules.update(previous_modules) +", + Some(&locals), + Some(&locals), + ) + .unwrap(); + }); + } } From c822c7fffa19b1292a789d147201c694b1b059da Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 02:42:08 +0000 Subject: [PATCH 09/10] ci: drop main and litellm_* branch filters from the CircleCI litellm-main workflows (#43272) Co-authored-by: yuneng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .circleci/config.yml | 121 ++++++++++++------------------------------- 1 file changed, 32 insertions(+), 89 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index 3ff8061fb48..d9c85cfa042 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -3329,22 +3329,12 @@ workflows: matrix: parameters: suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser] - filters: - branches: - only: - - main - - /litellm_.*/ - integration_contracts: name: integration-<< matrix.suite >>-replica matrix: parameters: suite: [management, database] mode: [replica] - filters: - branches: - only: - - main - - /litellm_.*/ build_and_test: unless: or: @@ -3352,101 +3342,60 @@ workflows: - not: equal: ["", << pipeline.parameters.routing_parity_base >>] jobs: - - using_litellm_on_windows: - filters: &main_branches - branches: - only: - - main - - /litellm_.*/ - - unit: - filters: *main_branches + - using_litellm_on_windows + - unit - provider_replay_harness - - base_sdk_install: - filters: *main_branches - - local_testing_part1: - filters: *main_branches - - local_testing_part2: - filters: *main_branches - - langfuse_logging_unit_tests: - filters: *main_branches - - litellm_assistants_api_testing: - filters: *main_branches - - litellm_router_testing: - filters: *main_branches - - litellm_router_unit_testing: - filters: *main_branches - - auth_ui_unit_tests: - filters: *main_branches - - build_docker_database_image: - filters: *main_branches - - e2e_ui_testing: - filters: *main_branches - - e2e_ui_testing_server_root_path: - filters: *main_branches + - base_sdk_install + - local_testing_part1 + - local_testing_part2 + - langfuse_logging_unit_tests + - litellm_assistants_api_testing + - litellm_router_testing + - litellm_router_unit_testing + - auth_ui_unit_tests + - build_docker_database_image + - e2e_ui_testing + - e2e_ui_testing_server_root_path - build_and_test: requires: - build_docker_database_image - filters: *main_branches - e2e_openai_endpoints: requires: - build_docker_database_image - filters: *main_branches - proxy_logging_guardrails_model_info_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_spend_accuracy_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_multi_instance_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_store_model_in_db_tests: requires: - build_docker_database_image - filters: *main_branches - - proxy_build_from_pip_tests: - filters: *main_branches + - proxy_build_from_pip_tests - proxy_pass_through_endpoint_tests: requires: - build_docker_database_image - filters: *main_branches - proxy_e2e_anthropic_messages_tests: requires: - build_docker_database_image - filters: *main_branches - - llm_translation_testing: - filters: *main_branches - - realtime_translation_testing: - filters: *main_branches - - agent_testing: - filters: *main_branches - - guardrails_testing: - filters: *main_branches - - google_generate_content_endpoint_testing: - filters: *main_branches - - llm_responses_api_testing: - filters: *main_branches - - ocr_testing: - filters: *main_branches - - search_testing: - filters: *main_branches - - batches_testing: - filters: *main_branches - - litellm_utils_testing: - filters: *main_branches - - pass_through_unit_testing: - filters: *main_branches - - image_gen_testing: - filters: *main_branches - - logging_testing: - filters: *main_branches - - audio_testing: - filters: *main_branches - - redis_caching_unit_tests: - filters: *main_branches + - llm_translation_testing + - realtime_translation_testing + - agent_testing + - guardrails_testing + - google_generate_content_endpoint_testing + - llm_responses_api_testing + - ocr_testing + - search_testing + - batches_testing + - litellm_utils_testing + - pass_through_unit_testing + - image_gen_testing + - logging_testing + - audio_testing + - redis_caching_unit_tests - upload-coverage: requires: - realtime_translation_testing @@ -3471,18 +3420,12 @@ workflows: - db_migration_disable_update_check: requires: - build_docker_database_image - filters: *main_branches - - installing_litellm_on_python: - filters: *main_branches - - installing_litellm_on_python_3_13: - filters: *main_branches - - installing_litellm_on_python_v2_migration_resolver: - filters: *main_branches + - installing_litellm_on_python + - installing_litellm_on_python_3_13 + - installing_litellm_on_python_v2_migration_resolver - helm_chart_testing: requires: - build_docker_database_image - filters: *main_branches - test_bad_database_url: requires: - build_docker_database_image - filters: *main_branches From d08746feb15d88d97f9d978d2dfba7ef4bb6d607 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 26 Sep 2026 02:42:45 +0000 Subject: [PATCH 10/10] feat(proxy): email alerts at configured percentages of a team member budget (#42665) * feat(proxy): email alerts at configured percentages of a team member budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(alerting): label team member budget crossings as team member budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(auth): cover the team member alert dispatch from _check_team_member_budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(email): drop the emoji from the team member budget alert template Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): ignore team member alert thresholds outside 1 to 100 on both the backend and the dashboard Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): bound team member alert threshold key length before int parsing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): drop the legacy covers marker from the team member alert test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(team): reject malformed team_member_max_budget_alert_emails on team writes Thresholds outside 1-100, non-list recipients, and invalid emails now return 422 on /team/new, /team/update and PATCH /team/{id} instead of being stored and silently ignored. The value is stored canonically. Read-side LiteLLM_TeamTable is unchanged, and the PATCH body stays a raw merge patch so a null threshold still deletes it. * fix(auth): enforce and alert on team member budgets only in common_checks The builder re-checked the team member budget inline before common_checks ran the same check, so one request that crossed a team_member_max_budget_alert_emails threshold dispatched two alerts. Drop the inline check; common_checks is the single authorization point and already covers per-member rows, the team default member budget, zero-cost skips and the cross-pod spend counter. Its 422 message now uses the TeamMember=user:team form the builder and budget reservation already returned. * Revert "fix(team): reject malformed team_member_max_budget_alert_emails on team writes" This reverts commit 703e754b4615a292b0a947b142751ae33524e329. * fix(alerting): keep BaseBudgetAlertType.get_event_message zero-arg Requiring user_info broke existing callers and out-of-tree subclasses. The team member label now comes from SlackAlerting.budget_alerts, so the interface and its Readme are unchanged from main. * fix(mcp): keep team member budget enforcement on the MCP OAuth auth dependency The MCP OAuth dependency stops at _user_api_key_auth_builder and never reaches common_checks, so removing the builder's inline member budget check would have let over-budget members through there. Enforce it explicitly for that caller. * fix(auth): keep main's team member budget enforcement, alert once per request Restore the builder's team member budget check and 422 message exactly as on main and drop the MCP-only gate. The builder sends the member alert only on the request it rejects; common_checks sends it for requests that get past the builder, so no request alerts twice. * test(integration): read team member alert deliveries without a shared accumulator * test(integration): match team member alert deliveries by subject so other alerts cannot race the count * refactor(proxy): build the team member alert threshold config without mutable collections Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): collapse the alert recipient isinstance checks into one call Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): read the SMTP sink through lock-guarded snapshots and assert the exact deliveries --------- Co-authored-by: ryan Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../send_emails/base_email.py | 53 +++++-- .../SlackAlerting/budget_alert_types.py | 2 + .../SlackAlerting/slack_alerting.py | 6 +- .../integrations/email_templates/templates.py | 22 +++ litellm/proxy/auth/auth_checks.py | 77 ++++++++- litellm/proxy/auth/user_api_key_auth.py | 14 ++ tests/integration/_support/mail.py | 127 +++++++++++++++ .../spend/test_team_member_budget_alerts.py | 93 +++++++++++ .../proxy/auth/test_auth_checks.py | 137 ++++++++++++++++ .../proxy/auth/test_user_api_key_auth.py | 147 ++++++++++++++++++ .../send_emails/test_base_email.py | 41 +++++ .../SlackAlerting/test_budget_alert_types.py | 33 +++- .../SlackAlerting/test_slack_alerting.py | 27 ++++ .../src/components/team/TeamInfo.test.tsx | 97 ++++++++++++ .../src/components/team/TeamInfo.tsx | 104 +++++++++++++ .../team/teamMemberBudgetAlertEmails.test.ts | 94 +++++++++++ .../team/teamMemberBudgetAlertEmails.ts | 57 +++++++ 17 files changed, 1113 insertions(+), 18 deletions(-) create mode 100644 tests/integration/_support/mail.py create mode 100644 tests/integration/spend/test_team_member_budget_alerts.py create mode 100644 ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts create mode 100644 ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py index 4be09670e92..6e33d9f1bf3 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/base_email.py @@ -32,6 +32,7 @@ from litellm.integrations.email_templates.key_rotated_email import ( from litellm.integrations.email_templates.templates import ( MAX_BUDGET_ALERT_EMAIL_TEMPLATE, SOFT_BUDGET_ALERT_EMAIL_TEMPLATE, + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE, TEAM_SOFT_BUDGET_ALERT_EMAIL_TEMPLATE, ) from litellm.integrations.email_templates.user_invitation_email import ( @@ -48,6 +49,12 @@ from litellm.secret_managers.main import get_secret_bool from litellm.types.integrations.slack_alerting import LITELLM_LOGO_URL +def _max_budget_alert_id(user_info: CallInfo) -> str: + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER: + return f"team_member:{user_info.user_id}:{user_info.team_id}" + return user_info.token or user_info.user_id or "default_id" + + def _parse_email_list(raw) -> List[str]: """Parse emails from a list or comma-separated string.""" if isinstance(raw, list): @@ -373,17 +380,31 @@ class BaseEmailLogger(CustomLogger): greeting = html.escape( event.user_email or event.key_alias or event.token or "" ) - email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format( - email_logo_url=email_params.logo_url, - recipient_email=greeting, - percentage=percentage, - spend=spend_str, - max_budget=max_budget_str, - alert_threshold=alert_threshold_str, - base_url=email_params.base_url, - email_support_contact=email_params.support_contact, - email_footer=email_params.signature, - ) + if event.event_group == Litellm_EntityType.TEAM_MEMBER: + email_html_content = TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format( + email_logo_url=email_params.logo_url, + member=html.escape(event.user_email or event.user_id or ""), + team_alias=html.escape(event.team_alias or event.team_id or ""), + percentage=percentage, + spend=spend_str, + max_budget=max_budget_str, + alert_threshold=alert_threshold_str, + base_url=email_params.base_url, + email_support_contact=email_params.support_contact, + email_footer=email_params.signature, + ) + else: + email_html_content = MAX_BUDGET_ALERT_EMAIL_TEMPLATE.format( + email_logo_url=email_params.logo_url, + recipient_email=greeting, + percentage=percentage, + spend=spend_str, + max_budget=max_budget_str, + alert_threshold=alert_threshold_str, + base_url=email_params.base_url, + email_support_contact=email_params.support_contact, + email_footer=email_params.signature, + ) await self.send_email( from_email=self.DEFAULT_LITELLM_EMAIL, to_email=recipient_emails, @@ -607,7 +628,7 @@ class BaseEmailLogger(CustomLogger): if user_info.spend < threshold_amount: continue - _id = user_info.token or user_info.user_id or "default_id" + _id = _max_budget_alert_id(user_info) _cache_key = ( f"email_budget_alerts:max_budget_alert:{threshold_pct}:{_id}" ) @@ -618,7 +639,7 @@ class BaseEmailLogger(CustomLogger): emails.append(user_info.user_email) if not emails: verbose_proxy_logger.warning( - "No recipients for %d%% threshold on key %s, skipping alert", + "No recipients for %d%% threshold on %s, skipping alert", threshold_pct, _id, ) @@ -633,7 +654,11 @@ class BaseEmailLogger(CustomLogger): if send_count is not None and send_count > 1: continue - event_message = f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached" + event_message = ( + f"Team Member Budget Alert - {threshold_pct}% of Team Member Budget Reached" + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER + else f"Max Budget Alert - {threshold_pct}% of Maximum Budget Reached" + ) webhook_event = WebhookEvent( event="max_budget_alert", event_message=event_message, diff --git a/litellm/integrations/SlackAlerting/budget_alert_types.py b/litellm/integrations/SlackAlerting/budget_alert_types.py index f35ff7b5f82..4fe833acecc 100644 --- a/litellm/integrations/SlackAlerting/budget_alert_types.py +++ b/litellm/integrations/SlackAlerting/budget_alert_types.py @@ -63,6 +63,8 @@ class TokenBudgetAlert(BaseBudgetAlertType): return "Key Budget: " def get_id(self, user_info: CallInfo) -> str: + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER: + return f"team_member:{user_info.user_id}:{user_info.team_id}" return user_info.token or "default_id" diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index 17ec3ed787d..7c608aac8d9 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -555,7 +555,11 @@ class SlackAlerting(CustomBatchLogger): budget_alert_class: Final = get_budget_alert_type(type) _id: Final = budget_alert_class.get_id(user_info) user_info_str: Final = self._get_user_info_str(user_info) - event_message = budget_alert_class.get_event_message() + event_message = ( + "Team Member Budget: " + if user_info.event_group == Litellm_EntityType.TEAM_MEMBER + else budget_alert_class.get_event_message() + ) # Set default event unless we're in projected_limit_exceeded event: ( diff --git a/litellm/integrations/email_templates/templates.py b/litellm/integrations/email_templates/templates.py index 935067c97fc..2bd079ef15d 100644 --- a/litellm/integrations/email_templates/templates.py +++ b/litellm/integrations/email_templates/templates.py @@ -131,3 +131,25 @@ MAX_BUDGET_ALERT_EMAIL_TEMPLATE: Final = """ {email_footer} """ + +TEAM_MEMBER_MAX_BUDGET_ALERT_EMAIL_TEMPLATE: Final = """ + LiteLLM Logo + +

Hi,
+ + Team member {member} has reached {percentage}% of their team member budget in team {team_alias}.

+ + Current Spend: {spend}
+ Team Member Budget: {max_budget}
+ Alert Threshold: {alert_threshold} ({percentage}%)
+ +

+ Warning: Once this member reaches their team member budget of {max_budget}, their requests in this team will be rejected. +

+ + You can view usage and manage team member budgets in the LiteLLM Dashboard.

+ + If you have any questions, please send an email to {email_support_contact}

+ + {email_footer} +""" diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index f19a8055ae6..12d420141f1 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -18,7 +18,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, Protocol, TypeAlias from fastapi import HTTPException, Request, status -from pydantic import BaseModel, TypeAdapter +from pydantic import BaseModel, TypeAdapter, ValidationError from typing_extensions import NotRequired, ReadOnly, Required, TypedDict, Unpack import litellm @@ -5682,6 +5682,64 @@ async def _virtual_key_max_budget_alert_check( ) +TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY: Final = "team_member_max_budget_alert_emails" +_TEAM_MEMBER_ALERT_CONFIG_ADAPTER: Final[TypeAdapter[Mapping[str, object]]] = TypeAdapter(Mapping[str, object]) + + +def _is_valid_alert_threshold_pct(pct: str) -> bool: + return pct.isdigit() and len(pct) <= 3 and 1 <= int(pct) <= 100 + + +def _alert_recipients(raw: object) -> Sequence[str] | None: + if isinstance(raw, (str, Sequence)): + return _parse_email_list(raw) + return None + + +def _valid_alert_threshold_config(raw_config: object) -> Mapping[str, str | Sequence[object] | None] | None: + try: + config: Final = _TEAM_MEMBER_ALERT_CONFIG_ADAPTER.validate_python(raw_config) + except ValidationError: + return None + return MappingProxyType( + {pct: _alert_recipients(emails) for pct, emails in config.items() if _is_valid_alert_threshold_pct(pct)} + ) + + +def _team_member_max_budget_alert_check( + team_id: str, + team_alias: str | None, + team_metadata: Mapping[str, object] | None, + organization_id: str | None, + user_id: str, + user_email: str | None, + proxy_logging_obj: ProxyLogging, + spend: float, + max_budget: float, +) -> None: + raw_config: Final = team_metadata.get(TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY) if team_metadata else None + alert_email_config: Final = _merge_budget_alert_email_configs( + global_cfg=None, per_key_cfg=_valid_alert_threshold_config(raw_config) + ) + if not alert_email_config or spend <= 0: + return + min_pct: Final = min(int(pct) for pct in alert_email_config) + if spend < max_budget * (min_pct / 100.0): + return + call_info: Final = CallInfo( + spend=spend, + max_budget=max_budget, + user_id=user_id, + team_id=team_id, + team_alias=team_alias, + organization_id=organization_id, + user_email=user_email, + event_group=Litellm_EntityType.TEAM_MEMBER, + max_budget_alert_emails=alert_email_config, + ) + asyncio.create_task(proxy_logging_obj.budget_alerts(type="max_budget_alert", user_info=call_info)) + + async def _check_team_member_budget( team_object: LiteLLM_TeamTable | None, user_object: LiteLLM_UserTable | None, @@ -5747,7 +5805,22 @@ async def _check_team_member_budget( max_budget=team_member_budget, ) - if math.isfinite(team_member_budget) and team_member_spend >= team_member_budget: + if not math.isfinite(team_member_budget): + return + + _team_member_max_budget_alert_check( + team_id=team_object.team_id, + team_alias=team_object.team_alias, + team_metadata=team_object.metadata, + organization_id=team_object.organization_id, + user_id=valid_token.user_id, + user_email=user_object.user_email if user_object is not None else None, + proxy_logging_obj=proxy_logging_obj, + spend=team_member_spend, + max_budget=team_member_budget, + ) + + if team_member_spend >= team_member_budget: raise litellm.BudgetExceededError( current_cost=team_member_spend, max_budget=team_member_budget, diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 22c3a248b9d..e3ce9bcd850 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -50,6 +50,7 @@ from litellm.proxy.auth.auth_checks import ( _get_user_role, _is_model_cost_zero, _is_user_proxy_admin, + _team_member_max_budget_alert_check, _virtual_key_max_budget_alert_check, _virtual_key_max_budget_check, _virtual_key_soft_budget_check, @@ -2287,6 +2288,19 @@ async def _user_api_key_auth_builder( max_budget=team_member_budget, ) if team_member_spend >= team_member_budget: + # common_checks sends this alert on requests that get past here, so only the + # request rejected here sends it from the builder. + _team_member_max_budget_alert_check( + team_id=_team_id, + team_alias=valid_token.team_alias, + team_metadata=valid_token.team_metadata, + organization_id=valid_token.org_id, + user_id=_user_id, + user_email=user_obj.user_email if user_obj is not None else None, + proxy_logging_obj=proxy_logging_obj, + spend=team_member_spend, + max_budget=team_member_budget, + ) _entity_id: Final = f"{valid_token.user_id}:{valid_token.team_id}" raise litellm.BudgetExceededError( current_cost=team_member_spend, diff --git a/tests/integration/_support/mail.py b/tests/integration/_support/mail.py new file mode 100644 index 00000000000..3894baeccc3 --- /dev/null +++ b/tests/integration/_support/mail.py @@ -0,0 +1,127 @@ +from __future__ import annotations + +import socketserver +import threading +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass +from email import message_from_bytes +from email.message import Message +from queue import SimpleQueue +from typing import Final + + +@dataclass(frozen=True, slots=True) +class Delivery: + sender: str + recipients: tuple[str, ...] + message: Message + + @property + def subject(self) -> str: + return str(self.message["Subject"]) + + @property + def html(self) -> str: + for part in self.message.walk(): + if part.get_content_type() == "text/html": + return part.get_payload(decode=True).decode() + return "" + + +class Mailbox: + def __init__(self, host: str, port: int) -> None: + self.host: Final = host + self.port: Final = port + self._lock: Final = threading.Lock() + self._deliveries: tuple[Delivery, ...] = () + + def record(self, delivery: Delivery) -> None: + with self._lock: + self._deliveries = (*self._deliveries, delivery) + + def deliveries(self) -> tuple[Delivery, ...]: + with self._lock: + return self._deliveries + + +def _address(argument: str) -> str: + return argument.split(":", 1)[1].strip().strip("<>") + + +@contextmanager +def smtp_sink() -> Generator[Mailbox, None, None]: + """Owned plaintext SMTP peer; deliveries traverse the proxy's real smtplib client.""" + errors: Final[SimpleQueue[Exception]] = SimpleQueue() + + class Handler(socketserver.StreamRequestHandler): + timeout = 5 + + def handle(self) -> None: + try: + self._session() + except Exception as error: + errors.put(error) + + def _reply(self, line: str) -> None: + self.wfile.write(f"{line}\r\n".encode()) + self.wfile.flush() + + def _session(self) -> None: + self._reply("220 integration-smtp ready") + # rebind-ok: the SMTP envelope is built across MAIL/RCPT lines and reset after DATA or RSET. + sender = "" + recipients: tuple[str, ...] = () + while True: + raw: Final = self.rfile.readline() + if not raw: + return + line: Final = raw.decode().rstrip("\r\n") + verb: Final = line.split(" ", 1)[0].upper() + if verb in {"EHLO", "HELO"}: + self._reply("250 integration-smtp") + elif verb == "MAIL": + sender = _address(line) + self._reply("250 OK") + elif verb == "RCPT": + recipients = (*recipients, _address(line)) + self._reply("250 OK") + elif verb == "DATA": + self._reply("354 End data with .") + body = bytearray() + while True: + chunk: Final = self.rfile.readline() + if not chunk or chunk == b".\r\n": + break + body.extend(chunk[1:] if chunk.startswith(b"..") else chunk) + mailbox.record(Delivery(sender, recipients, message_from_bytes(bytes(body)))) + sender, recipients = "", () + self._reply("250 OK queued") + elif verb == "RSET": + sender, recipients = "", () + self._reply("250 OK") + elif verb == "NOOP": + self._reply("250 OK") + elif verb == "QUIT": + self._reply("221 Bye") + return + else: + self._reply("502 Command not implemented") + + class OwnedServer(socketserver.ThreadingTCPServer): + allow_reuse_address = True + daemon_threads = False + + with OwnedServer(("127.0.0.1", 0), Handler) as server: + mailbox: Final = Mailbox("127.0.0.1", server.server_address[1]) + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield mailbox + finally: + server.shutdown() + thread.join(timeout=6) + assert not thread.is_alive(), "Owned SMTP server survived cleanup" + server.server_close() + failure: Final = None if errors.empty() else errors.get_nowait() + assert failure is None, f"Owned SMTP peer failed: {failure!r}" diff --git a/tests/integration/spend/test_team_member_budget_alerts.py b/tests/integration/spend/test_team_member_budget_alerts.py new file mode 100644 index 00000000000..f12bcb9748a --- /dev/null +++ b/tests/integration/spend/test_team_member_budget_alerts.py @@ -0,0 +1,93 @@ +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.mail import smtp_sink +from integration._support.process import owned_proxy + +MEMBER_BUDGET: Final = 0.10 +CALL_COST: Final = 20 * 0.001 + 20 * 0.002 + + +def _membership_spend(user_id: str, team_id: str) -> float: + rows: Final = read_rows( + 'SELECT spend FROM "LiteLLM_TeamMembership" WHERE user_id = %s AND team_id = %s', (user_id, team_id) + ) + return float(str(rows[0]["spend"])) if rows else 0.0 + + +def test_team_member_budget_thresholds_email_member_and_configured_recipients(gateway: Gateway, tmp_path: Path) -> None: + member_email: Final = f"member-{uuid.uuid4().hex}@integration.test" + finance_email: Final = f"finance-{uuid.uuid4().hex}@integration.test" + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["general_settings"]["alerting"] = ["email"] + path: Final = tmp_path / "email-alerting.yaml" + path.write_text(yaml.safe_dump(configuration)) + with smtp_sink() as mailbox: + overrides: Final = { + "SMTP_HOST": mailbox.host, + "SMTP_PORT": str(mailbox.port), + "SMTP_TLS": "False", + "SMTP_SENDER_EMAIL": "alerts@integration.test", + } + with owned_proxy(gateway, tmp_path, overrides, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + user_id: Final = scenario.user(user_email=member_email) + team_id: Final = scenario.team( + models=[model], + team_member_budget=MEMBER_BUDGET, + metadata={"team_member_max_budget_alert_emails": {"50": [], "100": [finance_email]}}, + ) + candidate.post("/team/member_add", {"team_id": team_id, "member": {"user_id": user_id, "role": "user"}}) + key: Final = scenario.key(team_id=team_id, user_id=user_id) + + first: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "first call"}]}, + key=key, + ) + assert first.status_code == 200, first.text + assert float(first.headers["x-litellm-response-cost"]) == pytest.approx(CALL_COST) + eventually( + lambda: _membership_spend(user_id, team_id), lambda spend: spend == pytest.approx(CALL_COST), seconds=70 + ) + assert mailbox.deliveries() == (), "no threshold is reached before the first call is recorded" + + second: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "second call"}]}, + key=key, + ) + assert second.status_code == 200, second.text + halfway: Final = eventually(mailbox.deliveries, lambda found: len(found) >= 1, seconds=30) + assert [delivery.recipients for delivery in halfway] == [(member_email,)], halfway + assert "50%" in halfway[0].subject, halfway[0].subject + assert f"${MEMBER_BUDGET}" in halfway[0].html, halfway[0].html + eventually( + lambda: _membership_spend(user_id, team_id), + lambda spend: spend == pytest.approx(2 * CALL_COST), + seconds=70, + ) + + third: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "third call"}]}, + key=key, + ) + assert third.status_code == 422 and third.json()["error"]["type"] == "budget_exceeded", third.text + capped: Final = eventually(mailbox.deliveries, lambda found: len(found) >= 3, seconds=30) + hundred: Final = capped[1:] + assert all("100%" in delivery.subject for delivery in hundred), capped + assert {recipient for delivery in hundred for recipient in delivery.recipients} == { + member_email, + finance_email, + }, capped + assert all(member_email in delivery.html and f"${MEMBER_BUDGET}" in delivery.html for delivery in hundred) + assert len(capped) == 3, capped diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index e42a47a1091..f014e9c26d1 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1,5 +1,6 @@ import asyncio import json +import sys import time from collections.abc import Iterator, Mapping from types import SimpleNamespace @@ -52,6 +53,7 @@ from litellm.proxy.auth.auth_checks import ( _log_budget_lookup_failure, _tag_max_budget_check, _team_max_budget_check, + _team_member_max_budget_alert_check, _virtual_key_max_budget_alert_check, _check_agent_caller_model_access, _virtual_key_max_budget_check, @@ -3774,6 +3776,141 @@ async def test_virtual_key_max_budget_alert_check_without_user_obj(): assert captured_call_info.user_email is None +@pytest.mark.parametrize( + "spend, team_metadata, expect_alert", + [ + (0.05, {"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, True), + (0.10, {"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, True), + (0.049, {"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, False), + (0.0, {"team_member_max_budget_alert_emails": {"50": []}}, False), + (0.10, {"team_member_max_budget_alert_emails": {"abc": []}}, False), + (0.05, {"team_member_max_budget_alert_emails": {"0": ["finance@co.com"], "100": []}}, False), + (0.10, {"team_member_max_budget_alert_emails": {"101": ["finance@co.com"]}}, False), + (0.10, {"team_member_max_budget_alert_emails": "50"}, False), + (0.10, {"soft_budget_alerting_emails": ["finance@co.com"]}, False), + (0.10, None, False), + ], +) +@pytest.mark.asyncio +async def test_team_member_max_budget_alert_check_dispatches_only_at_configured_thresholds( + spend, team_metadata, expect_alert +): + captured: list[tuple[str, CallInfo]] = [] + + class RecordingProxyLogging: + async def budget_alerts(self, type, user_info): + captured.append((type, user_info)) + + _team_member_max_budget_alert_check( + team_id="team-1", + team_alias="platform", + team_metadata=team_metadata, + organization_id="org-1", + user_id="user-1", + user_email="member@co.com", + proxy_logging_obj=RecordingProxyLogging(), + spend=spend, + max_budget=0.10, + ) + await asyncio.sleep(0) + + if not expect_alert: + assert captured == [], captured + return + assert [type for type, _ in captured] == ["max_budget_alert"], captured + call_info = captured[0][1] + assert call_info.event_group == Litellm_EntityType.TEAM_MEMBER + assert (call_info.spend, call_info.max_budget) == (spend, 0.10) + assert (call_info.user_id, call_info.user_email) == ("user-1", "member@co.com") + assert (call_info.team_id, call_info.team_alias, call_info.organization_id) == ("team-1", "platform", "org-1") + assert call_info.max_budget_alert_emails == {"50": [], "100": ["finance@co.com"]} + assert call_info.token is None + + +@pytest.mark.asyncio +async def test_team_member_max_budget_alert_check_drops_thresholds_outside_1_to_100(): + captured: list[CallInfo] = [] + + class RecordingProxyLogging: + async def budget_alerts(self, type, user_info): + captured.append(user_info) + + _team_member_max_budget_alert_check( + team_id="team-1", + team_alias="platform", + team_metadata={ + "team_member_max_budget_alert_emails": { + "0": ["a@co.com"], + "50": [], + "150": ["b@co.com"], + "1" * (sys.int_info.default_max_str_digits + 1): ["c@co.com"], + } + }, + organization_id="org-1", + user_id="user-1", + user_email="member@co.com", + proxy_logging_obj=RecordingProxyLogging(), + spend=0.05, + max_budget=0.10, + ) + await asyncio.sleep(0) + + assert [call_info.max_budget_alert_emails for call_info in captured] == [{"50": []}], captured + + +@pytest.mark.asyncio +async def test_check_team_member_budget_dispatches_the_configured_alert_before_the_hard_cap(): + from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership + + captured: list[tuple[str, CallInfo]] = [] + + class RecordingProxyLogging: + async def budget_alerts(self, type, user_info): + captured.append((type, user_info)) + + team_object = LiteLLM_TeamTable( + team_id="team-1", + team_alias="platform", + metadata={"team_member_max_budget_alert_emails": {"50": [], "100": ["finance@co.com"]}}, + ) + user_object = LiteLLM_UserTable(user_id="user-1", user_email="member@co.com") + valid_token = UserAPIKeyAuth(token="tok-1", user_id="user-1", team_id="team-1") + team_membership = LiteLLM_TeamMembership( + user_id="user-1", + team_id="team-1", + spend=0.10, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10), + ) + + async def spend_from_fallback(counter_key, fallback_spend, max_budget=None, **kwargs): + return fallback_spend + + with ( + patch("litellm.proxy.proxy_server.get_current_spend", spend_from_fallback), + patch( + "litellm.proxy.auth.auth_checks.get_team_membership", new_callable=AsyncMock, return_value=team_membership + ), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_team_member_budget( + team_object=team_object, + user_object=user_object, + valid_token=valid_token, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=RecordingProxyLogging(), + ) + await asyncio.sleep(0) + + assert (exc_info.value.entity_type, exc_info.value.entity_id) == ("team_member", "user-1:team-1") + assert [type for type, _ in captured] == ["max_budget_alert"], captured + call_info = captured[0][1] + assert call_info.event_group == Litellm_EntityType.TEAM_MEMBER + assert (call_info.spend, call_info.max_budget) == (0.10, 0.10) + assert (call_info.user_id, call_info.user_email, call_info.team_id) == ("user-1", "member@co.com", "team-1") + assert call_info.max_budget_alert_emails == {"50": [], "100": ["finance@co.com"]} + + @pytest.mark.parametrize( "spend, max_budget, expect_alert", [ diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index de669449f85..470db99108a 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -29,6 +29,7 @@ from litellm.proxy._types import ( LiteLLM_OrganizationTable, LiteLLM_TeamTableCachedObj, LiteLLM_UserTable, + Litellm_EntityType, LitellmUserRoles, ProxyErrorTypes, ProxyException, @@ -8055,6 +8056,152 @@ async def test_cached_key_team_member_budget_honours_temp_increase(expiry_offset assert "Max budget: 2.0" in exc_info.value.message +async def _authenticate_and_authorize(mock_request, api_key): + """Builder then the single common_checks gate, the same sequence user_api_key_auth runs.""" + from litellm.proxy.auth.user_api_key_auth import _authorize_authenticated_request + + request_data = {"model": "claude-sonnet-5", "messages": [{"role": "user", "content": "hi"}]} + auth_obj = await _user_api_key_auth_builder( + request=mock_request, + api_key=f"Bearer {api_key}", + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data=request_data, + ) + recovered = await _authorize_authenticated_request( + user_api_key_auth_obj=auth_obj, + request=mock_request, + request_data=request_data, + route="/v1/messages", + api_key=f"Bearer {api_key}", + ) + return recovered or auth_obj + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "team_member_spend, expect_blocked, expected_alerts", + [ + (1.1, False, 0), + (1.2, False, 1), + (2.4, True, 1), + ], +) +async def test_cached_key_team_member_budget_emails_configured_thresholds( + team_member_spend, expect_blocked, expected_alerts +): + """The team's team_member_max_budget_alert_emails thresholds fire from the cached-key auth path, + including on the request that trips the hard cap, and stay silent below the lowest threshold.""" + from litellm.proxy._types import LiteLLM_TeamMembership, LiteLLM_TeamTableCachedObj + from litellm.proxy.common_utils.user_api_key_cache import ( + team_membership_auth_cache_key, + team_membership_reservation_cache_key, + ) + from litellm.proxy.utils import hash_token + + api_key = "sk-team-member-alert-thresholds" + hashed_token = hash_token(api_key) + team_id = "team-alert-thresholds" + user_id = "user-alert-thresholds" + alert_emails = {"50": [], "100": ["finance@example.com"]} + + user_api_key_cache = DualCache() + await _cache_key_object( + hashed_token=hashed_token, + user_api_key_obj=UserAPIKeyAuth( + token=hashed_token, + team_id=team_id, + team_alias="platform", + team_metadata={"team_member_max_budget_alert_emails": alert_emails}, + user_id=user_id, + team_member_spend=team_member_spend, + ), + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=None, + ) + await user_api_key_cache.async_set_cache( + key=f"team_id:{team_id}", + value=LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="platform", + metadata={"team_member_max_budget_alert_emails": alert_emails}, + ), + ) + await user_api_key_cache.async_set_cache( + key=user_id, + value=LiteLLM_UserTable( + user_id=user_id, user_email="member@example.com", user_role=LitellmUserRoles.INTERNAL_USER + ), + ) + membership = LiteLLM_TeamMembership( + user_id=user_id, + team_id=team_id, + spend=team_member_spend, + budget_id="budget-alert-thresholds", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=2.4), + ) + # A live proxy holds the row under both keys, so any second team-member check in the + # auth flow would find it too and send a duplicate alert. + for membership_cache_key in ( + team_membership_reservation_cache_key(team_id=team_id, user_id=user_id), + team_membership_auth_cache_key(team_id=team_id, user_id=user_id), + ): + await user_api_key_cache.async_set_cache(key=membership_cache_key, value=membership) + + mock_request = MagicMock() + mock_request.url.path = "/v1/messages" + mock_request.method = "POST" + mock_request.headers = {"authorization": f"Bearer {api_key}"} + mock_request.query_params = {} + mock_request.state = SimpleNamespace() + + proxy_logging_obj = MagicMock() + proxy_logging_obj.budget_alerts = AsyncMock() + proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None) + proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock(return_value=None) + + async def _auth(): + return await _authenticate_and_authorize(mock_request, api_key) + + with ( + patch( # test-quality-ok: the builder reads proxy settings from module globals, no injection seam + "litellm.proxy.proxy_server.general_settings", {"disable_budget_reservation": True} + ), + patch("litellm.proxy.proxy_server.master_key", "sk-master"), # test-quality-ok: module-global proxy state + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), # test-quality-ok: module-global proxy state + patch( # test-quality-ok: seed the cached key, team and membership without a DB + "litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache + ), + patch( # test-quality-ok: module-global proxy state + "litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj + ), + patch( # test-quality-ok: the live counter needs Redis or a DB; pin the spend the check compares + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=team_member_spend), + ), + ): + if expect_blocked: + with pytest.raises(ProxyException) as exc_info: + await _auth() + assert exc_info.value.type == ProxyErrorTypes.budget_exceeded + else: + await _auth() + await asyncio.sleep(0) + + assert proxy_logging_obj.budget_alerts.await_count == expected_alerts + if expected_alerts == 0: + return + call_info = proxy_logging_obj.budget_alerts.await_args.kwargs["user_info"] + assert proxy_logging_obj.budget_alerts.await_args.kwargs["type"] == "max_budget_alert" + assert call_info.event_group == Litellm_EntityType.TEAM_MEMBER + assert (call_info.spend, call_info.max_budget) == (team_member_spend, 2.4) + assert (call_info.user_id, call_info.user_email) == (user_id, "member@example.com") + assert (call_info.team_id, call_info.team_alias) == (team_id, "platform") + assert call_info.max_budget_alert_emails == alert_emails + + async def _proxy_exception_for_key( api_key: str, general_settings: dict[str, bool], diff --git a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py index 8b89c592f02..52e44ca5448 100644 --- a/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py +++ b/tests/unit/enterprise/enterprise_callbacks/send_emails/test_base_email.py @@ -1090,6 +1090,47 @@ async def test_multi_threshold_empty_emails_only_owner( assert to_emails == ["owner@co.com"] +@pytest.mark.asyncio +async def test_multi_threshold_team_member_alert_renders_member_template_per_team( + base_email_logger, mock_send_email +): + """A team member budget alert is keyed per member and team, names the member and team, + and goes to the member plus the threshold's configured recipients""" + user_info = CallInfo( + user_id="member_1", + user_email="member@co.com", + team_id="team_a", + team_alias="Platform", + spend=0.10, + max_budget=0.10, + event_group=Litellm_EntityType.TEAM_MEMBER, + max_budget_alert_emails={"50": [], "100": ["finance@co.com"]}, + ) + + mock_cache = mock.AsyncMock() + mock_cache.async_increment_cache = mock.AsyncMock(return_value=1) + base_email_logger.internal_usage_cache = mock_cache + + with mock.patch.dict(os.environ, {"PROXY_BASE_URL": "http://test.com"}): + await base_email_logger.budget_alerts(type="max_budget_alert", user_info=user_info) + + cache_keys = sorted(c[1]["key"] for c in mock_cache.async_increment_cache.call_args_list) + assert cache_keys == [ + "email_budget_alerts:max_budget_alert:100:team_member:member_1:team_a", + "email_budget_alerts:max_budget_alert:50:team_member:member_1:team_a", + ] + assert mock_send_email.call_count == 2 + hundred = next( + c.kwargs for c in mock_send_email.call_args_list if "100%" in c.kwargs["subject"] + ) + assert hundred["subject"] == "LiteLLM: Team Member Budget Alert - 100% of Team Member Budget Reached" + assert sorted(hundred["to_email"]) == ["finance@co.com", "member@co.com"] + assert "member@co.com" in hundred["html_body"] and "Platform" in hundred["html_body"] + assert "team member budget" in hundred["html_body"] and "$0.1" in hundred["html_body"] + fifty = next(c.kwargs for c in mock_send_email.call_args_list if "50%" in c.kwargs["subject"]) + assert fifty["to_email"] == ["member@co.com"] + + @pytest.mark.asyncio async def test_no_map_preserves_old_single_threshold( base_email_logger, mock_send_email diff --git a/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py b/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py index 52b7cc983a7..f3199d9ebf9 100644 --- a/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py +++ b/tests/unit/integrations/SlackAlerting/test_budget_alert_types.py @@ -1,4 +1,7 @@ -from litellm.integrations.SlackAlerting.budget_alert_types import SoftBudgetAlert +from litellm.integrations.SlackAlerting.budget_alert_types import ( + SoftBudgetAlert, + TokenBudgetAlert, +) from litellm.proxy._types import CallInfo, Litellm_EntityType @@ -64,3 +67,31 @@ class TestSoftBudgetAlert: result = alert.get_id(user_info) assert result == "default_id" + + +class TestTokenBudgetAlert: + def test_get_id_dedupes_team_member_alerts_per_member_and_team(self): + alert = TokenBudgetAlert() + team_a = CallInfo( + spend=8.0, max_budget=10.0, user_id="member_1", team_id="team_a", event_group=Litellm_EntityType.TEAM_MEMBER + ) + team_b = CallInfo( + spend=8.0, max_budget=10.0, user_id="member_1", team_id="team_b", event_group=Litellm_EntityType.TEAM_MEMBER + ) + + assert alert.get_id(team_a) == "team_member:member_1:team_a" + assert alert.get_id(team_b) == "team_member:member_1:team_b" + + def test_get_id_uses_token_for_key_alerts(self): + alert = TokenBudgetAlert() + user_info = CallInfo( + spend=8.0, + max_budget=10.0, + token="hashed_key", + user_id="member_1", + team_id="team_a", + event_group=Litellm_EntityType.KEY, + ) + + assert alert.get_id(user_info) == "hashed_key" + assert alert.get_event_message() == "Key Budget: " diff --git a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py index b9e5ff2eeb7..0c2b95fd448 100644 --- a/tests/unit/integrations/SlackAlerting/test_slack_alerting.py +++ b/tests/unit/integrations/SlackAlerting/test_slack_alerting.py @@ -393,6 +393,33 @@ def _slack_alerting_with_env_resolution() -> SlackAlerting: return slack_alerting +@pytest.mark.asyncio +@pytest.mark.parametrize( + "event_group, expected_prefix", + [ + (Litellm_EntityType.TEAM_MEMBER, "Team Member Budget: Budget Crossed"), + (Litellm_EntityType.KEY, "Key Budget: Budget Crossed"), + ], +) +async def test_max_budget_alert_labels_team_member_budget(event_group, expected_prefix): + slack_alerting: Final = _slack_alerting_with_env_resolution() + slack_alerting.send_alert = AsyncMock() + + await slack_alerting.budget_alerts( + type="max_budget_alert", + user_info=CallInfo( + spend=10.5, + max_budget=10.0, + token="hashed_key", + user_id="member_1", + team_id="team_a", + event_group=event_group, + ), + ) + + assert slack_alerting.send_alert.await_args.kwargs["message"].startswith(expected_prefix) + + @pytest.mark.asyncio async def test_send_alert_falls_back_to_alerting_webhook_url_env(monkeypatch): monkeypatch.delenv("SLACK_WEBHOOK_URL", raising=False) diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index 37f63433d2a..a693ee971d4 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -2338,6 +2338,103 @@ describe("TeamInfoView - the exact bytes the update call sends", () => { expect(wireBody(payload)).toStrictEqual(expected); }); + const openEditorWithMemberBudgetAlerts = async (user: ReturnType) => { + vi.mocked(networking.teamInfoCall).mockResolvedValue( + createMockTeamData({ + models: ["gpt-4"], + team_member_budget_table: { max_budget: 42 }, + metadata: { team_member_max_budget_alert_emails: { "50": [], "100": ["finance@test.com"] } }, + }), + ); + vi.mocked(networking.teamUpdateCall).mockResolvedValue({ data: {}, team_id: "123" } as any); + + renderWithProviders(); + await waitFor(() => expect(screen.queryAllByText("Test Team").length).toBeGreaterThan(0)); + await user.click(screen.getByRole("tab", { name: "Settings" })); + await user.click(await screen.findByRole("button", { name: /edit settings/i })); + await screen.findByLabelText("Team Name"); + }; + + const memberBudgetAlertEmails = (payload: Record) => + (wireBody(payload).metadata as Record).team_member_max_budget_alert_emails; + + it("resends the stored team member budget alert thresholds when the section stays closed", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + const payload = await save(user); + + expect(memberBudgetAlertEmails(payload)).toStrictEqual({ "50": [], "100": ["finance@test.com"] }); + }); + + it("sends the edited team member budget alert thresholds as a percent to recipients map", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + const thresholds = screen.getAllByPlaceholderText("% of budget"); + const recipients = screen.getAllByPlaceholderText(/Additional recipients/); + expect(thresholds.map((input) => (input as HTMLInputElement).value)).toStrictEqual(["50", "100"]); + expect(recipients.map((input) => (input as HTMLInputElement).value)).toStrictEqual(["", "finance@test.com"]); + + fireEvent.change(thresholds[0], { target: { value: "75" } }); + fireEvent.change(recipients[0], { target: { value: " lead@test.com, finance@test.com " } }); + await user.click(screen.getByRole("button", { name: "Add Budget Alert Threshold" })); + fireEvent.change(screen.getAllByPlaceholderText("% of budget")[2], { target: { value: "90" } }); + + const payload = await save(user); + + expect(memberBudgetAlertEmails(payload)).toStrictEqual({ + "75": ["lead@test.com", "finance@test.com"], + "100": ["finance@test.com"], + "90": [], + }); + }); + + it("drops the team member budget alert thresholds key once every row is removed", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + const removeButtons = screen.getAllByRole("button", { name: "Remove budget alert threshold" }); + await user.click(removeButtons[1]); + await user.click(removeButtons[0]); + + const payload = await save(user); + + expect(memberBudgetAlertEmails(payload)).toBeUndefined(); + }); + + it("blocks the save when a team member budget alert threshold is above 100", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + const threshold = screen.getAllByPlaceholderText("% of budget")[0] as HTMLInputElement; + fireEvent.change(threshold, { target: { value: "150" } }); + expect(threshold.validity.rangeOverflow).toBe(true); + + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await waitFor(() => expect(networking.teamUpdateCall).not.toHaveBeenCalled()); + }); + + it("refuses to save a team member budget alert row with no threshold", async () => { + const user = userEvent.setup({ delay: null }); + await openEditorWithMemberBudgetAlerts(user); + + await user.click(screen.getByText("Team Member Settings")); + await screen.findByLabelText("Default Budget (USD)"); + await user.click(screen.getByRole("button", { name: "Add Budget Alert Threshold" })); + await user.click(screen.getByRole("button", { name: /save changes/i })); + + await screen.findByText("Enter a whole number from 1 to 100"); + expect(networking.teamUpdateCall).not.toHaveBeenCalled(); + }); + it("carries every typed value to the update payload at the type and shape antd sends today", async () => { const user = userEvent.setup({ delay: null }); await openEditor(user); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 78ed507216a..3845f94593d 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -118,6 +118,13 @@ import { TEAM_INFO_TAB_LABELS, } from "./tabVisibilityUtils"; import TeamMembersComponent from "./TeamMemberTab"; +import { + isValidThreshold, + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY, + teamMemberBudgetAlertEmailsFromRows, + teamMemberBudgetAlertRowsFromMetadata, + teamMemberBudgetAlertSummary, +} from "./teamMemberBudgetAlertEmails"; import { TeamVirtualKeysTable } from "./TeamVirtualKeysTable"; import ResetMemberBudgetsDialog from "./ResetMemberBudgetsDialog"; import { customBudgetMemberUserIds, shouldPromptMemberBudgetReset } from "./memberBudgetReset"; @@ -128,6 +135,7 @@ const UI_MANAGED_METADATA_KEYS: ReadonlySet = new Set([ "logging", "secret_manager_settings", "soft_budget_alerting_emails", + TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY, "model_tpm_limit", "model_rpm_limit", "default_estimated_output_tokens", @@ -355,6 +363,18 @@ const teamUpdateFieldsSchema = z.object({ team_member_key_duration: z.string().optional(), team_member_tpm_limit: numericInputSchema, team_member_rpm_limit: numericInputSchema, + team_member_max_budget_alert_emails: z + .array(z.object({ threshold: z.number().nullable(), emails: z.string() })) + .superRefine((rows, ctx) => { + rows.forEach((row, index) => { + if (!isValidThreshold(row.threshold)) { + ctx.addIssue({ code: "custom", message: "Enter a whole number from 1 to 100", path: [index, "threshold"] }); + } else if (rows.filter((other) => other.threshold === row.threshold).length > 1) { + ctx.addIssue({ code: "custom", message: "Duplicate threshold", path: [index, "threshold"] }); + } + }); + }) + .optional(), budget_duration: z.string().nullish(), tpm_limit: numericInputSchema, rpm_limit: numericInputSchema, @@ -422,6 +442,7 @@ const TEAM_MEMBER_SETTINGS_FIELDS = [ "team_member_key_duration", "team_member_tpm_limit", "team_member_rpm_limit", + "team_member_max_budget_alert_emails", ] as const; const SEARCH_TOOL_SETTINGS_FIELDS = ["object_permission_search_tools"] as const; @@ -437,6 +458,7 @@ const EMPTY_TEAM_UPDATE_VALUES: TeamUpdateFormValues = { team_member_key_duration: undefined, team_member_tpm_limit: undefined, team_member_rpm_limit: undefined, + team_member_max_budget_alert_emails: [], budget_duration: undefined, tpm_limit: undefined, rpm_limit: undefined, @@ -487,6 +509,7 @@ const toTeamFormValues = (info: TeamInfoRecord, effectiveGuardrails: string[]): team_member_key_duration: info.metadata?.team_member_key_duration, team_member_tpm_limit: info.team_member_budget_table?.tpm_limit, team_member_rpm_limit: info.team_member_budget_table?.rpm_limit, + team_member_max_budget_alert_emails: [...teamMemberBudgetAlertRowsFromMetadata(info.metadata)], budget_duration: info.budget_duration, tpm_limit: info.tpm_limit, rpm_limit: info.rpm_limit, @@ -572,6 +595,11 @@ const TeamInfoView: React.FC = ({ append: appendModelLimit, remove: removeModelLimit, } = useFieldArray({ control: form.control, name: "modelLimits" }); + const { + fields: memberBudgetAlertRows, + append: appendMemberBudgetAlertRow, + remove: removeMemberBudgetAlertRow, + } = useFieldArray({ control: form.control, name: "team_member_max_budget_alert_emails" }); const [teamMemberSettingsOpen, setTeamMemberSettingsOpen] = useState(false); const [searchToolSettingsOpen, setSearchToolSettingsOpen] = useState(false); const [isEditMemberModalVisible, setIsEditMemberModalVisible] = useState(false); @@ -994,6 +1022,15 @@ const TeamInfoView: React.FC = ({ ? { allowed_passthrough_routes: info.metadata.allowed_passthrough_routes } : {}; + const memberBudgetAlertEmails = + values.team_member_max_budget_alert_emails !== undefined + ? teamMemberBudgetAlertEmailsFromRows(values.team_member_max_budget_alert_emails) + : info.metadata?.[TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY]; + const memberBudgetAlertEmailsMetadata = + memberBudgetAlertEmails !== undefined && Object.keys(memberBudgetAlertEmails).length > 0 + ? { [TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY]: memberBudgetAlertEmails } + : {}; + const updateData: any = { team_id: teamId, team_alias: values.team_alias, @@ -1025,6 +1062,7 @@ const TeamInfoView: React.FC = ({ .filter((email: string) => email.length > 0) : values.soft_budget_alerting_emails || [], ...(secretManagerSettings !== undefined ? { secret_manager_settings: secretManagerSettings } : {}), + ...memberBudgetAlertEmailsMetadata, }, ...(values.policies?.length > 0 ? { policies: values.policies } : {}), ...(values.organization_id !== info.organization_id ? { organization_id: values.organization_id ?? null } : {}), @@ -1632,6 +1670,71 @@ const TeamInfoView: React.FC = ({ )} + + + {labelWithHint( + "Budget Alert Thresholds", + "Email each member when their spend reaches a percentage of their team member budget. The member is always notified; add comma-separated addresses to notify others as well. Requires email alerting to be configured on the proxy.", + )} + + {memberBudgetAlertRows.map((row, index) => ( +
+ + {({ ref, value, onChange, ...field }) => ( + ) => + onChange(event.target.value === "" ? null : Number(event.target.value)) + } + placeholder="% of budget" + min={1} + max={100} + step={1} + /> + )} + + + {({ ref, value, ...field }) => ( + + )} + + +
+ ))} + +
@@ -2202,6 +2305,7 @@ const TeamInfoView: React.FC = ({
Key Duration: {info.metadata?.team_member_key_duration || "No Limit"}
TPM Limit: {info.team_member_budget_table?.tpm_limit ?? "No Limit"}
RPM Limit: {info.team_member_budget_table?.rpm_limit ?? "No Limit"}
+
Budget Alert Thresholds: {teamMemberBudgetAlertSummary(info.metadata).join("; ") || "None"}

Router Settings

diff --git a/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts new file mode 100644 index 00000000000..3faa051047b --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.test.ts @@ -0,0 +1,94 @@ +import { describe, expect, it } from "vitest"; +import { + isValidThreshold, + teamMemberBudgetAlertEmailsFromRows, + teamMemberBudgetAlertRowsFromMetadata, + teamMemberBudgetAlertSummary, +} from "./teamMemberBudgetAlertEmails"; + +describe("teamMemberBudgetAlertRowsFromMetadata", () => { + it("turns the stored threshold map into rows sorted by threshold", () => { + const metadata = { + team_member_max_budget_alert_emails: { "100": ["finance@example.com", "cto@example.com"], "50": [] }, + }; + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([ + { threshold: 50, emails: "" }, + { threshold: 100, emails: "finance@example.com, cto@example.com" }, + ]); + }); + + it("drops non-numeric thresholds and non-list recipients instead of crashing", () => { + const metadata = { + team_member_max_budget_alert_emails: { fifty: [], "75": "finance@example.com", "90": [1], "100": ["a@b.c"] }, + }; + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([{ threshold: 100, emails: "a@b.c" }]); + }); + + it("drops API-stored thresholds outside 1 to 100 so they never block the form", () => { + const metadata = { + team_member_max_budget_alert_emails: { "0": ["a@b.c"], "50": [], "101": ["a@b.c"] }, + }; + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([{ threshold: 50, emails: "" }]); + }); + + it.each([undefined, null, "50", { team_member_max_budget_alert_emails: "50" }, { soft_budget_alerting_emails: [] }])( + "returns no rows for unrelated or malformed metadata %j", + (metadata) => { + expect(teamMemberBudgetAlertRowsFromMetadata(metadata)).toEqual([]); + }, + ); +}); + +describe("teamMemberBudgetAlertEmailsFromRows", () => { + it("builds the threshold map, splitting, trimming and deduplicating recipients", () => { + expect( + teamMemberBudgetAlertEmailsFromRows([ + { threshold: 50, emails: "" }, + { threshold: 100, emails: " finance@example.com,cto@example.com , finance@example.com, " }, + ]), + ).toEqual({ "50": [], "100": ["finance@example.com", "cto@example.com"] }); + }); + + it("skips rows without a valid threshold", () => { + expect( + teamMemberBudgetAlertEmailsFromRows([ + { threshold: null, emails: "finance@example.com" }, + { threshold: 0, emails: "" }, + { threshold: 101, emails: "" }, + { threshold: 12.5, emails: "" }, + { threshold: 80, emails: "" }, + ]), + ).toEqual({ "80": [] }); + }); + + it("round-trips the stored config", () => { + const stored = { team_member_max_budget_alert_emails: { "50": [], "100": ["finance@example.com"] } }; + expect(teamMemberBudgetAlertEmailsFromRows(teamMemberBudgetAlertRowsFromMetadata(stored))).toEqual( + stored.team_member_max_budget_alert_emails, + ); + }); +}); + +describe("isValidThreshold", () => { + it.each([ + [1, true], + [50, true], + [100, true], + [0, false], + [101, false], + [33.3, false], + [null, false], + ])("treats %s as valid=%s", (threshold, valid) => { + expect(isValidThreshold(threshold)).toBe(valid); + }); +}); + +describe("teamMemberBudgetAlertSummary", () => { + it("states that the member is always notified and lists extra recipients", () => { + expect( + teamMemberBudgetAlertSummary({ + team_member_max_budget_alert_emails: { "100": ["finance@example.com"], "50": [] }, + }), + ).toEqual(["50%: member", "100%: member, finance@example.com"]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts new file mode 100644 index 00000000000..36d5ddbac02 --- /dev/null +++ b/ui/litellm-dashboard/src/components/team/teamMemberBudgetAlertEmails.ts @@ -0,0 +1,57 @@ +export const TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY = "team_member_max_budget_alert_emails" as const; + +export interface TeamMemberBudgetAlertRow { + readonly threshold: number | null; + readonly emails: string; +} + +export type TeamMemberBudgetAlertEmails = Readonly>; + +const isEmailList = (value: unknown): value is readonly string[] => + Array.isArray(value) && value.every((email) => typeof email === "string"); + +const splitEmails = (emails: string): readonly string[] => + Array.from( + new Set( + emails + .split(",") + .map((email) => email.trim()) + .filter((email) => email.length > 0), + ), + ); + +const THRESHOLD_MIN = 1; +const THRESHOLD_MAX = 100; + +export const isValidThreshold = (threshold: number | null): threshold is number => { + const isWholeNumber = threshold !== null && Number.isInteger(threshold); + return isWholeNumber && threshold >= THRESHOLD_MIN && threshold <= THRESHOLD_MAX; +}; + +export const teamMemberBudgetAlertRowsFromMetadata = (metadata: unknown): readonly TeamMemberBudgetAlertRow[] => { + if (typeof metadata !== "object" || metadata === null) return []; + const config: unknown = (metadata as Record)[TEAM_MEMBER_MAX_BUDGET_ALERT_EMAILS_KEY]; + if (typeof config !== "object" || config === null || Array.isArray(config)) return []; + return Object.entries(config as Record) + .flatMap(([key, emails]) => { + const threshold = Number(key); + return /^\d+$/.test(key) && isValidThreshold(threshold) && isEmailList(emails) + ? [{ threshold, emails: emails.join(", ") }] + : []; + }) + .sort((a, b) => (a.threshold ?? 0) - (b.threshold ?? 0)); +}; + +export const teamMemberBudgetAlertEmailsFromRows = ( + rows: readonly TeamMemberBudgetAlertRow[], +): TeamMemberBudgetAlertEmails => + Object.fromEntries( + rows + .filter((row) => isValidThreshold(row.threshold)) + .map((row) => [String(row.threshold), splitEmails(row.emails)]), + ); + +export const teamMemberBudgetAlertSummary = (metadata: unknown): readonly string[] => + teamMemberBudgetAlertRowsFromMetadata(metadata).map((row) => + row.emails.length > 0 ? `${row.threshold}%: member, ${row.emails}` : `${row.threshold}%: member`, + );