local ffi = require("diffview.ffi")
local oop = require("diffview.oop")
local fmt = string.format
local uv = vim.loop
local DEFAULT_ERROR = "Unkown error."
local M = {}
M._watching = setmetatable({}, { __mode = "k" })
M._handles = {}
local function dstring(object)
if not DiffviewGlobal.logger then return "" end
dstring = DiffviewGlobal.logger.dstring
return dstring(object)
end
local function tbl_pack(...)
return { n = select("#", ...), ... }
end
local function tbl_unpack(t, i, j)
return unpack(t, i or 1, j or t.n or table.maxn(t))
end
local function current_thread()
local current, ismain = coroutine.running()
if type(ismain) == "boolean" then
return not ismain and current or nil
else
return current
end
end
local Waitable = oop.create_class("Waitable")
M.Waitable = Waitable
function Waitable:await() oop.abstract_stub() end
function Waitable:finally(callback)
(M.void(function()
local ret = tbl_pack(M.await(self))
callback(tbl_unpack(ret))
end))()
end
local Future = oop.create_class("Future", Waitable)
function Future:init(opt)
opt = opt or {}
if opt.thread then
self.thread = opt.thread
elseif opt.func then
self.thread = coroutine.create(opt.func)
else
error("Either 'thread' or 'func' must be specified!")
end
M._handles[self.thread] = self
self.listeners = {}
self.kind = opt.kind
self.started = false
self.awaiting_cb = false
self.done = false
self.has_raised = false
end
function Future:__tostring()
return dstring(self.thread)
end
function Future:destroy()
M._handles[self.thread] = nil
end
function Future:set_done(value)
self.done = value
if self:is_watching() then
self:dprint("done was set:", self.done)
end
end
function Future:is_done()
return not not self.done
end
function Future:get_returned()
if not self.return_values then return end
return unpack(self.return_values, 2, table.maxn(self.return_values))
end
function Future:dprint(...)
if not DiffviewGlobal.logger then return end
if DiffviewGlobal.debug_level >= 10 or M._watching[self] then
local t = { self, "::", ... }
for i = 1, table.maxn(t) do t[i] = dstring(t[i]) end
DiffviewGlobal.logger:debug(table.concat(t, " "))
end
end
function Future:dprintf(...)
self:dprint(fmt(...))
end
function Future:watch()
M._watching[self] = true
end
function Future:unwatch()
M._watching[self] = nil
end
function Future:is_watching()
return not not M._watching[self]
end
function Future:raise(force)
if self.has_raised and not force then return end
self.has_raised = true
error(self.err)
end
function Future:step(...)
self:dprint("step")
local ret = { coroutine.resume(self.thread, ...) }
local ok = ret[1]
if not ok then
local err = ret[2] or DEFAULT_ERROR
local func_info
if self.func then
func_info = debug.getinfo(self.func, "uS")
end
local msg = fmt(
"The coroutine failed with this message: \n"
.. "\tcontext: cur_thread=%s co_thread=%s %s\n%s",
dstring(current_thread() or "main"),
dstring(self.thread),
func_info and fmt("co_func=%s:%d", func_info.short_src, func_info.linedefined) or "",
debug.traceback(self.thread, err)
)
self:set_done(true)
self:notify_all(false, msg)
self:destroy()
self:raise()
return
end
if coroutine.status(self.thread) == "dead" then
self:dprint("handle dead")
self:set_done(true)
self:notify_all(true, unpack(ret, 2, table.maxn(ret)))
self:destroy()
return
end
end
function Future:notify_all(ok, ...)
local ret_values = tbl_pack(ok, ...)
if not ok then
self.err = ret_values[2] or DEFAULT_ERROR
end
local seen = {}
while next(self.listeners) do
local handle = table.remove(self.listeners, #self.listeners)
if handle and not seen[handle.thread] then
self:dprint("notifying:", handle)
seen[handle.thread] = true
handle:step(ret_values)
end
end
end
function Future:await()
if self.err then
self:raise(true)
return
end
if self:is_done() then
return self:get_returned()
end
local current = current_thread()
if not current then
return self:toplevel_await()
end
local parent_handle = M._handles[current]
if not parent_handle then
self:dprint("creating a wrapper around unmanaged thread")
self.parent = Future({
thread = current,
kind = "void",
})
else
self.parent = parent_handle
end
if current ~= self.thread then
table.insert(self.listeners, self.parent)
end
self:dprintf("awaiting: yielding=%s listeners=%s", dstring(current), dstring(self.listeners))
coroutine.yield()
local ok
if not self.return_values then
ok = self.err == nil
else
ok = self.return_values[1]
if not ok then
self.err = self.return_values[2] or DEFAULT_ERROR
end
end
if not ok then
self:raise(true)
return
end
return self:get_returned()
end
function Future:toplevel_await()
local ok, status
while true do
ok, status = vim.wait(1000 * 60, function()
return coroutine.status(self.thread) == "dead"
end, 1)
if status ~= -1 then break end
end
if not ok then
if status == -1 then
error("Async task timed out!")
elseif status == -2 then
error("Async task got interrupted!")
end
end
if self.err then
self:raise(true)
return
end
return self:get_returned()
end
function M._run(func, opt)
opt = opt or {}
local handle
local use_err_handler = not not current_thread()
local function wrapped_func(...)
if use_err_handler then
local ok = xpcall(func, function(err)
handle.err = debug.traceback(err, 2)
end, ...)
if not ok then
handle:dprint("an error was raised: terminating")
handle:set_done(true)
handle:destroy()
error(handle.err, 0)
return
end
else
func(...)
end
if opt.kind == "callback" and not handle:is_done() then
handle.awaiting_cb = true
handle:dprintf("yielding for cb: current=%s", dstring(current_thread()))
coroutine.yield()
handle:dprintf("resuming after cb: current=%s", dstring(current_thread()))
end
handle:set_done(true)
end
if opt.kind == "callback" then
local cur_cb = opt.args[opt.nparams]
local function wrapped_cb(...)
handle:set_done(true)
handle.return_values = { true, ... }
if cur_cb then cur_cb(...) end
if handle.awaiting_cb then
handle.awaiting_cb = false
handle:step()
end
handle:notify_all(true, ...)
end
opt.args[opt.nparams] = wrapped_cb
end
handle = Future({ func = wrapped_func, kind = opt.kind })
handle:dprint("created thread")
handle.func = func
handle.started = true
handle:step(tbl_unpack(opt.args))
return handle
end
function M.void(func)
return function(...)
return M._run(func, {
kind = "void",
args = { ... },
})
end
end
function M.wrap(func, nparams)
if not nparams then
local info = debug.getinfo(func, "uS")
assert(info.what == "Lua", "Parameter count can only be derived for Lua functions!")
nparams = info.nparams
end
return function(...)
return M._run(func, {
nparams = nparams,
kind = "callback",
args = { ... },
})
end
end
function M.await(waitable)
return waitable:await()
end
function M.pawait(x, ...)
local args = tbl_pack(...)
return pcall(function()
if type(x) == "function" then
return M.await(x(tbl_unpack(args)))
else
return x:await()
end
end)
end
local await = M.await
function M.sync_void(func)
local afunc = M.void(func)
return function(...)
return await(afunc(...))
end
end
function M.sync_wrap(func, nparams)
local afunc = M.wrap(func, nparams)
return function(...)
return await(afunc(...))
end
end
M.join = M.void(function(tasks)
local futures = {}
for _, cur in ipairs(tasks) do
if cur then
if type(cur) == "function" then
futures[#futures+1] = cur()
else
futures[#futures+1] = cur
end
end
end
for _, future in ipairs(futures) do
await(future)
end
end)
M.chain = M.void(function(tasks)
for _, task in ipairs(tasks) do
if type(task) == "function" then
await(task())
else
await(task)
end
end
end)
M.timeout = M.wrap(function(timeout, callback)
local timer = assert(uv.new_timer())
timer:start(
timeout,
0,
function()
if not timer:is_closing() then timer:close() end
callback()
end
)
end)
M.scheduler = M.wrap(function(fast_only, callback)
if (fast_only and not vim.in_fast_event()) or not ffi.nvim_is_locked() then
callback()
return
end
vim.schedule(callback)
end)
M.schedule_now = M.wrap(vim.schedule, 1)
return M