diff options
| author | Bogdan Nikolov <84232456+0xr3ngar@users.noreply.github.com> | 2026-02-19 20:30:08 +0100 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2026-02-19 20:30:08 +0100 |
| commit | 725bf38ffc08cb8b537cbff322ef2d889b197dad (patch) | |
| tree | 2fa53c6b40463b8a18d9e175617b9b413a844e70 /lua/99/test | |
| parent | 6d219f47aa29cafb1ec39c41d4dd0a4875a74506 (diff) | |
| parent | d34e4dcbcc1f95d8f31f851baad18da12c59e7df (diff) | |
| download | a4-725bf38ffc08cb8b537cbff322ef2d889b197dad.tar.xz a4-725bf38ffc08cb8b537cbff322ef2d889b197dad.zip | |
Merge branch 'master' into master
Diffstat (limited to 'lua/99/test')
| -rw-r--r-- | lua/99/test/marks_spec.lua | 2 | ||||
| -rw-r--r-- | lua/99/test/providers_spec.lua | 33 | ||||
| -rw-r--r-- | lua/99/test/request_spec.lua | 54 | ||||
| -rw-r--r-- | lua/99/test/test_utils.lua | 40 | ||||
| -rw-r--r-- | lua/99/test/visual_spec.lua | 26 |
5 files changed, 131 insertions, 24 deletions
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)) |
