diff --git a/litellm-rust/crates/host-python/src/argument.rs b/litellm-rust/crates/host-python/src/argument.rs index 0b89f94cd74..78879036f7e 100644 --- a/litellm-rust/crates/host-python/src/argument.rs +++ b/litellm-rust/crates/host-python/src/argument.rs @@ -19,6 +19,19 @@ pub fn present<'py>( Ok(lookup(kwargs, bound, name)?.filter(|value| !value.is_none())) } +/// What `original_function(*args, **kwargs)` sees: the signature base with the keyword +/// dict laid over it, so a rewritten keyword wins and a deleted keyword falls back to the +/// signature default. +pub fn effective_py_args<'py>( + base: &Bound<'py, PyDict>, + kwargs: &Bound<'py, PyDict>, +) -> PyResult> { + // Shallow copy, like the Python path: nested values stay shared with the caller. + let merged = base.copy()?; + merged.update(kwargs.as_mapping())?; + Ok(merged) +} + #[cfg(test)] mod tests { use rstest::rstest; @@ -95,4 +108,56 @@ mod tests { assert!(found.is(&document)); }); } + + 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::keyword_wins("{'api_key': 'base'}", "{'api_key': 'keyword'}", Some(Some("keyword")))] + #[case::explicit_none_wins("{'api_key': 'base'}", "{'api_key': None}", Some(None))] + #[case::base_default("{'api_key': 'base'}", "{}", Some(Some("base")))] + #[case::keyword_only("{}", "{'api_key': 'keyword'}", Some(Some("keyword")))] + #[case::missing("{}", "{}", None)] + fn effective_lays_the_keywords_over_the_base( + #[case] base: &str, + #[case] kwargs: &str, + #[case] expected: Option>, + ) { + crate::initialize_python(); + Python::attach(|py| { + let merged = effective_py_args(&dict(py, base), &dict(py, kwargs)).unwrap(); + let value = merged + .get_item("api_key") + .unwrap() + .map(|value| value.extract::>().unwrap()); + assert_eq!(value, expected.map(|value| value.map(str::to_owned))); + }); + } + + #[rstest] + fn effective_leaves_both_inputs_untouched_and_keeps_object_identity() { + crate::initialize_python(); + Python::attach(|py| { + let document = PyDict::new(py); + let base = dict(py, "{'model': 'base', 'pages': None}"); + let kwargs = PyDict::new(py); + kwargs.set_item("document", &document).unwrap(); + let merged = effective_py_args(&base, &kwargs).unwrap(); + merged.set_item("model", "merged").unwrap(); + assert_eq!( + base.get_item("model") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + "base" + ); + assert!(!kwargs.contains("model").unwrap()); + assert!(merged.get_item("document").unwrap().unwrap().is(&document)); + }); + } } diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 323f5b7a79b..95bc03c2174 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -19,7 +19,7 @@ mod owned; mod runtime; mod services; -pub use argument::{lookup, present}; +pub use argument::{effective_py_args, lookup, present}; pub use binding::PythonBinding; pub use conversion_cache::{FromPythonCache, ToPythonCache}; pub use driver::{CallOptions, run_call}; diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index 4c0c2760875..e62673d4d4f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -43,11 +43,11 @@ async fn execute( #[pyfunction] pub(crate) fn transcription(py: Python<'_>, call: NativeCall<'_>) -> PyResult> { + let arguments = call.resolved()?; let audio: Value = - litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; - let options = value_route_options(&call.bound)?; - let optional_params = - optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + litellm_host_python::from_py_argument(&required_field(&arguments, "audio")?)?; + let options = value_route_options(&arguments)?; + let optional_params = optional_object_field(&arguments, "optional_params")?.unwrap_or_default(); let http = crate::http::provider_client(py, &call.kwargs, false)?; let secrets = crate::secrets::source(py)?; run_sync( @@ -62,11 +62,11 @@ pub(crate) fn atranscription<'py>( py: Python<'py>, call: NativeCall<'py>, ) -> PyResult> { + let arguments = call.resolved()?; let audio: Value = - litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?; - let options = value_route_options(&call.bound)?; - let optional_params = - optional_object_field(&call.bound, "optional_params")?.unwrap_or_default(); + litellm_host_python::from_py_argument(&required_field(&arguments, "audio")?)?; + let options = value_route_options(&arguments)?; + let optional_params = optional_object_field(&arguments, "optional_params")?.unwrap_or_default(); let http = crate::http::provider_client(py, &call.kwargs, true)?; let secrets = crate::secrets::source(py)?; run_async( 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 1146436bfe6..2b69d0e6321 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.bound.clone().unbind(), + call.resolved()?.unbind(), "litellm.rust_bridge.chat_completions.route_host", ); run_inference::( diff --git a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs index 2b184274041..0f73eefedd5 100644 --- a/litellm-rust/crates/python-bridge/src/routes/embeddings.rs +++ b/litellm-rust/crates/python-bridge/src/routes/embeddings.rs @@ -32,7 +32,7 @@ mod tests { let locals = PyDict::new(py); py.run( c"from types import SimpleNamespace -call = SimpleNamespace(args=(), kwargs={}, bound={'model':'test-model','input':'hello'})", +call = SimpleNamespace(args=(), kwargs={'model':'test-model','input':'hello'}, base={})", Some(&locals), Some(&locals), ) 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 6ed6656cd21..5f69cf078b7 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.bound.unbind(), asynchronous), + MessagesPythonHost::new(call.resolved()?.unbind(), 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 deb29f12bab..803214656e9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -10,12 +10,16 @@ pub(crate) mod traces; use litellm_callbacks_legacy_python::{LegacyLogging, LoggingOperation, PublicCall}; use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol}; -use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls}; +use litellm_host_python::{ + HookChain, PythonBinding, PythonCallHooks, PythonHostCalls, effective_py_args, +}; use pyo3::{ prelude::*, types::{PyDict, PyMapping, PyTuple}, }; +/// The public call as Python bound it: `base` holds the positional arguments by name plus +/// the signature defaults, `kwargs` the caller's keyword dict. #[derive(FromPyObject)] pub(crate) struct NativeCall<'py> { #[pyo3(attribute)] @@ -23,7 +27,14 @@ pub(crate) struct NativeCall<'py> { #[pyo3(attribute, from_py_with = mapping_dict)] kwargs: Bound<'py, PyDict>, #[pyo3(attribute, from_py_with = mapping_dict)] - bound: Bound<'py, PyDict>, + base: Bound<'py, PyDict>, +} + +impl<'py> NativeCall<'py> { + /// The call before any hook ran, for the reads that admit or decline it. + fn resolved(&self) -> PyResult> { + effective_py_args(&self.base, &self.kwargs) + } } fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult> { @@ -42,7 +53,7 @@ fn call_hooks( call: &NativeCall<'_>, asynchronous: bool, ) -> PyResult<(Py, impl PythonCallHooks + use<>)> { - let call = PublicCall::capture(&call.bound, &call.args, &call.kwargs)?; + let call = PublicCall::capture(&call.base, &call.args, &call.kwargs)?; let arguments = call.arguments(py); Ok(( arguments, @@ -109,7 +120,7 @@ mod tests { .set_item("args", pyo3::types::PyTuple::empty(py)) .unwrap(); attributes.set_item("kwargs", &fields).unwrap(); - attributes.set_item("bound", &fields).unwrap(); + attributes.set_item("base", PyDict::new(py)).unwrap(); py.import("types") .unwrap() .getattr("SimpleNamespace") 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 4dcb2f4c53e..854cdb25f7b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -48,7 +48,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.bound.unbind()), + OcrPythonHost::new(call.resolved()?.unbind()), hooks, asynchronous, ) 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 a5b81b91752..50e5790a532 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses/mod.rs @@ -24,15 +24,16 @@ use crate::errors::RustBridgeDeclined; const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.responses.route_host"; fn run_responses(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult> { + let resolved = call.resolved()?; if let Some(reason) = py .import(ROUTE_HOST_MODULE)? .getattr("decline_reason")? - .call1((&call.bound,))? + .call1((&resolved,))? .extract::>()? { return Err(RustBridgeDeclined::new_err(reason)); } - let argument = |name: &str| present(&call.kwargs, &call.bound, name); + let argument = |name: &str| present(&call.kwargs, &call.base, name); let model = argument("model")? .ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))? .extract::()?; @@ -60,7 +61,7 @@ fn run_responses(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> Py "native Python responses streaming", )); } - let host = InferenceHost::new(call.bound.clone().unbind(), ROUTE_HOST_MODULE); + let host = InferenceHost::new(resolved.unbind(), ROUTE_HOST_MODULE); run_inference::(py, call, asynchronous, ResponsesPythonHost(host)) } diff --git a/litellm/chat_completions/dispatch.py b/litellm/chat_completions/dispatch.py index 2eee37d41ed..99900177820 100644 --- a/litellm/chat_completions/dispatch.py +++ b/litellm/chat_completions/dispatch.py @@ -58,14 +58,14 @@ def _public_request( messages: Final = optional_sequence(fields.get("messages")) if not isinstance(model, str) or messages is None: return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.CHAT_COMPLETIONS, - provider=optional_str(request.bound.get("custom_llm_provider")), - model=str(request.bound["model"]), + provider=optional_str(request.resolved.get("custom_llm_provider")), + model=str(request.resolved["model"]), ) diff --git a/litellm/embeddings/dispatch.py b/litellm/embeddings/dispatch.py index 54408691910..3fd8893464d 100644 --- a/litellm/embeddings/dispatch.py +++ b/litellm/embeddings/dispatch.py @@ -35,14 +35,14 @@ def _public_request( model: Final = fields.get("model") if not isinstance(model, str): return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.EMBEDDINGS, - provider=optional_str(request.bound.get("custom_llm_provider")), - model=str(request.bound["model"]), + provider=optional_str(request.resolved.get("custom_llm_provider")), + model=str(request.resolved["model"]), ) diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index 0afa5efc29d..bf3c5477886 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -63,7 +63,7 @@ class BedrockAudioTranscriptionRustDispatch: "optional_params": optional_params, "timeout_seconds": timeout_to_seconds(timeout), } - call: Final = NativeCall(args=(), kwargs=fields, bound=fields) + call: Final = NativeCall(args=(), kwargs=fields, base={}) return TranscriptionResponse(**rust(call)) return runtime.run( @@ -96,7 +96,7 @@ class BedrockAudioTranscriptionRustDispatch: "optional_params": optional_params, "timeout_seconds": timeout_to_seconds(timeout), } - call: Final = NativeCall(args=(), kwargs=fields, bound=fields) + call: Final = NativeCall(args=(), kwargs=fields, base={}) return TranscriptionResponse(**await rust(call)) return await runtime.arun( diff --git a/litellm/messages/dispatch.py b/litellm/messages/dispatch.py index 74736011f20..c7865132840 100644 --- a/litellm/messages/dispatch.py +++ b/litellm/messages/dispatch.py @@ -60,21 +60,23 @@ def _public_request( max_tokens: Final = fields.get("max_tokens") if not isinstance(model, str) or messages is None or not isinstance(max_tokens, int): return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _resolved_provider(request: NativeCall) -> str | None: try: - return get_llm_provider(str(request.bound["model"]), optional_str(request.bound.get("custom_llm_provider")))[1] + return get_llm_provider( + str(request.resolved["model"]), optional_str(request.resolved.get("custom_llm_provider")) + )[1] except BadRequestError: - return optional_str(request.bound.get("custom_llm_provider")) + return optional_str(request.resolved.get("custom_llm_provider")) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.MESSAGES, provider=_resolved_provider(request), - model=str(request.bound["model"]), + model=str(request.resolved["model"]), ) diff --git a/litellm/ocr/dispatch.py b/litellm/ocr/dispatch.py index a6a6c5d0c50..4fb4a6c709a 100644 --- a/litellm/ocr/dispatch.py +++ b/litellm/ocr/dispatch.py @@ -9,7 +9,7 @@ from litellm.rust_bridge import runtime from litellm.rust_bridge.catalog import Route, RouteContext from litellm.rust_bridge.dispatch import PublicDispatch from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR -from litellm.rust_bridge.public_call import NativeCall, native_call, native_call_hook, optional_str +from litellm.rust_bridge.public_call import NativeCall, native_call, native_call_hook, optional_str, signature __all__ = ("aocr", "ocr") @@ -38,18 +38,21 @@ def _bind_request( ) +_OCR: Final = signature(_bind_request) + + def _public_request(name: str, args: tuple[object, ...], kwargs: Mapping[str, object]) -> NativeCall: try: - fields: Final = _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation - return native_call(args, kwargs, fields) + _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds the public arguments before native validation except TypeError as error: raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None + return native_call(_OCR, args, kwargs) def _context(request: NativeCall) -> RouteContext: - prefix, separator, _ = str(request.bound["model"]).partition("/") - provider: Final = optional_str(request.bound.get("custom_llm_provider")) or (prefix if separator else None) - return RouteContext(Route.OCR, provider=provider, model=str(request.bound["model"])) + prefix, separator, _ = str(request.resolved["model"]).partition("/") + provider: Final = optional_str(request.resolved.get("custom_llm_provider")) or (prefix if separator else None) + return RouteContext(Route.OCR, provider=provider, model=str(request.resolved["model"])) _DISPATCH: Final = PublicDispatch( diff --git a/litellm/responses/dispatch.py b/litellm/responses/dispatch.py index c728dd0ea13..ad642fb427c 100644 --- a/litellm/responses/dispatch.py +++ b/litellm/responses/dispatch.py @@ -56,14 +56,14 @@ def _public_request( model: Final = fields.get("model") if not isinstance(model, str): return None - return native_call(args, kwargs, fields) + return native_call(legacy, args, kwargs) def _context(request: NativeCall) -> RouteContext: return RouteContext( Route.RESPONSES, - provider=optional_str(request.bound.get("custom_llm_provider")), - model=str(request.bound["model"]), + provider=optional_str(request.resolved.get("custom_llm_provider")), + model=str(request.resolved["model"]), ) diff --git a/litellm/rust_bridge/public_call.py b/litellm/rust_bridge/public_call.py index a25e9802593..e5700be89ad 100644 --- a/litellm/rust_bridge/public_call.py +++ b/litellm/rust_bridge/public_call.py @@ -84,17 +84,35 @@ def inference_decline_reason(parameters: tuple[str, ...], kwargs: Mapping[str, o return None +_BAGS: Final = frozenset({inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD}) + + @dataclass(frozen=True, slots=True) class NativeCall: + """A public call as Python bound it. + + ``base`` is the positional arguments by name plus the signature defaults, with no caller + keyword in it. ``kwargs`` is the caller's keyword dict, which callbacks may rewrite before + the request is decoded. The call the public function sees is ``kwargs`` laid over ``base``. + """ + args: tuple[object, ...] kwargs: Mapping[str, object] - bound: Mapping[str, object] + base: Mapping[str, object] + + @property + def resolved(self) -> Mapping[str, object]: + return MappingProxyType({**self.base, **self.kwargs}) -def native_call(args: tuple[object, ...], kwargs: Mapping[str, object], fields: Mapping[str, object]) -> NativeCall: - extra: Final = optional_mapping(fields.get("kwargs")) or MappingProxyType({}) - named: Final = {name: value for name, value in fields.items() if name != "kwargs"} - return NativeCall(args=args, kwargs=kwargs, bound=MappingProxyType({**named, **extra})) +def _without_bags(legacy: inspect.Signature, named: Mapping[str, object]) -> Mapping[str, object]: + return MappingProxyType({name: value for name, value in named.items() if legacy.parameters[name].kind not in _BAGS}) + + +def native_call(legacy: inspect.Signature, args: tuple[object, ...], kwargs: Mapping[str, object]) -> NativeCall: + positional: Final = legacy.bind_partial(*args) + positional.apply_defaults() + return NativeCall(args=args, kwargs=kwargs, base=_without_bags(legacy, positional.arguments)) NativeResultT: Final = TypeVar("NativeResultT") diff --git a/tests/test_litellm_rust/cache/test_python_cache.py b/tests/test_litellm_rust/cache/test_python_cache.py index 490643191ea..139859bdf55 100644 --- a/tests/test_litellm_rust/cache/test_python_cache.py +++ b/tests/test_litellm_rust/cache/test_python_cache.py @@ -446,7 +446,7 @@ def test_sync_rust_messages_calls_python_cache(recording_server: RecordingServer request: Final = NativeCall( args=(), kwargs=arguments, - bound={ + base={ "model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, diff --git a/tests/test_litellm_rust/messages/test_request_shaping.py b/tests/test_litellm_rust/messages/test_request_shaping.py index a1c99f245b4..3584a2a8593 100644 --- a/tests/test_litellm_rust/messages/test_request_shaping.py +++ b/tests/test_litellm_rust/messages/test_request_shaping.py @@ -251,7 +251,7 @@ async def test_native_messages_observes_runtime_capabilities_and_separate_caller request: Final = NativeCall( args=(), kwargs={"temperature": 0.2, "drop_params": True}, - bound={ + base={ "model": model, "messages": MESSAGES, "max_tokens": 16, diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 1eae999ff70..8767b273b68 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -624,7 +624,7 @@ def test_native_projection_errors_never_select_python( request: Final = NativeCall( args=(), kwargs={}, - bound={ + base={ "model": "mistral/mistral-ocr-latest", "document": OCR_DOCUMENT, "api_key": "test-key", diff --git a/tests/test_litellm_rust/support/cache.py b/tests/test_litellm_rust/support/cache.py index 352722f3269..5deacc7869f 100644 --- a/tests/test_litellm_rust/support/cache.py +++ b/tests/test_litellm_rust/support/cache.py @@ -51,7 +51,7 @@ async def invoke( request: Final = NativeCall( args=(), kwargs=arguments, - bound={ + base={ "model": RESPONSES_MODEL, "input": "hello", "stream": None, @@ -78,7 +78,7 @@ async def invoke( if route == "chat": if not native: return await litellm.acompletion(**parameters) - chat: Final = NativeCall(args=(), kwargs=parameters, bound=parameters) + chat: Final = NativeCall(args=(), kwargs=parameters, base={}) return await runtime.arun( RouteContext(Route.CHAT_COMPLETIONS), binding=NATIVE_ACOMPLETION, @@ -91,7 +91,7 @@ async def invoke( messages: Final = NativeCall( args=(), kwargs=parameters, - bound={ + base={ "model": MESSAGES_MODEL, "messages": list(MESSAGES), "max_tokens": 32, diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py index 5e4b7b284bb..8211b50cb8d 100644 --- a/tests/test_litellm_rust/test_inference.py +++ b/tests/test_litellm_rust/test_inference.py @@ -66,7 +66,7 @@ def native_call( "max_tokens": 32, **options, } - request: Final = NativeCall(args=(), kwargs=kwargs, bound=kwargs) + request: Final = NativeCall(args=(), kwargs=kwargs, base={}) return (_native.acompletion if asynchronous else _native.completion)(request) response_kwargs: Final = { "model": RESPONSES_MODEL, @@ -79,7 +79,7 @@ def native_call( response_request: Final = NativeCall( args=(), kwargs=response_kwargs, - bound={ + base={ "model": RESPONSES_MODEL, "input": "hello", "stream": None, diff --git a/tests/unit/chat_completions/test_dispatch.py b/tests/unit/chat_completions/test_dispatch.py index fbfc4875aa1..4ed75ea5f95 100644 --- a/tests/unit/chat_completions/test_dispatch.py +++ b/tests/unit/chat_completions/test_dispatch.py @@ -134,13 +134,13 @@ def test_native_receives_bound_request_and_original_call_shape() -> None: ) request, call_args, call_kwargs = captured[0] - assert request.bound["model"] == "anthropic/claude-sonnet-4-5" - assert request.bound["messages"] is MESSAGES - assert request.bound["stream"] is True - assert request.bound["api_key"] == "sk-test" - assert request.bound["base_url"] == "https://example.invalid" - assert request.bound["custom_llm_provider"] == "anthropic" - assert request.bound["extra_headers"] == {"x-test": "1"} + assert request.resolved["model"] == "anthropic/claude-sonnet-4-5" + assert request.resolved["messages"] is MESSAGES + assert request.resolved["stream"] is True + assert request.resolved["api_key"] == "sk-test" + assert request.resolved["base_url"] == "https://example.invalid" + assert request.resolved["custom_llm_provider"] == "anthropic" + assert request.resolved["extra_headers"] == {"x-test": "1"} assert request.kwargs is kwargs assert call_args == args assert call_kwargs == kwargs @@ -218,7 +218,7 @@ def test_public_completion_routes_through_dispatch(monkeypatch: pytest.MonkeyPat finally: NATIVE_COMPLETION.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -238,7 +238,7 @@ async def test_public_acompletion_routes_through_dispatch(monkeypatch: pytest.Mo finally: NATIVE_ACOMPLETION.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -257,10 +257,10 @@ def test_sync_completion_request_projects_public_arguments() -> None: expected: Final = ModelResponse() def native(request: NativeCall) -> ModelResponse: - assert request.bound["model"] == "test-model" - assert request.bound["messages"] == MESSAGES - assert request.bound["custom_llm_provider"] == "openai" - assert request.bound["stream"] is True + assert request.resolved["model"] == "test-model" + assert request.resolved["messages"] == MESSAGES + assert request.resolved["custom_llm_provider"] == "openai" + assert request.resolved["stream"] is True return expected binding: Final[NativeBinding[NativeCompletion]] = NativeBinding("completion", validate=lambda _: None) @@ -335,6 +335,6 @@ def test_internal_acompletion_marker_bypasses_native() -> None: def test_positional_parameters_remain_available_to_native_projection() -> None: request: Final = _DISPATCH.request(("anthropic/test-model", MESSAGES, 12.0, 0.25), {}) assert request is not None - assert request.bound["timeout"] == 12.0 - assert request.bound["temperature"] == 0.25 - assert request.bound["messages"] is MESSAGES + assert request.resolved["timeout"] == 12.0 + assert request.resolved["temperature"] == 0.25 + assert request.resolved["messages"] is MESSAGES diff --git a/tests/unit/embeddings/test_dispatch.py b/tests/unit/embeddings/test_dispatch.py index 88a2e7532c2..750c88cd235 100644 --- a/tests/unit/embeddings/test_dispatch.py +++ b/tests/unit/embeddings/test_dispatch.py @@ -33,9 +33,9 @@ def test_sync_embedding_request_projects_public_arguments() -> None: expected: Final = EmbeddingResponse(model="test-model", data=[]) def native(request: NativeCall) -> EmbeddingResponse: - assert request.bound["model"] == "test-model" - assert request.bound["input"] == "hello" - assert request.bound["custom_llm_provider"] == "openai" + assert request.resolved["model"] == "test-model" + assert request.resolved["input"] == "hello" + assert request.resolved["custom_llm_provider"] == "openai" return expected binding: Final[NativeBinding[Callable[[NativeCall], EmbeddingResponse]]] = NativeBinding( diff --git a/tests/unit/messages/test_dispatch.py b/tests/unit/messages/test_dispatch.py index 8668e11ee65..faaefa1437e 100644 --- a/tests/unit/messages/test_dispatch.py +++ b/tests/unit/messages/test_dispatch.py @@ -154,13 +154,13 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: ) assert result is expected request, call_args, call_kwargs = captured[0] - assert request.bound["model"] == "anthropic/claude-sonnet-4-5" - assert request.bound["messages"] is MESSAGES - assert request.bound["max_tokens"] == 16 - assert request.bound["stream"] is True - assert request.bound["api_key"] == "sk-test" - assert request.bound["api_base"] == "https://example.invalid" - assert request.bound["custom_llm_provider"] == "anthropic" + assert request.resolved["model"] == "anthropic/claude-sonnet-4-5" + assert request.resolved["messages"] is MESSAGES + assert request.resolved["max_tokens"] == 16 + assert request.resolved["stream"] is True + assert request.resolved["api_key"] == "sk-test" + assert request.resolved["api_base"] == "https://example.invalid" + assert request.resolved["custom_llm_provider"] == "anthropic" assert request.kwargs == kwargs assert request.kwargs["litellm_metadata"] is metadata assert call_args == args @@ -246,7 +246,7 @@ def test_anthropic_create_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc finally: NATIVE_MESSAGES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"] + assert [request.resolved["model"] for request in captured] == ["claude-sonnet-4-5"] @pytest.mark.asyncio @@ -268,7 +268,7 @@ async def test_anthropic_acreate_routes_through_dispatch(monkeypatch: pytest.Mon finally: NATIVE_AMESSAGES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["claude-sonnet-4-5"] + assert [request.resolved["model"] for request in captured] == ["claude-sonnet-4-5"] @pytest.mark.asyncio @@ -287,10 +287,10 @@ def test_sync_messages_request_projects_public_arguments() -> None: expected: Final = AnthropicMessagesResponse(model="claude-test") def native(request: NativeCall) -> AnthropicMessagesResponse: - assert request.bound["model"] == "claude-test" - assert request.bound["messages"] == MESSAGES - assert request.bound["max_tokens"] == 10 - assert request.bound["custom_llm_provider"] == "anthropic" + assert request.resolved["model"] == "claude-test" + assert request.resolved["messages"] == MESSAGES + assert request.resolved["max_tokens"] == 10 + assert request.resolved["custom_llm_provider"] == "anthropic" return expected binding: Final[NativeBinding[NativeMessages]] = NativeBinding("messages", validate=lambda _: None) diff --git a/tests/unit/ocr/test_dispatch.py b/tests/unit/ocr/test_dispatch.py index 3b0d75ae07a..34ee8e40552 100644 --- a/tests/unit/ocr/test_dispatch.py +++ b/tests/unit/ocr/test_dispatch.py @@ -72,13 +72,13 @@ def test_native_receives_normalized_positional_request_and_original_call_shape() request, call_args, call_kwargs = captured[0] assert result is expected - assert request.bound["model"] == "mistral/mistral-ocr-latest" - assert request.bound["document"] is document - assert request.bound["api_key"] == "test-key" - assert request.bound["api_base"] == "https://example.invalid" - assert request.bound["timeout"] is timeout - assert request.bound["custom_llm_provider"] == "mistral" - assert request.bound["extra_headers"] is extra_headers + assert request.resolved["model"] == "mistral/mistral-ocr-latest" + assert request.resolved["document"] is document + assert request.resolved["api_key"] == "test-key" + assert request.resolved["api_base"] == "https://example.invalid" + assert request.resolved["timeout"] is timeout + assert request.resolved["custom_llm_provider"] == "mistral" + assert request.resolved["extra_headers"] is extra_headers assert request.kwargs == kwargs assert request.kwargs["pages"] is pages assert call_args is args @@ -119,8 +119,8 @@ def test_native_preserves_keyword_model_and_document_in_original_call_shape() -> request, call_args, call_kwargs = captured[0] assert result is expected - assert request.bound["model"] == "mistral/mistral-ocr-latest" - assert request.bound["document"] is document + assert request.resolved["model"] == "mistral/mistral-ocr-latest" + assert request.resolved["document"] is document assert request.kwargs == kwargs assert call_args is args assert call_kwargs is kwargs @@ -271,7 +271,7 @@ def test_public_ocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPatch) -> finally: NATIVE_OCR.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"] + assert [request.resolved["model"] for request in captured] == ["mistral/mistral-ocr-latest"] @pytest.mark.asyncio @@ -297,4 +297,4 @@ async def test_public_aocr_routes_through_dispatch(monkeypatch: pytest.MonkeyPat finally: NATIVE_AOCR.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["mistral/mistral-ocr-latest"] + assert [request.resolved["model"] for request in captured] == ["mistral/mistral-ocr-latest"] diff --git a/tests/unit/responses/test_dispatch.py b/tests/unit/responses/test_dispatch.py index befe2d0000e..e6a0193a063 100644 --- a/tests/unit/responses/test_dispatch.py +++ b/tests/unit/responses/test_dispatch.py @@ -165,13 +165,13 @@ def test_native_receives_normalized_request_and_original_call_shape() -> None: request, call_args, call_kwargs = captured[0] assert result is response - assert request.bound["model"] == "anthropic/claude-sonnet-4-5" - assert request.bound["input"] is INPUT - assert request.bound["stream"] is True - assert request.bound["api_key"] == "sk-test" - assert request.bound["base_url"] == "https://example.invalid" - assert request.bound["custom_llm_provider"] == "anthropic" - assert request.bound["extra_headers"] is extra_headers + assert request.resolved["model"] == "anthropic/claude-sonnet-4-5" + assert request.resolved["input"] is INPUT + assert request.resolved["stream"] is True + assert request.resolved["api_key"] == "sk-test" + assert request.resolved["base_url"] == "https://example.invalid" + assert request.resolved["custom_llm_provider"] == "anthropic" + assert request.resolved["extra_headers"] is extra_headers assert request.kwargs == kwargs assert request.kwargs["litellm_metadata"] is metadata assert call_args == args @@ -262,7 +262,7 @@ def test_public_responses_routes_through_dispatch(monkeypatch: pytest.MonkeyPatc finally: NATIVE_RESPONSES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] @pytest.mark.asyncio @@ -284,7 +284,7 @@ async def test_public_aresponses_routes_through_dispatch(monkeypatch: pytest.Mon finally: NATIVE_ARESPONSES.reset() assert result is expected - assert [request.bound["model"] for request in captured] == ["gpt-4o"] + assert [request.resolved["model"] for request in captured] == ["gpt-4o"] def test_responses_with_retries_uses_the_dispatch_entrypoint(monkeypatch: pytest.MonkeyPatch) -> None: @@ -307,6 +307,6 @@ def test_positional_parameters_remain_available_to_native_projection() -> None: include: Final = ["reasoning.encrypted_content"] request: Final = _DISPATCH.request((INPUT, "openai/test-model", include, "Be brief", 16), {}) assert request is not None - assert request.bound["include"] is include - assert request.bound["instructions"] == "Be brief" - assert request.bound["max_output_tokens"] == 16 + assert request.resolved["include"] is include + assert request.resolved["instructions"] == "Be brief" + assert request.resolved["max_output_tokens"] == 16 diff --git a/tests/unit/rust_bridge/messages/test_route_host.py b/tests/unit/rust_bridge/messages/test_route_host.py index 16b9634beb9..e721b9624d2 100644 --- a/tests/unit/rust_bridge/messages/test_route_host.py +++ b/tests/unit/rust_bridge/messages/test_route_host.py @@ -81,7 +81,7 @@ def test_native_request_rejections_map_to_the_public_400() -> None: request: Final = NativeCall( args=(), kwargs=MappingProxyType({}), - bound={ + base={ "model": "anthropic/claude-sonnet-5", "messages": (), "max_tokens": 8, @@ -95,14 +95,14 @@ def test_native_request_rejections_map_to_the_public_400() -> None: rejected: Final = ValueError("claude-sonnet-5 does not support top_k=5") rejected.messages_request_error = True # pyright: ignore[reportAttributeAccessIssue] # marker the native host sets - mapped: Final = route_host.map_failure(rejected, request.bound, "anthropic") + mapped: Final = route_host.map_failure(rejected, request.resolved, "anthropic") assert isinstance(mapped, litellm.BadRequestError) assert mapped.status_code == 400 assert "does not support top_k=5" in mapped.message assert mapped.model == "claude-sonnet-5" assert not isinstance( - route_host.map_failure(ValueError("plain"), request.bound, "anthropic"), litellm.BadRequestError + route_host.map_failure(ValueError("plain"), request.resolved, "anthropic"), litellm.BadRequestError ) diff --git a/tests/unit/rust_bridge/messages/test_secrets.py b/tests/unit/rust_bridge/messages/test_secrets.py index 53e1376b978..1dab9604740 100644 --- a/tests/unit/rust_bridge/messages/test_secrets.py +++ b/tests/unit/rust_bridge/messages/test_secrets.py @@ -53,7 +53,7 @@ def _native_request() -> NativeCall: return NativeCall( args=(), kwargs=supplied, - bound={ + base={ "model": MESSAGES_MODEL, "messages": MESSAGES, "max_tokens": 8, diff --git a/tests/unit/rust_bridge/native_route_wheel_test.py b/tests/unit/rust_bridge/native_route_wheel_test.py index 268eca0f663..bc3269f5dd9 100644 --- a/tests/unit/rust_bridge/native_route_wheel_test.py +++ b/tests/unit/rust_bridge/native_route_wheel_test.py @@ -101,7 +101,7 @@ def load_native(native_path: Path) -> object: def route_call(route: str, api_base: str, outcome: str) -> SimpleNamespace: fields: Final = route_kwargs(route, api_base, outcome) - return SimpleNamespace(args=(), kwargs=fields, bound=fields) + return SimpleNamespace(args=(), kwargs=fields, base={}) def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: diff --git a/tests/unit/rust_bridge/ocr/test_route_host.py b/tests/unit/rust_bridge/ocr/test_route_host.py index 0b8af515a64..879edb951fc 100644 --- a/tests/unit/rust_bridge/ocr/test_route_host.py +++ b/tests/unit/rust_bridge/ocr/test_route_host.py @@ -10,7 +10,7 @@ from litellm.rust_bridge.public_call import NativeCall REQUEST: Final = NativeCall( args=(), kwargs={"req_format": "markdown"}, - bound={ + base={ "model": "mistral/mistral-ocr-latest", "document": {"type": "document_url", "document_url": "https://example.com/file.pdf"}, "api_key": "test-key", @@ -18,7 +18,6 @@ REQUEST: Final = NativeCall( "timeout": None, "custom_llm_provider": None, "extra_headers": None, - **{"req_format": "markdown"}, }, ) @@ -53,7 +52,7 @@ def test_rust_ocr_response_retains_provider_native_response(): def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> None: error: Final = RustUpstreamError(429, '{"message": "slow down"}', (("retry-after", "7"),)) - public_error: Final = map_failure(error, REQUEST.bound, "mistral") + public_error: Final = map_failure(error, REQUEST.resolved, "mistral") assert isinstance(public_error, litellm.RateLimitError) assert public_error.status_code == 429 @@ -66,7 +65,7 @@ def test_map_failure_builds_public_error_from_upstream_status_and_headers() -> N def test_map_failure_maps_upstream_401_to_authentication_error() -> None: error: Final = RustUpstreamError(401, '{"message": "Unauthorized"}', ()) - public_error: Final = map_failure(error, REQUEST.bound, "mistral") + public_error: Final = map_failure(error, REQUEST.resolved, "mistral") assert isinstance(public_error, litellm.AuthenticationError) assert public_error.status_code == 401 @@ -77,7 +76,7 @@ def test_map_failure_maps_upstream_401_to_authentication_error() -> None: def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: error: Final = RuntimeError("bridge exploded") - public_error: Final = map_failure(error, REQUEST.bound, "mistral") + public_error: Final = map_failure(error, REQUEST.resolved, "mistral") assert not isinstance(public_error, UpstreamFailure) assert isinstance(public_error, litellm.APIConnectionError) @@ -86,4 +85,4 @@ def test_map_failure_leaves_non_upstream_errors_unwrapped() -> None: def test_map_failure_reports_invalid_request_format_as_unsupported_params() -> None: with pytest.raises(litellm.UnsupportedParamsError, match="Invalid `req_format`: 'markdown'"): - raise map_failure(RustFormatError(), REQUEST.bound, "mistral") + raise map_failure(RustFormatError(), REQUEST.resolved, "mistral") diff --git a/tests/unit/rust_bridge/ocr/test_secrets.py b/tests/unit/rust_bridge/ocr/test_secrets.py index 5d051fbffe0..262fba1429b 100644 --- a/tests/unit/rust_bridge/ocr/test_secrets.py +++ b/tests/unit/rust_bridge/ocr/test_secrets.py @@ -64,7 +64,7 @@ def _native_request(api_base: str) -> NativeCall: return NativeCall( args=(), kwargs=supplied, - bound={ + base={ "model": OCR_MODEL, "document": OCR_DOCUMENT, "api_key": None, diff --git a/tests/unit/rust_bridge/test_public_call.py b/tests/unit/rust_bridge/test_public_call.py index 24d5db057d6..63029e288e2 100644 --- a/tests/unit/rust_bridge/test_public_call.py +++ b/tests/unit/rust_bridge/test_public_call.py @@ -3,7 +3,7 @@ from typing import Final import pytest -from litellm.rust_bridge.public_call import bind, native_call, signature +from litellm.rust_bridge.public_call import native_call, signature def _messages( @@ -17,45 +17,48 @@ def _messages( return None +_SIGNATURE: Final = signature(_messages) + + @pytest.mark.parametrize("supplied", ({}, {"api_key": None}, {"api_key": "explicit"})) -def test_native_call_preserves_omission_separately_from_bound_defaults(supplied: Mapping[str, object]) -> None: +def test_base_holds_positionals_and_defaults_and_never_a_keyword(supplied: Mapping[str, object]) -> None: messages: Final[Sequence[object]] = [{"role": "user", "content": "hello"}] args: Final = (128, messages, "model", 0.25) - fields: Final = bind(signature(_messages), args, supplied) - assert fields is not None - call: Final = native_call(args, supplied, fields) + call: Final = native_call(_SIGNATURE, args, supplied) assert call.args is args assert call.kwargs is supplied - assert call.bound == { + assert call.base == { "max_tokens": 128, "messages": messages, "model": "model", "temperature": 0.25, - "api_key": supplied.get("api_key"), + "api_key": None, } - assert call.bound["messages"] is messages + assert call.base["messages"] is messages + assert call.resolved["api_key"] == supplied.get("api_key") assert ("api_key" in call.kwargs) == ("api_key" in supplied) -def test_native_call_keeps_extra_option_objects_without_nested_kwargs() -> None: +def test_resolved_lays_the_keywords_over_the_base_and_equals_the_full_binding() -> None: messages: Final[Sequence[object]] = [] metadata: Final = {"trace": "caller"} - supplied: Final = {"metadata": metadata} - args: Final = (128, messages, "model") - fields: Final = bind(signature(_messages), args, supplied) - assert fields is not None + supplied: Final = {"model": "keyword-model", "metadata": metadata} + args: Final = (128, messages) - call: Final = native_call(args, supplied, fields) + call: Final = native_call(_SIGNATURE, args, supplied) - assert call.bound == { + assert "model" not in call.base + assert "kwargs" not in call.base + assert "metadata" not in call.base + assert call.resolved == { "max_tokens": 128, "messages": messages, - "model": "model", + "model": "keyword-model", "temperature": None, "api_key": None, "metadata": metadata, } - assert call.bound["metadata"] is metadata - assert supplied == {"metadata": metadata} + assert call.resolved["metadata"] is metadata + assert supplied == {"model": "keyword-model", "metadata": metadata} diff --git a/tests/unit/test_audio_transcription_rust_bridge.py b/tests/unit/test_audio_transcription_rust_bridge.py index 31ad5b72cc9..e3a6781dee7 100644 --- a/tests/unit/test_audio_transcription_rust_bridge.py +++ b/tests/unit/test_audio_transcription_rust_bridge.py @@ -99,7 +99,7 @@ def test_dispatch_marshals_audio_into_rust_call() -> None: "timeout_seconds": 5.0, } assert response.text == "rust" - assert bridge.calls == (NativeCall(args=(), kwargs=expected, bound=expected),) + assert bridge.calls == (NativeCall(args=(), kwargs=expected, base={}),) @pytest.mark.parametrize("disable", ("process", "environment")) @@ -164,7 +164,7 @@ async def test_async_dispatch_marshals_audio_into_rust_call() -> None: "timeout_seconds": 5.0, } assert response.text == "async rust" - assert bridge.calls == (NativeCall(args=(), kwargs=expected, bound=expected),) + assert bridge.calls == (NativeCall(args=(), kwargs=expected, base={}),) def test_bedrock_transcription_dispatches_to_rust_from_sdk_entrypoint() -> None: @@ -175,8 +175,8 @@ def test_bedrock_transcription_dispatches_to_rust_from_sdk_entrypoint() -> None: assert isinstance(response, litellm.TranscriptionResponse) assert response.text == "rust" - assert bridge.calls[0].bound["model"] == MODEL.removeprefix("bedrock/") - assert frozenset(bridge.calls[0].bound) == TRANSCRIPTION_FIELDS + assert bridge.calls[0].resolved["model"] == MODEL.removeprefix("bedrock/") + assert frozenset(bridge.calls[0].resolved) == TRANSCRIPTION_FIELDS @pytest.mark.asyncio @@ -187,8 +187,8 @@ async def test_bedrock_atranscription_dispatches_to_rust_from_sdk_entrypoint() - response: Final = await litellm.atranscription(model=MODEL, file=AUDIO_FILE) assert response.text == "async rust" - assert tuple(call.bound["model"] for call in bridge.calls) == (MODEL.removeprefix("bedrock/"),) - assert frozenset(bridge.calls[0].bound) == TRANSCRIPTION_FIELDS + assert tuple(call.resolved["model"] for call in bridge.calls) == (MODEL.removeprefix("bedrock/"),) + assert frozenset(bridge.calls[0].resolved) == TRANSCRIPTION_FIELDS @pytest.mark.asyncio