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,