Skip to content

Commit 9f629b9

Browse files
authored
Merge pull request #198 from frankroeder/xai-pipeline
feat(provider): support xAI /v1/responses for Grok
2 parents 964675c + 5ca3500 commit 9f629b9

6 files changed

Lines changed: 116 additions & 15 deletions

File tree

.github/workflows/release.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ jobs:
1818
manager: sudo apt-get
1919
packages: -y ripgrep
2020
- os: ubuntu-latest
21-
rev: v0.10.4/nvim-linux-x86_64.tar.gz
21+
rev: v0.12.4/nvim-linux-x86_64.tar.gz
2222
manager: sudo apt-get
2323
packages: -y ripgrep
2424
steps:

.github/workflows/test.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ jobs:
1515
manager: sudo apt-get
1616
packages: -y ripgrep
1717
- os: ubuntu-latest
18-
rev: v0.10.4/nvim-linux-x86_64.tar.gz
18+
rev: v0.12.4/nvim-linux-x86_64.tar.gz
1919
manager: sudo apt-get
2020
packages: -y ripgrep
2121
steps:

README.md

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -595,20 +595,21 @@ providers = {
595595
providers = {
596596
xai = {
597597
name = "xai",
598-
endpoint = "https://api.x.ai/v1/chat/completions",
598+
endpoint = "https://api.x.ai/v1/responses",
599599
model_endpoint = "https://api.x.ai/v1/language-models",
600600
api_key = os.getenv "XAI_API_KEY",
601601
params = {
602602
chat = { temperature = 1.1, top_p = 1 },
603603
command = { temperature = 1.1, top_p = 1 },
604604
},
605605
topic = {
606-
model = "grok-3-mini-beta",
607-
params = { max_completion_tokens = 64 },
606+
model = "grok-4.3",
607+
params = { max_output_tokens = 64 },
608608
},
609609
models = {
610-
"grok-3-beta",
611-
"grok-3-mini-beta",
610+
"grok-4.3",
611+
"grok-4.5",
612+
"grok-4.20",
612613
},
613614
},
614615
}

lua/parrot/provider/multi_provider.lua

Lines changed: 41 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,9 @@ local defaults = {
7777

7878
if decoded.choices and decoded.choices[1] and decoded.choices[1].delta and decoded.choices[1].delta.content then
7979
return decoded.choices[1].delta.content
80+
elseif decoded.type == "response.output_text.delta" and type(decoded.delta) == "string" then
81+
-- OpenAI / xAI Responses API streaming
82+
return decoded.delta
8083
elseif decoded.message and decoded.message.content then
8184
return decoded.message.content
8285
elseif decoded.delta and decoded.delta.type == "text_delta" and decoded.delta.text then
@@ -115,6 +118,21 @@ local defaults = {
115118
return decoded.message.content
116119
elseif decoded.content and decoded.content[1] and decoded.content[1].text then
117120
return decoded.content[1].text
121+
elseif decoded.output and type(decoded.output) == "table" then
122+
-- OpenAI / xAI Responses API non-streaming output
123+
local texts = {}
124+
for _, item in ipairs(decoded.output) do
125+
if item.type == "message" and type(item.content) == "table" then
126+
for _, part in ipairs(item.content) do
127+
if part.text and (part.type == "output_text" or part.type == "text") then
128+
table.insert(texts, part.text)
129+
end
130+
end
131+
end
132+
end
133+
if #texts > 0 then
134+
return table.concat(texts)
135+
end
118136
end
119137

120138
return nil
@@ -342,11 +360,32 @@ function MultiProvider:set_model(model)
342360
self._model = model
343361
end
344362

345-
-- Preprocesses the payload before sending to the API
363+
-- Resolve endpoint string (handles function endpoints)
364+
---@return string|nil
365+
function MultiProvider:get_endpoint()
366+
if type(self.endpoint) == "function" then
367+
local ok, result = pcall(self.endpoint, self)
368+
if not ok then
369+
logger.error("Error executing endpoint function for provider " .. self.name .. ": " .. tostring(result))
370+
return nil
371+
end
372+
return result
373+
end
374+
return self.endpoint
375+
end
376+
377+
-- Preprocesses the payload before sending to the API.
378+
-- /v1/responses expects `input` instead of `messages`.
346379
---@param payload table
347380
---@return table
348381
function MultiProvider:preprocess_payload(payload)
349-
return self.preprocess_payload_func(payload)
382+
local result = self.preprocess_payload_func(payload)
383+
local endp = self:get_endpoint()
384+
if type(endp) == "string" and endp:find("/responses", 1, true) and result.messages then
385+
result.input = result.messages
386+
result.messages = nil
387+
end
388+
return result
350389
end
351390

352391
-- Returns the curl parameters for the API request

tests/parrot/context_spec.lua

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -68,12 +68,13 @@ describe("context", function()
6868
end)
6969

7070
it("should correctly handle buffers", function()
71-
vim.fn.writefile(
72-
{ "local test_buffer_content = 'this is a buffer'", "print(test_buffer_content)" },
73-
"test/buffer.lua"
74-
)
75-
vim.cmd("edit test/buffer.lua")
76-
local buf_id = vim.api.nvim_get_current_buf()
71+
local lines = { "local test_buffer_content = 'this is a buffer'", "print(test_buffer_content)" }
72+
local path = vim.fn.fnamemodify("test/buffer.lua", ":p")
73+
vim.fn.writefile(lines, path)
74+
-- Create buffer without :edit to avoid FileType/treesitter ftplugin errors on nvim 0.12+
75+
local buf_id = vim.api.nvim_create_buf(true, false)
76+
vim.api.nvim_buf_set_name(buf_id, path)
77+
vim.api.nvim_buf_set_lines(buf_id, 0, -1, false, lines)
7778
local buf_name = vim.api.nvim_buf_get_name(buf_id)
7879
local result_current_buffer = context.insert_contexts("@buffer:" .. buf_name)
7980
local buf_content = table.concat(vim.api.nvim_buf_get_lines(buf_id, 0, -1, false), "\n")
@@ -82,6 +83,7 @@ describe("context", function()
8283
local non_existent_buffer = "@buffer:non_existent_buffer.lua"
8384
local result_buffer_not_found = context.insert_contexts(non_existent_buffer)
8485
assert.are.equal("", result_buffer_not_found)
86+
vim.api.nvim_buf_delete(buf_id, { force = true })
8587
end)
8688

8789
it("should handle multiple commands", function()

tests/parrot/provider/multi_provider_spec.lua

Lines changed: 59 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -281,6 +281,15 @@ describe("MultiProvider", function()
281281

282282
assert.is_nil(result)
283283
end)
284+
285+
it("should extract content from Responses API output_text delta", function()
286+
local input =
287+
'data: {"type":"response.output_text.delta","item_id":"msg_123","output_index":0,"content_index":0,"delta":" Hello"}'
288+
289+
local result = provider:process_stdout(input)
290+
291+
assert.equals(" Hello", result)
292+
end)
284293
end)
285294

286295
describe("preprocess_payload", function()
@@ -336,6 +345,56 @@ describe("MultiProvider", function()
336345
local result = provider:preprocess_payload(input)
337346
assert.are.same(result, expected)
338347
end)
348+
349+
it("should use input for /responses endpoints", function()
350+
local responses_provider = MultiProvider:new({
351+
name = "xai",
352+
endpoint = "https://api.x.ai/v1/responses",
353+
api_key = "test_api_key",
354+
model = { "grok-4.5" },
355+
})
356+
357+
local result = responses_provider:preprocess_payload({
358+
messages = {
359+
{ role = "system", content = " You are Grok. " },
360+
{ role = "user", content = " Hi " },
361+
},
362+
model = "grok-4.5",
363+
stream = true,
364+
temperature = 1.1,
365+
top_p = 1,
366+
max_output_tokens = 64,
367+
})
368+
369+
assert.is_nil(result.messages)
370+
assert.equals(64, result.max_output_tokens)
371+
assert.equals("You are Grok.", result.input[1].content)
372+
assert.equals("Hi", result.input[2].content)
373+
assert.equals("grok-4.5", result.model)
374+
end)
375+
end)
376+
377+
describe("process_onexit responses", function()
378+
it("should extract text from Responses API output", function()
379+
local input = vim.json.encode({
380+
id = "resp_123",
381+
object = "response",
382+
status = "completed",
383+
output = {
384+
{
385+
type = "message",
386+
role = "assistant",
387+
content = {
388+
{ type = "output_text", text = "Hello from responses" },
389+
},
390+
},
391+
},
392+
})
393+
394+
local result = provider:process_onexit(input)
395+
396+
assert.equals("Hello from responses", result)
397+
end)
339398
end)
340399

341400
describe("verify", function()

0 commit comments

Comments
 (0)