Release v2.29.0 #312
@@ -124,8 +124,12 @@ bool file_send_sendfile_with_skip(File* file, int file_descriptor, bool use_meta
|
||||
}
|
||||
|
||||
/* sendfile cannot encrypt TLS records. Keep the framing identical but
|
||||
route encrypted transfers through the deadline-aware IO layer. */
|
||||
if (io_get_ssl() != NULL) {
|
||||
route encrypted transfers through the deadline-aware IO layer. Resolve
|
||||
the transport from the bound session, not the thread-local io_ssl: a
|
||||
worker thread running a TLS transfer has its SSL only on the session it
|
||||
bound, so io_get_ssl() would be NULL there and the raw sendfile() path
|
||||
would be taken on an encrypted socket. */
|
||||
if (protocol_current_ssl() != NULL) {
|
||||
unsigned char buffer[64 * 1024];
|
||||
unsigned long long remaining = file_size;
|
||||
bool ok = true;
|
||||
|
||||
+160
-75
@@ -38,6 +38,122 @@ static atomic_ullong io_bytes_read = 0;
|
||||
|
||||
static unsigned long long global_bwlimit(void);
|
||||
|
||||
/* ------------------------------------------------------------------------- *
|
||||
* Transport vtable implementations.
|
||||
*
|
||||
* Each op performs exactly one transfer attempt. WANT_READ/WANT_WRITE and an
|
||||
* EINTR-interrupted syscall are reported as PROTOCOL_IO_RETRY (with
|
||||
* *wait_events set to the poll event the caller must wait on); a clean peer
|
||||
* close is PROTOCOL_IO_CLOSED and anything else is PROTOCOL_IO_ERROR. This
|
||||
* keeps every WANT_READ/WANT_WRITE and EINTR retry exactly where it was before
|
||||
* the vtable was introduced, just moved behind the function pointer.
|
||||
* ------------------------------------------------------------------------- */
|
||||
|
||||
static ssize_t plain_io_send(ProtocolSession* session, const void* data, size_t size,
|
||||
short* wait_events) {
|
||||
ssize_t written = write(session->write_fd, data, size);
|
||||
if (written < 0) {
|
||||
if (errno == EINTR)
|
||||
return PROTOCOL_IO_RETRY;
|
||||
return PROTOCOL_IO_ERROR;
|
||||
}
|
||||
if (written == 0)
|
||||
return PROTOCOL_IO_ERROR;
|
||||
*wait_events = POLLOUT;
|
||||
return written;
|
||||
}
|
||||
|
||||
static ssize_t plain_io_recv(ProtocolSession* session, void* data, size_t size,
|
||||
short* wait_events) {
|
||||
ssize_t received = read(session->read_fd, data, size);
|
||||
if (received < 0) {
|
||||
if (errno == EINTR)
|
||||
return PROTOCOL_IO_RETRY;
|
||||
return PROTOCOL_IO_ERROR;
|
||||
}
|
||||
if (received == 0)
|
||||
return PROTOCOL_IO_CLOSED;
|
||||
*wait_events = POLLIN;
|
||||
return received;
|
||||
}
|
||||
|
||||
static bool plain_io_has_pending(const ProtocolSession* session) {
|
||||
(void)session;
|
||||
return false;
|
||||
}
|
||||
|
||||
static ssize_t tls_io_send(ProtocolSession* session, const void* data, size_t size,
|
||||
short* wait_events) {
|
||||
/* SSL_write takes an int length; clamp a >INT_MAX request into chunks so the
|
||||
* size_t downcast can never truncate into a negative/partial write. */
|
||||
size_t chunk = size > (size_t)INT_MAX ? (size_t)INT_MAX : size;
|
||||
ssize_t written = SSL_write(session->ssl, data, (int)chunk);
|
||||
if (written <= 0) {
|
||||
int ssl_err = SSL_get_error(session->ssl, (int)written);
|
||||
if (ssl_err == SSL_ERROR_WANT_WRITE) {
|
||||
*wait_events = POLLOUT;
|
||||
return PROTOCOL_IO_RETRY;
|
||||
}
|
||||
if (ssl_err == SSL_ERROR_WANT_READ) {
|
||||
*wait_events = POLLIN;
|
||||
return PROTOCOL_IO_RETRY;
|
||||
}
|
||||
/* A signal (e.g. Ctrl-C) interrupts the blocking TLS write: retry so the
|
||||
* send loop can observe the abort flag at the next checkpoint. */
|
||||
if (ssl_err == SSL_ERROR_SYSCALL && errno == EINTR)
|
||||
return PROTOCOL_IO_RETRY;
|
||||
return PROTOCOL_IO_ERROR;
|
||||
}
|
||||
*wait_events = POLLOUT;
|
||||
return written;
|
||||
}
|
||||
|
||||
static ssize_t tls_io_recv(ProtocolSession* session, void* data, size_t size, short* wait_events) {
|
||||
/* SSL_read takes an int length; clamp a >INT_MAX request into chunks
|
||||
* (mirrors the send path) so the size_t downcast can never truncate into a
|
||||
* negative/partial read. */
|
||||
size_t chunk = size > (size_t)INT_MAX ? (size_t)INT_MAX : size;
|
||||
ssize_t received = SSL_read(session->ssl, data, (int)chunk);
|
||||
if (received <= 0) {
|
||||
int ssl_err = SSL_get_error(session->ssl, (int)received);
|
||||
if (ssl_err == SSL_ERROR_WANT_WRITE) {
|
||||
*wait_events = POLLOUT;
|
||||
return PROTOCOL_IO_RETRY;
|
||||
}
|
||||
if (ssl_err == SSL_ERROR_WANT_READ) {
|
||||
*wait_events = POLLIN;
|
||||
return PROTOCOL_IO_RETRY;
|
||||
}
|
||||
/* A signal interrupts the blocking TLS read: retry (mirrors the send path)
|
||||
* so the loop reaches its next abort/deadline checkpoint. */
|
||||
if (ssl_err == SSL_ERROR_SYSCALL && errno == EINTR)
|
||||
return PROTOCOL_IO_RETRY;
|
||||
/* A zero-length SSL_read is the peer's clean close_notify (or EOF without
|
||||
* one); report it distinctly so the caller can log it as a close. */
|
||||
if (received == 0)
|
||||
return PROTOCOL_IO_CLOSED;
|
||||
return PROTOCOL_IO_ERROR;
|
||||
}
|
||||
*wait_events = POLLIN;
|
||||
return received;
|
||||
}
|
||||
|
||||
static bool tls_io_has_pending(const ProtocolSession* session) {
|
||||
return session->ssl != NULL && SSL_pending(session->ssl) > 0;
|
||||
}
|
||||
|
||||
static const ProtocolIoOps plain_io_ops = {
|
||||
.send = plain_io_send,
|
||||
.recv = plain_io_recv,
|
||||
.has_pending = plain_io_has_pending,
|
||||
};
|
||||
|
||||
static const ProtocolIoOps tls_io_ops = {
|
||||
.send = tls_io_send,
|
||||
.recv = tls_io_recv,
|
||||
.has_pending = tls_io_has_pending,
|
||||
};
|
||||
|
||||
static bool protocol_reserve_memory(ProtocolSession* session, size_t charge) {
|
||||
unsigned long long allocated = atomic_load(&session->total_allocated_bytes);
|
||||
while (true) {
|
||||
@@ -79,6 +195,7 @@ void io_set_fds(int read_fd, int write_fd) {
|
||||
legacy_io_session.read_fd = read_fd;
|
||||
legacy_io_session.write_fd = write_fd;
|
||||
legacy_io_session.ssl = NULL;
|
||||
legacy_io_session.ops = &plain_io_ops;
|
||||
legacy_io_session.eight_bit_output = false;
|
||||
atomic_store(&legacy_io_session.total_allocated_bytes, 0);
|
||||
legacy_io_session.max_alloc = DEFAULT_MAX_ALLOC;
|
||||
@@ -91,6 +208,7 @@ void protocol_session_init(ProtocolSession* session, int read_fd, int write_fd)
|
||||
memset(session, 0, sizeof(*session));
|
||||
session->read_fd = read_fd;
|
||||
session->write_fd = write_fd;
|
||||
session->ops = &plain_io_ops;
|
||||
session->max_alloc = DEFAULT_MAX_ALLOC;
|
||||
session->io_timeout_sec = RECEIVE_TIMEOUT_SEC;
|
||||
atomic_init(&session->total_allocated_bytes, 0);
|
||||
@@ -158,8 +276,12 @@ void protocol_session_unbind(void) {
|
||||
}
|
||||
|
||||
void protocol_session_set_ssl(ProtocolSession* session, SSL* ssl) {
|
||||
if (session)
|
||||
if (!session)
|
||||
return;
|
||||
session->ssl = ssl;
|
||||
/* Select the transport dispatch once, here, instead of branching on the SSL
|
||||
* pointer inside every I/O loop. */
|
||||
session->ops = ssl ? &tls_io_ops : &plain_io_ops;
|
||||
}
|
||||
|
||||
static void bw_mutex_init(void) {
|
||||
@@ -270,6 +392,17 @@ SSL* io_get_ssl(void) {
|
||||
return io_ssl;
|
||||
}
|
||||
|
||||
SSL* protocol_current_ssl(void) {
|
||||
/* The bound session is the authoritative transport for a worker thread: it
|
||||
* was explicitly handed to protocol_session_bind() and carries its own SSL,
|
||||
* whereas io_ssl is thread-local and NULL in a thread that never performed
|
||||
* the handshake. With no session bound (the fd-shim path), fall back to the
|
||||
* legacy thread-local SSL. */
|
||||
if (bound_session)
|
||||
return bound_session->ssl;
|
||||
return io_ssl;
|
||||
}
|
||||
|
||||
unsigned long long protocol_bytes_written(void) {
|
||||
return atomic_load(&io_bytes_written);
|
||||
}
|
||||
@@ -298,6 +431,7 @@ static ProtocolSession* legacy_session(int read_fd, int write_fd) {
|
||||
protocol_session_set_bwlimit(&legacy_io_session, global_bwlimit());
|
||||
}
|
||||
legacy_io_session.ssl = io_ssl;
|
||||
legacy_io_session.ops = io_ssl ? &tls_io_ops : &plain_io_ops;
|
||||
return &legacy_io_session;
|
||||
}
|
||||
|
||||
@@ -335,8 +469,9 @@ bool protocol_send_n_data(ProtocolSession* session, const void* data, size_t dat
|
||||
if (!data && data_size != 0)
|
||||
return false;
|
||||
log_debug_message(LOG_DEBUG_IO, " Sending n Data: %zu", data_size);
|
||||
if (!session)
|
||||
if (!session || !session->ops)
|
||||
return false;
|
||||
const ProtocolIoOps* ops = session->ops;
|
||||
/* A non-positive session timeout disables the deadline entirely (rsync's
|
||||
* --timeout=0 default); poll then blocks until the socket becomes writable. */
|
||||
int timeout_sec = session->io_timeout_sec > 0 ? session->io_timeout_sec : 0;
|
||||
@@ -362,36 +497,16 @@ bool protocol_send_n_data(ProtocolSession* session, const void* data, size_t dat
|
||||
continue;
|
||||
if (pfd.revents & (POLLERR | POLLNVAL))
|
||||
return false;
|
||||
ssize_t bytes_send;
|
||||
if (session->ssl) {
|
||||
/* SSL_write takes an int length; clamp a >INT_MAX request into chunks so
|
||||
* the size_t downcast can never truncate into a negative/partial write. */
|
||||
size_t ssl_chunk = chunk > (size_t)INT_MAX ? (size_t)INT_MAX : chunk;
|
||||
bytes_send = SSL_write(session->ssl, (const char*)data + total_bytes_send, (int)ssl_chunk);
|
||||
} else {
|
||||
bytes_send = write(fd, (const char*)data + total_bytes_send, chunk);
|
||||
}
|
||||
ssize_t bytes_send =
|
||||
ops->send(session, (const char*)data + total_bytes_send, chunk, &wait_events);
|
||||
if (bytes_send == PROTOCOL_IO_RETRY)
|
||||
continue;
|
||||
if (bytes_send <= 0) {
|
||||
if (session->ssl) {
|
||||
int ssl_err = SSL_get_error(session->ssl, (int)bytes_send);
|
||||
if (ssl_err == SSL_ERROR_WANT_WRITE || ssl_err == SSL_ERROR_WANT_READ) {
|
||||
wait_events = ssl_err == SSL_ERROR_WANT_WRITE ? POLLOUT : POLLIN;
|
||||
continue;
|
||||
}
|
||||
/* A signal (e.g. Ctrl-C) interrupts the blocking TLS write: retry so
|
||||
the send loop can observe the abort flag at the next checkpoint. */
|
||||
if (ssl_err == SSL_ERROR_SYSCALL && errno == EINTR)
|
||||
continue;
|
||||
} else if (errno == EINTR) {
|
||||
continue;
|
||||
}
|
||||
log_message(LOG_LEVEL_ERROR, "Could not send data");
|
||||
return false;
|
||||
}
|
||||
bw_throttle_session(session, (size_t)bytes_send);
|
||||
total_bytes_send += bytes_send;
|
||||
if (session->ssl)
|
||||
wait_events = POLLOUT;
|
||||
}
|
||||
log_debug_message(LOG_DEBUG_IO, " Send n Data: %zd", total_bytes_send);
|
||||
atomic_fetch_add(&io_bytes_written, (unsigned long long)total_bytes_send);
|
||||
@@ -424,14 +539,15 @@ bool protocol_receive_n_data(ProtocolSession* session, void* data, size_t data_s
|
||||
static bool protocol_receive_n_data_until(ProtocolSession* session, void* data, size_t data_size,
|
||||
const struct timespec* deadline) {
|
||||
log_debug_message(LOG_DEBUG_IO, " Receiving n Data: %zu", data_size);
|
||||
if (!session)
|
||||
if (!session || !session->ops)
|
||||
return false;
|
||||
const ProtocolIoOps* ops = session->ops;
|
||||
int fd = session->read_fd;
|
||||
|
||||
size_t total_bytes_received = 0;
|
||||
short wait_events = POLLIN;
|
||||
while (total_bytes_received < data_size) {
|
||||
if (!session->ssl || SSL_pending(session->ssl) == 0) {
|
||||
if (!ops->has_pending(session)) {
|
||||
struct pollfd pfd = {.fd = fd, .events = wait_events};
|
||||
/* A NULL deadline means "wait indefinitely" (timeout disabled). */
|
||||
int poll_result = poll(&pfd, 1, deadline ? deadline_remaining_ms(deadline) : -1);
|
||||
@@ -449,43 +565,19 @@ static bool protocol_receive_n_data_until(ProtocolSession* session, void* data,
|
||||
return false;
|
||||
}
|
||||
|
||||
ssize_t bytes_received;
|
||||
if (session->ssl) {
|
||||
/* SSL_read takes an int length; clamp a >INT_MAX request into chunks
|
||||
* (mirrors the send path) so the size_t downcast can never truncate into
|
||||
* a negative/partial read. */
|
||||
size_t ssl_chunk = data_size - total_bytes_received > (size_t)INT_MAX
|
||||
? (size_t)INT_MAX
|
||||
: data_size - total_bytes_received;
|
||||
bytes_received = SSL_read(session->ssl, (char*)data + total_bytes_received, (int)ssl_chunk);
|
||||
} else {
|
||||
bytes_received =
|
||||
read(fd, (char*)data + total_bytes_received, data_size - total_bytes_received);
|
||||
ssize_t bytes_received = ops->recv(session, (char*)data + total_bytes_received,
|
||||
data_size - total_bytes_received, &wait_events);
|
||||
if (bytes_received == PROTOCOL_IO_RETRY)
|
||||
continue;
|
||||
if (bytes_received == PROTOCOL_IO_CLOSED) {
|
||||
log_message(LOG_LEVEL_ERROR, "Connection closed while receiving data");
|
||||
return false;
|
||||
}
|
||||
if (bytes_received <= 0) {
|
||||
if (session->ssl) {
|
||||
int ssl_err = SSL_get_error(session->ssl, (int)bytes_received);
|
||||
if (ssl_err == SSL_ERROR_WANT_WRITE || ssl_err == SSL_ERROR_WANT_READ) {
|
||||
wait_events = ssl_err == SSL_ERROR_WANT_WRITE ? POLLOUT : POLLIN;
|
||||
continue;
|
||||
}
|
||||
/* A signal interrupts the blocking TLS read: retry (mirrors the send
|
||||
path and protocol_read_status_until) so the loop reaches its next
|
||||
abort/deadline checkpoint instead of failing spuriously. */
|
||||
if (ssl_err == SSL_ERROR_SYSCALL && errno == EINTR)
|
||||
continue;
|
||||
} else if (errno == EINTR) {
|
||||
continue;
|
||||
}
|
||||
if (bytes_received == 0)
|
||||
log_message(LOG_LEVEL_ERROR, "Connection closed while receiving data");
|
||||
else
|
||||
log_message(LOG_LEVEL_ERROR, "Could not receive bytes");
|
||||
return false;
|
||||
}
|
||||
total_bytes_received += (size_t)bytes_received;
|
||||
if (session->ssl)
|
||||
wait_events = POLLIN;
|
||||
}
|
||||
log_debug_message(LOG_DEBUG_IO, " Received n Data: %zu", total_bytes_received);
|
||||
atomic_fetch_add(&io_bytes_read, (unsigned long long)total_bytes_received);
|
||||
@@ -845,11 +937,14 @@ bool protocol_receive_status_timed(ProtocolSession* session, Status* status, int
|
||||
* reply across a frame boundary. Returns false on timeout/EOF/error. */
|
||||
static bool protocol_read_status_until(ProtocolSession* session, Status* status,
|
||||
const struct timespec* deadline) {
|
||||
if (!session || !session->ops)
|
||||
return false;
|
||||
const ProtocolIoOps* ops = session->ops;
|
||||
Status received = STATUS_ERROR;
|
||||
size_t got = 0;
|
||||
short wait_events = POLLIN;
|
||||
while (got < sizeof(Status)) {
|
||||
if (!session->ssl || SSL_pending(session->ssl) == 0) {
|
||||
if (!ops->has_pending(session)) {
|
||||
int remaining_ms = deadline ? deadline_remaining_ms(deadline) : -1;
|
||||
if (remaining_ms == 0) {
|
||||
log_message(LOG_LEVEL_ERROR, "Receive timeout while reading status");
|
||||
@@ -869,21 +964,11 @@ static bool protocol_read_status_until(ProtocolSession* session, Status* status,
|
||||
if (pfd.revents & (POLLERR | POLLNVAL))
|
||||
return false;
|
||||
}
|
||||
ssize_t bytes_received;
|
||||
if (session->ssl)
|
||||
bytes_received = SSL_read(session->ssl, (char*)&received + got, sizeof(Status) - got);
|
||||
else
|
||||
bytes_received = read(session->read_fd, (char*)&received + got, sizeof(Status) - got);
|
||||
ssize_t bytes_received =
|
||||
ops->recv(session, (char*)&received + got, sizeof(Status) - got, &wait_events);
|
||||
if (bytes_received == PROTOCOL_IO_RETRY)
|
||||
continue;
|
||||
if (bytes_received <= 0) {
|
||||
if (session->ssl) {
|
||||
int ssl_err = SSL_get_error(session->ssl, (int)bytes_received);
|
||||
if (ssl_err == SSL_ERROR_WANT_READ || ssl_err == SSL_ERROR_WANT_WRITE) {
|
||||
wait_events = ssl_err == SSL_ERROR_WANT_WRITE ? POLLOUT : POLLIN;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
if (bytes_received < 0 && errno == EINTR)
|
||||
continue;
|
||||
log_message(LOG_LEVEL_ERROR, "Connection closed while receiving status");
|
||||
return false;
|
||||
}
|
||||
@@ -912,7 +997,7 @@ bool protocol_receive_status_keepalive(ProtocolSession* session, Status* status,
|
||||
while (true) {
|
||||
if (abort_check && abort_check())
|
||||
return false;
|
||||
if (!session->ssl || SSL_pending(session->ssl) == 0) {
|
||||
if (!session->ops || !session->ops->has_pending(session)) {
|
||||
int remaining_ms = deadline_remaining_ms(&deadline);
|
||||
if (remaining_ms <= 0) {
|
||||
log_message(LOG_LEVEL_ERROR, "Receive timeout after %ds", timeout_sec);
|
||||
|
||||
+43
-2
@@ -50,16 +50,48 @@
|
||||
|
||||
typedef struct ssl_st SSL;
|
||||
|
||||
typedef struct ProtocolSession ProtocolSession;
|
||||
|
||||
/*
|
||||
* Transport vtable: the per-session set of I/O primitives the three protocol
|
||||
* loops (send, receive, status-read) dispatch through. The ops are selected
|
||||
* once, when the session is initialized or its SSL is installed, so the loops
|
||||
* never branch on the transport at runtime. A plaintext session uses the
|
||||
* read()/write() ops; a TLS session uses the SSL_read()/SSL_write() ops.
|
||||
*
|
||||
* `send`/`recv` attempt exactly one transfer and return:
|
||||
* > 0 bytes transferred,
|
||||
* PROTOCOL_IO_RETRY no progress; poll on *wait_events and retry,
|
||||
* PROTOCOL_IO_CLOSED peer closed the stream,
|
||||
* PROTOCOL_IO_ERROR fatal transport error.
|
||||
* `has_pending` reports bytes already buffered by the transport (a TLS record
|
||||
* residue); the receive loops skip the poll() gate when it is true.
|
||||
*/
|
||||
typedef struct ProtocolIoOps {
|
||||
ssize_t (*send)(ProtocolSession* session, const void* data, size_t size, short* wait_events);
|
||||
ssize_t (*recv)(ProtocolSession* session, void* data, size_t size, short* wait_events);
|
||||
bool (*has_pending)(const ProtocolSession* session);
|
||||
} ProtocolIoOps;
|
||||
|
||||
/* Negative sentinels returned by ProtocolIoOps.send/recv (see above). */
|
||||
enum {
|
||||
PROTOCOL_IO_RETRY = -1,
|
||||
PROTOCOL_IO_CLOSED = -2,
|
||||
PROTOCOL_IO_ERROR = -3,
|
||||
};
|
||||
|
||||
/*
|
||||
* Explicit owner of protocol I/O. A session does not own the descriptors or
|
||||
* SSL object; it only describes the transport used by a transfer. This makes
|
||||
* it safe to pass the transport to a worker without relying on inherited
|
||||
* thread-local state.
|
||||
*/
|
||||
typedef struct ProtocolSession {
|
||||
struct ProtocolSession {
|
||||
int read_fd;
|
||||
int write_fd;
|
||||
SSL* ssl;
|
||||
/* Transport dispatch selected by protocol_session_init()/set_ssl(). */
|
||||
const ProtocolIoOps* ops;
|
||||
unsigned long long bwlimit;
|
||||
long long bw_tokens;
|
||||
long long bw_last_refill_sec;
|
||||
@@ -75,7 +107,7 @@ typedef struct ProtocolSession {
|
||||
* SO_RCVTIMEO/SO_SNDTIMEO. The server does not propagate a client 0 here: it
|
||||
* installs protocol_server_io_timeout_sec() so its sessions keep a floor. */
|
||||
int io_timeout_sec;
|
||||
} ProtocolSession;
|
||||
};
|
||||
|
||||
typedef int Status;
|
||||
enum NET_STATUS {
|
||||
@@ -216,6 +248,15 @@ void io_set_bwlimit(unsigned long long bytes_per_sec);
|
||||
unsigned long long io_get_bwlimit(void);
|
||||
void io_set_ssl(SSL* ssl);
|
||||
SSL* io_get_ssl(void);
|
||||
/* SSL object of the transport in effect on this thread: the currently bound
|
||||
* session's SSL when a session is bound, otherwise the legacy thread-local
|
||||
* io_ssl. NULL for a plaintext transport. Unlike io_get_ssl(), this resolves
|
||||
* worker threads that bound a TLS session via protocol_session_set_ssl()/
|
||||
* protocol_session_bind() but never called io_set_ssl() themselves (C11
|
||||
* _Thread_local state is not inherited by a new thread). Callers that must
|
||||
* choose a TLS-only code path (e.g. file_send.c's sendfile fallback) must use
|
||||
* this instead of io_get_ssl(). */
|
||||
SSL* protocol_current_ssl(void);
|
||||
|
||||
/* Process-wide wire byte counters. protocol_send_n_data/protocol_receive_n_data
|
||||
* update them; the zero-copy sendfile path reports through
|
||||
|
||||
@@ -146,7 +146,6 @@ class TestTLSBasic:
|
||||
assert not missing, f"Missing files: {missing}"
|
||||
assert not mismatches, f"Mismatched files: {mismatches}"
|
||||
|
||||
@pytest.mark.xfail(reason="TLS multithreading has architectural limitations with per-thread SSL context")
|
||||
def test_tls_with_multithreading(self, certs):
|
||||
"""TLS + multithreading."""
|
||||
clean_dir(DEST_DIR)
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
#include "protocol.h"
|
||||
#include "test_utils.h"
|
||||
#include <errno.h>
|
||||
#include <fcntl.h>
|
||||
#include <limits.h>
|
||||
#include <openssl/ssl.h>
|
||||
#include <poll.h>
|
||||
#include <stdlib.h>
|
||||
#include <string.h>
|
||||
#include <sys/socket.h>
|
||||
#include <time.h>
|
||||
#include <unistd.h>
|
||||
#include <threads.h>
|
||||
@@ -786,6 +790,129 @@ static void test_protocol_throttle_bytes_legacy_same_session() {
|
||||
io_set_fds(-1, -1);
|
||||
}
|
||||
|
||||
/* ------------------------------------------------------------------------- *
|
||||
* Transport-vtable dispatch tests.
|
||||
* ------------------------------------------------------------------------- */
|
||||
|
||||
static int dispatch_send_calls;
|
||||
static int dispatch_recv_calls;
|
||||
|
||||
static ssize_t counting_send(ProtocolSession* session, const void* data, size_t size,
|
||||
short* wait_events) {
|
||||
dispatch_send_calls++;
|
||||
ssize_t written = write(session->write_fd, data, size);
|
||||
if (written < 0)
|
||||
return errno == EINTR ? PROTOCOL_IO_RETRY : PROTOCOL_IO_ERROR;
|
||||
if (written == 0)
|
||||
return PROTOCOL_IO_ERROR;
|
||||
*wait_events = POLLOUT;
|
||||
return written;
|
||||
}
|
||||
|
||||
static ssize_t counting_recv(ProtocolSession* session, void* data, size_t size,
|
||||
short* wait_events) {
|
||||
dispatch_recv_calls++;
|
||||
ssize_t received = read(session->read_fd, data, size);
|
||||
if (received < 0)
|
||||
return errno == EINTR ? PROTOCOL_IO_RETRY : PROTOCOL_IO_ERROR;
|
||||
if (received == 0)
|
||||
return PROTOCOL_IO_CLOSED;
|
||||
*wait_events = POLLIN;
|
||||
return received;
|
||||
}
|
||||
|
||||
static bool counting_has_pending(const ProtocolSession* session) {
|
||||
(void)session;
|
||||
return false;
|
||||
}
|
||||
|
||||
static const ProtocolIoOps counting_ops = {
|
||||
.send = counting_send,
|
||||
.recv = counting_recv,
|
||||
.has_pending = counting_has_pending,
|
||||
};
|
||||
|
||||
/* A plain-TCP socketpair session must route every byte through the ops table:
|
||||
* installing a counting ops wrapper proves the send/receive loops dispatch via
|
||||
* session->ops instead of branching on session->ssl. */
|
||||
static void test_protocol_dispatch_via_ops() {
|
||||
int sv[2];
|
||||
EXPECT_EQ_INT(socketpair(AF_UNIX, SOCK_STREAM, 0, sv), 0);
|
||||
|
||||
ProtocolSession sender;
|
||||
ProtocolSession receiver;
|
||||
protocol_session_init(&sender, sv[0], sv[0]);
|
||||
protocol_session_set_bwlimit(&sender, 0);
|
||||
protocol_session_init(&receiver, sv[1], sv[1]);
|
||||
protocol_session_set_bwlimit(&receiver, 0);
|
||||
EXPECT_NOT_NULL(sender.ops);
|
||||
EXPECT_NOT_NULL(receiver.ops);
|
||||
|
||||
dispatch_send_calls = 0;
|
||||
dispatch_recv_calls = 0;
|
||||
sender.ops = &counting_ops;
|
||||
receiver.ops = &counting_ops;
|
||||
|
||||
const char payload[] = "dispatch-through-vtable";
|
||||
EXPECT_TRUE(protocol_send_n_data(&sender, payload, sizeof(payload)));
|
||||
char received[sizeof(payload)] = {0};
|
||||
EXPECT_TRUE(protocol_receive_n_data(&receiver, received, sizeof(received)));
|
||||
EXPECT_EQ_INT(memcmp(payload, received, sizeof(payload)), 0);
|
||||
EXPECT_TRUE(dispatch_send_calls > 0);
|
||||
EXPECT_TRUE(dispatch_recv_calls > 0);
|
||||
|
||||
close(sv[0]);
|
||||
close(sv[1]);
|
||||
}
|
||||
|
||||
typedef struct {
|
||||
ProtocolSession* session;
|
||||
SSL* expected_ssl;
|
||||
SSL* resolved_ssl;
|
||||
SSL* thread_local_ssl;
|
||||
} SslResolverWorkerArg;
|
||||
|
||||
static int ssl_resolver_worker(void* arg) {
|
||||
SslResolverWorkerArg* worker = arg;
|
||||
protocol_session_bind(worker->session);
|
||||
worker->resolved_ssl = protocol_current_ssl();
|
||||
worker->thread_local_ssl = io_get_ssl();
|
||||
protocol_session_unbind();
|
||||
return thrd_success;
|
||||
}
|
||||
|
||||
/* The worker-thread bug fix: a thread that bound a TLS session but never ran
|
||||
* the handshake has io_ssl == NULL, yet protocol_current_ssl() must return the
|
||||
* session's SSL so callers pick the TLS path. */
|
||||
static void test_protocol_current_ssl_prefers_bound_session() {
|
||||
SSL_CTX* ctx = SSL_CTX_new(TLS_method());
|
||||
EXPECT_NOT_NULL(ctx);
|
||||
SSL* ssl = SSL_new(ctx);
|
||||
EXPECT_NOT_NULL(ssl);
|
||||
|
||||
int sv[2];
|
||||
EXPECT_EQ_INT(socketpair(AF_UNIX, SOCK_STREAM, 0, sv), 0);
|
||||
ProtocolSession session;
|
||||
protocol_session_init(&session, sv[0], sv[0]);
|
||||
protocol_session_set_ssl(&session, ssl);
|
||||
|
||||
/* Clear the calling thread's legacy SSL: only the bound session carries it. */
|
||||
io_set_fds(-1, -1);
|
||||
|
||||
SslResolverWorkerArg arg = {
|
||||
.session = &session, .expected_ssl = ssl, .resolved_ssl = NULL, .thread_local_ssl = ssl};
|
||||
thrd_t worker;
|
||||
EXPECT_EQ_INT(thrd_create(&worker, ssl_resolver_worker, &arg), thrd_success);
|
||||
EXPECT_EQ_INT(thrd_join(worker, NULL), thrd_success);
|
||||
EXPECT_TRUE(arg.resolved_ssl == ssl);
|
||||
EXPECT_NULL(arg.thread_local_ssl);
|
||||
|
||||
close(sv[0]);
|
||||
close(sv[1]);
|
||||
SSL_free(ssl);
|
||||
SSL_CTX_free(ctx);
|
||||
}
|
||||
|
||||
void test_protocol() {
|
||||
test_send_receive_n_data();
|
||||
test_send_receive_n_data_zero();
|
||||
@@ -819,4 +946,6 @@ void test_protocol() {
|
||||
test_protocol_throttle_bytes_paces();
|
||||
test_protocol_throttle_bytes_unlimited();
|
||||
test_protocol_throttle_bytes_legacy_same_session();
|
||||
test_protocol_dispatch_via_ops();
|
||||
test_protocol_current_ssl_prefers_bound_session();
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user