diff --git a/src/client/client_send.c b/src/client/client_send.c index 41c54ed..bb293c1 100644 --- a/src/client/client_send.c +++ b/src/client/client_send.c @@ -458,8 +458,8 @@ 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) { 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/server/server.c b/src/server/server.c index d7221ef..a81f7f2 100644 --- a/src/server/server.c +++ b/src/server/server.c @@ -295,7 +295,8 @@ void handler(int file_descriptor) { return; } protocol_session_set_max_alloc(&context->session, config->max_alloc); - context->session.total_allocated_bytes = session.total_allocated_bytes; + 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/multiprocessing.c b/src/shared/multiprocessing.c index f58e946..aa8a1f9 100644 --- a/src/shared/multiprocessing.c +++ b/src/shared/multiprocessing.c @@ -181,6 +181,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); @@ -192,6 +193,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; } @@ -201,6 +203,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)) { @@ -212,6 +215,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/protocol.c b/src/shared/protocol.c index bb499fd..8624893 100644 --- a/src/shared/protocol.c +++ b/src/shared/protocol.c @@ -30,10 +30,12 @@ static unsigned long long global_bwlimit(void); 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; + 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 io_set_fds(int read_fd, int write_fd) { bound_session = NULL; @@ -45,7 +47,7 @@ void io_set_fds(int read_fd, int write_fd) { legacy_io_session.read_fd = read_fd; legacy_io_session.write_fd = write_fd; legacy_io_session.ssl = NULL; - legacy_io_session.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()); } @@ -57,6 +59,7 @@ void protocol_session_init(ProtocolSession* session, int read_fd, int write_fd) 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()); } @@ -179,7 +182,7 @@ 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()) { @@ -364,7 +367,7 @@ char* protocol_receive_str(ProtocolSession* session) { 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) { + size + 1 > MAX_CONNECTION_MEMORY - atomic_load(&session->total_allocated_bytes)) { log_message(LOG_LEVEL_ERROR, "String size %zu exceeds maximum %llu", size, (unsigned long long)MAX_STRING_SIZE); return NULL; @@ -382,7 +385,7 @@ char* protocol_receive_str(ProtocolSession* session) { return NULL; } data[size] = '\0'; - session->total_allocated_bytes += size + 1; + atomic_fetch_add(&session->total_allocated_bytes, size + 1); log_message(LOG_LEVEL_DEBUG, "Received String: %s", data); return data; } @@ -415,9 +418,9 @@ Data* protocol_receive_data_limited(ProtocolSession* session, unsigned long long 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 (allocation_size > MAX_CONNECTION_MEMORY - atomic_load(&session->total_allocated_bytes)) { 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; } @@ -428,11 +431,11 @@ Data* protocol_receive_data_limited(ProtocolSession* session, unsigned long long free(data); return NULL; } - session->total_allocated_bytes += allocation_size; + atomic_fetch_add(&session->total_allocated_bytes, allocation_size); log_message(LOG_LEVEL_DEBUG, "Received %lld data", size); Data* result = data_create(data, (size_t)size); if (!result) { - session->total_allocated_bytes -= allocation_size; + protocol_release_memory(allocation_size); return NULL; } result->protocol_charge = allocation_size; diff --git a/src/shared/protocol.h b/src/shared/protocol.h index 9022e22..7d67ac8 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) @@ -38,7 +39,7 @@ 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; unsigned long long max_alloc; } ProtocolSession; diff --git a/tests/test_protocol.c b/tests/test_protocol.c index fb325fa..53063d9 100644 --- a/tests/test_protocol.c +++ b/tests/test_protocol.c @@ -3,6 +3,40 @@ #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; + +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 void test_send_receive_n_data() { int p[2]; @@ -215,6 +249,53 @@ static void test_max_alloc_allows_configured_buffer() { 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); + } +} + void test_protocol() { test_send_receive_n_data(); test_send_receive_n_data_zero(); @@ -228,4 +309,6 @@ void test_protocol() { test_receive_str_truncated(); test_max_alloc_rejects_single_buffer(); test_max_alloc_allows_configured_buffer(); + test_max_alloc_is_bound_in_worker_threads(); + test_protocol_accounting_is_released_in_worker_threads(); } 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;