diff --git a/handlers/message_created.go b/handlers/message_created.go index 2c269ee..92055af 100644 --- a/handlers/message_created.go +++ b/handlers/message_created.go @@ -9,32 +9,66 @@ import ( "github.com/bwmarrin/discordgo" ) -func OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) { - fmt.Printf("\nmessage received: [%s] [%s]", m.ID, m.Message.Content) +const MAX_HISTORY = 50 - if m.Author.ID == s.State.User.ID { +var messagesHistory = []*openai.ChatMessage{} + +func OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) { + isHal := m.Author.ID == s.State.User.ID + + addMessageToHistory(m.Message, isHal) + + if isHal { return } if containUser(m.Mentions, s.State.User.ID) { aiclient := openai.NewClient(env.GetEnvVariables().OpenaiHalToken) - res, err := aiclient.Completions(m.Content) + res, err := aiclient.Chat(messagesHistory) if err != nil { + sendResponse( + s, + m.ChannelID, + fmt.Sprintf("X_X: %s", err.Error()), + ) + log.Panicf("failed to query open ai with the following prompt [%s]. Error: %s", m.Content, err.Error()) } if len(res.Choices) > 0 { - aiResponse := res.Choices[0].Text - if _, err = s.ChannelMessageSend(m.ChannelID, aiResponse); err != nil { - log.Panicf("failed to send the response [%s] to the discord channel [%s]", aiResponse, err.Error()) - } + aiResponse := res.Choices[0].Message.Content + sendResponse(s, m.ChannelID, aiResponse) } return } } +func addMessageToHistory(m *discordgo.Message, isHal bool) { + var role string + if isHal { + role = "system" + } else { + role = "user" + } + + messagesHistory = append(messagesHistory, &openai.ChatMessage{ + Role: role, + Content: m.Content, + }) + + if len(messagesHistory) > MAX_HISTORY { + messagesHistory = messagesHistory[1:] + } +} + +func sendResponse(s *discordgo.Session, channelID string, response string) { + if _, err := s.ChannelMessageSend(channelID, response); err != nil { + log.Panicf("failed to send the response [%s] to the discord channel [%s]", response, err.Error()) + } +} + func containUser(users []*discordgo.User, userID string) bool { for _, u := range users { if u.ID == userID { diff --git a/openai/client.go b/openai/client.go index 8bf25b8..ac28a06 100644 --- a/openai/client.go +++ b/openai/client.go @@ -8,6 +8,8 @@ import ( resty "github.com/go-resty/resty/v2" ) +const MAX_TOKENS = 250 + type Client struct { HttpClient *resty.Client } @@ -19,6 +21,34 @@ func NewClient(token string) *Client { return &Client{HttpClient: httpclient} } +func (c Client) Chat(messages []*ChatMessage) (*ChatResponse, error) { + body := &ChatPayload{ + Messages: messages, + Model: "gpt-4-1106-preview", + MaxTokens: MAX_TOKENS, + } + + b, err := json.Marshal(body) + if err != nil { + return nil, fmt.Errorf("failed to marshal the payload. %s", err.Error()) + } + + res, err := c.HttpClient.NewRequest(). + SetHeader("Content-Type", "application/json"). + SetBody(b). + Post("https://api.openai.com/v1/chat/completions") + if err != nil { + return nil, fmt.Errorf("failed to query openai. Error: %s", err.Error()) + } + + var r ChatResponse + if err := json.Unmarshal(res.Body(), &r); err != nil { + return nil, fmt.Errorf("failed to parse the openai response. Error: %s", err.Error()) + } + + return &r, nil +} + func (c Client) Completions(prompt string) (*CompletionResponse, error) { body := &CompletionPayload{ Model: "text-davinci-003", @@ -54,6 +84,32 @@ func cleanPrompt(p string) string { return regex.ReplaceAllString(p, "") } +type ChatResponse struct { + ID string `json:"id"` + Choices []ChatResponseChoice `json:"choices"` +} + +type ChatResponseChoice struct { + Message ChatResponseChoiceMessage `json:"message"` +} + +type ChatResponseChoiceMessage struct { + Content string `json:"content"` +} + +type ChatPayload struct { + Messages []*ChatMessage `json:"messages"` + Model string `json:"string"` + MaxTokens int `json:"max_tokens"` +} + +type ChatMessage struct { + Role string `json:"role"` + Content string `json:"content"` +} + +// ----- + type CompletionPayload struct { Model string `json:"model"` Prompt string `json:"prompt"`