mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(rust): avoid retained adapter initialization deadlocks
This commit is contained in:
parent
391bb4cc74
commit
415f6894a8
6 changed files with 218 additions and 16 deletions
|
|
@ -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())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<()> {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue