From 6830a5c33844afe78497b9f20dd53d5b9140b947 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Fri, 9 Oct 2026 19:36:49 -0700 Subject: [PATCH] fix(python-bridge): preserve callback keyword deletions in native requests (#45692) * fix(messages): preserve callback keyword deletions in native requests * fix(python-bridge): share hook-aware request view across native route hosts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../src/routes/chat_completions/host.rs | 10 +- .../src/routes/chat_completions/mod.rs | 2 +- .../python-bridge/src/routes/inference.rs | 32 ++++-- .../python-bridge/src/routes/messages/host.rs | 34 +++--- .../python-bridge/src/routes/messages/mod.rs | 2 +- .../crates/python-bridge/src/routes/mod.rs | 101 ++++++++++++++++++ .../python-bridge/src/routes/ocr/host.rs | 19 ++-- .../python-bridge/src/routes/ocr/mod.rs | 2 +- .../src/routes/responses/host.rs | 10 +- .../python-bridge/src/routes/responses/mod.rs | 2 +- .../messages/test_callbacks.py | 32 ++++++ tests/test_litellm_rust/ocr/test_callbacks.py | 19 +++- tests/test_litellm_rust/test_inference.py | 22 +++- 13 files changed, 242 insertions(+), 45 deletions(-) diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs index 1547a4daa93..5ad9f8df8b8 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs @@ -12,7 +12,7 @@ use pyo3::{ pub(super) struct ChatCompletionsPythonHost(pub InferenceHost); pub(super) fn project( - host: &InferenceHost, + host: &mut InferenceHost, py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult { @@ -38,7 +38,7 @@ impl PythonBinding for ChatCompletionsPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> Result> { - let call = project(&self.0, py, arguments).map_err(InvokeError::Python)?; + let call = project(&mut self.0, py, arguments).map_err(InvokeError::Python)?; if call .optional_params .get("stream") @@ -94,8 +94,10 @@ impl PythonHostCalls for ChatCompletionsPythonHost { } impl PythonOwned for ChatCompletionsPythonHost { - fn close(&mut self, _: Python<'_>) {} + fn close(&mut self, _: Python<'_>) { + self.0.close(); + } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.0.request) + self.0.traverse(visit) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs index 2b69d0e6321..00ca5e092a8 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/mod.rs @@ -22,7 +22,7 @@ fn run_chat_completions( asynchronous: bool, ) -> PyResult> { let host = InferenceHost::new( - call.resolved()?.unbind(), + call.view()?, "litellm.rust_bridge.chat_completions.route_host", ); run_inference::( diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs index 0333aaeb04a..bc442d1cf36 100644 --- a/litellm-rust/crates/python-bridge/src/routes/inference.rs +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -5,15 +5,20 @@ use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache}; use litellm_callbacks_legacy_python::LoggingOperation; use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; use litellm_host::{call::HostedMachine, protocol::Protocol}; -use litellm_host_python::{PythonBinding, PythonHostCalls, from_py, present}; +use litellm_host_python::{PythonBinding, PythonHostCalls, from_py}; use litellm_http::transport::Error as TransportError; use litellm_inference::RouteError; use litellm_secrets::source::SecretSource; -use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; +use pyo3::{ + exceptions::PyValueError, + gc::{PyTraverseError, PyVisit}, + prelude::*, + types::PyDict, +}; use serde::Serialize; use serde_json::{Map, Value}; -use super::NativeCall; +use super::{NativeCall, RequestView}; use crate::{ errors::{RustUpstreamError, route_error_to_pyerr}, marshal::{ @@ -23,7 +28,7 @@ use crate::{ }; pub(super) struct InferenceHost { - pub request: Py, + view: RequestView, module: &'static str, } @@ -34,16 +39,17 @@ pub(super) struct ProjectedCall { } impl InferenceHost { - pub fn new(request: Py, module: &'static str) -> Self { - Self { request, module } + pub fn new(view: RequestView, module: &'static str) -> Self { + Self { view, module } } pub fn project( - &self, + &mut self, py: Python<'_>, arguments: &Bound<'_, PyDict>, input: &str, ) -> PyResult { + self.view.resolve(py, arguments)?; let argument = |name: &str| self.argument(py, arguments, name); let string = |name: &str| -> PyResult> { argument(name)?.map(|value| value.extract()).transpose() @@ -93,7 +99,15 @@ impl InferenceHost { arguments: &Bound<'py, PyDict>, name: &str, ) -> PyResult>> { - present(arguments, self.request.bind(py), name) + self.view.argument(py, arguments, name) + } + + pub fn close(&mut self) { + self.view.close(); + } + + pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + self.view.traverse(visit) } pub fn parameters( @@ -130,7 +144,7 @@ impl InferenceHost { let mapped = py .import(self.module)? .getattr("map_failure")? - .call1((native.value(py), self.request.bind(py)))?; + .call1((native.value(py), self.view.request(py)?))?; Ok(PyErr::from_value(mapped)) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index d34047c735d..adeba397b26 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -2,7 +2,7 @@ use crate::cache::{CacheCall, Cached, PythonCache, Selection}; use litellm_host_python::{PythonHostCalls, PythonOwned}; use bytes::Bytes; -use litellm_host_python::{InvokeError, PythonBinding, from_py, present, to_py}; +use litellm_host_python::{InvokeError, PythonBinding, from_py, to_py}; use litellm_http::transport::Error as TransportError; use litellm_inference_messages::{ Error, MessagesCall, MessagesSettings, MessagesShaping, litellm_params, messages_body, @@ -23,6 +23,7 @@ use crate::{ errors::{RustUpstreamError, route_error_to_pyerr}, marshal::{optional_timeout, project_optional_fields, public_response, python_timeout_seconds}, python_settings::missing_module, + routes::RequestView, }; const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host"; @@ -118,14 +119,14 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult { /// The Python side of the Messages route: projects the prepared arguments and builds the /// public response, chunks and exceptions. pub(super) struct MessagesPythonHost { - request: Py, + view: RequestView, cache: PythonCache, } impl MessagesPythonHost { - pub(super) fn new(request: Py, asynchronous: bool) -> Self { + pub(super) fn new(view: RequestView, asynchronous: bool) -> Self { Self { - request, + view, cache: PythonCache::new(asynchronous), } } @@ -135,8 +136,7 @@ impl MessagesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult> { - let request = self.request.bind(py); - let argument = |name: &str| present(arguments, request, name); + let argument = |name: &str| self.view.argument(py, arguments, name); let string = |name: &str| -> PyResult> { argument(name)?.map(|value| value.extract()).transpose() }; @@ -182,9 +182,9 @@ impl MessagesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult>> { - let request = self.request.bind(py); let mapping = |name: &str| -> PyResult>> { - present(arguments, request, name)? + self.view + .argument(py, arguments, name)? .map(|value| from_py(&value)) .transpose() }; @@ -199,7 +199,8 @@ impl MessagesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult> { - present(arguments, self.request.bind(py), "provider_specific_header")? + self.view + .argument(py, arguments, "provider_specific_header")? .map(|value| from_py(&value)) .transpose() } @@ -226,9 +227,9 @@ impl MessagesPythonHost { } fn provider(&self, py: Python<'_>) -> String { - self.request - .bind(py) - .get_item("custom_llm_provider") + self.view + .request(py) + .and_then(|request| request.get_item("custom_llm_provider")) .ok() .flatten() .and_then(|value| value.extract::>().ok().flatten()) @@ -242,7 +243,7 @@ impl MessagesPythonHost { let mapped = py .import(ROUTE_HOST_MODULE) .and_then(|module| module.getattr("map_failure")) - .and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py)))) + .and_then(|map| map.call1((error.value(py), self.view.request(py)?, self.provider(py)))) .and_then(|mapped| { mapped .extract::>() @@ -264,6 +265,9 @@ impl PythonBinding for MessagesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> Result<(MessagesCall, Selection), InvokeError> { + self.view + .resolve(py, arguments) + .map_err(InvokeError::Python)?; let selection = crate::cache::configure(&mut self.cache, py, arguments, "anthropic_messages") .map_err(InvokeError::Python)?; @@ -341,16 +345,18 @@ impl PythonHostCalls> for MessagesPythonHost { impl PythonOwned for MessagesPythonHost { fn close(&mut self, _: Python<'_>) { + self.view.close(); self.cache.close(); } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.request)?; + self.view.traverse(visit)?; self.cache.traverse(visit) } } #[cfg(test)] mod tests { + use litellm_host_python::present; use rstest::rstest; use serde_json::json; 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 5f69cf078b7..c788760b268 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -44,7 +44,7 @@ fn run_messages(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyR }, )) }, - MessagesPythonHost::new(call.resolved()?.unbind(), asynchronous), + MessagesPythonHost::new(call.view()?, asynchronous), hooks, asynchronous, ) diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 7427ec8da7d..416b94db55b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -14,7 +14,9 @@ use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol} use litellm_host_python::{ HookChain, PythonBinding, PythonCallHooks, PythonHostCalls, effective_py_args, }; +use litellm_host_python::{missing_state, present}; use pyo3::{ + gc::{PyTraverseError, PyVisit}, prelude::*, types::{PyDict, PyMapping, PyTuple}, }; @@ -36,6 +38,69 @@ impl<'py> NativeCall<'py> { fn resolved(&self) -> PyResult> { effective_py_args(&self.base, &self.kwargs) } + + /// The view a route host reads its arguments from. + fn view(&self) -> PyResult { + RequestView::new(&self.base, &self.kwargs) + } +} + +/// The keyword view a route host reads: the signature base under the caller's keywords, +/// laid again once the hooks have rewritten them so a keyword a hook deleted stays deleted +/// instead of falling back to the caller's value. +pub(crate) struct RequestView(Option); + +struct RequestDicts { + base: Py, + request: Py, +} + +impl RequestView { + pub(crate) fn new(base: &Bound<'_, PyDict>, kwargs: &Bound<'_, PyDict>) -> PyResult { + Ok(Self(Some(RequestDicts { + base: base.clone().unbind(), + request: effective_py_args(base, kwargs)?.unbind(), + }))) + } + + fn dicts(&self) -> PyResult<&RequestDicts> { + self.0.as_ref().ok_or_else(missing_state) + } + + pub(crate) fn resolve( + &mut self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> PyResult<()> { + let dicts = self.0.as_mut().ok_or_else(missing_state)?; + dicts.request = effective_py_args(dicts.base.bind(py), arguments)?.unbind(); + Ok(()) + } + + pub(crate) fn request<'py>(&self, py: Python<'py>) -> PyResult<&Bound<'py, PyDict>> { + Ok(self.dicts()?.request.bind(py)) + } + + pub(crate) fn argument<'py>( + &self, + py: Python<'py>, + arguments: &Bound<'py, PyDict>, + name: &str, + ) -> PyResult>> { + present(arguments, self.request(py)?, name) + } + + pub(crate) fn close(&mut self) { + self.0 = None; + } + + pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + let Some(dicts) = &self.0 else { + return Ok(()); + }; + visit.call(&dicts.base)?; + visit.call(&dicts.request) + } } fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult> { @@ -130,6 +195,42 @@ mod tests { .unwrap() } + fn dict<'py>(py: Python<'py>, source: &str) -> Bound<'py, PyDict> { + py.eval(&std::ffi::CString::new(source).unwrap(), None, None) + .unwrap() + .cast_into::() + .unwrap() + } + + #[rstest] + #[case::deleted_keyword_stays_deleted("{'system': None}", "{'system': 'caller'}", "{}", None)] + #[case::deleted_keyword_without_default("{}", "{'system': 'caller'}", "{}", None)] + #[case::rewritten_keyword_wins( + "{'system': None}", + "{'system': 'caller'}", + "{'system': 'hook'}", + Some("hook") + )] + #[case::signature_default_survives("{'system': 'default'}", "{}", "{}", Some("default"))] + fn request_view_reads_the_hook_rewritten_keywords_over_the_base( + #[case] base: &str, + #[case] kwargs: &str, + #[case] prepared: &str, + #[case] expected: Option<&str>, + ) { + Python::initialize(); + Python::attach(|py| { + let mut view = super::RequestView::new(&dict(py, base), &dict(py, kwargs)).unwrap(); + let prepared = dict(py, prepared); + view.resolve(py, &prepared).unwrap(); + let value = view + .argument(py, &prepared, "system") + .unwrap() + .map(|value| value.extract::().unwrap()); + assert_eq!(value.as_deref(), expected); + }); + } + #[rstest] #[case::sync("transcription")] #[case::asynchronous("atranscription")] diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 5ae72acd104..dd3c648cc3e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -15,7 +15,7 @@ use super::{ errors::to_pyerr as ocr_error_to_pyerr, project::{OcrHostHandles, SETTINGS_ERROR_MARKER, project_request}, }; -use crate::marshal::public_response; +use crate::{marshal::public_response, routes::RequestView}; enum OcrHostData { Unprojected, @@ -27,14 +27,14 @@ enum OcrHostData { /// document as it goes), acquires Azure AD tokens, and builds the public response and /// exception. pub(super) struct OcrPythonHost { - request: Py, + view: RequestView, data: OcrHostData, } impl OcrPythonHost { - pub(super) fn new(request: Py) -> Self { + pub(super) fn new(view: RequestView) -> Self { Self { - request, + view, data: OcrHostData::Unprojected, } } @@ -58,7 +58,8 @@ impl OcrPythonHost { let OcrHostData::Unprojected = self.data else { return Err(missing_state()); }; - let (request, handles) = project_request(self.request.bind(py), arguments)?; + self.view.resolve(py, arguments)?; + let (request, handles) = project_request(self.view.request(py)?, arguments)?; let caller_token = handles.azure_ad_token_provider.is_some(); self.data = OcrHostData::Projected(Box::new(handles)); Ok(OcrCall { @@ -78,7 +79,7 @@ impl OcrPythonHost { let mapped = py .import("litellm.rust_bridge.ocr.route_host") .and_then(|module| module.getattr("map_failure")) - .and_then(|map| map.call1((error.value(py), self.request.bind(py), provider))) + .and_then(|map| map.call1((error.value(py), self.view.request(py)?, provider))) .and_then(|mapped| mapped.extract::>().map_err(PyErr::from)); match mapped { Ok(mapped) => PyErr::from_value(mapped.into_bound(py).into_any()), @@ -160,10 +161,11 @@ impl PythonHostCalls for OcrPythonHost { impl PythonOwned for OcrPythonHost { fn close(&mut self, _: Python<'_>) { + self.view.close(); self.data = OcrHostData::Released; } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.request)?; + self.view.traverse(visit)?; if let OcrHostData::Projected(handles) = &self.data && let Some(provider) = &handles.azure_ad_token_provider { @@ -218,7 +220,8 @@ del provider .unwrap() .cast_into::() .unwrap(); - let mut host = OcrPythonHost::new(PyDict::new(py).unbind()); + let empty = PyDict::new(py); + let mut host = OcrPythonHost::new(RequestView::new(&empty, &empty).unwrap()); assert!(host.decode_request(py, &kwargs).unwrap().caller_token); locals.del_item("kwargs").unwrap(); drop(kwargs); 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 fab9ecc13c1..ec7f2ce79d9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -44,7 +44,7 @@ fn run_ocr(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult let route = litellm_inference_ocr::OcrRoute::new(client); Ok(route.machine(request, None)) }, - OcrPythonHost::new(call.resolved()?.unbind()), + OcrPythonHost::new(call.view()?), hooks, asynchronous, ) diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs index b22bbc6192b..8261e66406b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs @@ -12,7 +12,7 @@ use pyo3::{ pub(super) struct ResponsesPythonHost(pub InferenceHost); pub(super) fn project( - host: &InferenceHost, + host: &mut InferenceHost, py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> PyResult { @@ -59,7 +59,7 @@ impl PythonBinding for ResponsesPythonHost { py: Python<'_>, arguments: &Bound<'_, PyDict>, ) -> Result> { - let call = project(&self.0, py, arguments).map_err(InvokeError::Python)?; + let call = project(&mut self.0, py, arguments).map_err(InvokeError::Python)?; if call .optional_params .get("stream") @@ -121,8 +121,10 @@ impl PythonHostCalls for ResponsesPythonHost { } impl PythonOwned for ResponsesPythonHost { - fn close(&mut self, _: Python<'_>) {} + fn close(&mut self, _: Python<'_>) { + self.0.close(); + } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.0.request) + self.0.traverse(visit) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs index 50e5790a532..1bf24813103 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs @@ -61,7 +61,7 @@ fn run_responses(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> Py "native Python responses streaming", )); } - let host = InferenceHost::new(resolved.unbind(), ROUTE_HOST_MODULE); + let host = InferenceHost::new(call.view()?, ROUTE_HOST_MODULE); run_inference::(py, call, asynchronous, ResponsesPythonHost(host)) } diff --git a/tests/test_litellm_rust/messages/test_callbacks.py b/tests/test_litellm_rust/messages/test_callbacks.py index dc66852d214..b3cd7bd01f9 100644 --- a/tests/test_litellm_rust/messages/test_callbacks.py +++ b/tests/test_litellm_rust/messages/test_callbacks.py @@ -239,3 +239,35 @@ async def test_native_messages_stream_success_log_carries_usage_rebuilt_from_the assert usage.completion_tokens == MESSAGES_EVENTS[4][1]["usage"]["output_tokens"] assert usage.prompt_tokens == MESSAGES_RESPONSE["usage"]["input_tokens"] assert success[0].response.choices[0].message.content == "Hello from native Messages" + + +@pytest.mark.asyncio +async def test_native_messages_does_not_restore_keywords_deleted_by_a_pre_call_hook( + messages_server: RecordingServer, +) -> None: + class DropKeywords(CustomLogger): + async def async_pre_call_deployment_hook( + self, kwargs: dict[str, object], call_type: object + ) -> dict[str, object]: + return {name: value for name, value in kwargs.items() if name not in ("system", "extra_headers")} + + call_arguments: Final = arguments( + messages_server, + system="remove this instruction", + extra_headers={"x-deleted": "remove this header"}, + headers={"x-kept": "keep this header"}, + stream=False, + ) + with rebound(litellm, "callbacks", [DropKeywords()]): + await litellm.anthropic.messages.acreate(**call_arguments) + + assert_served_natively(messages_server) + sent: Final = messages_server.requests[0] + assert sent.body == { + "model": MESSAGES_MODEL.removeprefix("anthropic/"), + "messages": call_arguments["messages"], + "max_tokens": call_arguments["max_tokens"], + "stream": call_arguments["stream"], + } + assert sent.headers.get("x-deleted") is None + assert sent.headers.get("x-kept") == "keep this header" diff --git a/tests/test_litellm_rust/ocr/test_callbacks.py b/tests/test_litellm_rust/ocr/test_callbacks.py index d3b04c8bb6f..9e3945f11b1 100644 --- a/tests/test_litellm_rust/ocr/test_callbacks.py +++ b/tests/test_litellm_rust/ocr/test_callbacks.py @@ -16,7 +16,7 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.integrations.custom_logger import CustomLogger from litellm.llms.base_llm.ocr.transformation import OCRResponse from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging -from tests.test_litellm_rust.support.isolation import isolated_callback_registries +from tests.test_litellm_rust.support.isolation import isolated_callback_registries, rebound from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec from tests.test_litellm_rust.support.requests import ( OCR_DOCUMENT, @@ -65,6 +65,23 @@ def test_native_ocr_pre_call_callback_receives_transformed_provider_request(ocr_ } +@pytest.mark.asyncio +async def test_native_ocr_does_not_restore_keywords_deleted_by_a_pre_call_hook(ocr_server: RecordingServer) -> None: + class DropKeywords(CustomLogger): + async def async_pre_call_deployment_hook( + self, kwargs: dict[str, object], call_type: object + ) -> dict[str, object]: + return {name: value for name, value in kwargs.items() if name not in ("pages", "extra_headers")} + + with rebound(litellm, "callbacks", [DropKeywords()]): + await call_native_aocr(ocr_server, pages=[0], extra_headers={"x-deleted": "remove this header"}) + + assert len(ocr_server.requests) == 1 + sent: Final = ocr_server.requests[0] + assert sent.body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT} + assert sent.headers.get("x-deleted") is None + + @pytest.mark.parametrize("raise_after_edit", [False, True], ids=["callback-returns", "callback-raises"]) def test_native_ocr_pre_call_body_edit_reaches_next_callback_and_provider( ocr_server: RecordingServer, raise_after_edit: bool diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py index 8211b50cb8d..83a01e8ad4d 100644 --- a/tests/test_litellm_rust/test_inference.py +++ b/tests/test_litellm_rust/test_inference.py @@ -87,7 +87,6 @@ def native_call( "api_base": server.base_url, "custom_llm_provider": "openai", "extra_headers": None, - **response_kwargs, }, ) return (_native.aresponses if asynchronous else _native.responses)(response_request) @@ -147,6 +146,27 @@ async def test_native_inference_pre_call_edits_reach_the_provider( assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.75 +@pytest.mark.asyncio +async def test_native_inference_does_not_restore_keywords_deleted_by_a_pre_call_hook( + route: Route, recording_server: RecordingServer +) -> None: + class DropKeywords(CustomLogger): + async def async_pre_call_deployment_hook( + self, kwargs: dict[str, object], call_type: CallTypes | None + ) -> dict[str, object]: + return {name: value for name, value in kwargs.items() if name not in ("temperature", "extra_headers")} + + litellm.callbacks.append(DropKeywords()) + await execute( + route, True, recording_server, {"temperature": 0.25, "extra_headers": {"x-deleted": "remove this header"}} + ) + + assert len(recording_server.requests) == 1 + sent: Final = recording_server.requests[0] + assert "temperature" not in _OBJECT.validate_python(sent.body) + assert sent.headers.get("x-deleted") is None + + @pytest.mark.asyncio @pytest.mark.parametrize("from_credentials", (False, True)) async def test_native_resource_setup_uses_deployment_hook_arguments(