refactor(rust): type request parameter projection at the Python boundary

This commit is contained in:
Yujong Lee 2026-09-15 13:27:11 -07:00
parent cee13d5d70
commit 62049309a3
12 changed files with 301 additions and 64 deletions

View file

@ -7,6 +7,7 @@ mod execution;
mod function_trace;
mod lifecycle;
mod marshal;
mod params;
mod routes;
mod token_counter;

View file

@ -90,14 +90,14 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyRe
pub(crate) fn project_optional_fields(
kwargs: &Bound<'_, PyDict>,
names: &[&str],
route: &crate::params::RouteParamSpec<'_>,
policy: &crate::params::RequestParamPolicy,
) -> PyResult<Map<String, Value>> {
let selected: Vec<String> = kwargs
.py()
.import("litellm.rust_bridge.params")?
.getattr("provider_param_names")?
.call1((kwargs, names.to_vec()))?
.extract()?;
let keys: Vec<String> = kwargs.keys().extract()?;
let selected: Vec<String> = keys
.into_iter()
.filter(|name| policy.includes(name, route))
.collect();
selected
.iter()
.filter_map(|name| match kwargs.get_item(name) {

View file

@ -0,0 +1,160 @@
use std::collections::HashSet;
use pyo3::prelude::*;
use pyo3::types::PyList;
pub(crate) struct RequestParamPolicy {
sdk_reserved_param_names: HashSet<String>,
}
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<Self> {
let names: Vec<String> = 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::<PyTypeError>(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::<PyTypeError>(py));
});
}
}

View file

@ -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<PyAny> },
Unprojected {
request: Py<PyAny>,
sdk_reserved_param_names: Py<PyList>,
},
Projected(Box<ProjectedOcrHost>),
Released,
}
@ -186,10 +189,17 @@ impl PythonRoute for PythonOcrHost {
fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult<OcrHostResult> {
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<Py<PyAny>> {
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)

View file

@ -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<PyAny>,
pub document: Py<PyAny>,
@ -114,6 +125,7 @@ pub(super) fn project_request(
py: Python<'_>,
request: &Bound<'_, PyAny>,
kwargs: &Bound<'_, PyDict>,
policy: &crate::params::RequestParamPolicy,
) -> PyResult<ProjectedOcrCall> {
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::<Vec<_>>();
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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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