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>
This commit is contained in:
yujonglee 2026-10-09 19:36:49 -07:00 • committed by GitHub
parent 546402c98c
commit 6830a5c338
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 242 additions and 45 deletions

View file

@ -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<ChatCompletionsCall> {
@ -38,7 +38,7 @@ impl PythonBinding for ChatCompletionsPythonHost {
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<ChatCompletionsCall, InvokeError<Error>> {
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<ChatCompletions> 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)
}
}

View file

@ -22,7 +22,7 @@ fn run_chat_completions(
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let host = InferenceHost::new(
call.resolved()?.unbind(),
call.view()?,
"litellm.rust_bridge.chat_completions.route_host",
);
run_inference::<ChatCompletionsRoute, _>(

View file

@ -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<PyDict>,
view: RequestView,
module: &'static str,
}
@ -34,16 +39,17 @@ pub(super) struct ProjectedCall {
}
impl InferenceHost {
pub fn new(request: Py<PyDict>, 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<ProjectedCall> {
self.view.resolve(py, arguments)?;
let argument = |name: &str| self.argument(py, arguments, name);
let string = |name: &str| -> PyResult<Option<String>> {
argument(name)?.map(|value| value.extract()).transpose()
@ -93,7 +99,15 @@ impl InferenceHost {
arguments: &Bound<'py, PyDict>,
name: &str,
) -> PyResult<Option<Bound<'py, PyAny>>> {
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))
}
}

View file

@ -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<PyErr> {
/// 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<PyDict>,
view: RequestView,
cache: PythonCache,
}
impl MessagesPythonHost {
pub(super) fn new(request: Py<PyDict>, 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<Result<MessagesCall, Error>> {
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<Option<String>> {
argument(name)?.map(|value| value.extract()).transpose()
};
@ -182,9 +182,9 @@ impl MessagesPythonHost {
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> PyResult<Option<Map<String, Value>>> {
let request = self.request.bind(py);
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
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<Option<ProviderSpecificHeaders>> {
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::<Option<String>>().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::<Py<pyo3::exceptions::PyBaseException>>()
@ -264,6 +265,9 @@ impl PythonBinding for MessagesPythonHost {
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<(MessagesCall, Selection), InvokeError<Error>> {
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<Cached<Messages>> 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;

View file

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

View file

@ -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<Bound<'py, PyDict>> {
effective_py_args(&self.base, &self.kwargs)
}
/// The view a route host reads its arguments from.
fn view(&self) -> PyResult<RequestView> {
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<RequestDicts>);
struct RequestDicts {
base: Py<PyDict>,
request: Py<PyDict>,
}
impl RequestView {
pub(crate) fn new(base: &Bound<'_, PyDict>, kwargs: &Bound<'_, PyDict>) -> PyResult<Self> {
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<Option<Bound<'py, PyAny>>> {
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<Bound<'py, PyDict>> {
@ -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::<PyDict>()
.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::<String>().unwrap());
assert_eq!(value.as_deref(), expected);
});
}
#[rstest]
#[case::sync("transcription")]
#[case::asynchronous("atranscription")]

View file

@ -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<PyDict>,
view: RequestView,
data: OcrHostData,
}
impl OcrPythonHost {
pub(super) fn new(request: Py<PyDict>) -> 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::<Py<PyBaseException>>().map_err(PyErr::from));
match mapped {
Ok(mapped) => PyErr::from_value(mapped.into_bound(py).into_any()),
@ -160,10 +161,11 @@ impl PythonHostCalls<Ocr> 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::<PyDict>()
.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);

View file

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

View file

@ -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<ResponsesCall> {
@ -59,7 +59,7 @@ impl PythonBinding for ResponsesPythonHost {
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<ResponsesCall, InvokeError<Error>> {
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<Responses> 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)
}
}

View file

@ -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::<ResponsesRoute, _>(py, call, asynchronous, ResponsesPythonHost(host))
}

View file

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

View file

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

View file

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