Skip to content
Draft
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
1 change: 1 addition & 0 deletions crates/core/src/host/host_controller.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1004,6 +1004,7 @@ async fn make_replica_ctx(
subscriptions,
module_instance_memory_tracker,
module_http,
module_http_client: Default::default(),
})
}

Expand Down
116 changes: 92 additions & 24 deletions crates/core/src/host/instance_env.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1039,22 +1039,14 @@ impl InstanceEnv {
return Err(NodesError::HttpError(BLOCKED_HTTP_ADDRESS_ERROR.to_string()));
}

let redirect_policy = reqwest::redirect::Policy::custom(|attempt| {
if is_blocked_ip_literal(attempt.url()) {
attempt.error(BLOCKED_HTTP_ADDRESS_ERROR)
} else {
reqwest::redirect::Policy::default().redirect(attempt)
}
});

// TODO(procedure-metrics): record size in bytes of response, time spent awaiting response.

// Actually execute the HTTP request!
// TODO(perf): Stash a long-lived `Client` in the env somewhere, rather than building a new one for each call.
let execute_fut = reqwest::Client::builder()
.dns_resolver(Arc::new(FilteredDnsResolver))
.redirect(redirect_policy)
.build()
// All of the replica's requests share one client, and so one connection pool.
let execute_fut = self
.replica_ctx
.module_http_client
.get_or_try_init(build_module_http_client)
.map_err(http_error)?
.execute(reqwest);

Expand Down Expand Up @@ -1131,6 +1123,23 @@ const HTTP_MAX_TIMEOUT: Duration = Duration::from_secs(180);
const MODULE_HTTP_DISABLED_ERROR: &str = "module outbound HTTP requests are disabled";
const BLOCKED_HTTP_ADDRESS_ERROR: &str = "refusing to connect to private or special-purpose addresses";

/// Build the client [`InstanceEnv::http_request`] sends requests with,
/// which refuses to connect to private or special-purpose addresses,
/// whether they come from DNS or from an IP literal in a redirect.
fn build_module_http_client() -> reqwest::Result<reqwest::Client> {
let redirect_policy = reqwest::redirect::Policy::custom(|attempt| {
if is_blocked_ip_literal(attempt.url()) {
attempt.error(BLOCKED_HTTP_ADDRESS_ERROR)
} else {
reqwest::redirect::Policy::default().redirect(attempt)
}
});
reqwest::Client::builder()
.dns_resolver(Arc::new(FilteredDnsResolver))
.redirect(redirect_policy)
.build()
}

struct FilteredDnsResolver;

impl reqwest::dns::Resolve for FilteredDnsResolver {
Expand Down Expand Up @@ -1497,6 +1506,7 @@ mod test {
subscriptions: subs,
module_instance_memory_tracker: ModuleInstanceMemoryTracker::new(Identity::ZERO, Arc::new(())),
module_http,
module_http_client: Default::default(),
},
runtime,
))
Expand Down Expand Up @@ -2586,7 +2596,7 @@ mod test {
/// otherwise blocked by [`is_blocked_ip`].
#[cfg(feature = "allow_loopback_http_for_tests")]
fn spawn_encoded_body_server(encoding: &str, body: Vec<u8>) -> (u16, std::thread::JoinHandle<String>) {
use std::io::{Read, Write};
use std::io::Write;
use std::net::TcpListener;

let listener = TcpListener::bind(("127.0.0.1", 0)).expect("failed to bind test server");
Expand All @@ -2604,16 +2614,7 @@ mod test {
);
let handle = std::thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("test server failed to accept");

// Read the request head. We never send a request body, so `\r\n\r\n` terminates it.
let mut request = Vec::new();
let mut byte = [0u8; 1];
while !request.ends_with(b"\r\n\r\n") {
match stream.read(&mut byte).expect("test server failed to read request") {
0 => break,
_ => request.push(byte[0]),
}
}
let request = read_request_head(&mut stream);

stream
.write_all(head.as_bytes())
Expand All @@ -2628,6 +2629,53 @@ mod test {
(port, handle)
}

/// Read a request head, which `\r\n\r\n` terminates as our test requests have no body.
/// Returns an empty head if the client closed the connection first.
#[cfg(feature = "allow_loopback_http_for_tests")]
fn read_request_head(stream: &mut impl std::io::Read) -> Vec<u8> {
let mut request = Vec::new();
let mut byte = [0u8; 1];
while !request.ends_with(b"\r\n\r\n") {
match stream.read(&mut byte).expect("test server failed to read request") {
0 => break,
_ => request.push(byte[0]),
}
}
request
}

/// Spin up a keep-alive HTTP/1.1 server on loopback which answers `requests` requests,
/// on as many connections as the client opens.
///
/// Returns the bound port, and a handle yielding how many connections it accepted.
#[cfg(feature = "allow_loopback_http_for_tests")]
fn spawn_keep_alive_server(requests: usize) -> (u16, std::thread::JoinHandle<usize>) {
use std::io::Write;
use std::net::TcpListener;

let listener = TcpListener::bind(("127.0.0.1", 0)).expect("failed to bind test server");
let port = listener
.local_addr()
.expect("failed to read test server address")
.port();
let handle = std::thread::spawn(move || {
let (mut served, mut accepted) = (0, 0);
while served < requests {
let (mut stream, _) = listener.accept().expect("test server failed to accept");
accepted += 1;
// Serve this connection until the client closes it or we're done.
while served < requests && !read_request_head(&mut stream).is_empty() {
stream
.write_all(b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok")
.expect("test server failed to write response");
served += 1;
}
}
accepted
});
(port, handle)
}

/// `GET http://127.0.0.1:{port}/` through [`InstanceEnv::http_request`],
/// i.e. the same path a procedure's `ctx.http` call takes.
#[cfg(feature = "allow_loopback_http_for_tests")]
Expand Down Expand Up @@ -2776,4 +2824,24 @@ mod test {

Ok(())
}

/// Consecutive requests from a replica reuse one pooled connection,
/// rather than paying for a new connection (and TLS handshake) each time.
#[test]
#[cfg(feature = "allow_loopback_http_for_tests")]
fn http_requests_reuse_connection() -> Result<()> {
let db = relational_db()?;
let (mut env, runtime) = instance_env(db)?;

let (port, server) = spawn_keep_alive_server(2);
for _ in 0..2 {
let (response, body) = get_via_http_request(&mut env, &runtime, port);
assert_eq!(response.code, 200);
assert_eq!(&body[..], b"ok");
}
let accepted = server.join().expect("test server thread panicked");

assert_eq!(accepted, 1, "both requests should share one connection");
Ok(())
}
}
4 changes: 4 additions & 0 deletions crates/core/src/replica_context.rs
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ use crate::error::DBError;
use crate::messages::control_db::Database;
use crate::resource::ModuleInstanceMemoryTracker;
use crate::subscription::module_subscription_actor::ModuleSubscriptions;
use once_cell::sync::OnceCell;
use std::io;
use std::ops::Deref;
use std::sync::Arc;
Expand All @@ -22,6 +23,9 @@ pub struct ReplicaContext {
pub subscriptions: ModuleSubscriptions,
pub module_instance_memory_tracker: ModuleInstanceMemoryTracker,
pub module_http: ModuleHttpConfig,
/// The client for procedure HTTP requests, built on first use.
/// Shared by all of this replica's module instances, so they reuse connections.
pub module_http_client: OnceCell<reqwest::Client>,
}

impl ReplicaContext {
Expand Down
15 changes: 15 additions & 0 deletions tools/ci/commands/test/src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,21 @@ fn main() -> Result<()> {
"--test-threads=2",
)
.run()?;
// Procedure HTTP tests that talk to a loopback server only build with this feature,
// which `cargo test --all` above doesn't enable.
cmd!(
"cargo",
"test",
"-p",
"spacetimedb-core",
"--lib",
"--features",
"allow_loopback_http_for_tests",
"--",
"--test-threads=2",
"host::instance_env::test::http_request",
)
.run()?;
// The SDK test harness uses the same child-process server guard as smoketests,
// which expects release CLI/standalone binaries to already exist.
if !use_prebuilt_runtime {
Expand Down
Loading