summaryrefslogtreecommitdiff
path: root/lua
diff options
context:
space:
mode:
authorBogdan Nikolov <84232456+0xr3ngar@users.noreply.github.com>2026-02-19 20:30:08 +0100
committerGitHub <noreply@github.com>2026-02-19 20:30:08 +0100
commit725bf38ffc08cb8b537cbff322ef2d889b197dad (patch)
tree2fa53c6b40463b8a18d9e175617b9b413a844e70 /lua
parent6d219f47aa29cafb1ec39c41d4dd0a4875a74506 (diff)
parentd34e4dcbcc1f95d8f31f851baad18da12c59e7df (diff)
downloada4-725bf38ffc08cb8b537cbff322ef2d889b197dad.tar.xz
a4-725bf38ffc08cb8b537cbff322ef2d889b197dad.zip
Merge branch 'master' into master
Diffstat (limited to 'lua')
-rw-r--r--lua/99/editor/lsp.lua5
-rw-r--r--lua/99/extensions/agents/init.lua27
-rw-r--r--lua/99/extensions/init.lua28
-rw-r--r--lua/99/extensions/work/worker.lua75
-rw-r--r--lua/99/init.lua271
-rw-r--r--lua/99/logger/logger.lua8
-rw-r--r--lua/99/ops/clean-up.lua52
-rw-r--r--lua/99/ops/implement-fn.lua92
-rw-r--r--lua/99/ops/init.lua2
-rw-r--r--lua/99/ops/make-prompt.lua32
-rw-r--r--lua/99/ops/over-range.lua36
-rw-r--r--lua/99/ops/search.lua82
-rw-r--r--lua/99/ops/tutorial.lua72
-rw-r--r--lua/99/prompt-settings.lua12
-rw-r--r--lua/99/providers.lua59
-rw-r--r--lua/99/request-context.lua85
-rw-r--r--lua/99/request/init.lua64
-rw-r--r--lua/99/test/marks_spec.lua2
-rw-r--r--lua/99/test/providers_spec.lua33
-rw-r--r--lua/99/test/request_spec.lua54
-rw-r--r--lua/99/test/test_utils.lua40
-rw-r--r--lua/99/test/visual_spec.lua26
-rw-r--r--lua/99/utils.lua9
-rw-r--r--lua/99/window/init.lua73
24 files changed, 855 insertions, 384 deletions
diff --git a/lua/99/editor/lsp.lua b/lua/99/editor/lsp.lua
index 8ef0751..85e17de 100644
--- a/lua/99/editor/lsp.lua
+++ b/lua/99/editor/lsp.lua
@@ -1041,11 +1041,6 @@ function Lsp._format_exports(
return table.concat(out, "\n")
end
-local imports = ts.imports(0)
-Lsp.stringify_definition_exports_from_node(0, imports[1], function(s)
- print("s", s.results)
-end)
-
return {
Lsp = Lsp,
}
diff --git a/lua/99/extensions/agents/init.lua b/lua/99/extensions/agents/init.lua
index 559f34b..b1fda0f 100644
--- a/lua/99/extensions/agents/init.lua
+++ b/lua/99/extensions/agents/init.lua
@@ -1,5 +1,4 @@
local helpers = require("99.extensions.agents.helpers")
-local Logger = require("99.logger.logger")
local M = {}
--- @class _99.Agents.Rule
@@ -110,6 +109,23 @@ function M.by_name(rules, prompt)
}
end
+--- @param rule _99.Agents.Rule
+--- @return string | nil
+function M.get_rule_content(rule)
+ local file_path = rule.absolute_path or rule.path
+ local ok, file = pcall(io.open, file_path, "r")
+ if not ok or not file then
+ return nil
+ end
+ local ok_read, content = pcall(file.read, file, "*a")
+ if not ok_read then
+ return nil
+ end
+
+ pcall(file.close, file)
+ return string.format("<%s>\n%s\n</%s>", rule.name, content, rule.name)
+end
+
--- @param _99 _99.State
--- @return _99.CompletionProvider
function M.completion_provider(_99)
@@ -140,14 +156,7 @@ function M.completion_provider(_99)
if not rule then
return nil
end
- local file_path = rule.absolute_path or rule.path
- local ok, file = pcall(io.open, file_path, "r")
- if not ok or not file then
- return nil
- end
- local content = file:read("*a")
- file:close()
- return string.format("<%s>\n%s\n</%s>", rule.name, content, rule.name)
+ return M.get_rule_content(rule)
end,
}
end
diff --git a/lua/99/extensions/init.lua b/lua/99/extensions/init.lua
index 95629a2..6aae2bb 100644
--- a/lua/99/extensions/init.lua
+++ b/lua/99/extensions/init.lua
@@ -13,7 +13,33 @@ local function get_source(completion)
end
local source = completion.source
if source == "cmp" then
- local cmp = require("99.extensions.cmp")
+ local ok, cmp = pcall(require, "99.extensions.cmp")
+ if not ok then
+ vim.notify(
+ '[99] nvim-cmp is not installed. Install hrsh7th/nvim-cmp or use source = "blink"',
+ vim.log.levels.WARN
+ )
+ return
+ end
+ return cmp
+ end
+ if source == "blink" then
+ local ok, _ = pcall(require, "blink.compat")
+ if not ok then
+ vim.notify(
+ "[99] blink.compat is required for blink source. Install: { 'saghen/blink.compat', version = '2.*' }",
+ vim.log.levels.ERROR
+ )
+ return
+ end
+ local cmp_ok, cmp = pcall(require, "99.extensions.cmp")
+ if not cmp_ok then
+ vim.notify(
+ "[99] 99 completion module failed to load",
+ vim.log.levels.ERROR
+ )
+ return
+ end
return cmp
end
end
diff --git a/lua/99/extensions/work/worker.lua b/lua/99/extensions/work/worker.lua
new file mode 100644
index 0000000..5cce21c
--- /dev/null
+++ b/lua/99/extensions/work/worker.lua
@@ -0,0 +1,75 @@
+local Window = require("99.window")
+
+--- @class _99.Extension.Worker
+local M = {}
+
+--- @class _99.WorkOpts
+--- @field description string | nil
+
+--- @param opts _99.WorkOpts | nil
+function M.set_work(opts)
+ opts = opts or {}
+ local description = opts.description
+ if description then
+ M.current_work_item = description
+ else
+ Window.capture_input(" Work ", {
+ cb = function(success, result)
+ if not success then
+ return
+ end
+ M.current_work_item = result
+ end,
+
+ content = { "Put in the description of the work you want to complete" },
+ })
+ end
+
+ -- i think this makes sense. last work search should be cleared
+ M.last_work_search = nil
+end
+
+--- craft_prompt can be overridden so you can create your own prompt
+--- @param worker _99.Extension.Worker
+--- @return string
+function M.craft_prompt(worker)
+ return string.format(
+ [[
+<YourGoal>
+You are to take the current git diff and git diff --staged and figure out what is
+left to change to complete the work item. The work item is described in <Description>
+
+Carefully review everything in git diff and git diff --staged and <Description> before you respond.
+respond with proper Search Format described in <Rule> and an example in <Output>
+
+If you see bugs, also report those
+</YourGoal>
+<Description>
+%s
+</Description>
+]],
+ worker.current_work_item
+ )
+end
+
+function M.work()
+ assert(
+ M.current_work_item,
+ 'you must call "set_work" and set your current work item before calling this'
+ )
+ local _99 = require("99")
+ M.last_work_search = _99.search({
+ additional_prompt = M.craft_prompt(M),
+ })
+end
+
+function M.last_search_results()
+ if M.last_work_search == nil then
+ print("no previous search results")
+ return
+ end
+
+ require("99").qfix_search_results(M.last_work_search)
+end
+
+return M
diff --git a/lua/99/init.lua b/lua/99/init.lua
index ad16465..a512ab6 100644
--- a/lua/99/init.lua
+++ b/lua/99/init.lua
@@ -43,19 +43,22 @@ end
--- @alias _99.Cleanup fun(): nil
+--- @class _99.RequestEntry.Data.Search
+--- @field type "search"
+--- @field qfix_items _99.Search.Result[]
+
+--- @class _99.RequestEntry.Data.Visual
+--- @field type "visual"
+
+-- luacheck: ignore
+--- @alias _99.RequestEntry.Data _99.RequestEntry.Data.Search | _99.RequestEntry.Data.Tutorial | _99.RequestEntry.Data.Visual
+
--- @class _99.RequestEntry
---- @field id number
---- @field operation string
---- @field status "running" | "success" | "failed" | "cancelled"
---- @field filename string
---- @field lnum number
---- @field col number
+--- @field context _99.RequestContext
+--- @field status _99.Request.State
+--- @field point _99.Point
--- @field started_at number
-
---- @class _99.ActiveRequest
---- @field clean_up _99.Cleanup
---- @field request_id number
---- @field name string
+--- @field operation_data _99.RequestEntry.Data | nil
--- @class _99.StateProps
--- @field model string
@@ -67,10 +70,10 @@ end
--- @field display_errors boolean
--- @field auto_add_skills boolean
--- @field provider_override _99.Providers.BaseProvider?
---- @field __active_requests table<number, _99.ActiveRequest>
--- @field __view_log_idx number
--- @field __request_history _99.RequestEntry[]
--- @field __request_by_id table<number, _99.RequestEntry>
+--- @field tmp_dir string | nil
--- @return _99.StateProps
local function create_99_state()
@@ -84,15 +87,15 @@ local function create_99_state()
display_errors = false,
provider_override = nil,
auto_add_skills = false,
- __active_requests = {},
__view_log_idx = 1,
__request_history = {},
__request_by_id = {},
+ tmp_dir = nil,
}
end
--- @class _99.Completion
---- @field source "cmp" | nil
+--- @field source "cmp" | "blink" | nil
--- @field custom_rules string[]
--- @field files _99.Files.Config?
@@ -106,6 +109,7 @@ end
--- @field display_errors? boolean
--- @field auto_add_skills? boolean
--- @field completion _99.Completion?
+--- @field tmp_dir? string
--- unanswered question -- will i need to queue messages one at a time or
--- just send them all... So to prepare ill be sending around this state object
@@ -123,11 +127,11 @@ end
--- @field provider_override _99.Providers.BaseProvider?
--- @field auto_add_skills boolean
--- @field rules _99.Agents.Rules
---- @field __active_requests table<number, _99.ActiveRequest>
--- @field __view_log_idx number
--- @field __request_history _99.RequestEntry[]
--- @field __request_by_id table<number, _99.RequestEntry>
--- @field __active_marks _99.Mark[]
+--- @field tmp_dir string | nil
local _99_State = {}
_99_State.__index = _99_State
@@ -152,49 +156,63 @@ function _99_State:refresh_rules()
Extensions.refresh(self)
end
+--- @param tutorial _99.RequestEntry.Data.Tutorial
+function _99_State:open_tutorial(tutorial) end
+
--- @param context _99.RequestContext
--- @return _99.RequestEntry
function _99_State:track_request(context)
+ assert(
+ context.operation,
+ "must have an operation defined to track the request"
+ )
+
local point = context.range and context.range.start or Point:from_cursor()
local entry = {
- id = context.xid,
- operation = context.operation or "request",
- status = "running",
- filename = context.full_path,
- lnum = point.row,
- col = point.col,
+ context = context,
+ status = "requesting",
+ point = point,
started_at = time.now(),
+ operation_data = nil,
}
table.insert(self.__request_history, entry)
- self.__request_by_id[entry.id] = entry
+ self.__request_by_id[context.xid] = entry
return entry
end
---- @param id number
---- @param status "success" | "failed" | "cancelled"
-function _99_State:finish_request(id, status)
+--- @param context _99.RequestContext
+--- @param status _99.Request.ResponseState
+function _99_State:finish_request(context, status)
+ local id = context.xid
local entry = self.__request_by_id[id]
- if entry then
- entry.status = status
+ if not entry then
+ return
end
+
+ entry.status = status
end
---- @param id number
-function _99_State:remove_request(id)
- for i, entry in ipairs(self.__request_history) do
- if entry.id == id then
- table.remove(self.__request_history, i)
- break
- end
+--- @param context _99.RequestContext
+---@param data _99.RequestEntry.Data
+function _99_State:add_data(context, data)
+ local id = context.xid
+ local entry = self.__request_by_id[id]
+ if not entry then
+ return
end
- self.__request_by_id[id] = nil
+ local logger = Logger:set_id(id)
+ logger:assert(
+ entry.context.operation == data.type,
+ "the data type is not the same as the operation"
+ )
+ entry.operation_data = data
end
--- @return number
function _99_State:previous_request_count()
local count = 0
for _, entry in ipairs(self.__request_history) do
- if entry.status ~= "running" then
+ if entry.status ~= "requesting" then
count = count + 1
end
end
@@ -204,31 +222,15 @@ end
function _99_State:clear_previous_requests()
local keep = {}
for _, entry in ipairs(self.__request_history) do
- if entry.status == "running" then
+ if entry.status == "requesting" then
table.insert(keep, entry)
else
- self.__request_by_id[entry.id] = nil
+ self.__request_by_id[entry.context.xid] = nil
end
end
self.__request_history = keep
end
-local _active_request_id = 0
----@param clean_up _99.Cleanup
----@param request_id number
----@param name string
----@return number
-function _99_State:add_active_request(clean_up, request_id, name)
- _active_request_id = _active_request_id + 1
- Logger:debug("adding active request", "id", _active_request_id)
- self.__active_requests[_active_request_id] = {
- clean_up = clean_up,
- request_id = request_id,
- name = name,
- }
- return _active_request_id
-end
-
--- @param mark _99.Mark
function _99_State:add_mark(mark)
table.insert(self.__active_marks, mark)
@@ -236,19 +238,25 @@ end
function _99_State:active_request_count()
local count = 0
- for _ in pairs(self.__active_requests) do
- count = count + 1
+ for _, r in pairs(self.__request_history) do
+ if r.status == "requesting" then
+ count = count + 1
+ end
end
return count
end
----@param id number
-function _99_State:remove_active_request(id)
- local logger = Logger:set_id(id)
- local r = self.__active_requests[id]
- logger:assert(r, "there is no active request for id. implementation broken")
- logger:debug("removing active request")
- self.__active_requests[id] = nil
+--- @param type "search" | "visual" | "tutorial"
+--- @return _99.RequestEntry.Data
+function _99_State:get_request_data_by_type(type)
+ local out = {}
+ for _, r in ipairs(self.__request_history) do
+ local data = r.operation_data
+ if data and data.type == type then
+ table.insert(out, data)
+ end
+ end
+ return out
end
local _99_state = _99_State.new()
@@ -275,8 +283,11 @@ end
--- @param name string
--- @param context _99.RequestContext
--- @param opts _99.ops.Opts
-local function capture_prompt(cb, name, context, opts)
+--- @param capture_content string[] | nil
+local function capture_prompt(cb, name, context, opts, capture_content)
Window.capture_input(name, {
+ content = capture_content,
+
--- @param ok boolean
--- @param response string
cb = function(ok, response)
@@ -333,6 +344,56 @@ function _99.info()
Window.display_centered_message(info)
end
+--- @param tutorials _99.RequestEntry.Data.Tutorial[]
+--- @return string[]
+local function tutorial_to_string(tutorials)
+ local out = {}
+ for _, t in ipairs(tutorials) do
+ table.insert(out, string.format("%d: %s", t.xid, t.tutorial[1]))
+ end
+ return out
+end
+
+--- @param xid number | nil
+--- @param opts? _99.window.SplitWindowOpts
+function _99.open_tutorial(xid, opts)
+ opts = opts or { split_direction = "vertical" }
+ if xid == nil then
+ local tutorials = _99_state:get_request_data_by_type("tutorial")
+ if #tutorials == 0 then
+ print("no tutorials available")
+ elseif #tutorials == 1 then
+ local data = tutorials[1].operation_data
+ assert(data, "tutorial is malformed")
+ Window.create_split(data.tutorial, data.buffer, opts)
+ else
+ local context = get_context("tutorial-lookup")
+ capture_prompt(function(_, o)
+ local response = o.additional_prompt
+ local lines = vim.split(response, "\n")
+ for _, l in ipairs(lines) do
+ local id = tonumber(vim.split(l, ":")[1])
+ if not id then
+ error(
+ "do not alter the tutoria lines, just delete the ones you dont want"
+ )
+ end
+ local tut = _99_state.__request_by_id[id]
+ local data = tut and tut.operation_data
+ assert(data and data.type == "tutorial", "invalid tutorial selected")
+ Window.create_split(data.tutorial, data.buffer, opts)
+ end
+ end, "Select Tutorial", context, {}, tutorial_to_string(tutorials))
+ end
+ return
+ end
+
+ local tutorial = _99_state.__request_by_id[xid]
+ local data = tutorial and tutorial.operation_data
+ assert(data and data.type == "tutorial", "cannot open a non tutorial")
+ Window.create_split(data.tutorial, data.buffer, opts)
+end
+
--- @param path string
function _99:rule_from_path(path)
_ = self
@@ -341,36 +402,27 @@ function _99:rule_from_path(path)
end
--- @param opts? _99.ops.SearchOpts
+--- @return number
function _99.search(opts)
local o = process_opts(opts) --[[ @as _99.ops.SearchOpts ]]
local context = get_context("search")
if o.additional_prompt then
ops.search(context, o)
- return
else
capture_prompt(ops.search, "Search", context, o)
end
+ return context.xid
end
--- @param opts _99.ops.Opts
-function _99.visual_prompt(opts)
- vim.notify(
- "use visual, visual_prompt has been deprecated",
- vim.log.levels.WARN
- )
- _99.visual(opts)
-end
-
-function _99.fill_in_function()
- error(
- "function has been removed. Just use visual. I really hate fill in function, sorry :)"
- )
-end
-
-function _99.fill_in_function_prompt()
- error(
- "function has been removed. Just use visual. I really hate fill in function, sorry :)"
- )
+function _99.tutorial(opts)
+ opts = process_opts(opts)
+ local context = get_context("tutorial")
+ if opts.additional_prompt then
+ ops.tutorial(context, opts)
+ else
+ capture_prompt(ops.tutorial, "Tutorial", context, opts)
+ end
end
--- @param opts _99.ops.Opts?
@@ -421,12 +473,18 @@ function _99.next_request_logs()
Window.display_full_screen_message(logs[_99_state.__view_log_idx])
end
+--- @class _99.QFixEntry
+--- @field filename string
+--- @field lnum number
+--- @field col number
+--- @field text string
+
function _99.stop_all_requests()
- for _, active in pairs(_99_state.__active_requests) do
- _99_state:remove_request(active.request_id)
- active.clean_up()
+ for _, request in pairs(_99_state.__request_by_id) do
+ if request.status == "requesting" then
+ request.context:stop()
+ end
end
- _99_state.__active_requests = {}
end
function _99.clear_all_marks()
@@ -436,17 +494,18 @@ function _99.clear_all_marks()
_99_state.__active_marks = {}
end
-function _99.previous_requests_to_qfix()
- local items = {}
- for _, entry in ipairs(_99_state.__request_history) do
- table.insert(items, {
- filename = entry.filename,
- lnum = entry.lnum,
- col = entry.col,
- text = string.format("[%s] %s", entry.status, entry.operation),
- })
- end
- vim.fn.setqflist({}, "r", { title = "99 Requests", items = items })
+--- @param xid number | nil
+function _99.qfix_search_results(xid)
+ --- @type _99.RequestEntry
+ local entry = _99_state.__request_by_id[xid]
+ assert(entry, "qfix_search_results could not find id: " .. xid)
+
+ local data = entry.operation_data
+ assert(data, "there must be data associated with request entry")
+ assert(data.type == "search", "the operation_data must be type search")
+
+ local items = data.qfix_items
+ vim.fn.setqflist({}, "r", { title = "99 Search Results", items = items })
vim.cmd("copen")
end
@@ -497,11 +556,15 @@ local function show_in_flight_requests()
return shut_down_in_flight_requests_window()
end
+ --- @type string[]
local lines = {
throb .. " requests(" .. tostring(count) .. ") " .. throb,
}
- for _, r in pairs(_99_state.__active_requests) do
- table.insert(lines, r.name)
+
+ for _, r in pairs(_99_state.__request_by_id) do
+ if r.status == "requesting" then
+ table.insert(lines, r.context.operation)
+ end
end
Window.resize(win, #lines[1], #lines)
@@ -562,6 +625,11 @@ function _99.setup(opts)
end
end
+ if opts.tmp_dir then
+ assert(type(opts.tmp_dir) == "string", "opts.tmp_dir must be a string")
+ end
+ _99_state.tmp_dir = opts.tmp_dir
+
_99_state.display_errors = opts.display_errors or false
_99_state:refresh_rules()
Languages.initialize(_99_state)
@@ -627,4 +695,7 @@ function _99.__debug()
end
_99.Providers = Providers
+_99.Extensions = {
+ Worker = require("99.extensions.work.worker"),
+}
return _99
diff --git a/lua/99/logger/logger.lua b/lua/99/logger/logger.lua
index b2dc37b..611b016 100644
--- a/lua/99/logger/logger.lua
+++ b/lua/99/logger/logger.lua
@@ -66,9 +66,13 @@ FileSink.__index = FileSink
--- @param path string
--- @return LoggerSink
function FileSink:new(path)
- local fd, err = vim.uv.fs_open(path, "w", 493)
+ -- Ensure the directory is already there (*thanks Windows*)
+ vim.fn.mkdir(vim.fn.fnamemodify(path, ":h"), "p")
+
+ -- 420 decimal == 644 octal (rw-r--r--)
+ local fd, err = vim.uv.fs_open(path, "w", 420)
if not fd then
- error("unable to file sink", err)
+ error("unable to file sink: " .. err)
end
return setmetatable({
diff --git a/lua/99/ops/clean-up.lua b/lua/99/ops/clean-up.lua
index ba76f01..cbf077a 100644
--- a/lua/99/ops/clean-up.lua
+++ b/lua/99/ops/clean-up.lua
@@ -1,20 +1,56 @@
----@param context _99.RequestContext
----@param name string
+local M = {}
+
+--- @alias _99.Providers.on_complete fun(status: _99.Request.ResponseState, response: string): nil
+--- @class _99.Providers.PartialObserver
+--- @field on_complete _99.Providers.on_complete
+--- @field on_stdout? fun(line: string): nil
+--- @field on_stderr? fun(line: string): nil
+--- @field on_start? fun(): nil
+
+--- @param clean_up fun(): nil
+--- @param obs_or_fn _99.Providers.PartialObserver | _99.Providers.on_complete
+--- @return _99.Providers.Observer
+M.make_observer = function(clean_up, obs_or_fn)
+ --- @type _99.Providers.PartialObserver
+ local obs = type(obs_or_fn) == "table" and obs_or_fn
+ or {
+ on_complete = obs_or_fn,
+ }
+ return {
+ on_start = function()
+ if obs.on_start then
+ obs.on_start()
+ end
+ end,
+ on_complete = function(status, res)
+ vim.schedule(clean_up)
+ obs.on_complete(status, res)
+ end,
+ on_stderr = function(line)
+ if obs.on_stderr then
+ obs.on_stderr(line)
+ end
+ end,
+ on_stdout = function(line)
+ if obs.on_stdout then
+ obs.on_stdout(line)
+ end
+ end,
+ } --[[@as _99.Providers.Observer ]]
+end
+
---@param clean_up_fn fun(): nil
---@return fun(): nil
-return function(context, name, clean_up_fn)
+M.make_clean_up = function(clean_up_fn)
local called = false
- local request_id = -1
local function clean_up()
if called then
return
end
-
called = true
clean_up_fn()
- context._99:remove_active_request(request_id)
end
- request_id = context._99:add_active_request(clean_up, context.xid, name)
-
return clean_up
end
+
+return M
diff --git a/lua/99/ops/implement-fn.lua b/lua/99/ops/implement-fn.lua
deleted file mode 100644
index efec8a3..0000000
--- a/lua/99/ops/implement-fn.lua
+++ /dev/null
@@ -1,92 +0,0 @@
-local Request = require("99.request")
-local editor = require("99.editor")
-local geo = require("99.geo")
-local Range = geo.Range
-local Point = geo.Point
-local Mark = require("99.ops.marks")
-local RequestStatus = require("99.ops.request_status")
-local make_clean_up = require("99.ops.clean-up")
-
---- @param context _99.RequestContext
---- @param response string
-local function update_code(context, response)
- local code_mark = context.marks.code_placement
- local logger = context.logger:set_area("implement_fn#update_code")
- local point = Point.from_mark(code_mark)
-
- logger:debug("setting text at mark", "Point", point)
- code_mark:set_text_at_mark("\n" .. response)
-end
-
---- @param context _99.RequestContext
-local function implement_fn(context)
- local ts = editor.treesitter
- local cursor = Point:from_cursor()
- local buffer = vim.api.nvim_get_current_buf()
- local fn_call = ts.fn_call(buffer, cursor)
- local logger = context.logger:set_area("implement_fn")
-
- if not fn_call then
- logger:fatal(
- "cannot implement function, cursor was not on an identifier that is a function call"
- )
- return
- end
-
- local range = Range:from_ts_node(fn_call, buffer)
- local request = Request.new(context)
-
- context.marks.end_of_fn_call = Mark.mark_end_of_range(buffer, range)
- local func = ts.containing_function(buffer, cursor)
- if func then
- context.marks.code_placement = Mark.mark_above_func(buffer, func)
- else
- context.marks.code_placement = Mark.mark_above_range(range)
- end
-
- local code_placement = RequestStatus.new(
- 250,
- context._99.ai_stdout_rows,
- "Loading",
- context.marks.code_placement
- )
- local at_call_site = RequestStatus.new(
- 250,
- 1,
- "Implementing Function",
- context.marks.end_of_fn_call
- )
-
- code_placement:start()
- at_call_site:start()
-
- local clean_up = make_clean_up(context, function()
- context:clear_marks()
- request:cancel()
- code_placement:stop()
- at_call_site:stop()
- end)
-
- request:add_prompt_content(context._99.prompts.prompts.implement_function)
- request:start({
- on_stdout = function(line)
- code_placement:push(line)
- end,
- on_complete = function(status, response)
- vim.schedule(clean_up)
- if status ~= "success" then
- logger:fatal(
- "unable to implement function, enable and check logger for more details"
- )
- end
- pcall(update_code, context, response)
- end,
- on_stderr = function(line)
- logger:error("stderr", "line", line)
- end,
- })
-
- return request
-end
-
-return implement_fn
diff --git a/lua/99/ops/init.lua b/lua/99/ops/init.lua
index 2e386d8..1f8e9f8 100644
--- a/lua/99/ops/init.lua
+++ b/lua/99/ops/init.lua
@@ -8,6 +8,6 @@
return {
search = require("99.ops.search"),
- implement_fn = require("99.ops.implement-fn"),
+ tutorial = require("99.ops.tutorial"),
over_range = require("99.ops.over-range"),
}
diff --git a/lua/99/ops/make-prompt.lua b/lua/99/ops/make-prompt.lua
new file mode 100644
index 0000000..4a17687
--- /dev/null
+++ b/lua/99/ops/make-prompt.lua
@@ -0,0 +1,32 @@
+local Completions = require("99.extensions.completions")
+local Agents = require("99.extensions.agents")
+
+--- @param context _99.RequestContext
+--- @param prompt string
+--- @param opts _99.ops.Opts
+--- @return string, _99.Reference[]
+return function(context, prompt, opts)
+ local user_prompt = opts.additional_prompt
+ assert(
+ user_prompt and type(user_prompt) == "string" and #user_prompt > 0,
+ "you must add a prompt to you request"
+ )
+
+ local full_prompt = prompt
+ full_prompt = context._99.prompts.prompts.prompt(user_prompt, full_prompt)
+
+ local refs = Completions.parse(user_prompt)
+ local additional_rules = opts.additional_rules
+ if additional_rules then
+ for _, r in ipairs(additional_rules) do
+ local content = Agents.get_rule_content(r)
+ if content then
+ table.insert(refs, {
+ content = content,
+ })
+ end
+ end
+ end
+
+ return full_prompt, refs
+end
diff --git a/lua/99/ops/over-range.lua b/lua/99/ops/over-range.lua
index fec2794..7070857 100644
--- a/lua/99/ops/over-range.lua
+++ b/lua/99/ops/over-range.lua
@@ -2,8 +2,11 @@ local Request = require("99.request")
local RequestStatus = require("99.ops.request_status")
local Mark = require("99.ops.marks")
local geo = require("99.geo")
-local make_clean_up = require("99.ops.clean-up")
-local Completions = require("99.extensions.completions")
+local make_prompt = require("99.ops.make-prompt")
+local CleanUp = require("99.ops.clean-up")
+
+local make_clean_up = CleanUp.make_clean_up
+local make_observer = CleanUp.make_observer
local Range = geo.Range
local Point = geo.Point
@@ -37,34 +40,24 @@ local function over_range(context, range, opts)
top_mark
)
local bottom_status = RequestStatus.new(250, 1, "Implementing", bottom_mark)
- local clean_up = make_clean_up(context, "Visual", function()
+ local clean_up = make_clean_up(function()
top_status:stop()
bottom_status:stop()
context:clear_marks()
request:cancel()
end)
- local full_prompt = context._99.prompts.prompts.visual_selection(range)
- local additional_prompt = opts.additional_prompt
- if additional_prompt then
- full_prompt =
- context._99.prompts.prompts.prompt(additional_prompt, full_prompt)
-
- local refs = Completions.parse(additional_prompt)
- context:add_references(refs)
- end
+ local system_cmd = context._99.prompts.prompts.visual_selection(range)
+ local prompt, refs = make_prompt(context, system_cmd, opts)
- local additional_rules = opts.additional_rules
- if additional_rules then
- context:add_agent_rules(additional_rules)
- end
+ request:add_prompt_content(prompt)
+ context:add_references(refs)
+ context:add_clean_up(clean_up)
- request:add_prompt_content(full_prompt)
top_status:start()
bottom_status:start()
- request:start({
+ request:start(make_observer(clean_up, {
on_complete = function(status, response)
- vim.schedule(clean_up)
if status == "cancelled" then
logger:debug("request cancelled for visual selection, removing marks")
elseif status == "failed" then
@@ -99,10 +92,7 @@ local function over_range(context, range, opts)
top_status:push(line)
end
end,
- on_stderr = function(line)
- logger:debug("visual_selection#on_stderr received", "line", line)
- end,
- })
+ }))
end
return over_range
diff --git a/lua/99/ops/search.lua b/lua/99/ops/search.lua
index a4f40b3..d6e5428 100644
--- a/lua/99/ops/search.lua
+++ b/lua/99/ops/search.lua
@@ -1,8 +1,9 @@
local Request = require("99.request")
-local make_clean_up = require("99.ops.clean-up")
-local Completions = require("99.extensions.completions")
-local Mark = require("99.ops.marks")
-local Point = require("99.geo").Point
+local make_prompt = require("99.ops.make-prompt")
+local CleanUp = require("99.ops.clean-up")
+
+local make_clean_up = CleanUp.make_clean_up
+local make_observer = CleanUp.make_observer
--- @class _99.Search.Result
--- @field filename string
@@ -35,10 +36,9 @@ local function parse_line(line)
}
end
---- @param _99 _99.State
+--- @param context _99.RequestContext
--- @param response string
-local function create_search_locations(_99, response)
- _ = _99
+local function create_search_locations(context, response)
local lines = vim.split(response, "\n")
local qf_list = {}
@@ -48,10 +48,13 @@ local function create_search_locations(_99, response)
table.insert(qf_list, res)
end
end
+ context._99:add_data(context, {
+ type = "search",
+ qfix_items = qf_list,
+ })
if #qf_list > 0 then
- vim.fn.setqflist(qf_list, "r")
- vim.cmd("copen")
+ require("99").qfix_search_results(context.xid)
else
vim.notify("No search results found", vim.log.levels.INFO)
end
@@ -61,58 +64,35 @@ end
---@param opts _99.ops.SearchOpts
local function search(context, opts)
opts = opts or {}
- local user_prompt = opts.additional_prompt
- assert(user_prompt, "search requires a prompt to run, please provide prompt")
local logger = context.logger:set_area("search")
local request = Request.new(context)
logger:debug("search", "with opts", opts.additional_prompt)
- -- TODO: How to surface progress.. I was thinking about a status line plugin
- -- local top_status = RequestStatus.new(
- -- 250,
- -- context._99.ai_stdout_rows or 1,
- -- "Implementing",
- -- top_mark
- -- )
- local clean_up = make_clean_up(context, "Search", function()
+ local clean_up = make_clean_up(function()
request:cancel()
end)
- local full_prompt = context._99.prompts.prompts.semantic_search()
- full_prompt = context._99.prompts.prompts.prompt(user_prompt, full_prompt)
- local refs = Completions.parse(user_prompt)
- context:add_references(refs)
+ local prompt, refs =
+ make_prompt(context, context._99.prompts.prompts.semantic_search(), opts)
- local additional_rules = opts.additional_rules
- if additional_rules then
- context:add_agent_rules(additional_rules)
- end
+ request:add_prompt_content(prompt)
+ context:add_references(refs)
+ context:add_clean_up(clean_up)
- request:add_prompt_content(full_prompt)
- request:start({
- on_complete = function(status, response)
- vim.schedule(clean_up)
- if status == "cancelled" then
- logger:debug("request cancelled for search")
- elseif status == "failed" then
- logger:error(
- "request failed for search",
- "error response",
- response or "no response provided"
- )
- elseif status == "success" then
- create_search_locations(context._99, response)
- end
- end,
- on_stdout = function(line)
- --- TODO: i need to figure out how to surface this information
- _ = line
- end,
- on_stderr = function(line)
- logger:debug("visual_selection#on_stderr received", "line", line)
- end,
- })
+ request:start(make_observer(clean_up, function(status, response)
+ if status == "cancelled" then
+ logger:debug("request cancelled for search")
+ elseif status == "failed" then
+ logger:error(
+ "request failed for search",
+ "error response",
+ response or "no response provided"
+ )
+ elseif status == "success" then
+ create_search_locations(context, response)
+ end
+ end))
end
return search
diff --git a/lua/99/ops/tutorial.lua b/lua/99/ops/tutorial.lua
new file mode 100644
index 0000000..320366f
--- /dev/null
+++ b/lua/99/ops/tutorial.lua
@@ -0,0 +1,72 @@
+local Request = require("99.request")
+local CleanUp = require("99.ops.clean-up")
+local Window = require("99.window")
+local make_prompt = require("99.ops.make-prompt")
+
+local make_clean_up = CleanUp.make_clean_up
+local make_observer = CleanUp.make_observer
+
+--- @class _99.RequestEntry.Data.Tutorial
+--- @field type "tutorial"
+--- @field buffer number
+--- @field window number
+--- @field xid number
+--- @field tutorial string[]
+
+--- @param context _99.RequestContext
+---@param response string
+---@return _99.RequestEntry.Data.Tutorial
+local function open_tutorial(context, response)
+ local content = vim.split(response, "\n")
+ local win = Window.create_split(content)
+
+ --- @type _99.RequestEntry.Data.Tutorial
+ local data = {
+ type = "tutorial",
+ buffer = win.buffer,
+ window = win.win,
+ xid = context.xid,
+ tutorial = content,
+ }
+ context._99:add_data(context, data)
+ return data
+end
+
+--- @param context _99.RequestContext
+---@param opts _99.ops.Opts
+local function tutorial(context, opts)
+ opts = opts or {}
+
+ local logger = context.logger:set_area("tutorial")
+ logger:debug("starting", "with opts", opts)
+
+ local request = Request.new(context)
+
+ local clean_up = make_clean_up(function()
+ request:cancel()
+ end)
+
+ local prompt, refs =
+ make_prompt(context, context._99.prompts.prompts.tutorial(), opts)
+
+ context:add_references(refs)
+ request:add_prompt_content(prompt)
+ context:add_clean_up(clean_up)
+
+ request:start(make_observer(clean_up, function(status, response)
+ vim.schedule(clean_up)
+ if status == "cancelled" then
+ logger:debug("cancelled")
+ elseif status == "failed" then
+ logger:error(
+ "failed",
+ "error response",
+ response or "no response provided"
+ )
+ elseif status == "success" then
+ local data = open_tutorial(context, response)
+ context._99:open_tutorial(data)
+ end
+ end))
+end
+return tutorial
diff --git a/lua/99/prompt-settings.lua b/lua/99/prompt-settings.lua
index 0bc564d..03f82a2 100644
--- a/lua/99/prompt-settings.lua
+++ b/lua/99/prompt-settings.lua
@@ -8,6 +8,7 @@ end
--- @class _99.Prompts.SpecificOperations
--- @field visual_selection fun(range: _99.Range): string
--- @field semantic_search fun(): string
+--- @field tutorial fun(): string
--- @field prompt fun(prompt: string, action: string, name?: string): string
--- @field role fun(): string
--- @field read_tmp fun(): string
@@ -15,6 +16,17 @@ local prompts = {
role = function()
return [[ You are a software engineering assistant mean to create robust and conanical code ]]
end,
+ tutorial = function()
+ return [[
+You are given a prompt and context and you must craft a tutorial. If a set of
+context has links, read through them thoroughly and decide which ones to retrieve.
+Once you have fetched all the relavent content, review it thoroughly before
+crafting the tutorial
+
+<Rule>The response format must be valid Markdown</Rule>
+<Rule>The first line of the response must be the title of the tutorial</Rule>
+]]
+ end,
semantic_search = function()
return [[
you are given a prompt and you must search through this project and return code that matches the description provided.
diff --git a/lua/99/providers.lua b/lua/99/providers.lua
index b26982e..3a2c3ec 100644
--- a/lua/99/providers.lua
+++ b/lua/99/providers.lua
@@ -2,14 +2,7 @@
--- @field on_stdout fun(line: string): nil
--- @field on_stderr fun(line: string): nil
--- @field on_complete fun(status: _99.Request.ResponseState, res: string): nil
-
---- @type _99.Providers.Observer
-local DevNullObserver = {
- name = "DevNullObserver",
- on_stdout = function() end,
- on_stderr = function() end,
- on_complete = function() end,
-}
+--- @field on_start fun(): nil
--- @param fn fun(...: any): nil
--- @return fun(...: any): nil
@@ -61,15 +54,21 @@ end
--- @param query string
--- @param request _99.Request
---- @param observer _99.Providers.Observer?
+--- @param observer _99.Providers.Observer
function BaseProvider:make_request(query, request, observer)
+ observer.on_start()
+
local logger = request.logger:set_area(self:_get_provider_name())
logger:debug("make_request", "tmp_file", request.context.tmp_file)
- observer = observer or DevNullObserver
- local once_complete = once(function(status, text)
- observer.on_complete(status, text)
- end)
+ local once_complete = once(
+ --- @param status "success" | "failed" | "cancelled"
+ ---@param text string
+ function(status, text)
+ request.state = status
+ observer.on_complete(status, text)
+ end
+ )
local command = self:_build_command(query, request)
logger:debug("make_request", "command", command)
@@ -294,9 +293,43 @@ function KiroProvider._get_default_model()
return "claude-sonnet-4.5"
end
+--- @class GeminiCLIProvider : _99.Providers.BaseProvider
+local GeminiCLIProvider = setmetatable({}, { __index = BaseProvider })
+
+--- @param query string
+--- @param request _99.Request
+--- @return string[]
+function GeminiCLIProvider._build_command(_, query, request)
+ return {
+ "gemini",
+ "--approval-mode",
+ -- Allow writing to temp files by default. See:
+ -- https://geminicli.com/docs/core/policy-engine/#default-policies
+ "auto_edit",
+ "--model",
+ request.context.model,
+ "--prompt",
+ query,
+ }
+end
+
+--- @return string
+function GeminiCLIProvider._get_provider_name()
+ return "GeminiCLIProvider"
+end
+
+--- @return string
+function GeminiCLIProvider._get_default_model()
+ -- Default to auto-routing between pro and flash. See:
+ -- https://geminicli.com/docs/cli/model/
+ return "auto"
+end
+
return {
+ BaseProvider = BaseProvider,
OpenCodeProvider = OpenCodeProvider,
ClaudeCodeProvider = ClaudeCodeProvider,
CursorAgentProvider = CursorAgentProvider,
KiroProvider = KiroProvider,
+ GeminiCLIProvider = GeminiCLIProvider,
}
diff --git a/lua/99/request-context.lua b/lua/99/request-context.lua
index 6671cab..48d4489 100644
--- a/lua/99/request-context.lua
+++ b/lua/99/request-context.lua
@@ -15,6 +15,7 @@ local random_file = utils.random_file
--- @field xid number
--- @field range _99.Range?
--- @field operation string?
+--- @field clean_ups (fun(): nil)[]
--- @field _99 _99.State
local RequestContext = {}
RequestContext.__index = RequestContext
@@ -36,11 +37,17 @@ function RequestContext.from_current_buffer(_99, xid)
table.insert(mds, md)
end
+ local tmp_dir = _99.tmp_dir
+ if tmp_dir then
+ tmp_dir = vim.fn.expand(tmp_dir)
+ end
+
return setmetatable({
_99 = _99,
+ clean_ups = {},
md_file_names = mds,
ai_context = {},
- tmp_file = random_file(),
+ tmp_file = random_file(tmp_dir),
buffer = buffer,
full_path = full_path,
file_type = file_type,
@@ -51,6 +58,17 @@ function RequestContext.from_current_buffer(_99, xid)
}, RequestContext)
end
+function RequestContext:stop()
+ for _, cb in ipairs(self.clean_ups) do
+ cb()
+ end
+end
+
+--- @param clean_up fun(): nil
+function RequestContext:add_clean_up(clean_up)
+ table.insert(self.clean_ups, clean_up)
+end
+
--- @param md_file_name string
--- @return self
function RequestContext:add_md_file_name(md_file_name)
@@ -58,40 +76,6 @@ function RequestContext:add_md_file_name(md_file_name)
return self
end
---- TODO: Dedupe any rules that have already been added
---- @param rules (_99.Agents.Rule | string)[]
-function RequestContext:add_agent_rules(rules)
- for _, rule in ipairs(rules) do
- -- Handle both string paths and rule objects
- self.logger:debug("adding custom rule to agent", "rule", rule)
- local file_path = rule.absolute_path or rule.path
- local ok, file = pcall(io.open, file_path, "r")
- if ok and file then
- local content = file:read("*a")
- file:close()
- self.logger:info(
- "Context#adding agent file to the context",
- "agent_path",
- rule.path
- )
- table.insert(
- self.ai_context,
- string.format(
- [[
-<%s>
-%s
-</%s>]],
- rule.name,
- content,
- rule.name
- )
- )
- else
- self.logger:debug("unable to read agent rule", "rule", rule)
- end
- end
-end
-
--- @param refs _99.Reference[]
function RequestContext:add_references(refs)
for _, ref in ipairs(refs) do
@@ -133,16 +117,34 @@ function RequestContext:content()
return self.ai_context
end
---- @param prompt string
-function RequestContext:save_prompt(prompt)
+--- @return boolean
+function RequestContext:_ready_request_files()
+ local response_file = self.tmp_file
local prompt_file = self.tmp_file .. "-prompt"
local dir = vim.fs.dirname(prompt_file)
if dir and not vim.uv.fs_stat(dir) then
- pcall(vim.uv.fs_mkdir, dir, 493)
+ vim.fn.mkdir(dir, "p")
+ end
+
+ local files = { prompt_file, response_file }
+ for _, f in ipairs(files) do
+ local file = io.open(f, "w")
+ if file then
+ file:write("")
+ file:close()
+ else
+ self.logger:error("unable to create prompt file")
+ return false
+ end
end
+ return true
+end
+--- @param prompt string
+function RequestContext:save_prompt(prompt)
+ local prompt_file = self.tmp_file .. "-prompt"
local file = io.open(prompt_file, "w")
if file then
file:write(prompt)
@@ -153,8 +155,11 @@ function RequestContext:save_prompt(prompt)
end
end
---- @return self
+--- @return boolean, self
function RequestContext:finalize()
+ if self:_ready_request_files() == false then
+ return false, self
+ end
self:_read_md_files()
if self.range then
table.insert(self.ai_context, self._99.prompts.get_file_location(self))
@@ -164,7 +169,7 @@ function RequestContext:finalize()
self.ai_context,
self._99.prompts.tmp_file_location(self.tmp_file)
)
- return self
+ return true, self
end
function RequestContext:clear_marks()
diff --git a/lua/99/request/init.lua b/lua/99/request/init.lua
index 371391d..8080211 100644
--- a/lua/99/request/init.lua
+++ b/lua/99/request/init.lua
@@ -1,5 +1,5 @@
---- @alias _99.Request.State "ready" | "calling-model" | "parsing-result" | "updating-file" | "cancelled"
--- @alias _99.Request.ResponseState "failed" | "success" | "cancelled"
+--- @alias _99.Request.State "ready" | "requesting" | _99.Request.ResponseState
local Providers = require("99.providers")
@@ -47,15 +47,21 @@ function Request:_set_process(proc)
end
function Request:cancel()
+ if self.state == "success" or self.state == "failed" then
+ return
+ end
+
self.logger:debug("cancel")
self.state = "cancelled"
+ local proc = self._proc
---@diagnostic disable-next-line: undefined-field
- if self._proc and self._proc.pid then
+ if proc and proc.pid then
+ self._proc = nil
pcall(function()
local sigterm = (vim.uv and vim.uv.constants and vim.uv.constants.SIGTERM)
or 15
---@diagnostic disable-next-line: undefined-field
- self._proc:kill(sigterm)
+ proc:kill(sigterm)
end)
end
end
@@ -71,18 +77,62 @@ function Request:add_prompt_content(content)
return self
end
+--- @param r _99.Request
+--- @param obs _99.Providers.Observer | nil
+local function observer_from_request(r, obs)
+ local context = r.context
+ return {
+ on_start = function()
+ r.state = "requesting"
+ context._99:track_request(context)
+ if obs then
+ obs.on_start()
+ end
+ end,
+ on_complete = function(status, res)
+ r.state = status
+ context._99:finish_request(context, status)
+ if obs then
+ obs.on_complete(status, res)
+ end
+ end,
+ on_stderr = function(line)
+ if obs then
+ obs.on_stderr(line)
+ end
+ end,
+ on_stdout = function(line)
+ if obs then
+ obs.on_stdout(line)
+ end
+ end,
+ }
+end
+
--- @param observer _99.Providers.Observer?
function Request:start(observer)
- self.context._99:track_request(self.context)
- self.context:finalize()
+ self.logger:assert(
+ self.state == "ready",
+ "request is not in state ready when attempting to start a request"
+ )
+ local ok = self.context:finalize()
+ self.logger:assert(
+ ok,
+ "request has failed due to context finalization: check logs for more details"
+ )
+
for _, content in ipairs(self.context.ai_context) do
self:add_prompt_content(content)
end
-
local prompt = table.concat(self._content, "\n")
+
self.context:save_prompt(prompt)
self.logger:debug("start", "prompt", prompt)
- self.provider:make_request(prompt, self, observer)
+ self.provider:make_request(
+ prompt,
+ self,
+ observer_from_request(self, observer)
+ )
end
return Request
diff --git a/lua/99/test/marks_spec.lua b/lua/99/test/marks_spec.lua
index cea43cb..701d4a7 100644
--- a/lua/99/test/marks_spec.lua
+++ b/lua/99/test/marks_spec.lua
@@ -27,7 +27,7 @@ describe("Mark", function()
end)
it("should get mark point from visual selection", function()
- local _, buf = test_utils.fif_setup({
+ local _, buf = test_utils.test_setup({
"local test_1 = 0",
"local test_2 = 0",
"local test_3 = 0",
diff --git a/lua/99/test/providers_spec.lua b/lua/99/test/providers_spec.lua
index 809080b..9c72d0e 100644
--- a/lua/99/test/providers_spec.lua
+++ b/lua/99/test/providers_spec.lua
@@ -66,6 +66,27 @@ describe("providers", function()
end)
end)
+ describe("GeminiCLIProvider", function()
+ it("builds correct command with model", function()
+ local request = { context = { model = "gemini-2.5-pro" } }
+ local cmd =
+ Providers.GeminiCLIProvider._build_command(nil, "test query", request)
+ eq({
+ "gemini",
+ "--approval-mode",
+ "auto_edit",
+ "--model",
+ "gemini-2.5-pro",
+ "--prompt",
+ "test query",
+ }, cmd)
+ end)
+
+ it("has correct default model", function()
+ eq("auto", Providers.GeminiCLIProvider._get_default_model())
+ end)
+ end)
+
describe("provider integration", function()
it("can be set as provider override", function()
local _99 = require("99")
@@ -108,6 +129,17 @@ describe("providers", function()
end
)
+ it(
+ "uses GeminiCLIProvider default model when provider specified but no model",
+ function()
+ local _99 = require("99")
+
+ _99.setup({ provider = Providers.GeminiCLIProvider })
+ local state = _99.__get_state()
+ eq("auto", state.model)
+ end
+ )
+
it("uses custom model when both provider and model specified", function()
local _99 = require("99")
@@ -125,6 +157,7 @@ describe("providers", function()
eq("function", type(Providers.OpenCodeProvider.make_request))
eq("function", type(Providers.ClaudeCodeProvider.make_request))
eq("function", type(Providers.CursorAgentProvider.make_request))
+ eq("function", type(Providers.GeminiCLIProvider.make_request))
end)
end)
end)
diff --git a/lua/99/test/request_spec.lua b/lua/99/test/request_spec.lua
new file mode 100644
index 0000000..3c58ecd
--- /dev/null
+++ b/lua/99/test/request_spec.lua
@@ -0,0 +1,54 @@
+-- luacheck: globals describe it assert
+local _99 = require("99")
+local test_utils = require("99.test.test_utils")
+local eq = assert.are.same
+
+local content = {
+ "local function foo()",
+ " -- TODO: implement",
+ "end",
+}
+
+describe("request test", function()
+ it("should replace visual selection with AI response", function()
+ local p = test_utils.test_setup(content, 2, 1, "lua")
+ local state = _99.__get_state()
+ local Request = require("99.request")
+ local RequestContext = require("99.request-context")
+
+ local context = RequestContext.from_current_buffer(state, 100)
+ context.operation = "test_request"
+ context:finalize()
+
+ local request = Request.new(context)
+
+ local finished_called = false
+ local finished_status = nil
+
+ eq("ready", request.state)
+
+ eq(0, state:active_request_count())
+ request:start({
+ on_start = function()
+ print("on_start")
+ end,
+ on_complete = function(status, _)
+ finished_called = true
+ finished_status = status
+ end,
+ on_stdout = function() end,
+ on_stderr = function() end,
+ })
+ test_utils.next_frame()
+ eq(1, state:active_request_count())
+
+ eq("requesting", request.state)
+
+ p:resolve("success", " return 'implemented!'")
+ assert.is_true(finished_called)
+
+ eq(0, state:active_request_count())
+ eq("success", request.state)
+ eq("success", finished_status)
+ end)
+end)
diff --git a/lua/99/test/test_utils.lua b/lua/99/test/test_utils.lua
index b3bd7db..a1459a4 100644
--- a/lua/99/test/test_utils.lua
+++ b/lua/99/test/test_utils.lua
@@ -1,6 +1,14 @@
local Levels = require("99.logger.level")
local M = {}
+--- @type _99.Providers.Observer
+local DevNullObserver = {
+ on_start = function() end,
+ on_complete = function() end,
+ on_stderr = function() end,
+ on_stdout = function() end,
+}
+
function M.next_frame()
local next = false
vim.schedule(function()
@@ -17,7 +25,7 @@ M.created_files = {}
--- @class _99.test.ProviderRequest
--- @field query string
--- @field request _99.Request
---- @field observer _99.Providers.Observer?
+--- @field observer _99.Providers.Observer
--- @field logger _99.Logger
--- @class _99.test.Provider : _99.Providers.BaseProvider
@@ -35,6 +43,10 @@ end
function TestProvider:make_request(query, request, observer)
local logger = request.context.logger:set_area("TestProvider")
logger:debug("make_request", "tmp_file", request.context.tmp_file)
+
+ observer = observer or DevNullObserver
+ observer.on_start()
+
self.request = {
query = query,
request = request,
@@ -47,34 +59,26 @@ end
--- @param result string
function TestProvider:resolve(status, result)
assert(self.request, "you cannot call resolve until make_request is called")
- local obs = self.request.observer
- if obs then
- --- to match the behavior expected from the OpenCodeProvider
- if self.request.request:is_cancelled() then
- obs.on_complete("cancelled", result)
- else
- obs.on_complete(status, result)
- end
+
+ if self.request.request:is_cancelled() then
+ self.request.observer.on_complete("cancelled", result)
+ else
+ self.request.observer.on_complete(status, result)
end
+
self.request = nil
end
--- @param line string
function TestProvider:stdout(line)
assert(self.request, "you cannot call stdout until make_request is called")
- local obs = self.request.observer
- if obs then
- obs.on_stdout(line)
- end
+ self.request.observer.on_stdout(line)
end
--- @param line string
function TestProvider:stderr(line)
assert(self.request, "you cannot call stderr until make_request is called")
- local obs = self.request.observer
- if obs then
- obs.on_stderr(line)
- end
+ self.request.observer.on_stderr(line)
end
M.TestProvider = TestProvider
@@ -109,7 +113,7 @@ end
--- @param col number
--- @param lang string?
--- @return _99.test.Provider, number
-function M.fif_setup(content, row, col, lang)
+function M.test_setup(content, row, col, lang)
assert(lang, "lang must be provided")
local provider = M.TestProvider.new()
require("99").setup({
diff --git a/lua/99/test/visual_spec.lua b/lua/99/test/visual_spec.lua
index a567a73..2100539 100644
--- a/lua/99/test/visual_spec.lua
+++ b/lua/99/test/visual_spec.lua
@@ -51,7 +51,11 @@ describe("visual", function()
local context =
require("99.request-context").from_current_buffer(state, 100)
- visual_fn(context, range)
+ context.operation = "test_op"
+
+ visual_fn(context, range, {
+ additional_prompt = "test prompt",
+ })
eq(1, state:active_request_count())
eq(content, r(buffer))
@@ -82,7 +86,10 @@ describe("visual", function()
local context =
require("99.request-context").from_current_buffer(state, 200)
- visual_fn(context, range)
+ context.operation = "test_op"
+ visual_fn(context, range, {
+ additional_prompt = "test prompt",
+ })
eq(1, state:active_request_count())
eq(multi_line_content, r(buffer))
@@ -107,8 +114,11 @@ describe("visual", function()
local state = _99.__get_state()
local context =
require("99.request-context").from_current_buffer(state, 300)
+ context.operation = "test_op"
- visual_fn(context, range)
+ visual_fn(context, range, {
+ additional_prompt = "test prompt",
+ })
eq(content, r(buffer))
@@ -134,8 +144,11 @@ describe("visual", function()
local state = _99.__get_state()
local context =
require("99.request-context").from_current_buffer(state, 400)
+ context.operation = "test_op"
- visual_fn(context, range)
+ visual_fn(context, range, {
+ additional_prompt = "test prompt",
+ })
eq(content, r(buffer))
@@ -152,8 +165,11 @@ describe("visual", function()
local state = _99.__get_state()
local context =
require("99.request-context").from_current_buffer(state, 500)
+ context.operation = "test_op"
- visual_fn(context, range)
+ visual_fn(context, range, {
+ additional_prompt = "test prompt",
+ })
eq(content, r(buffer))
diff --git a/lua/99/utils.lua b/lua/99/utils.lua
index c308624..0882489 100644
--- a/lua/99/utils.lua
+++ b/lua/99/utils.lua
@@ -3,12 +3,9 @@ local M = {}
--- directories. if this is still the case in neovim land, then we will need
--- to make the _99_state have the project directory.
--- @return string
-function M.random_file()
- return string.format(
- "%s/tmp/99-%d",
- vim.uv.cwd(),
- math.floor(math.random() * 10000)
- )
+function M.random_file(dir)
+ local directory = dir or (vim.uv.cwd() .. "/tmp")
+ return string.format("%s/99-%d", directory, math.floor(math.random() * 10000))
end
return M
diff --git a/lua/99/window/init.lua b/lua/99/window/init.lua
index b4b1cb4..8db5b2f 100644
--- a/lua/99/window/init.lua
+++ b/lua/99/window/init.lua
@@ -27,6 +27,10 @@ local nvim_buf_is_valid = vim.api.nvim_buf_is_valid
--- @field win_id number
--- @field buf_id number
+--- @class _99.window.SplitWindow
+--- @field win number
+--- @field buffer number
+
--- @param lines string[]
--- @return string[]
local function ensure_no_new_lines(lines)
@@ -279,9 +283,13 @@ local function set_defaul_win_options(win, name)
end
--- @param win _99.window.Window
---- @param rules _99.Agents.Rules
+--- @param rules _99.Agents.Rules?
--- @param group any
local function highlight_rules_found(win, rules, group)
+ if rules == nil then
+ return
+ end
+
local rule_nsid = vim.api.nvim_create_namespace("99.window.rules")
local function check_and_highlight_rules()
if not nvim_win_is_valid(win.win_id) then
@@ -347,7 +355,8 @@ end
--- @class _99.window.CaptureInputOpts
--- @field cb fun(success: boolean, result: string): nil
--- @field on_load? fun(): nil
---- @field rules _99.Agents.Rules
+--- @field content? string[]
+--- @field rules? _99.Agents.Rules
--- @param name string
--- @param opts _99.window.CaptureInputOpts
@@ -423,6 +432,16 @@ function M.capture_input(name, opts)
if opts.on_load then
vim.schedule(opts.on_load)
end
+
+ if opts.content then
+ vim.api.nvim_buf_set_lines(
+ win.buf_id,
+ 0,
+ -1,
+ false,
+ ensure_no_new_lines(opts.content)
+ )
+ end
end
function M.clear_active_popups()
@@ -493,4 +512,54 @@ function M.close(win)
end
end
end
+
+--- @class _99.window.SplitWindowOpts
+--- @field split_direction "vertical" | "horizontal" | nil
+--- @field filetype string
+
+--- @param content string[]
+---@param buffer number | nil
+---@param opts _99.window.SplitWindowOpts | nil
+--- @return _99.window.SplitWindow
+function M.create_split(content, buffer, opts)
+ opts = opts or { split_direction = "vertical" }
+
+ opts.split_direction = opts.split_direction or "vertical"
+ opts.filetype = opts.filetype or "markdown"
+
+ local split_direction = opts.split_direction
+ assert(
+ split_direction == "vertical" or split_direction == "horizontal",
+ "unknown split direction: "
+ .. vim.inspect(split_direction)
+ .. " : must be horizontal or vertical"
+ )
+
+ if split_direction == "horizontal" then
+ vim.cmd("split")
+ else
+ vim.cmd("vsplit")
+ end
+
+ local win_id = vim.api.nvim_get_current_win()
+ local buf_id = buffer
+ if not buf_id or not nvim_buf_is_valid(buf_id) then
+ buf_id = vim.api.nvim_create_buf(false, false)
+ vim.api.nvim_buf_set_lines(
+ buf_id,
+ 0,
+ -1,
+ false,
+ ensure_no_new_lines(content)
+ )
+ end
+
+ vim.api.nvim_win_set_buf(win_id, buf_id)
+ vim.bo[buf_id].filetype = opts.filetype
+
+ return {
+ win = win_id,
+ buffer = buf_id,
+ }
+end
return M