package llm
import (
"context"
"encoding/json"
"fmt"
"log"
"log/slog"
"github.com/firebase/genkit/go/ai"
"github.com/firebase/genkit/go/genkit"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"mimi/internal/bot/llm/agent"
"mimi/internal/bot/llm/agent/fallback"
"mimi/internal/bot/llm/agent/github"
"mimi/internal/bot/llm/agent/logseq"
"mimi/internal/bot/llm/agent/logseqquery"
"mimi/internal/bot/llm/agent/summary"
"mimi/internal/bot/llm/agent/telegram"
"mimi/internal/persist"
logseqscraper "mimi/internal/provider/logseq"
)
type LLM struct {
g *genkit.Genkit
q *persist.Queries
agents map[string]agent.Agent
router ai.Prompt
}
func New(ctx context.Context, pgPool *pgxpool.Pool, graph logseqscraper.RegexGraph, g *genkit.Genkit) LLM {
q := persist.New(pgPool)
ghOrg := "cyberia-to"
agents := []agent.Agent{
logseq.New(g, graph),
logseqquery.New(graph),
fallback.New(g),
github.New(g, ghOrg),
telegram.New(g, pgPool),
summary.New(ctx, g, pgPool, ghOrg, graph.Path),
}
mapped := make(map[string]agent.Agent, len(agents))
for _, agent := range agents {
mapped[agent.GetInfo().Name] = agent
}
router := genkit.LookupPrompt(g, "router")
if router == nil {
log.Fatal("no prompt named 'router' found")
}
return LLM{
g: g,
q: q,
agents: mapped,
router: router,
}
}
func (m LLM) Answer(ctx context.Context, id int64, query string) (result agent.Response, err error) {
resp, err := m.router.Execute(ctx, ai.WithInput(map[string]any{
"query": query,
"agents": m.getAgentsInfo(),
}))
if err != nil {
err = fmt.Errorf("initial LLM call failed with %w", err)
return
}
var output routerOutput
if err := resp.Output(&output); err != nil {
return result, fmt.Errorf("failed to parse router output with %w", err)
}
slog.Info("router answer", "agent", output.Agent)
a, ok := m.agents[output.Agent]
if !ok {
return result, fmt.Errorf("agent with name '%s' not found", output.Agent)
}
rows, err := m.q.FindChatMessages(ctx, id)
var messages []*ai.Message
switch err {
case pgx.ErrNoRows:
break
case nil:
err = json.Unmarshal(rows, &messages)
if err != nil {
err = fmt.Errorf("failed to unmarshal messages with %w", err)
return
}
default:
err = fmt.Errorf("failed to find message history with %w", err)
return
}
defer func() {
messages = append(messages, ai.NewTextMessage(ai.RoleUser, query))
if text, ok := result.Data.(agent.DataText); ok {
messages = append(messages, ai.NewTextMessage(ai.RoleModel, text.Text))
}
if len(messages) > 20 {
messages = messages[len(messages)-20:]
}
var toSave []*ai.Message
for _, m := range messages {
if len(m.Content) == 0 {
continue
}
if !m.Content[0].IsText() {
continue
}
toSave = append(toSave, ai.NewTextMessage(m.Role, m.Content[0].Text))
}
encoded, e := json.Marshal(toSave)
if e != nil {
err = fmt.Errorf("failed to marshal messages with %w", e)
return
}
e = m.q.SaveChatMessages(ctx, persist.SaveChatMessagesParams{
TelegramID: id,
Messages: encoded,
})
if e != nil {
err = fmt.Errorf("failed to save messages with %w", e)
return
}
}()
result, err = a.Run(ctx, query, messages...)
if err != nil {
return result, fmt.Errorf("failed to run agent %s with %w", output.Agent, err)
}
return
}
func (m LLM) getAgentsInfo() (info []agent.Info) {
for _, agent := range m.agents {
info = append(info, agent.GetInfo())
}
return info
}
type routerOutput struct {
Agent string `json:"agent"`
}