From 415f6894a8aa128097c11ada42649392dea0c67b Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sun, 6 Sep 2026 20:16:48 -0700 Subject: [PATCH] fix(rust): avoid retained adapter initialization deadlocks --- .../python-bridge/src/routes/retained_http.rs | 7 +- .../tests/fixtures/ocr_retained.py | 89 +++++++++++++++++++ .../tests/fixtures/retained_http_contract.py | 54 +++++++++++ .../python-bridge/tests/ocr_retained.rs | 12 +-- .../python-bridge/tests/retained_http.rs | 58 ++++++++++++ .../crates/python-interop/src/callback.rs | 14 +-- 6 files changed, 218 insertions(+), 16 deletions(-) diff --git a/litellm-rust/crates/python-bridge/src/routes/retained_http.rs b/litellm-rust/crates/python-bridge/src/routes/retained_http.rs index 350f1545d6f..53d4d6ced42 100644 --- a/litellm-rust/crates/python-bridge/src/routes/retained_http.rs +++ b/litellm-rust/crates/python-bridge/src/routes/retained_http.rs @@ -159,7 +159,7 @@ pub(crate) fn run_async(boundary: &Bound<'_, PyAny>) -> PyResult> { fn execution_module(py: Python<'_>) -> PyResult> { static MODULE: PyOnceLock> = 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()) } diff --git a/litellm-rust/crates/python-bridge/tests/fixtures/ocr_retained.py b/litellm-rust/crates/python-bridge/tests/fixtures/ocr_retained.py index 6c975c090ab..94daa1ce969 100644 --- a/litellm-rust/crates/python-bridge/tests/fixtures/ocr_retained.py +++ b/litellm-rust/crates/python-bridge/tests/fixtures/ocr_retained.py @@ -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" diff --git a/litellm-rust/crates/python-bridge/tests/fixtures/retained_http_contract.py b/litellm-rust/crates/python-bridge/tests/fixtures/retained_http_contract.py index 1b3c14995c0..f94db3386f0 100644 --- a/litellm-rust/crates/python-bridge/tests/fixtures/retained_http_contract.py +++ b/litellm-rust/crates/python-bridge/tests/fixtures/retained_http_contract.py @@ -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 = [] diff --git a/litellm-rust/crates/python-bridge/tests/ocr_retained.rs b/litellm-rust/crates/python-bridge/tests/ocr_retained.rs index c1c7d1e78c9..60128ff4a27 100644 --- a/litellm-rust/crates/python-bridge/tests/ocr_retained.rs +++ b/litellm-rust/crates/python-bridge/tests/ocr_retained.rs @@ -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(()) }) } diff --git a/litellm-rust/crates/python-bridge/tests/retained_http.rs b/litellm-rust/crates/python-bridge/tests/retained_http.rs index 1ff3514debe..2a4ea5c7892 100644 --- a/litellm-rust/crates/python-bridge/tests/retained_http.rs +++ b/litellm-rust/crates/python-bridge/tests/retained_http.rs @@ -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<()> { diff --git a/litellm-rust/crates/python-interop/src/callback.rs b/litellm-rust/crates/python-interop/src/callback.rs index bc22195f0d0..fbdf7962616 100644 --- a/litellm-rust/crates/python-interop/src/callback.rs +++ b/litellm-rust/crates/python-interop/src/callback.rs @@ -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)