diff --git a/Rakefile b/Rakefile index 537df87..a4b1941 100644 --- a/Rakefile +++ b/Rakefile @@ -20,6 +20,14 @@ Minitest::TestTask.create(:test) do |t| t.test_globs = ["test/**/*_test.rb"] end +namespace :test do + desc "Run native extension tests" + task :native do + sh "cargo test -p wreq_rb --lib" + end +end +task test: "test:native" + # Reset vendored submodules to clean state on rake clean task :reset_submodules do puts "Resetting vendored submodules..." diff --git a/ext/wreq_rb/Cargo.toml b/ext/wreq_rb/Cargo.toml index 9c8302f..c2fc0b6 100644 --- a/ext/wreq_rb/Cargo.toml +++ b/ext/wreq_rb/Cargo.toml @@ -31,6 +31,9 @@ serde_json = "1.0" bytes = "1" http = "1" +[dev-dependencies] +magnus = { version = "0.8", features = ["embed"] } + [target.'cfg(target_os = "linux")'.dependencies] wreq = { path = "../../vendor/wreq", version = "=0.16.1", features = [ "prefix-symbols", diff --git a/ext/wreq_rb/src/client.rs b/ext/wreq_rb/src/client.rs index 2d65eb2..9e76dda 100644 --- a/ext/wreq_rb/src/client.rs +++ b/ext/wreq_rb/src/client.rs @@ -47,7 +47,7 @@ fn runtime() -> &'static Runtime { /// # Safety /// The closure must NOT access any Ruby objects or call any Ruby C API. /// Extract all data from Ruby before calling this, convert results after. -unsafe fn without_gvl(f: F) -> R +unsafe fn without_gvl(f: F) -> Result where F: FnOnce(CancellationToken) -> R, { @@ -89,19 +89,22 @@ where let data_ptr = &mut data as *mut CallData as *mut c_void; unsafe { - rb_sys::rb_thread_call_without_gvl( - Some(call::), - data_ptr, - Some(ubf::), - data_ptr, - ); + magnus::rb_sys::protect(|| { + rb_sys::rb_thread_call_without_gvl( + Some(call::), + data_ptr, + Some(ubf::), + data_ptr, + ); + 0 + })?; } if let Some(payload) = data.panic_payload { panic::resume_unwind(payload); } - data.result.unwrap() + Ok(data.result.unwrap()) } /// Collected response data as pure Rust types (no Ruby objects). @@ -288,12 +291,13 @@ impl Client { builder = builder.default_headers(hmap); } - if let Some(t) = hash_get_float(&opts, "timeout")? { - builder = builder.timeout(Duration::from_secs_f64(t)); + let timeout = hash_get_float(&opts, "timeout")?; + if let Some(timeout) = timeout { + builder = builder.timeout(Duration::from_secs_f64(timeout)); } - if let Some(t) = hash_get_float(&opts, "connect_timeout")? { - builder = builder.connect_timeout(Duration::from_secs_f64(t)); + if let Some(connect_timeout) = hash_get_float(&opts, "connect_timeout")?.or(timeout) { + builder = builder.connect_timeout(Duration::from_secs_f64(connect_timeout)); } if let Some(t) = hash_get_float(&opts, "read_timeout")? { @@ -489,7 +493,7 @@ impl Client { } }) }) - }; + }?; let data = match outcome { RequestOutcome::Ok(d) => d, @@ -545,7 +549,7 @@ impl Client { } }) }) - }; + }?; let items = match outcome { BatchOutcome::Done(items) => items, @@ -877,3 +881,89 @@ pub fn init(_ruby: &magnus::Ruby, module: &magnus::RModule) -> Result<(), magnus Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering}; + + static READY: AtomicBool = AtomicBool::new(false); + static LIVE: AtomicUsize = AtomicUsize::new(0); + + struct DropProbe; + + impl DropProbe { + fn new() -> Self { + LIVE.fetch_add(1, Ordering::SeqCst); + Self + } + } + + impl Drop for DropProbe { + fn drop(&mut self) { + LIVE.fetch_sub(1, Ordering::SeqCst); + } + } + + fn probe_ready() -> bool { + READY.load(Ordering::SeqCst) + } + + fn drop_probe() -> Result<(), magnus::Error> { + let _caller_resource = DropProbe::new(); + let result = unsafe { + without_gvl(|token| { + let _callback_resource = DropProbe::new(); + READY.store(true, Ordering::SeqCst); + runtime().block_on(token.cancelled()); + DropProbe::new() + }) + }?; + drop(result); + Ok(()) + } + + #[test] + fn ruby_interrupts_drop_native_resources() { + let ruby = unsafe { magnus::embed::init() }; + let result = unsafe { without_gvl(|_| DropProbe::new()) }.unwrap(); + assert_eq!(LIVE.load(Ordering::SeqCst), 1); + drop(result); + assert_eq!(LIVE.load(Ordering::SeqCst), 0); + + let panicked = panic::catch_unwind(|| unsafe { + without_gvl::<_, ()>(|_| { + let _resource = DropProbe::new(); + panic!("probe panic"); + }) + }); + assert!(panicked.is_err()); + assert_eq!(LIVE.load(Ordering::SeqCst), 0); + + ruby.define_global_function("wreq_drop_probe", function!(drop_probe, 0)); + ruby.define_global_function("wreq_probe_ready", function!(probe_ready, 0)); + + for (interrupt, expected) in [ + ("worker.kill", None), + ("worker.raise(RuntimeError, 'stop request')", Some("stop request")), + ] { + READY.store(false, Ordering::SeqCst); + let result: Option = ruby.eval(&format!( + "worker = Thread.new do + begin + wreq_drop_probe + 'returned normally' + rescue RuntimeError => error + error.message + end + end + Thread.pass until wreq_probe_ready + {interrupt} + worker.value" + )).unwrap(); + + assert_eq!(result.as_deref(), expected); + assert_eq!(LIVE.load(Ordering::SeqCst), 0, "resources leaked after {interrupt}"); + } + } +} diff --git a/ext/wreq_rb/src/response.rs b/ext/wreq_rb/src/response.rs index 909af85..84e4952 100644 --- a/ext/wreq_rb/src/response.rs +++ b/ext/wreq_rb/src/response.rs @@ -5,7 +5,8 @@ use magnus::{ use crate::error::generic_error; /// Wraps a wreq::Response in a Ruby-accessible type. -#[magnus::wrap(class = "Wreq::Response", free_immediately)] +#[derive(magnus::TypedData)] +#[magnus(class = "Wreq::Response", free_immediately, size)] pub struct Response { status: u16, headers: Vec<(String, String)>, @@ -26,7 +27,7 @@ impl Response { content_length: Option, transfer_size: Option, ) -> Self { - Self { + let response = Self { status, headers, body, @@ -34,7 +35,18 @@ impl Response { version, content_length, transfer_size, - } + }; + unsafe { Ruby::get_unchecked() } + .gc_adjust_memory_usage(response.heap_size() as isize); + response + } + + fn heap_size(&self) -> usize { + self.body.capacity() + + self.url.capacity() + + self.version.capacity() + + self.headers.capacity() * std::mem::size_of::<(String, String)>() + + self.headers.iter().map(|(name, value)| name.capacity() + value.capacity()).sum::() } fn status(&self) -> u16 { @@ -119,6 +131,19 @@ impl Response { } } +impl magnus::DataTypeFunctions for Response { + fn size(&self) -> usize { + std::mem::size_of::() + self.heap_size() + } +} + +impl Drop for Response { + fn drop(&mut self) { + unsafe { Ruby::get_unchecked() } + .gc_adjust_memory_usage(-(self.heap_size() as isize)); + } +} + pub fn init(ruby: &magnus::Ruby, module: &magnus::RModule) -> Result<(), magnus::Error> { let class = module.define_class("Response", ruby.class_object())?; class.define_method("status", method!(Response::status, 0))?; diff --git a/test/client_test.rb b/test/client_test.rb index a23d4ab..c19fd79 100644 --- a/test/client_test.rb +++ b/test/client_test.rb @@ -86,8 +86,60 @@ def test_header_order_takes_precedence_over_emulation "got positions #{positions.inspect} in: #{received.inspect}" end + def test_connect_timeout_defaults_to_client_timeout + [{ timeout: 0.25 }, { timeout: 0.25, connect_timeout: nil }].each do |options| + elapsed = stalled_tls_request_duration(options) + assert_operator elapsed, :<, 2 + end + end + + def test_explicit_connect_timeout_overrides_client_timeout + elapsed = stalled_tls_request_duration({ timeout: 0.25, connect_timeout: 0.8 }) + assert_operator elapsed, :>=, 0.6 + assert_operator elapsed, :<, 2 + end + + def test_connect_timeout_without_client_timeout + elapsed = stalled_tls_request_duration({ connect_timeout: 0.25 }) + assert_operator elapsed, :<, 2 + end + + def test_request_timeout_without_client_timeouts + elapsed = stalled_tls_request_duration({}, request_timeout: 0.5) + assert_operator elapsed, :>=, 0.4 + end + private + def stalled_tls_request_duration(options, request_timeout: 3) + require "socket" + server = TCPServer.new("127.0.0.1", 0) + first_byte = nil + server_thread = Thread.new do + connection = server.accept + first_byte = connection.readpartial(4096).getbyte(0) + connection.read + rescue Errno::ECONNRESET + ensure + connection&.close + end + + client = Wreq::Client.new({ emulation: false, http1_only: true, no_proxy: true }.merge(options)) + started = Process.clock_gettime(Process::CLOCK_MONOTONIC) + assert_raises(Wreq::Error) do + client.get("https://127.0.0.1:#{server.addr[1]}/", timeout: request_timeout) + end + elapsed = Process.clock_gettime(Process::CLOCK_MONOTONIC) - started + assert server_thread.join(2), "Timed-out TLS connection did not close" + server_thread.value + assert_equal 22, first_byte, "Expected a TLS handshake record" + elapsed + ensure + server_thread&.kill + server_thread&.join + server&.close + end + # Spins up a local TCP server, yields the port formatted into a URL, captures # the header names from the raw HTTP/1.1 request, then tears down the server. def capture_wire_headers diff --git a/test/response_test.rb b/test/response_test.rb index 3a057e8..c86037a 100644 --- a/test/response_test.rb +++ b/test/response_test.rb @@ -1,8 +1,22 @@ # frozen_string_literal: true require_relative "test_helper" +require "objspace" +require "socket" class ResponseTest < Minitest::Test + def test_native_body_is_reported_to_object_space + with_large_response do |response, body_size, _increase| + assert_operator ObjectSpace.memsize_of(response), :>=, body_size + end + end + + def test_native_body_counts_toward_gc_threshold + with_large_response do |_response, body_size, increase| + assert_operator increase, :>=, body_size + end + end + def test_response_methods resp = Wreq.get("https://httpbun.com/get") assert_kind_of Integer, resp.status @@ -63,4 +77,31 @@ def test_headers_multiple_set_cookie assert cookies.length >= 2, "expected at least 2 set-cookie values, got #{cookies.length}: #{cookies.inspect}" end + + private + + def with_large_response + body = "a" * (4 * 1024 * 1024) + server = TCPServer.new("127.0.0.1", 0) + server_thread = Thread.new do + connection = server.accept + connection.gets("\r\n\r\n") + connection.write("HTTP/1.1 200 OK\r\nContent-Length: #{body.bytesize}\r\nConnection: close\r\n\r\n") + connection.write(body) + ensure + connection&.close + end + client = Wreq::Client.new(emulation: false, http1_only: true, no_proxy: true, timeout: 5) + was_disabled = GC.disable + before = GC.stat(:malloc_increase_bytes) + response = client.get("http://127.0.0.1:#{server.addr[1]}/") + increase = GC.stat(:malloc_increase_bytes) - before + server_thread.value + yield response, body.bytesize, increase + ensure + GC.enable unless was_disabled + server&.close + server_thread&.kill + server_thread&.join + end end