Files

296 lines
6.3 KiB
Go

package chat
import (
"context"
"fmt"
"log/slog"
"sync"
"time"
"trankilou.fr/lassistanoque/backend/internal/domain"
"trankilou.fr/lassistanoque/backend/internal/service/tool"
)
type Service struct {
repoUser domain.UserRepository
repoAgent domain.AgentRepository
repoProvider domain.ProviderRepository
repoChat domain.ChatRepository
llmEngine domain.LLMEngine
runningSessions map[string]*ChatSession
sessionCancels map[string]context.CancelFunc
toolService *tool.Service
mu sync.Mutex
}
func NewService(
repoUser domain.UserRepository,
repoAgent domain.AgentRepository,
repoProvider domain.ProviderRepository,
repoChat domain.ChatRepository,
llmEngine domain.LLMEngine,
toolService *tool.Service,
) *Service {
return &Service{
repoUser: repoUser,
repoAgent: repoAgent,
repoProvider: repoProvider,
repoChat: repoChat,
llmEngine: llmEngine,
toolService: toolService,
runningSessions: make(map[string]*ChatSession),
sessionCancels: make(map[string]context.CancelFunc),
}
}
func (s *Service) NewChat(
userID string,
teamID string,
agentID string,
taskID string,
) (*domain.Chat, error) {
team, err := s.repoUser.FindTeam(userID, teamID)
if err != nil {
return nil, err
}
start := time.Now()
chat := &domain.Chat{
UserID: userID,
AgentID: agentID,
TeamID: team.ID,
TaskID: taskID,
StartDatetime: &start,
Title: "Nouvelle conversation",
}
chat, err = s.repoChat.CreateChat(userID, chat)
if err != nil {
return nil, err
}
return chat, nil
}
func (s *Service) AddChatMessage(
ctx context.Context,
userID string,
teamID string,
chatID string,
agentID string,
prompt string,
params *domain.ChatParams,
) (*domain.Message, error) {
team, err := s.repoUser.FindTeam(userID, teamID)
if err != nil {
return nil, err
}
var chatAgentID string
if team.DefaultAgentID != nil {
chatAgentID = *team.DefaultAgentID
} else if agentID != "" {
chatAgentID = agentID
} else {
return nil, fmt.Errorf("no agent defined")
}
agent, err := s.repoAgent.GetAgent(userID, team.ID, chatAgentID)
if err != nil {
return nil, err
}
chatProviderID := agent.DefaultProviderID
chatModelID := agent.DefaultModelID
if params != nil && params.ProviderID != "" {
chatProviderID = params.ProviderID
}
if params != nil && params.ModelID != "" {
chatModelID = params.ModelID
}
provider, err := s.repoProvider.GetProvider(userID, teamID, chatProviderID)
if err != nil {
return nil, err
}
message, err := s.repoChat.CreateChatMessage(userID, &domain.Message{
ChatID: chatID,
TeamID: teamID,
Role: string(domain.RoleUser),
Content: prompt,
})
if err != nil {
return nil, err
}
messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID, true)
if err != nil {
return nil, err
}
var cb domain.StreamCallback
if params != nil && params.OnChunk != nil {
cb = params.OnChunk
}
contextMax := defaultContextMaxTokens
maxResponse := defaultMaxResponseTokens
if params != nil {
if params.ContextMaxTokens > 0 {
contextMax = params.ContextMaxTokens
}
if params.MaxResponseTokens > 0 {
maxResponse = params.MaxResponseTokens
}
}
s.runQuery(
ctx,
userID,
teamID,
chatID,
agent,
provider,
chatModelID,
contextMax,
maxResponse,
messages,
cb,
)
chat, _ := s.repoChat.GetChat(userID, teamID, chatID)
chat.FreshTitle = false
s.repoChat.UpdateChat(userID, chat)
go s.RefreshTitles(provider, chatModelID, teamID)
return message, nil
}
func (s *Service) ListChats(userID string, teamID string, page int, query string, taskID string) ([]*domain.Chat, error) {
return s.repoChat.ListChats(userID, teamID, page, query, taskID)
}
func (s *Service) GetChat(userID string, teamID, chatID string) (*domain.Chat, error) {
return s.repoChat.GetChat(userID, teamID, chatID)
}
func (s *Service) GetChatWithMessages(userID string, teamID, chatID string) (*domain.ChatWithMessages, error) {
chat, err := s.repoChat.GetChat(userID, teamID, chatID)
if err != nil {
return nil, err
}
messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID, false)
if err != nil {
return nil, err
}
return &domain.ChatWithMessages{
Chat: chat,
Messages: messages,
}, nil
}
func (s *Service) DeleteChat(userID string, teamID, chatID string) error {
// stoppe une éventuelle session en cours avant la suppression
s.mu.Lock()
if cancel, ok := s.sessionCancels[chatID]; ok {
cancel()
}
delete(s.runningSessions, chatID)
delete(s.sessionCancels, chatID)
s.mu.Unlock()
return s.repoChat.DeleteChat(userID, teamID, chatID)
}
func (s *Service) runQuery(
ctx context.Context,
userID string,
teamID string,
chatID string,
agent *domain.Agent,
provider *domain.Provider,
modelID string,
contextMaxTokens int,
maxResponseTokens int,
messages []*domain.Message,
callback domain.StreamCallback,
) error {
chatSession := &ChatSession{
userID: userID,
teamID: teamID,
chatID: chatID,
agent: agent,
provider: provider,
modelID: modelID,
messages: messages,
contextMaxTokens: contextMaxTokens,
maxResponseTokens: maxResponseTokens,
llmEngine: s.llmEngine,
toolService: s.toolService,
subscribers: make([]domain.StreamCallback, 0),
repoChat: s.repoChat,
}
if callback != nil {
chatSession.subscribers = append(chatSession.subscribers, callback)
}
// contexte annulable : permet d'arrêter la session (ex. suppression du chat)
ctx, cancel := context.WithCancel(ctx)
s.mu.Lock()
s.runningSessions[chatID] = chatSession
s.sessionCancels[chatID] = cancel
s.mu.Unlock()
go func() {
err := chatSession.run(ctx)
if err != nil {
slog.Info("chat session error", "error", err)
}
s.mu.Lock()
delete(s.runningSessions, chatID)
delete(s.sessionCancels, chatID)
s.mu.Unlock()
cancel()
}()
return nil
}
func (s *Service) Subscribe(
userID string,
teamID string,
chatID string,
callback domain.StreamCallback,
) error {
team, err := s.repoUser.FindTeam(userID, teamID)
if err != nil {
return err
}
_, err = s.repoChat.GetChat(userID, team.ID, chatID)
if err != nil {
return err
}
s.mu.Lock()
if chat, ok := s.runningSessions[chatID]; ok {
chat.subscribers = append(chat.subscribers, callback)
}
s.mu.Unlock()
return nil
}