From 69773bd16ad87f91ead7e53a79f4546f8b51cfe0 Mon Sep 17 00:00:00 2001 From: TapTap Date: Wed, 29 Jul 2026 18:52:08 +0200 Subject: [PATCH] fix: resolve all security issues (#154, #156, #157, #159, #160, #161, #162, #170) --- src/server/server.c | 2 +- src/shared/compression.c | 11 +++++-- src/shared/delta.c | 21 ++++++++++++ src/shared/file.c | 61 +++++++++++++++++++++++++++++++++-- src/shared/protocol.c | 6 ++-- src/shared/protocol.h | 7 ++-- src/shared/transport_tls.c | 43 ++++++++++++++++-------- src/shared/utils.c | 8 ++++- tests/integration/test_tls.py | 10 +++++- 9 files changed, 143 insertions(+), 26 deletions(-) diff --git a/src/server/server.c b/src/server/server.c index 2e045ca..832307d 100644 --- a/src/server/server.c +++ b/src/server/server.c @@ -69,7 +69,7 @@ int receive_files(Config* config, int fd) { } char* full_path = path_cat(config->receive_root_directory, check_path); struct stat st; - bool has_old = full_path && stat(full_path, &st) == 0; + bool has_old = full_path && lstat(full_path, &st) == 0; bool match = has_old && (unsigned long long)st.st_size == check_size && (long long)st.st_mtime == check_mtime; if (match) diff --git a/src/shared/compression.c b/src/shared/compression.c index 97ae61f..c1eb1f1 100644 --- a/src/shared/compression.c +++ b/src/shared/compression.c @@ -78,14 +78,21 @@ Data* data_decompress(Data* compressed_data) { return NULL; } + // ZSTD_CONTENTSIZE_UNKNOWN (~2^64) can cause massive allocation; + // fall back to a conservative estimate (3x compressed size) when unknown. + if (dst_size == ZSTD_CONTENTSIZE_UNKNOWN) { + dst_size = compressed_data->size * 3; + if (dst_size < INITIAL_DECOMPRESS_BUF_SIZE) + dst_size = INITIAL_DECOMPRESS_BUF_SIZE; + } + ZSTD_DCtx* dctx = ZSTD_createDCtx(); if (!dctx) { log_message(LOG_LEVEL_ERROR, "Failed to create ZSTD decompression context"); return NULL; } - size_t buf_size = - (!ZSTD_isError(dst_size) && dst_size > 0) ? (size_t)dst_size : INITIAL_DECOMPRESS_BUF_SIZE; + size_t buf_size = (dst_size > 0) ? (size_t)dst_size : INITIAL_DECOMPRESS_BUF_SIZE; Data* uncompressed_data = data_create_empty(buf_size); if (!uncompressed_data) { log_message(LOG_LEVEL_ERROR, "Failed to allocate decompression buffer"); diff --git a/src/shared/delta.c b/src/shared/delta.c index 206ee78..48a3f97 100644 --- a/src/shared/delta.c +++ b/src/shared/delta.c @@ -1,5 +1,6 @@ #include "delta.h" #include "log.h" +#include #include #include @@ -7,6 +8,10 @@ #define XXH_IMPLEMENTATION #include +/* Maximum number of blocks/instructions allowed from the wire to prevent OOM */ +#define MAX_DELTA_BLOCKS (1024U * 1024U) /* 1M signature blocks */ +#define MAX_DELTA_INSTRUCTIONS (1024U * 1024U) /* 1M delta instructions */ + uint32_t delta_adler32(const void* data, uint32_t len) { const uint8_t* p = (const uint8_t*)data; uint32_t s1 = 1; @@ -101,6 +106,14 @@ DeltaSignature* delta_signature_deserialize(const Data* data) { memcpy(&sig->block_count, buf + pos, sizeof(uint32_t)); pos += sizeof(uint32_t); + // Reject unreasonably large block counts to prevent OOM + if (sig->block_count > MAX_DELTA_BLOCKS) { + log_message(LOG_LEVEL_ERROR, "Delta signature block count %u exceeds maximum %u", + sig->block_count, MAX_DELTA_BLOCKS); + free(sig); + return NULL; + } + uint64_t expected = sizeof(uint64_t) + sizeof(uint32_t) + sizeof(uint32_t) + (uint64_t)sig->block_count * (sizeof(uint32_t) + sizeof(uint32_t)); if (data->size < expected) { @@ -340,6 +353,14 @@ Delta* delta_deserialize(const Data* data) { memcpy(&delta->instruction_count, buf + pos, sizeof(uint32_t)); pos += sizeof(uint32_t); + // Reject unreasonably large instruction counts to prevent OOM + if (delta->instruction_count > MAX_DELTA_INSTRUCTIONS) { + log_message(LOG_LEVEL_ERROR, "Delta instruction count %u exceeds maximum %u", + delta->instruction_count, MAX_DELTA_INSTRUCTIONS); + free(delta); + return NULL; + } + delta->instructions = malloc(delta->instruction_count * sizeof(DeltaInstruction)); if (!delta->instructions) { free(delta); diff --git a/src/shared/file.c b/src/shared/file.c index a531e3d..55ff441 100644 --- a/src/shared/file.c +++ b/src/shared/file.c @@ -134,9 +134,64 @@ bool file_save_to_disk(const char* root_directory, File* file, const Config* con log_message(LOG_LEVEL_ERROR, "Path traversal detected in file path: %s", file->path); return false; } - char* disk_path = path_cat((char*)root_directory, file->path); - if (disk_path == NULL) + + // Resolve the destination root to its real path, preventing symlink-based escapes. + // If the root does not yet exist, try to create it so realpath can succeed. + char* resolved_root = realpath(root_directory, NULL); + if (resolved_root == NULL) { + if (mkdir_r(root_directory)) { + resolved_root = realpath(root_directory, NULL); + } + } + if (resolved_root == NULL) { + log_message(LOG_LEVEL_ERROR, "Failed to resolve destination root: %s", root_directory); return false; + } + + char* disk_path = path_cat(resolved_root, file->path); + if (disk_path == NULL) { + free(resolved_root); + return false; + } + + // Ensure the target directory exists so the parent can be resolved for path safety. + char* dir_dup = str_dup(disk_path); + if (!dir_dup) { + free(resolved_root); + free(disk_path); + return false; + } + char* dir_str = dirname(dir_dup); + // Create the directory if needed (no-op if it already exists) so realpath can resolve it. + if (!mkdir_r(dir_str)) { + free(dir_dup); + free(resolved_root); + free(disk_path); + return false; + } + char* resolved_dir = realpath(dir_str, NULL); + free(dir_dup); + if (resolved_dir == NULL) { + log_message(LOG_LEVEL_ERROR, "Failed to resolve directory for: %s", disk_path); + free(resolved_root); + free(disk_path); + return false; + } + + // Verify that the resolved directory is inside the resolved root. + // Both are canonical absolute paths — this prevents symlink-based escapes. + size_t root_len = strlen(resolved_root); + if (strncmp(resolved_dir, resolved_root, root_len) != 0 || + (resolved_dir[root_len] != '\0' && resolved_dir[root_len] != '/')) { + log_message(LOG_LEVEL_ERROR, "Path escape detected: %s is outside %s", disk_path, root_directory); + free(resolved_dir); + free(resolved_root); + free(disk_path); + return false; + } + free(resolved_dir); + free(resolved_root); + bool ok = to_disk(disk_path, file->data->data, file->data->size); if (ok) file_restore_metadata(disk_path, file->metadata); @@ -341,7 +396,7 @@ File* receive_incremental_check(int fd, const Config* config, bool* skipped) { char* full_path = path_cat(config->receive_root_directory, check_path); struct stat st; - bool has_old_file = (full_path && stat(full_path, &st) == 0); + bool has_old_file = (full_path && lstat(full_path, &st) == 0); unsigned long long old_size = has_old_file ? (unsigned long long)st.st_size : 0; bool match = has_old_file && (unsigned long long)st.st_size == check_size && diff --git a/src/shared/protocol.c b/src/shared/protocol.c index 8bd7cfa..18f5a96 100644 --- a/src/shared/protocol.c +++ b/src/shared/protocol.c @@ -9,7 +9,7 @@ #include #include -#define MAX_DATA_SIZE (256ULL * 1024 * 1024) /* 256 MB max per message */ +#define MAX_DATA_SIZE (100ULL * 1024 * 1024) /* 100 MB max per data message */ #define RECEIVE_TIMEOUT_SEC 60 /* 60 second per-message timeout */ #define MAX_CONNECTION_MEMORY (1024ULL * 1024 * 1024) /* 1 GB total per connection */ @@ -186,9 +186,9 @@ char* receive_str(int file_descriptor) { size_t size; if (!receive_n_data(file_descriptor, &size, sizeof(size_t))) return NULL; - if (size > MAX_DATA_SIZE) { + if (size > MAX_STRING_SIZE) { log_message(LOG_LEVEL_ERROR, "String size %zu exceeds maximum %llu", size, - (unsigned long long)MAX_DATA_SIZE); + (unsigned long long)MAX_STRING_SIZE); return NULL; } char* data = (char*)malloc(size + 1); diff --git a/src/shared/protocol.h b/src/shared/protocol.h index 405832e..f4133c1 100644 --- a/src/shared/protocol.h +++ b/src/shared/protocol.h @@ -5,8 +5,11 @@ #include #include -/* Maximum allowed string size for receive_str (10 MB) */ -#define MAX_STRING_SIZE (10 * 1024 * 1024) +/* Maximum allowed string size for receive_str (64 KB) */ +#define MAX_STRING_SIZE (64 * 1024) + +/* Maximum allowed data payload size for receive_data (100 MB) */ +#define MAX_DATA_PAYLOAD_SIZE (100ULL * 1024 * 1024) typedef struct ssl_st SSL; diff --git a/src/shared/transport_tls.c b/src/shared/transport_tls.c index 7959ee7..f245a68 100644 --- a/src/shared/transport_tls.c +++ b/src/shared/transport_tls.c @@ -79,25 +79,38 @@ static SSL_CTX* create_ssl_ctx(bool is_server, const char* cert, const char* key return ctx; } -static SSL* wrap_fd_with_ssl(int fd, SSL_CTX* ctx, bool is_server) { +static SSL* wrap_fd_with_ssl(int fd, SSL_CTX* ctx, bool is_server, const char* hostname) { SSL* ssl = SSL_new(ctx); if (!ssl) { log_message(LOG_LEVEL_ERROR, "Failed to create SSL object"); return NULL; } SSL_set_fd(ssl, fd); - int ret; - if (is_server) - ret = SSL_accept(ssl); - else - ret = SSL_connect(ssl); - if (ret <= 0) { - log_message(LOG_LEVEL_ERROR, "SSL %s failed", is_server ? "accept" : "connect"); - log_ssl_errors(); - SSL_free(ssl); - return NULL; + // Enable hostname verification for client connections when a hostname is provided. + // Must be done before SSL_connect to take effect during the handshake. + if (!is_server && hostname) { + SSL_set1_host(ssl, hostname); } + + // Retry SSL_accept/SSL_connect on WANT_READ/WANT_WRITE (non-blocking handshake) + int ret; + do { + if (is_server) + ret = SSL_accept(ssl); + else + ret = SSL_connect(ssl); + + if (ret <= 0) { + int ssl_err = SSL_get_error(ssl, ret); + if (ssl_err == SSL_ERROR_WANT_READ || ssl_err == SSL_ERROR_WANT_WRITE) + continue; + log_message(LOG_LEVEL_ERROR, "SSL %s failed", is_server ? "accept" : "connect"); + log_ssl_errors(); + SSL_free(ssl); + return NULL; + } + } while (ret <= 0); return ssl; } @@ -117,7 +130,7 @@ struct tls_child_ctx { static void tls_child_fn(int fd, void* arg) { struct tls_child_ctx* ctx = (struct tls_child_ctx*)arg; - SSL* ssl = wrap_fd_with_ssl(fd, ctx->ssl_ctx, true); + SSL* ssl = wrap_fd_with_ssl(fd, ctx->ssl_ctx, true, NULL); if (!ssl) return; io_set_ssl(ssl); @@ -151,12 +164,16 @@ bool client_connect_tls(Client* client, char* host, int port, const char* cert_p return false; client->ssl_ctx = ctx; - SSL* ssl = wrap_fd_with_ssl(client->file_descriptor, ctx, false); + // Pass the server hostname for TLS hostname verification (SSL_set1_host + // is called inside wrap_fd_with_ssl before the handshake when ca_path is set). + const char* verify_host = ca_path ? host : NULL; + SSL* ssl = wrap_fd_with_ssl(client->file_descriptor, ctx, false, verify_host); if (!ssl) { SSL_CTX_free(ctx); client->ssl_ctx = NULL; return false; } + client->ssl = ssl; io_set_ssl(ssl); return true; diff --git a/src/shared/utils.c b/src/shared/utils.c index 8adc37b..755d695 100644 --- a/src/shared/utils.c +++ b/src/shared/utils.c @@ -126,7 +126,13 @@ static void delete_extras_walk(const char* abs_path, const char* rel_path, Array char* child_abs = path_cat((char*)abs_path, entry->d_name); char* child_rel = path_cat((char*)rel_path, entry->d_name); struct stat st; - if (stat(child_abs, &st) != 0) { + if (lstat(child_abs, &st) != 0) { + free(child_abs); + free(child_rel); + continue; + } + // Skip symlinks to prevent following them outside the destination tree + if (S_ISLNK(st.st_mode)) { free(child_abs); free(child_rel); continue; diff --git a/tests/integration/test_tls.py b/tests/integration/test_tls.py index 25d31cd..a59edc6 100644 --- a/tests/integration/test_tls.py +++ b/tests/integration/test_tls.py @@ -37,15 +37,23 @@ def _generate_certs(cert_dir): ], check=True, capture_output=True) # Server key + CSR + cert (signed by CA) + # Use a config file to include IP SAN 127.0.0.1 so hostname verification passes + san_config = os.path.join(cert_dir, "server_san.conf") + with open(san_config, "w") as f: + f.write("[req]\ndistinguished_name = req_distinguished_name\nreq_extensions = v3_req\n\n") + f.write("[req_distinguished_name]\nCN = localhost\n\n") + f.write("[v3_req]\nsubjectAltName = @alt_names\n\n") + f.write("[alt_names]\nDNS.1 = localhost\nIP.1 = 127.0.0.1\n") subprocess.run([ "openssl", "req", "-newkey", "rsa:2048", "-nodes", "-keyout", server_key, "-out", os.path.join(cert_dir, "server.csr"), - "-subj", "/CN=localhost", + "-subj", "/CN=localhost", "-config", san_config, ], check=True, capture_output=True) subprocess.run([ "openssl", "x509", "-req", "-in", os.path.join(cert_dir, "server.csr"), "-CA", ca_cert, "-CAkey", ca_key, "-CAcreateserial", "-out", server_cert, "-days", "1", + "-extfile", san_config, "-extensions", "v3_req", ], check=True, capture_output=True) # Client key + CSR + cert (signed by CA)