diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b25ade7..aaf2b14 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -5,6 +5,7 @@ on: paths: - "src/**" - "tests/**" + - "scripts/**" - "Cargo.toml" - "Cargo.lock" - "pyproject.toml" @@ -17,6 +18,7 @@ on: paths: - "src/**" - "tests/**" + - "scripts/**" - "Cargo.toml" - "Cargo.lock" - "pyproject.toml" @@ -59,6 +61,10 @@ jobs: - run: cargo fmt --check - run: cargo check --locked --all-targets - run: cargo clippy --locked --all-targets -- -D warnings + - name: Check rustdoc + env: + RUSTDOCFLAGS: "-D warnings" + run: cargo doc --locked --no-deps --document-private-items tests: runs-on: ubuntu-24.04 @@ -92,3 +98,37 @@ jobs: python -m pip install dist/*.whl - name: Test public Python API run: python -m unittest discover -s tests -v + + coverage: + name: Coverage (>=90%) + runs-on: ubuntu-24.04 + timeout-minutes: 25 + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-python@v7 + with: + python-version: "3.14" + - name: Install native build dependencies + run: sudo apt-get update && sudo apt-get install -y cmake clang libclang-dev + - name: Install coverage tooling + run: | + rustup toolchain install stable --profile minimal --component llvm-tools-preview + rustup default stable + cargo install cargo-llvm-cov --version 0.9.1 --locked + python -m pip install "maturin>=1.7,<2" + - uses: actions/cache@v5 + with: + path: | + ~/.cargo/registry + ~/.cargo/git + target/coverage + key: coverage-${{ runner.os }}-${{ hashFiles('Cargo.lock', 'Cargo.toml') }} + - name: Measure Rust + Python API coverage and enforce 90% + run: bash scripts/coverage.sh + - name: Upload coverage report and badge + if: always() + uses: actions/upload-artifact@v7 + with: + name: rust-coverage + path: target/coverage/report/ + if-no-files-found: warn diff --git a/CHANGELOG.md b/CHANGELOG.md index 1ec7364..6a15cd0 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,10 +10,15 @@ Unreleased entry. ## [Unreleased] ### Added +- Internal Rust module documentation and a rustdoc warning check. +- Unit and local Python integration tests for stream messages, error propagation, body handling, and bounded concurrency. +- A combined Rust/Python coverage check requiring at least 90% production Rust line coverage, with a CI progress bar and downloadable reports/badge. - CI for relevant pull requests, pushes to `master`, and manual runs, with path filters and cancellation of superseded runs. - Rust checks, regression tests, pedantic Clippy, and Python API tests on Python 3.10 and 3.15. ### Changed +- 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. - Update the Rust `rsloop` dependency to 0.1.56 and `wreq` to the stable 0.16.1 series. - Update PyO3 to 0.29.3 for compatibility with `rsloop`. - Require Python 3.10 or newer and Rust 1.98 or newer for the updated dependencies. diff --git a/Cargo.toml b/Cargo.toml index ccfce01..125b001 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,7 +12,7 @@ crate-type = ["cdylib"] [dependencies] async-std = "1" futures = "0.3" -pyo3 = { version = "0.29.3", features = ["extension-module"] } +pyo3 = { version = "0.29.3" } rsloop = { version = "0.1.56" } tokio = { version = "1", features = ["full"] } wreq = { version = "0.16.1", features = ["socks", "stream"] } diff --git a/README.md b/README.md index 576dc41..2cce61e 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,7 @@ # Rust Proxy Checker Example +[![CI](https://github.com/RustedBytes/proxychecker/actions/workflows/ci.yml/badge.svg)](https://github.com/RustedBytes/proxychecker/actions/workflows/ci.yml) + This example is a standalone PyO3 extension built on top of `rsloop::rust_async` that checks many proxies concurrently and yields Python-friendly result objects as soon as each proxy finishes. @@ -104,6 +106,35 @@ async for result in stream: failed.append(proxy_result) ``` +## Rust module layout + +- `api`: Python arguments, validation, and stream creation. +- `config`: shared settings and defaults. +- `checker`: proxy parsing, HTTP client, and single-proxy outcomes. +- `worker`: Tokio worker lifecycle and bounded concurrency. +- `stream`: Python queue delivery, result dictionaries, and async iteration. +- `lib`: module registration. + +Build the internal Rust documentation with `cargo doc --no-deps --document-private-items`. + +## Coverage + +The CI check **Coverage (>=90%)** requires at least 90% line coverage across all +production Rust modules, including the Python binding and worker code. It combines +Rust unit tests with Python tests against an instrumented wheel; tests and external +dependencies are excluded from the denominator. It measures line coverage, not branch +coverage. The CI summary shows a progress bar; the `rust-coverage` artifact includes +an HTML report, JSON metrics, and a badge with the measured percentage. + +To reproduce locally in an activated Python virtual environment: + +```bash +rustup component add llvm-tools-preview +cargo install cargo-llvm-cov --version 0.9.1 --locked +python -m pip install "maturin>=1.7,<2" +bash scripts/coverage.sh +``` + ## Supported proxy strings The example passes the proxy string directly to `wreq`, so support follows the diff --git a/pyproject.toml b/pyproject.toml index 2a11233..4679c6e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,6 +11,7 @@ dependencies = [ ] [tool.maturin] +features = ["pyo3/extension-module"] module-name = "rsloop_rust_proxychecker" [dependency-groups] diff --git a/scripts/coverage.sh b/scripts/coverage.sh new file mode 100644 index 0000000..6ba9b93 --- /dev/null +++ b/scripts/coverage.sh @@ -0,0 +1,21 @@ +#!/usr/bin/env bash +# Measure all production Rust modules using both unit and Python API tests. +set -euo pipefail +cd "$(dirname "$0")/.." + +export CARGO_TARGET_DIR="${CARGO_TARGET_DIR:-$PWD/target/coverage}" +# Rust unit tests embed Python; some distributions need its shared-library path. +python_libdir="$(python -c 'import sysconfig; print(sysconfig.get_config_var("LIBDIR") or "")')" +export LD_LIBRARY_PATH="$python_libdir${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" +source <(cargo llvm-cov show-env --sh) +cargo llvm-cov clean --workspace +cargo test --locked +maturin build --locked --out "$CARGO_TARGET_DIR/wheels" +python -m pip install --force-reinstall "$CARGO_TARGET_DIR"/wheels/*.whl +python -m unittest discover -s tests -v + +mkdir -p "$CARGO_TARGET_DIR/report" +cargo llvm-cov report --json --output-path "$CARGO_TARGET_DIR/report/coverage.json" +cargo llvm-cov report --html --output-dir "$CARGO_TARGET_DIR/report/html" +python scripts/coverage_summary.py "$CARGO_TARGET_DIR/report/coverage.json" +cargo llvm-cov report --fail-under-lines 90 diff --git a/scripts/coverage_summary.py b/scripts/coverage_summary.py new file mode 100644 index 0000000..f8e4d76 --- /dev/null +++ b/scripts/coverage_summary.py @@ -0,0 +1,31 @@ +"""Render the measured Rust line coverage as a badge and CI progress bar.""" + +import json +import os +import sys +from pathlib import Path + +report_path = Path(sys.argv[1]) +report = json.loads(report_path.read_text()) +lines = report["data"][0]["totals"]["lines"] +percent = lines["percent"] +covered, count = lines["covered"], lines["count"] +color = "#4c1" if percent >= 90 else "#e05d44" +label = f"{percent:.2f}%" +svg = f''' + + +Rust coverage{label}''' +report_path.with_name("coverage.svg").write_text(svg) +filled = min(20, round(percent / 5)) +bar = "▰" * filled + "▱" * (20 - filled) +summary = ( + f"### Rust line coverage: {label}\n\n" + f"{bar} **{label}** — required: **90%**\n\n" + f"{covered}/{count} production Rust lines covered by unit and Python API tests.\n" + "All production modules are included; test code and dependencies are excluded.\n" +) +print(summary) +if summary_path := os.environ.get("GITHUB_STEP_SUMMARY"): + with Path(summary_path).open("a") as output: + output.write(summary) diff --git a/src/api.rs b/src/api.rs new file mode 100644 index 0000000..a3006fc --- /dev/null +++ b/src/api.rs @@ -0,0 +1,76 @@ +//! Python entry point and synchronous argument validation. + +use std::time::Duration; + +use pyo3::exceptions::PyValueError; +use pyo3::prelude::*; + +use crate::config::{CheckerConfig, DEFAULT_CHECK_URL, DEFAULT_CONCURRENCY, DEFAULT_TIMEOUT_MS}; +use crate::stream::PyProxyCheckStream; +use crate::worker::spawn_proxy_checks; + +/// Check proxy URLs and resolve to an asynchronous stream of result dictionaries. +/// +/// Arguments are validated synchronously. A nonempty user agent, positive timeout, +/// positive concurrency limit, and parseable target URL are required. Network and +/// body-read errors are returned as failed results rather than raised exceptions. +/// `return_response` controls body collection; concurrency controls in-flight work. +/// +/// # Errors +/// +/// Raises `ValueError` for invalid arguments, or propagates Python event-loop and +/// queue setup errors. Iterating the returned stream can raise `RuntimeError` when +/// the worker cannot initialize or deliver results. +#[pyfunction] +#[pyo3(signature=(proxies, *, user_agent, check_url=None, timeout_ms=DEFAULT_TIMEOUT_MS, concurrency=DEFAULT_CONCURRENCY, return_response=false))] +pub(crate) fn check_proxies( + py: Python<'_>, + proxies: Vec, + user_agent: String, + check_url: Option, + timeout_ms: u64, + concurrency: usize, + return_response: bool, +) -> PyResult> { + if user_agent.trim().is_empty() { + return Err(PyValueError::new_err("user_agent must not be empty")); + } + if timeout_ms == 0 { + return Err(PyValueError::new_err("timeout_ms must be greater than 0")); + } + if concurrency == 0 { + return Err(PyValueError::new_err("concurrency must be greater than 0")); + } + + let check_url = check_url.unwrap_or_else(|| DEFAULT_CHECK_URL.to_string()); + url::Url::parse(&check_url) + .map_err(|err| PyValueError::new_err(format!("invalid check_url: {err}")))?; + + let config = CheckerConfig { + check_url, + user_agent, + timeout: Duration::from_millis(timeout_ms), + concurrency, + return_response, + }; + let locals = rsloop::rust_async::get_current_locals(py)?; + let queue = py.import("asyncio")?.getattr("Queue")?.call0()?.unbind(); + let stream = Py::new( + py, + PyProxyCheckStream { + queue: queue.clone_ref(py), + }, + )? + .into_any(); + + spawn_proxy_checks( + proxies, + config.clone(), + locals.event_loop(py).unbind(), + queue, + ); + rsloop::rust_async::future_into_py(py, async move { + let _ = config; + Ok(stream) + }) +} diff --git a/src/checker.rs b/src/checker.rs new file mode 100644 index 0000000..06f1f6a --- /dev/null +++ b/src/checker.rs @@ -0,0 +1,149 @@ +//! HTTP client configuration, proxy validation, and single-proxy checks. + +use std::time::Instant; + +use futures::StreamExt; +use wreq::header; + +use crate::config::{CheckerConfig, REQUEST_BODY}; + +/// Result of one proxy check before conversion to a Python dictionary. +pub(crate) struct ProxyOutcome { + pub(crate) proxy: String, + pub(crate) elapsed_ms: u128, + pub(crate) status: Option, + pub(crate) ok: bool, + pub(crate) error: Option, + pub(crate) response_text: Option, +} + +/// Build a Tokio HTTP client with the configured timeout and disabled idle pooling. +pub(crate) fn build_client(config: &CheckerConfig) -> Result { + let mut builder = wreq::Client::builder() + .user_agent(config.user_agent.clone()) + .no_proxy() + .pool_max_idle_per_host(0) + .timeout(config.timeout) + .read_timeout(config.timeout) + .connect_timeout(config.timeout) + .pool_idle_timeout(Some(config.timeout)) + .tcp_keepalive(Some(config.timeout)) + .tcp_keepalive_interval(Some(config.timeout)) + .tcp_user_timeout(Some(config.timeout)); + + builder = builder.tcp_nodelay(true); + + builder + .build() + .map_err(|err| format!("failed to build wreq client: {err}")) +} + +/// Parse HTTP/HTTPS/SOCKS proxies, rejecting schemes that would bypass the proxy. +fn parse_proxy(proxy: &str) -> Result { + // Preserve support for host:port inputs while rejecting schemes that wreq's + // matcher silently ignores (which would otherwise send the request directly). + let uri = if proxy.contains("://") { + url::Url::parse(proxy) + } else { + url::Url::parse(&format!("http://{proxy}")) + } + .map_err(|err| format!("invalid proxy URL: {err}"))?; + if !matches!( + uri.scheme(), + "http" | "https" | "socks4" | "socks4a" | "socks5" | "socks5h" + ) { + return Err(format!("unsupported proxy scheme: {}", uri.scheme())); + } + wreq::Proxy::all(uri.as_str()).map_err(|err| err.to_string()) +} + +/// POST through one proxy, drain its body, and record transport or HTTP failures. +/// +/// Body-read failures override status-based success; text is collected only when requested. +pub(crate) async fn check_one_proxy( + client: wreq::Client, + config: CheckerConfig, + proxy: String, +) -> ProxyOutcome { + let started = Instant::now(); + let request_proxy = match parse_proxy(&proxy) { + Ok(request_proxy) => request_proxy, + Err(err) => { + return ProxyOutcome { + proxy, + elapsed_ms: started.elapsed().as_millis(), + status: None, + ok: false, + error: Some(err), + response_text: None, + }; + } + }; + let request = client + .post(config.check_url.clone()) + .proxy(request_proxy) + .header(header::CONTENT_TYPE, "text/plain; charset=utf-8") + .timeout(config.timeout) + .read_timeout(config.timeout) + .body(REQUEST_BODY); + + match request.send().await { + Ok(response) => { + let status = response.status(); + let status_code = status.as_u16(); + let body = if config.return_response { + response + .bytes() + .await + .map(|bytes| Some(String::from_utf8_lossy(&bytes).into_owned())) + } else { + async { + let mut chunks = response.bytes_stream().boxed(); + while let Some(chunk) = chunks.next().await { + drop(chunk?); + } + Ok(None) + } + .await + }; + match body { + Ok(response_text) if status.is_success() => ProxyOutcome { + proxy, + elapsed_ms: started.elapsed().as_millis(), + status: Some(status_code), + ok: true, + error: None, + response_text, + }, + Ok(response_text) => ProxyOutcome { + proxy, + elapsed_ms: started.elapsed().as_millis(), + status: Some(status_code), + ok: false, + error: Some(format!("target returned HTTP {status_code}")), + response_text, + }, + Err(err) => ProxyOutcome { + proxy, + elapsed_ms: started.elapsed().as_millis(), + status: Some(status_code), + ok: false, + error: Some(format!("response body read failed: {err}")), + response_text: None, + }, + } + } + Err(err) => ProxyOutcome { + proxy, + elapsed_ms: started.elapsed().as_millis(), + status: None, + ok: false, + error: Some(err.to_string()), + response_text: None, + }, + } +} + +#[cfg(test)] +#[path = "checker_tests.rs"] +mod tests; diff --git a/src/checker_tests.rs b/src/checker_tests.rs new file mode 100644 index 0000000..356ad71 --- /dev/null +++ b/src/checker_tests.rs @@ -0,0 +1,167 @@ +use super::*; +use std::time::Duration; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::net::TcpListener; + +#[test] +fn proxy_parser_rejects_unsupported_schemes() { + for proxy in [ + "invalid://proxy", + "ftp://127.0.0.1:8080", + "file:///tmp/proxy", + ] { + assert!(parse_proxy(proxy).is_err()); + } +} + +#[test] +fn proxy_parser_preserves_supported_schemes_and_bare_addresses() { + for proxy in [ + "http://127.0.0.1:8080", + "https://127.0.0.1:8080", + "socks4://127.0.0.1:1080", + "socks4a://127.0.0.1:1080", + "socks5://user:pass@127.0.0.1:1080", + "socks5h://127.0.0.1:1080", + "127.0.0.1:8080", + ] { + assert!(parse_proxy(proxy).is_ok(), "failed to parse {proxy}"); + } +} + +async fn check_response(status: u16, return_response: bool, truncated: bool) -> ProxyOutcome { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let proxy = format!("http://{}", listener.local_addr().unwrap()); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = Vec::new(); + let mut buffer = [0; 1024]; + loop { + let count = socket.read(&mut buffer).await.unwrap(); + assert_ne!(count, 0, "client closed before sending request body"); + request.extend_from_slice(&buffer[..count]); + if let Some(end) = request.windows(4).position(|part| part == b"\r\n\r\n") { + if request.len() >= end + 4 + REQUEST_BODY.len() { + break; + } + } + } + assert!(request.starts_with(b"POST http://target.invalid/check HTTP/1.1\r\n")); + if truncated { + socket + .write_all( + format!( + "HTTP/1.1 {status} Test\r\nContent-Length: 100\r\nConnection: close\r\n\r\nshort" + ) + .as_bytes(), + ) + .await + .unwrap(); + } else { + socket + .write_all( + format!( + "HTTP/1.1 {status} Test\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n" + ) + .as_bytes(), + ) + .await + .unwrap(); + socket.write_all(b"3\r\nhel\r\n").await.unwrap(); + socket.write_all(b"3\r\nlo\xff\r\n0\r\n\r\n").await.unwrap(); + } + socket.shutdown().await.unwrap(); + }); + let config = CheckerConfig { + check_url: "http://target.invalid/check".into(), + user_agent: "proxychecker regression test".into(), + timeout: Duration::from_secs(5), + concurrency: 1, + return_response, + }; + let outcome = check_one_proxy(build_client(&config).unwrap(), config, proxy.clone()).await; + server.await.unwrap(); + assert_eq!(outcome.proxy, proxy); + assert_eq!(outcome.status, Some(status)); + outcome +} + +#[tokio::test] +async fn discards_body_for_success_and_non_success_status() { + for status in [200, 503] { + let outcome = check_response(status, false, false).await; + assert_eq!(outcome.ok, status == 200); + assert_eq!(outcome.response_text, None); + assert_eq!( + outcome.error, + (status != 200).then(|| format!("target returned HTTP {status}")) + ); + } +} + +#[tokio::test] +async fn body_read_errors_fail_for_both_modes_and_statuses() { + for return_response in [false, true] { + for status in [200, 503] { + let outcome = check_response(status, return_response, true).await; + assert!(!outcome.ok); + assert_eq!(outcome.response_text, None); + assert!(outcome + .error + .unwrap() + .starts_with("response body read failed: ")); + } + } +} + +#[tokio::test] +async fn return_response_preserves_full_body_and_lossy_utf8() { + for status in [200, 503] { + let outcome = check_response(status, true, false).await; + assert_eq!(outcome.ok, status == 200); + assert_eq!(outcome.response_text.as_deref(), Some("hello\u{fffd}")); + assert_eq!( + outcome.error, + (status != 200).then(|| format!("target returned HTTP {status}")) + ); + } +} + +fn config_for(check_url: &str) -> CheckerConfig { + CheckerConfig { + check_url: check_url.into(), + user_agent: "test".into(), + timeout: Duration::from_secs(1), + concurrency: 2, + return_response: false, + } +} + +#[test] +fn malformed_proxy_and_invalid_user_agent_are_rejected() { + for proxy in ["http://", "http://[", "", "socks5://"] { + assert!( + parse_proxy(proxy).is_err(), + "accepted malformed proxy {proxy}" + ); + } + let mut config = config_for("http://target.invalid/check"); + config.user_agent = "invalid\nheader".into(); + assert!(build_client(&config).is_err()); +} + +#[tokio::test] +async fn invalid_proxy_and_transport_errors_have_no_http_status() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let unavailable = format!("http://{}", listener.local_addr().unwrap()); + drop(listener); + for proxy in ["invalid://proxy".to_string(), unavailable] { + let config = config_for("http://target.invalid/check"); + let outcome = check_one_proxy(build_client(&config).unwrap(), config, proxy.clone()).await; + assert!(!outcome.ok); + assert_eq!(outcome.proxy, proxy); + assert!(outcome.error.is_some()); + assert_eq!(outcome.status, None); + assert_eq!(outcome.response_text, None); + } +} diff --git a/src/config.rs b/src/config.rs new file mode 100644 index 0000000..95ac252 --- /dev/null +++ b/src/config.rs @@ -0,0 +1,22 @@ +//! Validated request settings and Python API defaults. + +use std::time::Duration; + +/// Default target for proxy checks. +pub(crate) const DEFAULT_CHECK_URL: &str = "https://httpbin.org/post"; +/// Default per-proxy timeout in milliseconds. +pub(crate) const DEFAULT_TIMEOUT_MS: u64 = 5_000; +/// Default limit for in-flight proxy checks. +pub(crate) const DEFAULT_CONCURRENCY: usize = 64; +/// Plain-text payload sent to the check target. +pub(crate) const REQUEST_BODY: &str = "rsloop proxy checker"; + +/// Immutable settings shared by checks in one stream. +#[derive(Clone)] +pub(crate) struct CheckerConfig { + pub(crate) check_url: String, + pub(crate) user_agent: String, + pub(crate) timeout: Duration, + pub(crate) concurrency: usize, + pub(crate) return_response: bool, +} diff --git a/src/lib.rs b/src/lib.rs index 872b98e..9c2a41d 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,507 +1,22 @@ -use std::time::{Duration, Instant}; +//! Concurrent HTTP/SOCKS proxy checking exposed as a Python asynchronous iterator. +//! +//! `api` validates Python arguments, `worker` owns Tokio execution, `checker` +//! performs individual requests, and `stream` transfers results to the Python loop. +//! Request failures are data (`ok=false`); worker failures become Python exceptions. +//! Bodies are drained without retention unless `return_response=true`. + +mod api; +mod checker; +mod config; +mod stream; +mod worker; -use futures::stream::{self, StreamExt}; -use pyo3::exceptions::{PyRuntimeError, PyStopAsyncIteration, PyValueError}; use pyo3::prelude::*; -use pyo3::types::PyDict; -use tokio::runtime::Builder as TokioRuntimeBuilder; -use wreq::header; - -const DEFAULT_CHECK_URL: &str = "https://httpbin.org/post"; -const DEFAULT_TIMEOUT_MS: u64 = 5_000; -const DEFAULT_CONCURRENCY: usize = 64; -const REQUEST_BODY: &str = "rsloop proxy checker"; - -#[derive(Clone)] -struct CheckerConfig { - check_url: String, - user_agent: String, - timeout: Duration, - concurrency: usize, - return_response: bool, -} - -struct ProxyOutcome { - proxy: String, - elapsed_ms: u128, - status: Option, - ok: bool, - error: Option, - response_text: Option, -} - -#[pyclass(module = "rsloop_rust_proxychecker")] -struct PyProxyCheckStream { - queue: Py, -} - -#[pyfunction] -#[pyo3(signature=(proxies, *, user_agent, check_url=None, timeout_ms=DEFAULT_TIMEOUT_MS, concurrency=DEFAULT_CONCURRENCY, return_response=false))] -fn check_proxies( - py: Python<'_>, - proxies: Vec, - user_agent: String, - check_url: Option, - timeout_ms: u64, - concurrency: usize, - return_response: bool, -) -> PyResult> { - if user_agent.trim().is_empty() { - return Err(PyValueError::new_err("user_agent must not be empty")); - } - if timeout_ms == 0 { - return Err(PyValueError::new_err("timeout_ms must be greater than 0")); - } - if concurrency == 0 { - return Err(PyValueError::new_err("concurrency must be greater than 0")); - } - - let check_url = check_url.unwrap_or_else(|| DEFAULT_CHECK_URL.to_string()); - url::Url::parse(&check_url) - .map_err(|err| PyValueError::new_err(format!("invalid check_url: {err}")))?; - - let config = CheckerConfig { - check_url, - user_agent, - timeout: Duration::from_millis(timeout_ms), - concurrency, - return_response, - }; - let locals = rsloop::rust_async::get_current_locals(py)?; - let queue = py.import("asyncio")?.getattr("Queue")?.call0()?.unbind(); - let stream = Py::new( - py, - PyProxyCheckStream { - queue: queue.clone_ref(py), - }, - )? - .into_any(); - - spawn_proxy_checks( - proxies, - config.clone(), - locals.event_loop(py).unbind(), - queue, - ); - rsloop::rust_async::future_into_py(py, async move { - let _ = config; - Ok(stream) - }) -} - -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); - }); -} - -fn run_proxy_checks_blocking( - proxies: Vec, - config: CheckerConfig, - loop_obj: &Py, - queue: &Py, -) -> Result<(), String> { - let runtime = TokioRuntimeBuilder::new_multi_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)) -} - -async fn run_proxy_checks_async( - proxies: Vec, - config: CheckerConfig, - loop_obj: &Py, - queue: &Py, -) -> Result<(), String> { - let client = build_client(&config)?; - let concurrency = config.concurrency.min(proxies.len().max(1)); - - let mut outcomes = stream::iter(proxies.into_iter().map(|proxy| { - let client = client.clone(); - let config = config.clone(); - async move { check_one_proxy(client, config, proxy).await } - })) - .buffer_unordered(concurrency) - .boxed(); - - while let Some(outcome) = outcomes.next().await { - emit_stream_result(loop_obj, queue, outcome) - .map_err(|err| format!("failed to emit proxy result: {err}"))?; - } - - Ok(()) -} - -fn build_client(config: &CheckerConfig) -> Result { - let mut builder = wreq::Client::builder() - .user_agent(config.user_agent.clone()) - .no_proxy() - .pool_max_idle_per_host(0) - .timeout(config.timeout) - .read_timeout(config.timeout) - .connect_timeout(config.timeout) - .pool_idle_timeout(Some(config.timeout)) - .tcp_keepalive(Some(config.timeout)) - .tcp_keepalive_interval(Some(config.timeout)) - .tcp_user_timeout(Some(config.timeout)); - - builder = builder.tcp_nodelay(true); - - builder - .build() - .map_err(|err| format!("failed to build wreq client: {err}")) -} - -fn parse_proxy(proxy: &str) -> Result { - // Preserve support for host:port inputs while rejecting schemes that wreq's - // matcher silently ignores (which would otherwise send the request directly). - let uri = if proxy.contains("://") { - url::Url::parse(proxy) - } else { - url::Url::parse(&format!("http://{proxy}")) - } - .map_err(|err| format!("invalid proxy URL: {err}"))?; - if !matches!( - uri.scheme(), - "http" | "https" | "socks4" | "socks4a" | "socks5" | "socks5h" - ) { - return Err(format!("unsupported proxy scheme: {}", uri.scheme())); - } - wreq::Proxy::all(uri.as_str()).map_err(|err| err.to_string()) -} - -async fn check_one_proxy( - client: wreq::Client, - config: CheckerConfig, - proxy: String, -) -> ProxyOutcome { - let started = Instant::now(); - let request_proxy = match parse_proxy(&proxy) { - Ok(request_proxy) => request_proxy, - Err(err) => { - return ProxyOutcome { - proxy, - elapsed_ms: started.elapsed().as_millis(), - status: None, - ok: false, - error: Some(err), - response_text: None, - }; - } - }; - let request = client - .post(config.check_url.clone()) - .proxy(request_proxy) - .header(header::CONTENT_TYPE, "text/plain; charset=utf-8") - .timeout(config.timeout) - .read_timeout(config.timeout) - .body(REQUEST_BODY); - - match request.send().await { - Ok(response) => { - let status = response.status(); - let status_code = status.as_u16(); - let body = if config.return_response { - response - .bytes() - .await - .map(|bytes| Some(String::from_utf8_lossy(&bytes).into_owned())) - } else { - async { - let mut chunks = response.bytes_stream().boxed(); - while let Some(chunk) = chunks.next().await { - drop(chunk?); - } - Ok(None) - } - .await - }; - match body { - Ok(response_text) if status.is_success() => ProxyOutcome { - proxy, - elapsed_ms: started.elapsed().as_millis(), - status: Some(status_code), - ok: true, - error: None, - response_text, - }, - Ok(response_text) => ProxyOutcome { - proxy, - elapsed_ms: started.elapsed().as_millis(), - status: Some(status_code), - ok: false, - error: Some(format!("target returned HTTP {status_code}")), - response_text, - }, - Err(err) => ProxyOutcome { - proxy, - elapsed_ms: started.elapsed().as_millis(), - status: Some(status_code), - ok: false, - error: Some(format!("response body read failed: {err}")), - response_text: None, - }, - } - } - Err(err) => ProxyOutcome { - proxy, - elapsed_ms: started.elapsed().as_millis(), - status: None, - ok: false, - error: Some(err.to_string()), - response_text: None, - }, - } -} - -fn build_outcome_dict(py: Python<'_>, outcome: ProxyOutcome) -> PyResult> { - let item = PyDict::new(py); - item.set_item("proxy", outcome.proxy)?; - item.set_item("ok", outcome.ok)?; - item.set_item( - "elapsed_ms", - u64::try_from(outcome.elapsed_ms).unwrap_or(u64::MAX), - )?; - if let Some(status) = outcome.status { - item.set_item("status", status)?; - } - if let Some(error) = outcome.error { - item.set_item("error", error)?; - } - if let Some(response_text) = outcome.response_text { - item.set_item("response_text", response_text)?; - } - Ok(item.unbind().into_any()) -} - -fn emit_stream_result( - loop_obj: &Py, - queue: &Py, - outcome: ProxyOutcome, -) -> PyResult<()> { - Python::attach(|py| { - let value = build_outcome_dict(py, outcome)?; - let message = PyDict::new(py); - message.set_item("kind", "item")?; - message.set_item("value", value)?; - schedule_queue_put(py, loop_obj, queue, message.unbind().into_any()) - }) -} - -fn emit_stream_error(loop_obj: &Py, queue: &Py, error: String) -> PyResult<()> { - Python::attach(|py| { - let message = PyDict::new(py); - message.set_item("kind", "error")?; - message.set_item("error", error)?; - schedule_queue_put(py, loop_obj, queue, message.unbind().into_any()) - }) -} - -fn emit_stream_end(loop_obj: &Py, queue: &Py) -> PyResult<()> { - Python::attach(|py| { - let message = PyDict::new(py); - message.set_item("kind", "end")?; - schedule_queue_put(py, loop_obj, queue, message.unbind().into_any()) - }) -} - -fn schedule_queue_put( - py: Python<'_>, - loop_obj: &Py, - queue: &Py, - message: Py, -) -> PyResult<()> { - let queue_bound = queue.bind(py); - let put_nowait = queue_bound.getattr("put_nowait")?; - loop_obj - .bind(py) - .call_method1("call_soon_threadsafe", (put_nowait, message))?; - Ok(()) -} - -#[pymethods] -impl PyProxyCheckStream { - fn __aiter__(slf: Py) -> Py { - slf - } - - // 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 locals = rsloop::rust_async::get_current_locals(py)?; - - rsloop::rust_async::future_into_py_with_locals(py, locals.clone(), async move { - let queued = Python::attach(|py| { - let awaitable = queue.bind(py).call_method0("get")?; - rsloop::rust_async::into_future_with_locals(&locals, awaitable) - })? - .await?; - - Python::attach(|py| { - let message = queued.bind(py).cast::()?; - let kind = message - .get_item("kind")? - .ok_or_else(|| PyRuntimeError::new_err("stream message missing kind"))? - .extract::()?; - - match kind.as_str() { - "item" => { - let value = message - .get_item("value")? - .ok_or_else(|| PyRuntimeError::new_err("stream item missing value"))?; - Ok(value.unbind()) - } - "error" => { - let error = message - .get_item("error")? - .ok_or_else(|| PyRuntimeError::new_err("stream error missing payload"))? - .extract::()?; - Err(PyRuntimeError::new_err(error)) - } - "end" => Err(PyStopAsyncIteration::new_err(())), - other => Err(PyRuntimeError::new_err(format!( - "unexpected stream message kind: {other}" - ))), - } - }) - }) - } -} +/// Register the existing Python function and asynchronous stream class. #[pymodule(gil_used = false)] fn rsloop_rust_proxychecker(m: &Bound<'_, PyModule>) -> PyResult<()> { - m.add_class::()?; - m.add_function(wrap_pyfunction!(check_proxies, m)?)?; + m.add_class::()?; + m.add_function(wrap_pyfunction!(api::check_proxies, m)?)?; Ok(()) } - -#[cfg(test)] -mod tests { - use super::*; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - use tokio::net::TcpListener; - - #[test] - fn proxy_parser_rejects_unsupported_schemes() { - for proxy in [ - "invalid://proxy", - "ftp://127.0.0.1:8080", - "file:///tmp/proxy", - ] { - assert!(parse_proxy(proxy).is_err()); - } - } - - #[test] - fn proxy_parser_preserves_supported_schemes_and_bare_addresses() { - for proxy in [ - "http://127.0.0.1:8080", - "https://127.0.0.1:8080", - "socks4://127.0.0.1:1080", - "socks4a://127.0.0.1:1080", - "socks5://user:pass@127.0.0.1:1080", - "socks5h://127.0.0.1:1080", - "127.0.0.1:8080", - ] { - assert!(parse_proxy(proxy).is_ok(), "failed to parse {proxy}"); - } - } - - async fn check_response(status: u16, return_response: bool, truncated: bool) -> ProxyOutcome { - let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); - let proxy = format!("http://{}", listener.local_addr().unwrap()); - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.unwrap(); - let mut request = Vec::new(); - let mut buffer = [0; 1024]; - loop { - let count = socket.read(&mut buffer).await.unwrap(); - assert_ne!(count, 0, "client closed before sending request body"); - request.extend_from_slice(&buffer[..count]); - if let Some(end) = request.windows(4).position(|part| part == b"\r\n\r\n") { - if request.len() >= end + 4 + REQUEST_BODY.len() { - break; - } - } - } - assert!(request.starts_with(b"POST http://target.invalid/check HTTP/1.1\r\n")); - if truncated { - socket.write_all(format!( - "HTTP/1.1 {status} Test\r\nContent-Length: 100\r\nConnection: close\r\n\r\nshort" - ).as_bytes()).await.unwrap(); - } else { - socket.write_all(format!( - "HTTP/1.1 {status} Test\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n" - ).as_bytes()).await.unwrap(); - socket.write_all(b"3\r\nhel\r\n").await.unwrap(); - socket.write_all(b"3\r\nlo\xff\r\n0\r\n\r\n").await.unwrap(); - } - socket.shutdown().await.unwrap(); - }); - let config = CheckerConfig { - check_url: "http://target.invalid/check".into(), - user_agent: "proxychecker regression test".into(), - timeout: Duration::from_secs(5), - concurrency: 1, - return_response, - }; - let outcome = check_one_proxy(build_client(&config).unwrap(), config, proxy.clone()).await; - server.await.unwrap(); - assert_eq!(outcome.proxy, proxy); - assert_eq!(outcome.status, Some(status)); - outcome - } - - #[tokio::test] - async fn discards_body_for_success_and_non_success_status() { - for status in [200, 503] { - let outcome = check_response(status, false, false).await; - assert_eq!(outcome.ok, status == 200); - assert_eq!(outcome.response_text, None); - assert_eq!( - outcome.error, - (status != 200).then(|| format!("target returned HTTP {status}")) - ); - } - } - - #[tokio::test] - async fn body_read_errors_fail_for_both_modes_and_statuses() { - for return_response in [false, true] { - for status in [200, 503] { - let outcome = check_response(status, return_response, true).await; - assert!(!outcome.ok); - assert_eq!(outcome.response_text, None); - assert!(outcome - .error - .unwrap() - .starts_with("response body read failed: ")); - } - } - } - - #[tokio::test] - async fn return_response_preserves_full_body_and_lossy_utf8() { - for status in [200, 503] { - let outcome = check_response(status, true, false).await; - assert_eq!(outcome.ok, status == 200); - assert_eq!(outcome.response_text.as_deref(), Some("hello\u{fffd}")); - assert_eq!( - outcome.error, - (status != 200).then(|| format!("target returned HTTP {status}")) - ); - } - } -} diff --git a/src/stream.rs b/src/stream.rs new file mode 100644 index 0000000..964bfb5 --- /dev/null +++ b/src/stream.rs @@ -0,0 +1,146 @@ +//! Thread-safe Python queue delivery and the asynchronous iterator protocol. + +use pyo3::exceptions::{PyRuntimeError, PyStopAsyncIteration}; +use pyo3::prelude::*; +use pyo3::types::PyDict; + +use crate::checker::ProxyOutcome; + +/// Asynchronous iterator backed by a Python queue populated from a Rust worker. +#[pyclass(module = "rsloop_rust_proxychecker")] +pub(crate) struct PyProxyCheckStream { + pub(crate) queue: Py, +} + +/// Convert an outcome to the public dictionary, omitting absent optional fields. +fn build_outcome_dict(py: Python<'_>, outcome: ProxyOutcome) -> PyResult> { + let item = PyDict::new(py); + item.set_item("proxy", outcome.proxy)?; + item.set_item("ok", outcome.ok)?; + item.set_item( + "elapsed_ms", + u64::try_from(outcome.elapsed_ms).unwrap_or(u64::MAX), + )?; + if let Some(status) = outcome.status { + item.set_item("status", status)?; + } + if let Some(error) = outcome.error { + item.set_item("error", error)?; + } + if let Some(response_text) = outcome.response_text { + item.set_item("response_text", response_text)?; + } + Ok(item.unbind().into_any()) +} + +/// Schedule an item message on the Python event loop. +pub(crate) fn emit_stream_result( + loop_obj: &Py, + queue: &Py, + outcome: ProxyOutcome, +) -> PyResult<()> { + Python::attach(|py| { + let value = build_outcome_dict(py, outcome)?; + let message = PyDict::new(py); + message.set_item("kind", "item")?; + message.set_item("value", value)?; + schedule_queue_put(py, loop_obj, queue, message.unbind().into_any()) + }) +} + +/// Schedule a worker failure message on the Python event loop. +pub(crate) fn emit_stream_error( + loop_obj: &Py, + queue: &Py, + error: String, +) -> PyResult<()> { + Python::attach(|py| { + let message = PyDict::new(py); + message.set_item("kind", "error")?; + message.set_item("error", error)?; + schedule_queue_put(py, loop_obj, queue, message.unbind().into_any()) + }) +} + +/// Schedule the terminal message after all work is finished. +pub(crate) fn emit_stream_end(loop_obj: &Py, queue: &Py) -> PyResult<()> { + Python::attach(|py| { + let message = PyDict::new(py); + message.set_item("kind", "end")?; + schedule_queue_put(py, loop_obj, queue, message.unbind().into_any()) + }) +} + +/// Marshal a queue insertion through the event loop thread-safe callback API. +fn schedule_queue_put( + py: Python<'_>, + loop_obj: &Py, + queue: &Py, + message: Py, +) -> PyResult<()> { + let queue_bound = queue.bind(py); + let put_nowait = queue_bound.getattr("put_nowait")?; + loop_obj + .bind(py) + .call_method1("call_soon_threadsafe", (put_nowait, message))?; + Ok(()) +} + +#[pymethods] +impl PyProxyCheckStream { + /// Return this stream as its asynchronous iterator. + fn __aiter__(slf: Py) -> Py { + slf + } + + /// Await the next queued result or propagate a terminal/error message. + // 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 locals = rsloop::rust_async::get_current_locals(py)?; + + rsloop::rust_async::future_into_py_with_locals(py, locals.clone(), async move { + let queued = Python::attach(|py| { + let awaitable = queue.bind(py).call_method0("get")?; + rsloop::rust_async::into_future_with_locals(&locals, awaitable) + })? + .await?; + + Python::attach(|py| decode_message(queued.bind(py))) + }) + } +} + +/// Decode the internal queue protocol and reject malformed messages. +fn decode_message(queued: &Bound<'_, PyAny>) -> PyResult> { + let message = queued.cast::()?; + let kind = message + .get_item("kind")? + .ok_or_else(|| PyRuntimeError::new_err("stream message missing kind"))? + .extract::()?; + + match kind.as_str() { + "item" => { + let value = message + .get_item("value")? + .ok_or_else(|| PyRuntimeError::new_err("stream item missing value"))?; + Ok(value.unbind()) + } + "error" => { + let error = message + .get_item("error")? + .ok_or_else(|| PyRuntimeError::new_err("stream error missing payload"))? + .extract::()?; + Err(PyRuntimeError::new_err(error)) + } + "end" => Err(PyStopAsyncIteration::new_err(())), + other => Err(PyRuntimeError::new_err(format!( + "unexpected stream message kind: {other}" + ))), + } +} + +#[cfg(test)] +#[path = "stream_tests.rs"] +mod tests; diff --git a/src/stream_tests.rs b/src/stream_tests.rs new file mode 100644 index 0000000..17d7962 --- /dev/null +++ b/src/stream_tests.rs @@ -0,0 +1,155 @@ +use super::*; +use pyo3::ffi::c_str; + +fn with_python(test: impl for<'py> FnOnce(Python<'py>)) { + Python::initialize(); + Python::attach(test); +} + +#[test] +fn outcome_dict_preserves_optional_fields_and_saturates_elapsed() { + with_python(|py| { + let full = ProxyOutcome { + proxy: "http://proxy".into(), + elapsed_ms: u128::MAX, + status: Some(503), + ok: false, + error: Some("HTTP failure".into()), + response_text: Some("body".into()), + }; + let dict = build_outcome_dict(py, full).unwrap(); + let dict = dict.bind(py).cast::().unwrap(); + assert_eq!(dict.len(), 6); + assert_eq!( + dict.get_item("elapsed_ms") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + u64::MAX + ); + assert_eq!( + dict.get_item("status") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + 503 + ); + assert_eq!( + dict.get_item("error") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + "HTTP failure" + ); + assert_eq!( + dict.get_item("response_text") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + "body" + ); + let empty = ProxyOutcome { + proxy: "invalid://proxy".into(), + elapsed_ms: 7, + status: None, + ok: false, + error: None, + response_text: None, + }; + let dict = build_outcome_dict(py, empty).unwrap(); + let dict = dict.bind(py).cast::().unwrap(); + assert_eq!(dict.len(), 3); + assert_eq!( + dict.get_item("elapsed_ms") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + 7 + ); + }); +} + +#[test] +fn queue_messages_deliver_items_errors_and_end() { + with_python(|py| { + let queue = py + .eval(c_str!("__import__('asyncio').Queue()"), None, None) + .unwrap() + .unbind(); + let loop_obj = py + .eval( + c_str!( + "type('Loop', (), {'call_soon_threadsafe': lambda self, cb, msg: cb(msg)})()" + ), + None, + None, + ) + .unwrap() + .unbind(); + let outcome = ProxyOutcome { + proxy: "proxy".into(), + elapsed_ms: 1, + status: Some(200), + ok: true, + error: None, + response_text: None, + }; + emit_stream_result(&loop_obj, &queue, outcome).unwrap(); + let message = queue.bind(py).call_method0("get_nowait").unwrap(); + let item = decode_message(&message).unwrap(); + assert!(item + .bind(py) + .get_item("ok") + .unwrap() + .extract::() + .unwrap()); + emit_stream_error(&loop_obj, &queue, "worker failed".into()).unwrap(); + let message = queue.bind(py).call_method0("get_nowait").unwrap(); + let err = decode_message(&message).unwrap_err(); + assert!(err.is_instance_of::(py)); + assert_eq!(err.value(py).to_string(), "worker failed"); + emit_stream_end(&loop_obj, &queue).unwrap(); + let message = queue.bind(py).call_method0("get_nowait").unwrap(); + assert!(decode_message(&message) + .unwrap_err() + .is_instance_of::(py)); + }); +} + +#[test] +fn malformed_messages_raise_descriptive_errors() { + with_python(|py| { + for (kind, expected) in [ + (None, "stream message missing kind"), + (Some("item"), "stream item missing value"), + (Some("error"), "stream error missing payload"), + (Some("other"), "unexpected stream message kind: other"), + ] { + let message = PyDict::new(py); + if let Some(kind) = kind { + message.set_item("kind", kind).unwrap(); + } + let err = decode_message(message.as_any()).unwrap_err(); + assert!(err.is_instance_of::(py)); + assert_eq!(err.value(py).to_string(), expected); + } + assert!(decode_message(py.None().bind(py)).is_err()); + }); +} + +#[test] +fn queue_delivery_propagates_event_loop_errors() { + with_python(|py| { + let closed_loop = py.None(); + let queue = py + .eval(c_str!("__import__('asyncio').Queue()"), None, None) + .unwrap() + .unbind(); + assert!(emit_stream_end(&closed_loop, &queue).is_err()); + }); +} diff --git a/src/worker.rs b/src/worker.rs new file mode 100644 index 0000000..a472875 --- /dev/null +++ b/src/worker.rs @@ -0,0 +1,66 @@ +//! Worker-thread lifecycle and bounded concurrent request scheduling. + +use futures::stream::{self, StreamExt}; +use pyo3::prelude::*; +use tokio::runtime::Builder as TokioRuntimeBuilder; + +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. +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); + }); +} + +/// Create a dedicated Tokio runtime for HTTP work on the worker thread. +fn run_proxy_checks_blocking( + proxies: Vec, + config: CheckerConfig, + loop_obj: &Py, + queue: &Py, +) -> Result<(), String> { + let runtime = TokioRuntimeBuilder::new_multi_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)) +} + +/// 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, +) -> Result<(), String> { + let client = build_client(&config)?; + let concurrency = config.concurrency.min(proxies.len().max(1)); + + let mut outcomes = stream::iter(proxies.into_iter().map(|proxy| { + let client = client.clone(); + let config = config.clone(); + async move { check_one_proxy(client, config, proxy).await } + })) + .buffer_unordered(concurrency) + .boxed(); + + while let Some(outcome) = outcomes.next().await { + emit_stream_result(loop_obj, queue, outcome) + .map_err(|err| format!("failed to emit proxy result: {err}"))?; + } + + Ok(()) +} diff --git a/tests/test_api.py b/tests/test_api.py index 4b398ed..eca32f2 100644 --- a/tests/test_api.py +++ b/tests/test_api.py @@ -1,4 +1,5 @@ import asyncio +from contextlib import contextmanager import threading from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer import unittest @@ -65,3 +66,123 @@ def test_invalid_arguments(self): ]: with self.subTest(kwargs=kwargs), self.assertRaises(ValueError): checker.check_proxies([], **kwargs) + + +@contextmanager +def proxy_server(status=200, truncated=False, gate=None, state=None): + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + body = self.rfile.read(int(self.headers.get("Content-Length", "0"))) + if body != b"rsloop proxy checker": + self.send_error(400, "unexpected request body") + return + if state is not None: + with state["lock"]: + state["active"] += 1 + state["maximum"] = max(state["maximum"], state["active"]) + if state["active"] == 3: + state["started"].set() + try: + if gate is not None and not gate.wait(5): + return + self.send_response(status) + self.send_header("Content-Length", "100" if truncated else "6") + self.send_header("Connection", "close") + self.end_headers() + self.wfile.write(b"hello\xff") + self.wfile.flush() + self.close_connection = True + finally: + if state is not None: + with state["lock"]: + state["active"] -= 1 + + def log_message(self, *args): + pass + + with ThreadingHTTPServer(("127.0.0.1", 0), Handler) as server: + thread = threading.Thread(target=server.serve_forever) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_port}" + finally: + if gate is not None: + gate.set() + server.shutdown() + thread.join() + + +class IntegrationTests(unittest.TestCase): + def test_statuses_and_body_modes(self): + for status in [200, 503]: + with proxy_server(status=status) as proxy: + for return_response in [False, True]: + async def run(): + stream = await checker.check_proxies( + [proxy], user_agent="test", check_url="http://target.invalid/check", + return_response=return_response, + ) + self.assertIs(stream.__aiter__(), stream) + results = [item async for item in stream] + self.assertEqual(len(results), 1) + result = results[0] + self.assertEqual(result["proxy"], proxy) + self.assertEqual(result["status"], status) + self.assertEqual(result["ok"], status == 200) + self.assertIsInstance(result["elapsed_ms"], int) + if return_response: + self.assertEqual(result["response_text"], "hello\ufffd") + else: + self.assertNotIn("response_text", result) + if status == 200: + self.assertNotIn("error", result) + else: + self.assertEqual(result["error"], "target returned HTTP 503") + with self.subTest(status=status, return_response=return_response): + asyncio.run(run()) + + def test_truncated_body_is_failure_in_both_modes(self): + with proxy_server(truncated=True) as proxy: + for return_response in [False, True]: + async def run(): + stream = await checker.check_proxies( + [proxy], user_agent="test", check_url="http://target.invalid/check", + return_response=return_response, + ) + result, = [item async for item in stream] + self.assertFalse(result["ok"]) + self.assertEqual(result["status"], 200) + self.assertIn("response body read failed", result["error"]) + self.assertNotIn("response_text", result) + asyncio.run(run()) + + def test_worker_setup_error_raises_runtime_error(self): + async def run(): + stream = await checker.check_proxies([], user_agent="invalid\nheader") + with self.assertRaisesRegex(RuntimeError, "failed to build wreq client"): + [item async for item in stream] + asyncio.run(run()) + + def test_requires_running_event_loop(self): + with self.assertRaises(RuntimeError): + checker.check_proxies([], user_agent="test") + + def test_concurrency_limit_and_complete_delivery(self): + gate = threading.Event() + state = {"lock": threading.Lock(), "started": threading.Event(), "active": 0, "maximum": 0} + with proxy_server(gate=gate, state=state) as proxy: + async def run(): + stream = await checker.check_proxies( + [proxy] * 6, user_agent="test", check_url="http://target.invalid/check", concurrency=3, + ) + try: + self.assertTrue(await asyncio.to_thread(state["started"].wait, 5)) + with state["lock"]: + self.assertEqual(state["active"], 3) + finally: + gate.set() + results = [item async for item in stream] + self.assertEqual(len(results), 6) + self.assertTrue(all(item["ok"] for item in results)) + self.assertLessEqual(state["maximum"], 3) + asyncio.run(run())