mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
546402c98c
commit
6830a5c338
13 changed files with 242 additions and 45 deletions
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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, _>(
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue