From 3dd41fdc95e81d24325cb1dd256fa45721c7c63d Mon Sep 17 00:00:00 2001 From: zyc9012 Date: Fri, 28 Aug 2026 14:25:00 +0800 Subject: [PATCH] Support batch request --- Cargo.lock | 2 +- README.md | 28 ++++++ ext/wreq_rb/Cargo.toml | 2 +- ext/wreq_rb/src/client.rs | 176 +++++++++++++++++++++++++++++++++++- lib/wreq-rb/version.rb | 2 +- test/batch_test.rb | 182 ++++++++++++++++++++++++++++++++++++++ 6 files changed, 388 insertions(+), 4 deletions(-) create mode 100644 test/batch_test.rb diff --git a/Cargo.lock b/Cargo.lock index 9395d97..ab1d56a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4056,7 +4056,7 @@ dependencies = [ [[package]] name = "wreq_rb" -version = "0.5.1" +version = "0.6.0" dependencies = [ "bytes", "http", diff --git a/README.md b/README.md index 409e565..be0d1ce 100644 --- a/README.md +++ b/README.md @@ -114,6 +114,34 @@ All methods are available on both `Wreq` (module-level) and `Wreq::Client` (inst | `head(url, **opts)` | HEAD request | | `options(url, **opts)` | OPTIONS request | +### Batch Requests + +`Wreq::Client#request_batch` runs many requests concurrently inside a **single** +GVL release, multiplexed over the client's connection pool. This avoids one Ruby +thread per request. + +```ruby +client = Wreq::Client.new(redirect: false, pool_max_idle_per_host: 32) + +responses = client.request_batch(urls, concurrency: 128) +``` + +Each element of the array is either a URL string or a hash: + +```ruby +client.request_batch([ + "https://example.com/a", # GET + { method: :put, url: "https://example.com/b", body: "hi" } +], concurrency: 32) +``` + +- Results are returned **in input order**. +- Each element is a `Wreq::Response` **or** a `Wreq::Error` — a single failed + request never discards the rest of the batch. Errors are returned, not raised. +- `concurrency` caps the number of in-flight requests and defaults to `16`. +- `cancel` and Ruby thread interrupts abort the whole batch, raising + `Wreq::Error` with `"request interrupted"`. + ### Cancelling Requests Call `cancel` on a client to interrupt all in-flight requests immediately: diff --git a/ext/wreq_rb/Cargo.toml b/ext/wreq_rb/Cargo.toml index 8d49fa7..6e9b203 100644 --- a/ext/wreq_rb/Cargo.toml +++ b/ext/wreq_rb/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "wreq_rb" -version = "0.5.1" +version = "0.6.0" edition = "2021" publish = false diff --git a/ext/wreq_rb/src/client.rs b/ext/wreq_rb/src/client.rs index 80895e8..80bcbc0 100644 --- a/ext/wreq_rb/src/client.rs +++ b/ext/wreq_rb/src/client.rs @@ -2,6 +2,7 @@ use std::ffi::c_void; use std::panic::{self, AssertUnwindSafe}; use std::ptr; use std::any::Any; +use std::sync::Arc; use std::time::Duration; use magnus::{ @@ -9,13 +10,15 @@ use magnus::{ try_convert::TryConvert, Value, }; use tokio::runtime::Runtime; +use tokio::sync::Semaphore; +use tokio::task::JoinSet; use tokio_util::sync::CancellationToken; use std::net::IpAddr; use wreq::header::{HeaderMap, HeaderName, HeaderValue, OrigHeaderMap}; use wreq::tls::TlsVersion; use wreq_util::{Emulation as BrowserEmulation, Platform as EmulationPlatform, Profile as BrowserProfile}; -use crate::error::{generic_error, to_magnus_error}; +use crate::error::{generic_error, to_magnus_error, wreq_error}; use crate::response::Response; // -------------------------------------------------------------------------- @@ -137,6 +140,62 @@ async fn execute_request(req: wreq::RequestBuilder) -> Result), + Interrupted, +} + +/// Run every request concurrently with at most `concurrency` in flight, +/// returning results in input order. +async fn execute_batch(reqs: Vec, concurrency: usize) -> Vec { + let permits = Arc::new(Semaphore::new(concurrency)); + let mut set: JoinSet<(usize, BatchItem)> = JoinSet::new(); + + for (idx, req) in reqs.into_iter().enumerate() { + let permits = Arc::clone(&permits); + set.spawn(async move { + let _permit = match permits.acquire_owned().await { + Ok(p) => p, + Err(_) => return (idx, BatchItem::Err("batch semaphore closed".to_owned())), + }; + let item = match execute_request(req).await { + Ok(data) => BatchItem::Ok(data), + Err(e) => BatchItem::Err(e.to_string()), + }; + (idx, item) + }); + } + + let mut slots: Vec> = Vec::new(); + slots.resize_with(set.len(), || None); + while let Some(joined) = set.join_next().await { + // A JoinError means the task panicked or was aborted; its slot is left + // empty and filled with a generic error below. + if let Ok((idx, item)) = joined { + slots[idx] = Some(item); + } + } + + slots + .into_iter() + .map(|slot| slot.unwrap_or_else(|| BatchItem::Err("request task failed".to_owned()))) + .collect() +} + // -------------------------------------------------------------------------- // Emulation helpers // -------------------------------------------------------------------------- @@ -439,6 +498,120 @@ impl Client { }; Ok(Response::new(data.status, data.headers, data.body, data.url, data.version, data.content_length, data.transfer_size)) } + + /// Wreq::Client#request_batch(specs) or #request_batch(specs, options) + fn request_batch(&self, args: &[Value]) -> Result { + if args.is_empty() { + return Err(generic_error("an array of requests is required")); + } + let specs = RArray::try_convert(args[0])?; + + let opts: Option = if args.len() > 1 { + Some(RHash::try_convert(args[1])?) + } else { + None + }; + + let concurrency = match opts.as_ref() { + Some(o) => hash_get_usize(o, "concurrency")?.unwrap_or(DEFAULT_BATCH_CONCURRENCY), + None => DEFAULT_BATCH_CONCURRENCY, + }; + if concurrency == 0 { + return Err(generic_error("concurrency must be >= 1")); + } + + // All Ruby -> Rust conversion happens here, while we still hold the GVL. + let mut reqs: Vec = Vec::with_capacity(specs.len()); + for spec in specs.into_iter() { + reqs.push(self.build_request(spec, opts.as_ref())?); + } + + let ruby = unsafe { Ruby::get_unchecked() }; + if reqs.is_empty() { + return Ok(ruby.ary_new()); + } + + let client_token = self.cancel_token.lock().unwrap_or_else(|e| e.into_inner()).clone(); + + // One GVL release covering the whole batch. + let outcome: BatchOutcome = unsafe { + without_gvl(|thread_token| { + runtime().block_on(async { + tokio::select! { + biased; + _ = thread_token.cancelled() => BatchOutcome::Interrupted, + _ = client_token.cancelled() => BatchOutcome::Interrupted, + items = execute_batch(reqs, concurrency) => BatchOutcome::Done(items), + } + }) + }) + }; + + let items = match outcome { + BatchOutcome::Done(items) => items, + BatchOutcome::Interrupted => return Err(generic_error("request interrupted")), + }; + + // Back under the GVL: Ruby objects may be created again. + let results = ruby.ary_new_capa(items.len()); + for item in items { + match item { + BatchItem::Ok(d) => { + let resp = Response::new( + d.status, d.headers, d.body, d.url, d.version, d.content_length, + d.transfer_size, + ); + results.push(ruby.obj_wrap(resp))?; + } + BatchItem::Err(msg) => { + let err: Value = wreq_error().funcall("new", (msg,))?; + results.push(err)?; + } + } + } + Ok(results) + } + + /// Convert a single batch spec into a RequestBuilder. Accepted forms: + /// "https://example.com" -> GET + /// { method:, url:, **opts } + fn build_request( + &self, + spec: Value, + shared: Option<&RHash>, + ) -> Result { + let mut method = wreq::Method::GET; + let url: String; + let mut item_opts: Option = None; + + if let Some(hash) = RHash::from_value(spec) { + url = hash_get_string(&hash, "url")? + .ok_or_else(|| generic_error("each request hash requires a :url"))?; + if let Some(val) = hash_get_value(&hash, "method")? { + method = value_to_method(val)?; + } + item_opts = Some(hash); + } else { + url = TryConvert::try_convert(spec)?; + } + + let mut req = self.inner.request(method, &url); + if let Some(shared) = shared { + req = apply_request_options(req, shared)?; + } + if let Some(item_opts) = item_opts { + req = apply_request_options(req, &item_opts)?; + } + Ok(req) + } +} + +/// Parse a String or Symbol like "post" / :post into an HTTP method. +fn value_to_method(val: Value) -> Result { + let name: String = val.funcall("to_s", ())?; + name.to_uppercase() + .parse() + .map_err(|_| generic_error(format!("invalid HTTP method: {}", name))) } fn apply_request_options( @@ -692,6 +865,7 @@ pub fn init(_ruby: &magnus::Ruby, module: &magnus::RModule) -> Result<(), magnus client_class.define_method("delete", method!(Client::delete, -1))?; client_class.define_method("head", method!(Client::head, -1))?; client_class.define_method("options", method!(Client::options, -1))?; + client_class.define_method("request_batch", method!(Client::request_batch, -1))?; client_class.define_method("cancel", method!(Client::cancel, 0))?; module.define_module_function("get", function!(wreq_get, -1))?; diff --git a/lib/wreq-rb/version.rb b/lib/wreq-rb/version.rb index 6bba5b6..94a445d 100644 --- a/lib/wreq-rb/version.rb +++ b/lib/wreq-rb/version.rb @@ -1,5 +1,5 @@ # frozen_string_literal: true module Wreq - VERSION = "0.5.1" + VERSION = "0.6.0" end diff --git a/test/batch_test.rb b/test/batch_test.rb new file mode 100644 index 0000000..48ed0c8 --- /dev/null +++ b/test/batch_test.rb @@ -0,0 +1,182 @@ +# frozen_string_literal: true + +require_relative "test_helper" + +class BatchTest < Minitest::Test + def test_returns_responses_in_input_order + client = Wreq::Client.new + urls = (1..6).map { |i| "https://httpbun.com/status/20#{i % 5}" } + responses = client.request_batch(urls, concurrency: 6) + + assert_equal urls.size, responses.size + responses.each { |r| assert_kind_of Wreq::Response, r } + assert_equal [201, 202, 203, 204, 200, 201], responses.map(&:status) + end + + def test_empty_array + client = Wreq::Client.new + assert_equal [], client.request_batch([]) + end + + def test_defaults_to_low_concurrency + client = Wreq::Client.new + responses = client.request_batch(Array.new(3, "https://httpbun.com/get")) + assert_equal 3, responses.size + responses.each { |r| assert_equal 200, r.status } + end + + def test_collects_location_headers_without_following_redirects + client = Wreq::Client.new(redirect: false) + urls = %w[ + https://httpbun.com/redirect-to?url=https%3A%2F%2Fexample.com%2Fa + https://httpbun.com/redirect-to?url=https%3A%2F%2Fexample.com%2Fb + ] + responses = client.request_batch(urls, concurrency: 2) + + assert_equal [302, 302], responses.map(&:status) + assert_equal ["https://example.com/a", "https://example.com/b"], + responses.map { |r| r.headers["location"].first } + end + + def test_per_item_errors_do_not_fail_the_batch + client = Wreq::Client.new + responses = client.request_batch( + [ + "https://httpbun.com/get", + "https://this-host-does-not-exist.invalid/", + "https://httpbun.com/get" + ], + concurrency: 3 + ) + + assert_equal 3, responses.size + assert_kind_of Wreq::Response, responses[0] + assert_kind_of Wreq::Error, responses[1] + assert_kind_of Wreq::Response, responses[2] + refute_empty responses[1].message + end + + def test_applies_shared_request_options + client = Wreq::Client.new + responses = client.request_batch( + ["https://httpbun.com/headers"], + concurrency: 1, + headers: { "X-Batch" => "shared" } + ) + assert_equal "shared", responses[0].json["headers"]["X-Batch"] + end + + def test_mixed_specs + client = Wreq::Client.new + specs = [ + "https://httpbun.com/get", + { method: :post, url: "https://httpbun.com/post", json: { a: 1 } }, + { method: :put, url: "https://httpbun.com/put", body: "hello" } + ] + responses = client.request_batch(specs, concurrency: 3) + + assert_equal [200, 200, 200], responses.map(&:status) + assert_equal 1, responses[1].json["json"]["a"] + assert_equal "hello", responses[2].json["data"] + end + + def test_per_item_timeout_only_fails_that_item + client = Wreq::Client.new + responses = client.request_batch( + [ + "https://httpbun.com/get", + { url: "https://httpbun.com/delay/5", timeout: 1 }, + "https://httpbun.com/get" + ], + concurrency: 3 + ) + + assert_kind_of Wreq::Response, responses[0] + assert_kind_of Wreq::Error, responses[1] + assert_kind_of Wreq::Response, responses[2] + assert_match(/time/i, responses[1].message) + end + + def test_shared_timeout_applies_to_every_request + client = Wreq::Client.new + responses = client.request_batch( + Array.new(2, "https://httpbun.com/delay/5"), + concurrency: 2, + timeout: 1 + ) + + assert_equal 2, responses.size + responses.each { |r| assert_kind_of Wreq::Error, r } + end + + def test_per_item_timeout_overrides_the_shared_one + client = Wreq::Client.new + responses = client.request_batch( + [ + "https://httpbun.com/delay/3", + { url: "https://httpbun.com/delay/3", timeout: 10 } + ], + concurrency: 2, + timeout: 1 + ) + + assert_kind_of Wreq::Error, responses[0] + assert_kind_of Wreq::Response, responses[1] + end + + # The timeout clock must start when a request is dispatched, not when it is + # enqueued, or anything queued behind the semaphore would fail spuriously. + def test_timeout_clock_starts_at_dispatch_not_enqueue + client = Wreq::Client.new + responses = client.request_batch( + Array.new(4, "https://httpbun.com/delay/1"), + concurrency: 1, + timeout: 3 + ) + + assert_equal [200, 200, 200, 200], responses.map(&:status) + end + + def test_rejects_zero_concurrency + client = Wreq::Client.new + err = assert_raises(Wreq::Error) do + client.request_batch(["https://httpbun.com/get"], concurrency: 0) + end + assert_match(/concurrency/, err.message) + end + + def test_rejects_hash_without_url + client = Wreq::Client.new + assert_raises(Wreq::Error) do + client.request_batch([{ method: "GET" }]) + end + end + + def test_batch_is_cancellable + client = Wreq::Client.new + t = Thread.new do + client.request_batch(Array.new(4, "https://httpbun.com/delay/10"), concurrency: 4) + end + sleep 1 + client.cancel + + err = assert_raises(Wreq::Error) { t.value } + assert_match(/interrupted/, err.message) + end + + def test_batch_releases_the_gvl + client = Wreq::Client.new + ticks = 0 + ticker = Thread.new do + loop do + ticks += 1 + sleep 0.05 + end + end + + client.request_batch(Array.new(4, "https://httpbun.com/delay/1"), concurrency: 4) + ticker.kill + + assert_operator ticks, :>, 5, "GVL appears to have been held for the whole batch" + end +end