diff --git a/CHANGELOG.md b/CHANGELOG.md index 4aa085f..b3e132d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,7 @@ Unreleased entry. - Rust checks, regression tests, pedantic Clippy, and Python API tests on Python 3.10 and 3.15. ### Changed +- Use a current-thread Tokio runtime per batch and move Python result delivery to the blocking pool. Bound queued results by the batch concurrency. - **Breaking:** Rename the Python distribution, import module, and Rust crate to `proxyprobe`; update dependencies and imports from `rsloop-rust-proxychecker` / `rsloop_rust_proxychecker`. Function signatures and result dictionaries remain unchanged. - Split the extension implementation into API, configuration, checker, worker, and stream modules without changing the public Python API. - Enable PyO3 extension-module mode through Maturin so Rust unit tests can embed Python normally. @@ -26,6 +27,7 @@ Unreleased entry. - Adapt proxy parsing, idle connection configuration, and body streaming to the current `wreq` API. ### Fixed +- Cancel in-flight checks when the stream is dropped or a pending iteration is cancelled; propagate worker panics and thread startup failures instead of leaving consumers waiting. - Reject unsupported proxy schemes before sending requests, preventing silent direct-request fallback with the current `wreq`; preserve bare `host:port` proxy inputs. - Drain response bodies without accumulating them when `return_response=false`, preserving body-read failures and full body return when enabled (PR #1). diff --git a/Cargo.lock b/Cargo.lock index d600aaf..9ddb0af 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1123,6 +1123,19 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "proxyprobe" +version = "0.1.0" +dependencies = [ + "async-std", + "futures", + "pyo3", + "rsloop", + "tokio", + "url", + "wreq", +] + [[package]] name = "pyo3" version = "0.29.3" @@ -1292,19 +1305,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "proxyprobe" -version = "0.1.0" -dependencies = [ - "async-std", - "futures", - "pyo3", - "rsloop", - "tokio", - "url", - "wreq", -] - [[package]] name = "rustc-hash" version = "2.1.2" diff --git a/README.md b/README.md index 01e351e..586e376 100644 --- a/README.md +++ b/README.md @@ -12,6 +12,18 @@ It uses: - `wreq` for the HTTP client and proxy support - a dedicated Tokio runtime inside the Rust worker because `wreq` runs on Tokio +Each stream uses a current-thread Tokio runtime with I/O and timers enabled. +Python result delivery runs on Tokio’s blocking pool. Pending result delivery is +bounded by `concurrency`, so a slow consumer pauses further scheduling. At most +`concurrency` queued item messages, one pending delivery outcome, and the bounded +in-flight checks can retain results (plus terminal messages). + +Dropping the stream or cancelling a pending `__anext__` stops its Rust checks. +Keep consuming the stream if you want every result; cancelling an iteration ends +the whole batch. Worker panics are reported as `RuntimeError` when delivery remains +possible. Cancellation cannot interrupt a Python callback already running on the +blocking pool. + ## Migration The distribution and import name are now `proxyprobe` (previously diff --git a/src/api.rs b/src/api.rs index a3006fc..3043da9 100644 --- a/src/api.rs +++ b/src/api.rs @@ -2,8 +2,10 @@ use std::time::Duration; -use pyo3::exceptions::PyValueError; +use pyo3::exceptions::{PyRuntimeError, PyValueError}; use pyo3::prelude::*; +use std::sync::Arc; +use tokio::sync::{watch, Semaphore}; use crate::config::{CheckerConfig, DEFAULT_CHECK_URL, DEFAULT_CONCURRENCY, DEFAULT_TIMEOUT_MS}; use crate::stream::PyProxyCheckStream; @@ -55,10 +57,14 @@ pub(crate) fn check_proxies( }; let locals = rsloop::rust_async::get_current_locals(py)?; let queue = py.import("asyncio")?.getattr("Queue")?.call0()?.unbind(); + let (cancel, cancelled) = watch::channel(false); + let credits = Arc::new(Semaphore::new(config.concurrency.min(proxies.len().max(1)))); let stream = Py::new( py, PyProxyCheckStream { queue: queue.clone_ref(py), + cancel, + credits: credits.clone(), }, )? .into_any(); @@ -68,7 +74,10 @@ pub(crate) fn check_proxies( config.clone(), locals.event_loop(py).unbind(), queue, - ); + cancelled, + credits, + ) + .map_err(PyRuntimeError::new_err)?; rsloop::rust_async::future_into_py(py, async move { let _ = config; Ok(stream) diff --git a/src/stream.rs b/src/stream.rs index eea883a..05993ec 100644 --- a/src/stream.rs +++ b/src/stream.rs @@ -3,6 +3,8 @@ use pyo3::exceptions::{PyRuntimeError, PyStopAsyncIteration}; use pyo3::prelude::*; use pyo3::types::PyDict; +use std::sync::Arc; +use tokio::sync::{watch, Semaphore}; use crate::checker::ProxyOutcome; @@ -10,6 +12,31 @@ use crate::checker::ProxyOutcome; #[pyclass(module = "proxyprobe")] pub(crate) struct PyProxyCheckStream { pub(crate) queue: Py, + pub(crate) cancel: watch::Sender, + pub(crate) credits: Arc, +} + +impl Drop for PyProxyCheckStream { + fn drop(&mut self) { + let _ = self.cancel.send(true); + } +} + +/// Cancel a pending iteration if its Rust future is dropped before completion. +struct PendingIteration(Option>); + +impl PendingIteration { + fn complete(mut self) { + self.0 = None; + } +} + +impl Drop for PendingIteration { + fn drop(&mut self) { + if let Some(cancel) = &self.0 { + let _ = cancel.send(true); + } + } } /// Convert an outcome to the public dictionary, omitting absent optional fields. @@ -97,7 +124,11 @@ impl PyProxyCheckStream { // PyO3 requires an owned receiver for this async iterator slot. #[allow(clippy::needless_pass_by_value)] fn __anext__(slf: Py, py: Python<'_>) -> PyResult> { - let queue = slf.borrow(py).queue.clone_ref(py); + let stream = slf.borrow(py); + let queue = stream.queue.clone_ref(py); + let credits = stream.credits.clone(); + let pending = PendingIteration(Some(stream.cancel.clone())); + drop(stream); let locals = rsloop::rust_async::get_current_locals(py)?; rsloop::rust_async::future_into_py_with_locals(py, locals.clone(), async move { @@ -107,7 +138,18 @@ impl PyProxyCheckStream { })? .await?; - Python::attach(|py| decode_message(queued.bind(py))) + pending.complete(); + // Only item messages consumed a credit. Return it even if decoding fails. + Python::attach(|py| { + let message = queued.bind(py).cast::()?; + if message + .get_item("kind")? + .is_some_and(|kind| kind.extract::().is_ok_and(|kind| kind == "item")) + { + credits.add_permits(1); + } + decode_message(queued.bind(py)) + }) }) } } diff --git a/src/stream_tests.rs b/src/stream_tests.rs index 17d7962..d3d7977 100644 --- a/src/stream_tests.rs +++ b/src/stream_tests.rs @@ -153,3 +153,23 @@ fn queue_delivery_propagates_event_loop_errors() { assert!(emit_stream_end(&closed_loop, &queue).is_err()); }); } + +#[test] +fn stream_drop_and_pending_iteration_cancel_but_completion_does_not() { + let (cancel, cancelled) = tokio::sync::watch::channel(false); + PendingIteration(Some(cancel.clone())).complete(); + assert!(!*cancelled.borrow()); + drop(PendingIteration(Some(cancel.clone()))); + assert!(*cancelled.borrow()); + cancel.send(false).unwrap(); + Python::initialize(); + Python::attach(|py| { + let stream = PyProxyCheckStream { + queue: py.None(), + cancel, + credits: std::sync::Arc::new(tokio::sync::Semaphore::new(1)), + }; + drop(stream); + }); + assert!(*cancelled.borrow()); +} diff --git a/src/worker.rs b/src/worker.rs index a472875..987f9ea 100644 --- a/src/worker.rs +++ b/src/worker.rs @@ -2,26 +2,36 @@ use futures::stream::{self, StreamExt}; use pyo3::prelude::*; +use std::sync::Arc; use tokio::runtime::Builder as TokioRuntimeBuilder; +use tokio::sync::{watch, Semaphore}; use crate::checker::{build_client, check_one_proxy}; use crate::config::CheckerConfig; use crate::stream::{emit_stream_end, emit_stream_error, emit_stream_result}; -/// Start an independent worker and always schedule a terminal stream message. +/// Start an independent worker; turn initialization failures and panics into stream errors. pub(crate) fn spawn_proxy_checks( proxies: Vec, config: CheckerConfig, loop_obj: Py, queue: Py, -) { - std::thread::spawn(move || { - let result = run_proxy_checks_blocking(proxies, config.clone(), &loop_obj, &queue); - if let Err(err) = result { - let _ = emit_stream_error(&loop_obj, &queue, err); - } - let _ = emit_stream_end(&loop_obj, &queue); - }); + cancelled: watch::Receiver, + credits: Arc, +) -> Result<(), String> { + std::thread::Builder::new() + .name("proxyprobe-worker".into()) + .spawn(move || { + let result = catch_worker_failure(|| { + run_proxy_checks_blocking(proxies, config, &loop_obj, &queue, cancelled, credits) + }); + if let Err(err) = result { + let _ = emit_stream_error(&loop_obj, &queue, err); + } + let _ = emit_stream_end(&loop_obj, &queue); + }) + .map_err(|err| format!("failed to start proxy worker: {err}"))?; + Ok(()) } /// Create a dedicated Tokio runtime for HTTP work on the worker thread. @@ -30,21 +40,36 @@ fn run_proxy_checks_blocking( config: CheckerConfig, loop_obj: &Py, queue: &Py, + cancelled: watch::Receiver, + credits: Arc, ) -> Result<(), String> { - let runtime = TokioRuntimeBuilder::new_multi_thread() + let runtime = TokioRuntimeBuilder::new_current_thread() .enable_all() .build() .map_err(|err| format!("failed to build tokio runtime for wreq: {err}"))?; - runtime.block_on(run_proxy_checks_async(proxies, config, loop_obj, queue)) + let (loop_obj, queue) = Python::attach(|py| { + ( + Arc::new(loop_obj.clone_ref(py)), + Arc::new(queue.clone_ref(py)), + ) + }); + runtime.block_on(async { + tokio::select! { + biased; + () = wait_for_cancellation(cancelled) => Ok(()), + result = run_proxy_checks_async(proxies, config, loop_obj, queue, credits) => result, + } + }) } /// Check proxies with bounded concurrency and emit results in completion order. async fn run_proxy_checks_async( proxies: Vec, config: CheckerConfig, - loop_obj: &Py, - queue: &Py, + loop_obj: Arc>, + queue: Arc>, + credits: Arc, ) -> Result<(), String> { let client = build_client(&config)?; let concurrency = config.concurrency.min(proxies.len().max(1)); @@ -58,9 +83,34 @@ async fn run_proxy_checks_async( .boxed(); while let Some(outcome) = outcomes.next().await { - emit_stream_result(loop_obj, queue, outcome) + let permit = credits.acquire().await.map_err(|err| err.to_string())?; + // Python attachment can wait for the GIL; keep it off Tokio's I/O thread. + let (loop_obj, queue) = (loop_obj.clone(), queue.clone()); + tokio::task::spawn_blocking(move || emit_stream_result(&loop_obj, &queue, outcome)) + .await + .map_err(|err| format!("result delivery task failed: {err}"))? .map_err(|err| format!("failed to emit proxy result: {err}"))?; + permit.forget(); } Ok(()) } + +/// Wait for explicit cancellation or loss of every cancellation sender. +async fn wait_for_cancellation(mut cancelled: watch::Receiver) { + while !*cancelled.borrow_and_update() { + if cancelled.changed().await.is_err() { + break; + } + } +} + +/// Keep a worker panic from leaving the Python consumer awaiting an absent result. +fn catch_worker_failure(work: impl FnOnce() -> Result<(), String>) -> Result<(), String> { + std::panic::catch_unwind(std::panic::AssertUnwindSafe(work)) + .unwrap_or_else(|_| Err("proxy worker panicked".into())) +} + +#[cfg(test)] +#[path = "worker_tests.rs"] +mod tests; diff --git a/src/worker_tests.rs b/src/worker_tests.rs new file mode 100644 index 0000000..eec2420 --- /dev/null +++ b/src/worker_tests.rs @@ -0,0 +1,40 @@ +use super::*; + +#[test] +fn panic_and_errors_are_reported() { + assert_eq!( + catch_worker_failure(|| panic!("test")), + Err("proxy worker panicked".into()) + ); + assert_eq!( + catch_worker_failure(|| Err("setup failed".into())), + Err("setup failed".into()) + ); + assert!(catch_worker_failure(|| Ok(())).is_ok()); +} + +#[tokio::test] +async fn cancellation_interrupts_a_backpressure_wait() { + let (cancel, cancelled) = watch::channel(false); + let credits = Semaphore::new(0); + let waiting = async { + tokio::select! { + () = wait_for_cancellation(cancelled) => true, + _ = credits.acquire() => false, + } + }; + cancel.send(true).unwrap(); + assert!(waiting.await); +} + +#[tokio::test] +async fn cancellation_wakes_on_signal_and_sender_loss() { + let (cancel, cancelled) = watch::channel(false); + let task = tokio::spawn(wait_for_cancellation(cancelled)); + tokio::task::yield_now().await; + cancel.send(true).unwrap(); + task.await.unwrap(); + let (cancel, cancelled) = watch::channel(false); + drop(cancel); + wait_for_cancellation(cancelled).await; +} diff --git a/tests/test_api.py b/tests/test_api.py index 1cb277e..ed45479 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -192,3 +192,77 @@ async def run(): self.assertTrue(all(item["ok"] for item in results)) self.assertLessEqual(state["maximum"], 3) asyncio.run(run()) + + +class LifecycleTests(unittest.TestCase): + def test_cancelling_iteration_closes_active_request(self): + async def run(): + started = asyncio.Event() + disconnected = asyncio.Event() + + async def handle(reader, writer): + try: + await reader.readuntil(b"\r\n\r\n") + await reader.readexactly(len(b"rsloop proxy checker")) + started.set() + self.assertEqual(await reader.read(), b"") + disconnected.set() + finally: + writer.close() + await writer.wait_closed() + + server = await asyncio.start_server(handle, "127.0.0.1", 0) + async with server: + port = server.sockets[0].getsockname()[1] + stream = await checker.check_proxies( + [f"http://127.0.0.1:{port}"], user_agent="test", + check_url="http://target.invalid/check", timeout_ms=10000 + ) + pending = asyncio.ensure_future(anext(stream)) + await asyncio.wait_for(started.wait(), 3) + pending.cancel() + with self.assertRaises(asyncio.CancelledError): + await pending + await asyncio.wait_for(disconnected.wait(), 3) + + asyncio.run(run()) + + def test_slow_consumer_bounds_requests_and_drop_cancels(self): + async def run(): + requests = 0 + second = asyncio.Event() + third = asyncio.Event() + + async def handle(reader, writer): + nonlocal requests + try: + await reader.readuntil(b"\r\n\r\n") + await reader.readexactly(len(b"rsloop proxy checker")) + requests += 1 + if requests == 2: + second.set() + if requests == 3: + third.set() + writer.write(b"HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n") + await writer.drain() + finally: + writer.close() + await writer.wait_closed() + + server = await asyncio.start_server(handle, "127.0.0.1", 0) + async with server: + port = server.sockets[0].getsockname()[1] + stream = await checker.check_proxies( + [f"http://127.0.0.1:{port}"] * 20, + user_agent="test", check_url="http://target.invalid/check", concurrency=1, + ) + await asyncio.wait_for(second.wait(), 3) + with self.assertRaises(asyncio.TimeoutError): + await asyncio.wait_for(third.wait(), 0.2) + self.assertTrue((await anext(stream))["ok"]) + await asyncio.wait_for(third.wait(), 3) + del stream + await asyncio.sleep(0.2) + self.assertLessEqual(requests, 3) + + asyncio.run(run())