From 66aea6b4ae307d3162e6aa1bcf5bc420b1382275 Mon Sep 17 00:00:00 2001 From: theprimeagain Date: Sat, 14 Feb 2026 09:12:21 -0700 Subject: working through this idea of tutorials... --- lua/99/test/request_spec.lua | 49 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 49 insertions(+) create mode 100644 lua/99/test/request_spec.lua (limited to 'lua/99/test/request_spec.lua') diff --git a/lua/99/test/request_spec.lua b/lua/99/test/request_spec.lua new file mode 100644 index 0000000..5ad114a --- /dev/null +++ b/lua/99/test/request_spec.lua @@ -0,0 +1,49 @@ +-- 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:finalize() + + local request = Request.new(context) + + local finished_called = false + local finished_status = nil + + eq("ready", request.state) + + request:start({ + on_complete = function(status, _) + finished_called = true + finished_status = status + end, + on_stdout = function() end, + on_stderr = function() end, + }) + + eq("requesting", request.state) + + p:resolve("success", " return 'implemented!'") + test_utils.next_frame() + + assert.is_true(finished_called) + test_utils.next_frame() + + eq("success", request.state) + eq("success", finished_status) + end) +end) -- cgit v1.3-3-g829e From a65a39fc96b004d57a80691e090f8b8e5d6da508 Mon Sep 17 00:00:00 2001 From: theprimeagain Date: Sat, 14 Feb 2026 09:10:13 -0700 Subject: fixed the requesting logic to better use request for state management --- lua/99/providers.lua | 12 +----------- lua/99/request/init.lua | 27 +++++++++++++++++++++++++-- lua/99/test/request_spec.lua | 3 --- lua/99/test/test_utils.lua | 4 +++- 4 files changed, 29 insertions(+), 17 deletions(-) (limited to 'lua/99/test/request_spec.lua') diff --git a/lua/99/providers.lua b/lua/99/providers.lua index 080a909..ff1188b 100644 --- a/lua/99/providers.lua +++ b/lua/99/providers.lua @@ -3,14 +3,6 @@ --- @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, -} - --- @param fn fun(...: any): nil --- @return fun(...: any): nil local function once(fn) @@ -56,17 +48,15 @@ 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) 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( --- @param status "success" | "failed" | "cancelled" ---@param text string function(status, text) - print("setting status", status) request.state = status observer.on_complete(status, text) end diff --git a/lua/99/request/init.lua b/lua/99/request/init.lua index e0e86dd..327aa0f 100644 --- a/lua/99/request/init.lua +++ b/lua/99/request/init.lua @@ -47,7 +47,7 @@ function Request:_set_process(proc) end function Request:cancel() - if self.state == "finished" then + if self.state == "success" then return end @@ -77,6 +77,29 @@ 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) + return { + on_complete = function(status, res) + r.state = 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.logger:assert( @@ -94,7 +117,7 @@ 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) + 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 5ad114a..d2e0885 100644 --- a/lua/99/test/request_spec.lua +++ b/lua/99/test/request_spec.lua @@ -38,10 +38,7 @@ describe("request test", function() eq("requesting", request.state) p:resolve("success", " return 'implemented!'") - test_utils.next_frame() - assert.is_true(finished_called) - test_utils.next_frame() eq("success", request.state) eq("success", finished_status) diff --git a/lua/99/test/test_utils.lua b/lua/99/test/test_utils.lua index 29d12d0..8f417d6 100644 --- a/lua/99/test/test_utils.lua +++ b/lua/99/test/test_utils.lua @@ -23,7 +23,8 @@ M.created_files = {} --- @class _99.test.Provider : _99.Providers.BaseProvider --- @field request _99.test.ProviderRequest? -local TestProvider = setmetatable({}, { __index = BaseProvider }) +local TestProvider = {} +TestProvider.__index = TestProvider function TestProvider.new() return setmetatable({}, TestProvider) @@ -48,6 +49,7 @@ end 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 -- cgit v1.3-3-g829e From 3ade8bbbbcc4fe80a09f22eb59ce8539e4fe2987 Mon Sep 17 00:00:00 2001 From: theprimeagain Date: Sun, 15 Feb 2026 15:28:36 -0700 Subject: the final parts of tutorial --- lua/99/extensions/agents/init.lua | 1 - lua/99/init.lua | 92 ++++++++++++--------------------------- lua/99/ops/clean-up.lua | 56 ++++++++++++++++++++---- lua/99/ops/implement-fn.lua | 92 --------------------------------------- lua/99/ops/init.lua | 1 - lua/99/ops/make-prompt.lua | 10 ++++- lua/99/ops/over-range.lua | 18 ++++---- lua/99/ops/search.lua | 46 ++++++++------------ lua/99/ops/tutorial.lua | 48 +++++++++----------- lua/99/providers.lua | 3 ++ lua/99/request-context.lua | 36 +-------------- lua/99/request/init.lua | 9 +++- lua/99/test/request_spec.lua | 3 ++ 13 files changed, 147 insertions(+), 268 deletions(-) delete mode 100644 lua/99/ops/implement-fn.lua (limited to 'lua/99/test/request_spec.lua') 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 --- @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 @@ -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 --- @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 --- @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 -]], - 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) -- cgit v1.3-3-g829e From ab095297ecfe6dea1c035ac08f6f6b4eed2c62b5 Mon Sep 17 00:00:00 2001 From: theprimeagain Date: Wed, 18 Feb 2026 07:07:13 -0700 Subject: fixed request spec --- lua/99/test/request_spec.lua | 5 +++++ lua/99/test/test_utils.lua | 40 +++++++++++++++++++++------------------- 2 files changed, 26 insertions(+), 19 deletions(-) (limited to 'lua/99/test/request_spec.lua') diff --git a/lua/99/test/request_spec.lua b/lua/99/test/request_spec.lua index f6ce34d..3c58ecd 100644 --- a/lua/99/test/request_spec.lua +++ b/lua/99/test/request_spec.lua @@ -17,6 +17,7 @@ describe("request test", function() 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) @@ -28,6 +29,9 @@ describe("request test", function() 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 @@ -35,6 +39,7 @@ describe("request test", function() on_stdout = function() end, on_stderr = function() end, }) + test_utils.next_frame() eq(1, state:active_request_count()) eq("requesting", request.state) diff --git a/lua/99/test/test_utils.lua b/lua/99/test/test_utils.lua index 8f417d6..a1459a4 100644 --- a/lua/99/test/test_utils.lua +++ b/lua/99/test/test_utils.lua @@ -1,7 +1,14 @@ -local BaseProvider = require("99.providers").BaseProvider 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() @@ -18,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 @@ -36,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, @@ -48,35 +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 -- cgit v1.3-3-g829e