Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions Rakefile
Original file line number Diff line number Diff line change
Expand Up @@ -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..."
Expand Down
3 changes: 3 additions & 0 deletions ext/wreq_rb/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
118 changes: 104 additions & 14 deletions ext/wreq_rb/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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, R>(f: F) -> R
unsafe fn without_gvl<F, R>(f: F) -> Result<R, magnus::Error>
where
F: FnOnce(CancellationToken) -> R,
{
Expand Down Expand Up @@ -89,19 +89,22 @@ where
let data_ptr = &mut data as *mut CallData<F, R> as *mut c_void;

unsafe {
rb_sys::rb_thread_call_without_gvl(
Some(call::<F, R>),
data_ptr,
Some(ubf::<F, R>),
data_ptr,
);
magnus::rb_sys::protect(|| {
rb_sys::rb_thread_call_without_gvl(
Some(call::<F, R>),
data_ptr,
Some(ubf::<F, R>),
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).
Expand Down Expand Up @@ -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")? {
Expand Down Expand Up @@ -489,7 +493,7 @@ impl Client {
}
})
})
};
}?;

let data = match outcome {
RequestOutcome::Ok(d) => d,
Expand Down Expand Up @@ -545,7 +549,7 @@ impl Client {
}
})
})
};
}?;

let items = match outcome {
BatchOutcome::Done(items) => items,
Expand Down Expand Up @@ -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<String> = 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}");
}
}
}
31 changes: 28 additions & 3 deletions ext/wreq_rb/src/response.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)>,
Expand All @@ -26,15 +27,26 @@ impl Response {
content_length: Option<u64>,
transfer_size: Option<u64>,
) -> Self {
Self {
let response = Self {
status,
headers,
body,
url,
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::<usize>()
}

fn status(&self) -> u16 {
Expand Down Expand Up @@ -119,6 +131,19 @@ impl Response {
}
}

impl magnus::DataTypeFunctions for Response {
fn size(&self) -> usize {
std::mem::size_of::<Self>() + 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))?;
Expand Down
52 changes: 52 additions & 0 deletions test/client_test.rb
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
41 changes: 41 additions & 0 deletions test/response_test.rb
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Loading