diff --git a/README.md b/README.md index 8868dc8..1386d72 100644 --- a/README.md +++ b/README.md @@ -103,6 +103,7 @@ partial, alternate, and planned behavior. | `--include ` | Only transfer files matching glob pattern (repeatable, whitelist) | | `--max-size ` | Skip files larger than n bytes | | `--min-size ` | Skip files smaller than n bytes | +| `--max-alloc ` | Maximum single allocation (binary units: B, K, M, G, T, P, E; default 1G) | | `--incremental` | Skip files unchanged since last transfer (size + mtime). Auto-enables `--preserve`. Incompatible with `-s`. | | `--existing` | Skip files not already present at the destination; update existing files normally. | | `--bwlimit ` | Bandwidth limit in kilobytes per second | @@ -478,11 +479,11 @@ defaults to the current directory. | ## Protocol and Security -FastSync protocol version `2.3.0` is shared by the client and server. The +FastSync protocol version `2.4.0` is shared by the client and server. The current protocol is sender-driven and includes configuration negotiation, -incremental checks, checksums, manifests, keep-alives, abort handling, and -FastSync-native delta messages. Client and server versions must currently -match exactly. +including the maximum allocation limit, incremental checks, checksums, +manifests, keep-alives, abort handling, and FastSync-native delta messages. +Client and server versions must currently match exactly. TLS provides encrypted TCP transport. Supplying `--ca` enables certificate verification; without it, traffic is encrypted but peer identity is not diff --git a/src/client/client_cli.c b/src/client/client_cli.c index bce7d4c..5f0ae93 100644 --- a/src/client/client_cli.c +++ b/src/client/client_cli.c @@ -191,6 +191,56 @@ static int parse_ull_arg(const char* val, unsigned long long* out, const char* o return 0; } +static int parse_size_arg(const char* value, unsigned long long* out) { + if (!value || *value < '0' || *value > '9') + return -1; + char* end; + errno = 0; + unsigned long long number = strtoull(value, &end, 10); + if (errno != 0 || end == value) + return -1; + unsigned long long multiplier = 1; + if (*end != '\0') { + if (end[1] != '\0') + return -1; + switch (*end) { + case 'b': + case 'B': + break; + case 'k': + case 'K': + multiplier = 1024ULL; + break; + case 'm': + case 'M': + multiplier = 1024ULL * 1024; + break; + case 'g': + case 'G': + multiplier = 1024ULL * 1024 * 1024; + break; + case 't': + case 'T': + multiplier = 1024ULL * 1024 * 1024 * 1024; + break; + case 'p': + case 'P': + multiplier = 1024ULL * 1024 * 1024 * 1024 * 1024; + break; + case 'e': + case 'E': + multiplier = 1024ULL * 1024 * 1024 * 1024 * 1024 * 1024; + break; + default: + return -1; + } + } + if (number == 0 || number > ULLONG_MAX / multiplier) + return -1; + *out = number * multiplier; + return 0; +} + /* Append a duplicated pattern to a growable pattern array. Returns 0 on success, -1 on error. */ static int config_add_pattern(char*** patterns, int* count, const char* value, const char* optname) { @@ -369,7 +419,7 @@ int parse_args(Config* config, int argc, char* argv[], int* positional_args, bool verbose = false; protocol_set_8_bit_output(config->eight_bit_output); for (int i = 1; i < argc; i++) { - const char* modify_window_prefix = "--modify-window="; +const char* modify_window_prefix = "--modify-window="; if (strncmp(argv[i], modify_window_prefix, strlen(modify_window_prefix)) == 0) { if (set_nonneg_int_option(&config->modify_window, argv[i] + strlen(modify_window_prefix), "--modify-window") != 0) @@ -388,6 +438,22 @@ int parse_args(Config* config, int argc, char* argv[], int* positional_args, return -1; continue; } + if (strncmp(argv[i], "--max-alloc=", 12) == 0 || strcmp(argv[i], "--max-alloc") == 0) { + const char* value = strcmp(argv[i], "--max-alloc") == 0 ? "" : argv[i] + 12; + if (*value == '\0') { + if (i + 1 >= argc) { + log_message(LOG_LEVEL_ERROR, "missing argument for --max-alloc"); + return -1; + } + value = argv[++i]; + } + if (parse_size_arg(value, &config->max_alloc) != 0) { + log_message(LOG_LEVEL_ERROR, + "--max-alloc must be a positive size (B, K, M, G, T, P, or E)"); + return -1; + } + continue; + } const OptionEntry* entry = find_table_option(argv[i]); const char* inline_value = NULL; if (!entry) diff --git a/src/client/client_send.c b/src/client/client_send.c index 0f017c7..af2c0ec 100644 --- a/src/client/client_send.c +++ b/src/client/client_send.c @@ -612,14 +612,16 @@ static int send_chunks_multithreaded(void* pipeline_context) { static int scan_directory_multithreaded(void* pipeline_context) { PipelineContextSender* context = (PipelineContextSender*)pipeline_context; + protocol_session_bind(&context->allocation_session); ScannerOptions options = scanner_options_from_config(context->config, 4); - ParallelScanner* scanner = - parallel_scanner_create_with_options(context->config->send_directory, &options); + ParallelScanner* scanner = parallel_scanner_create_with_options( + context->config->send_directory, &options, &context->allocation_session); Chunk* current_chunk; if (scanner == NULL) { log_message(LOG_LEVEL_ERROR, "Failed to create parallel scanner"); pipeline_cancel(context); + protocol_session_unbind(); return thrd_error; } while ((current_chunk = parallel_scanner_next(scanner)) != NULL) { @@ -631,6 +633,7 @@ static int scan_directory_multithreaded(void* pipeline_context) { pipeline_cancel(context); chunk_destroy(current_chunk); parallel_scanner_destroy(scanner); + protocol_session_unbind(); return thrd_error; } } @@ -641,6 +644,7 @@ static int scan_directory_multithreaded(void* pipeline_context) { chunk_destroy(current_chunk); pipeline_cancel(context); parallel_scanner_destroy(scanner); + protocol_session_unbind(); return thrd_error; } } @@ -652,6 +656,7 @@ static int scan_directory_multithreaded(void* pipeline_context) { cnd_broadcast(&context->condition_not_full_scanner); mtx_unlock(&context->mutex_scanner); pipeline_cancel(context); + protocol_session_unbind(); return thrd_error; } mtx_lock(&context->mutex_scanner); @@ -660,11 +665,13 @@ static int scan_directory_multithreaded(void* pipeline_context) { mtx_unlock(&context->mutex_scanner); parallel_scanner_destroy(scanner); + protocol_session_unbind(); return thrd_success; } static int load_files_multithreaded(void* pipeline_context) { PipelineContextSender* context = (PipelineContextSender*)pipeline_context; + protocol_session_bind(&context->allocation_session); while (true) { Chunk* chunk = queue_dequeue_multithreaded( context->queue_scanner, &context->mutex_scanner, &context->condition_not_empty_scanner, @@ -674,6 +681,7 @@ static int load_files_multithreaded(void* pipeline_context) { context->loader_done = true; cnd_signal(&context->condition_not_empty_loader); mtx_unlock(&context->mutex_loader); + protocol_session_unbind(); return thrd_success; } if (!context->config->use_sendfile) { @@ -685,6 +693,7 @@ static int load_files_multithreaded(void* pipeline_context) { log_message(LOG_LEVEL_ERROR, "Failed to load file data"); chunk_destroy(chunk); pipeline_cancel(context); + protocol_session_unbind(); return thrd_error; } } @@ -695,6 +704,7 @@ static int load_files_multithreaded(void* pipeline_context) { &context->cancelled)) { chunk_destroy(chunk); pipeline_cancel(context); + protocol_session_unbind(); return thrd_error; } } diff --git a/src/client/scanner.c b/src/client/scanner.c index 46b70c7..276ddfb 100644 --- a/src/client/scanner.c +++ b/src/client/scanner.c @@ -351,10 +351,14 @@ typedef struct { char** dirs; int dir_count; ScannerOptions options; + ProtocolSession* allocation_session; } ParallelWorkerArg; static int parallel_worker_thread(void* arg) { ParallelWorkerArg* wa = (ParallelWorkerArg*)arg; + ProtocolSession* allocation_session = wa->allocation_session; + if (allocation_session) + protocol_session_bind(allocation_session); for (int i = 0; i < wa->dir_count; i++) { DirectoryScanner* ds = directory_scanner_create_with_options(wa->dirs[i], &wa->options); if (!ds) { @@ -398,6 +402,8 @@ static int parallel_worker_thread(void* arg) { cnd_signal(&ps->result_not_empty); } mtx_unlock(&ps->result_mutex); + if (allocation_session) + protocol_session_unbind(); return thrd_success; } @@ -633,6 +639,7 @@ static void spawn_parallel_workers(ParallelScanner* ps, ArrayList* subdirs, wa->dir_count = count; wa->options = *options; wa->options.chunk_size = cs; + wa->allocation_session = ps->allocation_session; start += count; if (thrd_create(&ps->threads[t], parallel_worker_thread, wa) != thrd_success) { for (int j = 0; j < count; j++) @@ -648,7 +655,8 @@ static void spawn_parallel_workers(ParallelScanner* ps, ArrayList* subdirs, } ParallelScanner* parallel_scanner_create_with_options(const char* root_directory, - const ScannerOptions* options) { + const ScannerOptions* options, + ProtocolSession* allocation_session) { if (!root_directory || !options) return NULL; ParallelScanner* ps = calloc(1, sizeof(ParallelScanner)); @@ -658,6 +666,7 @@ ParallelScanner* parallel_scanner_create_with_options(const char* root_directory free(ps); return NULL; } + ps->allocation_session = allocation_session; ArrayList* root_files = array_list_create(file_destroy); ArrayList* subdirs = array_list_create(free); diff --git a/src/client/scanner.h b/src/client/scanner.h index 9bb24f8..f699218 100644 --- a/src/client/scanner.h +++ b/src/client/scanner.h @@ -2,6 +2,7 @@ #define SCANNER_H #include "chunk.h" +#include "protocol.h" #include "queue.h" #include #include @@ -62,6 +63,7 @@ typedef struct { atomic_bool cancelled; int completed; Chunk* initial_chunk; + ProtocolSession* allocation_session; } ParallelScanner; DirectoryScanner* directory_scanner_create(const char* root_directory, bool use_metadata, @@ -78,7 +80,8 @@ bool directory_scanner_failed(const DirectoryScanner* scanner); void directory_scanner_destroy(DirectoryScanner* scanner); ParallelScanner* parallel_scanner_create_with_options(const char* root_directory, - const ScannerOptions* options); + const ScannerOptions* options, + ProtocolSession* allocation_session); Chunk* parallel_scanner_next(ParallelScanner* scanner); bool parallel_scanner_failed(const ParallelScanner* scanner); void parallel_scanner_destroy(ParallelScanner* scanner); diff --git a/src/client/usage.c b/src/client/usage.c index e38b003..93f2413 100644 --- a/src/client/usage.c +++ b/src/client/usage.c @@ -30,6 +30,7 @@ void print_usage(void) { printf(" --include-from Read include patterns from file\n"); printf(" --max-size Skip files larger than n bytes\n"); printf(" --min-size Skip files smaller than n bytes\n"); + printf(" --max-alloc Maximum single allocation (default: 1G)\n"); printf(" --incremental Skip files unchanged since last transfer\n"); printf(" --size-only Skip incremental files matching in size, ignoring mtime\n"); printf(" -I, --ignore-times Transfer files even when size and mtime match\n"); diff --git a/src/server/server.c b/src/server/server.c index 42351a9..c968115 100644 --- a/src/server/server.c +++ b/src/server/server.c @@ -311,7 +311,9 @@ void handler(int file_descriptor) { protocol_session_unbind(); return; } - context->session.total_allocated_bytes = session.total_allocated_bytes; + protocol_session_set_max_alloc(&context->session, config->max_alloc); + atomic_store(&context->session.total_allocated_bytes, + atomic_load(&session.total_allocated_bytes)); thrd_t receiver, writer; bool receiver_created = thrd_create(&receiver, receive_thread, context) == thrd_success; bool writer_created = false; diff --git a/src/shared/array_list.c b/src/shared/array_list.c index 0974fd4..95799b0 100644 --- a/src/shared/array_list.c +++ b/src/shared/array_list.c @@ -1,17 +1,18 @@ #include "log.h" #include "array_list.h" +#include "protocol.h" #include #include #include ArrayList* array_list_create(void (*item_destroyer)(void* item)) { - ArrayList* list = (ArrayList*)malloc(sizeof(ArrayList)); + ArrayList* list = (ArrayList*)protocol_alloc(sizeof(ArrayList)); if (list == NULL) { log_perror("ERROR: Could not allocate memory for array list struct"); return NULL; } - list->items = malloc(INITIAL_ARRAY_SIZE * sizeof(void*)); + list->items = protocol_alloc(INITIAL_ARRAY_SIZE * sizeof(void*)); if (list->items == NULL) { free(list); return NULL; @@ -41,7 +42,7 @@ static bool array_list_extend(ArrayList* array_list) { int new_capacity = array_list->capacity * 2; if (new_capacity == 0) new_capacity = INITIAL_ARRAY_SIZE; - void* new_items = realloc(array_list->items, new_capacity * sizeof(void*)); + void* new_items = protocol_realloc(array_list->items, new_capacity * sizeof(void*)); if (new_items == NULL) { log_perror("ERROR: Could not reallocate memory for array list items"); return false; @@ -67,7 +68,7 @@ void** array_list_to_array(const ArrayList* array_list) { if (array_list == NULL) { return NULL; } - void** array = malloc(array_list->size * sizeof(void*)); + void** array = protocol_alloc(array_list->size * sizeof(void*)); if (array == NULL) { log_perror("Could not malloc space for array from array list!"); return NULL; diff --git a/src/shared/chunk.c b/src/shared/chunk.c index 37eb2ff..7c258ba 100644 --- a/src/shared/chunk.c +++ b/src/shared/chunk.c @@ -22,7 +22,7 @@ Chunk* chunk_create(File** items, int element_count) { if (element_count < 0 || (element_count > 0 && items == NULL)) return NULL; - Chunk* chunk = (Chunk*)malloc(sizeof(Chunk)); + Chunk* chunk = (Chunk*)protocol_alloc(sizeof(Chunk)); if (chunk == NULL) { log_perror("ERROR: Could not allocate memory for chunk structure"); return NULL; @@ -35,7 +35,7 @@ Chunk* chunk_create(File** items, int element_count) { free(chunk); return NULL; } - chunk->items = (File**)malloc((size_t)element_count * sizeof(File*)); + chunk->items = (File**)protocol_alloc((size_t)element_count * sizeof(File*)); if (chunk->items == NULL) { free(chunk); return NULL; @@ -158,7 +158,7 @@ Chunk* chunk_deserialize(Data* data, bool use_metadata) { array_list_delete(files); return NULL; } - char* path = malloc(path_len + 1); + char* path = protocol_alloc(path_len + 1); if (path == NULL) { log_perror("Could not allocate memory for file path"); array_list_delete(files); @@ -245,7 +245,7 @@ Chunk* chunk_deserialize(Data* data, bool use_metadata) { } size_t allocation_size = file_data_size > 0 ? file_data_size : 1; - void* file_data = malloc(allocation_size); + void* file_data = protocol_alloc(allocation_size); if (file_data == NULL) { log_perror("Could not allocate memory for file data"); file_destroy(file); diff --git a/src/shared/compression.c b/src/shared/compression.c index 96221b3..e4d5b95 100644 --- a/src/shared/compression.c +++ b/src/shared/compression.c @@ -1,6 +1,7 @@ #include "compression.h" #include "data.h" #include "log.h" +#include "protocol.h" #include #include #include @@ -182,7 +183,7 @@ Data* data_decompress_limited(Data* compressed_data, size_t maximum_size) { buf_size *= 2; if (buf_size > hard_limit) buf_size = (size_t)hard_limit; - void* new_data = realloc(uncompressed_data->data, buf_size); + void* new_data = protocol_realloc(uncompressed_data->data, buf_size); if (!new_data) { log_message(LOG_LEVEL_ERROR, "Failed to grow decompression buffer"); ZSTD_freeDCtx(dctx); diff --git a/src/shared/config.c b/src/shared/config.c index e0cb796..1a0e317 100644 --- a/src/shared/config.c +++ b/src/shared/config.c @@ -37,6 +37,7 @@ static void config_set_defaults(Config* config) { config->include_count = 0; config->max_size = 0; config->min_size = 0; + config->max_alloc = DEFAULT_MAX_ALLOC; config->use_incremental = false; config->ignore_times = false; config->size_only = false; @@ -155,9 +156,9 @@ static bool validate_received_config(const Config* config) { config->chunk_size > 0 && config->chunk_size <= MAX_CHUNK_SIZE && config->delta_block_size >= DELTA_BLOCK_SIZE_MIN && config->delta_block_size <= DELTA_BLOCK_SIZE_MAX && - config->delta_max_file_size <= DELTA_MAX_FILE_SIZE && config->modify_window >= 0 && +config->delta_max_file_size <= DELTA_MAX_FILE_SIZE && config->modify_window >= 0 && config->max_delete >= 0 && config->skip_compress_count >= 0 && - config->skip_compress_count <= 10000 && + config->skip_compress_count <= 10000 && config->max_alloc > 0 && (!config->chmod_spec || !*config->chmod_spec || chmod_apply(0, config->chmod_spec, &(mode_t){0})); } @@ -252,6 +253,8 @@ static bool send_core_fields(int fd, const Config* c) { if (!send_str(fd, c->version) || !send_int(fd, c->eight_bit_output)) return false; protocol_set_8_bit_output(c->eight_bit_output); + if (!send_n_data(fd, &c->max_alloc, sizeof(c->max_alloc))) + return false; return send_str(fd, c->send_directory) && send_str(fd, c->receive_root_directory) && send_int(fd, c->save_to_disk) && send_int(fd, c->use_multithreading) && send_int(fd, c->use_chunk_serialization) && send_int(fd, c->use_compression) && @@ -311,6 +314,11 @@ static bool receive_core_fields(int fd, Config* c) { if (!receive_wire_bool(fd, &c->eight_bit_output)) return false; protocol_set_8_bit_output(c->eight_bit_output); + if (!receive_n_data(fd, &c->max_alloc, sizeof(c->max_alloc)) || c->max_alloc == 0) + return false; + if (c->max_alloc > MAX_SERVER_ALLOC) + c->max_alloc = MAX_SERVER_ALLOC; + protocol_session_set_max_alloc(NULL, c->max_alloc); c->send_directory = receive_str(fd); c->receive_root_directory = receive_str(fd); if (!c->send_directory || !c->receive_root_directory) @@ -413,6 +421,7 @@ static bool receive_resume_options(int fd, Config* c) { } bool config_send(int file_descriptor, const Config* config) { + protocol_session_set_max_alloc(NULL, config->max_alloc); if (!send_core_fields(file_descriptor, config) || !send_delta_fields(file_descriptor, config) || !send_file_options(file_descriptor, config) || !send_selection_options(file_descriptor, config) || diff --git a/src/shared/config.h b/src/shared/config.h index 1499efc..84994e1 100644 --- a/src/shared/config.h +++ b/src/shared/config.h @@ -36,6 +36,7 @@ typedef struct Config { int include_count; unsigned long long max_size; unsigned long long min_size; + unsigned long long max_alloc; bool use_incremental; bool ignore_times; bool size_only; @@ -144,7 +145,7 @@ typedef struct Config { bool skip_compress_set; } Config; -#define PROTOCOL_VERSION "2.3.0" +#define PROTOCOL_VERSION "2.4.0" #define DEFAULT_CHUNK_SIZE (10 * 1024 * 1024) Config* config_create(void); diff --git a/src/shared/data.c b/src/shared/data.c index dd43e15..55af503 100644 --- a/src/shared/data.c +++ b/src/shared/data.c @@ -1,11 +1,12 @@ #include "data.h" #include "log.h" +#include "protocol.h" #include Data* data_create_empty(size_t data_size) { /* malloc(0) is UB; allocate at least 1 byte but preserve requested size */ size_t alloc_size = data_size > 0 ? data_size : 1; - void* data = malloc(alloc_size); + void* data = protocol_alloc(alloc_size); if (data == NULL) { log_message(LOG_LEVEL_ERROR, "Could not allocate memory for empty data"); return NULL; @@ -14,7 +15,7 @@ Data* data_create_empty(size_t data_size) { } Data* data_create_reserve(size_t size) { - Data* d = malloc(sizeof(Data)); + Data* d = protocol_alloc(sizeof(Data)); if (d == NULL) { log_message(LOG_LEVEL_ERROR, "Could not allocate memory for data"); return NULL; @@ -26,7 +27,7 @@ Data* data_create_reserve(size_t size) { } Data* data_create(void* data, size_t data_size) { - Data* new_data = malloc(sizeof(Data)); + Data* new_data = protocol_alloc(sizeof(Data)); if (new_data == NULL) { log_message(LOG_LEVEL_ERROR, "Could not allocate memory for data"); free(data); diff --git a/src/shared/delta.c b/src/shared/delta.c index fc941fb..776ef90 100644 --- a/src/shared/delta.c +++ b/src/shared/delta.c @@ -1,5 +1,6 @@ #include "delta.h" #include "log.h" +#include "protocol.h" #include #include #include @@ -43,7 +44,7 @@ DeltaSignature* delta_signature_create(const void* old_file_data, uint64_t old_f uint32_t block_count = (uint32_t)((old_file_size + block_size - 1) / block_size); - DeltaSignature* sig = malloc(sizeof(DeltaSignature)); + DeltaSignature* sig = protocol_alloc(sizeof(DeltaSignature)); if (!sig) return NULL; @@ -54,7 +55,7 @@ DeltaSignature* delta_signature_create(const void* old_file_data, uint64_t old_f free(sig); return NULL; } - sig->blocks = malloc((size_t)block_count * sizeof(DeltaBlockSig)); + sig->blocks = protocol_alloc((size_t)block_count * sizeof(DeltaBlockSig)); if (!sig->blocks) { free(sig); return NULL; @@ -82,7 +83,7 @@ Data* delta_signature_serialize(const DeltaSignature* sig) { total > SIZE_MAX) return NULL; - uint8_t* buf = malloc((size_t)total); + uint8_t* buf = protocol_alloc((size_t)total); if (!buf) return NULL; @@ -111,7 +112,7 @@ DeltaSignature* delta_signature_deserialize(const Data* data) { const uint8_t* buf = (const uint8_t*)data->data; size_t pos = 0; - DeltaSignature* sig = malloc(sizeof(DeltaSignature)); + DeltaSignature* sig = protocol_alloc(sizeof(DeltaSignature)); if (!sig) return NULL; @@ -149,7 +150,7 @@ DeltaSignature* delta_signature_deserialize(const Data* data) { free(sig); return NULL; } - sig->blocks = malloc((size_t)blocks_size); + sig->blocks = protocol_alloc((size_t)blocks_size); if (!sig->blocks) { free(sig); return NULL; @@ -178,7 +179,7 @@ static bool ensure_capacity(DeltaInstruction** instrs, uint32_t* capacity, uint3 if (*capacity > MAX_DELTA_INSTRUCTIONS / 2) return false; uint32_t new_cap = *capacity * 2; - DeltaInstruction* tmp = realloc(*instrs, (size_t)new_cap * sizeof(DeltaInstruction)); + DeltaInstruction* tmp = protocol_realloc(*instrs, (size_t)new_cap * sizeof(DeltaInstruction)); if (!tmp) return false; *instrs = tmp; @@ -195,7 +196,7 @@ static bool flush_literal(DeltaInstruction** instrs, uint32_t* capacity, uint32_ uint32_t lit_len = (uint32_t)(end - start); if (!ensure_capacity(instrs, capacity, *count)) return false; - uint8_t* lit_data = malloc(lit_len); + uint8_t* lit_data = protocol_alloc(lit_len); if (!lit_data) return false; memcpy(lit_data, data + start, lit_len); @@ -225,7 +226,7 @@ Delta* delta_compute(const void* new_file_data, uint64_t new_file_size, const De uint32_t capacity = 64; uint32_t count = 0; - DeltaInstruction* instrs = malloc((size_t)capacity * sizeof(DeltaInstruction)); + DeltaInstruction* instrs = protocol_alloc((size_t)capacity * sizeof(DeltaInstruction)); if (!instrs) return NULL; @@ -309,7 +310,7 @@ Delta* delta_compute(const void* new_file_data, uint64_t new_file_size, const De } } - Delta* delta = malloc(sizeof(Delta)); + Delta* delta = protocol_alloc(sizeof(Delta)); if (!delta) { free_instructions(instrs, count); return NULL; @@ -355,7 +356,7 @@ Data* delta_serialize(const Delta* delta) { if (delta->delta_size > UINT64_MAX - header_size || header_size + delta->delta_size > SIZE_MAX) return NULL; uint64_t total = header_size + delta->delta_size; - uint8_t* buf = malloc((size_t)total); + uint8_t* buf = protocol_alloc((size_t)total); if (!buf) return NULL; @@ -395,7 +396,7 @@ Delta* delta_deserialize(const Data* data) { const uint8_t* buf = (const uint8_t*)data->data; size_t pos = 0; - Delta* delta = malloc(sizeof(Delta)); + Delta* delta = protocol_alloc(sizeof(Delta)); if (!delta) return NULL; @@ -412,9 +413,10 @@ Delta* delta_deserialize(const Data* data) { return NULL; } - delta->instructions = delta->instruction_count == 0 - ? NULL - : malloc((size_t)delta->instruction_count * sizeof(DeltaInstruction)); + delta->instructions = + delta->instruction_count == 0 + ? NULL + : protocol_alloc((size_t)delta->instruction_count * sizeof(DeltaInstruction)); if (delta->instruction_count > 0 && !delta->instructions) { free(delta); return NULL; @@ -465,7 +467,7 @@ Delta* delta_deserialize(const Data* data) { free(delta); return NULL; } - delta->instructions[i].literal.data = malloc(lit_len ? lit_len : 1); + delta->instructions[i].literal.data = protocol_alloc(lit_len ? lit_len : 1); if (!delta->instructions[i].literal.data) { log_message(LOG_LEVEL_ERROR, "Failed to allocate %u bytes for literal data", lit_len); free_instructions(delta->instructions, i); @@ -492,7 +494,7 @@ void* delta_apply(const void* old_data, uint64_t old_size, const Delta* delta, delta->new_file_size > DELTA_MAX_FILE_SIZE || delta->new_file_size > SIZE_MAX) return NULL; - void* output = malloc(delta->new_file_size ? (size_t)delta->new_file_size : 1); + void* output = protocol_alloc(delta->new_file_size ? (size_t)delta->new_file_size : 1); if (!output) return NULL; diff --git a/src/shared/file.c b/src/shared/file.c index 60fe37a..aa4b95c 100644 --- a/src/shared/file.c +++ b/src/shared/file.c @@ -14,6 +14,7 @@ #include "log.h" #include "metadata.h" #include "utils.h" +#include "protocol.h" static bool write_all(int fd, const void* data, unsigned long long size) { const unsigned char* p = data; @@ -45,14 +46,14 @@ bool file_checksum(File* file, uint64_t* checksum) { File* file_create(const char* path) { if (!path) return NULL; - File* file = (File*)malloc(sizeof(File)); + File* file = (File*)protocol_alloc(sizeof(File)); if (file == NULL) { log_perror("ERROR: Could not allocate memory for file struct"); return NULL; } size_t path_len = strlen(path); - file->path = (char*)malloc(path_len + 1); + file->path = (char*)protocol_alloc(path_len + 1); if (file->path == NULL) { free(file); return NULL; @@ -85,7 +86,7 @@ void file_destroy(void* item) { } FileMetadata* file_metadata_create(const struct stat* stats) { - FileMetadata* m = malloc(sizeof(FileMetadata)); + FileMetadata* m = protocol_alloc(sizeof(FileMetadata)); if (m == NULL) { log_perror("ERROR: Could not allocate memory for file metadata"); return NULL; @@ -112,7 +113,7 @@ bool file_load_data(File* file) { if (file->data->data == NULL) { if (file->data->size == 0) return true; - file->data->data = malloc(file->data->size); + file->data->data = protocol_alloc(file->data->size); if (file->data->data == NULL) { log_perror("Could not allocate memory for file data"); return false; diff --git a/src/shared/file_receive.c b/src/shared/file_receive.c index 04c6fda..dd4e1a8 100644 --- a/src/shared/file_receive.c +++ b/src/shared/file_receive.c @@ -421,7 +421,7 @@ File* receive_incremental_check(int fd, const Config* config, bool* skipped) { unsigned long long old_size = has_old_file ? (unsigned long long)st.st_size : 0; void* old_data = NULL; if (has_old_file && old_size > 0 && old_size <= MAX_RECEIVE_FILE_SIZE && old_size <= SIZE_MAX) { - old_data = malloc((size_t)old_size); + old_data = protocol_alloc((size_t)old_size); if (old_data) { size_t got = 0; while (got < (size_t)old_size) { diff --git a/src/shared/metadata.c b/src/shared/metadata.c index 6f91abd..d4a251e 100644 --- a/src/shared/metadata.c +++ b/src/shared/metadata.c @@ -77,7 +77,7 @@ FileMetadata* metadata_from_buf(char** buf) { return NULL; if (!present) return NULL; - FileMetadata* m = malloc(sizeof(FileMetadata)); + FileMetadata* m = protocol_alloc(sizeof(FileMetadata)); if (m == NULL) return NULL; int32_t mode; @@ -144,7 +144,7 @@ FileMetadata* metadata_receive(int file_descriptor, int* ok) { *ok = 0; return NULL; } - FileMetadata* m = malloc(sizeof(FileMetadata)); + FileMetadata* m = protocol_alloc(sizeof(FileMetadata)); if (m == NULL) { if (ok) *ok = 0; diff --git a/src/shared/multiprocessing.c b/src/shared/multiprocessing.c index 47dd5fc..5944ef6 100644 --- a/src/shared/multiprocessing.c +++ b/src/shared/multiprocessing.c @@ -30,6 +30,8 @@ PipelineContextSender* pipeline_context_sender_create(Config* config, Queue* que context->progress_bytes = 0; context->sender_done = false; atomic_init(&context->cancelled, false); + protocol_session_init(&context->allocation_session, -1, -1); + protocol_session_set_max_alloc(&context->allocation_session, config->max_alloc); int init = 0; if (mtx_init(&context->mutex_scanner, mtx_plain) != thrd_success) goto fail; @@ -182,6 +184,7 @@ int receive_thread(void* pipeline_context) { int write_thread(void* pipeline_context) { PipelineContextReceiver* context = (PipelineContextReceiver*)pipeline_context; + protocol_session_bind(&context->session); mtx_lock(&context->mutex); bool save_to_disk = context->config->save_to_disk; char* root_directory = str_dup(context->config->receive_root_directory); @@ -193,6 +196,7 @@ int write_thread(void* pipeline_context) { cnd_broadcast(&context->condition_not_full); cnd_broadcast(&context->condition_not_empty); mtx_unlock(&context->mutex); + protocol_session_unbind(); return thrd_error; } @@ -202,6 +206,7 @@ int write_thread(void* pipeline_context) { &context->condition_not_full, &context->receiver_done); if (file == NULL) { free(root_directory); + protocol_session_unbind(); return thrd_success; } if (save_to_disk && !file_save_to_disk(root_directory, file, context->config)) { @@ -213,6 +218,7 @@ int write_thread(void* pipeline_context) { cnd_broadcast(&context->condition_not_empty); mtx_unlock(&context->mutex); free(root_directory); + protocol_session_unbind(); return thrd_error; } file_destroy(file); diff --git a/src/shared/multiprocessing.h b/src/shared/multiprocessing.h index 6fd2f31..8a9b79b 100644 --- a/src/shared/multiprocessing.h +++ b/src/shared/multiprocessing.h @@ -29,6 +29,7 @@ typedef struct { unsigned long long progress_bytes; bool sender_done; atomic_bool cancelled; + ProtocolSession allocation_session; } PipelineContextSender; typedef struct PipelineContextReceiver { diff --git a/src/shared/protocol.c b/src/shared/protocol.c index aa79fa4..3b70e50 100644 --- a/src/shared/protocol.c +++ b/src/shared/protocol.c @@ -20,7 +20,8 @@ static __thread int io_read_fd = -1; static __thread int io_write_fd = -1; static __thread SSL* io_ssl; static __thread ProtocolSession* bound_session; -static __thread ProtocolSession legacy_io_session = {.read_fd = -1, .write_fd = -1}; +static __thread ProtocolSession legacy_io_session = { + .read_fd = -1, .write_fd = -1, .max_alloc = DEFAULT_MAX_ALLOC}; static unsigned long long io_bwlimit = 0; static mtx_t bw_mutex; @@ -28,12 +29,30 @@ static once_flag bw_mutex_once = ONCE_FLAG_INIT; static unsigned long long global_bwlimit(void); +static bool protocol_reserve_memory(ProtocolSession* session, size_t charge) { + unsigned long long allocated = atomic_load(&session->total_allocated_bytes); + while (true) { + if (allocated > MAX_CONNECTION_MEMORY || + (unsigned long long)charge > MAX_CONNECTION_MEMORY - allocated) + return false; + if (atomic_compare_exchange_weak(&session->total_allocated_bytes, &allocated, + allocated + (unsigned long long)charge)) + return true; + } +} + +static void protocol_release_memory_for_session(ProtocolSession* session, size_t charge) { + unsigned long long allocated = atomic_load(&session->total_allocated_bytes); + while (true) { + unsigned long long remaining = (unsigned long long)charge >= allocated ? 0 : allocated - charge; + if (atomic_compare_exchange_weak(&session->total_allocated_bytes, &allocated, remaining)) + break; + } +} + void protocol_release_memory(size_t charge) { ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; - if ((unsigned long long)charge >= session->total_allocated_bytes) - session->total_allocated_bytes = 0; - else - session->total_allocated_bytes -= charge; + protocol_release_memory_for_session(session, charge); } void io_set_fds(int read_fd, int write_fd) { bound_session = NULL; @@ -45,8 +64,9 @@ void io_set_fds(int read_fd, int write_fd) { legacy_io_session.read_fd = read_fd; legacy_io_session.write_fd = write_fd; legacy_io_session.ssl = NULL; - legacy_io_session.eight_bit_output = false; - legacy_io_session.total_allocated_bytes = 0; +legacy_io_session.eight_bit_output = false; + atomic_store(&legacy_io_session.total_allocated_bytes, 0); + legacy_io_session.max_alloc = DEFAULT_MAX_ALLOC; protocol_session_set_bwlimit(&legacy_io_session, global_bwlimit()); } @@ -56,9 +76,43 @@ void protocol_session_init(ProtocolSession* session, int read_fd, int write_fd) memset(session, 0, sizeof(*session)); session->read_fd = read_fd; session->write_fd = write_fd; + session->max_alloc = DEFAULT_MAX_ALLOC; + atomic_init(&session->total_allocated_bytes, 0); protocol_session_set_bwlimit(session, global_bwlimit()); } +void protocol_session_set_max_alloc(ProtocolSession* session, unsigned long long max_alloc) { + if (!session) + session = bound_session ? bound_session : &legacy_io_session; + session->max_alloc = max_alloc; +} + +static bool allocation_allowed(const ProtocolSession* session, size_t size) { + return (unsigned long long)size <= session->max_alloc; +} + +static void* protocol_alloc_for_session(const ProtocolSession* session, size_t size) { + if (!allocation_allowed(session, size)) + return NULL; + return malloc(size); +} + +static void* protocol_realloc_for_session(const ProtocolSession* session, void* ptr, size_t size) { + if (!allocation_allowed(session, size)) + return NULL; + return realloc(ptr, size); +} + +void* protocol_alloc(size_t size) { + const ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; + return protocol_alloc_for_session(session, size); +} + +void* protocol_realloc(void* ptr, size_t size) { + const ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; + return protocol_realloc_for_session(session, ptr, size); +} + void protocol_session_bind(ProtocolSession* session) { bound_session = session; log_set_8_bit_output(session && session->eight_bit_output); @@ -169,7 +223,8 @@ static ProtocolSession* legacy_session(int read_fd, int write_fd) { legacy_io_session.write_fd != target_write_fd) { legacy_io_session.read_fd = target_read_fd; legacy_io_session.write_fd = target_write_fd; - legacy_io_session.total_allocated_bytes = 0; + atomic_store(&legacy_io_session.total_allocated_bytes, 0); + legacy_io_session.max_alloc = DEFAULT_MAX_ALLOC; protocol_session_set_bwlimit(&legacy_io_session, global_bwlimit()); } else if (legacy_io_session.bwlimit != global_bwlimit()) { protocol_session_set_bwlimit(&legacy_io_session, global_bwlimit()); @@ -352,13 +407,12 @@ char* protocol_receive_str(ProtocolSession* session) { size_t size; if (!protocol_receive_n_data(session, &size, sizeof(size_t))) return NULL; - if (size > MAX_STRING_SIZE || size > SIZE_MAX - 1 || - size + 1 > MAX_CONNECTION_MEMORY - session->total_allocated_bytes) { + if (size > MAX_STRING_SIZE || size > SIZE_MAX - 1) { log_message(LOG_LEVEL_ERROR, "String size %zu exceeds maximum %llu", size, (unsigned long long)MAX_STRING_SIZE); return NULL; } - char* data = (char*)malloc(size + 1); + char* data = (char*)protocol_alloc_for_session(session, size + 1); if (data == NULL) return NULL; if (!protocol_receive_n_data(session, data, size)) { @@ -371,7 +425,6 @@ char* protocol_receive_str(ProtocolSession* session) { return NULL; } data[size] = '\0'; - session->total_allocated_bytes += size + 1; log_debug_message(LOG_DEBUG_PROTO, "Received String: %s", data); return data; } @@ -401,25 +454,29 @@ Data* protocol_receive_data_limited(ProtocolSession* session, unsigned long long (unsigned long long)MAX_DATA_PAYLOAD_SIZE); return NULL; } + if (size > SIZE_MAX) + return NULL; size_t allocation_size = size == 0 ? 1 : (size_t)size; - if (allocation_size > MAX_CONNECTION_MEMORY - session->total_allocated_bytes) { + if (!protocol_reserve_memory(session, allocation_size)) { log_message(LOG_LEVEL_ERROR, "Per-connection memory limit exceeded (%llu + %llu > %llu)", - (unsigned long long)session->total_allocated_bytes, size, + (unsigned long long)atomic_load(&session->total_allocated_bytes), size, (unsigned long long)MAX_CONNECTION_MEMORY); return NULL; } - void* data = malloc(allocation_size); - if (data == NULL) - return NULL; - if (!protocol_receive_n_data(session, data, (size_t)size)) { - free(data); + void* data = protocol_alloc_for_session(session, allocation_size); + if (data == NULL) { + protocol_release_memory_for_session(session, allocation_size); return NULL; } - session->total_allocated_bytes += allocation_size; - log_debug_message(LOG_DEBUG_PROTO, "Received %lld data", size); + if (!protocol_receive_n_data(session, data, (size_t)size)) { + free(data); + protocol_release_memory_for_session(session, allocation_size); + return NULL; + } +log_debug_message(LOG_DEBUG_PROTO, "Received %lld data", size); Data* result = data_create(data, (size_t)size); if (!result) { - session->total_allocated_bytes -= allocation_size; + protocol_release_memory_for_session(session, allocation_size); return NULL; } result->protocol_charge = allocation_size; diff --git a/src/shared/protocol.h b/src/shared/protocol.h index 8310692..54b3326 100644 --- a/src/shared/protocol.h +++ b/src/shared/protocol.h @@ -4,6 +4,7 @@ #include "data.h" #include #include +#include /* Maximum allowed string size for receive_str (64 KB) */ #define MAX_STRING_SIZE (64 * 1024) @@ -18,6 +19,9 @@ #define MAX_MANIFEST_ENTRIES (1024 * 1024) /* Aggregate bytes retained by one received deletion manifest. */ #define MAX_MANIFEST_BYTES (16ULL * 1024 * 1024) +#define DEFAULT_MAX_ALLOC (1ULL * 1024 * 1024 * 1024) +/* Server policy ceiling for a client-provided allocation limit. */ +#define MAX_SERVER_ALLOC (256ULL * 1024 * 1024) typedef struct ssl_st SSL; @@ -35,8 +39,9 @@ typedef struct ProtocolSession { long long bw_tokens; long long bw_last_refill_sec; long bw_last_refill_nsec; - unsigned long long total_allocated_bytes; + atomic_ullong total_allocated_bytes; bool eight_bit_output; + unsigned long long max_alloc; } ProtocolSession; typedef int Status; @@ -66,6 +71,9 @@ void protocol_session_bind(ProtocolSession* session); void protocol_session_unbind(void); void protocol_session_set_ssl(ProtocolSession* session, SSL* ssl); void protocol_session_set_bwlimit(ProtocolSession* session, unsigned long long bytes_per_sec); +void protocol_session_set_max_alloc(ProtocolSession* session, unsigned long long max_alloc); +void* protocol_alloc(size_t size); +void* protocol_realloc(void* ptr, size_t size); void protocol_session_set_8_bit_output(ProtocolSession* session, bool enabled); void protocol_set_8_bit_output(bool enabled); bool protocol_send_n_data(ProtocolSession* session, const void* data, size_t data_size); diff --git a/tests/test_client_cli.c b/tests/test_client_cli.c index 7c54a33..d8a8519 100644 --- a/tests/test_client_cli.c +++ b/tests/test_client_cli.c @@ -435,6 +435,48 @@ static void test_parse_args_rejects_invalid_modify_window() { } } +static void test_parse_args_max_alloc_sizes() { + const char* values[] = {"1", "4K", "2m", "3G", "1T", "1P", "1E", "512B"}; + const unsigned long long expected[] = {1, + 4ULL * 1024, + 2ULL * 1024 * 1024, + 3ULL * 1024 * 1024 * 1024, + 1ULL * 1024 * 1024 * 1024 * 1024, + 1ULL * 1024 * 1024 * 1024 * 1024 * 1024, + 1ULL * 1024 * 1024 * 1024 * 1024 * 1024 * 1024, + 512}; + for (size_t i = 0; i < sizeof(values) / sizeof(values[0]); i++) { + Config* cfg = config_create(); + char* argv[] = {"fastsync", "--max-alloc", (char*)values[i], "/src", "/dst"}; + int positional_args[2]; + int positional_count = 0; + EXPECT_EQ_INT(parse_args(cfg, 5, argv, positional_args, &positional_count), 0); + EXPECT_TRUE(cfg->max_alloc == expected[i]); + config_delete(cfg); + } + + Config* cfg = config_create(); + char* argv[] = {"fastsync", "--max-alloc=8M", "/src", "/dst"}; + int positional_args[2]; + int positional_count = 0; + EXPECT_EQ_INT(parse_args(cfg, 4, argv, positional_args, &positional_count), 0); + EXPECT_TRUE(cfg->max_alloc == 8ULL * 1024 * 1024); + config_delete(cfg); +} + +static void test_parse_args_rejects_invalid_max_alloc() { + const char* values[] = {"0", "-1", "+1", " 1", "1 ", + "1Z", "1K2", "1 K", "1\tK", "18446744073709551615K"}; + for (size_t i = 0; i < sizeof(values) / sizeof(values[0]); i++) { + Config* cfg = config_create(); + char* argv[] = {"fastsync", "--max-alloc", (char*)values[i], "/src", "/dst"}; + int positional_args[2]; + int positional_count = 0; + EXPECT_EQ_INT(parse_args(cfg, 5, argv, positional_args, &positional_count), -1); + config_delete(cfg); + } +} + static void test_parse_args_skip_compress() { Config* cfg = config_create(); char* argv[] = {"fastsync", "--skip-compress=.ZIP, .GZ", "/src", "/dst"}; @@ -867,7 +909,7 @@ void test_client_cli() { test_parse_args_invalid_server_port(); test_parse_args_invalid_compression_level(); test_parse_args_valid_compression_level(); - test_parse_args_debug_flags(); +test_parse_args_debug_flags(); test_parse_args_debug_help(); test_parse_args_debug_flags_validation(); test_parse_args_modify_window(); @@ -875,6 +917,8 @@ void test_client_cli() { test_parse_args_skip_compress(); test_parse_args_empty_skip_compress(); test_parse_args_compression_threads(); + test_parse_args_max_alloc_sizes(); + test_parse_args_rejects_invalid_max_alloc(); test_parse_args_unknown_option(); test_parse_args_rejects_unimplemented_options(); test_parse_args_quiet(); diff --git a/tests/test_config.c b/tests/test_config.c index be1c724..5792847 100644 --- a/tests/test_config.c +++ b/tests/test_config.c @@ -95,6 +95,7 @@ static void test_pipeline_sender_lifecycle() { EXPECT_EQ_INT(pcs->queue_loader->capacity, 15); EXPECT_FALSE(pcs->scanner_done); EXPECT_FALSE(pcs->loader_done); + EXPECT_EQ_INT((int)pcs->allocation_session.max_alloc, (int)cfg->max_alloc); pipeline_context_sender_destroy(pcs); } @@ -131,7 +132,7 @@ static void test_config_send_receive() { send_cfg->size_only = true; send_cfg->compression_level = 5; send_cfg->chunk_size = 1024; - send_cfg->eight_bit_output = true; +send_cfg->eight_bit_output = true; send_cfg->modify_window = 4; send_cfg->existing = true; send_cfg->ignore_existing = true; @@ -139,6 +140,7 @@ static void test_config_send_receive() { send_cfg->skip_compress_count = 1; send_cfg->skip_compress_suffixes = calloc(1, sizeof(char*)); send_cfg->skip_compress_suffixes[0] = str_dup(".zip"); + send_cfg->max_alloc = MAX_SERVER_ALLOC + 1; /* Use socketpair for bidirectional communication */ int p[2]; @@ -173,7 +175,7 @@ static void test_config_send_receive() { ok = false; if (recv_cfg->chunk_size != 1024) ok = false; - if (!recv_cfg->use_executability) +if (!recv_cfg->use_executability) ok = false; if (!recv_cfg->size_only) ok = false; @@ -183,8 +185,6 @@ static void test_config_send_receive() { ok = false; if (recv_cfg->use_delta) ok = false; - if (!recv_cfg->whole_file) - ok = false; if (recv_cfg->modify_window != 4) ok = false; if (!recv_cfg->existing) @@ -194,6 +194,8 @@ static void test_config_send_receive() { if (!recv_cfg->skip_compress_set || recv_cfg->skip_compress_count != 1 || strcmp(recv_cfg->skip_compress_suffixes[0], ".zip") != 0) ok = false; + if (recv_cfg->max_alloc != MAX_SERVER_ALLOC) + ok = false; } config_delete(recv_cfg); close(p[0]); @@ -223,7 +225,7 @@ static void test_config_send_receive_version_mismatch() { Config* cfg = config_create(); EXPECT_NOT_NULL(cfg); free(cfg->version); - cfg->version = str_dup("2.2.0"); +cfg->version = str_dup("2.3.0"); cfg->send_directory = str_dup("/src"); cfg->receive_root_directory = str_dup("/dst"); @@ -267,6 +269,8 @@ static void test_config_receive_truncated() { /* A valid prefix exercises cleanup after allocated wire strings and a * partially received scalar field. */ EXPECT_TRUE(send_str(p[1], PROTOCOL_VERSION)); + unsigned long long max_alloc = DEFAULT_MAX_ALLOC; + EXPECT_TRUE(send_n_data(p[1], &max_alloc, sizeof(max_alloc))); EXPECT_TRUE(send_str(p[1], "/src")); EXPECT_TRUE(send_str(p[1], "/dst")); EXPECT_TRUE(send_int(p[1], 1)); diff --git a/tests/test_protocol.c b/tests/test_protocol.c index 4c67fee..c6181f5 100644 --- a/tests/test_protocol.c +++ b/tests/test_protocol.c @@ -3,6 +3,60 @@ #include #include #include +#include + +typedef struct { + ProtocolSession* session; + bool allocation_allowed; +} AllocationWorkerArg; + +static int allocation_worker(void* arg) { + AllocationWorkerArg* worker = arg; + protocol_session_bind(worker->session); + void* allocation = protocol_alloc(8); + worker->allocation_allowed = allocation != NULL; + free(allocation); + protocol_session_unbind(); + return thrd_success; +} + +typedef struct { + ProtocolSession* session; + int read_fd; + bool released; +} AccountingWorkerArg; + +typedef struct { + ProtocolSession* session; + atomic_int* ready; + atomic_bool* release; + bool received; +} ConcurrentAccountingWorkerArg; + +static int accounting_worker(void* arg) { + AccountingWorkerArg* worker = arg; + protocol_session_bind(worker->session); + Data* data = protocol_receive_data_limited(worker->session, 8); + if (data) { + data_destroy(data); + worker->released = atomic_load(&worker->session->total_allocated_bytes) == 0; + } + protocol_session_unbind(); + return data ? thrd_success : thrd_error; +} + +static int concurrent_accounting_worker(void* arg) { + ConcurrentAccountingWorkerArg* worker = arg; + protocol_session_bind(worker->session); + Data* data = protocol_receive_data_limited(worker->session, 8); + worker->received = data != NULL; + atomic_fetch_add(worker->ready, 1); + while (!atomic_load(worker->release)) + thrd_yield(); + data_destroy(data); + protocol_session_unbind(); + return thrd_success; +} static void test_send_receive_n_data() { int p[2]; @@ -187,6 +241,177 @@ static void test_receive_str_truncated() { close(p[0]); } +static void test_max_alloc_rejects_single_buffer() { + int p[2]; + EXPECT_EQ_INT(pipe(p), 0); + ProtocolSession session; + protocol_session_init(&session, p[0], p[1]); + protocol_session_set_max_alloc(&session, 4); + protocol_session_bind(&session); + char payload[8] = {0}; + EXPECT_TRUE(write(p[1], &(size_t){sizeof(payload)}, sizeof(size_t)) == sizeof(size_t)); + EXPECT_NULL(protocol_receive_str(&session)); + protocol_session_unbind(); + close(p[0]); + close(p[1]); +} + +static void test_explicit_session_max_alloc_cannot_be_bypassed() { + int p[2]; + EXPECT_EQ_INT(pipe(p), 0); + ProtocolSession explicit_session; + ProtocolSession unrelated_session; + protocol_session_init(&explicit_session, p[0], p[1]); + protocol_session_init(&unrelated_session, p[0], p[1]); + protocol_session_set_max_alloc(&explicit_session, 4); + protocol_session_set_max_alloc(&unrelated_session, 64); + protocol_session_bind(&unrelated_session); + + unsigned long long size = 8; + EXPECT_EQ_INT((int)write(p[1], &size, sizeof(size)), (int)sizeof(size)); + EXPECT_EQ_INT((int)write(p[1], "12345678", 8), 8); + EXPECT_NULL(protocol_receive_data_limited(&explicit_session, 8)); + EXPECT_EQ_INT((int)atomic_load(&explicit_session.total_allocated_bytes), 0); + + protocol_session_unbind(); + close(p[0]); + close(p[1]); +} + +static void test_max_alloc_allows_configured_buffer() { + ProtocolSession session; + protocol_session_init(&session, -1, -1); + protocol_session_set_max_alloc(&session, 4); + protocol_session_bind(&session); + void* allowed = protocol_alloc(4); + const void* rejected = protocol_alloc(5); + EXPECT_NOT_NULL(allowed); + EXPECT_NULL(rejected); + free(allowed); + protocol_session_unbind(); +} + +static void test_max_alloc_is_bound_in_worker_threads() { + enum { WORKER_COUNT = 4 }; + ProtocolSession sessions[WORKER_COUNT]; + AllocationWorkerArg args[WORKER_COUNT] = {0}; + thrd_t threads[WORKER_COUNT]; + for (int i = 0; i < WORKER_COUNT; i++) { + protocol_session_init(&sessions[i], -1, -1); + protocol_session_set_max_alloc(&sessions[i], 4); + args[i].session = &sessions[i]; + EXPECT_EQ_INT(thrd_create(&threads[i], allocation_worker, &args[i]), thrd_success); + } + for (int i = 0; i < WORKER_COUNT; i++) { + int result; + EXPECT_EQ_INT(thrd_join(threads[i], &result), thrd_success); + EXPECT_EQ_INT(result, thrd_success); + EXPECT_FALSE(args[i].allocation_allowed); + } +} + +static void test_protocol_accounting_is_released_in_worker_threads() { + enum { WORKER_COUNT = 4 }; + ProtocolSession sessions[WORKER_COUNT]; + AccountingWorkerArg args[WORKER_COUNT] = {0}; + thrd_t threads[WORKER_COUNT]; + for (int i = 0; i < WORKER_COUNT; i++) { + int p[2]; + EXPECT_EQ_INT(pipe(p), 0); + protocol_session_init(&sessions[i], p[0], p[1]); + protocol_session_set_max_alloc(&sessions[i], 64); + unsigned long long size = 8; + EXPECT_EQ_INT((int)write(p[1], &size, sizeof(size)), (int)sizeof(size)); + EXPECT_EQ_INT((int)write(p[1], "12345678", 8), 8); + close(p[1]); + args[i].session = &sessions[i]; + args[i].read_fd = p[0]; + EXPECT_EQ_INT(thrd_create(&threads[i], accounting_worker, &args[i]), thrd_success); + } + for (int i = 0; i < WORKER_COUNT; i++) { + int result; + EXPECT_EQ_INT(thrd_join(threads[i], &result), thrd_success); + EXPECT_EQ_INT(result, thrd_success); + EXPECT_TRUE(args[i].released); + EXPECT_EQ_INT((int)atomic_load(&sessions[i].total_allocated_bytes), 0); + close(args[i].read_fd); + } +} + +static void test_protocol_accounting_reservation_is_atomic() { + enum { WORKER_COUNT = 8 }; + int p[2]; + EXPECT_EQ_INT(pipe(p), 0); + ProtocolSession session; + protocol_session_init(&session, p[0], p[1]); + protocol_session_set_max_alloc(&session, 64); + const unsigned long long budget_before = MAX_SERVER_ALLOC - 8; + atomic_store(&session.total_allocated_bytes, budget_before); + + for (int i = 0; i < WORKER_COUNT; i++) { + unsigned long long size = 8; + EXPECT_EQ_INT((int)write(p[1], &size, sizeof(size)), (int)sizeof(size)); + EXPECT_EQ_INT((int)write(p[1], "12345678", 8), 8); + } + close(p[1]); + + atomic_int ready; + atomic_bool release; + atomic_init(&ready, 0); + atomic_init(&release, false); + ConcurrentAccountingWorkerArg args[WORKER_COUNT] = {0}; + thrd_t threads[WORKER_COUNT]; + for (int i = 0; i < WORKER_COUNT; i++) { + args[i].session = &session; + args[i].ready = &ready; + args[i].release = &release; + EXPECT_EQ_INT(thrd_create(&threads[i], concurrent_accounting_worker, &args[i]), thrd_success); + } + while (atomic_load(&ready) != WORKER_COUNT) + thrd_yield(); + bool budget_ok = atomic_load(&session.total_allocated_bytes) == budget_before + 8; + atomic_store(&release, true); + int received = 0; + for (int i = 0; i < WORKER_COUNT; i++) { + int result; + EXPECT_EQ_INT(thrd_join(threads[i], &result), thrd_success); + EXPECT_EQ_INT(result, thrd_success); + received += args[i].received ? 1 : 0; + } + EXPECT_EQ_INT(received, 1); + EXPECT_TRUE(budget_ok); + EXPECT_EQ_INT((int)atomic_load(&session.total_allocated_bytes), (int)budget_before); + close(p[0]); +} + +static void test_protocol_string_accounting_is_transient() { + int p[2]; + EXPECT_EQ_INT(pipe(p), 0); + ProtocolSession session; + protocol_session_init(&session, p[0], p[1]); + protocol_session_set_max_alloc(&session, 64); + EXPECT_TRUE(protocol_send_str(&session, "temporary")); + char* received = protocol_receive_str(&session); + EXPECT_NOT_NULL(received); + EXPECT_EQ_STR(received, "temporary"); + EXPECT_EQ_INT((int)atomic_load(&session.total_allocated_bytes), 0); + free(received); + close(p[0]); + close(p[1]); +} + +static void test_protocol_accounting_release_does_not_underflow() { + ProtocolSession session; + protocol_session_init(&session, -1, -1); + atomic_store(&session.total_allocated_bytes, 4); + protocol_session_bind(&session); + protocol_release_memory(8); + EXPECT_EQ_INT((int)atomic_load(&session.total_allocated_bytes), 0); + protocol_release_memory(1); + EXPECT_EQ_INT((int)atomic_load(&session.total_allocated_bytes), 0); + protocol_session_unbind(); +} + void test_protocol() { test_send_receive_n_data(); test_send_receive_n_data_zero(); @@ -198,4 +423,12 @@ void test_protocol() { test_send_receive_status(); test_receive_n_data_truncated(); test_receive_str_truncated(); + test_max_alloc_rejects_single_buffer(); + test_explicit_session_max_alloc_cannot_be_bypassed(); + test_max_alloc_allows_configured_buffer(); + test_max_alloc_is_bound_in_worker_threads(); + test_protocol_accounting_is_released_in_worker_threads(); + test_protocol_accounting_reservation_is_atomic(); + test_protocol_string_accounting_is_transient(); + test_protocol_accounting_release_does_not_underflow(); } diff --git a/tests/test_robustness.c b/tests/test_robustness.c index a03f163..1e1cd09 100644 --- a/tests/test_robustness.c +++ b/tests/test_robustness.c @@ -109,6 +109,20 @@ static void test_delta_deserialize_garbage() { data_destroy(d); } +static void test_delta_deserialize_respects_max_alloc() { + unsigned char serialized[sizeof(uint64_t) + sizeof(uint32_t)] = {0}; + Data data = {.data = serialized, .size = sizeof(serialized)}; + ProtocolSession session; + protocol_session_init(&session, -1, -1); + protocol_session_set_max_alloc(&session, sizeof(Delta) - 1); + protocol_session_bind(&session); + + const Delta* result = delta_deserialize(&data); + EXPECT_NULL(result); + + protocol_session_unbind(); +} + static void test_delta_signature_deserialize_truncated() { char old_data[4096]; for (int i = 0; i < 4096; i++) @@ -206,6 +220,7 @@ void test_robustness() { test_delta_deserialize_truncated(); test_delta_deserialize_empty(); test_delta_deserialize_garbage(); + test_delta_deserialize_respects_max_alloc(); test_delta_deserialize_truncated_instructions(); test_delta_signature_deserialize_truncated(); test_delta_apply_null(); diff --git a/tests/test_scanner.c b/tests/test_scanner.c index 71d8a20..37a0565 100644 --- a/tests/test_scanner.c +++ b/tests/test_scanner.c @@ -396,7 +396,7 @@ static void test_parallel_scanner_root_chunks_without_workers() { ScannerOptions options = {false, 1, NULL, 0, NULL, 0, 0, 0, 0, 0, false, false, false, false, false}; - ParallelScanner* scanner = parallel_scanner_create_with_options(dir, &options); + ParallelScanner* scanner = parallel_scanner_create_with_options(dir, &options, NULL); EXPECT_NOT_NULL(scanner); int total_files = 0;