#define pr_fmt(fmt) KBUILD_MODNAME ": " fmt
#include <linux/cleanup.h>
#include <linux/err.h>
#include <linux/errno.h>
#include <linux/io.h>
#include <linux/kexec_handover.h>
#include <linux/kho/abi/luo.h>
#include <linux/list_private.h>
#include <linux/liveupdate.h>
#include <linux/module.h>
#include <linux/mutex.h>
#include <linux/slab.h>
#include "luo_internal.h"
#define LUO_FLB_PGCNT 1ul
#define LUO_FLB_MAX (((LUO_FLB_PGCNT << PAGE_SHIFT) - \
sizeof(struct luo_flb_header_ser)) / sizeof(struct luo_flb_ser))
struct luo_flb_header {
struct luo_flb_header_ser *header_ser;
struct luo_flb_ser *ser;
bool active;
};
struct luo_flb_global {
struct luo_flb_header incoming;
struct luo_flb_header outgoing;
struct list_head list;
long count;
};
static struct luo_flb_global luo_flb_global = {
.list = LIST_HEAD_INIT(luo_flb_global.list),
};
struct luo_flb_link {
struct liveupdate_flb *flb;
struct list_head list;
};
static struct luo_flb_private *luo_flb_get_private(struct liveupdate_flb *flb)
{
struct luo_flb_private *private = &ACCESS_PRIVATE(flb, private);
static DEFINE_SPINLOCK(luo_flb_init_lock);
if (smp_load_acquire(&private->initialized))
return private;
guard(spinlock)(&luo_flb_init_lock);
if (!private->initialized) {
mutex_init(&private->incoming.lock);
mutex_init(&private->outgoing.lock);
INIT_LIST_HEAD(&private->list);
private->users = 0;
smp_store_release(&private->initialized, true);
}
return private;
}
static int luo_flb_file_preserve_one(struct liveupdate_flb *flb)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
scoped_guard(mutex, &private->outgoing.lock) {
if (!refcount_read(&private->outgoing.count)) {
struct liveupdate_flb_op_args args = {0};
int err;
if (!try_module_get(flb->ops->owner))
return -ENODEV;
args.flb = flb;
err = flb->ops->preserve(&args);
if (err) {
module_put(flb->ops->owner);
return err;
}
private->outgoing.data = args.data;
private->outgoing.obj = args.obj;
refcount_set(&private->outgoing.count, 1);
} else {
refcount_inc(&private->outgoing.count);
}
}
return 0;
}
static void luo_flb_file_unpreserve_one(struct liveupdate_flb *flb)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
scoped_guard(mutex, &private->outgoing.lock) {
if (refcount_dec_and_test(&private->outgoing.count)) {
struct liveupdate_flb_op_args args = {0};
args.flb = flb;
args.data = private->outgoing.data;
args.obj = private->outgoing.obj;
if (flb->ops->unpreserve)
flb->ops->unpreserve(&args);
private->outgoing.data = 0;
private->outgoing.obj = NULL;
module_put(flb->ops->owner);
}
}
}
static int luo_flb_retrieve_one(struct liveupdate_flb *flb)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
struct luo_flb_header *fh = &luo_flb_global.incoming;
struct liveupdate_flb_op_args args = {0};
bool found = false;
int err;
lockdep_assert_held(&private->incoming.lock);
if (private->incoming.finished)
return -ENODATA;
if (private->incoming.retrieved)
return 0;
if (!fh->active)
return -ENODATA;
for (int i = 0; i < fh->header_ser->count; i++) {
if (!strcmp(fh->ser[i].name, flb->compatible)) {
private->incoming.data = fh->ser[i].data;
refcount_set(&private->incoming.count, fh->ser[i].count);
found = true;
break;
}
}
if (!found)
return -ENOENT;
if (!try_module_get(flb->ops->owner))
return -ENODEV;
args.flb = flb;
args.data = private->incoming.data;
err = flb->ops->retrieve(&args);
if (err) {
module_put(flb->ops->owner);
return err;
}
private->incoming.obj = args.obj;
private->incoming.retrieved = true;
return 0;
}
void liveupdate_flb_put_incoming(struct liveupdate_flb *flb)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
struct liveupdate_flb_op_args args = {0};
scoped_guard(mutex, &private->incoming.lock) {
if (!refcount_dec_and_test(&private->incoming.count))
return;
if (!private->incoming.retrieved) {
int err = luo_flb_retrieve_one(flb);
if (WARN_ON(err))
return;
}
args.flb = flb;
args.obj = private->incoming.obj;
flb->ops->finish(&args);
private->incoming.data = 0;
private->incoming.obj = NULL;
private->incoming.finished = true;
module_put(flb->ops->owner);
}
}
int luo_flb_file_preserve(struct liveupdate_file_handler *fh)
{
struct list_head *flb_list = &ACCESS_PRIVATE(fh, flb_list);
struct luo_flb_link *iter;
int err = 0;
down_read(&luo_register_rwlock);
list_for_each_entry(iter, flb_list, list) {
err = luo_flb_file_preserve_one(iter->flb);
if (err)
goto exit_err;
}
up_read(&luo_register_rwlock);
return 0;
exit_err:
list_for_each_entry_continue_reverse(iter, flb_list, list)
luo_flb_file_unpreserve_one(iter->flb);
up_read(&luo_register_rwlock);
return err;
}
void luo_flb_file_unpreserve(struct liveupdate_file_handler *fh)
{
struct list_head *flb_list = &ACCESS_PRIVATE(fh, flb_list);
struct luo_flb_link *iter;
guard(rwsem_read)(&luo_register_rwlock);
list_for_each_entry_reverse(iter, flb_list, list)
luo_flb_file_unpreserve_one(iter->flb);
}
void luo_flb_file_finish(struct liveupdate_file_handler *fh)
{
struct list_head *flb_list = &ACCESS_PRIVATE(fh, flb_list);
struct luo_flb_link *iter;
guard(rwsem_read)(&luo_register_rwlock);
list_for_each_entry_reverse(iter, flb_list, list)
liveupdate_flb_put_incoming(iter->flb);
}
static void luo_flb_unregister_one(struct liveupdate_file_handler *fh,
struct liveupdate_flb *flb)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
struct list_head *flb_list = &ACCESS_PRIVATE(fh, flb_list);
struct luo_flb_link *iter;
bool found = false;
list_for_each_entry(iter, flb_list, list) {
if (iter->flb == flb) {
list_del(&iter->list);
kfree(iter);
found = true;
break;
}
}
if (!found) {
pr_warn("Failed to unregister FLB '%s': not found in file handler '%s'\n",
flb->compatible, fh->compatible);
return;
}
private->users--;
if (!private->users) {
list_del_init(&private->list);
luo_flb_global.count--;
}
}
void luo_flb_unregister_all(struct liveupdate_file_handler *fh)
{
struct list_head *flb_list = &ACCESS_PRIVATE(fh, flb_list);
struct luo_flb_link *iter, *tmp;
if (!liveupdate_enabled())
return;
lockdep_assert_held_write(&luo_register_rwlock);
list_for_each_entry_safe(iter, tmp, flb_list, list)
luo_flb_unregister_one(fh, iter->flb);
}
int liveupdate_register_flb(struct liveupdate_file_handler *fh,
struct liveupdate_flb *flb)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
struct list_head *flb_list = &ACCESS_PRIVATE(fh, flb_list);
struct luo_flb_link *link __free(kfree) = NULL;
struct liveupdate_flb *gflb;
struct luo_flb_link *iter;
if (!liveupdate_enabled())
return -EOPNOTSUPP;
if (WARN_ON(!flb->ops->preserve || !flb->ops->unpreserve ||
!flb->ops->retrieve || !flb->ops->finish)) {
return -EINVAL;
}
if (WARN_ON(list_empty(&ACCESS_PRIVATE(fh, list))))
return -EINVAL;
link = kzalloc_obj(*link);
if (!link)
return -ENOMEM;
guard(rwsem_write)(&luo_register_rwlock);
list_for_each_entry(iter, flb_list, list) {
if (iter->flb == flb)
return -EEXIST;
}
if (!private->users) {
if (WARN_ON(!list_empty(&private->list)))
return -EINVAL;
if (luo_flb_global.count == LUO_FLB_MAX)
return -ENOSPC;
list_private_for_each_entry(gflb, &luo_flb_global.list, private.list) {
if (!strcmp(gflb->compatible, flb->compatible))
return -EEXIST;
}
list_add_tail(&private->list, &luo_flb_global.list);
luo_flb_global.count++;
}
private->users++;
link->flb = flb;
list_add_tail(&no_free_ptr(link)->list, flb_list);
return 0;
}
void liveupdate_unregister_flb(struct liveupdate_file_handler *fh,
struct liveupdate_flb *flb)
{
if (!liveupdate_enabled())
return;
guard(rwsem_write)(&luo_register_rwlock);
luo_flb_unregister_one(fh, flb);
}
int liveupdate_flb_get_incoming(struct liveupdate_flb *flb, void **objp)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
if (!liveupdate_enabled())
return -EOPNOTSUPP;
guard(mutex)(&private->incoming.lock);
if (!private->incoming.obj) {
int err = luo_flb_retrieve_one(flb);
if (err)
return err;
}
refcount_inc(&private->incoming.count);
*objp = private->incoming.obj;
return 0;
}
int liveupdate_flb_get_outgoing(struct liveupdate_flb *flb, void **objp)
{
struct luo_flb_private *private = luo_flb_get_private(flb);
if (!liveupdate_enabled())
return -EOPNOTSUPP;
guard(mutex)(&private->outgoing.lock);
*objp = private->outgoing.obj;
return 0;
}
int __init luo_flb_setup_outgoing(u64 *flbs_pa)
{
struct luo_flb_header_ser *header_ser;
header_ser = kho_alloc_preserve(LUO_FLB_PGCNT << PAGE_SHIFT);
if (IS_ERR(header_ser))
return PTR_ERR(header_ser);
*flbs_pa = virt_to_phys(header_ser);
header_ser->pgcnt = LUO_FLB_PGCNT;
luo_flb_global.outgoing.header_ser = header_ser;
luo_flb_global.outgoing.ser = (void *)(header_ser + 1);
luo_flb_global.outgoing.active = true;
return 0;
}
void __init luo_flb_setup_incoming(u64 flbs_pa)
{
struct luo_flb_header_ser *header_ser;
if (!flbs_pa)
return;
header_ser = phys_to_virt(flbs_pa);
luo_flb_global.incoming.header_ser = header_ser;
luo_flb_global.incoming.ser = (void *)(header_ser + 1);
luo_flb_global.incoming.active = true;
}
void luo_flb_serialize(void)
{
struct luo_flb_header *fh = &luo_flb_global.outgoing;
struct liveupdate_flb *gflb;
int i = 0;
guard(rwsem_read)(&luo_register_rwlock);
list_private_for_each_entry(gflb, &luo_flb_global.list, private.list) {
struct luo_flb_private *private = luo_flb_get_private(gflb);
long count = refcount_read(&private->outgoing.count);
if (count > 0) {
strscpy(fh->ser[i].name, gflb->compatible,
sizeof(fh->ser[i].name));
fh->ser[i].data = private->outgoing.data;
fh->ser[i].count = count;
i++;
}
}
fh->header_ser->count = i;
}