From e9762f68d7a432e6c6ad9bd5b5e8cd0be155978b Mon Sep 17 00:00:00 2001 From: unintendedfraud Date: Sat, 6 May 2023 11:24:14 +0200 Subject: [PATCH] add openai capabilities --- env/main.go | 41 ++++++++++++ go.mod | 10 ++- go.sum | 19 ++++++ handlers/message_created.go | 43 ++++++++++++ main.go | 128 +++++++++++++++++------------------- openai/client.go | 89 +++++++++++++++++++++++++ 6 files changed, 260 insertions(+), 70 deletions(-) create mode 100644 env/main.go create mode 100644 handlers/message_created.go create mode 100644 openai/client.go diff --git a/env/main.go b/env/main.go new file mode 100644 index 0000000..4be5b33 --- /dev/null +++ b/env/main.go @@ -0,0 +1,41 @@ +package env + +import ( + "os" + "strings" + + tempest "github.com/Amatsagu/Tempest" +) + +func GetEnvVariables() Env { + if os.Getenv("RAILWAY_ENVIRONMENT") == "production" { + ids := strings.Split(os.Getenv("SERVER_IDS"), ",") + + serverIDs := []tempest.Snowflake{} + for _, id := range ids { + serverIDs = append(serverIDs, tempest.StringToSnowflake(id)) + } + + return Env{ + AppID: tempest.StringToSnowflake(os.Getenv("APP_ID")), + PublicKey: os.Getenv("PUBLIC_KEY"), + Token: os.Getenv("TOKEN"), + Port: "8080", + Addr: os.Getenv("ADDR"), + ServerIDs: serverIDs, + OpenaiHalToken: os.Getenv("OPENAI_HAL"), + } + } + + return Env{} +} + +type Env struct { + AppID tempest.Snowflake + PublicKey string + Token string + Port string + Addr string + ServerIDs []tempest.Snowflake + OpenaiHalToken string +} diff --git a/go.mod b/go.mod index 61f0262..4daba18 100644 --- a/go.mod +++ b/go.mod @@ -2,4 +2,12 @@ module hal go 1.19 -require github.com/Amatsagu/Tempest v1.0.1 // indirect +require ( + github.com/Amatsagu/Tempest v1.0.1 // indirect + github.com/bwmarrin/discordgo v0.27.1 // indirect + github.com/go-resty/resty/v2 v2.7.0 // indirect + github.com/gorilla/websocket v1.4.2 // indirect + golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b // indirect + golang.org/x/net v0.0.0-20211029224645-99673261e6eb // indirect + golang.org/x/sys v0.0.0-20210423082822-04245dca01da // indirect +) diff --git a/go.sum b/go.sum index f64f810..ae46a4d 100644 --- a/go.sum +++ b/go.sum @@ -1,2 +1,21 @@ github.com/Amatsagu/Tempest v1.0.1 h1:PgOPFLNbBevmWwBZuckjuNDuf/WRhVqwp2rjY6qsflE= github.com/Amatsagu/Tempest v1.0.1/go.mod h1:xlvyMhNWe2t/cbyS2pcxb+WR9qtxMvVs2tjmtvLPyLQ= +github.com/bwmarrin/discordgo v0.27.1 h1:ib9AIc/dom1E/fSIulrBwnez0CToJE113ZGt4HoliGY= +github.com/bwmarrin/discordgo v0.27.1/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY= +github.com/go-resty/resty/v2 v2.7.0 h1:me+K9p3uhSmXtrBZ4k9jcEAfJmuC8IivWHwaLZwPrFY= +github.com/go-resty/resty/v2 v2.7.0/go.mod h1:9PWDzw47qPphMRFfhsyk0NnSgvluHcljSMVIq3w7q0I= +github.com/gorilla/websocket v1.4.2 h1:+/TMaTYc4QFitKJxsQ7Yye35DkWvkdLcvGKqM+x0Ufc= +github.com/gorilla/websocket v1.4.2/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= +golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b h1:7mWr3k41Qtv8XlltBkDkl8LoP3mpSgBW8BUoxtEdbXg= +golang.org/x/crypto v0.0.0-20210421170649-83a5a9bb288b/go.mod h1:T9bdIzuCu7OtxOm1hfPfRQxPLYneinmdGuTeoZ9dtd4= +golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= +golang.org/x/net v0.0.0-20211029224645-99673261e6eb h1:pirldcYWx7rx7kE5r+9WsOXPXK0+WH5+uZ7uPmJ44uM= +golang.org/x/net v0.0.0-20211029224645-99673261e6eb/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= +golang.org/x/sys v0.0.0-20201119102817-f84b799fce68 h1:nxC68pudNYkKU6jWhgrqdreuFiOQWj1Fs7T3VrH4Pjw= +golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/sys v0.0.0-20210423082822-04245dca01da h1:b3NXsE2LusjYGGjL5bxEVZZORm/YEFFrWFjR8eFrw/c= +golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= +golang.org/x/term v0.0.0-20201126162022-7de9c90e9dd1/go.mod h1:bj7SfCRtBDWHUb9snDiAeCFNEtKQo2Wmx5Cou7ajbmo= +golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= +golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ= diff --git a/handlers/message_created.go b/handlers/message_created.go new file mode 100644 index 0000000..1b81fe0 --- /dev/null +++ b/handlers/message_created.go @@ -0,0 +1,43 @@ +package handlers + +import ( + "hal/env" + "hal/openai" + "log" + + "github.com/bwmarrin/discordgo" +) + +func OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) { + if m.Author.ID == s.State.User.ID { + return + } + + if containUser(m.Mentions, s.State.User.ID) { + aiclient := openai.NewClient(env.GetEnvVariables().OpenaiHalToken) + + res, err := aiclient.ChatCompletions(m.Content) + if err != nil { + 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].Message.Content + 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()) + } + } + + return + } +} + +func containUser(users []*discordgo.User, userID string) bool { + for _, u := range users { + if u.ID == userID { + return true + } + } + + return false +} diff --git a/main.go b/main.go index 0aff18a..7cb860e 100644 --- a/main.go +++ b/main.go @@ -3,92 +3,82 @@ package main import ( "fmt" "hal/commands" + "hal/env" + "hal/handlers" "log" - "os" - "strings" "time" tempest "github.com/Amatsagu/Tempest" + discordbot "github.com/bwmarrin/discordgo" ) func main() { - env := getEnvVariables() + env := env.GetEnvVariables() - client := tempest.CreateClient(tempest.ClientOptions{ - ApplicationId: env.AppID, - PublicKey: env.PublicKey, - Token: env.Token, - PreCommandExecutionHandler: func(itx tempest.CommandInteraction) *tempest.ResponseData { - log.Printf("running [%s] slash command", itx.Data.Name) - return nil - }, - Cooldowns: &tempest.ClientCooldownOptions{ - Duration: 5 * time.Second, - Ephemeral: true, - CooldownResponse: func(user tempest.User, timeLeft time.Duration) tempest.ResponseData { - return tempest.ResponseData{ - Content: fmt.Sprintf("stop spamming, try again in %.2fs", timeLeft.Seconds()), - } - }, - }, - }) + dgclient, err := discordbot.New(env.Token) + if err != nil { + panic(err) + } + defer dgclient.Close() - if err := initialize(client, env.ServerIDs); err != nil { - panic(err) - } + dgclient.AddHandler(handlers.OnMessageCreated) - client.RegisterCommand(commands.Pinned) - client.RegisterCommand(commands.PsgRefreshFixtures) - client.RegisterCommand(commands.PsgNextMatch) + err = dgclient.Open() + if err != nil { + fmt.Println("error opening connection,", err) + return + } - client.SyncCommands(env.ServerIDs, nil, false) + client := tempest.CreateClient(tempest.ClientOptions{ + ApplicationId: env.AppID, + PublicKey: env.PublicKey, + Token: env.Token, + PreCommandExecutionHandler: func(itx tempest.CommandInteraction) *tempest.ResponseData { + log.Printf("running [%s] slash command", itx.Data.Name) + return nil + }, + Cooldowns: &tempest.ClientCooldownOptions{ + Duration: 5 * time.Second, + Ephemeral: true, + CooldownResponse: func(user tempest.User, timeLeft time.Duration) tempest.ResponseData { + return tempest.ResponseData{ + Content: fmt.Sprintf("stop spamming, try again in %.2fs", timeLeft.Seconds()), + } + }, + }, + }) + if err = initialize(client, env.ServerIDs); err != nil { + panic(err) + } - addr := fmt.Sprintf("%s:%s", env.Addr, env.Port) - fmt.Println("starting server at", addr) + if err = client.RegisterCommand(commands.Pinned); err != nil { + panic(err) + } + if err = client.RegisterCommand(commands.PsgRefreshFixtures); err != nil { + panic(err) + } + if err = client.RegisterCommand(commands.PsgNextMatch); err != nil { + panic(err) + } - if err := client.ListenAndServe(addr); err != nil { - panic(err) - } + if err = client.SyncCommands(env.ServerIDs, nil, false); err != nil { + panic(err) + } + + addr := fmt.Sprintf("%s:%s", env.Addr, env.Port) + fmt.Println("starting server at", addr) + + if err := client.ListenAndServe(addr); err != nil { + panic(err) + } } func initialize(c tempest.Client, serverIDs []tempest.Snowflake) error { - if err := commands.InitPinned(c, serverIDs); err != nil { - return err - } + if err := commands.InitPinned(c, serverIDs); err != nil { + return err + } - return nil -} - -func getEnvVariables() Env { - if os.Getenv("RAILWAY_ENVIRONMENT") == "production" { - ids := strings.Split(os.Getenv("SERVER_IDS"), ",") - - serverIDs := []tempest.Snowflake{} - for _, id := range ids { - serverIDs = append(serverIDs, tempest.StringToSnowflake(id)) - } - - return Env{ - AppID: tempest.StringToSnowflake(os.Getenv("APP_ID")), - PublicKey: os.Getenv("PUBLIC_KEY"), - Token: os.Getenv("TOKEN"), - Port: "8080", - Addr: os.Getenv("ADDR"), - ServerIDs: serverIDs, - } - } - - return Env{} -} - - -type Env struct { - AppID tempest.Snowflake - PublicKey string - Token string - Port string - Addr string - ServerIDs []tempest.Snowflake + return nil } diff --git a/openai/client.go b/openai/client.go new file mode 100644 index 0000000..5fec9f7 --- /dev/null +++ b/openai/client.go @@ -0,0 +1,89 @@ +package openai + +import ( + "encoding/json" + "errors" + "fmt" + + "github.com/go-resty/resty/v2" +) + +type Client struct { + Token string + HttpClient *resty.Client +} + +func NewClient(token string) *Client { + httpclient := resty.New() + httpclient.SetAuthToken(token) + + return &Client{ + Token: token, + HttpClient: httpclient, + } +} + +func (c Client) ChatCompletions(prompt string) (*CompletionResponse, error) { + body := &ChatCompletionPayload{ + Model: "gpt-3.5-turbo", + Prompt: prompt, + MaxTokens: 50, + Temperature: 0.2, + N: 1, + } + + 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 CompletionResponse + 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 +} + +type ChatCompletionPayload struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + MaxTokens int `json:"max_tokens"` + Temperature float32 `json:"temperature"` + TopP float32 `json:"top_p"` + N int `json:"n"` +} + +type CompletionResponse struct { + ID string `json:"id` + Object string `json:"object` + Created int64 `json:"created` + Model string `json:"model` + Usage CompletionUsage `json:"usage"` + Choices []CompletionChoice `json:"choices"` +} + +type CompletionUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +type CompletionChoice struct { + Message CompletionMessage `json:"message"` + FinishReason string `json:"finish_reason"` + Index int `json:"index"` +} + +type CompletionMessage struct { + Role string `json:"role"` + Content string `json:"content"` +}