summaryrefslogtreecommitdiff
path: root/lua
diff options
context:
space:
mode:
authortheprimeagain <the.primeagen@gmail.com>2026-02-15 15:28:36 -0700
committertheprimeagain <the.primeagen@gmail.com>2026-02-15 15:28:36 -0700
commit3ade8bbbbcc4fe80a09f22eb59ce8539e4fe2987 (patch)
treeb9bfdb5bd07aa1786de4af174a1ce270f43d557d /lua
parent8db9030ddcfcd227cd265397a3fa4aa521b0f2de (diff)
downloada4-3ade8bbbbcc4fe80a09f22eb59ce8539e4fe2987.tar.xz
a4-3ade8bbbbcc4fe80a09f22eb59ce8539e4fe2987.zip
the final parts of tutorial
Diffstat (limited to 'lua')
-rw-r--r--lua/99/extensions/agents/init.lua1
-rw-r--r--lua/99/init.lua92
-rw-r--r--lua/99/ops/clean-up.lua56
-rw-r--r--lua/99/ops/implement-fn.lua92
-rw-r--r--lua/99/ops/init.lua1
-rw-r--r--lua/99/ops/make-prompt.lua10
-rw-r--r--lua/99/ops/over-range.lua18
-rw-r--r--lua/99/ops/search.lua46
-rw-r--r--lua/99/ops/tutorial.lua48
-rw-r--r--lua/99/providers.lua3
-rw-r--r--lua/99/request-context.lua36
-rw-r--r--lua/99/request/init.lua9
-rw-r--r--lua/99/test/request_spec.lua3
13 files changed, 147 insertions, 268 deletions
diff --git a/lua/99/extensions/agents/init.lua b/lua/99/extensions/agents/init.lua
index 4002c75..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
diff --git a/lua/99/init.lua b/lua/99/init.lua
index 6f6010f..43daed1 100644
--- a/lua/99/init.lua
+++ b/lua/99/init.lua
@@ -54,17 +54,18 @@ end
--- @class _99.RequestEntry.Data.Visual
--- @field type "visual"
+---
--- @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 status _99.Request.State
--- @field filename string
--- @field lnum number
--- @field col number
--- @field started_at number
---- @field operation_data _99.RequestEntry.Data
+--- @field operation_data _99.RequestEntry.Data | nil
--- @class _99.ActiveRequest
--- @field clean_up _99.Cleanup
@@ -81,9 +82,9 @@ 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 __tutorials _99.RequestEntry.Data.Tutorial[]
+--- @field __searches _99.RequestEntry.Data.Search[]
--- @field __request_history _99.RequestEntry[]
--- @field __request_by_id table<number, _99.RequestEntry>
@@ -99,8 +100,8 @@ local function create_99_state()
display_errors = false,
provider_override = nil,
auto_add_skills = false,
- __active_requests = {},
__tutorials = {},
+ __searches = {},
__view_log_idx = 1,
__request_history = {},
__request_by_id = {},
@@ -139,10 +140,10 @@ 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 __tutorials _99.RequestEntry.Data.Tutorial[]
+--- @field __searches _99.RequestEntry.Data.Search[]
--- @field __request_by_id table<number, _99.RequestEntry>
--- @field __active_marks _99.Mark[]
local _99_State = {}
@@ -170,30 +171,42 @@ function _99_State:refresh_rules()
end
--- @param context _99.RequestContext
+--- @param clean_up fun(): nil
--- @return _99.RequestEntry
-function _99_State:track_request(context)
+function _99_State:track_request(context, clean_up)
local point = context.range and context.range.start or Point:from_cursor()
local entry = {
id = context.xid,
+ clean_up = clean_up,
operation = context.operation or "request",
status = "running",
filename = context.full_path,
lnum = point.row,
col = point.col,
started_at = time.now(),
+ operation_data = nil,
}
table.insert(self.__request_history, entry)
self.__request_by_id[entry.id] = entry
return entry
end
---- @param id number
+--- @param context _99.RequestContext
--- @param status "success" | "failed" | "cancelled"
-function _99_State:finish_request(id, status)
+function _99_State:finish_request(context, status)
+ local id = context.xid
local entry = self.__request_by_id[id]
if entry then
entry.status = status
end
+ if entry.operation == "success" and entry.operation_data then
+ local data = entry.operation_data
+ if data.type == "tutorial" then
+ table.insert(self.__tutorials, data)
+ elseif data.type == "search" then
+ table.insert(self.__searches, data)
+ end
+ end
end
--- @param id number
@@ -211,17 +224,6 @@ function _99_State:add_data(id, data)
entry.operation_data = data
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
- end
- self.__request_by_id[id] = nil
-end
-
--- @return number
function _99_State:previous_request_count()
local count = 0
@@ -243,22 +245,8 @@ function _99_State:clear_previous_requests()
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
+ self.__searches = {}
+ self.__tutorials = {}
end
--- @param mark _99.Mark
@@ -268,31 +256,14 @@ 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 == "running" 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
-
- local entry = self.__request_history[id]
- if entry.operation == "tutorial" and entry.status == "success" then
- local data = entry.operation_data
- logger:assert(
- data.type == "tutorial",
- "data type tutorial expected for request tutorial"
- )
- table.insert(self.__tutorials, data)
- end
-end
-
local _99_state = _99_State.new()
--- @class _99
@@ -464,11 +435,7 @@ function _99.next_request_logs()
end
function _99.stop_all_requests()
- for _, active in pairs(_99_state.__active_requests) do
- _99_state:remove_request(active.request_id)
- active.clean_up()
- end
- _99_state.__active_requests = {}
+ error("implement")
end
function _99.clear_all_marks()
@@ -542,9 +509,6 @@ local function show_in_flight_requests()
local lines = {
throb .. " requests(" .. tostring(count) .. ") " .. throb,
}
- for _, r in pairs(_99_state.__active_requests) do
- table.insert(lines, r.name)
- end
Window.resize(win, #lines[1], #lines)
vim.api.nvim_buf_set_lines(win.buf_id, 0, 1, false, lines)
diff --git a/lua/99/ops/clean-up.lua b/lua/99/ops/clean-up.lua
index ba76f01..2ac099c 100644
--- a/lua/99/ops/clean-up.lua
+++ b/lua/99/ops/clean-up.lua
@@ -1,20 +1,60 @@
----@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 context _99.RequestContext
+--- @param clean_up fun(): nil
+--- @param obs_or_fn _99.Providers.PartialObserver | _99.Providers.on_complete
+--- @return _99.Providers.Observer
+M.make_observer = function(context, 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()
+ context._99:track_request(context, clean_up)
+ if obs.on_start then
+ obs.on_start()
+ end
+ end,
+ on_complete = function(status, res)
+ vim.schedule(clean_up)
+ context._99:finish_request(context, status)
+ 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 247c903..1f8e9f8 100644
--- a/lua/99/ops/init.lua
+++ b/lua/99/ops/init.lua
@@ -9,6 +9,5 @@
return {
search = require("99.ops.search"),
tutorial = require("99.ops.tutorial"),
- implement_fn = require("99.ops.implement-fn"),
over_range = require("99.ops.over-range"),
}
diff --git a/lua/99/ops/make-prompt.lua b/lua/99/ops/make-prompt.lua
index 8b77256..4a17687 100644
--- a/lua/99/ops/make-prompt.lua
+++ b/lua/99/ops/make-prompt.lua
@@ -1,9 +1,10 @@
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.Completion
+--- @return string, _99.Reference[]
return function(context, prompt, opts)
local user_prompt = opts.additional_prompt
assert(
@@ -18,7 +19,12 @@ return function(context, prompt, opts)
local additional_rules = opts.additional_rules
if additional_rules then
for _, r in ipairs(additional_rules) do
- table.insert(refs, r)
+ local content = Agents.get_rule_content(r)
+ if content then
+ table.insert(refs, {
+ content = content,
+ })
+ end
end
end
diff --git a/lua/99/ops/over-range.lua b/lua/99/ops/over-range.lua
index ba4a7d3..a37aaf8 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 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,7 +40,7 @@ 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()
@@ -45,17 +48,15 @@ local function over_range(context, range, opts)
end)
local system_cmd = context._99.prompts.prompts.visual_selection(range)
- local prompt, rules, refs = make_prompt(context, system_cmd, opts)
+ local prompt, refs = make_prompt(context, system_cmd, opts)
- context:add_agent_rules(rules)
request:add_prompt_content(prompt)
context:add_references(refs)
top_status:start()
bottom_status:start()
- request:start({
+ request:start(make_observer(context, 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
@@ -90,10 +91,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 cc6752d..cde8f64 100644
--- a/lua/99/ops/search.lua
+++ b/lua/99/ops/search.lua
@@ -1,6 +1,9 @@
local Request = require("99.request")
-local make_clean_up = require("99.ops.clean-up")
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
@@ -65,39 +68,28 @@ local function search(context, opts)
logger:debug("search", "with opts", opts.additional_prompt)
- local clean_up = make_clean_up(context, "Search", function()
+ local clean_up = make_clean_up(function()
request:cancel()
end)
- local prompt, rules, refs =
+ local prompt, refs =
make_prompt(context, context._99.prompts.prompts.semantic_search(), opts)
- context:add_agent_rules(rules)
request:add_prompt_content(prompt)
context:add_references(refs)
- 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(context, 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._99, response)
+ end
+ end))
end
return search
diff --git a/lua/99/ops/tutorial.lua b/lua/99/ops/tutorial.lua
index 54e7804..bba56e6 100644
--- a/lua/99/ops/tutorial.lua
+++ b/lua/99/ops/tutorial.lua
@@ -1,7 +1,9 @@
local Request = require("99.request")
-local make_clean_up = require("99.ops.clean-up")
+local CleanUp = require("99.ops.clean-up")
local make_prompt = require("99.ops.make-prompt")
+local make_clean_up = CleanUp.make_clean_up
+local make_observer = CleanUp.make_observer
--- @class _99.Tutorial.Result
--- @param context _99.RequestContext
@@ -14,37 +16,29 @@ local function tutorial(context, opts)
local request = Request.new(context)
- local clean_up = make_clean_up(context, "Search", function()
+ local clean_up = make_clean_up(function()
request:cancel()
end)
- local prompt, rules =
+ local prompt, refs =
make_prompt(context, context._99.prompts.prompts.tutorial(), opts)
- context:add_agent_rules(rules)
+ context:add_references(refs)
request:add_prompt_content(prompt)
- request:start({
- on_complete = 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
- error("what the hell")
- 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("on_stderr", "line", line)
- end,
- })
+ request:start(make_observer(context, 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
+ error("what the hell")
+ end
+ end))
+
end
return tutorial
diff --git a/lua/99/providers.lua b/lua/99/providers.lua
index ff1188b..1b5c4d9 100644
--- a/lua/99/providers.lua
+++ b/lua/99/providers.lua
@@ -2,6 +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
+--- @field on_start fun(): nil
--- @param fn fun(...: any): nil
--- @return fun(...: any): nil
@@ -50,6 +51,8 @@ end
--- @param request _99.Request
--- @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)
diff --git a/lua/99/request-context.lua b/lua/99/request-context.lua
index 6671cab..ef10aff 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
@@ -38,6 +39,7 @@ function RequestContext.from_current_buffer(_99, xid)
return setmetatable({
_99 = _99,
+ clean_ups = {},
md_file_names = mds,
ai_context = {},
tmp_file = random_file(),
@@ -58,40 +60,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
diff --git a/lua/99/request/init.lua b/lua/99/request/init.lua
index 327aa0f..3cff512 100644
--- a/lua/99/request/init.lua
+++ b/lua/99/request/init.lua
@@ -81,8 +81,10 @@ end
--- @param obs _99.Providers.Observer | nil
local function observer_from_request(r, obs)
return {
+ on_start = obs and obs.on_start or function() end,
on_complete = function(status, res)
r.state = status
+ r.context._99:finish_request(r.context, status)
if obs then
obs.on_complete(status, res)
end
@@ -108,7 +110,6 @@ function Request:start(observer)
)
self.state = "requesting"
- self.context._99:track_request(self.context)
self.context:finalize()
for _, content in ipairs(self.context.ai_context) do
self:add_prompt_content(content)
@@ -117,7 +118,11 @@ function Request:start(observer)
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_from_request(self, observer))
+ self.provider:make_request(
+ prompt,
+ self,
+ observer_from_request(self, observer)
+ )
end
return Request
diff --git a/lua/99/test/request_spec.lua b/lua/99/test/request_spec.lua
index d2e0885..f6ce34d 100644
--- a/lua/99/test/request_spec.lua
+++ b/lua/99/test/request_spec.lua
@@ -26,6 +26,7 @@ describe("request test", function()
eq("ready", request.state)
+ eq(0, state:active_request_count())
request:start({
on_complete = function(status, _)
finished_called = true
@@ -34,12 +35,14 @@ describe("request test", function()
on_stdout = function() end,
on_stderr = function() end,
})
+ 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)