#include <linux/kernel.h>
#include <linux/errno.h>
#include <linux/sched/signal.h>
#include <linux/io_uring.h>
#include <linux/indirect_call_wrapper.h>
#include "io_uring.h"
#include "tctx.h"
#include "poll.h"
#include "rw.h"
#include "eventfd.h"
#include "wait.h"
#include "mpscq.h"
static void ctx_flush_and_put(struct io_ring_ctx *ctx, io_tw_token_t tw)
{
if (!ctx)
return;
if (ctx->flags & IORING_SETUP_TASKRUN_FLAG)
atomic_andnot(IORING_SQ_TASKRUN, &ctx->rings->sq_flags);
io_submit_flush_completions(ctx);
mutex_unlock(&ctx->uring_lock);
percpu_ref_put(&ctx->refs);
}
void io_tctx_fallback_work(struct work_struct *work)
{
struct io_uring_task *tctx = container_of(work, struct io_uring_task,
fallback_work);
unsigned int count = 0;
tctx_task_work_run(tctx, UINT_MAX, &count);
put_task_struct(tctx->task);
}
static void io_fallback_tw(struct io_uring_task *tctx)
{
get_task_struct(tctx->task);
if (!queue_work(system_dfl_wq, &tctx->fallback_work))
put_task_struct(tctx->task);
}
void tctx_task_work_run(struct io_uring_task *tctx, unsigned int max_entries,
unsigned int *count)
{
struct io_ring_ctx *ctx = NULL;
struct io_tw_state ts = { };
while (*count < max_entries) {
struct llist_node *node = mpscq_pop(&tctx->task_list,
&tctx->task_head);
struct io_kiocb *req;
if (!node) {
if (mpscq_empty(&tctx->task_list))
break;
ctx_flush_and_put(ctx, ts);
ctx = NULL;
cond_resched();
continue;
}
req = container_of(node, struct io_kiocb, io_task_work.node);
if (req->ctx != ctx) {
ctx_flush_and_put(ctx, ts);
ctx = req->ctx;
mutex_lock(&ctx->uring_lock);
percpu_ref_get(&ctx->refs);
ts.cancel = io_should_terminate_tw(ctx);
}
INDIRECT_CALL_2(req->io_task_work.func,
io_poll_task_func, io_req_rw_complete,
(struct io_tw_req){req}, ts);
(*count)++;
if (mpscq_pop_emptied(&tctx->task_list, tctx->task_head))
break;
if (unlikely(need_resched())) {
ctx_flush_and_put(ctx, ts);
ctx = NULL;
cond_resched();
}
}
ctx_flush_and_put(ctx, ts);
if (unlikely(atomic_read(&tctx->in_cancel)) &&
current->io_uring == tctx)
io_uring_drop_tctx_refs(current);
trace_io_uring_task_work_run(tctx, *count);
}
void tctx_task_work(struct callback_head *cb)
{
struct io_uring_task *tctx;
unsigned int count = 0;
tctx = container_of(cb, struct io_uring_task, task_work);
tctx_task_work_run(tctx, UINT_MAX, &count);
}
static void io_ctx_mark_taskrun(struct io_ring_ctx *ctx)
{
lockdep_assert_in_rcu_read_lock();
if (ctx->flags & IORING_SETUP_TASKRUN_FLAG) {
struct io_rings *rings = rcu_dereference(ctx->rings_rcu);
atomic_or(IORING_SQ_TASKRUN, &rings->sq_flags);
}
}
void io_req_local_work_add(struct io_kiocb *req, unsigned flags)
{
struct io_ring_ctx *ctx = req->ctx;
int nr_wait;
guard(rcu)();
if (req->flags & IO_REQ_LINK_FLAGS)
flags &= ~IOU_F_TWQ_LAZY_WAKE;
if (mpscq_push(&ctx->work_list, &req->io_task_work.node)) {
io_ctx_mark_taskrun(ctx);
if (data_race(ctx->int_flags) & IO_RING_F_HAS_EVFD)
io_eventfd_signal(ctx, false);
}
nr_wait = atomic_read(&ctx->cq_wait_nr);
if (nr_wait <= 0)
return;
if (flags & IOU_F_TWQ_LAZY_WAKE) {
if (!atomic_dec_and_test(&ctx->cq_wait_nr))
return;
} else if (atomic_xchg(&ctx->cq_wait_nr, IO_CQ_WAKE_INIT) <= 0) {
return;
}
wake_up_state(ctx->submitter_task, TASK_INTERRUPTIBLE);
}
void io_req_normal_work_add(struct io_kiocb *req)
{
struct io_uring_task *tctx = req->tctx;
struct io_ring_ctx *ctx = req->ctx;
if (!mpscq_push(&tctx->task_list, &req->io_task_work.node))
return;
if (ctx->flags & IORING_SETUP_TASKRUN_FLAG)
atomic_or(IORING_SQ_TASKRUN, &ctx->rings->sq_flags);
if (ctx->flags & IORING_SETUP_SQPOLL) {
__set_notify_signal(tctx->task);
return;
}
if (likely(!task_work_add(tctx->task, &tctx->task_work, ctx->notify_method)))
return;
io_fallback_tw(tctx);
}
void io_req_task_work_add_remote(struct io_kiocb *req, unsigned flags)
{
if (WARN_ON_ONCE(!(req->ctx->flags & IORING_SETUP_DEFER_TASKRUN)))
return;
__io_req_task_work_add(req, flags);
}
void __cold io_cancel_local_task_work(struct io_ring_ctx *ctx)
{
struct io_tw_state ts = { .cancel = true };
struct llist_node *node;
guard(mutex)(&ctx->uring_lock);
while (!mpscq_empty(&ctx->work_list)) {
struct io_kiocb *req;
node = mpscq_pop(&ctx->work_list, &ctx->work_head);
if (!node) {
cond_resched();
continue;
}
req = container_of(node, struct io_kiocb, io_task_work.node);
req->io_task_work.func((struct io_tw_req){req}, ts);
}
io_submit_flush_completions(ctx);
}
static bool io_run_local_work_continue(struct io_ring_ctx *ctx, int events,
int min_events)
{
if (!io_local_work_pending(ctx))
return false;
if (events < min_events)
return true;
if (ctx->flags & IORING_SETUP_TASKRUN_FLAG)
atomic_or(IORING_SQ_TASKRUN, &ctx->rings->sq_flags);
return false;
}
static int __io_run_local_work_loop(struct io_ring_ctx *ctx,
io_tw_token_t tw,
int events)
{
int ret = 0;
while (ret < events) {
struct llist_node *node = mpscq_pop(&ctx->work_list, &ctx->work_head);
struct io_kiocb *req;
if (!node)
break;
req = container_of(node, struct io_kiocb, io_task_work.node);
INDIRECT_CALL_2(req->io_task_work.func,
io_poll_task_func, io_req_rw_complete,
(struct io_tw_req){req}, tw);
ret++;
}
return ret;
}
static int __io_run_local_work(struct io_ring_ctx *ctx, io_tw_token_t tw,
int min_events, int max_events)
{
unsigned int loops = 0;
int ret = 0;
if (WARN_ON_ONCE(ctx->submitter_task != current))
return -EEXIST;
if (ctx->flags & IORING_SETUP_TASKRUN_FLAG)
atomic_andnot(IORING_SQ_TASKRUN, &ctx->rings->sq_flags);
again:
if (unlikely(loops && !ret))
cond_resched();
tw.cancel = io_should_terminate_tw(ctx);
min_events -= ret;
ret = __io_run_local_work_loop(ctx, tw, max_events);
loops++;
if (io_run_local_work_continue(ctx, ret, min_events))
goto again;
io_submit_flush_completions(ctx);
if (io_run_local_work_continue(ctx, ret, min_events))
goto again;
trace_io_uring_local_work_run(ctx, ret, loops);
return ret;
}
int io_run_local_work_locked(struct io_ring_ctx *ctx, int min_events)
{
struct io_tw_state ts = {};
if (!io_local_work_pending(ctx))
return 0;
return __io_run_local_work(ctx, ts, min_events,
max(IO_LOCAL_TW_DEFAULT_MAX, min_events));
}
int io_run_local_work(struct io_ring_ctx *ctx, int min_events, int max_events)
{
struct io_tw_state ts = {};
int ret;
mutex_lock(&ctx->uring_lock);
ret = __io_run_local_work(ctx, ts, min_events, max_events);
mutex_unlock(&ctx->uring_lock);
return ret;
}