diff --git a/.github/workflows/cmake.yml b/.github/workflows/cmake.yml index d218f15..d5c4386 100644 --- a/.github/workflows/cmake.yml +++ b/.github/workflows/cmake.yml @@ -44,6 +44,9 @@ jobs: - name: Test run: make -C ${{github.workspace}}/ test ARGS="--verbose" + - name: Test Linux interface binding + run: sudo env HDP_RUN_PRIVILEGED_TEST=1 ctest --test-dir "${{github.workspace}}" -R '^outbound_interface_linux$' --output-on-failure --no-tests=error + - uses: actions/upload-artifact@v7 if: ${{ success() || failure() }} with: diff --git a/CMakeLists.txt b/CMakeLists.txt index c525b31..69c2201 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -205,6 +205,7 @@ if(BUILD_TESTING) add_executable(dns_poller_watchers_test tests/unit/test_dns_poller.c src/logging.c + src/outbound_interface.c src/ring_buffer.c) set_property(SOURCE tests/unit/test_dns_poller.c APPEND PROPERTY COMPILE_DEFINITIONS __FILENAME__="dns_poller_test") @@ -225,6 +226,71 @@ if(BUILD_TESTING) add_test(NAME dns_truncate COMMAND dns_truncate_test) set_tests_properties(dns_truncate PROPERTIES LABELS unit) + add_executable(dns_listener_tcp_test + tests/unit/test_dns_listener_tcp.c + src/logging.c + src/ring_buffer.c) + set_property(SOURCE tests/unit/test_dns_listener_tcp.c APPEND + PROPERTY COMPILE_DEFINITIONS __FILENAME__="dns_listener_tcp_test") + target_link_libraries(dns_listener_tcp_test ev) + set_property(TARGET dns_listener_tcp_test PROPERTY C_STANDARD 11) + add_test(NAME dns_listener_tcp COMMAND dns_listener_tcp_test) + set_tests_properties(dns_listener_tcp PROPERTIES LABELS unit TIMEOUT 20) + + add_executable(https_client_limits_test + tests/unit/test_https_client_limits.c + src/logging.c + src/outbound_interface.c + src/ring_buffer.c + src/stat.c) + set_property(SOURCE tests/unit/test_https_client_limits.c APPEND + PROPERTY COMPILE_DEFINITIONS __FILENAME__="https_client_limits_test") + target_link_libraries(https_client_limits_test curl ev m) + set_property(TARGET https_client_limits_test PROPERTY C_STANDARD 11) + add_test(NAME https_client_limits COMMAND https_client_limits_test) + set_tests_properties(https_client_limits PROPERTIES LABELS unit) + + add_executable(doh_proxy_resolver_test + tests/unit/test_doh_proxy_resolver.c + src/logging.c + src/ring_buffer.c) + set_property(SOURCE tests/unit/test_doh_proxy_resolver.c APPEND + PROPERTY COMPILE_DEFINITIONS __FILENAME__="doh_proxy_resolver_test") + target_link_libraries(doh_proxy_resolver_test curl ev) + set_property(TARGET doh_proxy_resolver_test PROPERTY C_STANDARD 11) + add_test(NAME doh_proxy_resolver COMMAND doh_proxy_resolver_test) + set_tests_properties(doh_proxy_resolver PROPERTIES LABELS unit) + + add_executable(outbound_interface_test + tests/unit/test_outbound_interface.c) + set_property(SOURCE tests/unit/test_outbound_interface.c APPEND + PROPERTY COMPILE_DEFINITIONS __FILENAME__="outbound_interface_test") + set_property(TARGET outbound_interface_test PROPERTY C_STANDARD 11) + add_test(NAME outbound_interface COMMAND outbound_interface_test) + set_tests_properties(outbound_interface PROPERTIES LABELS unit) + + add_executable(options_test + tests/unit/test_options.c) + set_property(SOURCE tests/unit/test_options.c APPEND + PROPERTY COMPILE_DEFINITIONS __FILENAME__="options_test") + set_property(TARGET options_test PROPERTY C_STANDARD 11) + add_test(NAME options COMMAND options_test) + set_tests_properties(options PROPERTIES LABELS unit) + + if(CMAKE_SYSTEM_NAME STREQUAL "Linux") + add_executable(outbound_interface_linux_test + tests/linux/test_outbound_interface_linux.c + src/outbound_interface.c) + target_include_directories(outbound_interface_linux_test PRIVATE tests/unit) + set_property(SOURCE tests/linux/test_outbound_interface_linux.c APPEND + PROPERTY COMPILE_DEFINITIONS __FILENAME__="outbound_interface_linux_test") + set_property(TARGET outbound_interface_linux_test PROPERTY C_STANDARD 11) + add_test(NAME outbound_interface_linux COMMAND outbound_interface_linux_test) + set_tests_properties(outbound_interface_linux PROPERTIES + SKIP_RETURN_CODE 77 + LABELS linux-privileged) + endif() + find_program(VALGRIND_EXE NAMES valgrind) if(VALGRIND_EXE) add_test(NAME dns_truncate_valgrind diff --git a/README.md b/README.md index fc61f0e..312a45c 100644 --- a/README.md +++ b/README.md @@ -162,6 +162,7 @@ Just run it as a daemon and point traffic at it. Commandline flags are: Usage: ./https_dns_proxy [-a ] [-p ] [-T ] [-b ] [-i ] [-4] [-r ] [-t ] [-S ] [-x] [-q] [-C ] [-c ] + [--outbound-interface ] [-d] [-u ] [-g ] [-v]+ [-l ] [-s ] [-F ] [-V] [-h] @@ -189,6 +190,9 @@ Usage: ./https_dns_proxy [-a ] [-p ] [-T ] [-p ] [-T ] [-p ] [-T +#include #include // A DNS listener accepts requests on some transport (UDP, TCP, ...) and routes @@ -23,16 +24,20 @@ typedef enum { // Invoked once per fully-received DNS request. `dns_req` is heap-allocated // and ownership transfers to the callee. `listener` is a back-pointer the // callee uses later to deliver the matching response. +// connection_id identifies the TCP connection; UDP uses zero. typedef void (*dns_request_fn)(void *ctx, dns_listener_t *listener, - struct sockaddr *raddr, + struct sockaddr *raddr, uint64_t connection_id, char *dns_req, size_t dns_req_len); struct dns_listener { // Send `dns_resp` to `raddr`. UDP listeners may EDNS-truncate the response // in place using `dns_req`; TCP listeners ignore the request bytes. - void (*respond)(dns_listener_t *self, struct sockaddr *raddr, + void (*respond)(dns_listener_t *self, struct sockaddr *raddr, uint64_t connection_id, const char *dns_req, size_t dns_req_len, char *dns_resp, size_t dns_resp_len); + // Release any transport-level capacity reserved for this request. The + // proxy core calls this exactly once whether the request succeeds or fails. + void (*request_complete)(dns_listener_t *self, struct sockaddr *raddr, uint64_t connection_id); // Stop accepting new requests. Existing per-client state (TCP) is retained // so any in-flight DoH responses can still be delivered during graceful // drain. diff --git a/src/dns_listener_tcp.c b/src/dns_listener_tcp.c index 3b6d101..d6a3785 100644 --- a/src/dns_listener_tcp.c +++ b/src/dns_listener_tcp.c @@ -31,13 +31,23 @@ enum { LISTEN_BACKLOG = 5, IDLE_TIMEOUT_S = 120, // "two minutes" according to RFC1035 4.2.2 - RESPONSE_SEND_ATTEMPTS = 50, // 0.025 sec max wait - RESPONSE_SEND_DELAY_US = 500, // 0.0005 sec - TCP_DNS_MAX_PAYLOAD = UINT16_MAX - sizeof(uint16_t), // Max after 2-byte length prefix + SHUTDOWN_TIMEOUT_S = 15, + TCP_DNS_MAX_PAYLOAD = UINT16_MAX, + TCP_DNS_MAX_FRAME = TCP_DNS_MAX_PAYLOAD + sizeof(uint16_t), + TCP_CLIENT_REQUEST_LIMIT = 16, + TCP_CLIENT_OUTPUT_LIMIT = 2 * UINT16_MAX, + TCP_WRITE_BUDGET = UINT16_MAX, }; typedef struct dns_listener_tcp_s dns_listener_tcp_t; +struct tcp_response_s { + struct tcp_response_s *next; + size_t length; + size_t sent; + char data[]; +}; + struct tcp_client_s { dns_listener_tcp_t * d; @@ -50,10 +60,17 @@ struct tcp_client_s { char * input_buffer; uint32_t input_buffer_size; uint32_t input_buffer_used; + uint16_t pending_requests; + int read_eof; ev_io read_watcher; + ev_io write_watcher; ev_timer timer_watcher; + struct tcp_response_s *output_head; + struct tcp_response_s *output_tail; + size_t output_bytes; + struct tcp_client_s * next; }; @@ -72,24 +89,47 @@ struct dns_listener_tcp_s { uint64_t client_id; uint16_t client_count; uint16_t client_limit; + int stopping; struct tcp_client_s * clients; }; +static struct tcp_client_s *find_client(dns_listener_tcp_t *d, + struct sockaddr *raddr, uint64_t connection_id) { + for (struct tcp_client_s *client = d->clients; client != NULL; client = client->next) { + if (client->id == connection_id && + memcmp(raddr, &client->raddr, client->addr_len) == 0) { + return client; + } + } + return NULL; +} + +static void free_output_queue(struct tcp_client_s *client) { + while (client->output_head != NULL) { + struct tcp_response_s *response = client->output_head; + client->output_head = response->next; + free(response); + } + client->output_tail = NULL; + client->output_bytes = 0; +} static void remove_client(struct tcp_client_s * client) { dns_listener_tcp_t *d = client->d; DLOG_CLIENT("Removing client, socket %d", client->sock); - if (d->client_count == d->client_limit) { + if (!d->stopping && d->client_count == d->client_limit) { ev_io_start(d->loop, &d->accept_watcher); // continue accepting new client connections } d->client_count--; ev_io_stop(d->loop, &client->read_watcher); + ev_io_stop(d->loop, &client->write_watcher); ev_timer_stop(d->loop, &client->timer_watcher); free(client->input_buffer); + free_output_queue(client); close(client->sock); @@ -110,12 +150,24 @@ static void remove_client(struct tcp_client_s * client) { static int get_dns_request(struct tcp_client_s *client, char ** dns_req, uint16_t * req_size) { + if (client->input_buffer_used < sizeof(uint16_t)) { + return 0; // Partial length prefix + } + // check if whole request is available - *req_size = ntohs(*((uint16_t*)client->input_buffer)); - uint16_t data_size = sizeof(uint16_t) + *req_size; + uint16_t network_req_size = 0; + memcpy(&network_req_size, client->input_buffer, sizeof(network_req_size)); + *req_size = ntohs(network_req_size); + const size_t data_size = sizeof(uint16_t) + (size_t)*req_size; if (data_size > client->input_buffer_used) { return 0; // Partial request } + if (*req_size < DNS_HEADER_LENGTH) { + client->input_buffer_used -= (uint32_t)data_size; + memmove(client->input_buffer, client->input_buffer + data_size, client->input_buffer_used); + *dns_req = NULL; + return 1; + } // copy whole request *dns_req = (char *)malloc(*req_size); // freed when DoH request completes if (*dns_req == NULL) { @@ -123,84 +175,142 @@ static int get_dns_request(struct tcp_client_s *client, } memcpy(*dns_req, client->input_buffer + sizeof(uint16_t), *req_size); // move down data of next request(s) if any - client->input_buffer_used -= data_size; + client->input_buffer_used -= (uint32_t)data_size; memmove(client->input_buffer, client->input_buffer + data_size, client->input_buffer_used); return 1; } +static int dispatch_client_requests(struct tcp_client_s *client) { + dns_listener_tcp_t *d = client->d; + char *dns_req = NULL; + uint16_t req_size = 0; + int request_received = 0; + while (client->pending_requests < TCP_CLIENT_REQUEST_LIMIT && + get_dns_request(client, &dns_req, &req_size)) { + if (req_size < DNS_HEADER_LENGTH) { + WLOG_CLIENT("Malformed request received, too short: %u, dropping client", req_size); + free(dns_req); + remove_client(client); + return -1; + } + client->pending_requests++; + if (client->pending_requests == TCP_CLIENT_REQUEST_LIMIT) { + ev_io_stop(d->loop, &client->read_watcher); + } + d->cb(d->cb_data, &d->base, (struct sockaddr *)&client->raddr, + client->id, dns_req, req_size); + request_received = 1; + } + if (request_received) { + ev_timer_again(d->loop, &client->timer_watcher); + } + return client->pending_requests == TCP_CLIENT_REQUEST_LIMIT; +} + static void read_cb(struct ev_loop __attribute__((unused)) *loop, ev_io *w, int __attribute__((unused)) revents) { struct tcp_client_s *client = (struct tcp_client_s *)w->data; - dns_listener_tcp_t *d = client->d; - - // Receive data - char buf[DNS_REQUEST_BUFFER_SIZE]; // if there would be more data, callback will be called again - ssize_t len = recv(w->fd, buf, DNS_REQUEST_BUFFER_SIZE, 0); - if (len <= 0) { - if (len == 0 || errno == ECONNRESET) { - DLOG_CLIENT("TCP client closed connection"); - } else if (errno == EAGAIN || errno == EWOULDBLOCK) { - return; - } else { - WLOG_CLIENT("Read error: %s (%d), dropping client", strerror(errno), errno); - } - remove_client(client); + if (client->read_eof) { + return; + } + if (dispatch_client_requests(client) != 0) { return; } - // Append data into input buffer - // Check for integer overflow and maximum message size - if (len > UINT16_MAX || client->input_buffer_used > UINT16_MAX - (uint32_t)len) { + // Receive data + char buf[DNS_REQUEST_BUFFER_SIZE]; // if there would be more data, callback will be called again + const size_t remaining = TCP_DNS_MAX_FRAME - client->input_buffer_used; + if (remaining == 0) { WLOG_CLIENT("Request too large, dropping client"); remove_client(client); return; } - const uint32_t free_space = client->input_buffer_size - client->input_buffer_used; - const uint32_t needed_space = client->input_buffer_used + (uint32_t)len; - // Limit buffer size to prevent memory exhaustion attacks - if (needed_space > TCP_DNS_MAX_PAYLOAD) { - WLOG_CLIENT("Request too large, dropping client"); - remove_client(client); + const size_t read_size = remaining < sizeof(buf) ? remaining : sizeof(buf); + ssize_t len = recv(w->fd, buf, read_size, 0); + if (len == 0) { + client->read_eof = 1; + ev_io_stop(client->d->loop, &client->read_watcher); + if (client->pending_requests == 0 && client->output_head == NULL) { + remove_client(client); + } return; } - DLOG_CLIENT("Received %d byte, free: %u", len, free_space); - if (free_space < len) { - for (client->input_buffer_size = 64; // lower value does not make much sense - client->input_buffer_size < needed_space; - client->input_buffer_size *= 2) { - if (client->input_buffer_size > TCP_DNS_MAX_PAYLOAD) { - FLOG_CLIENT("Unrealistic input buffer size: %u", client->input_buffer_size); - } + if (len < 0) { + if (errno == ECONNRESET) { + DLOG_CLIENT("TCP client closed connection"); + } else if (errno == EAGAIN || errno == EWOULDBLOCK) { + len = 0; + } else { + WLOG_CLIENT("Read error: %s (%d), dropping client", strerror(errno), errno); } - DLOG_CLIENT("Resize input buffer to %u", client->input_buffer_size); - client->input_buffer = (char *) realloc((void*) client->input_buffer, // NOLINT(bugprone-suspicious-realloc-usage) if realloc fails, program stops - client->input_buffer_size); - if (client->input_buffer == NULL) { - FLOG_CLIENT("Out of mem"); + if (len < 0) { + remove_client(client); + return; } } - memcpy(client->input_buffer + client->input_buffer_used, buf, (size_t)len); - client->input_buffer_used = needed_space; - // Split requests - char *dns_req = NULL; - uint16_t req_size = 0; - uint8_t request_received = 0; - while (get_dns_request(client, &dns_req, &req_size)) { - if (req_size < DNS_HEADER_LENGTH) { - WLOG_CLIENT("Malformed request received, too short: %u, dropping client", req_size); - free(dns_req); + if (len > 0) { + // Append data into input buffer + // Check for integer overflow and maximum message size + if ((size_t)len > TCP_DNS_MAX_FRAME - client->input_buffer_used) { + WLOG_CLIENT("Request too large, dropping client"); + remove_client(client); + return; + } + const uint32_t free_space = client->input_buffer_size - client->input_buffer_used; + const uint32_t needed_space = client->input_buffer_used + (uint32_t)len; + // Limit buffer size to prevent memory exhaustion attacks + if (needed_space > TCP_DNS_MAX_FRAME) { + WLOG_CLIENT("Request too large, dropping client"); remove_client(client); return; } + DLOG_CLIENT("Received %d byte, free: %u", len, free_space); + if (free_space < len) { + for (client->input_buffer_size = 64; // lower value does not make much sense + client->input_buffer_size < needed_space; + client->input_buffer_size *= 2) { + } + if (client->input_buffer_size > TCP_DNS_MAX_FRAME) { + client->input_buffer_size = TCP_DNS_MAX_FRAME; + } + DLOG_CLIENT("Resize input buffer to %u", client->input_buffer_size); + client->input_buffer = (char *) realloc((void*) client->input_buffer, // NOLINT(bugprone-suspicious-realloc-usage) if realloc fails, program stops + client->input_buffer_size); + if (client->input_buffer == NULL) { + FLOG_CLIENT("Out of mem"); + } + } + memcpy(client->input_buffer + client->input_buffer_used, buf, (size_t)len); + client->input_buffer_used = needed_space; + } - DLOG_CLIENT("Requested %04hX", ntohs(*((uint16_t*)dns_req))); - d->cb(d->cb_data, &d->base, (struct sockaddr*)&client->raddr, dns_req, req_size); - request_received = 1; + (void)dispatch_client_requests(client); +} + +static void tcp_request_complete(dns_listener_t *self, struct sockaddr *raddr, + uint64_t connection_id) { + dns_listener_tcp_t *d = (dns_listener_tcp_t *)self; + struct tcp_client_s *client = find_client(d, raddr, connection_id); + if (client == NULL) { + return; + } + if (client->pending_requests == 0) { + WLOG_CLIENT("Completed request with no request pending"); + return; } - if (request_received) { - ev_timer_again(d->loop, &client->timer_watcher); + const int resume_reads = client->pending_requests == TCP_CLIENT_REQUEST_LIMIT; + client->pending_requests--; + if (client->read_eof && client->pending_requests == 0 && client->output_head == NULL) { + remove_client(client); + return; + } + if (resume_reads && !client->read_eof) { + ev_io_start(d->loop, &client->read_watcher); + if (client->input_buffer_used > 0) { + ev_feed_event(d->loop, &client->read_watcher, EV_READ); + } } } @@ -211,6 +321,62 @@ static void timer_cb(struct ev_loop __attribute__((unused)) *loop, remove_client(client); } +static int flush_client_output(struct tcp_client_s *client, int *made_progress) { + size_t budget = TCP_WRITE_BUDGET; + *made_progress = 0; + + while (client->output_head != NULL && budget > 0) { + struct tcp_response_s *response = client->output_head; + const size_t remaining = response->length - response->sent; + const size_t send_size = remaining < budget ? remaining : budget; + const ssize_t sent = send(client->sock, response->data + response->sent, + send_size, MSG_NOSIGNAL); + if (sent > 0) { + const size_t sent_size = (size_t)sent; + response->sent += sent_size; + client->output_bytes -= sent_size; + budget -= sent_size; + *made_progress = 1; + if (response->sent == response->length) { + client->output_head = response->next; + if (client->output_head == NULL) { + client->output_tail = NULL; + } + free(response); + } + continue; + } + if (sent < 0 && (errno == EAGAIN || errno == EWOULDBLOCK)) { + break; + } + WLOG_CLIENT("Send error: %s (%d), dropping client", strerror(errno), errno); + return -1; + } + + if (client->output_head == NULL) { + ev_io_stop(client->d->loop, &client->write_watcher); + } else { + ev_io_start(client->d->loop, &client->write_watcher); + } + return 0; +} + +static void write_cb(struct ev_loop __attribute__((unused)) *loop, + ev_io *w, int __attribute__((unused)) revents) { + struct tcp_client_s *client = (struct tcp_client_s *)w->data; + int made_progress = 0; + if (flush_client_output(client, &made_progress) != 0) { + remove_client(client); + return; + } + if (made_progress && !client->d->stopping) { + ev_timer_again(client->d->loop, &client->timer_watcher); + } + if (client->read_eof && client->pending_requests == 0 && client->output_head == NULL) { + remove_client(client); + } +} + static void accept_cb(struct ev_loop __attribute__((unused)) *loop, ev_io *w, int __attribute__((unused)) revents) { dns_listener_tcp_t *d = (dns_listener_tcp_t *)w->data; @@ -265,6 +431,9 @@ static void accept_cb(struct ev_loop __attribute__((unused)) *loop, client->read_watcher.data = client; ev_io_start(d->loop, &client->read_watcher); + ev_io_init(&client->write_watcher, write_cb, client->sock, EV_WRITE); + client->write_watcher.data = client; + ev_init(&client->timer_watcher, timer_cb); client->timer_watcher.repeat = IDLE_TIMEOUT_S; client->timer_watcher.data = client; @@ -324,13 +493,12 @@ static int get_tcp_listen_sock(struct addrinfo *listen_addrinfo) { } static void tcp_respond(dns_listener_t *self, struct sockaddr *raddr, + uint64_t connection_id, const char __attribute__((unused)) *dns_req, size_t __attribute__((unused)) dns_req_len, char *resp, size_t resp_len) { dns_listener_tcp_t *d = (dns_listener_tcp_t *)self; - // Limit response size to prevent overflow when accounting for the 2-byte - // length prefix. The total on-wire size would be resp_len + sizeof(uint16_t). if (resp_len < DNS_HEADER_LENGTH || resp_len > TCP_DNS_MAX_PAYLOAD) { WLOG("Malformed response received, invalid length: %zu", resp_len); return; @@ -338,71 +506,73 @@ static void tcp_respond(dns_listener_t *self, struct sockaddr *raddr, const uint16_t response_id = ntohs(*((uint16_t*)resp)); // find client data - struct tcp_client_s *client = NULL; - for (struct tcp_client_s * cur = d->clients; cur != NULL; cur = cur->next) { - if (memcmp(raddr, &(cur->raddr), cur->addr_len) == 0) { - client = cur; - break; - } - } + struct tcp_client_s *client = find_client(d, raddr, connection_id); if (client == NULL) { WLOG("Could not find client, can not send DNS response: %04hX", response_id); return; } - // NOTE: Single-threaded libev event loop ensures no TOCTOU race here. - // No other callbacks can execute while this function runs, and usleep() - // below is a blocking syscall (not an event loop yield). If remove_client() - // is called due to send errors, the function returns immediately. - - DLOG_CLIENT("Sending %zu bytes", resp_len); - - // send length of response - uint16_t resp_size = htons((uint16_t)resp_len); - ssize_t len = send(client->sock, &resp_size, sizeof(uint16_t), MSG_MORE | MSG_NOSIGNAL); - if (len != sizeof(uint16_t)) { - WLOG_CLIENT("Send error: %s (%d), len: %d, dropping client", strerror(errno), errno, len); + const size_t frame_len = sizeof(uint16_t) + resp_len; + if (frame_len > TCP_CLIENT_OUTPUT_LIMIT - client->output_bytes) { + WLOG_CLIENT("Output queue limit exceeded, dropping client"); remove_client(client); return; } - // send the response - ssize_t sent = 0; - int attempts = 0; - for (; attempts < RESPONSE_SEND_ATTEMPTS; ++attempts) - { - len = send(client->sock, resp + sent, resp_len - (size_t)sent, MSG_NOSIGNAL); - if (len > 0) { - sent += len; - if (sent == (ssize_t)resp_len) { - break; - } - } else if (len < 0) { - if (errno != EAGAIN && errno != EWOULDBLOCK) { - WLOG_CLIENT("Send error: %s (%d), dropping client", strerror(errno), errno); - remove_client(client); - return; - } - } - usleep(RESPONSE_SEND_DELAY_US); + struct tcp_response_s *response = + (struct tcp_response_s *)malloc(sizeof(*response) + frame_len); + if (response == NULL) { + FLOG_CLIENT("Out of mem"); + } + response->next = NULL; + response->length = frame_len; + response->sent = 0; + const uint16_t response_size = htons((uint16_t)resp_len); + memcpy(response->data, &response_size, sizeof(response_size)); + memcpy(response->data + sizeof(response_size), resp, resp_len); + + if (client->output_tail != NULL) { + client->output_tail->next = response; + } else { + client->output_head = response; } - if (sent != (ssize_t)resp_len) { - WLOG_CLIENT("Send timeout after %d attempts, sent %zd/%zu bytes, dropping client", - attempts, sent, resp_len); + client->output_tail = response; + client->output_bytes += frame_len; + + DLOG_CLIENT("Queued %zu bytes, %zu bytes pending", resp_len, client->output_bytes); + int made_progress = 0; + if (flush_client_output(client, &made_progress) != 0) { remove_client(client); return; } - DLOG_CLIENT("Responded %04hX", response_id); - - ev_timer_again(d->loop, &client->timer_watcher); + if (made_progress) { + DLOG_CLIENT("Responded %04hX", response_id); + if (!d->stopping) { + ev_timer_again(d->loop, &client->timer_watcher); + } + } } static void tcp_stop(dns_listener_t *self) { dns_listener_tcp_t *d = (dns_listener_tcp_t *)self; - while (d->clients) { - remove_client(d->clients); //NOLINT(clang-analyzer-unix.Malloc) false use after free detection + if (d->stopping) { + return; } + d->stopping = 1; ev_io_stop(d->loop, &d->accept_watcher); + for (struct tcp_client_s *client = d->clients; client != NULL;) { + struct tcp_client_s *next = client->next; + client->read_eof = 1; + ev_io_stop(d->loop, &client->read_watcher); + if (client->pending_requests == 0 && client->output_head == NULL) { + remove_client(client); + } else { + ev_timer_stop(d->loop, &client->timer_watcher); + ev_timer_set(&client->timer_watcher, SHUTDOWN_TIMEOUT_S, 0); + ev_timer_start(d->loop, &client->timer_watcher); + } + client = next; + } } static void tcp_destroy(dns_listener_t *self) { @@ -420,6 +590,7 @@ dns_listener_t * dns_tcp_listener_create(struct ev_loop *loop, FLOG("Out of mem"); } d->base.respond = tcp_respond; + d->base.request_complete = tcp_request_complete; d->base.stop = tcp_stop; d->base.destroy = tcp_destroy; d->base.transport = DNS_TRANSPORT_TCP; diff --git a/src/dns_listener_udp.c b/src/dns_listener_udp.c index 96cef08..49dad2e 100644 --- a/src/dns_listener_udp.c +++ b/src/dns_listener_udp.c @@ -82,10 +82,11 @@ static void watcher_cb(struct ev_loop __attribute__((unused)) *loop, } memcpy(dns_req, tmp_buf, (size_t)len); - d->cb(d->cb_data, &d->base, (struct sockaddr*)&tmp_raddr, dns_req, (size_t)len); + d->cb(d->cb_data, &d->base, (struct sockaddr*)&tmp_raddr, 0, dns_req, (size_t)len); } static void udp_respond(dns_listener_t *self, struct sockaddr *raddr, + uint64_t __attribute__((unused)) connection_id, const char *dns_req, size_t dns_req_len, char *dns_resp, size_t dns_resp_len) { dns_listener_udp_t *d = (dns_listener_udp_t *)self; @@ -102,6 +103,11 @@ static void udp_respond(dns_listener_t *self, struct sockaddr *raddr, } } +static void udp_request_complete(dns_listener_t __attribute__((unused)) *self, + struct sockaddr __attribute__((unused)) *raddr, + uint64_t __attribute__((unused)) connection_id) { +} + static void udp_stop(dns_listener_t *self) { dns_listener_udp_t *d = (dns_listener_udp_t *)self; ev_io_stop(d->loop, &d->watcher); @@ -121,6 +127,7 @@ dns_listener_t * dns_udp_listener_create(struct ev_loop *loop, FLOG("Out of mem"); } d->base.respond = udp_respond; + d->base.request_complete = udp_request_complete; d->base.stop = udp_stop; d->base.destroy = udp_destroy; d->base.transport = DNS_TRANSPORT_UDP; diff --git a/src/dns_poller.c b/src/dns_poller.c index b594288..e5c92da 100644 --- a/src/dns_poller.c +++ b/src/dns_poller.c @@ -1,9 +1,11 @@ #include +#include #include #include #include "dns_poller.h" #include "logging.h" +#include "outbound_interface.h" static void sock_cb(struct ev_loop __attribute__((unused)) *loop, ev_io *w, int revents) { @@ -166,6 +168,19 @@ static void set_bootstrap_source_addr(ares_channel channel, } } +static int configure_bootstrap_socket(ares_socket_t socket_fd, + int __attribute__((unused)) type, + void *userdata) { + dns_poller_t *d = (dns_poller_t *)userdata; + if (outbound_interface_bind_socket(socket_fd, d->outbound_interface) != 0) { + const int bind_errno = errno; + ELOG("Failed to bind bootstrap DNS socket to interface '%s': %s (%d)", + d->outbound_interface, strerror(bind_errno), bind_errno); + return -1; + } + return ARES_SUCCESS; +} + static ev_tstamp get_timeout(dns_poller_t *d) { static struct timeval max_tv = {.tv_sec = 5, .tv_usec = 0}; @@ -219,6 +234,7 @@ void dns_poller_init(dns_poller_t *d, struct ev_loop *loop, const char *bootstrap_dns, int bootstrap_dns_polling_interval, const char *source_addr, + const char *outbound_interface, const char *hostname, int family, dns_poller_cb cb, void *cb_data) { d->io_events = NULL; @@ -250,8 +266,19 @@ void dns_poller_init(dns_poller_t *d, struct ev_loop *loop, if (d->hostname == NULL) { FLOG("Out of mem"); } + d->outbound_interface = NULL; + if (outbound_interface != NULL) { + d->outbound_interface = strdup(outbound_interface); + if (d->outbound_interface == NULL) { + FLOG("Out of mem"); + } + } d->family = family; set_bootstrap_source_addr(d->ares, source_addr, family); + if (d->outbound_interface != NULL) { + ares_set_socket_configure_callback(d->ares, configure_bootstrap_socket, + d); + } d->cb = cb; d->polling_interval = bootstrap_dns_polling_interval; d->request_ongoing = 0; @@ -272,4 +299,5 @@ void dns_poller_cleanup(dns_poller_t *d) { free(event); } free(d->hostname); + free(d->outbound_interface); } diff --git a/src/dns_poller.h b/src/dns_poller.h index 2fa1c62..f1d1fa9 100644 --- a/src/dns_poller.h +++ b/src/dns_poller.h @@ -27,6 +27,7 @@ typedef struct { ares_channel ares; struct ev_loop *loop; char *hostname; // owned; strdup'd in init, freed in cleanup + char *outbound_interface; // owned; optional interface used by socket callback int family; // AF_UNSPEC for IPv4 or IPv6, AF_INET for IPv4 only. dns_poller_cb cb; int polling_interval; @@ -41,14 +42,17 @@ typedef struct { // provided ev_loop. `bootstrap_dns` is a comma-separated list of DNS servers to // use for the lookup `hostname` every `interval_seconds`. For each successful // lookup, `cb` will be called with the resolved address. -// `source_addr` optionally binds bootstrap DNS lookups to a specific IP. +// `source_addr` optionally binds bootstrap DNS lookups to a specific IP and +// `outbound_interface` optionally binds them to a Linux network interface. // `family` should be AF_INET for IPv4 or AF_UNSPEC for both IPv4 and IPv6. // -// Note: hostname is copied; the caller's buffer need not outlive this call. +// Note: hostname and outbound_interface are copied; the caller's buffers need +// not outlive this call. void dns_poller_init(dns_poller_t *d, struct ev_loop *loop, const char *bootstrap_dns, int bootstrap_dns_polling_interval, const char *source_addr, + const char *outbound_interface, const char *hostname, int family, dns_poller_cb cb, void *cb_data); diff --git a/src/doh_proxy.c b/src/doh_proxy.c index 3b485eb..b955486 100644 --- a/src/doh_proxy.c +++ b/src/doh_proxy.c @@ -34,6 +34,7 @@ typedef struct { dns_listener_t *listener; ev_tstamp start_tstamp; struct sockaddr_storage raddr; + uint64_t connection_id; char *dns_req; size_t dns_req_len; } doh_request_t; @@ -73,11 +74,25 @@ void doh_proxy_set_resolv(doh_proxy_t *p, const char *buf) { } } -// Returns 1 if `addr_list` is a (possibly equal, possibly proper) subset of -// `full_list`, where both are comma-separated IP literals. Used to decide -// whether a fresh poll result actually changed anything; if every IP in the -// new list is already in the old list, we skip the curl reset. -static int addr_list_reduced(const char* full_list, const char* list) { +static int addr_list_contains(const char *list, const char *address, + size_t address_len) { + const char *pos = list; + const char *end = list + strlen(list); + while (pos < end) { + const char *comma = strchr(pos, ','); + const size_t token_len = (size_t)(comma ? comma - pos : end - pos); + if (token_len == address_len && memcmp(pos, address, address_len) == 0) { + return 1; + } + pos += token_len + 1; + } + return 0; +} + +// Returns 1 if `list` contains an address missing from `full_list`. +// Both lists contain comma-separated IP literals. Checking both directions +// detects additions and removals while ignoring address order. +static int addr_list_reduced(const char *full_list, const char *list) { const char *pos = list; const char *end = list + strlen(list); while (pos < end) { @@ -91,10 +106,7 @@ static int addr_list_reduced(const char* full_list, const char* list) { strncpy(current, pos, ip_len); current[ip_len] = '\0'; - const char *match_begin = strstr(full_list, current); - if (!match_begin || - !(match_begin == full_list || *(match_begin - 1) == ',') || - !(*(match_begin + ip_len) == ',' || *(match_begin + ip_len) == '\0')) { + if (!addr_list_contains(full_list, current, ip_len)) { DLOG("IP address missing: %s", current); return 1; } @@ -128,7 +140,8 @@ void doh_proxy_handle_resolver_update(const char *hostname, void *ctx, char * old_addr_list = strstr(p->resolv->data, port_colon); if (old_addr_list) { old_addr_list += strlen(port_colon); - if (!addr_list_reduced(addr_list, old_addr_list)) { + if (!addr_list_reduced(old_addr_list, addr_list) && + !addr_list_reduced(addr_list, old_addr_list)) { DLOG("DNS server IP address unchanged (%s).", buf + ip_start); free((void*)addr_list); return; @@ -137,12 +150,18 @@ void doh_proxy_handle_resolver_update(const char *hostname, void *ctx, } free((void*)addr_list); DLOG("Received new DNS server IP '%s'", buf + ip_start); - curl_slist_free_all(p->resolv); - p->resolv = curl_slist_append(NULL, buf); + struct curl_slist *new_resolv = curl_slist_append(NULL, buf); + if (new_resolv == NULL) { + FLOG("Out of mem updating DNS server IP"); + } // Reset libcurl: in-flight connections were aimed at the old IP, and curl // gets confused if we leave them around with a different CURLOPT_RESOLVE. + // The old list must remain alive until every easy handle referencing it has + // been removed and cleaned up by the reset. https_client_reset(p->client); + curl_slist_free_all(p->resolv); + p->resolv = new_resolv; if (p->awaiting_bootstrap) { p->awaiting_bootstrap = 0; @@ -172,6 +191,7 @@ static void doh_response_cb(void *data, char *buf, size_t buflen) { req_id, resp_id); } else { req->listener->respond(req->listener, (struct sockaddr*)&req->raddr, + req->connection_id, req->dns_req, req->dns_req_len, buf, buflen); if (p->stat) { stat_request_end(p->stat, buflen, @@ -182,12 +202,14 @@ static void doh_response_cb(void *data, char *buf, size_t buflen) { } } + req->listener->request_complete(req->listener, (struct sockaddr*)&req->raddr, + req->connection_id); free((void*)req->dns_req); free(req); } void doh_proxy_handle_request(void *ctx, dns_listener_t *listener, - struct sockaddr *raddr, + struct sockaddr *raddr, uint64_t connection_id, char *dns_req, size_t dns_req_len) { doh_proxy_t *p = (doh_proxy_t *)ctx; @@ -196,6 +218,7 @@ void doh_proxy_handle_request(void *ctx, dns_listener_t *listener, if (p->awaiting_bootstrap) { WLOG("%04hX: Query received before bootstrapping is completed, discarding.", req_id); + listener->request_complete(listener, raddr, connection_id); free(dns_req); return; } @@ -206,6 +229,7 @@ void doh_proxy_handle_request(void *ctx, dns_listener_t *listener, } req->proxy = p; req->listener = listener; + req->connection_id = connection_id; req->dns_req = dns_req; req->dns_req_len = dns_req_len; // raddr length depends on family; sockaddr_storage holds either. Copy what diff --git a/src/doh_proxy.h b/src/doh_proxy.h index a0dd2f4..cdf96ae 100644 --- a/src/doh_proxy.h +++ b/src/doh_proxy.h @@ -42,7 +42,7 @@ void doh_proxy_set_resolv(doh_proxy_t *p, const char *buf); // dns_request_fn — pass to dns_*_listener_create as the request callback. // `ctx` must be a doh_proxy_t *. void doh_proxy_handle_request(void *ctx, dns_listener_t *listener, - struct sockaddr *raddr, + struct sockaddr *raddr, uint64_t connection_id, char *dns_req, size_t dns_req_len); // dns_poller_cb — pass to dns_poller_init as the resolver-update callback. diff --git a/src/https_client.c b/src/https_client.c index fdafb03..b4a9c7e 100644 --- a/src/https_client.c +++ b/src/https_client.c @@ -12,6 +12,7 @@ #include "https_client.h" #include "logging.h" #include "options.h" +#include "outbound_interface.h" #include "stat.h" #define DOH_CONTENT_TYPE "application/dns-message" @@ -55,6 +56,12 @@ static void https_fetch_ctx_cleanup(https_client_t *client, struct https_fetch_ctx *ctx, int curl_result_code); +static void schedule_client_reset(https_client_t *client, ev_tstamp delay) { + ev_timer_stop(client->loop, &client->reset_timer); + ev_timer_set(&client->reset_timer, delay, 0); + ev_timer_start(client->loop, &client->reset_timer); +} + static size_t write_buffer(void *buf, size_t size, size_t nmemb, void *userp) { GET_PTR(struct https_fetch_ctx, ctx, userp); size_t write_size = size * nmemb; @@ -90,6 +97,15 @@ static curl_socket_t opensocket_callback(void *clientp, curlsocktype purpose, ELOG("Could not open curl socket %d:%s", errno, strerror(errno)); return CURL_SOCKET_BAD; } + if (client->opt->outbound_interface != NULL && + outbound_interface_bind_socket(sock, + client->opt->outbound_interface) != 0) { + const int bind_errno = errno; + ELOG("Failed to bind HTTPS socket to interface '%s': %s (%d)", + client->opt->outbound_interface, strerror(bind_errno), bind_errno); + (void)close(sock); + return CURL_SOCKET_BAD; + } DLOG("curl opened socket: %d", sock); client->connections++; @@ -287,6 +303,7 @@ static void https_fetch_ctx_init(https_client_t *client, ctx->buflen = 0; ctx->next = client->fetches; client->fetches = ctx; + client->fetch_count++; ASSERT_CURL_EASY_SETOPT(ctx, CURLOPT_RESOLVE, resolv); @@ -322,7 +339,14 @@ static void https_fetch_ctx_init(https_client_t *client, } if (client->opt->source_addr) { DLOG_REQ("Using source address: %s", client->opt->source_addr); - ASSERT_CURL_EASY_SETOPT(ctx, CURLOPT_INTERFACE, client->opt->source_addr); + if (client->opt->outbound_interface != NULL) { + char source_address[sizeof("host!") + INET6_ADDRSTRLEN]; + (void)snprintf(source_address, sizeof(source_address), "host!%s", + client->opt->source_addr); + ASSERT_CURL_EASY_SETOPT(ctx, CURLOPT_INTERFACE, source_address); + } else { + ASSERT_CURL_EASY_SETOPT(ctx, CURLOPT_INTERFACE, client->opt->source_addr); + } } if (client->opt->ca_info) { ASSERT_CURL_EASY_SETOPT(ctx, CURLOPT_CAINFO, client->opt->ca_info); @@ -372,7 +396,7 @@ static int https_fetch_ctx_process_response(https_client_t *client, recoverable = 1; if (!ev_is_active(&client->reset_timer)) { ILOG_REQ("Client reset timer started"); - ev_timer_start(client->loop, &client->reset_timer); + schedule_client_reset(client, (ev_tstamp)client->opt->conn_loss_time); } ILOG_REQ("curl request failed with %d: %s (recoverable, reset timer will recycle connection)", curl_result_code, curl_easy_strerror(curl_result_code)); @@ -555,6 +579,10 @@ static void https_fetch_ctx_cleanup(https_client_t *client, } else { client->fetches = ctx->next; } + if (client->fetch_count == 0) { + FLOG_REQ("HTTPS fetch count underflow"); + } + client->fetch_count--; free(ctx); } @@ -580,6 +608,19 @@ static void check_multi_info(https_client_t *c) { } } +static int handle_multi_action_result(https_client_t *client, CURLMcode code) { + if (code == CURLM_OK) { + return 0; + } + if (code == CURLM_ABORTED_BY_CALLBACK) { + WLOG("curl_multi_socket_action aborted by callback; scheduling HTTPS client reset"); + schedule_client_reset(client, 0); + return -1; + } + FLOG("curl_multi_socket_action error %d: %s", code, curl_multi_strerror(code)); + return -1; +} + static void sock_cb(struct ev_loop __attribute__((unused)) *loop, struct ev_io *w, int revents) { GET_PTR(https_client_t, c, w->data); @@ -588,16 +629,10 @@ static void sock_cb(struct ev_loop __attribute__((unused)) *loop, c->curlm, w->fd, (revents & EV_READ ? CURL_CSELECT_IN : 0) | (revents & EV_WRITE ? CURL_CSELECT_OUT : 0), &ignore); - if (code == CURLM_OK) { - check_multi_info(c); - } - else { - FLOG("curl_multi_socket_action error %d: %s", code, curl_multi_strerror(code)); - if (code == CURLM_ABORTED_BY_CALLBACK) { - WLOG("Resetting HTTPS client to recover from faulty state!"); - https_client_reset(c); - } + if (handle_multi_action_result(c, code) != 0) { + return; } + check_multi_info(c); } static void timer_cb(struct ev_loop __attribute__((unused)) *loop, @@ -606,27 +641,41 @@ static void timer_cb(struct ev_loop __attribute__((unused)) *loop, int ignore = 0; CURLMcode code = curl_multi_socket_action(c->curlm, CURL_SOCKET_TIMEOUT, 0, &ignore); - if (code != CURLM_OK) { - ELOG("curl_multi_socket_action error %d: %s", code, curl_multi_strerror(code)); + if (handle_multi_action_result(c, code) != 0) { + return; } check_multi_info(c); } -static struct ev_io * get_io_event(struct ev_io io_events[], curl_socket_t sock) { - for (int i = 0; i < HTTPS_SOCKET_LIMIT; i++) { - if (io_events[i].fd == sock) { - return &io_events[i]; +static struct ev_io * get_io_event(struct https_io_event *io_events, + curl_socket_t sock) { + for (struct https_io_event *event = io_events; + event != NULL; event = event->next) { + if (event->watcher.fd == sock) { + return &event->watcher; } } return NULL; } -static void dump_io_events(struct ev_io io_events[]) { - for (int i = 0; i < HTTPS_SOCKET_LIMIT; i++) { +static void dump_io_events(struct https_io_event *io_events) { + int index = 0; + for (struct https_io_event *event = io_events; + event != NULL; event = event->next) { + const ev_io *watcher = &event->watcher; ILOG("IO event #%d: fd=%d, events=%d/%s%s", - i+1, io_events[i].fd, io_events[i].events, - (io_events[i].events & EV_READ ? "R" : ""), - (io_events[i].events & EV_WRITE ? "W" : "")); + ++index, watcher->fd, watcher->events, + (watcher->events & EV_READ ? "R" : ""), + (watcher->events & EV_WRITE ? "W" : "")); + } +} + +static void free_io_events(https_client_t *client) { + while (client->io_events != NULL) { + struct https_io_event *event = client->io_events; + client->io_events = event->next; + ev_io_stop(client->loop, &event->watcher); + free(event); } } @@ -640,19 +689,26 @@ static int multi_sock_cb(CURL *curl, curl_socket_t sock, int what, struct ev_io *io_event_ptr = get_io_event(c->io_events, sock); if (io_event_ptr) { ev_io_stop(c->loop, io_event_ptr); - io_event_ptr->fd = 0; + io_event_ptr->fd = CURL_SOCKET_BAD; DLOG("Released used io event: %p", io_event_ptr); } if (what == CURL_POLL_REMOVE) { return 0; } // reserve and start new event on unused slot - io_event_ptr = get_io_event(c->io_events, 0); + io_event_ptr = get_io_event(c->io_events, CURL_SOCKET_BAD); if (!io_event_ptr) { - ELOG("curl needed more IO event handler, than the number of maximum sockets: %d", HTTPS_SOCKET_LIMIT); - dump_io_events(c->io_events); - logging_flight_recorder_dump(); - return -1; + struct https_io_event *event = + (struct https_io_event *)calloc(1, sizeof(*event)); + if (event == NULL) { + dump_io_events(c->io_events); + FLOG("Out of mem allocating curl socket watcher"); + } + event->watcher.fd = CURL_SOCKET_BAD; + event->watcher.data = c; + event->next = c->io_events; + c->io_events = event; + io_event_ptr = &event->watcher; } DLOG("Reserved new io event: %p", io_event_ptr); ev_io_init(io_event_ptr, sock_cb, sock, @@ -699,9 +755,6 @@ void https_client_init(https_client_t *c, options_t *opt, c->loop = loop; c->fetches = NULL; c->timer.data = c; - for (int i = 0; i < HTTPS_SOCKET_LIMIT; i++) { - c->io_events[i].data = c; - } c->opt = opt; c->stat = stat; @@ -718,6 +771,11 @@ void https_client_fetch(https_client_t *c, const char *url, const char* postdata, size_t postdata_len, struct curl_slist *resolv, uint16_t id, https_response_cb cb, void *data) { + if (c->fetch_count >= HTTPS_FETCH_LIMIT) { + WLOG("%04hX: Too many HTTPS requests in flight, dropping request", id); + cb(data, NULL, 0); + return; + } struct https_fetch_ctx *ctx = (struct https_fetch_ctx *)calloc(1, sizeof(struct https_fetch_ctx)); if (!ctx) { @@ -730,6 +788,7 @@ void https_client_reset(https_client_t *c) { struct curl_slist *header_list = c->header_list; c->header_list = NULL; https_client_cleanup(c); + ev_timer_set(&c->reset_timer, (ev_tstamp)c->opt->conn_loss_time, 0); https_client_multi_init(c, header_list); } @@ -739,5 +798,8 @@ void https_client_cleanup(https_client_t *c) { } curl_slist_free_all(c->header_list); curl_multi_cleanup(c->curlm); + c->curlm = NULL; + ev_timer_stop(c->loop, &c->timer); ev_timer_stop(c->loop, &c->reset_timer); + free_io_events(c); } diff --git a/src/https_client.h b/src/https_client.h index de60bbb..81d608b 100644 --- a/src/https_client.h +++ b/src/https_client.h @@ -9,6 +9,7 @@ enum { HTTPS_SOCKET_LIMIT = 12, HTTPS_CONNECTION_LIMIT = 8, + HTTPS_FETCH_LIMIT = 64, }; // Callback type for receiving data when a transfer finishes. @@ -30,15 +31,21 @@ struct https_fetch_ctx { struct https_fetch_ctx *next; }; +struct https_io_event { + ev_io watcher; + struct https_io_event *next; +}; + // Holds state on the whole multiplexed CURL machine. typedef struct { struct ev_loop *loop; CURLM *curlm; struct curl_slist *header_list; struct https_fetch_ctx *fetches; + size_t fetch_count; ev_timer timer; - ev_io io_events[HTTPS_SOCKET_LIMIT]; + struct https_io_event *io_events; int connections; options_t *opt; diff --git a/src/main.c b/src/main.c index 95aafd7..ccbbfc7 100644 --- a/src/main.c +++ b/src/main.c @@ -22,6 +22,7 @@ #include "https_client.h" #include "logging.h" #include "options.h" +#include "outbound_interface.h" #include "stat.h" static int is_ipv4_address(const char *str) { @@ -161,6 +162,30 @@ static const char * sw_version(void) { #endif } +static void validate_outbound_settings(const options_t *opt) { + if (opt->outbound_interface == NULL) { + return; + } + if (outbound_interface_validate(opt->outbound_interface) != 0) { + const int bind_errno = errno; + switch (bind_errno) { + case ENODEV: + FLOG("Outbound interface '%s' does not exist", opt->outbound_interface); + break; + case EPERM: + case EACCES: + FLOG("Permission denied binding to outbound interface '%s'; " + "start as the final service user with CAP_NET_RAW", + opt->outbound_interface); + break; + default: + FLOG("Cannot bind outbound sockets to interface '%s': %s (%d)", + opt->outbound_interface, strerror(bind_errno), bind_errno); + break; + } + } +} + int main(int argc, char *argv[]) { struct Options opt; options_init(&opt); @@ -278,6 +303,8 @@ int main(int argc, char *argv[]) { FLOG("Failed to set uid"); } + validate_outbound_settings(&opt); + if (opt.daemonize) { // daemon() is non-standard. If needed, see OpenSSH openbsd-compat/daemon.c if (daemon(0, 0) == -1) { @@ -316,6 +343,7 @@ int main(int argc, char *argv[]) { doh_proxy_await_bootstrap(proxy, systemd_notify_ready); dns_poller_init(dns_poller, loop, opt.bootstrap_dns, opt.bootstrap_dns_polling_interval, opt.source_addr, + opt.outbound_interface, hostname, opt.ipv4 ? AF_INET : AF_UNSPEC, doh_proxy_handle_resolver_update, proxy); ILOG("DNS polling initialized for '%s'", hostname); diff --git a/src/options.c b/src/options.c index ef3d9bb..b3e2830 100644 --- a/src/options.c +++ b/src/options.c @@ -1,5 +1,8 @@ +#include #include +#include #include +#include #include #include #include @@ -41,6 +44,7 @@ void options_init(struct Options *opt) { opt->resolver_ip = NULL; opt->curl_proxy = NULL; opt->source_addr = NULL; + opt->outbound_interface = NULL; opt->use_http_version = DEFAULT_HTTP_VERSION; opt->max_idle_time = 118; opt->conn_loss_time = 15; @@ -59,8 +63,14 @@ int parse_int(char * str) { } enum OptionsParseResult options_parse_args(struct Options *opt, int argc, char **argv) { + static const struct option long_options[] = { + {"outbound-interface", required_argument, NULL, 'I'}, + {NULL, 0, NULL, 0} + }; int c = 0; - while ((c = getopt(argc, argv, "a:c:p:T:du:g:b:i:4r:R:e:t:l:vxqm:L:s:S:C:F:hV")) != -1) { + while ((c = getopt_long(argc, argv, + "a:c:p:T:du:g:b:i:4r:R:e:t:l:vxqm:L:s:S:I:C:F:hV", + long_options, NULL)) != -1) { switch (c) { case 'a': // listen_addr opt->listen_addr = optarg; @@ -131,6 +141,9 @@ enum OptionsParseResult options_parse_args(struct Options *opt, int argc, char * case 'S': // source address opt->source_addr = optarg; break; + case 'I': // outbound interface + opt->outbound_interface = optarg; + break; case 'C': // CA info opt->ca_info = optarg; break; @@ -154,6 +167,9 @@ enum OptionsParseResult options_parse_args(struct Options *opt, int argc, char * return OPR_OPTION_ERROR; } opt->uid = p->pw_uid; + if (!opt->group && geteuid() == 0) { + opt->gid = p->pw_gid; + } } if (opt->group) { struct group *g = getgrnam(opt->group); @@ -169,7 +185,8 @@ enum OptionsParseResult options_parse_args(struct Options *opt, int argc, char * } opt->dscp <<= 2; // Get noisy about bad security practices. - if (getuid() == 0 && (!opt->user || !opt->group)) { + if (getuid() == 0 && + (opt->uid == (uid_t)-1 || opt->gid == (gid_t)-1)) { printf("----------------------------\n" "WARNING: Running as root without dropping privileges " "is NOT recommended.\n" @@ -222,6 +239,32 @@ enum OptionsParseResult options_parse_args(struct Options *opt, int argc, char * printf("TCP client limit must be between 0 and %u.\n", MAX_TCP_CLIENTS); return OPR_OPTION_ERROR; } + if (opt->outbound_interface != NULL) { + if (opt->source_addr != NULL) { + struct in_addr address_v4; + struct in6_addr address_v6; + if (inet_pton(AF_INET, opt->source_addr, &address_v4) != 1 && + inet_pton(AF_INET6, opt->source_addr, &address_v6) != 1) { + printf("-S must be an IP literal when using --outbound-interface.\n"); + return OPR_OPTION_ERROR; + } + } +#if !defined(__linux__) && !defined(OUTBOUND_INTERFACE_TEST) + printf("Outbound interface binding is supported only on Linux.\n"); + return OPR_OPTION_ERROR; +#endif + const size_t interface_length = strlen(opt->outbound_interface); + if (interface_length == 0 || interface_length >= IFNAMSIZ) { + printf("Outbound interface name must contain between 1 and %u characters.\n", + (unsigned)(IFNAMSIZ - 1)); + return OPR_OPTION_ERROR; + } + if (opt->user != NULL) { + printf("--outbound-interface cannot be combined with -u. Start directly " + "as the target user with CAP_NET_RAW supplied by the service manager.\n"); + return OPR_OPTION_ERROR; + } + } return OPR_SUCCESS; } @@ -231,6 +274,7 @@ void options_show_usage(int __attribute__((unused)) argc, char **argv) { printf("Usage: %s [-a ] [-p ] [-T ]\n", argv[0]); printf(" [-b ] [-i ] [-4]\n"); printf(" [-r ] [-R ] [-t ] [-S ]\n"); + printf(" [--outbound-interface ]\n"); printf(" [-x] [-q] [-C ] [-c ]\n"); printf(" [-d] [-u ] [-g ] \n"); printf(" [-v]+ [-l ] [-s ] [-F ] [-V] [-h]\n"); @@ -262,6 +306,9 @@ void options_show_usage(int __attribute__((unused)) argc, char **argv) { printf(" bootstrap DNS servers.\n"); printf(" -S source_addr Source IPv4/v6 address for outbound HTTPS and bootstrap DNS.\n"); printf(" (Default: system default)\n"); + printf(" -I, --outbound-interface interface\n"); + printf(" Linux network interface for outbound HTTPS and bootstrap DNS.\n"); + printf(" Incompatible with -u; requires CAP_NET_RAW.\n"); printf(" -x Use HTTP/1.1 instead of HTTP/2. Useful with broken\n" " or limited builds of libcurl.\n"); printf(" -q Use HTTP/3 (QUIC) only.\n"); @@ -278,6 +325,7 @@ void options_show_usage(int __attribute__((unused)) argc, char **argv) { printf("\n Process\n"); printf(" -d Daemonize.\n"); printf(" -u user Optional user to drop to if launched as root.\n"); + printf(" Also uses the user's primary group unless -g is set.\n"); printf(" -g group Optional group to drop to if launched as root.\n"); printf("\n Logging\n"); printf(" -v Increase logging verbosity. (Default: error)\n"); diff --git a/src/options.h b/src/options.h index 4e2aea4..7b6cf81 100644 --- a/src/options.h +++ b/src/options.h @@ -46,9 +46,12 @@ struct Options { // e.g. "socks5://127.0.0.1:1080" const char *curl_proxy; - // Source address for outbound HTTPS connections + // Source address for outbound HTTPS and bootstrap DNS sockets. const char *source_addr; + // Linux network interface for outbound HTTPS and bootstrap DNS sockets. + const char *outbound_interface; + // 1 = Use only HTTP/1.1 for limited OpenWRT libcurl (which is not built with HTTP/2 support) // 2 = Use only HTTP/2 default // 3 = Use only HTTP/3 QUIC diff --git a/src/outbound_interface.c b/src/outbound_interface.c new file mode 100644 index 0000000..17172da --- /dev/null +++ b/src/outbound_interface.c @@ -0,0 +1,41 @@ +#include +#include +#include +#include +#include + +#include "outbound_interface.h" + +int outbound_interface_bind_socket(int socket_fd, const char *interface_name) { + if (interface_name == NULL || interface_name[0] == '\0') { + errno = EINVAL; + return -1; + } + const size_t interface_length = strlen(interface_name); + if (interface_length >= IFNAMSIZ) { + errno = ENAMETOOLONG; + return -1; + } + +#if defined(__linux__) || defined(OUTBOUND_INTERFACE_TEST) + return setsockopt(socket_fd, SOL_SOCKET, SO_BINDTODEVICE, + interface_name, (socklen_t)(interface_length + 1)); +#else + (void)socket_fd; + errno = ENOTSUP; + return -1; +#endif +} + +int outbound_interface_validate(const char *interface_name) { + int socket_fd = socket(AF_INET, SOCK_DGRAM, 0); + if (socket_fd < 0) { + return -1; + } + + const int result = outbound_interface_bind_socket(socket_fd, interface_name); + const int saved_errno = errno; + (void)close(socket_fd); + errno = saved_errno; + return result; +} diff --git a/src/outbound_interface.h b/src/outbound_interface.h new file mode 100644 index 0000000..ac78068 --- /dev/null +++ b/src/outbound_interface.h @@ -0,0 +1,7 @@ +#ifndef _OUTBOUND_INTERFACE_H_ +#define _OUTBOUND_INTERFACE_H_ + +int outbound_interface_bind_socket(int socket_fd, const char *interface_name); +int outbound_interface_validate(const char *interface_name); + +#endif diff --git a/tests/linux/test_outbound_interface_linux.c b/tests/linux/test_outbound_interface_linux.c new file mode 100644 index 0000000..6532f58 --- /dev/null +++ b/tests/linux/test_outbound_interface_linux.c @@ -0,0 +1,27 @@ +#include +#include +#include + +#include "test_harness.h" +#include "outbound_interface.h" + +static void setUp(void) { +} + +static void tearDown(void) { +} + +static void test_real_interface_binding(void) { + TEST_ASSERT(geteuid() == 0); + TEST_ASSERT(outbound_interface_validate("lo") == 0); + TEST_ASSERT(outbound_interface_validate("hdp-no-such-if") == -1); + TEST_ASSERT(errno == ENODEV); +} + +int main(void) { + if (getenv("HDP_RUN_PRIVILEGED_TEST") == NULL) { + return 77; + } + TEST_RUN(test_real_interface_binding); + return test_summary(); +} diff --git a/tests/unit/test_dns_listener_tcp.c b/tests/unit/test_dns_listener_tcp.c new file mode 100644 index 0000000..5cacdff --- /dev/null +++ b/tests/unit/test_dns_listener_tcp.c @@ -0,0 +1,451 @@ +#include +#include +#include +#include +#include + +#include "test_harness.h" +#include "../../src/dns_listener_tcp.c" + +static struct tcp_client_s client; +static dns_listener_tcp_t listener; +static char *dns_request; +static int sockets[2]; + +static void setUp(void) { + memset(&client, 0, sizeof(client)); + memset(&listener, 0, sizeof(listener)); + sockets[0] = -1; + sockets[1] = -1; + client.input_buffer_size = 64; + client.input_buffer = (char *)calloc(1, client.input_buffer_size); + dns_request = NULL; + + listener.loop = ev_loop_new(0); + listener.clients = &client; + client.d = &listener; + if (listener.loop == NULL || socketpair(AF_UNIX, SOCK_STREAM, 0, sockets) != 0) { + test_fail(__FILE__, __LINE__, "test socket setup"); + return; + } + client.sock = sockets[0]; + client.addr_len = sizeof(struct sockaddr); + client.raddr.ss_family = AF_UNIX; + ev_io_init(&client.read_watcher, read_cb, client.sock, EV_READ); + client.read_watcher.data = &client; + ev_io_init(&client.write_watcher, write_cb, client.sock, EV_WRITE); + client.write_watcher.data = &client; + ev_timer_init(&client.timer_watcher, timer_cb, IDLE_TIMEOUT_S, IDLE_TIMEOUT_S); + client.timer_watcher.data = &client; +} + +static void tearDown(void) { + if (listener.clients != NULL && listener.clients != &client) { + remove_client(listener.clients); + sockets[0] = -1; + } + free(dns_request); + free(client.input_buffer); + free_output_queue(&client); + if (listener.loop != NULL) { + ev_io_stop(listener.loop, &client.read_watcher); + ev_io_stop(listener.loop, &client.write_watcher); + ev_timer_stop(listener.loop, &client.timer_watcher); + ev_loop_destroy(listener.loop); + } + if (sockets[0] >= 0) { + close(sockets[0]); + } + if (sockets[1] >= 0) { + close(sockets[1]); + } +} + +static void test_waits_for_complete_length_prefix(void) { + client.input_buffer[0] = (char)UINT8_MAX; + client.input_buffer_used = 1; + uint16_t request_size = 0; + + const int received = get_dns_request(&client, &dns_request, &request_size); + + TEST_ASSERT(received == 0); + TEST_ASSERT(dns_request == NULL); + TEST_ASSERT(client.input_buffer_used == 1); +} + +static void test_large_length_does_not_wrap_frame_size(void) { + client.input_buffer[0] = (char)UINT8_MAX; + client.input_buffer[1] = (char)UINT8_MAX; + client.input_buffer_used = 2; + uint16_t request_size = 0; + + const int received = get_dns_request(&client, &dns_request, &request_size); + + TEST_ASSERT(received == 0); + TEST_ASSERT(request_size == UINT16_MAX); + TEST_ASSERT(dns_request == NULL); + TEST_ASSERT(client.input_buffer_used == 2); +} + +static void test_rejects_zero_length_request_without_allocating(void) { + client.input_buffer[0] = 0; + client.input_buffer[1] = 0; + client.input_buffer_used = sizeof(uint16_t); + uint16_t request_size = UINT16_MAX; + + const int received = get_dns_request(&client, &dns_request, &request_size); + + TEST_ASSERT(received == 1); + TEST_ASSERT(request_size == 0); + TEST_ASSERT(dns_request == NULL); + TEST_ASSERT(client.input_buffer_used == 0); +} + +static void test_extracts_complete_request(void) { + enum { REQUEST_SIZE = DNS_HEADER_LENGTH }; + client.input_buffer[0] = 0; + client.input_buffer[1] = REQUEST_SIZE; + for (size_t i = 0; i < REQUEST_SIZE; i++) { + client.input_buffer[sizeof(uint16_t) + i] = (char)i; + } + client.input_buffer_used = sizeof(uint16_t) + REQUEST_SIZE; + uint16_t request_size = 0; + + const int received = get_dns_request(&client, &dns_request, &request_size); + + TEST_ASSERT(received == 1); + TEST_ASSERT(request_size == REQUEST_SIZE); + TEST_ASSERT(client.input_buffer_used == 0); + for (size_t i = 0; i < REQUEST_SIZE; i++) { + TEST_ASSERT((uint8_t)dns_request[i] == i); + } +} + +static void test_large_partial_request_grows_without_exiting(void) { + client.input_buffer_size = 32768; + client.input_buffer = realloc(client.input_buffer, client.input_buffer_size); + TEST_ASSERT(client.input_buffer != NULL); + client.input_buffer_used = client.input_buffer_size; + const uint16_t length = htons(60000); + memcpy(client.input_buffer, &length, sizeof(length)); + char extra[1024] = {0}; + TEST_ASSERT(send(sockets[1], extra, sizeof(extra), 0) == sizeof(extra)); + read_cb(listener.loop, &client.read_watcher, EV_READ); + TEST_ASSERT(client.input_buffer_size >= client.input_buffer_used); + TEST_ASSERT(client.input_buffer_size <= TCP_DNS_MAX_FRAME); + TEST_ASSERT(client.input_buffer_used == 32768 + sizeof(extra)); +} + +static void test_request_completion_resumes_client_reads(void) { + client.pending_requests = TCP_CLIENT_REQUEST_LIMIT; + + tcp_request_complete(&listener.base, (struct sockaddr*)&client.raddr, client.id); + + TEST_ASSERT(client.pending_requests == TCP_CLIENT_REQUEST_LIMIT - 1); + TEST_ASSERT(ev_is_active(&client.read_watcher)); +} + +static void consume_request(void *ctx, dns_listener_t *unused_listener, + struct sockaddr *unused_address, + uint64_t unused_id, char *request, size_t unused_length) { + (void)unused_listener; + (void)unused_address; + (void)unused_id; + (void)unused_length; + unsigned *requests = ctx; + (*requests)++; + free(request); +} + +static void test_resume_drains_buffer_before_reading_socket(void) { + char frame[sizeof(uint16_t) + DNS_HEADER_LENGTH] = {0}; + const uint16_t length = htons(DNS_HEADER_LENGTH); + memcpy(frame, &length, sizeof(length)); + memcpy(client.input_buffer, frame, sizeof(frame)); + memcpy(client.input_buffer + sizeof(frame), frame, sizeof(frame)); + client.input_buffer_used = 2 * sizeof(frame); + client.pending_requests = TCP_CLIENT_REQUEST_LIMIT; + unsigned requests = 0; + listener.cb = consume_request; + listener.cb_data = &requests; + TEST_ASSERT(send(sockets[1], frame, sizeof(frame), 0) == sizeof(frame)); + + tcp_request_complete(&listener.base, (struct sockaddr *)&client.raddr, client.id); + read_cb(listener.loop, &client.read_watcher, EV_READ); + + TEST_ASSERT(requests == 1); + TEST_ASSERT(client.pending_requests == TCP_CLIENT_REQUEST_LIMIT); + TEST_ASSERT(client.input_buffer_used == sizeof(frame)); + TEST_ASSERT(!ev_is_active(&client.read_watcher)); + char unread[sizeof(frame)]; + TEST_ASSERT(recv(sockets[0], unread, sizeof(unread), MSG_DONTWAIT) == sizeof(frame)); + TEST_ASSERT(memcmp(unread, frame, sizeof(frame)) == 0); +} + +static void test_read_respects_remaining_input_capacity(void) { + client.input_buffer_size = TCP_DNS_MAX_FRAME; + client.input_buffer = realloc(client.input_buffer, client.input_buffer_size); + TEST_ASSERT(client.input_buffer != NULL); + client.input_buffer_used = TCP_DNS_MAX_FRAME - 8; + const uint16_t length = htons(UINT16_MAX); + memcpy(client.input_buffer, &length, sizeof(length)); + unsigned requests = 0; + listener.cb = consume_request; + listener.cb_data = &requests; + char extra[16] = {0}; + TEST_ASSERT(send(sockets[1], extra, sizeof(extra), 0) == sizeof(extra)); + read_cb(listener.loop, &client.read_watcher, EV_READ); + TEST_ASSERT(requests == 1); + TEST_ASSERT(client.input_buffer_used == 0); + TEST_ASSERT(recv(sockets[0], extra, sizeof(extra), MSG_DONTWAIT) == 8); +} + +static void test_old_connection_cannot_complete_or_respond_to_new_client(void) { + client.id = 2; + client.pending_requests = TCP_CLIENT_REQUEST_LIMIT; + tcp_request_complete(&listener.base, (struct sockaddr *)&client.raddr, 1); + TEST_ASSERT(client.pending_requests == TCP_CLIENT_REQUEST_LIMIT); + TEST_ASSERT(!ev_is_active(&client.read_watcher)); + + char response[DNS_HEADER_LENGTH] = {0}; + tcp_respond(&listener.base, (struct sockaddr *)&client.raddr, + 1, NULL, 0, response, sizeof(response)); + TEST_ASSERT(client.output_head == NULL); + TEST_ASSERT(client.output_bytes == 0); + char received[DNS_HEADER_LENGTH]; + TEST_ASSERT(recv(sockets[1], received, sizeof(received), MSG_DONTWAIT) == -1); + TEST_ASSERT(errno == EAGAIN || errno == EWOULDBLOCK); +} + +static void test_queues_response_without_blocking_event_loop(void) { + const int flags = fcntl(sockets[0], F_GETFL, 0); + TEST_ASSERT(flags >= 0); + TEST_ASSERT(fcntl(sockets[0], F_SETFL, flags | O_NONBLOCK) == 0); + + char fill[4096] = {0}; + while (send(sockets[0], fill, sizeof(fill), MSG_NOSIGNAL) > 0) { + } + TEST_ASSERT(errno == EAGAIN || errno == EWOULDBLOCK); + + char response_one[DNS_HEADER_LENGTH]; + char response_two[DNS_HEADER_LENGTH]; + memset(response_one, 0x11, sizeof(response_one)); + memset(response_two, 0x22, sizeof(response_two)); + response_one[0] = 0x12; + response_one[1] = 0x34; + response_two[0] = 0x56; + response_two[1] = 0x78; + tcp_respond(&listener.base, (struct sockaddr *)&client.raddr, + client.id, NULL, 0, response_one, sizeof(response_one)); + tcp_respond(&listener.base, (struct sockaddr *)&client.raddr, + client.id, NULL, 0, response_two, sizeof(response_two)); + + TEST_ASSERT(client.output_head != NULL); + TEST_ASSERT(client.output_bytes == + 2 * (sizeof(uint16_t) + sizeof(response_one))); + TEST_ASSERT(ev_is_active(&client.write_watcher)); + + const int peer_flags = fcntl(sockets[1], F_GETFL, 0); + TEST_ASSERT(peer_flags >= 0); + TEST_ASSERT(fcntl(sockets[1], F_SETFL, peer_flags | O_NONBLOCK) == 0); + while (recv(sockets[1], fill, sizeof(fill), 0) > 0) { + } + TEST_ASSERT(errno == EAGAIN || errno == EWOULDBLOCK); + + write_cb(listener.loop, &client.write_watcher, EV_WRITE); + + char wire_data[2 * (sizeof(uint16_t) + DNS_HEADER_LENGTH)]; + size_t received = 0; + while (received < sizeof(wire_data)) { + const ssize_t bytes = recv(sockets[1], wire_data + received, + sizeof(wire_data) - received, 0); + TEST_ASSERT(bytes > 0); + received += (size_t)bytes; + } + + uint16_t frame_size = 0; + memcpy(&frame_size, wire_data, sizeof(frame_size)); + TEST_ASSERT(ntohs(frame_size) == sizeof(response_one)); + TEST_ASSERT(memcmp(wire_data + sizeof(frame_size), + response_one, sizeof(response_one)) == 0); + const size_t second_frame = sizeof(frame_size) + sizeof(response_one); + memcpy(&frame_size, wire_data + second_frame, sizeof(frame_size)); + TEST_ASSERT(ntohs(frame_size) == sizeof(response_two)); + TEST_ASSERT(memcmp(wire_data + second_frame + sizeof(frame_size), + response_two, sizeof(response_two)) == 0); + + TEST_ASSERT(client.output_head == NULL); + TEST_ASSERT(client.output_tail == NULL); + TEST_ASSERT(client.output_bytes == 0); + TEST_ASSERT(!ev_is_active(&client.write_watcher)); +} + +static struct tcp_client_s *heap_client(void) { + struct tcp_client_s *owned = malloc(sizeof(*owned)); + if (owned == NULL) { + return NULL; + } + *owned = client; + client.input_buffer = NULL; + listener.clients = owned; + listener.client_count = 1; + listener.client_limit = 2; + ev_io_init(&owned->read_watcher, read_cb, owned->sock, EV_READ); + owned->read_watcher.data = owned; + ev_io_start(listener.loop, &owned->read_watcher); + ev_io_init(&owned->write_watcher, write_cb, owned->sock, EV_WRITE); + owned->write_watcher.data = owned; + ev_timer_init(&owned->timer_watcher, timer_cb, IDLE_TIMEOUT_S, IDLE_TIMEOUT_S); + owned->timer_watcher.data = owned; + return owned; +} + +static void check_queued_response_drains(int stopping) { + struct tcp_client_s *owned = heap_client(); + TEST_ASSERT(owned != NULL); + TEST_ASSERT(fcntl(sockets[0], F_SETFL, O_NONBLOCK) == 0); + TEST_ASSERT(fcntl(sockets[1], F_SETFL, O_NONBLOCK) == 0); + char fill[4096] = {0}; + while (send(sockets[0], fill, sizeof(fill), MSG_NOSIGNAL) > 0) { + } + TEST_ASSERT(errno == EAGAIN || errno == EWOULDBLOCK); + char response[DNS_HEADER_LENGTH] = {0}; + tcp_respond(&listener.base, (struct sockaddr *)&owned->raddr, + owned->id, NULL, 0, response, sizeof(response)); + TEST_ASSERT(owned->output_head != NULL); + if (stopping) { + tcp_stop(&listener.base); + TEST_ASSERT(ev_is_active(&owned->timer_watcher)); + TEST_ASSERT(owned->timer_watcher.repeat == 0); + } else { + TEST_ASSERT(shutdown(sockets[1], SHUT_WR) == 0); + read_cb(listener.loop, &owned->read_watcher, EV_READ); + } + TEST_ASSERT(listener.clients == owned); + TEST_ASSERT(owned->read_eof); + TEST_ASSERT(!ev_is_active(&owned->read_watcher)); + while (recv(sockets[1], fill, sizeof(fill), 0) > 0) { + } + write_cb(listener.loop, &owned->write_watcher, EV_WRITE); + TEST_ASSERT(listener.clients == NULL); + sockets[0] = -1; + char frame[sizeof(uint16_t) + DNS_HEADER_LENGTH]; + TEST_ASSERT(recv(sockets[1], frame, sizeof(frame), 0) == sizeof(frame)); + uint16_t length; + memcpy(&length, frame, sizeof(length)); + TEST_ASSERT(ntohs(length) == sizeof(response)); + TEST_ASSERT(memcmp(frame + sizeof(length), response, sizeof(response)) == 0); + TEST_ASSERT(recv(sockets[1], frame, sizeof(frame), 0) == 0); +} + +static void test_half_close_keeps_queued_response_until_drained(void) { + check_queued_response_drains(0); +} + +static void test_shutdown_keeps_queued_response_until_drained(void) { + check_queued_response_drains(1); +} + +static void test_shutdown_waits_for_pending_response(void) { + struct tcp_client_s *owned = heap_client(); + TEST_ASSERT(owned != NULL); + owned->pending_requests = TCP_CLIENT_REQUEST_LIMIT; + listener.client_limit = 1; + tcp_stop(&listener.base); + TEST_ASSERT(listener.clients == owned); + TEST_ASSERT(!ev_is_active(&owned->read_watcher)); + TEST_ASSERT(!ev_is_active(&listener.accept_watcher)); + const ev_tstamp deadline = owned->timer_watcher.at; + char response[DNS_HEADER_LENGTH] = {0}; + tcp_respond(&listener.base, (struct sockaddr *)&owned->raddr, + owned->id, NULL, 0, response, sizeof(response)); + tcp_request_complete(&listener.base, (struct sockaddr *)&owned->raddr, owned->id); + TEST_ASSERT(!ev_is_active(&owned->read_watcher)); + TEST_ASSERT(ev_is_active(&owned->timer_watcher)); + TEST_ASSERT(owned->timer_watcher.at == deadline); + TEST_ASSERT(owned->timer_watcher.repeat == 0); + tcp_stop(&listener.base); + TEST_ASSERT(owned->timer_watcher.at == deadline); + char frame[sizeof(uint16_t) + DNS_HEADER_LENGTH]; + TEST_ASSERT(recv(sockets[1], frame, sizeof(frame), 0) == sizeof(frame)); + timer_cb(listener.loop, &owned->timer_watcher, EV_TIMER); + sockets[0] = -1; + TEST_ASSERT(listener.clients == NULL); + TEST_ASSERT(!ev_is_active(&listener.accept_watcher)); +} + +static void test_shutdown_removes_idle_client(void) { + TEST_ASSERT(heap_client() != NULL); + listener.client_limit = 1; + tcp_stop(&listener.base); + sockets[0] = -1; + TEST_ASSERT(listener.clients == NULL); + TEST_ASSERT(!ev_is_active(&listener.accept_watcher)); +} + +static void test_half_close_waits_for_pending_request_completion(void) { + struct tcp_client_s *owned = heap_client(); + TEST_ASSERT(owned != NULL); + owned->pending_requests = 1; + TEST_ASSERT(shutdown(sockets[1], SHUT_WR) == 0); + read_cb(listener.loop, &owned->read_watcher, EV_READ); + TEST_ASSERT(listener.clients == owned); + TEST_ASSERT(owned->read_eof); + struct sockaddr_storage address = owned->raddr; + const uint64_t connection_id = owned->id; + tcp_request_complete(&listener.base, (struct sockaddr *)&address, connection_id); + TEST_ASSERT(listener.clients == NULL); + sockets[0] = -1; +} + +static void test_accepts_maximum_payload_request_and_response(void) { + client.input_buffer_size = TCP_DNS_MAX_FRAME; + client.input_buffer = realloc(client.input_buffer, client.input_buffer_size); + TEST_ASSERT(client.input_buffer != NULL); + memset(client.input_buffer, 0, client.input_buffer_size); + const uint16_t length = htons(UINT16_MAX); + memcpy(client.input_buffer, &length, sizeof(length)); + client.input_buffer_used = TCP_DNS_MAX_FRAME; + uint16_t request_size = 0; + TEST_ASSERT(get_dns_request(&client, &dns_request, &request_size) == 1); + TEST_ASSERT(request_size == UINT16_MAX); + TEST_ASSERT(client.input_buffer_used == 0); + TEST_ASSERT(fcntl(sockets[0], F_SETFL, O_NONBLOCK) == 0); + tcp_respond(&listener.base, (struct sockaddr *)&client.raddr, + client.id, NULL, 0, dns_request, request_size); + size_t received = 0; + char buffer[4096]; + while (received < TCP_DNS_MAX_FRAME) { + const size_t remaining = TCP_DNS_MAX_FRAME - received; + const ssize_t count = recv(sockets[1], buffer, + remaining < sizeof(buffer) ? remaining : sizeof(buffer), 0); + TEST_ASSERT(count > 0); + if (received == 0) { + TEST_ASSERT(count >= (ssize_t)sizeof(length)); + TEST_ASSERT(memcmp(buffer, &length, sizeof(length)) == 0); + } + received += (size_t)count; + write_cb(listener.loop, &client.write_watcher, EV_WRITE); + } + TEST_ASSERT(client.output_head == NULL); +} + +int main(void) { + TEST_RUN(test_waits_for_complete_length_prefix); + TEST_RUN(test_large_length_does_not_wrap_frame_size); + TEST_RUN(test_rejects_zero_length_request_without_allocating); + TEST_RUN(test_extracts_complete_request); + TEST_RUN(test_large_partial_request_grows_without_exiting); + TEST_RUN(test_request_completion_resumes_client_reads); + TEST_RUN(test_resume_drains_buffer_before_reading_socket); + TEST_RUN(test_read_respects_remaining_input_capacity); + TEST_RUN(test_old_connection_cannot_complete_or_respond_to_new_client); + TEST_RUN(test_queues_response_without_blocking_event_loop); + TEST_RUN(test_half_close_keeps_queued_response_until_drained); + TEST_RUN(test_shutdown_keeps_queued_response_until_drained); + TEST_RUN(test_shutdown_waits_for_pending_response); + TEST_RUN(test_shutdown_removes_idle_client); + TEST_RUN(test_half_close_waits_for_pending_request_completion); + TEST_RUN(test_accepts_maximum_payload_request_and_response); + return test_summary(); +} diff --git a/tests/unit/test_dns_poller.c b/tests/unit/test_dns_poller.c index b28b523..b86d74c 100644 --- a/tests/unit/test_dns_poller.c +++ b/tests/unit/test_dns_poller.c @@ -25,7 +25,7 @@ static void setUp(void) { test_fail(__FILE__, __LINE__, "ev_loop_new(0)"); return; } - dns_poller_init(&poller, loop, "127.0.0.1,127.0.0.2", 60, NULL, + dns_poller_init(&poller, loop, "127.0.0.1,127.0.0.2", 60, NULL, NULL, "example.com", AF_INET, NULL, NULL); poller_initialized = 1; ev_timer_stop(loop, &poller.timer); @@ -126,10 +126,17 @@ static void test_cleanup_stops_active_watchers(void) { TEST_ASSERT(ev_run(loop, EVRUN_NOWAIT) == 0); } +static void test_rejects_socket_when_interface_binding_fails(void) { + poller.outbound_interface = strdup(""); + TEST_ASSERT(poller.outbound_interface != NULL); + TEST_ASSERT(configure_bootstrap_socket(ARES_SOCKET_BAD, 0, &poller) == -1); +} + int main(void) { TEST_RUN(test_allocates_watchers_beyond_nameserver_count); TEST_RUN(test_reuses_released_watcher); TEST_RUN(test_fd_zero_does_not_alias_free_watcher); TEST_RUN(test_cleanup_stops_active_watchers); + TEST_RUN(test_rejects_socket_when_interface_binding_fails); return test_summary(); } diff --git a/tests/unit/test_doh_proxy_resolver.c b/tests/unit/test_doh_proxy_resolver.c new file mode 100644 index 0000000..1383751 --- /dev/null +++ b/tests/unit/test_doh_proxy_resolver.c @@ -0,0 +1,172 @@ +#include +#include + +#include "test_harness.h" +#include "../../src/doh_proxy.c" + +static https_client_t client; +static doh_proxy_t *proxy; +static unsigned reset_calls; +static https_response_cb pending_response_cb; +static void *pending_response_data; +static uint64_t response_connection_id; +static uint64_t completion_connection_id; +static unsigned response_calls; +static unsigned completion_calls; +static const char *expected_resolver_during_reset; +static int reset_saw_expected_resolver; + +void https_client_reset(https_client_t *unused_client) { + (void)unused_client; + reset_calls++; + if (expected_resolver_during_reset != NULL && + proxy != NULL && proxy->resolv != NULL && + strcmp(proxy->resolv->data, expected_resolver_during_reset) == 0) { + reset_saw_expected_resolver = 1; + } +} + +void https_client_fetch(https_client_t *unused_client, const char *url, + const char *postdata, size_t postdata_len, + struct curl_slist *resolv, uint16_t id, + https_response_cb cb, void *data) { + (void)unused_client; + (void)url; + (void)postdata; + (void)postdata_len; + (void)resolv; + (void)id; + pending_response_cb = cb; + pending_response_data = data; +} + +void stat_request_begin(stat_t *stat, size_t request_len, uint8_t is_tcp) { + (void)stat; + (void)request_len; + (void)is_tcp; +} + +void stat_request_end(stat_t *stat, size_t response_len, + ev_tstamp latency, uint8_t is_tcp) { + (void)stat; + (void)response_len; + (void)latency; + (void)is_tcp; +} + +static void setUp(void) { + memset(&client, 0, sizeof(client)); + reset_calls = 0; + pending_response_cb = NULL; + pending_response_data = NULL; + response_connection_id = 0; + completion_connection_id = 0; + response_calls = 0; + completion_calls = 0; + expected_resolver_during_reset = NULL; + reset_saw_expected_resolver = 0; + proxy = doh_proxy_create(NULL, &client, "https://resolver.test/dns-query", NULL); + doh_proxy_set_resolv(proxy, "resolver.test:443:192.0.2.1,192.0.2.2"); +} + +static void tearDown(void) { + doh_proxy_destroy(proxy); + proxy = NULL; +} + +static void update_resolver(const char *addresses) { + char *owned_addresses = strdup(addresses); + if (owned_addresses == NULL) { + test_fail(__FILE__, __LINE__, "strdup(addresses)"); + return; + } + doh_proxy_handle_resolver_update("resolver.test", proxy, owned_addresses); +} + +static void test_new_address_triggers_resolver_update(void) { + expected_resolver_during_reset = + "resolver.test:443:192.0.2.1,192.0.2.2"; + update_resolver("192.0.2.1,192.0.2.2,192.0.2.3"); + + TEST_ASSERT(reset_calls == 1); + TEST_ASSERT(reset_saw_expected_resolver); + TEST_ASSERT(strcmp(proxy->resolv->data, + "resolver.test:443:192.0.2.1,192.0.2.2,192.0.2.3") == 0); +} + +static void test_removed_address_triggers_resolver_update(void) { + update_resolver("192.0.2.1"); + + TEST_ASSERT(reset_calls == 1); + TEST_ASSERT(strcmp(proxy->resolv->data, + "resolver.test:443:192.0.2.1") == 0); +} + +static void test_exact_address_match_after_prefix_match(void) { + curl_slist_free_all(proxy->resolv); + proxy->resolv = NULL; + doh_proxy_set_resolv(proxy, "resolver.test:443:192.0.2.10,192.0.2.1"); + + update_resolver("192.0.2.1,192.0.2.10"); + + TEST_ASSERT(reset_calls == 0); + TEST_ASSERT(strcmp(proxy->resolv->data, + "resolver.test:443:192.0.2.10,192.0.2.1") == 0); +} + +static void record_response(dns_listener_t *unused_listener, + struct sockaddr *unused_address, + uint64_t connection_id, + const char *unused_request, size_t unused_request_len, + char *unused_response, size_t unused_response_len) { + (void)unused_listener; + (void)unused_address; + (void)unused_request; + (void)unused_request_len; + (void)unused_response; + (void)unused_response_len; + response_connection_id = connection_id; + response_calls++; +} + +static void record_completion(dns_listener_t *unused_listener, + struct sockaddr *unused_address, + uint64_t connection_id) { + (void)unused_listener; + (void)unused_address; + completion_connection_id = connection_id; + completion_calls++; +} + +static void test_preserves_connection_identity_through_proxy_callbacks(void) { + dns_listener_t listener = { + .respond = record_response, + .request_complete = record_completion, + .transport = DNS_TRANSPORT_TCP + }; + struct sockaddr_in address = {.sin_family = AF_INET}; + for (unsigned attempt = 0; attempt < 2; attempt++) { + char *request = calloc(1, DNS_HEADER_LENGTH); + TEST_ASSERT(request != NULL); + const uint64_t connection_id = 100 + attempt; + doh_proxy_handle_request(proxy, &listener, (struct sockaddr *)&address, + connection_id, request, DNS_HEADER_LENGTH); + TEST_ASSERT(pending_response_cb != NULL); + TEST_ASSERT(completion_calls == attempt); + char response[DNS_HEADER_LENGTH] = {0}; + pending_response_cb(pending_response_data, attempt == 0 ? response : NULL, + attempt == 0 ? sizeof(response) : 0); + TEST_ASSERT(completion_calls == attempt + 1); + TEST_ASSERT(completion_connection_id == connection_id); + TEST_ASSERT(response_calls == 1); + TEST_ASSERT(response_connection_id == 100); + } +} + +int main(void) { + TEST_RUN(test_new_address_triggers_resolver_update); + TEST_RUN(test_removed_address_triggers_resolver_update); + TEST_RUN(test_exact_address_match_after_prefix_match); + TEST_RUN(test_preserves_connection_identity_through_proxy_callbacks); + return test_summary(); +} diff --git a/tests/unit/test_https_client_limits.c b/tests/unit/test_https_client_limits.c new file mode 100644 index 0000000..8b41617 --- /dev/null +++ b/tests/unit/test_https_client_limits.c @@ -0,0 +1,132 @@ +#include +#include +#include +#include + +#include "dns_common.h" +#include "https_client.h" +#include "test_harness.h" +#include "../../src/https_client.c" + +typedef struct { + size_t calls; + size_t failures; +} callback_state_t; + +static struct ev_loop *loop; +static https_client_t client; +static options_t options; +static callback_state_t callback_state; + +static void response_cb(void *data, char *buf, size_t buflen) { + callback_state_t *state = (callback_state_t *)data; + state->calls++; + if (buf == NULL && buflen == 0) { + state->failures++; + } +} + +static void setUp(void) { + memset(&client, 0, sizeof(client)); + memset(&options, 0, sizeof(options)); + memset(&callback_state, 0, sizeof(callback_state)); + options.use_http_version = 1; + options.max_idle_time = 118; + options.conn_loss_time = 15; + loop = ev_loop_new(0); + if (loop == NULL) { + test_fail(__FILE__, __LINE__, "ev_loop_new(0)"); + return; + } + if (curl_global_init(CURL_GLOBAL_DEFAULT) != CURLE_OK) { + test_fail(__FILE__, __LINE__, "curl_global_init"); + return; + } + https_client_init(&client, &options, NULL, loop); +} + +static void tearDown(void) { + if (client.curlm != NULL) { + https_client_cleanup(&client); + } + curl_global_cleanup(); + if (loop != NULL) { + ev_loop_destroy(loop); + loop = NULL; + } +} + +static void test_rejects_requests_above_fetch_limit(void) { + static const char request[DNS_HEADER_LENGTH] = {0}; + + for (size_t i = 0; i < HTTPS_FETCH_LIMIT; i++) { + https_client_fetch(&client, "https://127.0.0.1/dns-query", + request, sizeof(request), NULL, (uint16_t)i, + response_cb, &callback_state); + } + + TEST_ASSERT(client.fetch_count == HTTPS_FETCH_LIMIT); + TEST_ASSERT(callback_state.calls == 0); + + https_client_fetch(&client, "https://127.0.0.1/dns-query", + request, sizeof(request), NULL, UINT16_MAX, + response_cb, &callback_state); + + TEST_ASSERT(client.fetch_count == HTTPS_FETCH_LIMIT); + TEST_ASSERT(callback_state.calls == 1); + TEST_ASSERT(callback_state.failures == 1); +} + +static void test_socket_watchers_grow_and_support_fd_zero(void) { + CURL *fake_easy_handle = (CURL *)(uintptr_t)1; + enum { WATCHER_COUNT = 16 }; + + for (curl_socket_t fd = 0; fd < WATCHER_COUNT; fd++) { + TEST_ASSERT(multi_sock_cb(fake_easy_handle, fd, CURL_POLL_IN, + &client, NULL) == 0); + struct ev_io *watcher = get_io_event(client.io_events, fd); + TEST_ASSERT(watcher != NULL); + TEST_ASSERT(watcher->fd == fd); + TEST_ASSERT(ev_is_active(watcher)); + } + + struct ev_io *fd_zero_watcher = get_io_event(client.io_events, 0); + TEST_ASSERT(fd_zero_watcher != NULL); + TEST_ASSERT(multi_sock_cb(fake_easy_handle, 0, CURL_POLL_REMOVE, + &client, NULL) == 0); + TEST_ASSERT(get_io_event(client.io_events, 0) == NULL); + TEST_ASSERT(!ev_is_active(fd_zero_watcher)); + + TEST_ASSERT(multi_sock_cb(fake_easy_handle, WATCHER_COUNT, + CURL_POLL_OUT, &client, NULL) == 0); + TEST_ASSERT(get_io_event(client.io_events, WATCHER_COUNT) == fd_zero_watcher); +} + +static void test_aborted_multi_action_schedules_client_reset(void) { + TEST_ASSERT(!ev_is_active(&client.reset_timer)); + + TEST_ASSERT(handle_multi_action_result( + &client, CURLM_ABORTED_BY_CALLBACK) == -1); + + TEST_ASSERT(ev_is_active(&client.reset_timer)); +} + +static void test_rejects_socket_when_interface_binding_fails(void) { + struct curl_sockaddr address; + memset(&address, 0, sizeof(address)); + address.family = AF_INET; + address.socktype = SOCK_STREAM; + options.outbound_interface = ""; + + TEST_ASSERT(opensocket_callback(&client, CURLSOCKTYPE_IPCXN, &address) == + CURL_SOCKET_BAD); + TEST_ASSERT(client.connections == 0); +} + +int main(void) { + TEST_RUN(test_rejects_requests_above_fetch_limit); + TEST_RUN(test_socket_watchers_grow_and_support_fd_zero); + TEST_RUN(test_aborted_multi_action_schedules_client_reset); + TEST_RUN(test_rejects_socket_when_interface_binding_fails); + return test_summary(); +} diff --git a/tests/unit/test_options.c b/tests/unit/test_options.c new file mode 100644 index 0000000..a51cb10 --- /dev/null +++ b/tests/unit/test_options.c @@ -0,0 +1,113 @@ +#include +#include + +#include "test_harness.h" + +static uid_t effective_uid; + +static uid_t test_geteuid(void) { + return effective_uid; +} + +static uid_t test_getuid(void) { + return 1000; +} + +static struct passwd *test_getpwnam(const char __attribute__((unused)) *name) { + static struct passwd user = { + .pw_name = "proxy", + .pw_uid = 1000, + .pw_gid = 1000 + }; + return &user; +} + +#define getuid test_getuid +#define geteuid test_geteuid +#define getpwnam test_getpwnam +#define OUTBOUND_INTERFACE_TEST +#include "../../src/options.c" +#undef OUTBOUND_INTERFACE_TEST +#undef getpwnam +#undef getuid +#undef geteuid + +static options_t options; + +static void setUp(void) { + options_init(&options); + effective_uid = 1000; + optind = 1; +#ifdef __APPLE__ + optreset = 1; +#endif +} + +static void tearDown(void) { + options_cleanup(&options); +} + +static void test_parses_outbound_interface(void) { + char *arguments[] = { + "https_dns_proxy", "-S", "192.0.2.10", + "--outbound-interface", "wg0", NULL + }; + + TEST_ASSERT(options_parse_args(&options, 5, arguments) == OPR_SUCCESS); + TEST_ASSERT(options.source_addr != NULL); + TEST_ASSERT(strcmp(options.source_addr, "192.0.2.10") == 0); + TEST_ASSERT(options.outbound_interface != NULL); + TEST_ASSERT(strcmp(options.outbound_interface, "wg0") == 0); +} + +static void test_rejects_interface_with_internal_user_drop(void) { + char *arguments[] = { + "https_dns_proxy", "-u", "proxy", "-I", "wg0", NULL + }; + + TEST_ASSERT(options_parse_args(&options, 5, arguments) == OPR_OPTION_ERROR); +} + +static void test_allows_interface_with_proxy(void) { + char *arguments[] = { + "https_dns_proxy", "-t", "socks5h://proxy.example", "-I", "wg0", NULL + }; + + TEST_ASSERT(options_parse_args(&options, 5, arguments) == OPR_SUCCESS); +} + +static void test_interface_rejects_source_interface_or_hostname(void) { + char *arguments[] = {"https_dns_proxy", "-I", "wg0", "-S", "eth0", NULL}; + TEST_ASSERT(options_parse_args(&options, 5, arguments) == OPR_OPTION_ERROR); +} + +static void test_interface_accepts_ipv6_source(void) { + char *arguments[] = {"https_dns_proxy", "-I", "wg0", "-S", "2001:db8::1", NULL}; + TEST_ASSERT(options_parse_args(&options, 5, arguments) == OPR_SUCCESS); +} + +static void test_nonroot_user_does_not_request_group_reset(void) { + char *arguments[] = {"https_dns_proxy", "-u", "proxy", NULL}; + TEST_ASSERT(options_parse_args(&options, 3, arguments) == OPR_SUCCESS); + TEST_ASSERT(options.uid == 1000); + TEST_ASSERT(options.gid == (gid_t)-1); +} + +static void test_root_user_uses_primary_group(void) { + effective_uid = 0; + char *arguments[] = {"https_dns_proxy", "-u", "proxy", NULL}; + TEST_ASSERT(options_parse_args(&options, 3, arguments) == OPR_SUCCESS); + TEST_ASSERT(options.uid == 1000); + TEST_ASSERT(options.gid == 1000); +} + +int main(void) { + TEST_RUN(test_parses_outbound_interface); + TEST_RUN(test_rejects_interface_with_internal_user_drop); + TEST_RUN(test_allows_interface_with_proxy); + TEST_RUN(test_interface_rejects_source_interface_or_hostname); + TEST_RUN(test_interface_accepts_ipv6_source); + TEST_RUN(test_nonroot_user_does_not_request_group_reset); + TEST_RUN(test_root_user_uses_primary_group); + return test_summary(); +} diff --git a/tests/unit/test_outbound_interface.c b/tests/unit/test_outbound_interface.c new file mode 100644 index 0000000..e2fa87c --- /dev/null +++ b/tests/unit/test_outbound_interface.c @@ -0,0 +1,91 @@ +#include +#include +#include +#include +#include + +#include "test_harness.h" + +static int setsockopt_calls; +static int setsockopt_result; +static int captured_socket; +static int captured_level; +static int captured_option; +static char captured_interface[IFNAMSIZ]; + +#ifndef SO_BINDTODEVICE +#define SO_BINDTODEVICE 25 +#endif + +static int test_setsockopt(int socket_fd, int level, int option, + const void *value, socklen_t value_length) { + setsockopt_calls++; + captured_socket = socket_fd; + captured_level = level; + captured_option = option; + const size_t copy_length = value_length < sizeof(captured_interface) + ? value_length : sizeof(captured_interface); + memcpy(captured_interface, value, copy_length); + if (setsockopt_result != 0) { + errno = EPERM; + } + return setsockopt_result; +} + +#define setsockopt test_setsockopt +#define OUTBOUND_INTERFACE_TEST +#include "../../src/outbound_interface.c" +#undef OUTBOUND_INTERFACE_TEST +#undef setsockopt + +static void setUp(void) { + setsockopt_calls = 0; + setsockopt_result = 0; + captured_socket = -1; + captured_level = -1; + captured_option = -1; + memset(captured_interface, 0, sizeof(captured_interface)); +} + +static void tearDown(void) { +} + +static void test_rejects_empty_interface(void) { + TEST_ASSERT(outbound_interface_bind_socket(7, "") == -1); + TEST_ASSERT(errno == EINVAL); + TEST_ASSERT(setsockopt_calls == 0); +} + +static void test_rejects_overlong_interface(void) { + char interface_name[IFNAMSIZ + 1]; + memset(interface_name, 'a', sizeof(interface_name)); + interface_name[sizeof(interface_name) - 1] = '\0'; + + TEST_ASSERT(outbound_interface_bind_socket(7, interface_name) == -1); + TEST_ASSERT(errno == ENAMETOOLONG); + TEST_ASSERT(setsockopt_calls == 0); +} + +static void test_binds_socket_to_exact_interface(void) { + TEST_ASSERT(outbound_interface_bind_socket(7, "wg0") == 0); + TEST_ASSERT(setsockopt_calls == 1); + TEST_ASSERT(captured_socket == 7); + TEST_ASSERT(captured_level == SOL_SOCKET); + TEST_ASSERT(captured_option == SO_BINDTODEVICE); + TEST_ASSERT(strcmp(captured_interface, "wg0") == 0); +} + +static void test_propagates_bind_permission_failure(void) { + setsockopt_result = -1; + TEST_ASSERT(outbound_interface_bind_socket(7, "wg0") == -1); + TEST_ASSERT(errno == EPERM); + TEST_ASSERT(setsockopt_calls == 1); +} + +int main(void) { + TEST_RUN(test_rejects_empty_interface); + TEST_RUN(test_rejects_overlong_interface); + TEST_RUN(test_binds_socket_to_exact_interface); + TEST_RUN(test_propagates_bind_permission_failure); + return test_summary(); +}