summaryrefslogtreecommitdiff
path: root/lua/99/test
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/99/test
parent6d219f47aa29cafb1ec39c41d4dd0a4875a74506 (diff)
parentd34e4dcbcc1f95d8f31f851baad18da12c59e7df (diff)
downloada4-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.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
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))