From ec6692ac411cc28a0e744139b32ec9cbd2e3843c Mon Sep 17 00:00:00 2001 From: TapTap Date: Wed, 23 Sep 2026 00:18:36 +0200 Subject: [PATCH] protocol: add transport I/O vtable over TCP/TLS primitives Introduce ProtocolIoOps (send/recv/has_pending), selected once by protocol_session_init() and protocol_session_set_ssl(), and dispatch the send, receive and status-read loops through session->ops instead of branching on session->ssl at runtime. Each op performs one transfer attempt and classifies the result (PROTOCOL_IO_RETRY/CLOSED/ERROR), preserving the WANT_READ/WANT_WRITE wait_events switching, the SSL_ERROR_SYSCALL/EINTR retry, the SSL_pending poll gating and the deadline handling. The raw read()/write() fallback lives in the plaintext ops. Add unit tests: a socketpair session with a counting ops wrapper proving the loops dispatch through the vtable, and a worker-thread test that protocol_current_ssl() resolves the bound session's SSL when io_ssl is NULL. --- src/shared/protocol.c | 228 ++++++++++++++++++++++++++++-------------- src/shared/protocol.h | 36 ++++++- tests/test_protocol.c | 129 ++++++++++++++++++++++++ 3 files changed, 314 insertions(+), 79 deletions(-) diff --git a/src/shared/protocol.c b/src/shared/protocol.c index eba3a61..69dec04 100644 --- a/src/shared/protocol.c +++ b/src/shared/protocol.c @@ -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) - session->ssl = ssl; + 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) { @@ -309,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; } @@ -346,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; @@ -373,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); @@ -435,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); @@ -460,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"); + 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); @@ -856,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"); @@ -880,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; } @@ -923,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); diff --git a/src/shared/protocol.h b/src/shared/protocol.h index ff4ec0d..bd6c752 100644 --- a/src/shared/protocol.h +++ b/src/shared/protocol.h @@ -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 { diff --git a/tests/test_protocol.c b/tests/test_protocol.c index 3992256..365c82d 100644 --- a/tests/test_protocol.c +++ b/tests/test_protocol.c @@ -1,9 +1,13 @@ #include "protocol.h" #include "test_utils.h" +#include #include #include +#include +#include #include #include +#include #include #include #include @@ -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(); }