fix(rust): avoid retained adapter initialization deadlocks

This commit is contained in:
Yujong Lee 2026-09-06 20:16:48 -07:00
parent 391bb4cc74
commit 415f6894a8
6 changed files with 218 additions and 16 deletions

View file

@ -159,7 +159,7 @@ pub(crate) fn run_async(boundary: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
fn execution_module(py: Python<'_>) -> PyResult<Bound<'_, PyModule>> {
static MODULE: PyOnceLock<Py<PyModule>> = PyOnceLock::new();
let module = MODULE.get_or_try_init(py, || {
if MODULE.get(py).is_none() {
let module = PyModule::from_code(
py,
c"class RetainedExecution:
@ -196,7 +196,8 @@ async def drive(execution):
module.add("_encode", wrap_pyfunction!(encode, &module)?)?;
module.add("_finish", wrap_pyfunction!(finish, &module)?)?;
module.add("_send", wrap_pyfunction!(send, &module)?)?;
Ok::<_, PyErr>(module.unbind())
})?;
let _ = MODULE.set(py, module.unbind());
}
let module = MODULE.get(py).unwrap();
Ok(module.bind(py).clone())
}

View file

@ -14,6 +14,8 @@ from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from types import ModuleType
from unittest.mock import patch
import httpx
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging
@ -308,6 +310,83 @@ class RealBoundaryTests(unittest.TestCase):
asyncio.run(exercise())
@unittest.expectedFailure
def test_known_gap_send_time_auth_and_request_hook(self):
"""UC-HTTPX-SEND: pre_call, client auth, then request hook must determine the provider's wire headers."""
observations = []
for mode in ("python-sync", "native-sync"):
events = []
def request_hook(request, events=events):
events.append(("request_hook", request.headers["Authorization"]))
request.headers["X-Send-Hook"] = "present"
callback = Callback(
lambda view, events=events: events.append(("pre_call", view["headers"]["Authorization"]))
)
with httpx.Client(auth=("local-user", "local-password"), event_hooks={"request": [request_hook]}) as client:
before = len(self.server.requests)
response = invoke(mode, inputs(self.server, [callback], client=HTTPHandler(client=client)))
self.check_callbacks([callback])
self.assertEqual(response.pages[0].markdown, "local OCR")
self.assertEqual(len(self.server.requests), before + 1)
headers = dict(self.server.requests[-1][1])
observations.append((events, headers["authorization"], headers.get("x-send-hook")))
baseline, retained = observations
self.assertEqual(baseline[0][0], ("pre_call", "Bearer local-test-key"))
self.assertEqual(baseline[0][1], ("request_hook", baseline[1]))
self.assertTrue(baseline[1].startswith("Basic "))
self.assertEqual(baseline[2], "present")
self.assertEqual(retained, baseline, "UC-HTTPX-SEND: retained send omitted client auth/request hook")
@unittest.expectedFailure
def test_known_gap_custom_transport(self):
"""UC-HTTPX-SEND: after pre_call, the client's transport must select the response without a network POST."""
observations = []
for mode in ("python-sync", "native-sync"):
events = []
def transport(request, events=events):
events.append("transport")
return httpx.Response(200, json={**RESPONSE, "pages": [{"index": 0, "markdown": "custom transport"}]})
callback = Callback(lambda view, events=events: events.append("pre_call"))
with httpx.Client(transport=httpx.MockTransport(transport)) as client:
before = len(self.server.requests)
response = invoke(mode, inputs(self.server, [callback], client=HTTPHandler(client=client)))
self.check_callbacks([callback])
observations.append((events, len(self.server.requests) - before, response.pages[0].markdown))
baseline, retained = observations
self.assertEqual(baseline, (["pre_call", "transport"], 0, "custom transport"))
self.assertEqual(retained, baseline, "UC-HTTPX-SEND: retained send bypassed the custom transport")
@unittest.expectedFailure
def test_known_gap_pre_call_timeout_mutation(self):
"""UC-TIMEOUT-MUTATION: a caller's timeout is changed by pre_call; encoding must accept read=None."""
observations = []
for mode in ("python-sync", "native-sync"):
timeout = httpx.Timeout(5.0)
def mutate_timeout(view, timeout=timeout):
timeout.read = None
callback = Callback(mutate_timeout)
kwargs = inputs(self.server, [callback], client=self.sync_client)
kwargs["timeout"] = timeout
before = len(self.server.requests)
try:
response = invoke(mode, kwargs)
except ValueError as error:
outcome = ("error", str(error))
else:
outcome = ("success", response.pages[0].markdown)
self.check_callbacks([callback])
self.assertIsNone(timeout.read)
observations.append((outcome, len(self.server.requests) - before))
baseline, retained = observations
self.assertEqual(baseline, (("success", "local OCR"), 1))
self.assertEqual(retained, baseline, "UC-TIMEOUT-MUTATION: retained encoding rejected the callback's timeout")
def check_negative_control(self, boundary_factory, failure, expected_requests):
async def exercise():
baseline = await self.differential("python-sync")
@ -633,3 +712,13 @@ class RealBoundaryTests(unittest.TestCase):
self.assertEqual(len(self.server.requests), 1)
asyncio.run(exercise())
def run_case(scenario):
suite = unittest.TestSuite([RealBoundaryTests("test_" + scenario)])
result = unittest.TextTestRunner(verbosity=2).run(suite)
for _, traceback in result.expectedFailures:
assert "AssertionError:" in traceback and (
"UC-HTTPX-SEND:" in traceback or "UC-TIMEOUT-MUTATION:" in traceback
), traceback
assert result.wasSuccessful(), "OCR retained parity test failed or unexpectedly succeeded"

View file

@ -4,6 +4,7 @@ import gc
import http.client
import http.server
import inspect
import sys
import threading
import weakref
from copy import deepcopy
@ -146,6 +147,59 @@ def collected(boundary):
assert boundary.view == {"headers": {"replacement": True}, "body": {"replacement": True}}
def run_cold_cache_reentry(filename):
"""UC-COLD-REENTRY: cold compilation permits nested native calls and subsequent transport."""
events = []
error = LookupError("cold-cache preparation error")
class FailingBoundary:
async def aprepare(self):
raise error
def invoke_failure():
pending = native.aocr_retained(FailingBoundary())
observed = None
try:
pending.send(None)
except LookupError as caught:
observed = caught
finally:
pending.close()
error.__traceback__ = None
assert observed is error
def audit(event, args):
if event == "compile" and args[1] == filename and not events:
events.append("entered")
print(f"UC-COLD-REENTRY: entering {filename}", flush=True) # noqa: T201 # diagnose a deadlocked child process
invoke_failure()
events.append("nested completed")
async def successful_async():
context.set("initial")
boundary = Boundary(asynchronous=True, nested=True)
assert await native.aocr_retained(boundary) is boundary.result
assert boundary.events == ["prepare", "encode", "finish"]
collected(boundary)
try:
if filename == "retained_callback.py":
native.aocr_retained(object()).close()
sys.addaudithook(audit)
invoke_failure()
events.append("outer completed")
assert events == ["entered", "nested completed", "outer completed"]
assert requests == []
boundary = Boundary(nested=True)
assert native.ocr_retained(boundary) is boundary.result
assert boundary.events == ["prepare", "encode", "finish"]
collected(boundary)
asyncio.run(successful_async())
assert len(requests) == 2
finally:
stop_server()
def check_error(boundary, error, phase):
assert error is boundary.error
names = []

View file

@ -9,6 +9,9 @@ use support::native::{native_globals, run_fixture};
#[rstest]
#[case::differential_callbacks_wire("differential_callbacks_wire")]
#[case::known_gap_send_time_auth_and_request_hook("known_gap_send_time_auth_and_request_hook")]
#[case::known_gap_custom_transport("known_gap_custom_transport")]
#[case::known_gap_pre_call_timeout_mutation("known_gap_pre_call_timeout_mutation")]
#[case::negative_control_copied_caller_document("negative_control_copied_caller_document")]
#[case::negative_control_rebound_logging_body("negative_control_rebound_logging_body")]
#[case::negative_control_rebound_logging_headers("negative_control_rebound_logging_headers")]
@ -42,14 +45,7 @@ fn retained_real_production_boundary_differential_and_lifecycle(
"/tests/fixtures/ocr_retained.py"
),
)?;
let case = globals
.get_item("RealBoundaryTests")?
.unwrap()
.call1((format!("test_{scenario}"),))?;
let outcome = case.call_method0("debug");
let cleanup = case.call_method0("doCleanups");
outcome?;
cleanup?;
globals.get_item("run_case")?.unwrap().call1((scenario,))?;
Ok(())
})
}

View file

@ -1,3 +1,6 @@
use std::process::Command;
use std::time::{Duration, Instant};
use pyo3::prelude::*;
use rstest::rstest;
use serial_test::serial;
@ -7,6 +10,61 @@ mod support;
use support::native::{native_globals, run_fixture};
#[test]
fn uc_cold_reentry_execution_cache() -> PyResult<()> {
cold_cache_reentry("retained_execution.py", "uc_cold_reentry_execution_cache")
}
#[test]
fn uc_cold_reentry_callback_cache() -> PyResult<()> {
cold_cache_reentry("retained_callback.py", "uc_cold_reentry_callback_cache")
}
fn cold_cache_reentry(filename: &str, test: &str) -> PyResult<()> {
if std::env::var("LITELLM_COLD_REENTRY_CHILD").as_deref() != Ok(filename) {
let mut child = Command::new(std::env::current_exe().unwrap())
.args(["--exact", test, "--nocapture"])
.env("LITELLM_COLD_REENTRY_CHILD", filename)
.spawn()
.unwrap();
let deadline = Instant::now() + Duration::from_secs(15);
loop {
if let Some(status) = child.try_wait().unwrap() {
assert!(
status.success(),
"{filename} reentry child failed: {status}"
);
return Ok(());
}
if Instant::now() >= deadline {
child.kill().unwrap();
child.wait().unwrap();
panic!("{filename} reentry did not complete within 15 seconds");
}
std::thread::sleep(Duration::from_millis(10));
}
}
Python::initialize();
Python::attach(|py| {
let globals = native_globals(py)?;
run_fixture(
py,
&globals,
include_str!("fixtures/retained_http_contract.py"),
concat!(
env!("CARGO_MANIFEST_DIR"),
"/tests/fixtures/retained_http_contract.py"
),
)?;
globals
.get_item("run_cold_cache_reentry")?
.unwrap()
.call1((filename,))?;
Ok(())
})
}
#[test]
#[serial(python_interpreter)]
fn retained_routes_preserve_callbacks_context_wire_and_ownership() -> PyResult<()> {

View file

@ -50,8 +50,8 @@ impl PreparedCall {
)
.map(InvocationOutcome::Returned),
InvocationMode::Await => {
let adapter = AWAIT_CALL.get_or_try_init(py, || {
PyModule::from_code(
if AWAIT_CALL.get(py).is_none() {
let adapter = PyModule::from_code(
py,
c"async def invoke_awaited(callable, positional, keywords):
if keywords is None:
@ -61,9 +61,13 @@ impl PreparedCall {
c"retained_callback.py",
c"_retained_callback",
)?
.getattr("invoke_awaited")
.map(Bound::unbind)
})?;
.getattr("invoke_awaited")?
.unbind();
let _ = AWAIT_CALL.set(py, adapter);
}
let adapter = AWAIT_CALL.get(py).unwrap();
adapter
.call1(py, (&self.callable, &self.positional, &self.keywords))
.map(InvocationOutcome::Awaitable)