locked rpc for thread safety

This commit is contained in:
Sparchatus 2022-04-06 14:05:44 +00:00
parent 3c8987f5f3
commit 0c224f9066
2 changed files with 79 additions and 14 deletions

View File

@ -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;
};
/**

View File

@ -15,7 +15,7 @@
#include <aos/aos.h>
#include <aos/aos_rpc.h>
#include <aos/dispatcher_arch.h>
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;
}