diff --git a/local/ai-provider/lua/ai-provider/core.lua b/local/ai-provider/lua/ai-provider/core.lua index 67d8b41..47eee0e 100644 --- a/local/ai-provider/lua/ai-provider/core.lua +++ b/local/ai-provider/lua/ai-provider/core.lua @@ -1,4 +1,5 @@ local M = {} +local log = require 'ai-provider.log' local providers = { ollama = require 'ai-provider.providers.ollama', @@ -215,6 +216,7 @@ function M.chat(first, second) local provider, request = normalize_chat_args(first, second) local implementation = M.get_provider(provider) if not implementation or not implementation.chat then + log.error('chat requested unavailable provider: ' .. tostring(provider)) if request.callback then request.callback(nil, nil) end @@ -224,6 +226,7 @@ function M.chat(first, second) request.model = request.model or M.get_selected_model(provider) request.provider_config = get_provider_config(provider) if not request.model then + log.error('chat requested without selected model for provider: ' .. tostring(provider)) if request.callback then request.callback(nil, nil) end @@ -231,6 +234,17 @@ function M.chat(first, second) return nil end + log.info( + string.format( + 'chat start provider=%s model=%s prompt_chars=%d max_tokens=%s context_size=%s stream=%s', + provider, + request.model, + type(request.prompt) == 'string' and #request.prompt or 0, + tostring(request.max_tokens), + tostring(request.context_size), + tostring(request.stream ~= false) + ) + ) return implementation.chat(request) end @@ -334,6 +348,45 @@ function M.select_model(provider) end, { force = true }) end +function M.select_helper(opts, callback) + if type(opts) == 'function' then + callback = opts + opts = {} + end + opts = opts or {} + + collect_models(function(models) + if #models == 0 then + vim.notify('No AI provider models are available.', vim.log.levels.WARN) + if callback then + callback(nil) + end + return + end + + local current = opts.current + vim.ui.select(models, { + prompt = opts.prompt or 'Select AI provider model:', + format_item = function(item) + local selected = current and item.provider == current.provider and item.model == current.model + local marker = selected and '✓ ' or ' ' + return marker .. item.label + end, + }, function(choice) + if not choice then + if callback then + callback(nil) + end + return + end + + if callback then + callback({ provider = choice.provider, model = choice.model, label = choice.label }) + end + end) + end) +end + local global_actions = { 'default', 'model', 'models' } local provider_actions = { 'auth', 'check', 'model', 'models' } diff --git a/local/ai-provider/lua/ai-provider/init.lua b/local/ai-provider/lua/ai-provider/init.lua index 4a6b021..1b91665 100644 --- a/local/ai-provider/lua/ai-provider/init.lua +++ b/local/ai-provider/lua/ai-provider/init.lua @@ -9,11 +9,23 @@ ---@field timeout? integer Optional provider timeout in milliseconds. ---@field context_size? integer Optional provider-wide default context size. ---@field keep_alive? string|integer Optional provider-wide keep-alive/unload timeout. Ollama accepts values like `4h`, `10m`, or `0`. +---@field load_timeout? integer Optional model preload timeout in milliseconds. Used by Ollama before the normal chat timeout starts. +---@field think? boolean Optional Ollama thinking mode override for reasoning models. ---@field models? table Optional logical model profiles. Keys are selectable model names. +---@class AiProviderHelperConfig +---@field provider string Provider name used by the helper. +---@field model string Model/profile name returned by the helper. This is a reference to the configured model name, not copied options. +---@field label? string Display label, usually `provider/model`. + +---@class AiProviderSelectHelperOptions +---@field prompt? string Picker prompt. +---@field current? AiProviderHelperConfig Current selection used for picker highlighting. + ---@class AiProviderModelConfig ---@field model string Underlying provider model name. ---@field context_size? integer Optional model-specific context size override. +---@field think? boolean Optional Ollama thinking mode override for this logical model profile. ---@class AiProviderConfig ---@field default_provider string Default provider used when no persisted default provider exists. @@ -23,6 +35,23 @@ ---@field requested_model string Model requested by the caller. ---@field used_model string Model reported by the provider response. ---@field elapsed_ms number Request duration in milliseconds. +---@field done_reason? string Provider stop reason, if reported. +---@field error? string Provider/runtime error when `message` is nil. +---@field total_duration? integer Provider total duration in nanoseconds, when reported. +---@field load_duration? integer Initial model load duration in nanoseconds, when reported. +---@field prompt_eval_count? integer Prompt tokens evaluated, when reported. +---@field prompt_eval_duration? integer Prompt evaluation duration in nanoseconds, when reported. +---@field eval_count? integer Generated tokens evaluated, when reported. +---@field eval_duration? integer Generation duration in nanoseconds, when reported. + +---@class AiProviderStatus +---@field provider string Provider name. +---@field phase string Standard phase, for example `loading`, `loaded`, `thinking`, `generating`, `done`, or `error`. +---@field message string Human-readable status message. +---@field model? string Model/profile name. +---@field used_model? string Raw provider model name. +---@field tokens? integer Token count when the provider reports one. +---@field elapsed_ms? number Elapsed duration in milliseconds. ---@class AiProviderChatRequest ---@field provider? string Provider name. Defaults to `get_default_provider()`. @@ -32,7 +61,11 @@ ---@field max_tokens? integer Maximum generated tokens/provider equivalent. ---@field context_size? integer Per-request context size override. ---@field keep_alive? string|integer Per-request keep-alive/unload timeout override. +---@field load_timeout? integer Per-request model preload timeout in milliseconds. +---@field preload? boolean Whether to preload the model before chat. Defaults to false; callers such as ai-commit can enable it to exclude load time from chat timeout. +---@field think? boolean Per-request Ollama thinking mode override. ---@field on_chunk? fun(chunk: string, raw: table) Called for each streamed text chunk. +---@field on_status? fun(status: AiProviderStatus) Called with standardized provider progress updates. ---@field callback? fun(message: string|nil, meta: AiProviderChatMeta|nil) Called once the request finishes. ---@field is_cancelled? fun(): boolean Optional cancellation predicate. ---@field register_http_job? fun(job: table) Receives the provider job/process for external cancellation. @@ -166,6 +199,14 @@ function M.select_model(provider) return core.select_model(provider) end +---Open a model picker for feature-specific callers. This does not persist state; +---callers should save the returned provider/model reference themselves. +---@param opts? AiProviderSelectHelperOptions|fun(selection: AiProviderHelperConfig|nil) +---@param callback? fun(selection: AiProviderHelperConfig|nil) Called with the selected provider/model reference. +function M.select_helper(opts, callback) + return core.select_helper(opts, callback) +end + ---@param arglead string Current command-line argument prefix. ---@param cmdline string Full command-line. ---@return string[] completions Matching command completions. diff --git a/local/ai-provider/lua/ai-provider/providers/ollama.lua b/local/ai-provider/lua/ai-provider/providers/ollama.lua index cea9c83..274a4cc 100644 --- a/local/ai-provider/lua/ai-provider/providers/ollama.lua +++ b/local/ai-provider/lua/ai-provider/providers/ollama.lua @@ -1,8 +1,10 @@ local M = {} local curl = require 'ai-provider.curl' +local log = require 'ai-provider.log' local DEFAULT_ENDPOINT = 'http://127.0.0.1:11434' local HEALTH_CACHE_TTL = 30 +local DEFAULT_LOAD_TIMEOUT = 120000 local state = { health = nil, @@ -45,6 +47,17 @@ local function model_config(provider_config, model) return { model = model } end +local function elapsed_ms_since(started_at) + return (vim.uv.hrtime() - started_at) / 1e6 +end + +local function tokens_per_second(count, duration_ns) + if type(count) ~= 'number' or type(duration_ns) ~= 'number' or duration_ns <= 0 then + return nil + end + return count / (duration_ns / 1e9) +end + function M.check(callback, opts) opts = opts or {} local now = os.time() @@ -59,6 +72,7 @@ function M.check(callback, opts) callback = function(response) state.health = response.status == 200 state.health_checked_at = os.time() + log.debug('ollama check status=' .. tostring(response.status) .. ' working=' .. tostring(state.health)) callback(state.health) end, } @@ -109,7 +123,53 @@ 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 think = request.think + if think == nil then + think = selected_config.think + end + if think == nil then + think = provider_config.think + end local final_model = raw_model + local done_reason = nil + local provider_error = nil + local metrics = {} + local thinking_chars = 0 + local last_status_key = nil + + local function emit_status(status) + if not request.on_status then + return + end + + status.provider = 'ollama' + status.model = status.model or selected_model + status.used_model = status.used_model or final_model + status.elapsed_ms = status.elapsed_ms or elapsed_ms_since(started_at) + local key = table.concat({ status.phase or '', status.message or '', tostring(status.tokens), tostring(status.used_model) }, '|') + if key == last_status_key then + return + end + last_status_key = key + vim.schedule(function() + request.on_status(status) + end) + end + + log.info( + string.format( + 'ollama request selected_model=%s raw_model=%s prompt_chars=%d context_size=%s max_tokens=%s keep_alive=%s timeout=%s stream=%s think=%s', + tostring(selected_model), + tostring(raw_model), + type(request.prompt) == 'string' and #request.prompt or 0, + tostring(context_size), + tostring(request.max_tokens), + tostring(keep_alive), + tostring(request.timeout or 30000), + tostring(request.stream ~= false), + tostring(think) + ) + ) local body = { model = raw_model, @@ -121,46 +181,285 @@ function M.chat(request) }, keep_alive = keep_alive, } + if think ~= nil then + body.think = think + end - local job = curl.stream_json_lines { - url = endpoint() .. '/api/chat', - body = body, - timeout = request.timeout or 30000, - is_cancelled = request.is_cancelled, - on_json_line = function(data) - if is_cancelled(request) then - return - end + local function run_chat() + return curl.stream_json_lines { + url = endpoint() .. '/api/chat', + body = body, + timeout = request.timeout or 30000, + is_cancelled = request.is_cancelled, + on_json_line = function(data) + if is_cancelled(request) then + return + end - final_model = data.model or final_model - local chunk = data.message and data.message.content or '' - if chunk ~= '' then - table.insert(chunks, chunk) - if request.on_chunk then - vim.schedule(function() - request.on_chunk(chunk, data) - end) + final_model = data.model or final_model + done_reason = data.done_reason or done_reason + if type(data.error) == 'string' and data.error ~= '' then + provider_error = data.error + end + metrics.total_duration = data.total_duration or metrics.total_duration + metrics.load_duration = data.load_duration or metrics.load_duration + metrics.prompt_eval_count = data.prompt_eval_count or metrics.prompt_eval_count + metrics.prompt_eval_duration = data.prompt_eval_duration or metrics.prompt_eval_duration + metrics.eval_count = data.eval_count or metrics.eval_count + metrics.eval_duration = data.eval_duration or metrics.eval_duration + local thinking = data.message and data.message.thinking or '' + if thinking ~= '' then + thinking_chars = thinking_chars + #thinking + emit_status { + phase = 'thinking', + message = 'Thinking', + tokens = data.eval_count, + } + end + local chunk = data.message and data.message.content or '' + if chunk ~= '' then + table.insert(chunks, chunk) + emit_status { + phase = thinking_chars > 0 and 'generating' or 'generating', + message = 'Generating response', + tokens = data.eval_count, + } + if request.on_chunk then + vim.schedule(function() + request.on_chunk(chunk, data) + end) + end + end + end, + callback = function(code, error_message) + if is_cancelled(request) then + return + end + + if code ~= 0 then + log.error( + 'ollama request process failed code=' + .. tostring(code) + .. ' model=' + .. tostring(raw_model) + .. ' error=' + .. tostring(error_message) + ) + if request.callback then + request.callback(nil, { + requested_model = selected_model, + used_model = final_model, + elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6, + error = error_message or 'ollama request failed', + }) + end + emit_status { + phase = 'error', + message = error_message or 'Ollama request failed', + } + return + end + + local elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6 + local meta = vim.tbl_extend('force', { + requested_model = selected_model, + used_model = final_model, + elapsed_ms = elapsed_ms, + done_reason = done_reason, + }, metrics) + local message = table.concat(chunks, '') + local prompt_tokens_per_second = tokens_per_second(metrics.prompt_eval_count, metrics.prompt_eval_duration) + local eval_tokens_per_second = tokens_per_second(metrics.eval_count, metrics.eval_duration) + log.info( + string.format( + 'ollama response requested_model=%s used_model=%s done_reason=%s elapsed_ms=%.0f output_chars=%d context_size=%s max_tokens=%s load_ms=%s prompt_eval_count=%s prompt_eval_ms=%s prompt_tokens_per_second=%s eval_count=%s eval_ms=%s tokens_per_second=%s total_ms=%s', + tostring(selected_model), + tostring(final_model), + tostring(done_reason), + elapsed_ms, + #message, + tostring(context_size), + tostring(request.max_tokens), + metrics.load_duration and string.format('%.0f', metrics.load_duration / 1e6) or 'nil', + tostring(metrics.prompt_eval_count), + metrics.prompt_eval_duration and string.format('%.0f', metrics.prompt_eval_duration / 1e6) or 'nil', + prompt_tokens_per_second and string.format('%.2f', prompt_tokens_per_second) or 'nil', + tostring(metrics.eval_count), + metrics.eval_duration and string.format('%.0f', metrics.eval_duration / 1e6) or 'nil', + eval_tokens_per_second and string.format('%.2f', eval_tokens_per_second) or 'nil', + metrics.total_duration and string.format('%.0f', metrics.total_duration / 1e6) or 'nil' + ) + ) + if provider_error then + meta.error = provider_error + log.error('ollama provider error requested_model=' .. tostring(selected_model) .. ' error=' .. provider_error) + if request.callback then + request.callback(nil, meta) + end + emit_status { + phase = 'error', + message = provider_error, + tokens = metrics.eval_count, + } + return + end + if done_reason == 'length' then + meta.error = 'ollama stopped because the context or generation length limit was reached' + local prompt_tokens_per_second = tokens_per_second(metrics.prompt_eval_count, metrics.prompt_eval_duration) + local eval_tokens_per_second = tokens_per_second(metrics.eval_count, metrics.eval_duration) + log.error( + string.format( + 'ollama length stop requested_model=%s used_model=%s prompt_chars=%d output_chars=%d context_size=%s max_tokens=%s load_ms=%s prompt_eval_count=%s prompt_eval_ms=%s prompt_tokens_per_second=%s eval_count=%s eval_ms=%s tokens_per_second=%s', + tostring(selected_model), + tostring(final_model), + type(request.prompt) == 'string' and #request.prompt or 0, + #message, + tostring(context_size), + tostring(request.max_tokens), + metrics.load_duration and string.format('%.0f', metrics.load_duration / 1e6) or 'nil', + tostring(metrics.prompt_eval_count), + metrics.prompt_eval_duration and string.format('%.0f', metrics.prompt_eval_duration / 1e6) or 'nil', + prompt_tokens_per_second and string.format('%.2f', prompt_tokens_per_second) or 'nil', + tostring(metrics.eval_count), + metrics.eval_duration and string.format('%.0f', metrics.eval_duration / 1e6) or 'nil', + eval_tokens_per_second and string.format('%.2f', eval_tokens_per_second) or 'nil' + ) + ) + if request.callback then + request.callback(nil, meta) + end + emit_status { + phase = 'error', + message = meta.error, + tokens = metrics.eval_count, + } + return end - end - end, - callback = function(code) - if is_cancelled(request) then - return - end - if code ~= 0 then + if message == '' then + meta.error = 'ollama returned no content' + log.error('ollama returned no content requested_model=' .. tostring(selected_model) .. ' done_reason=' .. tostring(done_reason)) + if request.callback then + request.callback(nil, meta) + end + emit_status { + phase = 'error', + message = meta.error, + tokens = metrics.eval_count, + } + return + end + + emit_status { + phase = 'done', + message = 'Response complete', + tokens = metrics.eval_count, + } if request.callback then - request.callback(nil, nil) + request.callback(message, meta) end - return - end + end, + } + end - local elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6 - if request.callback then - request.callback(table.concat(chunks, ''), { requested_model = selected_model, used_model = final_model, elapsed_ms = elapsed_ms }) - end - end, - } + local job = nil + local load_timeout = request.load_timeout or provider_config.load_timeout or DEFAULT_LOAD_TIMEOUT + if request.preload == true or provider_config.preload == true then + emit_status { + phase = 'loading', + message = 'Loading model', + } + log.info( + string.format( + 'ollama preload start selected_model=%s raw_model=%s context_size=%s keep_alive=%s load_timeout=%s', + tostring(selected_model), + tostring(raw_model), + tostring(context_size), + tostring(keep_alive), + tostring(load_timeout) + ) + ) + local preload_body = { + model = raw_model, + messages = { { role = 'user', content = 'ok' } }, + stream = false, + keep_alive = keep_alive, + options = { + num_ctx = context_size, + num_predict = 16, + }, + } + if think ~= nil then + preload_body.think = think + end + + job = curl.json { + method = 'POST', + url = endpoint() .. '/api/chat', + timeout = load_timeout, + body = preload_body, + callback = function(response) + if is_cancelled(request) then + return + end + + if response.status ~= 200 then + local error_message = response.error or response.body or 'ollama preload failed' + log.error( + 'ollama preload failed status=' + .. tostring(response.status) + .. ' model=' + .. tostring(raw_model) + .. ' error=' + .. tostring(error_message) + ) + if request.callback then + request.callback(nil, { + requested_model = selected_model, + used_model = raw_model, + elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6, + error = error_message, + }) + end + emit_status { + phase = 'error', + message = error_message, + } + return + end + + local load_duration = type(response.json) == 'table' and response.json.load_duration or nil + emit_status { + phase = 'loaded', + message = 'Model loaded', + elapsed_ms = load_duration and (load_duration / 1e6) or elapsed_ms_since(started_at), + } + log.info( + string.format( + 'ollama preload complete selected_model=%s raw_model=%s status=%s load_ms=%s', + tostring(selected_model), + tostring(raw_model), + tostring(response.status), + load_duration and string.format('%.0f', load_duration / 1e6) or 'nil' + ) + ) + emit_status { + phase = 'generating', + message = 'Generating response', + } + job = run_chat() + if request.register_http_job then + request.register_http_job(job) + end + end, + } + else + emit_status { + phase = 'generating', + message = 'Generating response', + } + job = run_chat() + end 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 7cbe480..20d4588 100644 --- a/local/ai-provider/tests/ai_provider/core_spec.lua +++ b/local/ai-provider/tests/ai_provider/core_spec.lua @@ -26,4 +26,35 @@ describe('ai-provider core', function() assert.is_false(ok) assert.matches('default_provider must be configured', err) end) + + it('select helper returns a provider/model reference without saving it globally', function() + ai_provider.setup { + default_provider = 'ollama', + providers = { + ollama = { default_model = 'gemma4:e2b' }, + }, + } + + local core = require 'ai-provider.core' + local original_list_models = core.list_models + local original_select = vim.ui.select + local selected = nil + + rawset(core, 'list_models', function(provider, callback) + callback(provider == 'ollama' and { 'gemma4:e2b 64k' } or {}) + end) + rawset(vim.ui, 'select', function(items, _, on_choice) + on_choice(items[1]) + end) + + ai_provider.select_helper(function(choice) + selected = choice + end) + + rawset(core, 'list_models', original_list_models) + rawset(vim.ui, 'select', original_select) + + assert.are.same({ provider = 'ollama', model = 'gemma4:e2b 64k', label = 'ollama/gemma4:e2b 64k' }, selected) + assert.are.same('gemma4:e2b', ai_provider.get_selected_model 'ollama') + 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 eca1248..e4bbd81 100644 --- a/local/ai-provider/tests/ai_provider/ollama_options_spec.lua +++ b/local/ai-provider/tests/ai_provider/ollama_options_spec.lua @@ -30,6 +30,7 @@ describe('ollama provider options', function() ollama.chat { model = 'gemma4:e2b 64k', prompt = 'Reply with exactly: ok', + preload = false, max_tokens = 4, provider_config = { context_size = 1024 * 8, @@ -74,6 +75,7 @@ describe('ollama provider options', function() ollama.chat { model = 'gemma4:e2b 32k', prompt = 'Reply with exactly: ok', + preload = false, context_size = 1024 * 16, keep_alive = '10m', provider_config = { @@ -95,4 +97,107 @@ describe('ollama provider options', function() assert.are.same(1024 * 16, captured.options.num_ctx) assert.are.same('10m', captured.keep_alive) end) + + it('sends configured thinking mode for reasoning model profiles', function() + ---@type table|nil + local captured = nil + package.loaded['ai-provider.curl'] = { + stream_json_lines = function(request) + captured = request.body + request.on_json_line { model = 'qwen3.5:4b', message = { content = 'ok' }, done_reason = 'stop' } + request.callback(0) + return { shutdown = function() end } + end, + } + + local ollama = require 'ai-provider.providers.ollama' + + ollama.chat { + model = 'qwen3.5:4b 256k', + prompt = 'Reply with exactly: ok', + preload = false, + provider_config = { + models = { + ['qwen3.5:4b 256k'] = { + model = 'qwen3.5:4b', + context_size = 1024 * 256, + think = false, + }, + }, + }, + callback = function() end, + } + + assert.is_table(captured) + ---@cast captured table + assert.are.same('qwen3.5:4b', captured.model) + assert.are.same(1024 * 256, captured.options.num_ctx) + assert.are.same(false, captured.think) + end) + + it('returns an error instead of partial output when ollama stops for length', function() + package.loaded['ai-provider.curl'] = { + stream_json_lines = function(request) + request.on_json_line { model = 'gemma4:e2b', message = { content = 'Ref' }, done_reason = 'length' } + request.callback(0) + return { shutdown = function() end } + end, + } + + local ollama = require 'ai-provider.providers.ollama' + local message = 'unset' + local meta = nil + + ollama.chat { + model = 'gemma4:e2b', + prompt = 'large prompt', + preload = false, + callback = function(result, result_meta) + message = result + meta = result_meta + end, + } + + assert.is_nil(message) + assert.is_table(meta) + ---@cast meta table + assert.are.same('length', meta.done_reason) + assert.matches('length limit', meta.error) + end) + + it('emits standardized thinking and generating status events', function() + package.loaded['ai-provider.curl'] = { + stream_json_lines = function(request) + request.on_json_line { model = 'gemma4:e2b', message = { thinking = 'thinking...' }, eval_count = 7 } + request.on_json_line { model = 'gemma4:e2b', message = { content = 'ok' }, eval_count = 9 } + request.on_json_line { model = 'gemma4:e2b', done_reason = 'stop', eval_count = 9 } + request.callback(0) + return { shutdown = function() end } + end, + } + + local ollama = require 'ai-provider.providers.ollama' + local statuses = {} + + ollama.chat { + model = 'gemma4:e2b', + prompt = 'Reply with exactly: ok', + preload = false, + on_status = function(status) + table.insert(statuses, status) + end, + callback = function() end, + } + + vim.wait(1000, function() + return #statuses >= 4 + end, 10) + + assert.are.same('generating', statuses[1].phase) + assert.are.same('thinking', statuses[2].phase) + assert.are.same(7, statuses[2].tokens) + assert.are.same('generating', statuses[3].phase) + assert.are.same(9, statuses[3].tokens) + assert.are.same('done', statuses[4].phase) + end) end) diff --git a/local/ai-provider/tests/ai_provider/ollama_spec.lua b/local/ai-provider/tests/ai_provider/ollama_spec.lua index a23cd8d..e55fb86 100644 --- a/local/ai-provider/tests/ai_provider/ollama_spec.lua +++ b/local/ai-provider/tests/ai_provider/ollama_spec.lua @@ -28,6 +28,27 @@ local function wait_for(done, timeout) assert.is_true(ok) end +local function run_command(command) + local result = vim.system(command, { text = true }):wait() + assert.are.same(0, result.code) + return result.stdout or '' +end + +local function unload_ollama_model(model) + run_command { + 'curl', + '--silent', + '--show-error', + '--request', + 'POST', + '--header', + 'Content-Type: application/json', + '--data', + string.format('{"model":"%s","keep_alive":0}', model), + 'http://127.0.0.1:11434/api/generate', + } +end + describe('ollama provider integration', function() before_each(function() ai_provider.setup(config) @@ -78,7 +99,7 @@ describe('ollama provider integration', function() ai_provider.chat('ollama', { model = 'gemma4:e2b 32k', prompt = 'Reply with exactly: ok', - max_tokens = 4, + max_tokens = 256, timeout = 10000, callback = function(result, result_meta) message = result @@ -105,7 +126,7 @@ describe('ollama provider integration', function() ai_provider.chat('ollama', { model = 'gemma4:e2b', prompt = 'Reply with exactly: ok', - max_tokens = 4, + max_tokens = 16, timeout = 10000, on_chunk = function(chunk) table.insert(chunks, chunk) @@ -123,4 +144,41 @@ describe('ollama provider integration', function() assert.is_true(#chunks > 0) assert.are.same('ok', table.concat(chunks, '')) end) + + it('returns an error when the prompt exceeds a small loaded context window', function() + unload_ollama_model 'gemma4:e2b' + + local done = false + ---@type string|nil + local message = 'unset' + ---@type table|nil + local meta = nil + local prompt = table.concat(vim.fn['repeat']({ 'context-overflow-token' }, 5000), ' ') + + ai_provider.chat('ollama', { + model = 'gemma4:e2b', + prompt = prompt, + context_size = 2048, + max_tokens = 32, + timeout = 60000, + callback = function(result, result_meta) + message = result + meta = result_meta + done = true + end, + }) + + wait_for(function() + return done + end, 70000) + + local ps = run_command { 'ollama', 'ps' } + assert.matches('gemma4:e2b', ps) + assert.matches('%s2048%s', ps) + assert.is_nil(message) + assert.is_table(meta) + ---@cast meta table + assert.are.same('length', meta.done_reason) + assert.matches('length limit', meta.error) + end) end) diff --git a/lua/plugins/ai-commit.lua b/lua/plugins/ai-commit.lua index 4eae0d7..a6b931c 100644 --- a/lua/plugins/ai-commit.lua +++ b/lua/plugins/ai-commit.lua @@ -8,6 +8,9 @@ local CONFIG = { -- Use :AIProvider to configure the local provider; this is only the remote fallback model. model = nil, -- nil = auto (use Copilot's default) model_name = nil, -- Friendly display name for selected model + local_provider = 'ollama', + local_model = 'gemma4:e2b 64k', + local_model_name = 'gemma4:e2b 64k', openrouter = { endpoint = 'https://openrouter.ai/api/v1', @@ -116,11 +119,20 @@ local function load_preferences() if ok and type(prefs) == 'table' then CONFIG.provider = prefs.provider or 'copilot' CONFIG.model = prefs.model + CONFIG.local_provider = prefs.local_provider or CONFIG.local_provider + CONFIG.local_model = prefs.local_model or CONFIG.local_model if type(prefs.model_name) == 'string' and prefs.model_name ~= '' then CONFIG.model_name = prefs.model_name else CONFIG.model_name = nil end + if type(prefs.local_model_name) == 'string' and prefs.local_model_name ~= '' then + CONFIG.local_model_name = prefs.local_model_name + elseif type(CONFIG.local_model) == 'string' and CONFIG.local_model ~= '' then + CONFIG.local_model_name = CONFIG.local_model + else + CONFIG.local_model_name = nil + end end end end @@ -136,6 +148,9 @@ local function save_preferences() provider = CONFIG.provider, model = CONFIG.model, model_name = CONFIG.model_name, + local_provider = CONFIG.local_provider, + local_model = CONFIG.local_model, + local_model_name = CONFIG.local_model_name, } local file = io.open(prefs_file, 'w') if file then @@ -225,11 +240,12 @@ end local function complete_ollama(full_prompt, callback, status_callback, request_context) local log = setup_logger() local ai_provider = require 'ai-provider' - local model = ai_provider.get_selected_model 'ollama' + local provider = CONFIG.local_provider or 'ollama' + local model = CONFIG.local_model or ai_provider.get_selected_model(provider) if not model then - log.error 'Ollama is reachable but no model is selected' - vim.notify('No Ollama model selected. Run :AIProvider ollama model first.', vim.log.levels.ERROR) + log.error(provider .. ' is reachable but no model is selected') + vim.notify('No ' .. provider .. ' model selected. Run AI commit model selection first.', vim.log.levels.ERROR) callback(nil, nil) return end @@ -238,21 +254,60 @@ local function complete_ollama(full_prompt, callback, status_callback, request_c status_callback('Waiting for response from ' .. model) end - ai_provider.chat('ollama', { + local function report_provider_status(status) + if not status_callback or type(status) ~= 'table' then + return + end + + local status_model = status.model or model + if status.phase == 'loading' then + status_callback('Loading model ' .. status_model) + elseif status.phase == 'loaded' then + status_callback('Loaded model ' .. status_model) + elseif status.phase == 'thinking' then + local token_suffix = status.tokens and (' (' .. status.tokens .. ' tokens)') or '' + status_callback('Thinking with ' .. status_model .. token_suffix) + elseif status.phase == 'generating' then + local token_suffix = status.tokens and (' (' .. status.tokens .. ' tokens)') or '' + status_callback('Generating response with ' .. status_model .. token_suffix) + elseif status.phase == 'error' then + status_callback('Provider error from ' .. status_model) + end + end + + ai_provider.chat(provider, { model = model, prompt = full_prompt, stream = true, + preload = true, max_tokens = CONFIG.max_tokens, is_cancelled = request_context and request_context.is_cancelled, register_http_job = request_context and request_context.register_http_job, + on_status = report_provider_status, callback = function(message, meta) if request_context and request_context.is_cancelled and request_context.is_cancelled() then return end if not message then - log.error 'Ollama chat request failed' - vim.notify('Ollama request failed. See ai-commit logs.', vim.log.levels.ERROR) + local error_message = meta and meta.error or 'unknown error' + local details = '' + if meta then + details = string.format( + ' (requested_model=%s used_model=%s done_reason=%s elapsed_ms=%s load_ms=%s prompt_eval_count=%s prompt_eval_ms=%s eval_count=%s eval_ms=%s)', + tostring(meta.requested_model), + tostring(meta.used_model), + tostring(meta.done_reason), + tostring(meta.elapsed_ms), + meta.load_duration and string.format('%.0f', meta.load_duration / 1e6) or 'nil', + tostring(meta.prompt_eval_count), + meta.prompt_eval_duration and string.format('%.0f', meta.prompt_eval_duration / 1e6) or 'nil', + tostring(meta.eval_count), + meta.eval_duration and string.format('%.0f', meta.eval_duration / 1e6) or 'nil' + ) + end + log.error(provider .. ' chat request failed: ' .. error_message .. details) + vim.notify(provider .. ' request failed: ' .. error_message .. '. See ai-commit logs.', vim.log.levels.ERROR) callback(nil, nil) return end @@ -260,7 +315,7 @@ local function complete_ollama(full_prompt, callback, status_callback, request_c local cleaned = clean_message(message) local used_model = meta and meta.used_model or model local elapsed_ms = meta and meta.elapsed_ms or 0 - log.info(string.format('Generated commit message with Ollama (took %.0fms, model=%s)', elapsed_ms, used_model)) + log.info(string.format('Generated commit message with %s (took %.0fms, model=%s)', provider, elapsed_ms, used_model)) callback(cleaned, { requested_model = model, used_model = used_model }) end, }) @@ -804,17 +859,18 @@ local function generate_commit_message_async(branch, recent_commits, diff, callb log.debug(string.format('Prompt built (branch=%s, commits=%d chars, diff=%d chars)', branch, #recent_commits, #diff)) local ai_provider = require 'ai-provider' + local local_provider = CONFIG.local_provider or 'ollama' if status_callback then - status_callback 'Checking Ollama' + status_callback('Checking ' .. local_provider) end - ai_provider.check('ollama', function(working) + ai_provider.check(local_provider, function(working) if request_context and request_context.is_cancelled and request_context.is_cancelled() then return end if working then - log.info 'Routing commit message generation to Ollama' + log.info('Routing commit message generation to ' .. local_provider) complete_ollama(full_prompt, callback, status_callback, request_context) return end @@ -871,6 +927,26 @@ local function get_copilot_models() return fallback_models end +local function select_local_model() + local log = setup_logger() + local ai_provider = require 'ai-provider' + + ai_provider.select_helper({ + prompt = 'Select AI commit model:', + current = { provider = CONFIG.local_provider, model = CONFIG.local_model }, + }, function(choice) + if not choice then + return + end + + CONFIG.local_provider = choice.provider + CONFIG.local_model = choice.model + CONFIG.local_model_name = choice.model + log.info('AI commit local model changed to: ' .. choice.label) + save_preferences() + end) +end + --- Select and save a model from multiple providers local function select_model() local log = setup_logger() @@ -1320,6 +1396,16 @@ end return { 'nvim-lua/plenary.nvim', ft = 'gitcommit', + cmd = { 'AICommit', 'AICommitModel' }, + keys = { + { + 'psc', + function() + select_local_model() + end, + desc = 'AI [P]rovider [S]elect [C]ommit', + }, + }, config = function() -- Load saved model preference @@ -1349,5 +1435,9 @@ return { desc = 'Generate AI commit message', }) + vim.api.nvim_create_user_command('AICommitModel', select_local_model, { + desc = 'Select AI commit model', + }) + end, } diff --git a/lua/plugins/ai-provider.lua b/lua/plugins/ai-provider.lua index 4b3825e..9a96306 100644 --- a/lua/plugins/ai-provider.lua +++ b/lua/plugins/ai-provider.lua @@ -11,6 +11,7 @@ return { ollama = { default_model = 'gemma4:e2b', context_size = 1024 * 8, + load_timeout = 120000, keep_alive = '4h', models = { ['gemma4:e2b 32k'] = { @@ -21,6 +22,15 @@ return { model = 'gemma4:e2b', context_size = 1024 * 64, }, + ['qwen3.5:4b 128k'] = { + model = 'qwen3.5:4b', + context_size = 1024 * 128, + }, + ['qwen3.5:4b 256k'] = { + model = 'qwen3.5:4b', + context_size = 1024 * 256, + think = false, + }, }, }, },