basic llm request stream
This commit is contained in:
1 parent
f74e4d1043
commit
a92cc76a9a
66 files changed
+1496
-158
No files matched your search
@@ -0,0 +1,246 @@
|
||||
package chat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"trankilou.fr/lassistanoque/backend/internal/domain"
|
||||
)
|
||||
|
||||
type Service struct {
|
||||
repoUser domain.UserRepository
|
||||
repoAgent domain.AgentRepository
|
||||
repoProvider domain.ProviderRepository
|
||||
repoChat domain.ChatRepository
|
||||
llmEngine domain.LLMEngine
|
||||
runningSessions map[string]*ChatSession
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
func NewService(
|
||||
repoUser domain.UserRepository,
|
||||
repoAgent domain.AgentRepository,
|
||||
repoProvider domain.ProviderRepository,
|
||||
repoChat domain.ChatRepository,
|
||||
llmEngine domain.LLMEngine,
|
||||
) *Service {
|
||||
return &Service{
|
||||
repoUser: repoUser,
|
||||
repoAgent: repoAgent,
|
||||
repoProvider: repoProvider,
|
||||
repoChat: repoChat,
|
||||
llmEngine: llmEngine,
|
||||
runningSessions: make(map[string]*ChatSession),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Service) NewChat(
|
||||
userID string,
|
||||
teamID string,
|
||||
agentID 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,
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var cb domain.StreamCallback
|
||||
if params != nil && params.OnChunk != nil {
|
||||
cb = params.OnChunk
|
||||
}
|
||||
|
||||
s.runQuery(
|
||||
ctx,
|
||||
userID,
|
||||
teamID,
|
||||
chatID,
|
||||
agent,
|
||||
provider,
|
||||
chatModelID,
|
||||
messages,
|
||||
cb,
|
||||
)
|
||||
|
||||
return message, nil
|
||||
}
|
||||
|
||||
func (s *Service) ListChats(userID string, teamID string, page int) ([]*domain.Chat, error) {
|
||||
return s.repoChat.ListChats(userID, teamID, page)
|
||||
}
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &domain.ChatWithMessages{
|
||||
Chat: chat,
|
||||
Messages: messages,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) DeleteChat(userID string, teamID, chatID string) error {
|
||||
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,
|
||||
messages []*domain.Message,
|
||||
callback domain.StreamCallback,
|
||||
) error {
|
||||
|
||||
chatSession := &ChatSession{
|
||||
userID: userID,
|
||||
teamID: teamID,
|
||||
chatID: chatID,
|
||||
agent: agent,
|
||||
provider: provider,
|
||||
modelID: modelID,
|
||||
messages: messages,
|
||||
llmEngine: s.llmEngine,
|
||||
subscribers: make([]domain.StreamCallback, 0),
|
||||
repoChat: s.repoChat,
|
||||
}
|
||||
|
||||
if callback != nil {
|
||||
chatSession.subscribers = append(chatSession.subscribers, callback)
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
s.runningSessions[chatID] = chatSession
|
||||
s.mu.Unlock()
|
||||
|
||||
go func() {
|
||||
_, err := chatSession.run(ctx)
|
||||
if err != nil {
|
||||
slog.Info("chat session error", "error", err)
|
||||
}
|
||||
delete(s.runningSessions, chatID)
|
||||
}()
|
||||
|
||||
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
|
||||
}
|
||||
Reference in new issue
Block a user