diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 12bc57a8931..f47ee5d743d 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -7,6 +7,7 @@ mod execution; mod function_trace; mod lifecycle; mod marshal; +mod params; mod routes; mod token_counter; diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index fdfecdafe61..ed276437e38 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -90,14 +90,14 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py) -> PyRe pub(crate) fn project_optional_fields( kwargs: &Bound<'_, PyDict>, - names: &[&str], + route: &crate::params::RouteParamSpec<'_>, + policy: &crate::params::RequestParamPolicy, ) -> PyResult> { - let selected: Vec = kwargs - .py() - .import("litellm.rust_bridge.params")? - .getattr("provider_param_names")? - .call1((kwargs, names.to_vec()))? - .extract()?; + let keys: Vec = kwargs.keys().extract()?; + let selected: Vec = keys + .into_iter() + .filter(|name| policy.includes(name, route)) + .collect(); selected .iter() .filter_map(|name| match kwargs.get_item(name) { diff --git a/litellm-rust/crates/python-bridge/src/params.rs b/litellm-rust/crates/python-bridge/src/params.rs new file mode 100644 index 00000000000..485cd48cea0 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/params.rs @@ -0,0 +1,160 @@ +use std::collections::HashSet; + +use pyo3::prelude::*; +use pyo3::types::PyList; + +pub(crate) struct RequestParamPolicy { + sdk_reserved_param_names: HashSet, +} + +pub(crate) struct RouteParamSpec<'a> { + pub bound: &'a [&'a str], + pub consumed: &'a [&'a str], +} + +impl RequestParamPolicy { + pub(crate) fn extract(names: &Bound<'_, PyList>) -> PyResult { + let names: Vec = names.extract()?; + Ok(Self { + sdk_reserved_param_names: names.into_iter().collect(), + }) + } + + pub(crate) fn includes(&self, name: &str, route: &RouteParamSpec<'_>) -> bool { + route.consumed.contains(&name) + || (!self.sdk_reserved_param_names.contains(name) && !route.bound.contains(&name)) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::marshal::project_optional_fields; + use pyo3::exceptions::PyTypeError; + use pyo3::types::PyDict; + use rstest::rstest; + use serde_json::json; + + #[rstest] + #[case::unknown("future", &[], true)] + #[case::sdk_control("callbacks", &[], false)] + #[case::bound_argument("document", &[], false)] + #[case::consumed_control("callbacks", &["callbacks"], true)] + #[case::consumed_bound_argument("document", &["document"], true)] + #[case::overrides("extra_body", &[], true)] + #[case::new_sdk_control("new_control", &[], false)] + fn selection(#[case] name: &str, #[case] consumed: &[&str], #[case] expected: bool) { + let policy = RequestParamPolicy { + sdk_reserved_param_names: ["callbacks", "new_control", "callbacks"] + .into_iter() + .map(str::to_owned) + .collect(), + }; + assert_eq!( + policy.includes( + name, + &RouteParamSpec { + bound: &["document"], + consumed + } + ), + expected + ); + } + + #[test] + fn registry_extraction_rejects_non_strings() { + Python::initialize(); + Python::attach(|py| { + let names = PyList::new(py, [42]).unwrap(); + let error = RequestParamPolicy::extract(&names).err().unwrap(); + assert!(error.is_instance_of::(py)); + }); + } + + #[rstest] + #[case::ocr("document", &["document"], false)] + #[case::chat("messages", &["messages"], false)] + #[case::transcription("audio", &["audio"], false)] + #[case::no_implicit_ocr_rule("document", &["messages"], true)] + #[case::no_implicit_chat_rule("messages", &["audio"], true)] + fn route_owns_bound_argument_names( + #[case] name: &str, + #[case] bound: &[&str], + #[case] expected: bool, + ) { + let policy = RequestParamPolicy { + sdk_reserved_param_names: HashSet::new(), + }; + assert_eq!( + policy.includes( + name, + &RouteParamSpec { + bound, + consumed: &[] + } + ), + expected + ); + } + + #[test] + fn registry_is_read_at_projection_not_capture() { + Python::initialize(); + Python::attach(|py| { + let names = PyList::empty(py); + let retained = names.clone().unbind(); + names.append("future").unwrap(); + let policy = RequestParamPolicy::extract(retained.bind(py)).unwrap(); + assert!(!policy.includes( + "future", + &RouteParamSpec { + bound: &[], + consumed: &[] + } + )); + }); + } + + #[test] + fn projection_skips_host_objects_and_preserves_unknown_values() { + Python::initialize(); + Python::attach(|py| { + let kwargs = PyDict::new(py); + let host = py.eval(c"object()", None, None).unwrap(); + kwargs.set_item("callbacks", &host).unwrap(); + kwargs.set_item("future", py.None()).unwrap(); + kwargs.set_item("enabled", false).unwrap(); + let policy = + RequestParamPolicy::extract(&PyList::new(py, ["callbacks"]).unwrap()).unwrap(); + assert_eq!( + project_optional_fields( + &kwargs, + &RouteParamSpec { + bound: &[], + consumed: &[] + }, + &policy + ) + .unwrap(), + json!({"future":null,"enabled":false}) + .as_object() + .unwrap() + .clone() + ); + assert!(kwargs.get_item("callbacks").unwrap().unwrap().is(&host)); + assert_eq!(kwargs.len(), 3); + kwargs.set_item("unknown", &host).unwrap(); + let error = project_optional_fields( + &kwargs, + &RouteParamSpec { + bound: &[], + consumed: &[], + }, + &policy, + ) + .unwrap_err(); + assert!(error.is_instance_of::(py)); + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index 12d902a3544..a5250486857 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -1,5 +1,5 @@ use pyo3::prelude::*; -use pyo3::types::{PyDict, PyTuple}; +use pyo3::types::{PyDict, PyList, PyTuple}; use litellm_core::auth::ResolvedCredential; use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest}; @@ -21,7 +21,10 @@ struct PythonOcrHost { } enum OcrHostData { - Unprojected { request: Py }, + Unprojected { + request: Py, + sdk_reserved_param_names: Py, + }, Projected(Box), Released, } @@ -186,10 +189,17 @@ impl PythonRoute for PythonOcrHost { fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult { Ok(match operation { OcrHostOperation::ProjectRequest => { - let OcrHostData::Unprojected { request } = &self.data else { + let OcrHostData::Unprojected { + request, + sdk_reserved_param_names, + } = &self.data + else { return Err(missing_state()); }; - let projected = project_request(py, request.bind(py), self.state.kwargs.bind(py))?; + let policy = + crate::params::RequestParamPolicy::extract(sdk_reserved_param_names.bind(py))?; + let projected = + project_request(py, request.bind(py), self.state.kwargs.bind(py), &policy)?; let has_token_provider = projected.fields.azure_ad_token_provider.is_some(); let request = projected.request; self.data = OcrHostData::Projected(Box::new(ProjectedOcrHost { @@ -227,7 +237,7 @@ impl PythonRoute for PythonOcrHost { } let error = self.state.error.as_ref().ok_or_else(missing_state)?; let (request, provider) = match &self.data { - OcrHostData::Unprojected { request } => (request.bind(py), ""), + OcrHostData::Unprojected { request, .. } => (request.bind(py), ""), OcrHostData::Projected(projected) => ( projected.fields.boundary_request.bind(py), projected.fields.provider, @@ -250,7 +260,13 @@ impl PythonRoute for PythonOcrHost { } fn traverse(&self, visit: &pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { match &self.data { - OcrHostData::Unprojected { request } => visit.call(request), + OcrHostData::Unprojected { + request, + sdk_reserved_param_names, + } => { + visit.call(request)?; + visit.call(sdk_reserved_param_names) + } OcrHostData::Projected(projected) => { visit.call(&projected.fields.boundary_request)?; visit.call(&projected.fields.document)?; @@ -282,6 +298,7 @@ fn _ocr_lifecycle( args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, asynchronous: bool, + sdk_reserved_param_names: Bound<'_, PyList>, ) -> PyResult> { let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?; let call = admitted_call(OcrCall::admit( @@ -301,6 +318,7 @@ fn _ocr_lifecycle( )?, data: OcrHostData::Unprojected { request: request.unbind(), + sdk_reserved_param_names: sdk_reserved_param_names.unbind(), }, }; run_call(py, call, host) diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 8b6a1b02e19..a5a77a916b4 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -15,6 +15,17 @@ use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; use crate::errors::RustBridgeDeclined; use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}; +const BOUND_PARAM_NAMES: &[&str] = &[ + "model", + "document", + "api_key", + "api_base", + "timeout", + "extra_headers", + "custom_llm_provider", + "input_sources", +]; + pub(super) struct ProjectedOcrFields { pub boundary_request: Py, pub document: Py, @@ -114,6 +125,7 @@ pub(super) fn project_request( py: Python<'_>, request: &Bound<'_, PyAny>, kwargs: &Bound<'_, PyDict>, + policy: &crate::params::RequestParamPolicy, ) -> PyResult { let boundary_request = request.clone().unbind(); let arguments = OcrArguments { request, kwargs }; @@ -125,7 +137,11 @@ pub(super) fn project_request( let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) .map_err(ocr_error_to_pyerr)?; let names = specs.iter().map(|spec| spec.name).collect::>(); - let optional_params = project_optional_fields(kwargs, &names)?; + let route = crate::params::RouteParamSpec { + bound: BOUND_PARAM_NAMES, + consumed: &names, + }; + let optional_params = project_optional_fields(kwargs, &route, policy)?; let input_sources = request_input_sources( kwargs, names diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 382c5d6aae4..21132956de4 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -9,7 +9,7 @@ from litellm.ocr.input import convert_file_document_to_url_document, get_mime_ty from litellm.rust_bridge.bindings import native_exception_types from litellm.rust_bridge.configuration import rust_ocr_enabled from litellm.rust_bridge.ocr import LiteLLMOcrRequest -from litellm.rust_bridge.ocr_lifecycle import select +from litellm.rust_bridge.ocr_lifecycle import invoke, select __all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr") @@ -52,7 +52,7 @@ def ocr( if native is not None: try: return cast( # cast-ok: False selects the synchronous result - OCRResponse, native(request, args, kwargs, False) + OCRResponse, invoke(native, request, args, kwargs, False) ) except _decline_types(): pass @@ -68,7 +68,7 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr if native is not None: try: return await cast( # cast-ok: True selects the asynchronous result - Awaitable[OCRResponse], native(request, args, kwargs, True) + Awaitable[OCRResponse], invoke(native, request, args, kwargs, True) ) except _decline_types(): pass diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index e62c85f4599..8004ced8367 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -47,6 +47,7 @@ def _ocr_lifecycle( args: tuple[object, ...], kwargs: dict[str, object], asynchronous: bool, + sdk_reserved_param_names: list[str], ) -> OCRResponse | Coroutine[object, object, OCRResponse]: ... def transcription( model: str, diff --git a/litellm/rust_bridge/ocr_lifecycle.py b/litellm/rust_bridge/ocr_lifecycle.py index 5ca584e1c11..ddf88e5708d 100644 --- a/litellm/rust_bridge/ocr_lifecycle.py +++ b/litellm/rust_bridge/ocr_lifecycle.py @@ -1,21 +1,23 @@ from __future__ import annotations -from collections.abc import Awaitable, Mapping, Sequence +from collections.abc import Awaitable, Mapping from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.ocr import LiteLLMOcrRequest +from litellm.types.utils import all_litellm_params class NativeOcrLifecycle(Protocol): def __call__( self, request: LiteLLMOcrRequest, - args: Sequence[object], - kwargs: Mapping[str, object], + args: tuple[object, ...], + kwargs: dict[str, object], asynchronous: bool, + sdk_reserved_param_names: list[str], ) -> OCRResponse | Awaitable[OCRResponse]: ... @@ -40,6 +42,16 @@ def _binding(value: object) -> NativeOcrLifecycle | None: NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding) +def invoke( + native: NativeOcrLifecycle, + request: LiteLLMOcrRequest, + args: tuple[object, ...], + kwargs: dict[str, object], + asynchronous: bool, +) -> OCRResponse | Awaitable[OCRResponse]: + return native(request, args, kwargs, asynchronous, all_litellm_params) + + def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None: if request.kwargs.get("aocr"): return None diff --git a/litellm/rust_bridge/params.py b/litellm/rust_bridge/params.py deleted file mode 100644 index 83626d01b3a..00000000000 --- a/litellm/rust_bridge/params.py +++ /dev/null @@ -1,19 +0,0 @@ -"""Select SDK extension fields without copying or converting their values.""" - -from collections.abc import Mapping, Sequence -from typing import Final - -from litellm.types.utils import all_litellm_params - - -def provider_param_names(kwargs: Mapping[str, object], consumed: Sequence[str]) -> tuple[str, ...]: - sdk_fields: Final = frozenset(all_litellm_params) | { - "model", - "document", - "timeout", - "extra_headers", - "custom_llm_provider", - "input_sources", - } - consumed_fields: Final = frozenset(consumed) - return tuple(name for name in kwargs if name in consumed_fields or name not in sdk_fields) diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 501a4e986c0..f9230c47930 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -64,7 +64,11 @@ def test_public_binding_keeps_positional_fields_and_defaults_out_of_native_hook_ args: tuple[object, ...], kwargs: Mapping[str, object], asynchronous: bool, + sdk_reserved_param_names: list[str], ) -> OCRResponse: + from litellm.types.utils import all_litellm_params + + assert sdk_reserved_param_names is all_litellm_params captured.append((request, args, kwargs, asynchronous)) return OCRResponse(pages=[], model=request.model) @@ -94,6 +98,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() args: tuple[object, ...], kwargs: Mapping[str, object], asynchronous: bool, + sdk_reserved_param_names: list[str], ) -> OCRResponse: assert args == () captured.append(kwargs) diff --git a/tests/test_litellm/rust_bridge/test_params.py b/tests/test_litellm/rust_bridge/test_params.py deleted file mode 100644 index 6cb89841a9d..00000000000 --- a/tests/test_litellm/rust_bridge/test_params.py +++ /dev/null @@ -1,24 +0,0 @@ -from typing import Final - -from litellm.rust_bridge.params import provider_param_names - - -def test_selects_unknown_fields_and_consumed_controls_without_reading_values() -> None: - callback: Final = object() - future: Final = {"nested": [None, False, 0]} - kwargs: Final = { - "model": "mistral/model", - "litellm_logging_obj": callback, - "metadata": callback, - "vertex_credentials": "credentials", - "future": future, - "extra_body": {"future": None}, - } - - assert provider_param_names(kwargs, ("vertex_credentials",)) == ("vertex_credentials", "future", "extra_body") - assert kwargs["future"] is future - assert kwargs["litellm_logging_obj"] is callback - - -def test_route_can_explicitly_consume_a_name_also_used_by_sdk() -> None: - assert provider_param_names({"metadata": {"provider": True}}, ("metadata",)) == ("metadata",) diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 9173394b697..d97599e13ed 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -23,6 +23,73 @@ from tests.test_litellm_rust.support.requests import OCR_RESPONSE, call_aocr, ca pytestmark = pytest.mark.requires_rust_extension +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"]) +async def test_provider_selection_preserves_route_hook_timing_and_reads_sdk_registry( + ocr_server: RecordingServer, asynchronous: bool +) -> None: + from litellm.types.utils import all_litellm_params + + control_name: Final = "test_projection_host_control" + host_object: Final = object() + + class ChangeInputs(CustomLogger): + async def async_pre_call_deployment_hook(self, kwargs, call_type): + all_litellm_params.append(control_name) + kwargs[control_name] = host_object + kwargs["future_added"] = False + kwargs["future_replaced"] = {"after": 0} + kwargs.pop("future_removed") + + litellm.callbacks.append(ChangeInputs()) + try: + if asynchronous: + await call_aocr(ocr_server, future_replaced="before", future_removed=True) + else: + call_ocr(ocr_server, future_replaced="before", future_removed=True) + finally: + if control_name in all_litellm_params: + all_litellm_params.remove(control_name) + + assert len(ocr_server.requests) == 1 + request: Final = ocr_server.requests[0] + assert not request.headers.get("user-agent", "").startswith("python-httpx") + if asynchronous: + assert request.body["future_added"] is False + assert request.body["future_replaced"] == {"after": 0} + assert "future_removed" not in request.body + else: + assert "future_added" not in request.body + assert request.body["future_replaced"] == "before" + assert request.body["future_removed"] is True + assert control_name not in request.body + + +def test_unstarted_native_call_releases_registry_cycle( + ocr_server: RecordingServer, +) -> None: + from litellm.ocr.main import _public_request + from litellm.rust_bridge import _native + + ocr_server.expected_requests = 0 + + class Registry(list[str]): + owner: object + + def create() -> weakref.ReferenceType[Registry]: + registry: Final = Registry(["callbacks"]) + kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {}} + coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True, registry) + registry.owner = coroutine + return weakref.ref(registry) + + reference: Final = create() + with pytest.warns(RuntimeWarning, match="coroutine 'drive' was never awaited"): + gc.collect() + assert reference() is None + assert ocr_server.requests == [] + + @pytest.mark.asyncio @pytest.mark.parametrize("phase", ["deployment", "failure"]) async def test_cancellation_during_failure_obeys_phase_policy(ocr_server: RecordingServer, phase: str) -> None: @@ -594,7 +661,7 @@ def test_unstarted_native_coroutine_releases_input_without_reading_file(ocr_serv def create(): file: Final = File() kwargs: Final = {"model": "mistral/mistral-ocr-latest", "document": {"type": "file", "file": file}} - coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True) + coroutine: Final = _native._ocr_lifecycle(_public_request("aocr", (), kwargs), (), kwargs, True, []) file.owner = coroutine coroutine.close() return weakref.ref(file)