From c7d3f9123203907c2fb88ebaa946040b3e231b62 Mon Sep 17 00:00:00 2001 From: Meisterlala <6453306+Meisterlala@users.noreply.github.com> Date: Sun, 7 Jun 2026 17:46:18 +0200 Subject: [PATCH] feat(ai-commit): implement heuristic-based auto-refinement for poorly formatted or overly long commit messages Add detection thresholds for body character count (800) and line count (8), enforce structural validation, and introduce optional recent commits context to improve refinement accuracy while preventing infinite loops --- local/ai-commit/lua/ai-commit/config.lua | 15 +- local/ai-commit/lua/ai-commit/generator.lua | 41 ++- local/ai-commit/lua/ai-commit/git.lua | 9 +- local/ai-commit/lua/ai-commit/heuristics.lua | 249 +++++++++++++++++++ local/ai-commit/lua/ai-commit/prompts.lua | 46 +++- local/ai-commit/lua/ai-commit/providers.lua | 204 +++++++++++++-- local/ai-commit/lua/ai-commit/util.lua | 21 -- lua/plugins/ai-commit.lua | 6 +- 8 files changed, 536 insertions(+), 55 deletions(-) create mode 100644 local/ai-commit/lua/ai-commit/heuristics.lua diff --git a/local/ai-commit/lua/ai-commit/config.lua b/local/ai-commit/lua/ai-commit/config.lua index 51b3b07..4245c3a 100644 --- a/local/ai-commit/lua/ai-commit/config.lua +++ b/local/ai-commit/lua/ai-commit/config.lua @@ -13,16 +13,27 @@ M.values = { spinner_interval = 80, preview_lines = 5, max_diff_chars = 100000, + refinement = { + enabled = true, + max_iterations = 2, + include_context = { + recent_commits = false, + staged_files = true, + staged_changes = false, + session_context = false, + }, + recent_commits_with_body = false, + }, diff_context = { small_changed_lines = 100, medium_changed_lines = 500, }, model_highlight_group = 'Special', - prompt_dump_path = vim.fn.stdpath 'log' .. '/ai-commit-last-prompt.md', + prompt_dump_dir = vim.fn.stdpath 'log' .. '/ai-commit-promts', opencode_context = { db_path = vim.fn.expand '~/.local/share/opencode/opencode.db', recent_ms = 60 * 60 * 1000, - assistant_messages = 4, + assistant_messages = 8, max_message_chars = 5000, max_transcript_chars = 30000, }, diff --git a/local/ai-commit/lua/ai-commit/generator.lua b/local/ai-commit/lua/ai-commit/generator.lua index 33a5c45..1c89488 100644 --- a/local/ai-commit/lua/ai-commit/generator.lua +++ b/local/ai-commit/lua/ai-commit/generator.lua @@ -38,6 +38,9 @@ local function set_stage_status(spinner, stage_text, diff_meta, last_status_line if not active_prefix then active_prefix, active_model, active_suffix = stage_text:match '^(Generating commit message with )(.+)( %(.- t/s%))$' end + if not active_prefix then + active_prefix, active_model, active_suffix = stage_text:match '^(Refining commit message %d+ with )(.+)( %(.- t/s%))$' + end if not active_prefix then active_prefix, active_model, active_suffix = stage_text:match '^(Summarizing OpenCode session with )(.+)( %(.- t/s%))$' end @@ -50,6 +53,9 @@ local function set_stage_status(spinner, stage_text, diff_meta, last_status_line if not active_prefix then active_prefix, active_model = stage_text:match '^(Generating commit message with )(.+)$' end + if not active_prefix then + active_prefix, active_model = stage_text:match '^(Refining commit message %d+ with )(.+)$' + end if not active_prefix then active_prefix, active_model = stage_text:match '^(Summarizing OpenCode session with )(.+)$' end @@ -106,6 +112,11 @@ function M.insert() local include_recent_commits = context_opts.recent_commits ~= false local include_opencode = context_opts.opencode ~= false local include_staged_changes = context_opts.staged_changes ~= false + local refinement_opts = config.values.refinement or {} + local refinement_context_opts = refinement_opts.include_context or {} + local include_refinement_commits = refinement_opts.enabled ~= false + and refinement_opts.recent_commits_with_body ~= false + and refinement_context_opts.recent_commits ~= false local context_total = 2 if include_recent_commits then context_total = context_total + 1 @@ -116,6 +127,9 @@ function M.insert() if include_staged_changes then context_total = context_total + 1 end + if include_refinement_commits then + context_total = context_total + 1 + end local context_done = 0 local pending = context_total local failed = false @@ -204,9 +218,19 @@ function M.insert() status(string.format('Preparing context (%d/%d)', context_done, context_total)) pending = pending - 1 if pending == 0 and not failed then - providers.generate_commit_message(context.branch, context.recent_commits, context.session_summary, context.diff_stat, context.diff, function(message) - finalize(message) - end, status, vim.tbl_extend('force', request_context, { status_action = 'Generating commit message' })) + providers.generate_commit_message( + context.branch, + context.recent_commits, + context.session_summary, + context.diff_stat, + context.diff, + context.refinement_recent_commits, + function(message) + finalize(message) + end, + status, + vim.tbl_extend('force', request_context, { status_action = 'Generating commit message' }) + ) elseif failed then finalize(nil) end @@ -234,6 +258,17 @@ function M.insert() end) end + if include_refinement_commits then + git.recent_commits(5, function(commits) + if done then + return + end + context.refinement_recent_commits = commits + logger.debug('Refinement recent commit context ready (chars=' .. tostring(#commits) .. ')') + mark_done() + end, true) + end + if include_opencode then session_context.get_recent(function(session) if done then diff --git a/local/ai-commit/lua/ai-commit/git.lua b/local/ai-commit/lua/ai-commit/git.lua index e140bb0..365a817 100644 --- a/local/ai-commit/lua/ai-commit/git.lua +++ b/local/ai-commit/lua/ai-commit/git.lua @@ -30,11 +30,16 @@ end ---@param count integer ---@param callback function(string) -function M.recent_commits(count, callback) +---@param include_body boolean|nil +function M.recent_commits(count, callback, include_body) local Job = require 'plenary.job' + local args = { 'log', '-n', tostring(count or 5), '--format=%h %s' } + if include_body then + args = { 'log', '-n', tostring(count or 5), '--format=%h %s%n%b%n---' } + end Job:new({ command = 'git', - args = { 'log', '-n', tostring(count or 5), '--format=%h %s' }, + args = args, on_exit = vim.schedule_wrap(function(job, code) if code ~= 0 then callback 'No recent commits available' diff --git a/local/ai-commit/lua/ai-commit/heuristics.lua b/local/ai-commit/lua/ai-commit/heuristics.lua new file mode 100644 index 0000000..0752d1f --- /dev/null +++ b/local/ai-commit/lua/ai-commit/heuristics.lua @@ -0,0 +1,249 @@ +local M = {} + +local DESCRIPTION_TARGET_MAX = 100 +local DESCRIPTION_HARD_MAX = 120 +local BODY_MAX = 80 +local BODY_CHAR_HARD_MAX = 800 +local BODY_LINE_HARD_MAX = 8 + +local function valid_footer(line) + if line:match '^BREAKING CHANGE: .+$' then + return true + end + return line:match '^[A-Za-z0-9-]+: .+$' ~= nil or line:match '^[A-Za-z0-9-]+ #%S.*$' ~= nil +end + +local function looks_like_footer(line) + return line:match '^[A-Z][A-Z ]*: ?' ~= nil or line:match '^[A-Za-z0-9-]+[:#]' ~= nil or line:match '^[A-Za-z0-9-]+ #%S*' ~= nil +end + +local function wrap_line(line, limit) + if #line <= limit or line:match '^%s*$' then + return { line } + end + + local wrapped = {} + local current = '' + for word in line:gmatch '%S+' do + if current == '' then + current = word + elseif #current + 1 + #word <= limit then + current = current .. ' ' .. word + else + table.insert(wrapped, current) + current = word + end + end + if current ~= '' then + table.insert(wrapped, current) + end + return wrapped +end + +local function split_one_line_message(message) + local prefix, description = message:match '^([^:]+: )(.+)$' + if not prefix or not description or #description <= DESCRIPTION_HARD_MAX then + return message + end + + local first_sentence, rest = description:match '^([^%.]+)%.%s+(.+)$' + if not first_sentence or not rest then + return message + end + + return prefix .. first_sentence:gsub('%s+$', '') .. '\n\n' .. rest:gsub('%.%s*$', '') +end + +local function parse_first_line(first_line) + local prefix, description = first_line:match '^([^:]+): (.+)$' + if not prefix or not description then + return nil + end + + local type_name, scope = prefix:match '^([a-z][a-z0-9-]*)%(([a-zA-Z0-9_.-]+)%)!?$' + if type_name and scope then + return { prefix = prefix, description = description } + end + + type_name = prefix:match '^([a-z][a-z0-9-]*)!?$' + if type_name then + return { prefix = prefix, description = description } + end + + return nil +end + +---@param failures string[] +---@param warnings string[] +---@param message string +---@param max_body integer +local function check_lines(failures, warnings, message, max_body) + local lines = vim.split(message or '', '\n', { plain = true }) + local first_line = lines[1] or '' + local parsed_first_line = parse_first_line(first_line) + local first_line_has_breaking_marker = first_line:match '^[^:]+!:' ~= nil + local has_breaking_footer = false + + if first_line == '' then + table.insert(failures, 'missing first line with [optional scope]: ') + elseif not parsed_first_line then + table.insert(failures, 'first line must match Conventional Commits: [optional scope][optional !]: ') + end + + if parsed_first_line then + local description = parsed_first_line.description + local description_length_message = string.format('description is %d chars; max is %d', #description, DESCRIPTION_TARGET_MAX) + if #description > DESCRIPTION_HARD_MAX then + table.insert(failures, description_length_message) + elseif #description > DESCRIPTION_TARGET_MAX then + table.insert(warnings, description_length_message) + end + + if description:match 'BREAKING CHANGE' then + table.insert(failures, 'BREAKING CHANGE must be a separate footer, not part of the description') + end + if #description <= DESCRIPTION_HARD_MAX then + if description:match '%.$' then + table.insert(failures, 'description must not end with a period') + end + if description:match '^%u' then + table.insert(failures, 'description should start lowercase') + end + end + end + + if #lines > 1 then + if lines[2] ~= '' then + table.insert(failures, 'body or footer(s) must be separated from the first line by one blank line') + end + if #lines == 2 then + table.insert(failures, 'message must not end with a dangling blank line') + end + if lines[3] == '' then + table.insert(failures, 'message must not contain multiple blank lines after the first line') + end + end + + local blank_run = 0 + local possible_footer_block = #lines > 2 and lines[2] == '' + local body_chars = 0 + local body_lines = 0 + local in_footer = false + for idx = 3, #lines do + local line = lines[idx] + if #line > max_body then + table.insert(failures, string.format('line %d is %d chars; max is %d', idx, #line, max_body)) + end + if line == '' then + blank_run = blank_run + 1 + if blank_run > 1 then + table.insert(failures, 'message must not contain repeated blank lines') + end + possible_footer_block = true + else + blank_run = 0 + if possible_footer_block and valid_footer(line) then + in_footer = true + if line:match '^BREAKING CHANGE: ' then + has_breaking_footer = true + end + possible_footer_block = true + elseif possible_footer_block and looks_like_footer(line) then + table.insert(failures, string.format('line %d looks like an invalid footer', idx)) + possible_footer_block = false + else + possible_footer_block = false + end + if line:match '^%s*[-*+]%s+' then + table.insert(failures, string.format('line %d must not use markdown list formatting', idx)) + end + if not in_footer then + body_chars = body_chars + #line + body_lines = body_lines + 1 + end + end + end + if body_chars > BODY_CHAR_HARD_MAX then + table.insert(failures, string.format('body is %d chars; max is %d', body_chars, BODY_CHAR_HARD_MAX)) + end + if body_lines > BODY_LINE_HARD_MAX then + table.insert(failures, string.format('body is %d lines; max is %d', body_lines, BODY_LINE_HARD_MAX)) + end + if first_line_has_breaking_marker and not has_breaking_footer and #lines == 1 then + table.insert(warnings, 'breaking-change marker used without body or BREAKING CHANGE footer') + end +end + +---@param message string|nil +---@return table +function M.validate(message) + local failures = {} + local warnings = {} + if type(message) ~= 'string' or message:match '^%s*$' then + return { valid = false, failures = { 'message is empty' }, warnings = warnings } + end + + if message:find '```' then + table.insert(failures, 'message must not include markdown code fences') + end + if message:find '^%s' or message:find '%s$' then + table.insert(failures, 'message must not have leading or trailing whitespace') + end + + check_lines(failures, warnings, message, BODY_MAX) + return { valid = #failures == 0, failures = failures, warnings = warnings } +end + +---@param message string|nil +---@return string|nil +function M.normalize(message) + if type(message) ~= 'string' then + return nil + end + + local trimmed = vim.trim(message) + if not trimmed:find('\n', 1, true) then + trimmed = split_one_line_message(trimmed) + end + + local lines = vim.split(trimmed, '\n', { plain = true }) + lines[1] = (lines[1] or ''):gsub('%.%s*$', '') + if #lines <= 2 then + return table.concat(lines, '\n') + end + + local normalized = { lines[1], lines[2] } + if #lines == 3 then + local line = lines[3]:gsub('^%s*[-*+]%s+', '') + for _, wrapped in ipairs(wrap_line(line, BODY_MAX)) do + table.insert(normalized, wrapped) + end + return table.concat(normalized, '\n') + end + + for idx = 3, #lines do + local line = lines[idx]:gsub('^%s*[-*+]%s+', '') + table.insert(normalized, line) + end + + return table.concat(normalized, '\n') +end + +---@param result table +---@return string +function M.format_failures(result) + if not result or not result.failures or #result.failures == 0 then + return 'No heuristic failures.' + end + + local lines = {} + for _, failure in ipairs(result.failures) do + table.insert(lines, '- ' .. failure) + end + for _, warning in ipairs(result.warnings or {}) do + table.insert(lines, '- warning: ' .. warning) + end + return table.concat(lines, '\n') +end + +return M diff --git a/local/ai-commit/lua/ai-commit/prompts.lua b/local/ai-commit/lua/ai-commit/prompts.lua index a4cd667..7a9b070 100644 --- a/local/ai-commit/lua/ai-commit/prompts.lua +++ b/local/ai-commit/lua/ai-commit/prompts.lua @@ -27,11 +27,16 @@ SPECIFICATION (https://www.conventionalcommits.org/en/v1.0.0/): 14. Types other than feat and fix MAY be used in your commit messages, e.g., docs: update ref docs. ADDITIONAL GUIDELINES: -- Description: Use lowercase, imperative mood, no ending period, max 50 chars +- Description: Use lowercase, imperative mood, no ending period, target max 100 chars - Header Only: Most of the time, ONLY output the single header line (type[scope]: description). +- Header: DO NOT use vague descriptions like `update stuff`, `fix bug`, or `misc changes`. - Body: FORBIDDEN for 75% of commits. DO NOT include a body for small changes, simple fixes, or minor features. - Body: ONLY include a body if the change is a massive architectural shift, highly complex, or a BREAKING CHANGE. -- Body Formatting: If a body is absolutely necessary, wrap at 72 chars, explain WHAT and WHY (not HOW). DO NOT ramble or over-explain. +- Body: DO mention user/API/data/security/deployment impact when relevant. +- Body: DO mention tradeoffs or known limitations only if future maintainers need them. +- Body: DO use neutral factual tone; DO NOT use jokes, apologies, blame, or uncertainty. +- Body: DO NOT list changed files/functions unless the location itself matters. +- Body Formatting: If a body is absolutely necessary, wrap at 80 chars, explain WHAT and WHY (not HOW). DO NOT ramble or over-explain. - Type casing: Any casing may be used, but be consistent (prefer lowercase) - SemVer relationship: fix = PATCH, feat = MINOR, BREAKING CHANGE = MAJOR - Revert commits: Use "revert" type with footer referencing commit SHAs @@ -58,6 +63,43 @@ function M.commit(branch, sections) return table.concat(parts, '\n\n') end +---@param branch string +---@param message string +---@param failures string +---@param sections table[] +---@return string +function M.refine_commit(branch, message, failures, sections) + message = vim.trim(message or '') + message = message:gsub('```', '` ` `') + local parts = { + M.commit_header, + 'Current branch: ' .. branch, + 'Previously generated commit message:', + '```\n' .. message .. '\n```', + 'Heuristic failures:', + failures or 'Unknown failure.', + table.concat({ + 'Refinement constraints:', + 'Return a commit message, not an explanation. Do not copy these constraints into the commit message.', + 'Prefer one short first line only. If a body is necessary, keep it plain text, wrapped at 80 chars, and without bullets.', + 'Do not use ! or BREAKING CHANGE unless the staged files clearly show removed or changed public behavior.', + }, '\n'), + } + + for _, section in ipairs(sections or {}) do + if section.body and section.body ~= '' then + if section.fenced then + table.insert(parts, section.title .. ':\n```\n' .. section.body .. '\n```') + else + table.insert(parts, section.title .. ':\n' .. section.body) + end + end + end + + table.insert(parts, 'Output ONLY the corrected commit message:') + return table.concat(parts, '\n\n') +end + M.session_summary = [[Summarize the following %s session for a git commit message generator. Focus on user intent, important design decisions, problems encountered, and why the final change was made. Ignore tool output details unless they explain the intent or a fix. Keep it around 200 words. Do not invent facts. diff --git a/local/ai-commit/lua/ai-commit/providers.lua b/local/ai-commit/lua/ai-commit/providers.lua index 8186f01..df7a8f5 100644 --- a/local/ai-commit/lua/ai-commit/providers.lua +++ b/local/ai-commit/lua/ai-commit/providers.lua @@ -1,30 +1,68 @@ local config = require 'ai-commit.config' +local heuristics = require 'ai-commit.heuristics' local log = require('ai-commit.log').get local prompts = require 'ai-commit.prompts' local util = require 'ai-commit.util' local M = {} -local function dump_prompt(prompt) +local function sanitize_filename_part(value) + value = tostring(value or 'unknown'):gsub('[^%w_.-]+', '-') + value = value:gsub('^-+', ''):gsub('-+$', '') + return value ~= '' and value or 'unknown' +end + +local function git_short_hash() + local hash = vim.fn.system({ 'git', 'rev-parse', '--short', 'HEAD' }):gsub('%s+$', '') + if vim.v.shell_error ~= 0 or hash == '' then + return 'no-head' + end + return hash +end + +local function git_current_branch() + local branch = vim.fn.system({ 'git', 'branch', '--show-current' }):gsub('%s+$', '') + if vim.v.shell_error ~= 0 or branch == '' then + return 'unknown' + end + return branch +end + +local function dump_prompt(prompt, reason, branch) if config.values.log_level ~= 'debug' then return end - local path = config.values.prompt_dump_path - if type(path) ~= 'string' or path == '' then + local paths = {} + local dump_dir = config.values.prompt_dump_dir + if type(dump_dir) == 'string' and dump_dir ~= '' then + local timestamp = os.date '%Y%m%d-%H%M%S' + local filename = string.format( + '%s-%s-%s-%s.md', + timestamp, + sanitize_filename_part(branch or git_current_branch()), + sanitize_filename_part(git_short_hash()), + sanitize_filename_part(reason or 'prompt') + ) + table.insert(paths, dump_dir .. '/' .. filename) + end + + if #paths == 0 then return end local logger = log() - local ok, err = pcall(function() - vim.fn.mkdir(vim.fn.fnamemodify(path, ':h'), 'p') - vim.fn.writefile(vim.split(prompt, '\n', { plain = true }), path) - end) + for _, path in ipairs(paths) do + local ok, err = pcall(function() + vim.fn.mkdir(vim.fn.fnamemodify(path, ':h'), 'p') + vim.fn.writefile(vim.split(prompt, '\n', { plain = true }), path) + end) - if ok then - logger.debug('Commit prompt dumped to ' .. path) - else - logger.warn('Failed to dump commit prompt to ' .. path .. ': ' .. tostring(err)) + if ok then + logger.debug('Commit prompt dumped to ' .. path) + else + logger.warn('Failed to dump commit prompt to ' .. path .. ': ' .. tostring(err)) + end end end @@ -54,7 +92,15 @@ local function complete_ai_provider(source_id, full_prompt, callback, status_cal local logger = log() local ai_provider = require 'ai-provider' local selection = ai_provider.get_source_selection(source_id) - local provider = selection and selection.provider or ai_provider.get_default_provider() or 'ollama' + local provider = selection and selection.provider or ai_provider.get_default_provider() + + if not provider then + logger.error('No AI provider selected for source=' .. source_id) + vim.notify('No AI provider selected. Check :AIProvider.', vim.log.levels.ERROR) + callback(nil, nil) + return + end + local model = selection and selection.model or ai_provider.get_selected_model(provider, source_id) if not model then @@ -126,7 +172,7 @@ function M.complete_prompt(full_prompt, callback, status_callback, request_conte local ai_provider = require 'ai-provider' local source_id = request_context and request_context.source_id or config.message_source_id local selection = ai_provider.get_source_selection(source_id) - local provider = selection and selection.provider or ai_provider.get_default_provider() or 'ollama' + local provider = selection and selection.provider or ai_provider.get_default_provider() logger.debug( string.format( 'Completing AI prompt (source=%s chars=%d provider=%s model=%s)', @@ -140,11 +186,20 @@ function M.complete_prompt(full_prompt, callback, status_callback, request_conte local action = request_context and request_context.status_action if action and selection and selection.model then status_callback(action .. ' with ' .. selection.model) + elseif not provider then + status_callback 'No AI provider selected' else status_callback('Checking ' .. provider) end end + if not provider then + logger.error('No AI provider selected for source=' .. source_id) + vim.notify('No AI provider selected. Check :AIProvider.', vim.log.levels.ERROR) + callback(nil, nil) + return + end + ai_provider.check(provider, function(working) if request_context and request_context.is_cancelled and request_context.is_cancelled() then return @@ -186,6 +241,7 @@ function M.summarize_session(session, callback, status_callback, request_context #prompt ) ) + dump_prompt(prompt, 'opencode-summary') local summary_context = child_request_context(request_context, { source_id = config.summary_source_id, status_action = 'Summarizing ' .. session_label .. ' session', @@ -202,15 +258,7 @@ function M.summarize_session(session, callback, status_callback, request_context end, status_callback, summary_context) end ----@param branch string ----@param recent_commits string ----@param session_summary string|nil ----@param diff_stat string ----@param diff string ----@param callback function(string|nil, table|nil) ----@param status_callback function(string)|nil ----@param request_context table|nil -function M.generate_commit_message(branch, recent_commits, session_summary, diff_stat, diff, callback, status_callback, request_context) +local function generation_sections(recent_commits, session_summary, diff_stat, diff) local sections = {} if type(diff_stat) == 'string' and not diff_stat:match '^%s*$' then table.insert(sections, { title = 'Staged files', body = diff_stat, fenced = true }) @@ -224,6 +272,106 @@ function M.generate_commit_message(branch, recent_commits, session_summary, diff if type(diff) == 'string' and not diff:match '^%s*$' then table.insert(sections, { title = 'Staged changes', body = diff, fenced = true }) end + return sections +end + +local function refinement_sections(context) + local refinement = config.values.refinement or {} + local include = refinement.include_context or {} + local sections = {} + + if include.staged_files ~= false and type(context.diff_stat) == 'string' and not context.diff_stat:match '^%s*$' then + table.insert(sections, { title = 'Staged files', body = context.diff_stat, fenced = true }) + end + if include.recent_commits ~= false then + local commits = context.refinement_recent_commits or context.recent_commits + if type(commits) == 'string' and not commits:match '^%s*$' then + table.insert(sections, { title = 'Recent commits with bodies', body = commits }) + end + end + if include.session_context ~= false and type(context.session_summary) == 'string' and not context.session_summary:match '^%s*$' then + table.insert(sections, { title = 'Recent assistant session context', body = context.session_summary }) + end + if include.staged_changes ~= false and type(context.diff) == 'string' and not context.diff:match '^%s*$' then + table.insert(sections, { title = 'Staged changes', body = context.diff, fenced = true }) + end + + return sections +end + +local function maybe_refine_message(context, message, iteration, callback, status_callback, request_context) + message = heuristics.normalize(message) or message + local refinement = config.values.refinement or {} + if refinement.enabled == false then + callback(message) + return + end + + local validation = heuristics.validate(message) + if validation.valid then + if validation.warnings and #validation.warnings > 0 then + log().debug('Commit message passed heuristics with warnings: ' .. table.concat(validation.warnings, '; ')) + end + log().debug('Commit message passed heuristics after ' .. tostring(iteration) .. ' refinement(s)') + callback(message) + return + end + + local max_iterations = tonumber(refinement.max_iterations) or 0 + if iteration >= max_iterations then + log().warn('Commit message failed heuristics after max refinements: ' .. heuristics.format_failures(validation):gsub('\n', '; ')) + callback(message) + return + end + + local next_iteration = iteration + 1 + local failures = heuristics.format_failures(validation) + local prompt = prompts.refine_commit(context.branch or 'unknown', message or '', failures, refinement_sections(context)) + log().debug(string.format('Refinement prompt built (iteration=%d prompt_chars=%d failures=%d)', next_iteration, #prompt, #validation.failures)) + dump_prompt(prompt, 'refinement-' .. tostring(next_iteration), context.branch) + + local refinement_context = child_request_context(request_context, { + source_id = config.message_source_id, + status_action = 'Refining commit message ' .. tostring(next_iteration), + }) + M.complete_prompt(prompt, function(refined_message) + if not refined_message then + callback(nil) + return + end + maybe_refine_message(context, refined_message, next_iteration, callback, status_callback, request_context) + end, status_callback, refinement_context) +end + +---@param branch string +---@param recent_commits string +---@param session_summary string|nil +---@param diff_stat string +---@param diff string +---@param refinement_recent_commits string|nil +---@param callback function(string|nil, table|nil) +---@param status_callback function(string)|nil +---@param request_context table|nil +function M.generate_commit_message( + branch, + recent_commits, + session_summary, + diff_stat, + diff, + refinement_recent_commits, + callback, + status_callback, + request_context +) + local context = { + branch = branch, + recent_commits = recent_commits, + refinement_recent_commits = refinement_recent_commits, + session_summary = session_summary, + diff_stat = diff_stat, + diff = diff, + } + local sections = generation_sections(recent_commits, session_summary, diff_stat, diff) local prompt = prompts.commit(branch, sections) log().debug( string.format( @@ -237,11 +385,19 @@ function M.generate_commit_message(branch, recent_commits, session_summary, diff session_summary and 'yes' or 'no' ) ) - dump_prompt(prompt) + dump_prompt(prompt, 'generate-message', branch) request_context = child_request_context(request_context, { source_id = config.message_source_id, }) - M.complete_prompt(prompt, callback, status_callback, request_context) + M.complete_prompt(prompt, function(message, meta) + if not message then + callback(nil, meta) + return + end + maybe_refine_message(context, message, 0, function(final_message) + callback(final_message, meta) + end, status_callback, request_context) + end, status_callback, request_context) end return M diff --git a/local/ai-commit/lua/ai-commit/util.lua b/local/ai-commit/lua/ai-commit/util.lua index df4f9f0..37f9d34 100644 --- a/local/ai-commit/lua/ai-commit/util.lua +++ b/local/ai-commit/lua/ai-commit/util.lua @@ -6,27 +6,6 @@ function M.clean_message(message) return message:gsub('^%s*```.-\n', ''):gsub('\n```%s*$', ''):gsub('^%s+', ''):gsub('%s+$', '') end ----@param body string|nil ----@param max_len integer|nil ----@return string -function M.format_body_for_log(body, max_len) - if type(body) ~= 'string' then - return '' - end - - local compact = body:gsub('%s+', ' '):gsub('^%s+', ''):gsub('%s+$', '') - if compact == '' then - return '' - end - - local limit = max_len or 300 - if #compact <= limit then - return compact - end - - return compact:sub(1, limit) .. '...' -end - ---@param text string|nil ---@param max_chars integer ---@return string diff --git a/lua/plugins/ai-commit.lua b/lua/plugins/ai-commit.lua index c6a99c1..0f165a4 100644 --- a/lua/plugins/ai-commit.lua +++ b/lua/plugins/ai-commit.lua @@ -5,7 +5,11 @@ return { ft = 'gitcommit', cmd = { 'AICommit', 'AICommitModel' }, dependencies = { 'nvim-lua/plenary.nvim', 'ai-provider' }, - opts = {}, + opts = { + refinement = { + max_iterations = 5, + }, + }, config = function(_, opts) require('ai-commit').setup(opts) end, -- 2.51.2