166 lines
4.0 KiB
Go
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
|
|
}
|