diff --git a/local/ai-commit/lua/ai-commit/spinner.lua b/local/ai-commit/lua/ai-commit/spinner.lua index 9fd4f29..37692dc 100644 --- a/local/ai-commit/lua/ai-commit/spinner.lua +++ b/local/ai-commit/lua/ai-commit/spinner.lua @@ -197,6 +197,12 @@ function M.append_stream(spinner, text) if not spinner or not text or text == '' then return end + if vim.in_fast_event() then + vim.schedule(function() + M.append_stream(spinner, text) + end) + return + end local chunks = vim.split(text, '\n', { plain = true }) for index, chunk in ipairs(chunks) do diff --git a/local/ai-provider/lua/ai-provider/curl.lua b/local/ai-provider/lua/ai-provider/curl.lua index 7a39954..7d6f151 100644 --- a/local/ai-provider/lua/ai-provider/curl.lua +++ b/local/ai-provider/lua/ai-provider/curl.lua @@ -100,7 +100,7 @@ end ---@field timeout? integer Curl max-time in milliseconds. ---@field is_cancelled? fun(): boolean Optional cancellation predicate. ---@field on_json_line? fun(data: table, line: string) Called for each decoded JSON line. ----@field callback fun(code: integer, error: string|nil) Called when curl exits. +---@field callback fun(code: integer, error: string|nil, status: integer|nil) Called when curl exits. ---Run a streaming curl request and decode each stdout line as JSON. ---@param request AiProviderCurlStreamRequest @@ -110,6 +110,7 @@ function M.stream_json_lines(request) local body = encode_body(request.body) local method = request.method or 'POST' local timeout = request.timeout or DEFAULT_TIMEOUT + local status_marker = '__AI_PROVIDER_HTTP_STATUS__:' local args = { '--silent', '--show-error', @@ -118,6 +119,8 @@ function M.stream_json_lines(request) tostring(math.ceil(timeout / 1000)), '--request', method, + '--write-out', + '\n' .. status_marker .. '%{http_code}', } for name, value in pairs(json_headers(request.headers)) do @@ -133,21 +136,43 @@ function M.stream_json_lines(request) table.insert(args, request.url) local stderr = {} + local stdout_lines = {} + local http_status = nil local job = Job:new { command = 'curl', args = args, - on_stdout = function(_, line) - if not line or line == '' or (request.is_cancelled and request.is_cancelled()) then + on_stdout = function(_, output) + if not output or output == '' or (request.is_cancelled and request.is_cancelled()) then return end - local ok, data = pcall(vim.json.decode, line) - if not ok or type(data) ~= 'table' then - return - end - - if request.on_json_line then - request.on_json_line(data, line) + for line in tostring(output):gmatch('[^\r\n]+') do + line = line:gsub('^%s+', ''):gsub('%s+$', '') + if line ~= '' then + local status = line:match('^' .. status_marker .. '(%d+)$') + if status then + http_status = tonumber(status) + else + local payload = line + local event_data = line:match '^data:%s*(.*)$' + if event_data then + if event_data == '[DONE]' then + goto continue + end + payload = event_data + end + + if #stdout_lines < 20 then + table.insert(stdout_lines, payload) + end + + local ok, data = pcall(vim.json.decode, payload) + if ok and type(data) == 'table' and request.on_json_line then + schedule(request.on_json_line, data, line) + end + end + end + ::continue:: end end, on_stderr = function(_, line) @@ -156,7 +181,14 @@ function M.stream_json_lines(request) end end, on_exit = function(_, code) - schedule(request.callback, code, #stderr > 0 and table.concat(stderr, '\n') or nil) + local error_parts = {} + if #stderr > 0 then + table.insert(error_parts, table.concat(stderr, '\n')) + end + if http_status and http_status >= 400 and #stdout_lines > 0 then + table.insert(error_parts, table.concat(stdout_lines, '\n')) + end + schedule(request.callback, code, #error_parts > 0 and table.concat(error_parts, '\n') or nil, http_status) end, } diff --git a/local/ai-provider/lua/ai-provider/providers/copilot.lua b/local/ai-provider/lua/ai-provider/providers/copilot.lua index fe2715a..fc0135a 100644 --- a/local/ai-provider/lua/ai-provider/providers/copilot.lua +++ b/local/ai-provider/lua/ai-provider/providers/copilot.lua @@ -97,16 +97,20 @@ local function register_job(request, job) end end +local function elapsed_ms_since(started_at) + return (vim.uv.hrtime() - started_at) / 1e6 +end + local function get_api_token(request, callback) if state.api_token and state.api_token.expires_at and state.api_token.expires_at > os.time() then callback(state.api_token.token) - return + return nil end local oauth_token = get_oauth_token() if not oauth_token then callback(nil) - return + return nil end local job = curl.json { @@ -138,6 +142,7 @@ local function get_api_token(request, callback) end, } register_job(request, job) + return job end function M.check(callback, opts) @@ -187,6 +192,15 @@ end function M.chat(request) local started_at = vim.uv.hrtime() + local active_job = nil + local proxy_job = { + shutdown = function() + if active_job and active_job.shutdown then + active_job:shutdown() + end + end, + } + register_job(request, proxy_job) local function emit_status(phase, message) if request.on_status then @@ -195,13 +209,13 @@ function M.chat(request) phase = phase, message = message, model = request.model, - elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6, + elapsed_ms = elapsed_ms_since(started_at), } end end emit_status('authenticating', 'Authenticating with Copilot') - get_api_token(request, function(token) + local auth_job = get_api_token(request, function(token) if is_cancelled(request) then return end @@ -211,7 +225,7 @@ function M.chat(request) request.callback(nil, { requested_model = request.model, used_model = request.model, - elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6, + elapsed_ms = elapsed_ms_since(started_at), error = 'copilot authentication failed', }) end @@ -224,12 +238,14 @@ function M.chat(request) headers['Content-Type'] = 'application/json' local requested_model = request.model or AUTO_MODEL - local active_job = nil local function send_chat(model, retried_auto) + local chunks = {} + local used_model = model or AUTO_MODEL + local stream_error = nil local body = { messages = { { role = 'user', content = request.prompt } }, - stream = false, + stream = request.stream ~= false, max_tokens = request.max_tokens, } if model and model ~= AUTO_MODEL then @@ -237,6 +253,91 @@ function M.chat(request) end emit_status('generating', 'Generating response') + if body.stream then + active_job = curl.stream_json_lines { + method = 'POST', + url = endpoint .. '/chat/completions', + headers = headers, + body = body, + timeout = request.timeout or 30000, + is_cancelled = request.is_cancelled, + on_json_line = function(data) + if is_cancelled(request) then + return + end + + used_model = data.model or used_model + if data.error then + stream_error = data.error + end + local choice = data.choices and data.choices[1] + local delta = choice and choice.delta and choice.delta.content + if type(delta) ~= 'string' or delta == '' then + return + end + + table.insert(chunks, delta) + if request.on_chunk then + request.on_chunk(delta, data, 'message') + end + end, + callback = function(code, error_message, status) + if is_cancelled(request) then + return + end + + local elapsed_ms = elapsed_ms_since(started_at) + if code ~= 0 or (status and status >= 400) or stream_error then + local error_message_with_status = 'copilot api request failed: ' .. tostring(status or code) + local error_code = type(stream_error) == 'table' and stream_error.code or nil + local error_body = error_message or (stream_error and vim.json.encode(stream_error)) or nil + log.error(string.format('%s elapsed_ms=%.0f body=%s', error_message_with_status, elapsed_ms, format_body_for_log(error_body))) + if status == 400 and (error_code == 'unsupported_api_for_model' or (error_body and error_body:match 'unsupported_api_for_model')) and body.model and not retried_auto then + log.warn('copilot model unsupported by streaming chat completions, retrying with auto: ' .. tostring(body.model)) + send_chat(AUTO_MODEL, true) + return + end + + if request.callback then + request.callback(nil, { + requested_model = requested_model, + used_model = used_model, + elapsed_ms = elapsed_ms, + error = error_message_with_status, + }) + end + return + end + + local message = table.concat(chunks, '') + if message == '' then + log.error('copilot streaming response missing message') + if request.callback then + request.callback(nil, { + requested_model = requested_model, + used_model = used_model, + elapsed_ms = elapsed_ms, + error = 'copilot returned no content', + }) + end + return + end + + log.info(string.format('copilot streaming response requested_model=%s used_model=%s elapsed_ms=%.0f', requested_model, used_model, elapsed_ms)) + emit_status('done', 'Response complete') + if request.callback then + request.callback(message, { + requested_model = requested_model, + used_model = used_model, + elapsed_ms = elapsed_ms, + }) + end + end, + } + register_job(request, active_job) + return + end + active_job = curl.json { method = 'POST', url = endpoint .. '/chat/completions', @@ -248,7 +349,7 @@ function M.chat(request) return end - local elapsed_ms = (vim.uv.hrtime() - started_at) / 1e6 + local elapsed_ms = elapsed_ms_since(started_at) if response.status ~= 200 then local error_code = response.json and response.json.error and response.json.error.code if response.status == 400 and error_code == 'unsupported_api_for_model' and body.model and not retried_auto then @@ -301,8 +402,11 @@ function M.chat(request) end send_chat(request.model or AUTO_MODEL, false) - return active_job end) + if auth_job then + active_job = auth_job + end + return proxy_job end return M diff --git a/local/ai-provider/lua/ai-provider/providers/ollama.lua b/local/ai-provider/lua/ai-provider/providers/ollama.lua index d20d084..88d623d 100644 --- a/local/ai-provider/lua/ai-provider/providers/ollama.lua +++ b/local/ai-provider/lua/ai-provider/providers/ollama.lua @@ -392,24 +392,25 @@ function M.chat(request) queue_chunk(chunk, data, 'message') end end, - callback = function(code, error_message) + callback = function(code, error_message, status) 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 code ~= 0 or (status and status >= 400) then + local error_detail = error_message or (status and ('HTTP ' .. status) or 'ollama request failed') + log.error('ollama request process failed code=' .. tostring(code) .. ' status=' .. tostring(status) .. ' model=' .. tostring(raw_model) .. ' error=' .. tostring(error_detail)) 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', + error = error_detail, }) end emit_status { phase = 'error', - message = error_message or 'Ollama request failed', + message = error_detail, } return end @@ -536,6 +537,10 @@ function M.chat(request) phase = 'loading', message = 'Loading model', } + emit_status { + phase = 'context', + message = 'Loading prompt context', + } else emit_status { phase = 'context', diff --git a/local/ai-provider/tests/ai_provider/copilot_spec.lua b/local/ai-provider/tests/ai_provider/copilot_spec.lua new file mode 100644 index 0000000..8af3750 --- /dev/null +++ b/local/ai-provider/tests/ai_provider/copilot_spec.lua @@ -0,0 +1,94 @@ +describe('copilot provider', function() + local original_curl + local original_path + local original_log + + before_each(function() + original_curl = package.loaded['ai-provider.curl'] + original_path = package.loaded['plenary.path'] + original_log = package.loaded['ai-provider.log'] + package.loaded['ai-provider.providers.copilot'] = nil + package.loaded['ai-provider.log'] = { + debug = function() end, + info = function() end, + warn = function() end, + error = function() end, + } + + local fake_path = {} + function fake_path:joinpath() + return self + end + function fake_path:exists() + return true + end + function fake_path:read() + return vim.json.encode { ['github.com'] = { oauth_token = 'oauth-token' } } + end + + package.loaded['plenary.path'] = { + new = function() + return fake_path + end, + } + end) + + after_each(function() + package.loaded['ai-provider.curl'] = original_curl + package.loaded['plenary.path'] = original_path + package.loaded['ai-provider.log'] = original_log + package.loaded['ai-provider.providers.copilot'] = nil + end) + + it('streams chat completion deltas through provider chunks', function() + local captured_body = nil + package.loaded['ai-provider.curl'] = { + json = function(request) + assert.matches('/copilot_internal/v2/token$', request.url) + request.callback { + status = 200, + json = { + token = 'api-token', + expires_at = os.time() + 60, + endpoints = { api = 'https://copilot.example.test' }, + }, + } + return { shutdown = function() end } + end, + stream_json_lines = function(request) + captured_body = request.body + request.on_json_line { model = 'gpt-4o', choices = { { delta = { content = 'he' } } } } + request.on_json_line { model = 'gpt-4o', choices = { { delta = { content = 'llo' } } } } + request.callback(0, nil, 200) + return { shutdown = function() end } + end, + } + + local copilot = require 'ai-provider.providers.copilot' + local chunks = {} + local message = nil + local meta = nil + + local job = copilot.chat { + model = 'gpt-4o', + prompt = 'Say hello', + on_chunk = function(chunk) + table.insert(chunks, chunk) + end, + callback = function(result, result_meta) + message = result + meta = result_meta + end, + } + + assert.is_table(job) + assert.is_table(captured_body) + ---@cast captured_body table + assert.are.same(true, captured_body.stream) + assert.are.same('hello', table.concat(chunks, '')) + assert.are.same('hello', message) + assert.is_table(meta) + ---@cast meta table + assert.are.same('gpt-4o', meta.used_model) + 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 4783e92..4cc3225 100644 --- a/local/ai-provider/tests/ai_provider/ollama_options_spec.lua +++ b/local/ai-provider/tests/ai_provider/ollama_options_spec.lua @@ -318,11 +318,11 @@ describe('ollama provider options', function() } vim.wait(1000, function() - return #statuses >= 3 + return #statuses >= 4 end, 10) assert.are.same('loading', statuses[1]) - assert.are.same('generating', statuses[2]) - assert.is_false(vim.tbl_contains(statuses, 'context')) + assert.are.same('context', statuses[2]) + assert.are.same('generating', statuses[3]) end) end)