diff --git a/include/memory.h b/include/memory.h new file mode 100644 index 0000000..d6cd54d --- /dev/null +++ b/include/memory.h @@ -0,0 +1,11 @@ +#ifndef RPC_MEMORY_H +#define RPC_MEMORY_H + +#include + +void *rpc_mem_alloc(size_t size); +void *rpc_mem_calloc(size_t count, size_t size); +void *rpc_mem_realloc(void *ptr, size_t size); +void rpc_mem_free(void *ptr); + +#endif diff --git a/include/routes.h b/include/routes.h index 8c003c7..00d3d70 100644 --- a/include/routes.h +++ b/include/routes.h @@ -15,12 +15,16 @@ typedef struct rpc_route { uint64_t proc_id; rpc_handler_fn handler; void *user_data; + rpc_route_finalizer_fn finalizer; int is_async; struct rpc_route *next; } rpc_route; typedef struct rpc_retired_route { rpc_route *route; + uint64_t finalize_proc_id; + int finalize_all; + int finalize_one; struct rpc_retired_route *next; } rpc_retired_route; @@ -38,7 +42,8 @@ typedef struct rpc_routes { int rpc_routes_init(rpc_routes *routes); void rpc_routes_destroy(rpc_routes *routes); int rpc_routes_add(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handler, void *user_data); -int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handler, void *user_data, int is_async); +int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handler, void *user_data, + rpc_route_finalizer_fn finalizer, int is_async); int rpc_routes_remove(rpc_routes *routes, uint64_t proc_id); int rpc_routes_lookup(rpc_routes *routes, uint64_t proc_id, rpc_route *out); diff --git a/include/rpc/protocol.h b/include/rpc/protocol.h index 438af56..f384d01 100644 --- a/include/rpc/protocol.h +++ b/include/rpc/protocol.h @@ -53,6 +53,13 @@ typedef struct rpc_value { } as; } rpc_value; +typedef struct rpc_allocator { + void *ctx; + void *(*alloc)(void *ctx, size_t size); + void *(*realloc)(void *ctx, void *ptr, size_t size); + void (*free)(void *ctx, void *ptr); +} rpc_allocator; + typedef struct rpc_header { rpc_op op; uint8_t flags; @@ -67,6 +74,10 @@ typedef struct rpc_writer { size_t cap; } rpc_writer; +/* Process-wide allocator hook. Set before creating/freeing RPC objects or payload buffers; do not swap it live. */ +int rpc_set_allocator(const rpc_allocator *allocator); +void rpc_get_allocator(rpc_allocator *out_allocator); + int rpc_header_encode(const rpc_header *header, uint8_t out[RPC_HEADER_SIZE]); int rpc_header_decode(const uint8_t in[RPC_HEADER_SIZE], rpc_header *out); diff --git a/include/rpc/server.h b/include/rpc/server.h index 8ecb67c..ded154b 100644 --- a/include/rpc/server.h +++ b/include/rpc/server.h @@ -11,6 +11,7 @@ typedef struct rpc_ctx rpc_ctx; typedef struct rpc_server rpc_server; typedef int (*rpc_handler_fn)(rpc_ctx *ctx, const rpc_value *args, size_t argc, rpc_writer *out, void *user_data); +typedef void (*rpc_route_finalizer_fn)(void *user_data); uint64_t rpc_ctx_call_id(const rpc_ctx *ctx); uint64_t rpc_ctx_proc_id(const rpc_ctx *ctx); @@ -28,6 +29,10 @@ void rpc_server_destroy(rpc_server *server); int rpc_server_add_route_name(rpc_server *server, const char *proc_name, rpc_handler_fn handler, void *user_data); int rpc_server_add_async_route_name(rpc_server *server, const char *proc_name, rpc_handler_fn handler, void *user_data); +int rpc_server_add_route_name_ex(rpc_server *server, const char *proc_name, rpc_handler_fn handler, void *user_data, + rpc_route_finalizer_fn finalizer); +int rpc_server_add_async_route_name_ex(rpc_server *server, const char *proc_name, rpc_handler_fn handler, + void *user_data, rpc_route_finalizer_fn finalizer); int rpc_server_remove_route_name(rpc_server *server, const char *proc_name); #ifdef __cplusplus diff --git a/meson.build b/meson.build index 02f6ab9..7128512 100644 --- a/meson.build +++ b/meson.build @@ -14,6 +14,7 @@ clang_format = find_program('clang-format', required: false) core_sources = files( 'src/backend/kqueue.c', 'src/client.c', + 'src/memory.c', 'src/payload.c', 'src/protocol.c', 'src/routes.c', diff --git a/src/backend/kqueue.c b/src/backend/kqueue.c index 2c804fc..29ea7e4 100644 --- a/src/backend/kqueue.c +++ b/src/backend/kqueue.c @@ -1,7 +1,8 @@ #include "backend.h" +#include "memory.h" + #include -#include #include #include #include @@ -38,11 +39,11 @@ static int set_interest(rpc_backend *backend, int fd, uint32_t events, uintptr_t int rpc_backend_kqueue_create(rpc_backend **out) { if (!out) { return -1; } - rpc_backend *backend = calloc(1, sizeof(*backend)); + rpc_backend *backend = rpc_mem_calloc(1, sizeof(*backend)); if (!backend) { return -1; } backend->kq = kqueue(); if (backend->kq < 0) { - free(backend); + rpc_mem_free(backend); return -1; } @@ -50,7 +51,7 @@ int rpc_backend_kqueue_create(rpc_backend **out) { EV_SET(&wake, 1, EVFILT_USER, EV_ADD | EV_CLEAR, 0, 0, NULL); if (kevent(backend->kq, &wake, 1, NULL, 0, NULL) != 0) { close(backend->kq); - free(backend); + rpc_mem_free(backend); return -1; } @@ -61,7 +62,7 @@ int rpc_backend_kqueue_create(rpc_backend **out) { void rpc_backend_destroy(rpc_backend *backend) { if (backend) { if (backend->kq >= 0) { close(backend->kq); } - free(backend); + rpc_mem_free(backend); } } diff --git a/src/client.c b/src/client.c index e140b39..9058b76 100644 --- a/src/client.c +++ b/src/client.c @@ -1,6 +1,7 @@ #include "rpc/client.h" #include "rpc/trace.h" +#include "memory.h" #include "proc.h" #include @@ -80,7 +81,7 @@ static int client_read_reserve(rpc_client *client, size_t need) { while (next_cap < need) { next_cap *= 2u; } - uint8_t *next = realloc(client->read_buf, next_cap); + uint8_t *next = rpc_mem_realloc(client->read_buf, next_cap); if (!next) { return -1; } client->read_buf = next; client->read_cap = next_cap; @@ -160,7 +161,7 @@ static int recv_packet(rpc_client *client, rpc_header *header, uint8_t **body) { return -1; } - *body = malloc(header->size ? header->size : 1); + *body = rpc_mem_alloc(header->size ? header->size : 1); if (!*body) { set_error(client, "out of memory"); rpc_trace_end(RPC_TRACE_CLIENT_RECV, trace); @@ -181,7 +182,7 @@ static int client_send_call_id(rpc_client *client, uint64_t proc_id, const rpc_w int rpc_client_connect(rpc_client **out_client, const char *host, const char *port) { if (!out_client || !port) { return -1; } - rpc_client *client = calloc(1, sizeof(*client)); + rpc_client *client = rpc_mem_calloc(1, sizeof(*client)); if (!client) { return -1; } client->fd = -1; client->next_call_id = 1; @@ -193,7 +194,7 @@ int rpc_client_connect(rpc_client **out_client, const char *host, const char *po struct addrinfo *res = NULL; if (getaddrinfo(host, port, &hints, &res) != 0) { - free(client); + rpc_mem_free(client); return -1; } @@ -213,7 +214,7 @@ int rpc_client_connect(rpc_client **out_client, const char *host, const char *po freeaddrinfo(res); if (client->fd < 0) { - free(client); + rpc_mem_free(client); return -1; } @@ -227,8 +228,8 @@ void rpc_client_close(rpc_client *client) { (void)send_packet(client, RPC_OP_DISCONNECT, 0, client->next_call_id++, NULL); close(client->fd); } - free(client->read_buf); - free(client); + rpc_mem_free(client->read_buf); + rpc_mem_free(client); } int rpc_client_ping(rpc_client *client) { @@ -239,7 +240,7 @@ int rpc_client_ping(rpc_client *client) { rpc_header header; uint8_t *body = NULL; if (recv_packet(client, &header, &body) != 0) { return -1; } - free(body); + rpc_mem_free(body); if (header.op != RPC_OP_RESPONSE || header.call_id != call_id || header.size != 0) { set_error(client, "unexpected ping response"); @@ -333,7 +334,7 @@ int rpc_client_recv_response(rpc_client *client, uint64_t *out_call_id, rpc_valu set_error(client, "unexpected response op"); } - free(body); + rpc_mem_free(body); return rc; } diff --git a/src/memory.c b/src/memory.c new file mode 100644 index 0000000..eaad0cc --- /dev/null +++ b/src/memory.c @@ -0,0 +1,65 @@ +#include "memory.h" +#include "rpc/protocol.h" + +#include +#include +#include + +static rpc_allocator g_allocator; + +static int allocator_valid(const rpc_allocator *allocator) { + if (!allocator) { return 0; } + return allocator->alloc && allocator->realloc && allocator->free; +} + +int rpc_set_allocator(const rpc_allocator *allocator) { + if (allocator) { + if (!allocator_valid(allocator)) { + return -1; + } + g_allocator = *allocator; + } else { + memset(&g_allocator, 0, sizeof(g_allocator)); + } + return 0; +} + +void rpc_get_allocator(rpc_allocator *out_allocator) { + if (!out_allocator) { return; } + *out_allocator = g_allocator; +} + +void *rpc_mem_alloc(size_t size) { + if (size == 0) { size = 1; } + rpc_allocator allocator; + rpc_get_allocator(&allocator); + if (allocator.alloc) { return allocator.alloc(allocator.ctx, size); } + return malloc(size); +} + +void *rpc_mem_calloc(size_t count, size_t size) { + if (count != 0 && size > SIZE_MAX / count) { return NULL; } + size_t bytes = count * size; + void *ptr = rpc_mem_alloc(bytes); + if (ptr) { memset(ptr, 0, bytes); } + return ptr; +} + +void *rpc_mem_realloc(void *ptr, size_t size) { + if (size == 0) { size = 1; } + rpc_allocator allocator; + rpc_get_allocator(&allocator); + if (allocator.realloc) { return allocator.realloc(allocator.ctx, ptr, size); } + return realloc(ptr, size); +} + +void rpc_mem_free(void *ptr) { + if (!ptr) { return; } + rpc_allocator allocator; + rpc_get_allocator(&allocator); + if (allocator.free) { + allocator.free(allocator.ctx, ptr); + } else { + free(ptr); + } +} diff --git a/src/payload.c b/src/payload.c index bece0cc..584eef9 100644 --- a/src/payload.c +++ b/src/payload.c @@ -1,8 +1,9 @@ #include "rpc/protocol.h" #include "rpc/trace.h" +#include "memory.h" + #include -#include #include static void put_u32(uint8_t *out, uint32_t value) { @@ -43,7 +44,7 @@ static int writer_reserve(rpc_writer *writer, size_t extra) { } cap *= 2u; } - uint8_t *next = realloc(writer->data, cap); + uint8_t *next = rpc_mem_realloc(writer->data, cap); if (!next) { return -1; } writer->data = next; writer->cap = cap; @@ -67,7 +68,7 @@ void rpc_writer_reset(rpc_writer *writer) { void rpc_writer_free(rpc_writer *writer) { if (writer) { - free(writer->data); + rpc_mem_free(writer->data); memset(writer, 0, sizeof(*writer)); } } @@ -144,9 +145,9 @@ int rpc_payload_decode(const uint8_t *data, size_t len, rpc_value **out_values, while (off < len) { if (count == cap) { size_t next_cap = cap ? cap * 2u : 4u; - rpc_value *next = realloc(values, next_cap * sizeof(*values)); + rpc_value *next = rpc_mem_realloc(values, next_cap * sizeof(*values)); if (!next) { - free(values); + rpc_mem_free(values); rpc_trace_end(RPC_TRACE_PAYLOAD_DECODE, trace); return -1; } @@ -211,7 +212,7 @@ int rpc_payload_decode(const uint8_t *data, size_t len, rpc_value **out_values, continue; malformed: - free(values); + rpc_mem_free(values); rpc_trace_end(RPC_TRACE_PAYLOAD_DECODE, trace); return -1; } @@ -223,5 +224,5 @@ int rpc_payload_decode(const uint8_t *data, size_t len, rpc_value **out_values, } void rpc_values_free(rpc_value *values) { - free(values); + rpc_mem_free(values); } diff --git a/src/routes.c b/src/routes.c index 6a17f90..170808f 100644 --- a/src/routes.c +++ b/src/routes.c @@ -1,6 +1,8 @@ #include "routes.h" #include "rpc/trace.h" +#include "memory.h" + #include #include #include @@ -15,10 +17,15 @@ static rpc_route_slot *route_page_alloc(const rpc_routes *routes) { return route_mmap(routes->page_bytes); } -static void route_chain_free(rpc_route *route) { +static void route_finalize(const rpc_route *route) { + if (route && route->finalizer) { route->finalizer(route->user_data); } +} + +static void route_chain_free(rpc_route *route, int finalize_all, int finalize_one, uint64_t finalize_proc_id) { while (route) { rpc_route *next = route->next; - free(route); + if (finalize_all || (finalize_one && route->proc_id == finalize_proc_id)) { route_finalize(route); } + rpc_mem_free(route); route = next; } } @@ -40,19 +47,25 @@ static void free_retired(rpc_routes *routes) { while (node) { rpc_retired_route *next = node->next; - route_chain_free(node->route); - free(node); + route_chain_free(node->route, node->finalize_all, node->finalize_one, node->finalize_proc_id); + rpc_mem_free(node); node = next; } } -static int retire_route(rpc_routes *routes, rpc_route *route) { +static int retire_route(rpc_routes *routes, rpc_route *route, int finalize_one, uint64_t finalize_proc_id) { if (!route) return 0; + if (atomic_load_explicit(&routes->active_readers, memory_order_acquire) == 0) { + route_chain_free(route, 0, finalize_one, finalize_proc_id); + return 0; + } - rpc_retired_route *node = calloc(1, sizeof(*node)); + rpc_retired_route *node = rpc_mem_calloc(1, sizeof(*node)); if (!node) return -1; node->route = route; + node->finalize_proc_id = finalize_proc_id; + node->finalize_one = finalize_one; node->next = routes->retired; routes->retired = node; free_retired(routes); @@ -93,7 +106,7 @@ void rpc_routes_destroy(rpc_routes *routes) { for (size_t i = 0; i < RPC_ROUTE_PAGE_SIZE; ++i) { rpc_route *route = atomic_load_explicit(&page[i], memory_order_relaxed); - route_chain_free(route); + route_chain_free(route, 1, 0, 0); } munmap(page, routes->page_bytes); } @@ -104,8 +117,8 @@ retired: while (node) { rpc_retired_route *next = node->next; - route_chain_free(node->route); - free(node); + route_chain_free(node->route, node->finalize_all, node->finalize_one, node->finalize_proc_id); + rpc_mem_free(node); node = next; } @@ -115,16 +128,21 @@ retired: } int rpc_routes_add(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handler, void *user_data) { - return rpc_routes_add_ex(routes, proc_id, handler, user_data, 0); + return rpc_routes_add_ex(routes, proc_id, handler, user_data, NULL, 0); } -int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handler, void *user_data, int is_async) { +int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handler, void *user_data, + rpc_route_finalizer_fn finalizer, int is_async) { if (!routes || !handler) return -1; - rpc_route *route = calloc(1, sizeof(*route)); + rpc_route *route = rpc_mem_calloc(1, sizeof(*route)); if (!route) return -1; - *route = (rpc_route){.proc_id = proc_id, .handler = handler, .user_data = user_data, .is_async = is_async ? 1 : 0}; + *route = (rpc_route){.proc_id = proc_id, + .handler = handler, + .user_data = user_data, + .finalizer = finalizer, + .is_async = is_async ? 1 : 0}; pthread_mutex_lock(&routes->mutate_lock); uint32_t index = route_index(proc_id); @@ -135,7 +153,7 @@ int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handl page = route_page_alloc(routes); if (!page) { pthread_mutex_unlock(&routes->mutate_lock); - free(route); + rpc_mem_free(route); return -1; } atomic_store_explicit(&routes->pages[page_idx], page, memory_order_release); @@ -143,11 +161,15 @@ int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handl rpc_route *old = atomic_load_explicit(&page[slot_idx], memory_order_acquire); rpc_route *copy_head = route; rpc_route **copy_tail = &route->next; + int replaced = 0; for (rpc_route *it = old; it; it = it->next) { - if (it->proc_id == proc_id) continue; - rpc_route *copy = calloc(1, sizeof(*copy)); + if (it->proc_id == proc_id) { + replaced = 1; + continue; + } + rpc_route *copy = rpc_mem_calloc(1, sizeof(*copy)); if (!copy) { - route_chain_free(copy_head); + route_chain_free(copy_head, 0, 0, 0); pthread_mutex_unlock(&routes->mutate_lock); return -1; } @@ -158,7 +180,7 @@ int rpc_routes_add_ex(rpc_routes *routes, uint64_t proc_id, rpc_handler_fn handl } atomic_store_explicit(&page[slot_idx], copy_head, memory_order_release); - int rc = retire_route(routes, old); + int rc = retire_route(routes, old, replaced, proc_id); pthread_mutex_unlock(&routes->mutate_lock); return rc; @@ -181,9 +203,9 @@ int rpc_routes_remove(rpc_routes *routes, uint64_t proc_id) { removed = 1; continue; } - rpc_route *copy = calloc(1, sizeof(*copy)); + rpc_route *copy = rpc_mem_calloc(1, sizeof(*copy)); if (!copy) { - route_chain_free(copy_head); + route_chain_free(copy_head, 0, 0, 0); pthread_mutex_unlock(&routes->mutate_lock); return -1; } @@ -195,10 +217,10 @@ int rpc_routes_remove(rpc_routes *routes, uint64_t proc_id) { if (page && removed) { atomic_store_explicit(&page[slot_idx], copy_head, memory_order_release); } else { - route_chain_free(copy_head); + route_chain_free(copy_head, 0, 0, 0); } - int rc = removed ? retire_route(routes, old) : -1; + int rc = removed ? retire_route(routes, old, 1, proc_id) : -1; pthread_mutex_unlock(&routes->mutate_lock); return rc; diff --git a/src/scheduler.c b/src/scheduler.c index ec19f03..d996e37 100644 --- a/src/scheduler.c +++ b/src/scheduler.c @@ -1,11 +1,14 @@ #include "scheduler.h" #include "arena.h" +#include "memory.h" #include "rpc/trace.h" #define MCO_USE_VMEM_ALLOCATOR #define MCO_ZERO_MEMORY #define MCO_DEFAULT_STACK_SIZE (1024 * 1024) +#define MCO_ALLOC(size) rpc_mem_calloc(1, size) +#define MCO_DEALLOC(ptr, size) rpc_mem_free(ptr) #define MINICORO_IMPL #include "minicoro.h" @@ -44,7 +47,7 @@ static void call_free(rpc_scheduler *scheduler, rpc_call *call) { if (call->co) { (void)mco_destroy(call->co); } rpc_values_free(call->args); rpc_writer_free(&call->response); - free(call->payload); + rpc_mem_free(call->payload); rpc_fixed_arena_free(&scheduler->call_arena, call); } @@ -97,10 +100,10 @@ static int call_resume(rpc_call *call) { int rpc_scheduler_init(rpc_scheduler **out) { if (!out) { return -1; } - rpc_scheduler *scheduler = calloc(1, sizeof(*scheduler)); + rpc_scheduler *scheduler = rpc_mem_calloc(1, sizeof(*scheduler)); if (!scheduler) { return -1; } if (rpc_fixed_arena_init(&scheduler->call_arena, sizeof(rpc_call), RPC_CALL_ARENA_CAPACITY) != 0) { - free(scheduler); + rpc_mem_free(scheduler); return -1; } *out = scheduler; @@ -114,7 +117,7 @@ void rpc_scheduler_destroy(rpc_scheduler *scheduler) { call_free(scheduler, call); } rpc_fixed_arena_destroy(&scheduler->call_arena); - free(scheduler); + rpc_mem_free(scheduler); } int rpc_scheduler_submit(rpc_scheduler *scheduler, uint64_t call_id, uint64_t proc_id, rpc_handler_fn handler, @@ -141,7 +144,7 @@ int rpc_scheduler_submit(rpc_scheduler *scheduler, uint64_t call_id, uint64_t pr rpc_writer_init(&call->response); if (payload_len > 0) { - call->payload = malloc(payload_len); + call->payload = rpc_mem_alloc(payload_len); if (!call->payload) { call_free(scheduler, call); rpc_trace_end(RPC_TRACE_SCHED_SUBMIT, trace_submit); diff --git a/src/server.c b/src/server.c index a836845..15d01c9 100644 --- a/src/server.c +++ b/src/server.c @@ -3,6 +3,7 @@ #include "arena.h" #include "backend.h" +#include "memory.h" #include "proc.h" #include "routes.h" #include "scheduler.h" @@ -133,7 +134,7 @@ static int append_bytes(uint8_t **buf, size_t *len, size_t *cap, const void *dat while (next_cap < need) { next_cap *= 2u; } - uint8_t *next = realloc(*buf, next_cap); + uint8_t *next = rpc_mem_realloc(*buf, next_cap); if (!next) { return -1; } *buf = next; *cap = next_cap; @@ -149,7 +150,7 @@ static int reserve_bytes(uint8_t **buf, size_t *cap, size_t need) { while (next_cap < need) { next_cap *= 2u; } - uint8_t *next = realloc(*buf, next_cap); + uint8_t *next = rpc_mem_realloc(*buf, next_cap); if (!next) { return -1; } *buf = next; *cap = next_cap; @@ -191,8 +192,8 @@ static void connection_destroy(rpc_connection *conn) { rpc_worker *worker = conn->worker; (void)rpc_backend_remove(worker->backend, conn->fd); close(conn->fd); - free(conn->read_buf); - free(conn->write_buf); + rpc_mem_free(conn->read_buf); + rpc_mem_free(conn->write_buf); rpc_connection **link = &worker->connections; while (*link && *link != conn) { @@ -429,7 +430,7 @@ static int worker_add_connection(rpc_worker *worker, int fd) { } static void worker_enqueue_connection(rpc_worker *worker, int fd) { - rpc_pending_fd *pending = malloc(sizeof(*pending)); + rpc_pending_fd *pending = rpc_mem_alloc(sizeof(*pending)); if (!pending) { close(fd); return; @@ -459,7 +460,7 @@ static void worker_drain_pending(rpc_worker *worker) { while (pending) { rpc_pending_fd *next = pending->next; (void)worker_add_connection(worker, pending->fd); - free(pending); + rpc_mem_free(pending); pending = next; } } @@ -527,7 +528,7 @@ static void worker_destroy(rpc_worker *worker) { while (pending) { rpc_pending_fd *next = pending->next; close(pending->fd); - free(pending); + rpc_mem_free(pending); pending = next; } pthread_mutex_destroy(&worker->pending_mutex); @@ -541,14 +542,14 @@ static void worker_destroy(rpc_worker *worker) { static int server_ensure_workers(rpc_server *server) { if (server->workers_ready) { return 0; } if (server->worker_count == 0) { server->worker_count = cpu_count(); } - server->workers = calloc(server->worker_count, sizeof(*server->workers)); + server->workers = rpc_mem_calloc(server->worker_count, sizeof(*server->workers)); if (!server->workers) { return -1; } for (uint32_t i = 0; i < server->worker_count; ++i) { if (worker_init(server, &server->workers[i], i) != 0) { for (uint32_t j = 0; j <= i; ++j) { worker_destroy(&server->workers[j]); } - free(server->workers); + rpc_mem_free(server->workers); server->workers = NULL; return -1; } @@ -559,7 +560,7 @@ static int server_ensure_workers(rpc_server *server) { int rpc_server_init(rpc_server **out_server) { if (!out_server) { return -1; } - rpc_server *server = calloc(1, sizeof(*server)); + rpc_server *server = rpc_mem_calloc(1, sizeof(*server)); if (!server) { return -1; } atomic_init(&server->stopping, false); atomic_init(&server->next_worker, 0); @@ -735,27 +736,39 @@ void rpc_server_destroy(rpc_server *server) { } worker_destroy(&server->workers[i]); } - free(server->workers); + rpc_mem_free(server->workers); } if (server->routes_ready) { rpc_routes_destroy(&server->routes); } - free(server); + rpc_mem_free(server); } -static int server_add_route_id(rpc_server *server, uint64_t proc_id, rpc_handler_fn handler, void *user_data) { - return server ? rpc_routes_add(&server->routes, proc_id, handler, user_data) : -1; +static int server_add_route_id(rpc_server *server, uint64_t proc_id, rpc_handler_fn handler, void *user_data, + rpc_route_finalizer_fn finalizer) { + return server ? rpc_routes_add_ex(&server->routes, proc_id, handler, user_data, finalizer, 0) : -1; } int rpc_server_add_route_name(rpc_server *server, const char *proc_name, rpc_handler_fn handler, void *user_data) { - return server_add_route_id(server, rpc_proc_id(proc_name), handler, user_data); + return rpc_server_add_route_name_ex(server, proc_name, handler, user_data, NULL); } -static int server_add_async_route_id(rpc_server *server, uint64_t proc_id, rpc_handler_fn handler, void *user_data) { - return server ? rpc_routes_add_ex(&server->routes, proc_id, handler, user_data, 1) : -1; +int rpc_server_add_route_name_ex(rpc_server *server, const char *proc_name, rpc_handler_fn handler, void *user_data, + rpc_route_finalizer_fn finalizer) { + return server_add_route_id(server, rpc_proc_id(proc_name), handler, user_data, finalizer); +} + +static int server_add_async_route_id(rpc_server *server, uint64_t proc_id, rpc_handler_fn handler, void *user_data, + rpc_route_finalizer_fn finalizer) { + return server ? rpc_routes_add_ex(&server->routes, proc_id, handler, user_data, finalizer, 1) : -1; } int rpc_server_add_async_route_name(rpc_server *server, const char *proc_name, rpc_handler_fn handler, void *user_data) { - return server_add_async_route_id(server, rpc_proc_id(proc_name), handler, user_data); + return rpc_server_add_async_route_name_ex(server, proc_name, handler, user_data, NULL); +} + +int rpc_server_add_async_route_name_ex(rpc_server *server, const char *proc_name, rpc_handler_fn handler, + void *user_data, rpc_route_finalizer_fn finalizer) { + return server_add_async_route_id(server, rpc_proc_id(proc_name), handler, user_data, finalizer); } static int server_remove_route_id(rpc_server *server, uint64_t proc_id) { diff --git a/tests/test_protocol.c b/tests/test_protocol.c index 60f7dd7..0a3abac 100644 --- a/tests/test_protocol.c +++ b/tests/test_protocol.c @@ -2,8 +2,33 @@ #include #include +#include #include +typedef struct alloc_stats { + size_t allocs; + size_t reallocs; + size_t frees; +} alloc_stats; + +static void *test_alloc(void *ctx, size_t size) { + alloc_stats *stats = ctx; + stats->allocs++; + return malloc(size); +} + +static void *test_realloc(void *ctx, void *ptr, size_t size) { + alloc_stats *stats = ctx; + stats->reallocs++; + return realloc(ptr, size); +} + +static void test_free(void *ctx, void *ptr) { + alloc_stats *stats = ctx; + stats->frees++; + free(ptr); +} + int main(void) { uint8_t buf[RPC_HEADER_SIZE]; rpc_header h = { @@ -25,6 +50,15 @@ int main(void) { buf[0] = 99; assert(rpc_header_decode(buf, &out) != 0); + alloc_stats stats = {0}; + rpc_allocator allocator = { + .ctx = &stats, + .alloc = test_alloc, + .realloc = test_realloc, + .free = test_free, + }; + assert(rpc_set_allocator(&allocator) == 0); + rpc_writer w; rpc_writer_init(&w); assert(rpc_writer_null(&w) == 0); @@ -53,5 +87,8 @@ int main(void) { uint8_t malformed[] = {RPC_TYPE_STRING, 0, 0, 0, 10, 'x'}; assert(rpc_payload_decode(malformed, sizeof(malformed), &values, &count) != 0); rpc_writer_free(&w); + assert(stats.reallocs > 0); + assert(stats.frees > 0); + assert(rpc_set_allocator(NULL) == 0); return 0; } diff --git a/tests/test_routes.c b/tests/test_routes.c index 01bcd1c..1b2440a 100644 --- a/tests/test_routes.c +++ b/tests/test_routes.c @@ -20,21 +20,32 @@ static int handler_b(rpc_ctx *ctx, const rpc_value *args, size_t argc, rpc_write return 0; } +static void count_finalizer(void *user_data) { + int *count = user_data; + (*count)++; +} + int main(void) { rpc_routes routes; rpc_route route; + int finalized_a = 0; + int finalized_b = 0; assert(rpc_routes_init(&routes) == 0); assert(rpc_routes_lookup(&routes, 10, &route) != 0); - assert(rpc_routes_add(&routes, 10, handler_a, (void *)1) == 0); + assert(rpc_routes_add_ex(&routes, 10, handler_a, &finalized_a, count_finalizer, 0) == 0); assert(rpc_routes_lookup(&routes, 10, &route) == 0); assert(route.handler == handler_a); - assert(route.user_data == (void *)1); - assert(rpc_routes_add(&routes, 10, handler_b, (void *)2) == 0); + assert(route.user_data == &finalized_a); + assert(rpc_routes_add_ex(&routes, 10, handler_b, &finalized_b, count_finalizer, 0) == 0); + assert(finalized_a == 1); assert(rpc_routes_lookup(&routes, 10, &route) == 0); assert(route.handler == handler_b); - assert(route.user_data == (void *)2); + assert(route.user_data == &finalized_b); assert(rpc_routes_remove(&routes, 10) == 0); + assert(finalized_b == 1); assert(rpc_routes_lookup(&routes, 10, &route) != 0); rpc_routes_destroy(&routes); + assert(finalized_a == 1); + assert(finalized_b == 1); return 0; }