From 0c224f9066d908eafefc71ebc80132d3f9044468 Mon Sep 17 00:00:00 2001 From: Sparchatus Date: Wed, 6 Apr 2022 14:05:44 +0000 Subject: [PATCH] locked rpc for thread safety --- include/aos/aos_rpc.h | 3 ++ lib/aos/aos_rpc.c | 90 ++++++++++++++++++++++++++++++++++++------- 2 files changed, 79 insertions(+), 14 deletions(-) diff --git a/include/aos/aos_rpc.h b/include/aos/aos_rpc.h index 85cebd3..8d4609f 100644 --- a/include/aos/aos_rpc.h +++ b/include/aos/aos_rpc.h @@ -41,6 +41,9 @@ struct aos_rpc { struct lmp_chan *chan; struct waitset ws; void *shared_mem; + + // mutually exclusive access for threads accessing rpc + struct thread_mutex lock; }; /** diff --git a/lib/aos/aos_rpc.c b/lib/aos/aos_rpc.c index f587b44..50b421f 100644 --- a/lib/aos/aos_rpc.c +++ b/lib/aos/aos_rpc.c @@ -15,7 +15,7 @@ #include #include - +#include static void noop_callback(void *arg) { } @@ -27,6 +27,14 @@ static errval_t do_aos_rpc( ) { errval_t err; + // the calling thread MUST have the rpc lock at this point + // we check to be sure we did not forget to lock in any of the rpc functions! + assert(rpc->lock.locked > 0); + dispatcher_handle_t handle = disp_disable(); + struct dispatcher_generic *disp_gen = get_dispatcher_generic(handle); + assert(rpc->lock.holder == disp_gen->current); + disp_enable(handle); + // allocate capability if we expect one and the recv slot is empty if (ret_cap != NULL && rpc->chan->endpoint->k.recv_cptr == 0){ lmp_chan_alloc_recv_slot(rpc->chan); @@ -67,25 +75,37 @@ errval_t aos_rpc_send_number(struct aos_rpc *rpc, uintptr_t num) { // Implement functionality to send a number over the channel // given channel and wait until the ack gets returned. - return do_aos_rpc( + errval_t err; + + thread_mutex_lock(&rpc->lock); + err = do_aos_rpc( rpc, RPC_MTYPE_SEND_NUMBER, NULL_CAP, 0, num, 0, NULL, NULL, NULL, NULL); + thread_mutex_unlock(&rpc->lock); + + return err; } errval_t aos_rpc_send_string(struct aos_rpc *rpc, const char *string) { // Implement functionality to send a string over the given channel // and wait for a response. + errval_t err; + size_t string_size = strlen(string) + 1; if (string_size > RPC_SHARED_SIZE) return AOS_ERR_RPC_ARG_TOO_BIG; - memcpy(rpc->shared_mem, string, string_size); - return do_aos_rpc( + thread_mutex_lock(&rpc->lock); + memcpy(rpc->shared_mem, string, string_size); + err = do_aos_rpc( rpc, RPC_MTYPE_SEND_STRING, NULL_CAP, string_size, 0, 0, NULL, NULL, NULL, NULL ); + thread_mutex_unlock(&rpc->lock); + + return err; } errval_t @@ -93,11 +113,18 @@ aos_rpc_get_ram_cap(struct aos_rpc *rpc, size_t bytes, size_t alignment, struct capref *ret_cap, size_t *ret_bytes) { // Implement functionality to request a RAM capability over the // given channel and wait until it is delivered. - return do_aos_rpc( + + errval_t err; + + thread_mutex_lock(&rpc->lock); + err = do_aos_rpc( rpc, RPC_MTYPE_GET_RAM_CAP, NULL_CAP, 0, bytes, alignment, ret_cap, NULL, ret_bytes, NULL ); + thread_mutex_unlock(&rpc->lock); + + return err; } @@ -105,11 +132,16 @@ errval_t aos_rpc_serial_getchar(struct aos_rpc *rpc, char *retc) { // Implement functionality to request a character from // the serial driver. + errval_t err; uintptr_t retval; - errval_t err = do_aos_rpc( + + thread_mutex_lock(&rpc->lock); + err = do_aos_rpc( rpc, RPC_MTYPE_SERIAL_GETCHAR, NULL_CAP, 0, 0, 0, NULL, NULL, &retval, NULL); + thread_mutex_unlock(&rpc->lock); + *retc = retval; return err; } @@ -119,11 +151,17 @@ errval_t aos_rpc_serial_putchar(struct aos_rpc *rpc, char c) { // Implement functionality to send a character to the // serial port. - return do_aos_rpc( + errval_t err; + + thread_mutex_lock(&rpc->lock); + err = do_aos_rpc( rpc, RPC_MTYPE_SERIAL_PUTCHAR, NULL_CAP, 0, c, 0, NULL, NULL, NULL, NULL ); + thread_mutex_unlock(&rpc->lock); + + return err; } // putchar for terminal printing is too slow so we write buffers instead @@ -136,15 +174,17 @@ aos_rpc_serial_write(struct aos_rpc *rpc, const char *buf, size_t buf_len) { size_t total_bytes = 0; while(buf_len > 0) { size_t chunk_len = MIN(buf_len, RPC_SHARED_SIZE); - - memcpy(rpc->shared_mem, buf, chunk_len); - size_t written_bytes; + + thread_mutex_lock(&rpc->lock); + memcpy(rpc->shared_mem, buf, chunk_len); err = do_aos_rpc( rpc, RPC_MTYPE_SERIAL_WRITE, NULL_CAP, chunk_len, 0, 0, NULL, NULL, &written_bytes, NULL ); + thread_mutex_unlock(&rpc->lock); + if (err_is_fail(err)) { DEBUG_ERR(err, "Failed to write the full buffer"); return total_bytes; @@ -175,8 +215,9 @@ aos_rpc_serial_read(struct aos_rpc *rpc, char *buf, size_t buf_len) { size_t total_bytes = 0; while(buf_len > 0) { size_t chunk_len = MIN(buf_len, RPC_SHARED_SIZE); - size_t read_bytes; + + thread_mutex_lock(&rpc->lock); err = do_aos_rpc( rpc, RPC_MTYPE_SERIAL_READ, NULL_CAP, 0, chunk_len, 0, @@ -189,6 +230,8 @@ aos_rpc_serial_read(struct aos_rpc *rpc, char *buf, size_t buf_len) { assert(read_bytes <= chunk_len); memcpy(buf, rpc->shared_mem, read_bytes); + thread_mutex_unlock(&rpc->lock); + total_bytes += read_bytes; // if the receiver was not able to write all we have given, return for now @@ -209,14 +252,17 @@ aos_rpc_process_spawn(struct aos_rpc *rpc, char *cmdline, // implement spawn new process rpc size_t cmdline_size = strlen(cmdline) + 1; if (cmdline_size > RPC_SHARED_SIZE) return AOS_ERR_RPC_ARG_TOO_BIG; - memcpy(rpc->shared_mem, cmdline, cmdline_size); + thread_mutex_lock(&rpc->lock); + memcpy(rpc->shared_mem, cmdline, cmdline_size); uintptr_t retval; errval_t err = do_aos_rpc( rpc, RPC_MTYPE_PROCESS_SPAWN, NULL_CAP, cmdline_size, core, 0, NULL, NULL, &retval, NULL ); + thread_mutex_unlock(&rpc->lock); + *newpid = retval; return err; } @@ -227,15 +273,22 @@ errval_t aos_rpc_process_get_name(struct aos_rpc *rpc, domainid_t pid, char **name) { // implement name lookup for process given a process id size_t name_len; - errval_t err = do_aos_rpc( + errval_t err; + + thread_mutex_lock(&rpc->lock); + err = do_aos_rpc( rpc, RPC_MTYPE_PROCESS_GET_NAME, NULL_CAP, 0, pid, 0, NULL, &name_len, NULL, NULL ); if (err_is_fail(err)) return err; + *name = malloc(name_len); if (*name == NULL) return LIB_ERR_MALLOC_FAIL; + memcpy(*name, rpc->shared_mem, name_len); + thread_mutex_unlock(&rpc->lock); + assert(name_len != 0 && (*name)[name_len - 1] == '\0'); return err; } @@ -246,16 +299,24 @@ aos_rpc_process_get_all_pids(struct aos_rpc *rpc, domainid_t **pids, size_t *pid_count) { // implement process id discovery size_t ret_len; - errval_t err = do_aos_rpc( + errval_t err; + + thread_mutex_lock(&rpc->lock); + err = do_aos_rpc( rpc, RPC_MTYPE_PROCESS_GET_ALL_PIDS, NULL_CAP, 0, 0, 0, NULL, &ret_len, NULL, NULL ); if (err_is_fail(err)) return err; + *pid_count = ret_len / sizeof(**pids); *pids = malloc(ret_len); if (*pids == NULL) return LIB_ERR_MALLOC_FAIL; + memcpy(*pids, rpc->shared_mem, ret_len); + + thread_mutex_unlock(&rpc->lock); + return err; } @@ -265,6 +326,7 @@ errval_t aos_rpc_init(struct aos_rpc *rpc, struct lmp_chan *chan, void *shared_m rpc->chan = chan; rpc->shared_mem = shared_mem; waitset_init(&rpc->ws); + thread_mutex_init(&rpc->lock); return SYS_ERR_OK; }