diff --git a/src/shared/protocol.c b/src/shared/protocol.c index 8624893..b246b32 100644 --- a/src/shared/protocol.c +++ b/src/shared/protocol.c @@ -28,6 +28,18 @@ 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; + } +} + void protocol_release_memory(size_t charge) { ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; unsigned long long allocated = atomic_load(&session->total_allocated_bytes); @@ -366,8 +378,7 @@ 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 - atomic_load(&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; @@ -385,7 +396,6 @@ char* protocol_receive_str(ProtocolSession* session) { return NULL; } data[size] = '\0'; - atomic_fetch_add(&session->total_allocated_bytes, size + 1); log_message(LOG_LEVEL_DEBUG, "Received String: %s", data); return data; } @@ -418,20 +428,22 @@ 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 - atomic_load(&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)atomic_load(&session->total_allocated_bytes), size, (unsigned long long)MAX_CONNECTION_MEMORY); return NULL; } void* data = protocol_alloc(allocation_size); - if (data == NULL) + if (data == NULL) { + protocol_release_memory(allocation_size); return NULL; + } if (!protocol_receive_n_data(session, data, (size_t)size)) { free(data); + protocol_release_memory(allocation_size); return NULL; } - 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) { diff --git a/tests/test_protocol.c b/tests/test_protocol.c index 53063d9..476dfe5 100644 --- a/tests/test_protocol.c +++ b/tests/test_protocol.c @@ -26,6 +26,13 @@ typedef struct { 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); @@ -38,6 +45,19 @@ static int accounting_worker(void* arg) { 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]; EXPECT_EQ_INT(pipe(p), 0); @@ -296,6 +316,80 @@ static void test_protocol_accounting_is_released_in_worker_threads() { } } +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(); @@ -311,4 +405,7 @@ void test_protocol() { 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(); }