Luigit
repositories / dotfiles

dotfiles

bugabingas dorkfiles

owned by admin

neovim/lua/bugabinga/pivi/assemble.lua

Raw
local TEXT = 'text'
local THINKING = 'thinking'
local TOOL_CALL = 'toolcall'

local new = function()
  return {
    parts = {},
    order = {},
    complete = false,
  }
end

local part_at = function( state, index, kind )
  local part = state.parts[index]
  if not part then
    part = { kind = kind, index = index, text = '', }
    state.parts[index] = part
    table.insert( state.order, index )
    table.sort( state.order )
  end
  return part
end

local ordered_parts = function( state )
  local parts = {}
  for _, index in ipairs( state.order ) do
    table.insert( parts, state.parts[index] )
  end
  return parts
end

local from_message = function( message )
  local state = new()
  local content = message and message.content
  if type( content ) ~= 'table' then return state end
  for offset, block in ipairs( content ) do
    local index = offset - 1
    if block.type == TEXT then
      local part = part_at( state, index, TEXT )
      part.text = block.text or ''
    elseif block.type == THINKING then
      local part = part_at( state, index, THINKING )
      part.text = block.thinking or ''
    elseif block.type == 'toolCall' then
      local part = part_at( state, index, TOOL_CALL )
      part.id = block.id
      part.name = block.name
      part.arguments = block.arguments
      part.complete = true
    end
  end
  state.complete = true
  return state
end

--- Applies one streaming delta to the assembled message state.
--- `message_end` is authoritative and replaces everything accumulated so far.
--- @param state table assembled state
--- @param event table one RPC event
--- @return table state the updated assembled state
local apply = function( state, event )
  if type( event ) ~= 'table' then return state end

  if event.type == 'message_start' then
    if event.message and event.message.role ~= 'assistant' then return state end
    return new()
  end

  if event.type == 'message_end' then
    if event.message and event.message.role ~= 'assistant' then return state end
    return from_message( event.message )
  end

  if event.type ~= 'message_update' then return state end

  local delta = event.assistantMessageEvent
  if type( delta ) ~= 'table' then return state end

  local index = delta.contentIndex or 0

  if delta.type == 'text_start' then
    part_at( state, index, TEXT )
  elseif delta.type == 'text_delta' then
    local part = part_at( state, index, TEXT )
    part.text = part.text .. ( delta.delta or '' )
  elseif delta.type == 'text_end' then
    local part = part_at( state, index, TEXT )
    if delta.content then part.text = delta.content end
  elseif delta.type == 'thinking_start' then
    part_at( state, index, THINKING )
  elseif delta.type == 'thinking_delta' then
    local part = part_at( state, index, THINKING )
    part.text = part.text .. ( delta.delta or '' )
  elseif delta.type == 'thinking_end' then
    local part = part_at( state, index, THINKING )
    if delta.content then part.text = delta.content end
  elseif delta.type == 'toolcall_start' then
    local part = part_at( state, index, TOOL_CALL )
    part.id = delta.id
    part.name = delta.toolName
  elseif delta.type == 'toolcall_delta' then
    local part = part_at( state, index, TOOL_CALL )
    part.text = part.text .. ( delta.delta or '' )
  elseif delta.type == 'toolcall_end' then
    local part = part_at( state, index, TOOL_CALL )
    local call = delta.toolCall
    if call then
      part.id = call.id or part.id
      part.name = call.name or part.name
      part.arguments = call.arguments
    end
    part.complete = true
  end

  return state
end

--- @return string text visible assistant text, thinking excluded
local text = function( state )
  local pieces = {}
  for _, part in ipairs( ordered_parts( state ) ) do
    if part.kind == TEXT then table.insert( pieces, part.text ) end
  end
  return table.concat( pieces )
end

--- @return string thinking accumulated reasoning text
local thinking = function( state )
  local pieces = {}
  for _, part in ipairs( ordered_parts( state ) ) do
    if part.kind == THINKING then table.insert( pieces, part.text ) end
  end
  return table.concat( pieces )
end

--- @return table[] calls tool calls in content order
local tool_calls = function( state )
  local calls = {}
  for _, part in ipairs( ordered_parts( state ) ) do
    if part.kind == TOOL_CALL then
      table.insert( calls, {
        index = part.index,
        id = part.id,
        name = part.name,
        arguments = part.arguments,
        arguments_text = part.text,
        complete = part.complete == true,
      } )
    end
  end
  return calls
end

return {
  new = new,
  apply = apply,
  text = text,
  thinking = thinking,
  tool_calls = tool_calls,
}