Files
lassistanoque/backend/internal/service/chat/session.go
T

166 lines
4.0 KiB
Go

package chat
import (
"context"
_ "embed"
"fmt"
"log/slog"
"sync"
"trankilou.fr/lassistanoque/backend/internal/domain"
"trankilou.fr/lassistanoque/backend/internal/service/tool"
)
//go:embed system_prompt.md
var systemSystemPrompt string
type ChatSession struct {
userID string
teamID string
chatID string
agent *domain.Agent
provider *domain.Provider
modelID string
messages []*domain.Message
contextMaxTokens int
maxResponseTokens int
llmEngine domain.LLMEngine
toolService *tool.Service
repoChat domain.ChatRepository
subscribers []domain.StreamCallback
mu sync.Mutex
}
func (s *ChatSession) run(ctx context.Context) error {
params := &domain.LLMParams{}
tools := make([]*domain.ToolDefinition, 0)
for _, t := range s.toolService.GetAllToolImpl() {
tools = append(tools, t.Definition(ctx))
}
params.Tools = tools
s.mu.Lock()
if len(s.subscribers) > 0 {
params.OnChunk = func(chunk *domain.Chunk) {
for _, cb := range s.subscribers {
cb(chunk)
}
}
}
s.mu.Unlock()
systemPrompt := &domain.Message{
Role: string(domain.RoleSystem),
Content: systemSystemPrompt + "\n\n## Customization\n\n" + s.agent.SystemPrompt,
}
continue_loop := true
// LOOP on TOOLS
for continue_loop {
// compaction du contexte : il doit rester assez de place pour la réponse
if s.shouldCompact(systemPrompt.Content) {
s.compactContext(ctx)
}
msg := s.llmEngine.Stream(
ctx,
s.provider,
s.modelID,
params,
append([]*domain.Message{systemPrompt}, s.messages...),
)
msg.ChatID = s.chatID
msg.TeamID = s.teamID
msg, err := s.repoChat.CreateChatMessage(s.userID, msg)
if err != nil {
return err
}
s.messages = append(s.messages, msg)
// TOOL
if len(msg.ToolCalls) > 0 {
continue_loop = true
for _, tc := range msg.ToolCalls {
toolImpl := s.toolService.GetToolImpl(tc.Function.Name)
var toolRecord *domain.Tool
var executeErr error
if toolImpl == nil {
executeErr = fmt.Errorf("unknown tool %s", tc.Function.Name)
slog.Info("Execute tool", "error", executeErr)
} else {
slog.Info("Execute tool", "name", tc.Function.Name, "impl", toolImpl)
var err error
toolRecord, err = s.toolService.GetToolByName(s.userID, s.teamID, toolImpl.Name())
if err != nil {
// l'outil n'est pas encore enregistré pour l'équipe :
// auto-enregistrement (comme le fait la liste des outils)
toolRecord, err = s.toolService.CreateTool(s.userID, &domain.Tool{
TeamID: s.teamID,
Name: toolImpl.Name(),
Enabled: true,
})
}
if err != nil {
executeErr = err
slog.Info("Execute tool", "error", err)
}
}
var toolResponse *domain.Message
if executeErr != nil {
toolResponse = &domain.Message{
ChatID: s.chatID,
TeamID: s.teamID,
ToolCallID: tc.ID,
Role: string(domain.RoleTool),
Content: "ERROR: " + executeErr.Error(),
}
} else {
var output []byte
output, executeErr = toolImpl.Execute(
domain.WithToolContext(ctx, s.userID, s.teamID, s.modelID),
[]byte(tc.Function.Arguments),
toolRecord.Configuration,
)
if executeErr != nil {
toolResponse = &domain.Message{
ChatID: s.chatID,
TeamID: s.teamID,
ToolCallID: tc.ID,
Role: string(domain.RoleTool),
Content: "ERROR: " + executeErr.Error(),
}
} else {
toolResponse = &domain.Message{
ChatID: s.chatID,
TeamID: s.teamID,
ToolCallID: tc.ID,
Role: string(domain.RoleTool),
Content: string(output),
}
}
}
persisted, err := s.repoChat.CreateChatMessage(s.userID, toolResponse)
if err != nil {
return err
}
s.messages = append(s.messages, persisted)
}
} else {
continue_loop = false
}
// END TOOL
}
// END LOOP
return nil
}