diff --git a/handlers/message_created.go b/handlers/message_created.go index 1bb646c..2b05647 100644 --- a/handlers/message_created.go +++ b/handlers/message_created.go @@ -4,55 +4,88 @@ import ( "fmt" "hal/openai" "log" + "math/rand" "regexp" + "time" "github.com/bwmarrin/discordgo" ) const MAX_HISTORY = 50 -var messagesHistory = []*openai.ChatMessage{} +const SPAM_PERIOD = 5 * time.Minute + +var spams = []string{ + ":GroGroDebile:", + "ArrĂȘte de spam putain!", + "Wesh...", + "...", + "Flemme.", + "Laisse-moi tranquille!", + "Fdr", + ":soxx:", +} + +type userHistoryCount struct { + date time.Time + bannedUntilAt time.Time + count int +} type Handler struct { + messagesHistory []*openai.ChatMessage + usersHistoryCount map[string]userHistoryCount + client *openai.Client } func Init(token string) Handler { return Handler{ client: openai.NewClient(token), + + messagesHistory: []*openai.ChatMessage{}, + usersHistoryCount: map[string]userHistoryCount{}, } } +// Tramp: 161970441441902592 + func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) { isHal := m.Author.ID == s.State.User.ID - addMessageToHistory(m.Message, isHal) + h.addMessageToHistory(m.Message, isHal) - if isHal { + if isHal || !containHal(m.Mentions, s.State.User.ID) { return } - if containUser(m.Mentions, s.State.User.ID) { - res, err := h.client.Chat(messagesHistory) - if err != nil { - sendResponse( - s, - m.ChannelID, - fmt.Sprintf("X_X: %s", err.Error()), - ) + userSpamTooMuch := h.updateUserHistoryCount(m.Author.ID) - log.Printf("\nfailed to query open ai with the following prompt [%s]. Error: %s", m.Content, err.Error()) - return - } + if userSpamTooMuch { + sendResponse(s, m.ChannelID, getRandomSpam()) + } - if len(res.Choices) > 0 { - aiResponse := res.Choices[0].Message.Content - sendResponse(s, m.ChannelID, aiResponse) - } + res, err := h.client.Chat(h.messagesHistory) + if err != nil { + sendResponse( + s, + m.ChannelID, + fmt.Sprintf("X_X: %s", err.Error()), + ) + + log.Printf("\nfailed to query open ai with the following prompt [%s]. Error: %s", m.Content, err.Error()) + return + } + + if len(res.Choices) > 0 { + aiResponse := res.Choices[0].Message.Content + sendResponse(s, m.ChannelID, aiResponse) + } else { + sendResponse(s, m.ChannelID, "la fatigue") } } -func addMessageToHistory(m *discordgo.Message, isHal bool) { +func (h Handler) addMessageToHistory(m *discordgo.Message, isHal bool) { var role string if isHal { role = "system" @@ -60,13 +93,13 @@ func addMessageToHistory(m *discordgo.Message, isHal bool) { role = "user" } - messagesHistory = append(messagesHistory, &openai.ChatMessage{ + h.messagesHistory = append(h.messagesHistory, &openai.ChatMessage{ Role: role, Content: cleanMessage(m.Content), }) - if len(messagesHistory) > MAX_HISTORY { - messagesHistory = messagesHistory[1:] + if len(h.messagesHistory) > MAX_HISTORY { + h.messagesHistory = h.messagesHistory[1:] } } @@ -81,7 +114,7 @@ func sendResponse(s *discordgo.Session, channelID string, response string) { } } -func containUser(users []*discordgo.User, userID string) bool { +func containHal(users []*discordgo.User, userID string) bool { for _, u := range users { if u.ID == userID { return true @@ -90,3 +123,40 @@ func containUser(users []*discordgo.User, userID string) bool { return false } + +func (h Handler) updateUserHistoryCount(userID string) bool { + now := time.Now() + + u, ok := h.usersHistoryCount[userID] + if !ok { + h.usersHistoryCount[userID] = userHistoryCount{ + date: now, + count: 1, + } + + return false + } + + u.count++ + + if now.Before(u.bannedUntilAt) { + return true + } + + if u.date.Add(SPAM_PERIOD).After(now) { + u.count = 1 + u.date = now + + return false + } + + if u.count > 5 { + return true + } + + return false +} + +func getRandomSpam() string { + return spams[rand.Intn(len(spams))] +}