add openai capabilities
This commit is contained in:
Vendored
+41
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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=
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
Reference in New Issue
Block a user