diff --git a/README.md b/README.md index be0d1ce..4e03ae5 100644 --- a/README.md +++ b/README.md @@ -91,6 +91,10 @@ client = Wreq::Client.new( tcp_nodelay: true, # disable Nagle algorithm (default: true) tcp_keepalive: 15, # SO_KEEPALIVE interval in seconds (default: 15) local_address: "1.2.3.4", # bind outgoing connections to this source IP + resolve: { # override DNS resolution per domain + "example.com" => "1.2.3.4", + "api.example.com" => ["1.2.3.4:8443", "[::1]:8443"], + }, tls_sni: true, # send SNI in TLS handshake (default: true) min_tls_version: "tls1.2", # minimum TLS version: tls1.0, tls1.1, tls1.2, tls1.3 max_tls_version: "tls1.3", # maximum TLS version @@ -100,6 +104,40 @@ resp = client.get("https://api.example.com/data") resp = client.post("https://api.example.com/data", json: { key: "value" }) ``` +## Overriding DNS + +Pass `resolve` to point specific domains at IP addresses of your choice, +bypassing DNS for them (the equivalent of curl's `--resolve`). Useful for +hitting a staging host by IP while keeping the real `Host` header and TLS SNI, +or for pinning a domain to one node of a load balancer. + +```ruby +client = Wreq::Client.new( + resolve: { + "example.com" => "127.0.0.1", # one address + "api.example.com" => ["10.0.0.1", "10.0.0.2"], # tried in order + } +) + +resp = client.get("https://example.com/") # connects to 127.0.0.1:443 +``` + +Addresses are IPv4 or IPv6, with an optional port: + +| Value | Connects to | +|-------|-------------| +| `"1.2.3.4"` | `1.2.3.4` on the URL's port, or the scheme default | +| `"1.2.3.4:8443"` | `1.2.3.4:8443` | +| `"::1"` / `"[::1]"` | `[::1]` on the URL's port, or the scheme default | +| `"[::1]:8443"` | `[::1]:8443` | + +An IPv6 address with a port **must** be bracketed: `"::1:8080"` is itself a +valid IPv6 address, so it is parsed as one (with no port) rather than as `::1` +on port 8080. + +An explicit port in the request URL always wins over a port in the override. +Domains without an entry resolve normally. + ## HTTP Methods All methods are available on both `Wreq` (module-level) and `Wreq::Client` (instance-level): diff --git a/ext/wreq_rb/src/client.rs b/ext/wreq_rb/src/client.rs index 9e76dda..1b5dc2f 100644 --- a/ext/wreq_rb/src/client.rs +++ b/ext/wreq_rb/src/client.rs @@ -13,7 +13,7 @@ use tokio::runtime::Runtime; use tokio::sync::Semaphore; use tokio::task::JoinSet; use tokio_util::sync::CancellationToken; -use std::net::IpAddr; +use std::net::{IpAddr, SocketAddr}; use wreq::header::{HeaderMap, HeaderName, HeaderValue, OrigHeaderMap}; use wreq::tls::TlsVersion; use wreq_util::{Emulation as BrowserEmulation, Platform as EmulationPlatform, Profile as BrowserProfile}; @@ -393,6 +393,12 @@ impl Client { builder = builder.local_address(addr); } + if let Some(resolve_hash) = hash_get_hash(&opts, "resolve")? { + for (domain, addrs) in hash_to_dns_overrides(&resolve_hash)? { + builder = builder.resolve_to_addrs(domain, addrs); + } + } + if let Some(v) = hash_get_bool(&opts, "tls_sni")? { builder = builder.tls_sni(v); } @@ -838,6 +844,59 @@ fn hash_to_header_map(hash: &RHash) -> Result { Ok(hmap) } +/// Parse a DNS override target into a SocketAddr. Accepts a bare IP +/// ("1.2.3.4", "::1", "[::1]") or an IP with a port ("1.2.3.4:8080", +/// "[::1]:8080"). A missing port is stored as 0, which makes wreq fall back to +/// the port in the request URL (or the scheme's default port). +/// +/// An IPv6 address with a port must be bracketed: "::1:8080" is a valid IPv6 +/// address in its own right and parses as one, not as "::1" on port 8080. +fn parse_socket_addr(s: &str) -> Result { + if let Ok(addr) = s.parse::() { + return Ok(addr); + } + let bare = s.strip_prefix('[').and_then(|s| s.strip_suffix(']')).unwrap_or(s); + match bare.parse::() { + Ok(ip) => Ok(SocketAddr::new(ip, 0)), + Err(_) => Err(generic_error(format!( + "invalid resolve address: '{}'. Use an IP ('1.2.3.4', '::1') \ + or an IP with a port ('1.2.3.4:8080', '[::1]:8080')", s + ))), + } +} + +/// Convert a `resolve:` hash of domain => address(es) into DNS overrides. Each +/// value is a single address string or an array of them. +fn hash_to_dns_overrides(hash: &RHash) -> Result)>, magnus::Error> { + let mut overrides: Vec<(String, Vec)> = Vec::new(); + hash.foreach(|k: Value, v: Value| { + let ruby = unsafe { Ruby::get_unchecked() }; + let domain: String = if k.is_kind_of(ruby.class_symbol()) { + k.funcall("to_s", ())? + } else { + TryConvert::try_convert(k)? + }; + let mut addrs: Vec = Vec::new(); + if v.is_kind_of(ruby.class_array()) { + for elem in RArray::try_convert(v)?.into_iter() { + let addr_str: String = TryConvert::try_convert(elem)?; + addrs.push(parse_socket_addr(&addr_str)?); + } + } else { + let addr_str: String = TryConvert::try_convert(v)?; + addrs.push(parse_socket_addr(&addr_str)?); + } + if addrs.is_empty() { + return Err(generic_error(format!( + "resolve entry for '{}' has no addresses", domain + ))); + } + overrides.push((domain, addrs)); + Ok(magnus::r_hash::ForEach::Continue) + })?; + Ok(overrides) +} + fn hash_to_pairs(hash: &RHash) -> Result, magnus::Error> { let mut pairs: Vec<(String, String)> = Vec::new(); hash.foreach(|k: Value, v: Value| { diff --git a/test/client_test.rb b/test/client_test.rb index c19fd79..9f280c6 100644 --- a/test/client_test.rb +++ b/test/client_test.rb @@ -109,8 +109,94 @@ def test_request_timeout_without_client_timeouts assert_operator elapsed, :>=, 0.4 end + def test_resolve_overrides_dns_for_a_domain + host = with_local_http_server do |port| + client = resolving_client("wreq-resolve.invalid" => "127.0.0.1") + assert_equal 200, client.get("http://wreq-resolve.invalid:#{port}/").status + end + assert_match(/\Awreq-resolve\.invalid:/, host) + end + + def test_resolve_honours_the_port_in_the_override + # The URL carries no explicit port, so the port from the override is used + # instead of the scheme default (80). + host = with_local_http_server do |port| + client = resolving_client("wreq-resolve.invalid" => "127.0.0.1:#{port}") + assert_equal 200, client.get("http://wreq-resolve.invalid/").status + end + assert_equal "wreq-resolve.invalid", host + end + + def test_resolve_accepts_multiple_addresses + dead_port = free_port + host = with_local_http_server do |port| + # The first address refuses connections, so wreq falls through to the second. + client = resolving_client( + "wreq-resolve.invalid" => ["127.0.0.1:#{dead_port}", "127.0.0.1:#{port}"] + ) + assert_equal 200, client.get("http://wreq-resolve.invalid/", timeout: 10).status + end + assert_equal "wreq-resolve.invalid", host + end + + def test_resolve_leaves_other_domains_alone + client = resolving_client("wreq-resolve.invalid" => "127.0.0.1") + assert_equal 200, client.get("https://httpbun.com/get").status + end + + def test_resolve_rejects_an_invalid_address + error = assert_raises(Wreq::Error) do + Wreq::Client.new(resolve: { "example.com" => "not-an-ip" }) + end + assert_match(/invalid resolve address/, error.message) + end + private + def resolving_client(overrides) + Wreq::Client.new( + emulation: false, + http1_only: true, + no_proxy: true, + resolve: overrides + ) + end + + def free_port + server = TCPServer.new("127.0.0.1", 0) + port = server.addr[1] + server.close + port + end + + # Spins up a local HTTP/1.1 server, yields its port, and returns the Host + # header of the request it received. + def with_local_http_server + require "socket" + server = TCPServer.new("127.0.0.1", 0) + host = nil + t = Thread.new do + conn = server.accept + conn.gets # skip request line + loop do + line = conn.gets&.chomp + break if line.nil? || line.empty? + name, value = line.split(":", 2) + host = value.strip if name.downcase == "host" + end + conn.write "HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: close\r\n\r\n" + conn.close + rescue + conn&.close + end + yield server.addr[1] + t.join(5) + host + ensure + t&.kill + server&.close + end + def stalled_tls_request_duration(options, request_timeout: 3) require "socket" server = TCPServer.new("127.0.0.1", 0)