diff options
Diffstat (limited to 'lua/clever_tee/host_adapter.lua')
| -rw-r--r-- | lua/clever_tee/host_adapter.lua | 1165 |
1 files changed, 1165 insertions, 0 deletions
diff --git a/lua/clever_tee/host_adapter.lua b/lua/clever_tee/host_adapter.lua new file mode 100644 index 0000000..c746f60 --- /dev/null +++ b/lua/clever_tee/host_adapter.lua @@ -0,0 +1,1165 @@ +local capabilities = require("clever_tee.capabilities") +local domain = require("clever_tee.domain") + +local M = {} +local unpack_values = table.unpack or unpack +local HostAdapter = {} +HostAdapter.__index = HostAdapter +M.HostAdapter = HostAdapter + +M.ActionEffect = { + NONE = "none", + ESCAPE = "escape", + ERROR = "error", +} +M.CONFIGURATION_PREFIX = "clever_tee_" +M.CONFIGURATION_GLOBALS = { + suppress_default_mappings = "clever_tee_not_overwrites_standard_mappings", +} + +local BOOLEAN_CONFIGURATION = { + search_current_line_only = true, + ignore_case = true, + smart_case = true, + use_migemo = true, + fix_key_direction = true, + show_prompt = true, + mark_cursor = true, + hide_cursor_on_cmdline = true, + mark_char = true, + mark_direct = true, + clean_labels_eagerly = true, +} + +local adapter_records = setmetatable({}, { __mode = "k" }) + +local function fail(message, level) + error(message, (level or 1) + 1) +end + +local function current_runtime(options) + if type(options) == "table" then + if options.runtime ~= nil then + return options.runtime + end + if options.api ~= nil then + return options + end + return rawget(_G, "vim") + end + if options ~= nil and options ~= HostAdapter then + return options + end + return rawget(_G, "vim") +end + +local function require_runtime(runtime) + if type(runtime) ~= "table" or type(runtime.api) ~= "table" then + fail("HostAdapter requires the Nvim Lua runtime", 3) + end + return runtime +end + +function HostAdapter.new(options) + if HostAdapter.is(options) then + return options + end + local adapter = setmetatable({}, HostAdapter) + local record = { + runtime = require_runtime(current_runtime(options)), + next_identity = 1, + highlights = {}, + timers = {}, + cursor_presentations = {}, + events = {}, + actions = {}, + mappings = {}, + dot_repeat = nil, + dot_bridge = nil, + event_order = {}, + action_diagnostics = nil, + augroup = nil, + } + record.event_queue = capabilities.EventQueue.new(function(_, delivery) + delivery.callback(delivery.name, delivery.payload) + end) + adapter_records[adapter] = record + return adapter +end + +function HostAdapter.is(value) + return type(value) == "table" and adapter_records[value] ~= nil +end + +function HostAdapter:runtime() + return adapter_records[self].runtime +end + +function HostAdapter:read_text() + local lines = self:runtime().api.nvim_buf_get_lines(0, 0, -1, true) + if #lines == 0 then + lines = { "" } + end + return domain.TextSnapshot.new(lines) +end + +function HostAdapter:read_buffer() + return self:runtime().api.nvim_get_current_buf() +end + +function HostAdapter:read_window() + return self:runtime().api.nvim_get_current_win() +end + +function HostAdapter:read_cursor() + local position = self:runtime().api.nvim_win_get_cursor(0) + return domain.Position.new(position[1], position[2] + 1) +end + +function HostAdapter:read_mode() + return self:runtime().api.nvim_get_mode().mode +end + +local function selection_option(runtime) + local value = runtime.api.nvim_get_option_value( + "selection", + { scope = "global" } + ) + if value == "exclusive" then + return domain.SelectionOption.EXCLUSIVE + end + return domain.SelectionOption.INCLUSIVE +end + +function HostAdapter:read_selection() + local runtime = self:runtime() + local context = domain.ModeContext.from_full_mode(self:read_mode()) + local kind = context.visual_kind or context.select_kind + local option = selection_option(runtime) + if kind == nil then + return domain.Selection.inactive(option) + end + local raw_anchor = runtime.fn.getpos("v") + local focus = self:read_cursor() + local anchor + if type(raw_anchor) == "table" + and type(raw_anchor[2]) == "number" + and raw_anchor[2] > 0 + and type(raw_anchor[3]) == "number" + and raw_anchor[3] > 0 + then + anchor = domain.Position.new(raw_anchor[2], raw_anchor[3]) + else + anchor = focus + end + return domain.Selection.active(kind, anchor, focus, option) +end + +function HostAdapter:read_count() + local runtime = self:runtime() + local count = runtime.v.count1 + if type(count) ~= "number" or count < 1 then + count = 1 + end + return domain.Count.new(count) +end + +local function configuration_global(name) + if type(name) ~= "string" or name == "" then + fail("configuration name must be a nonempty string", 3) + end + return M.CONFIGURATION_GLOBALS[name] or M.CONFIGURATION_PREFIX .. name +end + +function M.configuration_global(name) + return configuration_global(name) +end + +function HostAdapter:configuration_present(name) + local global = configuration_global(name) + return self:runtime().fn.exists("g:" .. global) == 1 +end + +local function normalize_configuration(name, value) + if BOOLEAN_CONFIGURATION[name] and type(value) == "number" then + return value ~= 0 + end + return value +end + +function HostAdapter:read_configuration(name) + local global = configuration_global(name) + return normalize_configuration(name, self:runtime().g[global]) +end + +function HostAdapter:write_configuration(name, value) + local global = configuration_global(name) + self:runtime().g[global] = value +end + +function HostAdapter:read_encoding() + local runtime = self:runtime() + return runtime.api.nvim_get_option_value( + "encoding", + { scope = "global" } + ) +end + +function HostAdapter:lowercase(value) + if type(value) ~= "string" then + fail("case conversion value must be a string", 2) + end + local result = self:runtime().fn.tolower(value) + if type(result) ~= "string" then + fail("Nvim case conversion must return a string", 2) + end + return result +end + +function HostAdapter:read_macro_state() + local register = self:runtime().fn.reg_executing() + return domain.MacroState.new(register ~= "" and register or nil) +end + +local function fold_open_policies(value) + local result = {} + for item in tostring(value):gmatch("[^,]+") do + if item == "hor" then + item = "horizontal" + end + result[#result + 1] = item + end + return result +end + +function HostAdapter:read_fold_state() + local runtime = self:runtime() + local foldopen = runtime.api.nvim_get_option_value( + "foldopen", + { scope = "global" } + ) + local line = self:read_cursor().line + local closed_levels = runtime.fn.foldclosed(line) == -1 and 0 or 1 + return domain.FoldState.new( + fold_open_policies(foldopen), + closed_levels + ) +end + +function HostAdapter:read_time_ms() + local runtime = self:runtime() + local uv = runtime.uv or runtime.loop + if type(uv) ~= "table" or type(uv.hrtime) ~= "function" then + fail("HostAdapter runtime must provide a monotonic clock", 2) + end + return uv.hrtime() / 1000000 +end + +function HostAdapter:read_pending_operator() + local operator = self:runtime().v.operator + if operator == nil then + return "" + end + return operator +end + +function HostAdapter:apply_cursor(position) + position = domain.Position.coerce(position) + self:runtime().api.nvim_win_set_cursor( + 0, + { position.line, position.byte_column - 1 } + ) +end + +function HostAdapter:apply_selection(position) + if domain.Selection.is(position) then + if not position.active then + fail("selection movement requires an active selection", 2) + end + position = position.focus + end + return self:apply_cursor(position) +end + +function HostAdapter:set_operator_inclusive(enabled) + if type(enabled) ~= "boolean" then + fail("operator inclusivity must be a Boolean", 2) + end + if not enabled then + return + end + + local runtime = self:runtime() + local api = runtime.api + local selection + if type(api.nvim_get_option_value) == "function" then + selection = api.nvim_get_option_value("selection", { scope = "global" }) + end + if selection ~= "exclusive" then + api.nvim_cmd({ + cmd = "normal", + bang = true, + args = { "v" }, + }, {}) + return + end + if type(api.nvim_set_option_value) ~= "function" + or type(runtime.schedule) ~= "function" + then + fail("HostAdapter cannot preserve exclusive selection during an operator", 2) + end + + api.nvim_set_option_value("selection", "inclusive", { scope = "global" }) + local ok, command_error = pcall(api.nvim_cmd, { + cmd = "normal", + bang = true, + args = { "v" }, + }, {}) + if not ok then + api.nvim_set_option_value("selection", selection, { scope = "global" }) + error(command_error, 0) + end + runtime.schedule(function() + api.nvim_set_option_value("selection", selection, { scope = "global" }) + end) +end + +local function string_bytes(value) + local bytes = {} + for index = 1, #value do + bytes[index] = string.byte(value, index) + end + return bytes +end + +function HostAdapter:read_input() + local runtime = self:runtime() + local value = runtime.fn.getcharstr() + if type(value) ~= "string" or value == "" then + fail("Nvim target input must be a nonempty string", 2) + end + local bytes = string_bytes(value) + if #bytes == 3 + and bytes[1] == 0x80 + and bytes[2] == 0xfd + and bytes[3] == 0x60 + then + return domain.InputPacket.raw_bytes(bytes) + end + if #bytes == 1 and bytes[1] == 27 then + return domain.InputPacket.special_key("Escape", bytes) + end + if bytes[1] == 0x80 then + local name = type(runtime.fn.keytrans) == "function" + and runtime.fn.keytrans(value) + or "Special" + return domain.InputPacket.special_key(name, value) + end + return domain.InputPacket.text(value) +end + +function HostAdapter:open_fold(position) + position = position and domain.Position.coerce(position) or self:read_cursor() + local runtime = self:runtime() + if runtime.fn.foldclosed(position.line) == -1 then + return false + end + runtime.api.nvim_cmd({ + cmd = "normal", + bang = true, + args = { "zo" }, + }, {}) + return true +end + +function HostAdapter:show_prompt(text) + if type(text) ~= "string" then + fail("prompt must be a string", 2) + end + self:runtime().api.nvim_echo({ { text } }, false, {}) +end + +function HostAdapter:redraw(kind) + if kind == "suppressed" then + return false + end + if kind ~= "screen" and kind ~= "full" then + fail("redraw kind must be screen, full, or suppressed", 2) + end + self:runtime().api.nvim_cmd({ + cmd = "redraw", + bang = kind == "full", + }, {}) + return true +end + +local DIAGNOSTIC_LEVELS = { + error = "ERROR", + warning = "WARN", + info = "INFO", +} + +function HostAdapter:emit_diagnostic(level, text) + local level_name = DIAGNOSTIC_LEVELS[level] + if level_name == nil then + fail("diagnostic level must be error, warning, or info", 2) + end + if type(text) ~= "string" or text == "" then + fail("diagnostic text must be a nonempty string", 2) + end + local record = adapter_records[self] + local runtime = record.runtime + if type(runtime.notify) ~= "function" then + fail("HostAdapter runtime must provide notify", 2) + end + local levels = type(runtime.log) == "table" and runtime.log.levels or {} + runtime.notify(text, levels[level_name], { title = "clever-tee" }) + if record.action_diagnostics ~= nil then + record.action_diagnostics[level .. "\0" .. text] = true + end +end + +local install_dot_bridge + +function HostAdapter:register_dot_repeat(payload, callback) + if not domain.DotPayload.is(payload) then + fail("dot-repeat payload must be a DotPayload", 2) + end + if callback ~= nil and type(callback) ~= "function" then + fail("dot-repeat callback must be a function", 2) + end + local record = adapter_records[self] + record.dot_repeat = { + payload = payload, + callback = callback, + operator = self:read_pending_operator(), + } + if install_dot_bridge ~= nil then + install_dot_bridge(self) + end + if record.dot_bridge ~= nil then + record.dot_bridge.awaiting_change = true + end + return payload +end + +function HostAdapter:dot_repeat_payload() + local registration = adapter_records[self].dot_repeat + return registration and registration.payload or nil +end + +function HostAdapter:replay_dot(count) + local registration = adapter_records[self].dot_repeat + if registration == nil or registration.callback == nil then + fail("dot repeat is not executable", 2) + end + return registration.callback( + registration.payload, + domain.Count.new(count) + ) +end + +local function next_identity(adapter, prefix) + local record = adapter_records[adapter] + local identity = prefix .. "-" .. tostring(record.next_identity) + record.next_identity = record.next_identity + 1 + return identity +end + +local function highlight_exists(runtime, name) + if type(runtime.fn) == "table" and type(runtime.fn.hlexists) == "function" then + return runtime.fn.hlexists(name) == 1 + end + local definition = runtime.api.nvim_get_hl(0, { + name = name, + link = true, + create = false, + }) + return next(definition) ~= nil +end + +function HostAdapter:read_highlight_group(name) + if type(name) ~= "string" or name == "" then + fail("highlight group name must be a nonempty string", 2) + end + local runtime = self:runtime() + if not highlight_exists(runtime, name) then + return nil + end + return runtime.api.nvim_get_hl(0, { + name = name, + link = true, + create = false, + }) +end + +local function native_highlight_definition(definition, options) + if type(definition) ~= "table" then + fail("highlight group definition must be a table", 3) + end + options = options or {} + if type(options) ~= "table" then + fail("highlight group options must be a table", 3) + end + local native = {} + for key, value in pairs(definition) do + if key ~= "guifg" and key ~= "guibg" and key ~= "gui" then + native[key] = value + end + end + if definition.guifg ~= nil then + native.fg = definition.guifg + end + if definition.guibg ~= nil then + native.bg = definition.guibg + end + for key, value in pairs(definition.gui or {}) do + native[key] = value + end + if options.default ~= nil then + native.default = options.default + end + if options.force ~= nil then + native.force = options.force + end + return native +end + +function HostAdapter:define_highlight_group(name, definition, options) + if type(name) ~= "string" or name == "" then + fail("highlight group name must be a nonempty string", 2) + end + options = options or {} + local runtime = self:runtime() + if options.default and highlight_exists(runtime, name) then + return false + end + runtime.api.nvim_set_hl( + 0, + name, + native_highlight_definition(definition, options) + ) + return true +end + +local function overlay_positions(specification) + local positions = specification.positions + if positions == nil and specification.position ~= nil then + positions = { specification.position } + end + if type(positions) ~= "table" then + fail("highlight positions must be a list", 3) + end + local native = {} + for index, position in ipairs(positions) do + position = domain.Position.coerce(position) + native[index] = { position.line, position.byte_column } + end + if #native == 0 then + native[1] = { 0 } + end + return native +end + +local function overlay_priority(value) + if value == "high" then + return 100 + end + if value == "ordinary" or value == nil then + return 10 + end + if type(value) == "number" then + return value + end + fail("highlight priority must be high, ordinary, or numeric", 3) +end + +function HostAdapter:create_highlight(specification) + if type(specification) ~= "table" then + fail("highlight specification must be a table", 2) + end + if type(specification.group) ~= "string" or specification.group == "" then + fail("highlight group must be a nonempty string", 2) + end + if specification.window == nil then + fail("highlight window must identify its Nvim window", 2) + end + local record = adapter_records[self] + local identity = specification.identity or next_identity(self, "highlight") + if record.highlights[identity] ~= nil then + fail("highlight identity is already active", 2) + end + local match_id = record.runtime.fn.matchaddpos( + specification.group, + overlay_positions(specification), + overlay_priority(specification.priority), + -1, + { window = specification.window } + ) + if type(match_id) ~= "number" or match_id < 0 then + fail("Nvim could not create the window-local highlight", 2) + end + record.highlights[identity] = { + match_id = match_id, + window = specification.window, + } + return identity +end + +function HostAdapter:remove_highlight(identity) + local record = adapter_records[self] + local resource = record.highlights[identity] + if resource == nil then + return false + end + record.highlights[identity] = nil + record.runtime.fn.matchdelete(resource.match_id, resource.window) + return true +end + +local function nonnegative_integer(value, name) + if type(value) ~= "number" + or value < 0 + or value ~= math.floor(value) + or value == math.huge + then + fail((name or "value") .. " must be a nonnegative integer", 3) + end + return value +end + +function HostAdapter:supports_timers() + local fn = self:runtime().fn + return type(fn) == "table" + and type(fn.timer_start) == "function" + and type(fn.timer_stop) == "function" +end + +function HostAdapter:start_timer(delay_ms, callback) + nonnegative_integer(delay_ms, "timer delay") + if type(callback) ~= "function" then + fail("timer callback must be a function", 2) + end + if not self:supports_timers() then + return nil + end + local record = adapter_records[self] + local identity = next_identity(self, "timer") + local timer_id = record.runtime.fn.timer_start(delay_ms, function() + local resource = record.timers[identity] + if resource == nil or not resource.active then + return + end + record.timers[identity] = nil + callback(identity) + end) + if type(timer_id) ~= "number" or timer_id < 0 then + fail("Nvim could not start the timer", 2) + end + record.timers[identity] = { + timer_id = timer_id, + active = true, + } + return identity +end + +function HostAdapter:stop_timer(identity) + local record = adapter_records[self] + local resource = record.timers[identity] + if resource == nil or not resource.active then + return false + end + record.timers[identity] = nil + record.runtime.fn.timer_stop(resource.timer_id) + return true +end + +local function event_names(value) + if type(value) == "string" then + value = { value } + end + if type(value) ~= "table" or #value == 0 then + fail("event names must be a nonempty list", 3) + end + local names = {} + local set = {} + for index, name in ipairs(value) do + if type(name) ~= "string" or name == "" then + fail("event name must be a nonempty string", 3) + end + if not set[name] then + names[#names + 1] = name + set[name] = true + end + end + return names, set +end + +local function event_payload(adapter, event) + local payload = { + buffer = event.buf, + file = event.file, + match = event.match, + data = event.data, + } + local api = adapter:runtime().api + if type(api.nvim_get_current_win) == "function" then + payload.window = api.nvim_get_current_win() + end + return payload +end + +local function event_augroup(record) + if record.augroup == nil then + record.augroup = record.runtime.api.nvim_create_augroup( + "clever_tee", + { clear = true } + ) + end + return record.augroup +end + +local function queue_event(record, name, payload, callback) + return record.event_queue:emit(name, { + name = name, + payload = payload, + callback = callback, + }) +end + +function HostAdapter:register_events(names, callback, options) + local name_set + names, name_set = event_names(names) + if type(callback) ~= "function" then + fail("event callback must be a function", 2) + end + options = options or {} + if type(options) ~= "table" then + fail("event registration options must be a table", 2) + end + local record = adapter_records[self] + local identity = next_identity(self, "event-registration") + local autocmd_options = { + group = event_augroup(record), + desc = "clever-tee " .. table.concat(names, "/"), + callback = function(event) + local resource = record.events[identity] + if resource ~= nil and resource.active then + queue_event( + record, + event.event, + event_payload(self, event), + resource.callback + ) + end + end, + } + if options.buffer ~= nil then + autocmd_options.buffer = options.buffer + end + local autocmd_id = record.runtime.api.nvim_create_autocmd( + names, + autocmd_options + ) + record.events[identity] = { + autocmd_id = autocmd_id, + names = names, + name_set = name_set, + callback = callback, + buffer = options.buffer, + active = true, + } + record.event_order[#record.event_order + 1] = identity + return identity +end + +function HostAdapter:remove_event_registration(identity) + local record = adapter_records[self] + local resource = record.events[identity] + if resource == nil or not resource.active then + return false + end + resource.active = false + record.runtime.api.nvim_del_autocmd(resource.autocmd_id) + return true +end + +local DOT_MOTION_MAPPING = "<Plug>(clever-tee-dot-motion)" + +local function dot_bridge_supported(runtime) + return type(runtime.keymap) == "table" + and type(runtime.keymap.set) == "function" + and type(runtime.keymap.del) == "function" + and type(runtime.fn.maparg) == "function" + and type(runtime.fn.mapset) == "function" + and type(runtime.api.nvim_feedkeys) == "function" +end + +local function restore_dot_mapping(record) + local bridge = record.dot_bridge + if bridge == nil or not bridge.active then + return false + end + bridge.active = false + pcall(record.runtime.keymap.del, "n", ".") + if type(bridge.previous_mapping) == "table" + and next(bridge.previous_mapping) ~= nil + then + record.runtime.fn.mapset("n", false, bridge.previous_mapping) + end + return true +end + +local function dot_replay_keys(runtime, count, operator) + local prefix = count > 0 and tostring(count) or "" + local keys = prefix .. operator .. DOT_MOTION_MAPPING + if type(runtime.keycode) == "function" then + return runtime.keycode(keys) + end + return runtime.api.nvim_replace_termcodes(keys, true, false, true) +end + +install_dot_bridge = function(adapter) + local record = adapter_records[adapter] + local runtime = record.runtime + if not dot_bridge_supported(runtime) then + return nil + end + local bridge = record.dot_bridge + if bridge == nil then + bridge = { + active = false, + awaiting_change = false, + previous_mapping = nil, + } + record.dot_bridge = bridge + runtime.keymap.set("o", DOT_MOTION_MAPPING, function() + local registration = record.dot_repeat + if registration == nil or registration.callback == nil then + return + end + local outcome = registration.callback( + registration.payload, + adapter:read_count() + ) + if domain.ActionOutcome.is(outcome) then + adapter:translate_action_outcome(outcome) + end + end, { + silent = true, + remap = false, + desc = "clever-tee dot motion", + }) + local has_cmd_atom = type(runtime.fn.exists) == "function" + and runtime.fn.exists("##CmdAtom") == 1 + local ownership_events = has_cmd_atom + and "CmdAtom" + or { "TextChanged", "TextChangedI", "TextChangedP" } + runtime.api.nvim_create_autocmd(ownership_events, { + group = event_augroup(record), + desc = "clever-tee dot ownership", + callback = function(event) + if has_cmd_atom and not (event.data and event.data.changed) then + return + end + if bridge.awaiting_change then + bridge.awaiting_change = false + return + end + restore_dot_mapping(record) + end, + }) + end + if not bridge.active then + bridge.previous_mapping = runtime.fn.maparg(".", "n", false, true) + runtime.keymap.set("n", ".", function() + local registration = record.dot_repeat + if registration == nil or registration.operator == "" then + restore_dot_mapping(record) + runtime.api.nvim_feedkeys(".", "n", false) + return + end + bridge.awaiting_change = true + local count = runtime.v.count or 0 + runtime.api.nvim_feedkeys( + dot_replay_keys(runtime, count, registration.operator), + "n", + false + ) + end, { + silent = true, + remap = false, + desc = "clever-tee dot repeat", + }) + bridge.active = true + end + return bridge +end + +function HostAdapter:deliver_event(name, payload) + if type(name) ~= "string" or name == "" then + fail("event name must be a nonempty string", 2) + end + payload = payload or {} + local record = adapter_records[self] + local event_buffer = payload.buffer + for _, identity in ipairs(record.event_order) do + local resource = record.events[identity] + if resource.active + and resource.name_set[name] + and (resource.buffer == nil + or event_buffer == nil + or resource.buffer == event_buffer) + then + queue_event(record, name, payload, resource.callback) + end + end +end + +function HostAdapter:begin_action_transition() + local record = adapter_records[self] + record.action_diagnostics = {} + return record.event_queue:begin_transition() +end + +function HostAdapter:commit_action_transition(token) + local record = adapter_records[self] + local result = record.event_queue:commit_transition(token) + record.action_diagnostics = nil + return result +end + +local function terminal_cursor_option(runtime) + return runtime.fn.eval("&t_ve") +end + +local function set_terminal_cursor_option(runtime, value) + runtime.api.nvim_cmd({ + cmd = "let", + args = { "&t_ve", "=", runtime.fn.string(value) }, + }, {}) +end + +function HostAdapter:supports_cursor_presentation() + local runtime = self:runtime() + return type(runtime.api.nvim_get_option_value) == "function" + and type(runtime.api.nvim_set_option_value) == "function" + and type(runtime.api.nvim_cmd) == "function" + and type(runtime.fn) == "table" + and type(runtime.fn.exists) == "function" + and runtime.fn.exists("+t_ve") == 1 + and type(runtime.fn.eval) == "function" + and type(runtime.fn.string) == "function" +end + +function HostAdapter:suppress_cursor_presentation() + if not self:supports_cursor_presentation() then + return nil + end + local record = adapter_records[self] + local runtime = record.runtime + local identity = next_identity(self, "cursor-presentation") + local saved = { + guicursor = runtime.api.nvim_get_option_value( + "guicursor", + { scope = "global" } + ), + terminal_cursor = terminal_cursor_option(runtime), + } + runtime.api.nvim_set_option_value( + "guicursor", + "a:ver1", + { scope = "global" } + ) + local ok, failure = pcall(set_terminal_cursor_option, runtime, "") + if not ok then + runtime.api.nvim_set_option_value( + "guicursor", + saved.guicursor, + { scope = "global" } + ) + error(failure, 0) + end + record.cursor_presentations[identity] = saved + return identity +end + +function HostAdapter:restore_cursor_presentation(identity) + local record = adapter_records[self] + local saved = record.cursor_presentations[identity] + if saved == nil then + return false + end + record.cursor_presentations[identity] = nil + record.runtime.api.nvim_set_option_value( + "guicursor", + saved.guicursor, + { scope = "global" } + ) + set_terminal_cursor_option(record.runtime, saved.terminal_cursor) + return true +end + +local function escape_key(runtime) + if type(runtime.keycode) == "function" then + return runtime.keycode("<Esc>") + end + if type(runtime.api.nvim_replace_termcodes) == "function" then + return runtime.api.nvim_replace_termcodes("<Esc>", true, false, true) + end + return string.char(27) +end + +function HostAdapter:return_escape() + local runtime = self:runtime() + if type(runtime.api.nvim_feedkeys) ~= "function" then + fail("HostAdapter runtime must provide nvim_feedkeys", 2) + end + runtime.api.nvim_feedkeys(escape_key(runtime), "n", false) +end + +function HostAdapter:emit_action_error(text) + return self:emit_diagnostic("error", text) +end + +function HostAdapter:translate_action_outcome(outcome) + if not domain.ActionOutcome.is(outcome) then + fail("host action translation requires an ActionOutcome", 2) + end + if outcome.kind == domain.ActionKind.ESCAPE then + self:return_escape() + return M.ActionEffect.ESCAPE + end + if outcome.kind == domain.ActionKind.ERROR then + local diagnostics = adapter_records[self].action_diagnostics + local key = "error\0" .. outcome.diagnostic + if diagnostics == nil or not diagnostics[key] then + self:emit_action_error(outcome.diagnostic) + end + return M.ActionEffect.ERROR + end + return M.ActionEffect.NONE +end + +local function packed(...) + return { n = select("#", ...), ... } +end + +local function invoke_callback(adapter, callback, ...) + local arguments = packed(...) + local token = adapter:begin_action_transition() + local results = packed(pcall(function() + local values = packed(callback(unpack_values(arguments, 1, arguments.n))) + if domain.ActionOutcome.is(values[1]) then + adapter:translate_action_outcome(values[1]) + end + return unpack_values(values, 1, values.n) + end)) + local commit = packed(pcall(adapter.commit_action_transition, adapter, token)) + if not results[1] then + error(results[2], 0) + end + if not commit[1] then + error(commit[2], 0) + end + return unpack_values(results, 2, results.n) +end + +function HostAdapter:register_action(name, callback) + if type(name) ~= "string" or name == "" then + fail("action name must be a nonempty string", 2) + end + if type(callback) ~= "function" then + fail("action callback must be a function", 2) + end + local actions = adapter_records[self].actions + if actions[name] ~= nil then + fail("action is already registered", 2) + end + actions[name] = callback + return name +end + +function HostAdapter:invoke_action(name, ...) + local callback = adapter_records[self].actions[name] + if callback == nil then + fail("action is not registered", 2) + end + return invoke_callback(self, callback, ...) +end + +function HostAdapter:invoke_callback(callback, ...) + if type(callback) ~= "function" then + fail("action callback must be a function", 2) + end + return invoke_callback(self, callback, ...) +end + +local function mapping_modes(value) + if type(value) == "string" then + value = { value } + end + if type(value) ~= "table" or #value == 0 then + fail("mapping modes must be a nonempty list", 3) + end + local result = {} + for index, mode in ipairs(value) do + if type(mode) ~= "string" or mode == "" then + fail("mapping mode must be a nonempty string", 3) + end + result[index] = mode + end + return result +end + +function HostAdapter:register_mapping(modes, lhs, action, options) + modes = mapping_modes(modes) + if type(lhs) ~= "string" or lhs == "" then + fail("mapping lhs must be a nonempty string", 2) + end + if type(action) ~= "string" and type(action) ~= "function" then + fail("mapping action must be an action name or function", 2) + end + options = options or {} + if type(options) ~= "table" then + fail("mapping options must be a table", 2) + end + local callback + if type(action) == "string" then + callback = function() + return self:invoke_action(action) + end + else + callback = function(...) + return invoke_callback(self, action, ...) + end + end + local native_options = { + silent = options.silent == true, + remap = options.remap == true, + desc = options.desc + or ("clever-tee " .. (type(action) == "string" and action or lhs)), + } + self:runtime().keymap.set(modes, lhs, callback, native_options) + local identity = next_identity(self, "mapping") + adapter_records[self].mappings[identity] = { + modes = modes, + lhs = lhs, + action = action, + options = options, + callback = callback, + } + return identity +end + +function M.new(options) + return HostAdapter.new(options) +end + +setmetatable(M, { + __call = function(_, options) + return HostAdapter.new(options) + end, +}) + +return M |
