test(rust): property and parametrized tests for legacy callback contracts

The payload boundary of callbacks-legacy gets a model-based proptest: for any
JSON body, caller keywords and callback edit, keywords the route sends unchanged
reach pre_call as the caller's own objects, and the wire is the body pre_call
received as the callback left it. A parametrized test pins that a keyword the
bridge never reads keeps its identity through setup, the deployment hook,
check_limits and prepare.

Behaviour owned by the real Logging object is pinned end to end in the OCR
tests: a hypothesis version of the body property over HTTP, sync hooks seeing no
running event loop, retained payloads staying intact after the call, success
callbacks sharing one standard logging payload, and state stashed before a
blocking deployment hook raises reaching both failure callback families
This commit is contained in:
Yujong Lee 2026-09-18 14:41:50 -07:00
parent b4bfd92a2a
commit a737e3430a
5 changed files with 370 additions and 7 deletions

View file

@ -26,6 +26,7 @@ litellm-token-counter = { path = "crates/token-counter" }
litellm-host-python = { path = "crates/host-python" }
bytes = "1"
proptest = "1.7.0"
pyo3 = "0.29.2"
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"

View file

@ -16,5 +16,6 @@ serde_json.workspace = true
[dev-dependencies]
litellm-auth.workspace = true
proptest.workspace = true
rstest.workspace = true
serde_json.workspace = true

View file

@ -97,6 +97,43 @@ assert checked is prepared
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_keyword_the_bridge_never_reads_reaches_every_reader_as_the_callers_object(
#[case] asynchronous: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
opaque = object()
hooked = []
logger.hooks = {'pre': lambda kwargs: hooked.append(kwargs['vendor_extension']) or kwargs}
kwargs = {'logger': logger, 'vendor_extension': opaque}
",
);
let (mut logging, step) = begin(py, &locals, asynchronous);
let step = match step {
LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(),
step => step,
};
locals.set_item("prepared", arguments(py, step)).unwrap();
locals.set_item("asynchronous", asynchronous).unwrap();
run(
py,
&locals,
c"
assert prepared['vendor_extension'] is opaque
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked['vendor_extension'] is opaque
assert hooked == ([opaque] if asynchronous else []), hooked
",
);
});
}
#[test]
fn response_returned_by_the_post_call_hook_is_finalized_and_returned() {
Python::initialize();

View file

@ -2,10 +2,11 @@ use std::ffi::CStr;
use litellm_auth::SecretValue;
use litellm_callbacks::event::{CallEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{LifecycleStep, PythonLifecycle};
use litellm_host_python::{LifecycleStep, PythonLifecycle, to_py};
use proptest::prelude::*;
use pyo3::prelude::*;
use rstest::rstest;
use serde_json::{Value, json};
use serde_json::{Map, Value, json};
use super::LegacyLogging;
use crate::PythonLogger;
@ -57,10 +58,24 @@ fn before_send_with_secrets(
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
before_send_bound(&[], script, optional_params, body, secret_fields)
}
/// [`before_send_with_secrets`] with `bindings` placed in the namespace before `script` runs.
fn before_send_bound(
bindings: &[(&str, &Value)],
script: &CStr,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
for &(name, value) in bindings {
locals.set_item(name, to_py(py, value).unwrap()).unwrap();
}
run(py, &locals, script);
let mut logging = LegacyLogging {
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
@ -346,3 +361,159 @@ def check():
json!({"document": document(DOCUMENT), "include_image_base64": true})
);
}
/// What one pre-call callback does to the payload it is handed.
#[derive(Clone, Debug)]
enum Edit {
Nothing,
Set(String, Value),
Remove(String),
Rebind(Value),
RebindThenSetRetained(String, Value),
}
impl Edit {
fn script(&self) -> Value {
match self {
Self::Nothing => json!({"kind": "nothing"}),
Self::Set(key, value) => json!({"kind": "set", "key": key, "value": value}),
Self::Remove(key) => json!({"kind": "remove", "key": key}),
Self::Rebind(value) => json!({"kind": "rebind", "value": value}),
Self::RebindThenSetRetained(key, value) => {
json!({"kind": "rebind_then_set_retained", "key": key, "value": value})
}
}
}
/// The legacy contract: the provider is sent the body object `pre_call` received, as
/// the callback left it. Rebinding the envelope's key points the envelope elsewhere and
/// leaves that object alone.
fn sent(&self, body: &Map<String, Value>) -> Value {
let mut sent = body.clone();
match self {
Self::Nothing | Self::Rebind(_) => {}
Self::Set(key, value) | Self::RebindThenSetRetained(key, value) => {
sent.insert(key.clone(), value.clone());
}
Self::Remove(key) => {
sent.remove(key);
}
}
Value::Object(sent)
}
}
/// How the caller's keyword for a body key relates to what the route sends under it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Caller {
PassedUnchanged,
RewrittenByTheRoute,
NotPassed,
}
const MODEL: &CStr = c"
aliased = {}
def on_pre_call(args):
body = args['complete_input_dict']
aliased.update({name: body[name] is kwargs[name] for name in unchanged})
kind = edit['kind']
if kind == 'set':
body[edit['key']] = edit['value']
elif kind == 'remove':
body.pop(edit['key'], None)
elif kind == 'rebind':
args['complete_input_dict'] = edit['value']
elif kind == 'rebind_then_set_retained':
args['complete_input_dict'] = {}
body[edit['key']] = edit['value']
def check():
assert aliased == {name: True for name in unchanged}, aliased
assert logger.names() == ['pre_call', 'post_call'], logger.calls
";
fn json_value() -> impl Strategy<Value = Value> {
let leaf = prop_oneof![
Just(Value::Null),
any::<bool>().prop_map(Value::from),
any::<i64>().prop_map(Value::from),
any::<f64>()
.prop_filter("JSON has no NaN or infinity", |number| number.is_finite())
.prop_map(Value::from),
".{0,8}".prop_map(Value::from),
];
leaf.prop_recursive(3, 24, 4, |inner| {
prop_oneof![
prop::collection::vec(inner.clone(), 0..4).prop_map(Value::from),
prop::collection::btree_map(key(), inner, 0..4)
.prop_map(|fields| Value::Object(fields.into_iter().collect())),
]
})
}
fn key() -> impl Strategy<Value = String> {
"[a-z]{1,6}"
}
fn caller() -> impl Strategy<Value = Caller> {
prop_oneof![
Just(Caller::PassedUnchanged),
Just(Caller::RewrittenByTheRoute),
Just(Caller::NotPassed),
]
}
fn edit() -> impl Strategy<Value = Edit> {
prop_oneof![
Just(Edit::Nothing),
(key(), json_value()).prop_map(|(key, value)| Edit::Set(key, value)),
key().prop_map(Edit::Remove),
json_value().prop_map(Edit::Rebind),
(key(), json_value()).prop_map(|(key, value)| Edit::RebindThenSetRetained(key, value)),
]
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(128))]
/// For any body, any caller keywords and any callback edit: every keyword the route
/// sends unchanged reaches `pre_call` as the caller's own object, and the provider is
/// sent exactly what the model says, so a callback that edits nothing changes nothing.
#[test]
fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it(
fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5),
edit in edit(),
) {
let body: Map<String, Value> = fields
.iter()
.map(|(name, (value, _))| (name.clone(), value.clone()))
.collect();
let kwargs: Map<String, Value> = fields
.iter()
.filter_map(|(name, (value, caller))| match caller {
Caller::PassedUnchanged => Some((name.clone(), value.clone())),
Caller::RewrittenByTheRoute => Some((name.clone(), json!([value]))),
Caller::NotPassed => None,
})
.collect();
let unchanged: Value = fields
.iter()
.filter(|(_, (_, caller))| *caller == Caller::PassedUnchanged)
.map(|(name, _)| Value::from(name.clone()))
.collect();
let wire = before_send_bound(
&[
("kwargs", &Value::Object(kwargs)),
("unchanged", &unchanged),
("edit", &edit.script()),
],
MODEL,
json!({}),
Value::Object(body.clone()),
&[],
);
prop_assert_eq!(wire.body, edit.sent(&body));
prop_assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
}

View file

@ -1,24 +1,28 @@
import asyncio
import copy
import gc
import queue
import threading
from typing import Final
import pytest
from hypothesis import HealthCheck, given, settings
from hypothesis import strategies as st
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger
from tests.test_litellm_rust.support.callback_recorder import RecordingLogger, drain_logging
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
from tests.test_litellm_rust.support.requests import (
OCR_DOCUMENT,
OCR_RESPONSE,
call_native,
call_native_aocr,
call_native_ocr,
request_body,
request_headers,
)
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
pytestmark = pytest.mark.requires_rust_extension
@ -123,9 +127,7 @@ async def test_native_ocr_pre_call_nested_document_edit_updates_caller_callback_
"callbacks": [Retain(), Edit()],
}
response: Final = (
await call_native_aocr(ocr_server, **arguments)
if asynchronous
else call_native_ocr(ocr_server, **arguments)
await call_native_aocr(ocr_server, **arguments) if asynchronous else call_native_ocr(ocr_server, **arguments)
)
assert aliases == [True]
@ -291,6 +293,157 @@ def test_native_ocr_dispatches_each_callback_phase_once_when_logger_is_registere
assert "log_failure_event" not in recorder.names
JSON_SCALARS: Final = (
st.none()
| st.booleans()
| st.integers(min_value=-(2**63), max_value=2**63 - 1)
| st.floats(allow_nan=False, allow_infinity=False)
| st.text(max_size=8)
)
JSON_VALUES: Final = st.recursive(
JSON_SCALARS,
lambda children: st.lists(children, max_size=3) | st.dictionaries(st.text(max_size=6), children, max_size=3),
max_leaves=8,
)
LATEST_EDITS: Final[list[dict[str, object]]] = []
class ApplyLatestEdits(CustomLogger):
"""Registrations can outlive one hypothesis example, so every instance applies the current example's edits."""
def __init__(self, latest: list[dict[str, object]]) -> None:
super().__init__()
self.latest = latest
def log_pre_api_call(self, model, messages, kwargs):
request_body(kwargs).update(copy.deepcopy(self.latest[-1]))
@settings(max_examples=25, deadline=None, suppress_health_check=[HealthCheck.function_scoped_fixture])
@given(edits=st.dictionaries(st.from_regex(r"x_[a-z]{1,6}", fullmatch=True), JSON_VALUES, max_size=3))
def test_native_ocr_provider_receives_the_body_exactly_as_pre_call_callbacks_left_it(
ocr_server: RecordingServer, edits: dict[str, object]
) -> None:
LATEST_EDITS.append(edits)
call_native_ocr_with_callbacks(ocr_server, [ApplyLatestEdits(LATEST_EDITS)])
assert ocr_server.requests[-1].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT, **edits}
@pytest.mark.parametrize("hook", ["log_pre_api_call", "logging_hook", "log_success_event"])
def test_native_ocr_sync_hooks_see_no_running_event_loop(ocr_server: RecordingServer, hook: str) -> None:
recorder: Final = RecordingLogger()
call_native_ocr_with_callbacks(ocr_server, [recorder])
[event] = recorder.wait_for(hook)
assert event.loop is None
assert (event.thread is threading.current_thread()) == (hook == "log_pre_api_call")
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_ocr_payload_a_callback_retains_outlives_the_call_intact(
ocr_server: RecordingServer, asynchronous: bool
) -> None:
retained: Final = []
class Retain(CustomLogger):
def log_pre_api_call(self, model, messages, kwargs):
retained.append((kwargs, request_body(kwargs), request_headers(kwargs)))
await call_native(ocr_server, asynchronous, callbacks=[Retain()])
await drain_logging()
gc.collect()
[(details, body, headers)] = retained
assert body == ocr_server.requests[0].body
assert headers
assert all(ocr_server.requests[0].headers[name] == value for name, value in headers.items())
assert details["additional_args"]["complete_input_dict"] is body
assert details["additional_args"]["headers"] is headers
@pytest.mark.asyncio
@pytest.mark.parametrize("family", ["sync", "async"])
async def test_native_ocr_success_callbacks_share_one_logging_payload(ocr_server: RecordingServer, family: str) -> None:
queued: Final = []
finished: Final = threading.Event()
def queue_payload(kwargs: dict[str, object]) -> None:
queued.append(kwargs["standard_logging_object"])
def strip_payload(kwargs: dict[str, object]) -> None:
payload: Final = kwargs["standard_logging_object"]
assert isinstance(payload, dict)
payload["stripped-by-a-later-callback"] = True
finished.set()
class QueuePayload(CustomLogger):
if family == "sync":
def log_success_event(self, kwargs, response_obj, start_time, end_time):
queue_payload(kwargs)
else:
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
queue_payload(kwargs)
class StripPayload(CustomLogger):
if family == "sync":
def log_success_event(self, kwargs, response_obj, start_time, end_time):
strip_payload(kwargs)
else:
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
strip_payload(kwargs)
await call_native(ocr_server, family == "async", callbacks=[QueuePayload(), StripPayload()])
await drain_logging()
assert await asyncio.to_thread(finished.wait, 10)
assert [payload["stripped-by-a-later-callback"] for payload in queued] == [True]
@pytest.mark.asyncio
async def test_native_aocr_state_stashed_before_a_blocking_hook_raises_reaches_failure_callbacks(
ocr_server: RecordingServer,
) -> None:
token: Final = object()
observed: Final = []
class Blocked(Exception):
pass
class Block(CustomLogger):
async def async_post_call_success_deployment_hook(self, request_data, response, call_type):
request_data["litellm_logging_obj"].model_call_details["blocked-by"] = token
raise Blocked("blocked after the provider answered")
def log_success_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("success", None, None))
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("sync", kwargs.get("blocked-by"), kwargs["exception"]))
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
observed.append(("async", kwargs.get("blocked-by"), kwargs["exception"]))
litellm.callbacks.append(Block())
with pytest.raises(Blocked) as raised:
await call_native_aocr(ocr_server)
await drain_logging()
assert observed == [("sync", token, raised.value), ("async", token, raised.value)]
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True], ids=["sync", "async"])
async def test_native_azure_ocr_resolves_token_before_pre_call_on_caller_context(