From 4204f1dfb42b85e17df84f336895f67a050589e6 Mon Sep 17 00:00:00 2001 From: Meisterlala <6453306+Meisterlala@users.noreply.github.com> Date: Sun, 7 Jun 2026 19:19:13 +0200 Subject: [PATCH] feat(ai-commit): align streaming status with Ollama API and refine error reporting for timeout detection --- local/ai-commit/lua/ai-commit/config.lua | 3 +- local/ai-commit/lua/ai-commit/init.lua | 9 +++ local/ai-commit/lua/ai-commit/providers.lua | 61 +++++++++++++++++-- local/ai-provider/lua/ai-provider/core.lua | 56 ++++++++++------- local/ai-provider/lua/ai-provider/init.lua | 2 + .../lua/ai-provider/providers/ollama.lua | 19 +++--- .../tests/ai_provider/core_spec.lua | 41 +++++++++++++ .../tests/ai_provider/ollama_options_spec.lua | 28 ++++++++- lua/plugins/ai-provider.lua | 12 +++- 9 files changed, 194 insertions(+), 37 deletions(-) diff --git a/local/ai-commit/lua/ai-commit/config.lua b/local/ai-commit/lua/ai-commit/config.lua index f4ee05b..f4dbc52 100644 --- a/local/ai-commit/lua/ai-commit/config.lua +++ b/local/ai-commit/lua/ai-commit/config.lua @@ -2,6 +2,7 @@ local M = {} M.summary_source_id = 'ai-commit-summarize' M.message_source_id = 'ai-commit-message' +M.refine_source_id = 'ai-commit-refine' M.values = { context = { @@ -9,7 +10,7 @@ M.values = { recent_commits = true, staged_changes = true, }, - max_tokens = 10000, + max_tokens = 32768, spinner_interval = 80, preview_lines = 5, max_diff_chars = 100000, diff --git a/local/ai-commit/lua/ai-commit/init.lua b/local/ai-commit/lua/ai-commit/init.lua index 94cf296..7aa2b47 100644 --- a/local/ai-commit/lua/ai-commit/init.lua +++ b/local/ai-commit/lua/ai-commit/init.lua @@ -16,6 +16,7 @@ function M.setup(opts) local sources = { { id = config.summary_source_id, name = 'AI Commit: OpenCode Summary' }, { id = config.message_source_id, name = 'AI Commit: Commit Message' }, + { id = config.refine_source_id, name = 'AI Commit: Refinement' }, } for _, source in ipairs(sources) do ai_provider.register_source(source.id, { name = source.name }) @@ -53,6 +54,10 @@ function M.setup(opts) vim.api.nvim_create_user_command('AICommitSummaryModel', providers.select_summary_model, { desc = 'Select AI commit summary model', }) + + vim.api.nvim_create_user_command('AICommitRefinementModel', providers.select_refinement_model, { + desc = 'Select AI commit refinement model', + }) end function M.insert() @@ -63,4 +68,8 @@ function M.select_model() providers.select_model() end +function M.select_refinement_model() + providers.select_refinement_model() +end + return M diff --git a/local/ai-commit/lua/ai-commit/providers.lua b/local/ai-commit/lua/ai-commit/providers.lua index 5f148de..99c1897 100644 --- a/local/ai-commit/lua/ai-commit/providers.lua +++ b/local/ai-commit/lua/ai-commit/providers.lua @@ -77,6 +77,51 @@ local function provider_label(provider) return provider:sub(1, 1):upper() .. provider:sub(2) end +local function format_error_detail(meta) + if type(meta) ~= 'table' then + return 'unknown error' + end + + local parts = { tostring(meta.error or 'unknown error') } + if meta.done_reason then + table.insert(parts, 'done_reason=' .. tostring(meta.done_reason)) + end + if meta.requested_model then + table.insert(parts, 'requested_model=' .. tostring(meta.requested_model)) + end + if meta.used_model then + table.insert(parts, 'used_model=' .. tostring(meta.used_model)) + end + if meta.elapsed_ms then + table.insert(parts, string.format('elapsed=%.1fs', meta.elapsed_ms / 1000)) + end + if meta.load_duration then + table.insert(parts, string.format('load=%.1fs', meta.load_duration / 1e9)) + end + if meta.prompt_eval_count then + table.insert(parts, 'prompt_tokens=' .. tostring(meta.prompt_eval_count)) + end + if meta.eval_count then + table.insert(parts, 'eval_tokens=' .. tostring(meta.eval_count)) + end + + return table.concat(parts, ' | ') +end + +local function format_error_summary(provider, meta) + local error = type(meta) == 'table' and tostring(meta.error or '') or '' + if error:match 'timed out' or error:match 'Operation timed out' then + return provider .. ' request timed out while waiting for the model.' + end + if type(meta) == 'table' and meta.done_reason == 'length' then + return provider .. ' stopped because the output or context limit was reached.' + end + if error ~= '' and error ~= 'unknown error' then + return provider .. ' request failed: ' .. error + end + return provider .. ' request failed. See ai-commit logs for details.' +end + function M.select_model() local logger = log() local ai_provider = require 'ai-provider' @@ -91,6 +136,13 @@ function M.select_summary_model() ai_provider.select_source_model(config.summary_source_id) end +function M.select_refinement_model() + local logger = log() + local ai_provider = require 'ai-provider' + logger.info('Opening AI provider model picker for source=' .. config.refine_source_id) + ai_provider.select_source_model(config.refine_source_id) +end + ---@param full_prompt string ---@param callback function(string|nil, table|nil) ---@param status_callback function(string)|nil @@ -161,9 +213,10 @@ local function complete_ai_provider(source_id, full_prompt, callback, status_cal return end if not message then - logger.error(provider .. ' chat request failed: ' .. tostring(meta and meta.error or 'unknown error')) - vim.notify(provider .. ' request failed. See ai-commit logs.', vim.log.levels.ERROR) - callback(nil, nil) + local detail = format_error_detail(meta) + logger.error(provider .. ' chat request failed: ' .. detail) + vim.notify(format_error_summary(provider, meta), vim.log.levels.ERROR) + callback(nil, meta) return end callback(util.clean_message(message), { requested_model = model, used_model = meta and meta.used_model or model }) @@ -336,7 +389,7 @@ local function maybe_refine_message(context, message, iteration, callback, statu dump_prompt(prompt, 'refinement-' .. tostring(next_iteration), context.branch) local refinement_context = child_request_context(request_context, { - source_id = config.message_source_id, + source_id = config.refine_source_id, status_action = tostring(next_iteration) .. '. Refinement', }) M.complete_prompt(prompt, function(refined_message) diff --git a/local/ai-provider/lua/ai-provider/core.lua b/local/ai-provider/lua/ai-provider/core.lua index 10fa7e9..bb70e43 100644 --- a/local/ai-provider/lua/ai-provider/core.lua +++ b/local/ai-provider/lua/ai-provider/core.lua @@ -34,12 +34,13 @@ local function is_configured_provider(provider) end local function source_name(source_id) - local source = state.sources[source_id] + local prefs = M.load_preferences() + local source = prefs.sources and prefs.sources[source_id] if type(source) == 'table' and type(source.name) == 'string' and source.name ~= '' then return source.name end - local prefs = M.load_preferences() - source = prefs.sources and prefs.sources[source_id] + + source = state.sources[source_id] if type(source) == 'table' and type(source.name) == 'string' and source.name ~= '' then return source.name end @@ -47,18 +48,35 @@ local function source_name(source_id) end local function source_registered_name(source_id) - local source = state.sources[source_id] + local prefs = M.load_preferences() + local source = prefs.sources and prefs.sources[source_id] if type(source) == 'table' and type(source.name) == 'string' and source.name ~= '' then return source.name end - local prefs = M.load_preferences() - source = prefs.sources and prefs.sources[source_id] + + source = state.sources[source_id] if type(source) == 'table' and type(source.name) == 'string' and source.name ~= '' then return source.name end return nil end +local function registered_source_metadata(source_id) + local metadata = {} + local prefs = M.load_preferences() + local source = prefs.sources and prefs.sources[source_id] + if type(source) == 'table' and type(source.name) == 'string' and source.name ~= '' then + metadata.name = source.name + end + + source = state.sources[source_id] + if type(source) == 'table' and type(source.name) == 'string' and source.name ~= '' then + metadata.name = metadata.name or source.name + end + + return metadata +end + local function source_display(source_id) return source_name(source_id) end @@ -66,10 +84,7 @@ end local function has_source_model_preference(source_id) local prefs = M.load_preferences() local source = prefs.sources and prefs.sources[source_id] - return type(source) == 'table' - and is_configured_provider(source.provider) - and type(source.model) == 'string' - and source.model ~= '' + return type(source) == 'table' and is_configured_provider(source.provider) and type(source.model) == 'string' and source.model ~= '' end local function source_selection_display(source_id) @@ -199,11 +214,7 @@ function M.get_selected_model(provider, source_id) local prefs = M.load_preferences() if valid_source_id(source_id) then local source = prefs.sources and prefs.sources[source_id] - if type(source) == 'table' - and source.provider == provider - and type(source.model) == 'string' - and source.model ~= '' - then + if type(source) == 'table' and source.provider == provider and type(source.model) == 'string' and source.model ~= '' then return source.model end end @@ -260,10 +271,11 @@ function M.set_source_selection(source_id, provider, model) local prefs = M.load_preferences() prefs.sources = prefs.sources or {} local source = type(prefs.sources[source_id]) == 'table' and prefs.sources[source_id] or {} + local metadata = registered_source_metadata(source_id) source.provider = provider source.model = model source.label = nil - source.name = nil + source.name = source.name or metadata.name prefs.sources[source_id] = source return M.save_preferences(prefs) end @@ -275,7 +287,9 @@ function M.register_source(source_id, opts) opts = opts or {} state.sources[source_id] = state.sources[source_id] or {} - state.sources[source_id].name = opts.name + if type(opts.name) == 'string' and opts.name ~= '' then + state.sources[source_id].name = opts.name + end local prefs = M.load_preferences() prefs.sources = prefs.sources or {} @@ -293,7 +307,7 @@ function M.register_source(source_id, opts) local provider = opts.provider local model = opts.model if provider and model and is_configured_provider(provider) then - prefs.sources[source_id] = { provider = provider, model = model } + prefs.sources[source_id] = { provider = provider, model = model, name = source.name } if type(opts.name) == 'string' and opts.name ~= '' then prefs.sources[source_id].name = opts.name end @@ -408,9 +422,7 @@ function M.chat(first, second) request.model = request.model or M.get_selected_model(provider, request.source_id) request.provider_config = get_provider_config(provider) if not request.model then - log.error( - string.format('chat requested without selected model source=%s provider=%s', tostring(request.source_id), tostring(provider)) - ) + log.error(string.format('chat requested without selected model source=%s provider=%s', tostring(request.source_id), tostring(provider))) if request.callback then request.callback(nil, nil) end @@ -567,7 +579,7 @@ local function select_helper(opts, callback) end if callback then - callback({ provider = choice.provider, model = choice.model, label = choice.label }) + callback { provider = choice.provider, model = choice.model, label = choice.label } end end) end) diff --git a/local/ai-provider/lua/ai-provider/init.lua b/local/ai-provider/lua/ai-provider/init.lua index b9f925d..315da6c 100644 --- a/local/ai-provider/lua/ai-provider/init.lua +++ b/local/ai-provider/lua/ai-provider/init.lua @@ -23,6 +23,7 @@ ---@class AiProviderModelConfig ---@field model string Underlying provider model name. ---@field context_size? integer Optional model-specific context size override. +---@field timeout? integer Optional model-specific request timeout in milliseconds. ---@field think? boolean Optional Ollama thinking mode override for this logical model profile. ---@class AiProviderConfig @@ -60,6 +61,7 @@ ---@field stream? boolean Whether the provider should stream chunks. Defaults to true when supported. ---@field max_tokens? integer Maximum generated tokens/provider equivalent. ---@field context_size? integer Per-request context size override. +---@field timeout? integer Per-request timeout in milliseconds. ---@field keep_alive? string|integer Per-request keep-alive/unload timeout override. ---@field ps_timeout? integer Per-request loaded-model inspection timeout in milliseconds. ---@field think? boolean Per-request Ollama thinking mode override. diff --git a/local/ai-provider/lua/ai-provider/providers/ollama.lua b/local/ai-provider/lua/ai-provider/providers/ollama.lua index 09d3aa4..48fd785 100644 --- a/local/ai-provider/lua/ai-provider/providers/ollama.lua +++ b/local/ai-provider/lua/ai-provider/providers/ollama.lua @@ -167,6 +167,7 @@ function M.chat(request) local raw_model = selected_config.model or selected_model local context_size = request.context_size or selected_config.context_size or provider_config.context_size local keep_alive = request.keep_alive or provider_config.keep_alive + local timeout = request.timeout or selected_config.timeout or provider_config.timeout or 120000 local think = request.think if think == nil then think = selected_config.think @@ -246,7 +247,7 @@ function M.chat(request) tostring(context_size), tostring(request.max_tokens), tostring(keep_alive), - tostring(request.timeout or 30000), + tostring(timeout), tostring(request.stream ~= false), tostring(think) ) @@ -270,7 +271,7 @@ function M.chat(request) return curl.stream_json_lines { url = endpoint() .. '/api/chat', body = body, - timeout = request.timeout or 30000, + timeout = timeout, is_cancelled = request.is_cancelled, on_json_line = function(data) if is_cancelled(request) then @@ -463,17 +464,21 @@ function M.chat(request) phase = 'loaded', message = 'Model loaded', } + emit_status { + phase = 'context', + message = 'Loading prompt context', + } elseif model_loaded == false then emit_status { phase = 'loading', message = 'Loading model', } + else + emit_status { + phase = 'context', + message = 'Loading prompt context', + } end - - emit_status { - phase = 'context', - message = 'Loading prompt context', - } job = run_chat() if request.register_http_job then request.register_http_job(job) diff --git a/local/ai-provider/tests/ai_provider/core_spec.lua b/local/ai-provider/tests/ai_provider/core_spec.lua index 140aec3..832e1d8 100644 --- a/local/ai-provider/tests/ai_provider/core_spec.lua +++ b/local/ai-provider/tests/ai_provider/core_spec.lua @@ -97,4 +97,45 @@ describe('ai-provider core', function() rawset(core, 'load_preferences', original_load_preferences) rawset(core, 'save_preferences', original_save_preferences) end) + + it('persists registered source names across model selection updates', function() + ai_provider.setup { + default_provider = 'ollama', + providers = { + ollama = { default_model = 'gemma4:e2b' }, + }, + } + + local core = require 'ai-provider.core' + local original_load_preferences = core.load_preferences + local original_save_preferences = core.save_preferences + local prefs = {} + + rawset(core, 'load_preferences', function() + return vim.deepcopy(prefs) + end) + rawset(core, 'save_preferences', function(next_prefs) + prefs = vim.deepcopy(next_prefs) + return true + end) + + assert.is_true(ai_provider.register_source('ai-commit-refine', { name = 'AI Commit: Refinement' })) + assert.are.same('AI Commit: Refinement', ai_provider.get_source_name 'ai-commit-refine') + assert.is_true(ai_provider.set_source_selection('ai-commit-refine', 'ollama', 'gemma4:e2b 64k')) + + assert.are.same({ + name = 'AI Commit: Refinement', + provider = 'ollama', + model = 'gemma4:e2b 64k', + }, prefs.sources['ai-commit-refine']) + assert.are.same({ + provider = 'ollama', + model = 'gemma4:e2b 64k', + label = 'ollama/gemma4:e2b 64k', + name = 'AI Commit: Refinement', + }, ai_provider.get_source_selection 'ai-commit-refine') + + rawset(core, 'load_preferences', original_load_preferences) + rawset(core, 'save_preferences', original_save_preferences) + end) end) diff --git a/local/ai-provider/tests/ai_provider/ollama_options_spec.lua b/local/ai-provider/tests/ai_provider/ollama_options_spec.lua index e72a4c8..4783e92 100644 --- a/local/ai-provider/tests/ai_provider/ollama_options_spec.lua +++ b/local/ai-provider/tests/ai_provider/ollama_options_spec.lua @@ -96,6 +96,30 @@ describe('ollama provider options', function() assert.are.same('10m', captured.keep_alive) end) + it('uses configured request timeout for streaming chat', function() + local captured_timeout = nil + package.loaded['ai-provider.curl'] = { + stream_json_lines = function(request) + captured_timeout = request.timeout + request.callback(0) + return { shutdown = function() end } + end, + } + + local ollama = require 'ai-provider.providers.ollama' + + ollama.chat { + model = 'gemma4:e2b', + prompt = 'Reply with exactly: ok', + provider_config = { + timeout = 180000, + }, + callback = function() end, + } + + assert.are.same(180000, captured_timeout) + end) + it('sends configured thinking mode for reasoning model profiles', function() ---@type table|nil local captured = nil @@ -298,7 +322,7 @@ describe('ollama provider options', function() end, 10) assert.are.same('loading', statuses[1]) - assert.are.same('context', statuses[2]) - assert.are.same('generating', statuses[3]) + assert.are.same('generating', statuses[2]) + assert.is_false(vim.tbl_contains(statuses, 'context')) end) end) diff --git a/lua/plugins/ai-provider.lua b/lua/plugins/ai-provider.lua index 9cbdc54..b97484e 100644 --- a/lua/plugins/ai-provider.lua +++ b/lua/plugins/ai-provider.lua @@ -14,17 +14,27 @@ return { ollama = { default_model = 'gemma4:e4b', context_size = 1024 * 8, - load_timeout = 120000, + timeout = 180000, keep_alive = '1h', models = { ['gemma4:e2b 32k'] = { model = 'gemma4:e2b', context_size = 1024 * 32, }, + ['gemma4:e2b 32k fast'] = { + model = 'gemma4:e2b', + context_size = 1024 * 32, + think = false, + }, ['gemma4:e2b 64k'] = { model = 'gemma4:e2b', context_size = 1024 * 64, }, + ['gemma4:e2b 64k fast'] = { + model = 'gemma4:e2b', + context_size = 1024 * 64, + think = false, + }, ['qwen3.5:4b 64k'] = { model = 'qwen3.5:4b', context_size = 1024 * 64, -- 2.51.2