diff --git a/src/shared/protocol.c b/src/shared/protocol.c index b246b32..226d098 100644 --- a/src/shared/protocol.c +++ b/src/shared/protocol.c @@ -40,8 +40,7 @@ static bool protocol_reserve_memory(ProtocolSession* session, size_t charge) { } } -void protocol_release_memory(size_t charge) { - ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; +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; @@ -49,6 +48,11 @@ void protocol_release_memory(size_t charge) { break; } } + +void protocol_release_memory(size_t charge) { + ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; + protocol_release_memory_for_session(session, charge); +} void io_set_fds(int read_fd, int write_fd) { bound_session = NULL; io_read_fd = read_fd; @@ -81,23 +85,32 @@ void protocol_session_set_max_alloc(ProtocolSession* session, unsigned long long session->max_alloc = max_alloc; } -static bool allocation_allowed(size_t size) { - const ProtocolSession* session = bound_session ? bound_session : &legacy_io_session; +static bool allocation_allowed(const ProtocolSession* session, size_t size) { return (unsigned long long)size <= session->max_alloc; } -void* protocol_alloc(size_t size) { - if (!allocation_allowed(size)) +static void* protocol_alloc_for_session(const ProtocolSession* session, size_t size) { + if (!allocation_allowed(session, size)) return NULL; return malloc(size); } -void* protocol_realloc(void* ptr, size_t size) { - if (!allocation_allowed(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; } @@ -383,7 +396,7 @@ char* protocol_receive_str(ProtocolSession* session) { (unsigned long long)MAX_STRING_SIZE); return NULL; } - char* data = (char*)protocol_alloc(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)) { @@ -434,20 +447,20 @@ Data* protocol_receive_data_limited(ProtocolSession* session, unsigned long long (unsigned long long)MAX_CONNECTION_MEMORY); return NULL; } - void* data = protocol_alloc(allocation_size); + void* data = protocol_alloc_for_session(session, allocation_size); if (data == NULL) { - protocol_release_memory(allocation_size); + protocol_release_memory_for_session(session, allocation_size); return NULL; } if (!protocol_receive_n_data(session, data, (size_t)size)) { free(data); - protocol_release_memory(allocation_size); + protocol_release_memory_for_session(session, allocation_size); return NULL; } log_message(LOG_LEVEL_DEBUG, "Received %lld data", size); Data* result = data_create(data, (size_t)size); if (!result) { - protocol_release_memory(allocation_size); + protocol_release_memory_for_session(session, allocation_size); return NULL; } result->protocol_charge = allocation_size; diff --git a/tests/test_protocol.c b/tests/test_protocol.c index 476dfe5..c6181f5 100644 --- a/tests/test_protocol.c +++ b/tests/test_protocol.c @@ -256,6 +256,28 @@ static void test_max_alloc_rejects_single_buffer() { 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); @@ -402,6 +424,7 @@ void test_protocol() { 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();