diff --git a/local/ai-provider/lua/ai-provider/init.lua b/local/ai-provider/lua/ai-provider/init.lua index 1b91665..7cb9e95 100644 --- a/local/ai-provider/lua/ai-provider/init.lua +++ b/local/ai-provider/lua/ai-provider/init.lua @@ -51,6 +51,7 @@ ---@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 tokens_per_second? number Generation throughput when the provider reports enough timing data. ---@field elapsed_ms? number Elapsed duration in milliseconds. ---@class AiProviderChatRequest diff --git a/local/ai-provider/lua/ai-provider/providers/ollama.lua b/local/ai-provider/lua/ai-provider/providers/ollama.lua index b29357a..3be98f8 100644 --- a/local/ai-provider/lua/ai-provider/providers/ollama.lua +++ b/local/ai-provider/lua/ai-provider/providers/ollama.lua @@ -136,6 +136,7 @@ function M.chat(request) local metrics = {} local thinking_chars = 0 local last_status_key = nil + local generation_started_at = nil local function emit_status(status) if not request.on_status then @@ -146,7 +147,7 @@ function M.chat(request) 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) }, '|') + local key = table.concat({ status.phase or '', status.message or '', tostring(status.tokens_per_second), tostring(status.used_model) }, '|') if key == last_status_key then return end @@ -207,13 +208,20 @@ function M.chat(request) 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 + if not generation_started_at and data.eval_count then + generation_started_at = vim.uv.hrtime() + end + local status_tokens_per_second = tokens_per_second(data.eval_count, data.eval_duration) + if not status_tokens_per_second and generation_started_at and data.eval_count and data.eval_count > 0 then + status_tokens_per_second = data.eval_count / ((vim.uv.hrtime() - generation_started_at) / 1e9) + end 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, + tokens_per_second = status_tokens_per_second, } end local chunk = data.message and data.message.content or '' @@ -222,7 +230,7 @@ function M.chat(request) emit_status { phase = thinking_chars > 0 and 'generating' or 'generating', message = 'Generating response', - tokens = data.eval_count, + tokens_per_second = status_tokens_per_second, } if request.on_chunk then vim.schedule(function() 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 e4bbd81..4aff6a6 100644 --- a/local/ai-provider/tests/ai_provider/ollama_options_spec.lua +++ b/local/ai-provider/tests/ai_provider/ollama_options_spec.lua @@ -168,9 +168,9 @@ describe('ollama provider options', function() 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.on_json_line { model = 'gemma4:e2b', message = { thinking = 'thinking...' }, eval_count = 7, eval_duration = 1000000000 } + request.on_json_line { model = 'gemma4:e2b', message = { content = 'ok' }, eval_count = 9, eval_duration = 1000000000 } + request.on_json_line { model = 'gemma4:e2b', done_reason = 'stop', eval_count = 9, eval_duration = 1000000000 } request.callback(0) return { shutdown = function() end } end, @@ -195,9 +195,9 @@ describe('ollama provider options', function() assert.are.same('generating', statuses[1].phase) assert.are.same('thinking', statuses[2].phase) - assert.are.same(7, statuses[2].tokens) + assert.are.same(7, statuses[2].tokens_per_second) assert.are.same('generating', statuses[3].phase) - assert.are.same(9, statuses[3].tokens) + assert.are.same(9, statuses[3].tokens_per_second) assert.are.same('done', statuses[4].phase) end) end) diff --git a/lua/plugins/ai-commit.lua b/lua/plugins/ai-commit.lua index 324bc63..114557e 100644 --- a/lua/plugins/ai-commit.lua +++ b/lua/plugins/ai-commit.lua @@ -265,11 +265,11 @@ local function complete_ollama(full_prompt, callback, status_callback, request_c 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) + local speed_suffix = status.tokens_per_second and string.format(' (%.1f t/s)', status.tokens_per_second) or '' + status_callback('Thinking with ' .. status_model .. speed_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) + local speed_suffix = status.tokens_per_second and string.format(' (%.1f t/s)', status.tokens_per_second) or '' + status_callback('Generating response with ' .. status_model .. speed_suffix) elseif status.phase == 'error' then status_callback('Provider error from ' .. status_model) end