[PATCH 2/4] RISC-V: Factor-out per-hart MPXY shared-memory acquisition

From: Anup Patel

Date: Wed Sep 30 2026 - 11:28:02 EST


From: Amirreza Zarrabi <amirreza.zarrabi@xxxxxxxxxxxxxxxx>

Most of the SBI MPXY functions need to access MPXY shared-memory so
they have to acquire underlying host CPU before accessing the MPXY
shared-memory and release the host CPU after the work is done.

Factor-out the above mentioned per-hart MPXY shared-memory and host
CPU acquisition into mpxy_local_get()/put() functions.

Signed-off-by: Amirreza Zarrabi <amirreza.zarrabi@xxxxxxxxxxxxxxxx>
Signed-off-by: Anup Patel <anup.patel@xxxxxxxxxxxxxxxx>
---
arch/riscv/kernel/sbi_mpxy.c | 159 ++++++++++++++++++++++-------------
1 file changed, 100 insertions(+), 59 deletions(-)

diff --git a/arch/riscv/kernel/sbi_mpxy.c b/arch/riscv/kernel/sbi_mpxy.c
index 17a2cee21311..b2b34991fd42 100644
--- a/arch/riscv/kernel/sbi_mpxy.c
+++ b/arch/riscv/kernel/sbi_mpxy.c
@@ -40,6 +40,26 @@ static DEFINE_PER_CPU(struct mpxy_local, mpxy_local);
static unsigned long mpxy_shmem_size;
static bool mpxy_shmem_init_done;

+static int mpxy_local_get(struct mpxy_local **out)
+{
+ struct mpxy_local *mpxy;
+
+ get_cpu();
+ mpxy = this_cpu_ptr(&mpxy_local);
+ if (!mpxy->shmem_active) {
+ put_cpu();
+ return -ENODEV;
+ }
+
+ *out = mpxy;
+ return 0;
+}
+
+static void mpxy_local_put(void)
+{
+ put_cpu();
+}
+
unsigned long sbi_mpxy_shmem_size(void)
{
if (!mpxy_shmem_init_done)
@@ -50,53 +70,61 @@ EXPORT_SYMBOL_GPL(sbi_mpxy_shmem_size);

int sbi_mpxy_get_channel_count(u32 *channel_count)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
- struct sbi_mpxy_channel_ids_data *sdata = mpxy->shmem;
+ struct sbi_mpxy_channel_ids_data *sdata;
+ struct mpxy_local *mpxy;
u32 remaining, returned;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!channel_count)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;
+ sdata = mpxy->shmem;

/* Get the remaining and returned fields to calculate total */
sret = sbi_ecall(SBI_EXT_MPXY, SBI_EXT_MPXY_GET_CHANNEL_IDS,
0, 0, 0, 0, 0, 0);
- if (sret.error)
- goto err_put_cpu;
+ if (sret.error) {
+ rc = sbi_err_map_linux_errno(sret.error);
+ goto out;
+ }

remaining = le32_to_cpu(sdata->remaining);
returned = le32_to_cpu(sdata->returned);
*channel_count = remaining + returned;

-err_put_cpu:
- put_cpu();
- return sbi_err_map_linux_errno(sret.error);
+out:
+ mpxy_local_put();
+ return rc;
}
EXPORT_SYMBOL_GPL(sbi_mpxy_get_channel_count);

int sbi_mpxy_get_channel_ids(u32 channel_count, u32 *channel_ids)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
- struct sbi_mpxy_channel_ids_data *sdata = mpxy->shmem;
u32 remaining, returned, count, start_index = 0;
+ struct sbi_mpxy_channel_ids_data *sdata;
+ struct mpxy_local *mpxy;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!channel_count || !channel_ids)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;
+ sdata = mpxy->shmem;

do {
sret = sbi_ecall(SBI_EXT_MPXY, SBI_EXT_MPXY_GET_CHANNEL_IDS,
start_index, 0, 0, 0, 0, 0);
- if (sret.error)
- goto err_put_cpu;
+ if (sret.error) {
+ rc = sbi_err_map_linux_errno(sret.error);
+ goto out;
+ }

remaining = le32_to_cpu(sdata->remaining);
returned = le32_to_cpu(sdata->returned);
@@ -107,56 +135,60 @@ int sbi_mpxy_get_channel_ids(u32 channel_count, u32 *channel_ids)
start_index += count;
} while (remaining && start_index < channel_count);

-err_put_cpu:
- put_cpu();
- return sbi_err_map_linux_errno(sret.error);
+out:
+ mpxy_local_put();
+ return rc;
}
EXPORT_SYMBOL_GPL(sbi_mpxy_get_channel_ids);

int sbi_mpxy_read_attrs(u32 channel_id, u32 base_attrid, u32 attr_count,
u32 *attrs_buf)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
+ struct mpxy_local *mpxy;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!attr_count || !attrs_buf)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;

sret = sbi_ecall(SBI_EXT_MPXY, SBI_EXT_MPXY_READ_ATTRS,
channel_id, base_attrid, attr_count, 0, 0, 0);
- if (sret.error)
- goto err_put_cpu;
+ if (sret.error) {
+ rc = sbi_err_map_linux_errno(sret.error);
+ goto out;
+ }

memcpy_from_le32(attrs_buf, (__le32 *)mpxy->shmem, attr_count);

-err_put_cpu:
- put_cpu();
- return sbi_err_map_linux_errno(sret.error);
+out:
+ mpxy_local_put();
+ return rc;
}
EXPORT_SYMBOL_GPL(sbi_mpxy_read_attrs);

int sbi_mpxy_write_attrs(u32 channel_id, u32 base_attrid, u32 attr_count,
u32 *attrs_buf)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
+ struct mpxy_local *mpxy;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!attr_count || !attrs_buf)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;

memcpy_to_le32((__le32 *)mpxy->shmem, attrs_buf, attr_count);
sret = sbi_ecall(SBI_EXT_MPXY, SBI_EXT_MPXY_WRITE_ATTRS,
channel_id, base_attrid, attr_count, 0, 0, 0);

- put_cpu();
+ mpxy_local_put();
return sbi_err_map_linux_errno(sret.error);
}
EXPORT_SYMBOL_GPL(sbi_mpxy_write_attrs);
@@ -166,16 +198,17 @@ int sbi_mpxy_send_message_with_resp(u32 channel_id, u32 msg_id,
void *rx, unsigned long max_rx_len,
unsigned long *rx_len)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
+ struct mpxy_local *mpxy;
unsigned long rx_bytes;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!tx && tx_len)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;

/* Message protocols allowed to have no data in messages */
if (tx_len)
@@ -186,8 +219,8 @@ int sbi_mpxy_send_message_with_resp(u32 channel_id, u32 msg_id,
if (rx && !sret.error) {
rx_bytes = sret.value;
if (rx_bytes > max_rx_len) {
- put_cpu();
- return -ENOSPC;
+ rc = -ENOSPC;
+ goto out;
}

memcpy(rx, mpxy->shmem, rx_bytes);
@@ -195,23 +228,26 @@ int sbi_mpxy_send_message_with_resp(u32 channel_id, u32 msg_id,
*rx_len = rx_bytes;
}

- put_cpu();
- return sbi_err_map_linux_errno(sret.error);
+ rc = sbi_err_map_linux_errno(sret.error);
+out:
+ mpxy_local_put();
+ return rc;
}
EXPORT_SYMBOL_GPL(sbi_mpxy_send_message_with_resp);

int sbi_mpxy_send_message_without_resp(u32 channel_id, u32 msg_id,
void *tx, unsigned long tx_len)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
+ struct mpxy_local *mpxy;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!tx && tx_len)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;

/* Message protocols allowed to have no data in messages */
if (tx_len)
@@ -220,8 +256,9 @@ int sbi_mpxy_send_message_without_resp(u32 channel_id, u32 msg_id,
sret = sbi_ecall(SBI_EXT_MPXY, SBI_EXT_MPXY_SEND_MSG_WITHOUT_RESP,
channel_id, msg_id, tx_len, 0, 0, 0);

- put_cpu();
- return sbi_err_map_linux_errno(sret.error);
+ rc = sbi_err_map_linux_errno(sret.error);
+ mpxy_local_put();
+ return rc;
}
EXPORT_SYMBOL_GPL(sbi_mpxy_send_message_without_resp);

@@ -229,32 +266,36 @@ int sbi_mpxy_get_notifications(u32 channel_id,
struct sbi_mpxy_notification_data *notif_data,
unsigned long *events_data_len)
{
- struct mpxy_local *mpxy = this_cpu_ptr(&mpxy_local);
+ struct mpxy_local *mpxy;
struct sbiret sret;
+ int rc = 0;

- if (!mpxy->shmem_active)
- return -ENODEV;
if (!notif_data || !events_data_len)
return -EINVAL;

- get_cpu();
+ rc = mpxy_local_get(&mpxy);
+ if (rc)
+ return rc;

sret = sbi_ecall(SBI_EXT_MPXY, SBI_EXT_MPXY_GET_NOTIFICATION_EVENTS,
channel_id, 0, 0, 0, 0, 0);
- if (sret.error)
- goto err_put_cpu;
+ if (sret.error) {
+ rc = sbi_err_map_linux_errno(sret.error);
+ goto out;
+ }
if (sret.value < 0 || mpxy_shmem_size < sizeof(*notif_data) ||
sret.value > mpxy_shmem_size - sizeof(*notif_data)) {
- put_cpu();
- return -EOVERFLOW;
+ rc = -EOVERFLOW;
+ goto out;
}

memcpy(notif_data, mpxy->shmem, sret.value + sizeof(*notif_data));
*events_data_len = sret.value;

-err_put_cpu:
- put_cpu();
- return sbi_err_map_linux_errno(sret.error);
+ rc = sbi_err_map_linux_errno(sret.error);
+out:
+ mpxy_local_put();
+ return rc;
}
EXPORT_SYMBOL_GPL(sbi_mpxy_get_notifications);

--
2.43.0