diff --git a/lua/nvim-treesitter/async.lua b/lua/nvim-treesitter/async.lua index dee60525e..0539470b5 100644 --- a/lua/nvim-treesitter/async.lua +++ b/lua/nvim-treesitter/async.lua @@ -16,14 +16,14 @@ local function unpack_len(t, first) end end ---- @class nvim-treesitter.async +--- @class async local M = {} --- Weak table to keep track of running tasks ---- @type table +--- @type table local threads = setmetatable({}, { __mode = 'k' }) ---- @return nvim-treesitter.async.Task? +--- @return async.Task? local function running() local task = threads[coroutine.running()] if task and not (task:_completed() or task._closing) then @@ -33,21 +33,22 @@ end --- Base class for async tasks. Async functions should return a subclass of --- this. This is designed specifically to be a base class of uv_handle_t ---- @class nvim-treesitter.async.Handle ---- @field close fun(self: nvim-treesitter.async.Handle, callback?: fun()) ---- @field is_closing? fun(self: nvim-treesitter.async.Handle): boolean +--- @class async.Handle +--- @field close fun(self: async.Handle, callback?: fun()) +--- @field is_closing? fun(self: async.Handle): boolean ---- @alias vim.async.CallbackFn ---- | fun(callback: fun(...: any)): nvim-treesitter.async.Handle? +--- @alias async.CallbackFn +--- | fun(callback: fun(...: any)): async.Handle? ---- @class nvim-treesitter.async.Task : nvim-treesitter.async.Handle ---- @field private _callbacks table +--- @class async.Task : async.Handle +--- @field package _callbacks table +--- @field package _callback_pos integer --- @field private _thread thread --- --- Tasks can call other async functions (task of callback functions) --- when we are waiting on a child, we store the handle to it here so we can --- cancel it. ---- @field private _current_child? nvim-treesitter.async.Handle +--- @field private _current_child? async.Handle --- --- Error result of the task is an error occurs. --- Must use `await` to get the result. @@ -61,7 +62,7 @@ Task.__index = Task --- @private --- @param func function ---- @return nvim-treesitter.async.Task +--- @return async.Task function Task._new(func) local thread = coroutine.create(func) @@ -69,6 +70,7 @@ function Task._new(func) _closing = false, _thread = thread, _callbacks = {}, + _callback_pos = 1, }, Task) threads[thread] = self @@ -84,7 +86,8 @@ function Task:await(callback) -- Already finished or closed callback(self._err, unpack_len(self._result)) else - table.insert(self._callbacks, callback) + self._callbacks[self._callback_pos] = callback + self._callback_pos = self._callback_pos + 1 end end @@ -168,7 +171,7 @@ function Task:_traceback(msg, _lvl) local child = self._current_child if getmetatable(child) == Task then - --- @cast child nvim-treesitter.async.Task + --- @cast child async.Task msg = child:_traceback(msg, _lvl + 1) end @@ -284,7 +287,7 @@ local function is_async_handle(obj) end function Task:_resume(...) - --- @type [boolean, string|vim.async.CallbackFn] + --- @type [boolean, string|async.CallbackFn] local ret = pack_len(coroutine.resume(self._thread, ...)) local stat = ret[1] @@ -303,7 +306,7 @@ function Task:_resume(...) local ok, r ok, r = pcall(fn, function(...) if is_async_handle(r) then - --- @cast r nvim-treesitter.async.Handle + --- @cast r async.Handle -- We must close children before we resume to ensure -- all resources are collected. local args = pack_len(...) @@ -344,14 +347,14 @@ end --- ``` --- @param func function --- @param ... any ---- @return nvim-treesitter.async.Task +--- @return async.Task function M.arun(func, ...) local task = Task._new(func) task:_resume(...) return task end ---- @class nvim-treesitter.async.TaskFun +--- @class async.TaskFun --- @field package _fun fun(...: any): any local TaskFun = {} TaskFun.__index = TaskFun @@ -365,43 +368,9 @@ function M.async(fun) return setmetatable({ _fun = fun }, TaskFun) end ---- Use this to create a function which executes in an async context but ---- called from a non-async context. ---- ---- The returned function will take the same arguments as the original function. ---- If argc is provided, the function will have an additional callback function ---- as the last argument which will be called when the function completes. ---- ---- @generic F: function ---- @param argc integer ---- @param func F ---- @return F -function M.create(argc, func) - assert(type(argc) == 'number') - assert(type(func) == 'function') - - --- @param ... any - --- @return any ... - return function(...) - local task = Task._new(func) - - task:raise_on_error() - - --- @type fun(err:string?, ...:any) - local callback = argc and select(argc + 1, ...) or nil - if callback and type(callback) == 'function' then - task:await(callback) - end - - task:_resume(unpack({ ... }, 1, argc)) - - return task - end -end - --- Returns the status of a task’s thread. --- ---- @param task? nvim-treesitter.async.Task +--- @param task? async.Task --- @return 'running'|'suspended'|'normal'|'dead'? function M.status(task) task = task or running() @@ -419,7 +388,7 @@ local function yield(fun) return coroutine.yield(fun) end ---- @param task nvim-treesitter.async.Task +--- @param task async.Task --- @return any ... local function await_task(task) local res = pack_len(yield(function(callback) @@ -439,7 +408,7 @@ end --- Asynchronous blocking wait --- @param argc integer ---- @param fun vim.async.CallbackFn +--- @param fun async.CallbackFn --- @param ... any func arguments --- @return any ... local function await_cbfun(argc, fun, ...) @@ -454,7 +423,7 @@ local function await_cbfun(argc, fun, ...) end) end ---- @param taskfun nvim-treesitter.async.TaskFun +--- @param taskfun async.TaskFun --- @param ... any --- @return any ... local function await_taskfun(taskfun, ...) @@ -498,9 +467,9 @@ end --- end --- end) --- ``` ---- @overload fun(argc: integer, func: vim.async.CallbackFn, ...:any): any ... ---- @overload fun(task: nvim-treesitter.async.Task): any ... ---- @overload fun(taskfun: nvim-treesitter.async.TaskFun): any ... +--- @overload fun(argc: integer, func: async.CallbackFn, ...:any): any ... +--- @overload fun(task: async.Task): any ... +--- @overload fun(taskfun: async.TaskFun): any ... function M.await(...) assert(running(), 'Not in async context') @@ -541,7 +510,7 @@ end --- --- local atimer = async.awrap( --- @param argc integer ---- @param func vim.async.CallbackFn +--- @param func async.CallbackFn --- @return async function function M.awrap(argc, func) assert(type(argc) == 'number') @@ -557,6 +526,34 @@ if vim.schedule then M.schedule = M.awrap(1, vim.schedule) end +--- Create a function that runs a function when it is garbage collected. +--- @generic F +--- @param f F +--- @param gc fun() +--- @return F +local function gc_fun(f, gc) + local proxy = newproxy(true) + local proxy_mt = getmetatable(proxy) + proxy_mt.__gc = gc + proxy_mt.__call = function(_, ...) + return f(...) + end + + return proxy +end + +--- @param task_cbs table +local function gc_cbs(task_cbs) + for task, tcb in pairs(task_cbs) do + for j, cb in pairs(task._callbacks) do + if cb == tcb then + task._callbacks[j] = nil + break + end + end + end +end + --- @async --- Example: --- ```lua @@ -586,7 +583,7 @@ end --- 3 'task3 error' nil nil --- ``` --- ---- @param tasks nvim-treesitter.async.Task[] +--- @param tasks async.Task[] --- @return fun(): (integer?, any?, ...) function M.iter(tasks) assert(running(), 'Not in async context') @@ -595,10 +592,22 @@ function M.iter(tasks) -- Iter blocks in an async context so only one waiter is needed local waiter = nil - + local task_cbs = {} --- @type table local remaining = #tasks + + --- If can_gc_cbs is true, then the iterator function has been garbage + --- collected and means any awaiters can also be garbage collected. The + --- only time we can't do this is if with the special case when iter() is + --- called anonymously (`local i = async.iter(tasks)()`), so we should not + --- garbage collect the callbacks until at least one awaiter is called. + local can_gc_cbs = false + for i, task in ipairs(tasks) do - task:await(function(err, ...) + local cb = function(err, ...) + if can_gc_cbs == true then + gc_cbs(task_cbs) + end + local callback = waiter -- Clear waiter before calling it @@ -612,21 +621,29 @@ function M.iter(tasks) -- Task finished before Iterator was called. Store results. table.insert(results, pack_len(i, err, ...)) end - end) + end + + task_cbs[task] = cb + task:await(cb) end - --- @param callback fun(i?: integer, err?: any, ...: any) - return M.awrap(1, function(callback) - if next(results) then - local res = table.remove(results, 1) - callback(unpack_len(res)) - elseif remaining == 0 then - callback() -- finish - else - assert(not waiter, 'internal error: waiter already set') - waiter = callback + return gc_fun( + M.awrap(1, function(callback) + if next(results) then + local res = table.remove(results, 1) + callback(unpack_len(res)) + elseif remaining == 0 then + callback() -- finish + else + assert(not waiter, 'internal error: waiter already set') + waiter = callback + end + end), + function() + -- Don't gc callbacks just yet. Wait until at least one of them is called. + can_gc_cbs = true end - end) + ) end do -- join() @@ -681,7 +698,7 @@ do -- join() --- [3] = { 'task2 error' }, --- } --- ``` - --- @param tasks nvim-treesitter.async.Task[] + --- @param tasks async.Task[] --- @return table function M.join(tasks) assert(running(), 'Not in async context') diff --git a/lua/nvim-treesitter/install.lua b/lua/nvim-treesitter/install.lua index 5df85f7b1..c70fc33f8 100644 --- a/lua/nvim-treesitter/install.lua +++ b/lua/nvim-treesitter/install.lua @@ -28,19 +28,27 @@ local INSTALL_TIMEOUT = 60000 --- @async --- @param max_jobs integer ---- @param tasks nvim-treesitter.async.TaskFun -local function join(max_jobs, tasks) - while #tasks > 0 do - local tasks_to_run = {} --- @type nvim-treesitter.async.TaskFun[] - for i = 1, max_jobs do - if #tasks == 0 then - break - end - local task = table.remove(tasks, 1) - tasks_to_run[i] = task() - end - a.join(tasks_to_run) +--- @param task_funs async.TaskFun[] +local function join(max_jobs, task_funs) + task_funs = vim.deepcopy(task_funs) + max_jobs = math.min(max_jobs, #task_funs) + + local running = {} --- @type async.TaskFun[] + + -- Start the first batch of tasks + for i = 1, max_jobs do + running[i] = table.remove(task_funs, 1)() end + + -- As tasks finish, add new ones + for _, task in ipairs(task_funs) do + local finished = a.iter(running)() + table.remove(running, finished) + table.insert(running, task()) + end + + -- Wait for all tasks to finish + a.join(running) end ---@async @@ -429,10 +437,10 @@ local function install(languages, options) local cache_dir = fs.normalize(fn.stdpath('cache')) local install_dir = config.get_install_dir('parser') - local tasks = {} ---@type nvim-treesitter.async.TaskFun[] + local task_funs = {} ---@type async.TaskFun[] local done = 0 for _, lang in ipairs(languages) do - tasks[#tasks + 1] = a.async(function() + task_funs[#task_funs + 1] = a.async(function() a.schedule() local status = install_lang(lang, cache_dir, install_dir, options.force, options.generate) if status ~= 'failed' then @@ -440,19 +448,19 @@ local function install(languages, options) end end) end - local numtasks = #tasks - join(options and options.max_jobs or MAX_JOBS, tasks) - if numtasks > 1 then + join(options and options.max_jobs or MAX_JOBS, task_funs) + if #task_funs > 1 then a.schedule() - log.info('Installed %d/%d languages', done, numtasks) + log.info('Installed %d/%d languages', done, #task_funs) end - return done == numtasks + return done == #task_funs end ---@param languages string[]|string ---@param options? InstallOptions -M.install = a.create(2, function(languages, options) +---@param _callback? fun(success: boolean) +M.install = a.async(function(languages, options, _callback) reload_parsers() languages = config.norm_languages(languages, { unsupported = true }) return install(languages, options) @@ -460,7 +468,7 @@ end) ---@param languages? string[]|string ---@param _options? table -M.update = a.create(2, function(languages, _options) +M.update = a.async(function(languages, _options) reload_parsers() if not languages or #languages == 0 then languages = 'all' @@ -510,14 +518,14 @@ local function uninstall_lang(logger, lang, parser, queries) end ---@param languages string[]|string -M.uninstall = a.create(1, function(languages) +M.uninstall = a.async(function(languages) languages = config.norm_languages(languages or 'all', { missing = true, dependencies = true }) local parser_dir = config.get_install_dir('parser') local query_dir = config.get_install_dir('queries') local installed = config.installed_parsers() - local tasks = {} ---@type nvim-treesitter.async.TaskFun[] + local task_funs = {} ---@type async.TaskFun[] local done = 0 for _, lang in ipairs(languages) do local logger = log.new('uninstall/' .. lang) @@ -526,7 +534,7 @@ M.uninstall = a.create(1, function(languages) else local parser = fs.joinpath(parser_dir, lang) .. '.so' local queries = fs.joinpath(query_dir, lang) - tasks[#tasks + 1] = a.async(function() + task_funs[#task_funs + 1] = a.async(function() local err = uninstall_lang(logger, lang, parser, queries) if not err then done = done + 1 @@ -535,10 +543,10 @@ M.uninstall = a.create(1, function(languages) end end - join(MAX_JOBS, tasks) - if #tasks > 1 then + join(MAX_JOBS, task_funs) + if #task_funs > 1 then a.schedule() - log.info('Uninstalled %d/%d languages', done, #tasks) + log.info('Uninstalled %d/%d languages', done, #task_funs) end end) diff --git a/scripts/install-parsers.lua b/scripts/install-parsers.lua index 991e53373..c551d2de1 100755 --- a/scripts/install-parsers.lua +++ b/scripts/install-parsers.lua @@ -21,11 +21,25 @@ vim.opt.runtimepath:append('.') -- needed on CI vim.fn.mkdir(vim.fn.stdpath('cache'), 'p') -local ok = update and require('nvim-treesitter.install').update('all'):wait(6000000) - or require('nvim-treesitter.install') - .install(#parsers > 0 and parsers or 'all', { force = true, generate = generate, max_jobs = max_jobs }) - :wait(6000000) +local nvts_install = require('nvim-treesitter.install') +local task = nil +if update then + --- @type async.Task + task = nvts_install.update('all') +else + --- @type async.Task + task = nvts_install.install( + #parsers > 0 and parsers or 'all', + { force = true, generate = generate, max_jobs = max_jobs } + ) +end + +local ok, err_or_ok = task:pwait(6000000) if not ok then + print('ERROR: ', err_or_ok) + vim.cmd.cq() +elseif not err_or_ok then + print('Did not install all parsers') vim.cmd.cq() end