58 lines
1.1 KiB
Go
58 lines
1.1 KiB
Go
package chat
|
|
|
|
import (
|
|
"context"
|
|
"sync"
|
|
|
|
"trankilou.fr/lassistanoque/backend/internal/domain"
|
|
)
|
|
|
|
type ChatSession struct {
|
|
userID string
|
|
teamID string
|
|
chatID string
|
|
agent *domain.Agent
|
|
provider *domain.Provider
|
|
modelID string
|
|
messages []*domain.Message
|
|
llmEngine domain.LLMEngine
|
|
repoChat domain.ChatRepository
|
|
subscribers []domain.StreamCallback
|
|
mu sync.Mutex
|
|
}
|
|
|
|
func (s *ChatSession) run(ctx context.Context) (*domain.Message, error) {
|
|
|
|
params := &domain.LLMParams{}
|
|
|
|
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: s.agent.SystemPrompt,
|
|
}
|
|
|
|
msg, err := s.llmEngine.Stream(
|
|
ctx,
|
|
s.provider,
|
|
s.modelID,
|
|
params,
|
|
append([]*domain.Message{systemPrompt}, s.messages...),
|
|
)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
msg.ChatID = s.chatID
|
|
msg.TeamID = s.teamID
|
|
|
|
return s.repoChat.CreateChatMessage(s.userID, msg)
|
|
}
|