Implement send_data_sync() for RPMI MPXY channels, reusing the existing RPMI message dispatch and returning its status directly instead of going through mbox_chan_txdone().
Factor per-hart shared-memory acquisition into mpxy_shmem_get()/ mpxy_shmem_put() so both the queued and the new synchronous path pin the CPU around the same per-hart buffer, keeping a request and its response on one hart. Calls on separate harts remain independent, subject to firmware support.
Add rpmi_mbox_send_message_sync() as the RPMI counterpart to rpmi_mbox_send_message(), for use by the upcoming TEE transport.
Signed-off-by: Amirreza Zarrabi amirreza.zarrabi@oss.qualcomm.com --- drivers/mailbox/riscv-sbi-mpxy-mbox.c | 196 ++++++++++++++++++----------- include/linux/mailbox/riscv-rpmi-message.h | 13 ++ 2 files changed, 139 insertions(+), 70 deletions(-)
diff --git a/drivers/mailbox/riscv-sbi-mpxy-mbox.c b/drivers/mailbox/riscv-sbi-mpxy-mbox.c index ea69c6b6b4f9..5ca1b6d87f5c 100644 --- a/drivers/mailbox/riscv-sbi-mpxy-mbox.c +++ b/drivers/mailbox/riscv-sbi-mpxy-mbox.c @@ -125,54 +125,83 @@ static DEFINE_PER_CPU(struct mpxy_local, mpxy_local); static unsigned long mpxy_shmem_size; static bool mpxy_shmem_init_done;
+static int mpxy_shmem_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_shmem_put(void) +{ + put_cpu(); +} + static int 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 mpxy_local *mpxy; + struct sbi_mpxy_channel_ids_data *sdata; u32 remaining, returned; struct sbiret sret; + int rc;
- if (!mpxy->shmem_active) - return -ENODEV; if (!channel_count) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_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; + rc = 0;
-err_put_cpu: - put_cpu(); - return sbi_err_map_linux_errno(sret.error); +out: + mpxy_shmem_put(); + return rc; }
static int 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; + struct mpxy_local *mpxy; + struct sbi_mpxy_channel_ids_data *sdata; u32 remaining, returned, count, start_index = 0; struct sbiret sret; + int rc;
- if (!mpxy->shmem_active) - return -ENODEV; if (!channel_count || !channel_ids) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_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); @@ -182,55 +211,61 @@ static int mpxy_get_channel_ids(u32 channel_count, u32 *channel_ids) memcpy_from_le32(&channel_ids[start_index], sdata->channel_array, count); start_index += count; } while (remaining && start_index < channel_count); + rc = 0;
-err_put_cpu: - put_cpu(); - return sbi_err_map_linux_errno(sret.error); +out: + mpxy_shmem_put(); + return rc; }
static int 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;
- if (!mpxy->shmem_active) - return -ENODEV; if (!attr_count || !attrs_buf) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_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); + rc = 0;
-err_put_cpu: - put_cpu(); - return sbi_err_map_linux_errno(sret.error); +out: + mpxy_shmem_put(); + return rc; }
static int 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;
- if (!mpxy->shmem_active) - return -ENODEV; if (!attr_count || !attrs_buf) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_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_shmem_put(); return sbi_err_map_linux_errno(sret.error); }
@@ -239,16 +274,17 @@ static int 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;
- if (!mpxy->shmem_active) - return -ENODEV; if (!tx && tx_len) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_get(&mpxy); + if (rc) + return rc;
/* Message protocols allowed to have no data in messages */ if (tx_len) @@ -259,8 +295,8 @@ static int 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); @@ -268,22 +304,25 @@ static int 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_shmem_put(); + return rc; }
static int 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;
- if (!mpxy->shmem_active) - return -ENODEV; if (!tx && tx_len) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_get(&mpxy); + if (rc) + return rc;
/* Message protocols allowed to have no data in messages */ if (tx_len) @@ -292,40 +331,45 @@ static int 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_shmem_put(); + return rc; }
static int 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;
- if (!mpxy->shmem_active) - return -ENODEV; if (!notif_data || !events_data_len) return -EINVAL;
- get_cpu(); + rc = mpxy_shmem_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_shmem_put(); + return rc; }
static int mpxy_get_shmem_size(unsigned long *shmem_size) @@ -402,8 +446,8 @@ struct mpxy_mbox {
/* ====== MPXY RPMI processing ====== */
-static void mpxy_mbox_send_rpmi_data(struct mpxy_mbox_channel *mchan, - struct rpmi_mbox_message *msg) +static int mpxy_mbox_send_rpmi_data(struct mpxy_mbox_channel *mchan, + struct rpmi_mbox_message *msg) { msg->error = 0; switch (msg->type) { @@ -474,6 +518,8 @@ static void mpxy_mbox_send_rpmi_data(struct mpxy_mbox_channel *mchan, msg->error = -EOPNOTSUPP; break; } + + return msg->error; }
static void mpxy_mbox_peek_rpmi_data(struct mbox_chan *chan, @@ -516,12 +562,21 @@ static int mpxy_mbox_send_data(struct mbox_chan *chan, void *data) { struct mpxy_mbox_channel *mchan = chan->con_priv;
- if (mchan->attrs.msg_proto_id == SBI_MPXY_MSGPROTO_RPMI_ID) { - mpxy_mbox_send_rpmi_data(mchan, data); - return 0; - } + if (mchan->attrs.msg_proto_id != SBI_MPXY_MSGPROTO_RPMI_ID) + return -EOPNOTSUPP; + + mpxy_mbox_send_rpmi_data(mchan, data); + return 0; +} + +static int mpxy_mbox_send_data_sync(struct mbox_chan *chan, void *data) +{ + struct mpxy_mbox_channel *mchan = chan->con_priv; + + if (mchan->attrs.msg_proto_id != SBI_MPXY_MSGPROTO_RPMI_ID) + return -EOPNOTSUPP;
- return -EOPNOTSUPP; + return mpxy_mbox_send_rpmi_data(mchan, data); }
static bool mpxy_mbox_peek_data(struct mbox_chan *chan) @@ -713,10 +768,11 @@ static void mpxy_mbox_shutdown(struct mbox_chan *chan) }
static const struct mbox_chan_ops mpxy_mbox_ops = { - .send_data = mpxy_mbox_send_data, - .peek_data = mpxy_mbox_peek_data, - .startup = mpxy_mbox_startup, - .shutdown = mpxy_mbox_shutdown, + .send_data = mpxy_mbox_send_data, + .send_data_sync = mpxy_mbox_send_data_sync, + .peek_data = mpxy_mbox_peek_data, + .startup = mpxy_mbox_startup, + .shutdown = mpxy_mbox_shutdown, };
/* ====== MPXY platform driver ===== */ diff --git a/include/linux/mailbox/riscv-rpmi-message.h b/include/linux/mailbox/riscv-rpmi-message.h index d5362b5821f9..16fa85f65d80 100644 --- a/include/linux/mailbox/riscv-rpmi-message.h +++ b/include/linux/mailbox/riscv-rpmi-message.h @@ -251,4 +251,17 @@ static inline int rpmi_mbox_send_message(struct mbox_chan *chan, return ret; }
+static inline int rpmi_mbox_send_message_sync(struct mbox_chan *chan, + struct rpmi_mbox_message *msg) +{ + int ret; + + /* Send message for the underlying mailbox channel synchronously */ + ret = mbox_send_message_sync(chan, msg); + if (ret) + return ret; + + return msg->error; +} + #endif /* _LINUX_RISCV_RPMI_MESSAGE_H_ */