diff --git a/local/ai-commit/lua/ai-commit/config.lua b/local/ai-commit/lua/ai-commit/config.lua index 7bf110f..b5068fd 100644 --- a/local/ai-commit/lua/ai-commit/config.lua +++ b/local/ai-commit/lua/ai-commit/config.lua @@ -15,6 +15,7 @@ M.values = { preview_lines = 5, preview_max_chars = 4000, max_diff_chars = 100000, + prompt_context_ratio = 0.8, refinement = { enabled = true, max_iterations = 2, diff --git a/local/ai-commit/lua/ai-commit/providers.lua b/local/ai-commit/lua/ai-commit/providers.lua index 99c1897..53d25bd 100644 --- a/local/ai-commit/lua/ai-commit/providers.lua +++ b/local/ai-commit/lua/ai-commit/providers.lua @@ -333,6 +333,99 @@ local function generation_sections(recent_commits, session_summary, diff_stat, d return sections end +local function truncate_middle(text, max_chars, label) + if type(text) ~= 'string' or #text <= max_chars then + return text, false + end + + max_chars = math.max(0, max_chars) + local marker = string.format('\n\n[... %s truncated by ai-commit: original=%d chars, kept=%d chars ...]\n\n', label or 'content', #text, max_chars) + if max_chars <= #marker then + return marker, true + end + + local remaining = max_chars - #marker + local head_len = math.floor(remaining * 0.7) + local tail_len = remaining - head_len + local tail_start = math.max(1, #text - tail_len + 1) + return text:sub(1, head_len) .. marker .. text:sub(tail_start), true +end + +local function selected_context_size(source_id) + local ok, ai_provider = pcall(require, 'ai-provider') + if not ok then + return nil + end + + local selection = ai_provider.get_source_selection(source_id) + local provider = selection and selection.provider or ai_provider.get_default_provider() + if not provider then + return nil + end + + local model = selection and selection.model or ai_provider.get_selected_model(provider, source_id) + local provider_config = ai_provider.get_provider_config(provider) + if type(provider_config) ~= 'table' then + return nil + end + + local model_config = provider_config.models and provider_config.models[model] + if type(model_config) == 'table' and type(model_config.context_size) == 'number' then + return model_config.context_size + end + if type(provider_config.context_size) == 'number' then + return provider_config.context_size + end + return nil +end + +local function prompt_budget_chars(source_id) + local context_size = selected_context_size(source_id) + if not context_size then + return nil, nil + end + + local ratio = tonumber(config.values.prompt_context_ratio) or 0.8 + ratio = math.max(0.1, math.min(ratio, 1)) + return math.floor(context_size * ratio), context_size +end + +local function commit_prompt_with_budget(branch, recent_commits, session_summary, diff_stat, diff) + local sections = generation_sections(recent_commits, session_summary, diff_stat, diff) + local prompt = prompts.commit(branch, sections) + local max_prompt_chars, context_size = prompt_budget_chars(config.message_source_id) + if not max_prompt_chars or max_prompt_chars <= 0 or #prompt <= max_prompt_chars or type(diff) ~= 'string' or diff == '' then + return prompt, sections, diff, false, context_size, max_prompt_chars + end + + local excess = #prompt - max_prompt_chars + local diff_budget = math.max(0, #diff - excess - 512) + local truncated_diff, truncated = truncate_middle(diff, diff_budget, 'staged changes') + if not truncated then + return prompt, sections, diff, false, context_size, max_prompt_chars + end + + sections = generation_sections(recent_commits, session_summary, diff_stat, truncated_diff) + prompt = prompts.commit(branch, sections) + while #prompt > max_prompt_chars and diff_budget > 0 do + diff_budget = math.floor(diff_budget * 0.75) + truncated_diff = truncate_middle(diff, diff_budget, 'staged changes') + sections = generation_sections(recent_commits, session_summary, diff_stat, truncated_diff) + prompt = prompts.commit(branch, sections) + end + + log().warn( + string.format( + 'Commit prompt exceeded max size, trimmed staged changes (max_prompt_chars=%d prompt_chars=%d original_diff_chars=%d sent_diff_chars=%d)', + max_prompt_chars, + #prompt, + #diff, + #truncated_diff + ) + ) + return prompt, sections, truncated_diff, true, context_size, max_prompt_chars +end + local function refinement_sections(context) local refinement = config.values.refinement or {} local include = refinement.include_context or {} @@ -429,18 +522,21 @@ function M.generate_commit_message( diff_stat = diff_stat, diff = diff, } - local sections = generation_sections(recent_commits, session_summary, diff_stat, diff) - local prompt = prompts.commit(branch, sections) + local prompt, _, prompt_diff, prompt_truncated, context_size, max_prompt_chars = commit_prompt_with_budget(branch, recent_commits, session_summary, diff_stat, diff) + context.diff = prompt_diff log().debug( string.format( - 'Commit prompt built (branch=%s commits_chars=%d session_context_chars=%d diff_stat_chars=%d diff_chars=%d prompt_chars=%d has_session_context=%s)', + 'Commit prompt built (branch=%s commits_chars=%d session_context_chars=%d diff_stat_chars=%d diff_chars=%d prompt_chars=%d context_size=%s max_prompt_chars=%s has_session_context=%s prompt_truncated=%s)', branch, type(recent_commits) == 'string' and #recent_commits or 0, type(session_summary) == 'string' and #session_summary or 0, type(diff_stat) == 'string' and #diff_stat or 0, - type(diff) == 'string' and #diff or 0, + type(prompt_diff) == 'string' and #prompt_diff or 0, #prompt, - session_summary and 'yes' or 'no' + tostring(context_size), + tostring(max_prompt_chars), + session_summary and 'yes' or 'no', + tostring(prompt_truncated) ) ) dump_prompt(prompt, 'generate-message', branch) diff --git a/local/ai-commit/lua/ai-commit/session_context/opencode.lua b/local/ai-commit/lua/ai-commit/session_context/opencode.lua index 690a5f2..a87b2fa 100644 --- a/local/ai-commit/lua/ai-commit/session_context/opencode.lua +++ b/local/ai-commit/lua/ai-commit/session_context/opencode.lua @@ -165,13 +165,9 @@ local function parse_messages(rows) for _, row in ipairs(rows) do local message = by_id[row.message_id] if not message then - local ok, message_data = pcall(vim.json.decode, row.message_data or '{}') - if not ok or type(message_data) ~= 'table' then - message_data = {} - end message = { id = row.message_id, - role = message_data.role, + role = row.message_role, time_created = row.message_time, parts = {}, } @@ -179,12 +175,9 @@ local function parse_messages(rows) table.insert(messages, message) end - if row.part_data then - local ok, part_data = pcall(vim.json.decode, row.part_data) - if ok and type(part_data) == 'table' and part_data.type == 'text' and type(part_data.text) == 'string' then - table.insert(message.parts, part_data.text) - text_parts = text_parts + 1 - end + if type(row.part_text) == 'string' and row.part_text ~= '' then + table.insert(message.parts, row.part_text) + text_parts = text_parts + 1 end end @@ -199,14 +192,16 @@ end local function read_session_messages(db_path, session, callback, status_callback) local logger = log() local Job = require 'plenary.job' + local opts = config.values.opencode_context or {} + local max_part_chars = math.max(1000, math.floor((opts.max_message_chars or 5000) * 1.2)) if status_callback then status_callback 'OpenCode: Loading session context' end logger.debug('Reading OpenCode session messages for session=' .. tostring(session.id)) local sql = table.concat({ - 'select m.id as message_id, m.time_created as message_time, m.data as message_data,', - 'p.id as part_id, p.time_created as part_time, p.data as part_data', - 'from message m left join part p on p.message_id = m.id', + "select m.id as message_id, m.time_created as message_time, json_extract(m.data, '$.role') as message_role,", + "p.id as part_id, p.time_created as part_time, substr(json_extract(p.data, '$.text'), 1, " .. tostring(max_part_chars) .. ') as part_text', + "from message m left join part p on p.message_id = m.id and json_extract(p.data, '$.type') = 'text'", 'where m.session_id = ' .. sql_quote(session.id), 'order by m.time_created asc, p.time_created asc, p.id asc;', }, ' ') diff --git a/local/ai-provider/lua/ai-provider/providers/ollama.lua b/local/ai-provider/lua/ai-provider/providers/ollama.lua index 48fd785..d20d084 100644 --- a/local/ai-provider/lua/ai-provider/providers/ollama.lua +++ b/local/ai-provider/lua/ai-provider/providers/ollama.lua @@ -6,6 +6,7 @@ local DEFAULT_ENDPOINT = 'http://127.0.0.1:11434' local HEALTH_CACHE_TTL = 60 local MODEL_CACHE_TTL = 60 local STATUS_THROTTLE_MS = 100 +local MAX_PENDING_CHUNK_CHARS = 4000 local state = { health = nil, @@ -186,6 +187,73 @@ function M.chat(request) local status_throttle_ms = request.status_interval or STATUS_THROTTLE_MS local generation_started_at = nil local streamed_token_estimate = 0 + local pending_statuses = {} + local pending_status_flush_scheduled = false + local pending_chunks = {} + local pending_chunk_chars = 0 + local pending_chunk_flush_scheduled = false + + local function flush_pending_statuses() + pending_status_flush_scheduled = false + if not request.on_status or is_cancelled(request) then + pending_statuses = {} + return + end + + local statuses = pending_statuses + pending_statuses = {} + for _, status in ipairs(statuses) do + request.on_status(status) + end + end + + local function flush_pending_chunks() + pending_chunk_flush_scheduled = false + if not request.on_chunk or is_cancelled(request) then + pending_chunks = {} + pending_chunk_chars = 0 + return + end + + local chunks_to_flush = pending_chunks + pending_chunks = {} + pending_chunk_chars = 0 + for _, item in ipairs(chunks_to_flush) do + request.on_chunk(item.chunk, item.data, item.kind) + end + end + + local function schedule_status_flush() + if pending_status_flush_scheduled then + return + end + pending_status_flush_scheduled = true + vim.schedule(flush_pending_statuses) + end + + local function schedule_chunk_flush() + if pending_chunk_flush_scheduled then + return + end + pending_chunk_flush_scheduled = true + vim.schedule(flush_pending_chunks) + end + + local function queue_chunk(chunk, data, kind) + if not request.on_chunk or not chunk or chunk == '' then + return + end + + local item = { chunk = chunk, data = data, kind = kind } + table.insert(pending_chunks, item) + pending_chunk_chars = pending_chunk_chars + #chunk + + while pending_chunk_chars > MAX_PENDING_CHUNK_CHARS and #pending_chunks > 1 do + local removed = table.remove(pending_chunks, 1) + pending_chunk_chars = pending_chunk_chars - #(removed.chunk or '') + end + schedule_chunk_flush() + end local function approximate_stream_tokens_per_second() if not generation_started_at or streamed_token_estimate == 0 then @@ -233,9 +301,12 @@ function M.chat(request) status.elapsed_ms ) ) - vim.schedule(function() - request.on_status(status) - end) + if #pending_statuses > 0 and pending_statuses[#pending_statuses].phase == phase and not important then + pending_statuses[#pending_statuses] = status + else + table.insert(pending_statuses, status) + end + schedule_status_flush() end log.info( @@ -309,11 +380,7 @@ function M.chat(request) message = 'Thinking', tokens_per_second = status_tokens_per_second, } - if request.on_chunk then - vim.schedule(function() - request.on_chunk(thinking, data, 'thinking') - end) - end + queue_chunk(thinking, data, 'thinking') end if chunk ~= '' then table.insert(chunks, chunk) @@ -322,11 +389,7 @@ function M.chat(request) message = 'Generating response', tokens_per_second = status_tokens_per_second, } - if request.on_chunk then - vim.schedule(function() - request.on_chunk(chunk, data, 'message') - end) - end + queue_chunk(chunk, data, 'message') end end, callback = function(code, error_message) diff --git a/local/ai-provider/tests/ai_provider/core_spec.lua b/local/ai-provider/tests/ai_provider/core_spec.lua index 832e1d8..5e0c913 100644 --- a/local/ai-provider/tests/ai_provider/core_spec.lua +++ b/local/ai-provider/tests/ai_provider/core_spec.lua @@ -26,116 +26,4 @@ describe('ai-provider core', function() assert.matches('default_provider must be configured', err) end) - it('selects and saves a model for a source id', 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 original_list_models = core.list_models - local original_select = vim.ui.select - 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) - - 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) - - assert.is_true(ai_provider.register_source 'ai-commit') - ai_provider.select_source_model 'ai-commit' - - rawset(core, 'load_preferences', original_load_preferences) - rawset(core, 'save_preferences', original_save_preferences) - rawset(core, 'list_models', original_list_models) - rawset(vim.ui, 'select', original_select) - - assert.are.same({ provider = 'ollama', model = 'gemma4:e2b 64k' }, prefs.sources['ai-commit']) - end) - - it('stores model preferences per source id', 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') - assert.are.same({ 'ai-commit' }, ai_provider.list_sources()) - assert.is_true(ai_provider.set_source_selection('ai-commit', 'ollama', 'gemma4:e2b 64k')) - assert.are.same('gemma4:e2b 64k', ai_provider.get_selected_model('ollama', 'ai-commit')) - assert.are.same({ provider = 'ollama', model = 'gemma4:e2b 64k', label = 'ollama/gemma4:e2b 64k' }, ai_provider.get_source_selection 'ai-commit') - - 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)