296 lines
6.3 KiB
Go
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
|
|
}
|