mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
abee1c14e7
commit
1bc85fad40
33 changed files with 247 additions and 145 deletions
|
|
@ -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));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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, _>(
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -53,7 +53,7 @@ def _native_request() -> NativeCall:
|
|||
return NativeCall(
|
||||
args=(),
|
||||
kwargs=supplied,
|
||||
bound={
|
||||
base={
|
||||
"model": MESSAGES_MODEL,
|
||||
"messages": MESSAGES,
|
||||
"max_tokens": 8,
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue