refactor(python-bridge): ship a signature base and read the resolved call (#45450)

* refactor(python-bridge): ship a signature base and read the resolved call

NativeCall carries base (positionals by name plus signature defaults) instead
of the fully bound dict. resolved lays kwargs over base, which is what bound
held, so every pre-hook read and every route host keeps seeing the same values.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* fix(bedrock): build the transcription NativeCall with an empty base

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* test(dispatch): read the resolved call instead of bound

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

* refactor(host-python): rename effective to effective_py_args and note the shallow copy

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yujonglee 2026-10-08 17:20:29 -07:00 • committed by GitHub
parent abee1c14e7
commit 1bc85fad40
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
33 changed files with 247 additions and 145 deletions

View file

@ -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<Bound<'py, PyDict>> {
// 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::<PyDict>()
.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<Option<&str>>,
) {
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::<Option<String>>().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::<String>()
.unwrap(),
"base"
);
assert!(!kwargs.contains("model").unwrap());
assert!(merged.get_item("document").unwrap().unwrap().is(&document));
});
}
}

View file

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

View file

@ -43,11 +43,11 @@ async fn execute(
#[pyfunction]
pub(crate) fn transcription(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
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<Bound<'py, PyAny>> {
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(

View file

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

View file

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

View file

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

View file

@ -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<Bound<'py, PyDict>> {
effective_py_args(&self.base, &self.kwargs)
}
}
fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult<Bound<'py, PyDict>> {
@ -42,7 +53,7 @@ fn call_hooks(
call: &NativeCall<'_>,
asynchronous: bool,
) -> PyResult<(Py<PyDict>, 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")

View file

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

View file

@ -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<Py<PyAny>> {
let resolved = call.resolved()?;
if let Some(reason) = py
.import(ROUTE_HOST_MODULE)?
.getattr("decline_reason")?
.call1((&call.bound,))?
.call1((&resolved,))?
.extract::<Option<String>>()?
{
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::<String>()?;
@ -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::<ResponsesRoute, _>(py, call, asynchronous, ResponsesPythonHost(host))
}

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -53,7 +53,7 @@ def _native_request() -> NativeCall:
return NativeCall(
args=(),
kwargs=supplied,
bound={
base={
"model": MESSAGES_MODEL,
"messages": MESSAGES,
"max_tokens": 8,

View file

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

View file

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

View file

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

View file

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

View file

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