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
|
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 h1:PgOPFLNbBevmWwBZuckjuNDuf/WRhVqwp2rjY6qsflE=
|
||||||
github.com/Amatsagu/Tempest v1.0.1/go.mod h1:xlvyMhNWe2t/cbyS2pcxb+WR9qtxMvVs2tjmtvLPyLQ=
|
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 (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
"hal/commands"
|
"hal/commands"
|
||||||
|
"hal/env"
|
||||||
|
"hal/handlers"
|
||||||
"log"
|
"log"
|
||||||
"os"
|
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
tempest "github.com/Amatsagu/Tempest"
|
tempest "github.com/Amatsagu/Tempest"
|
||||||
|
discordbot "github.com/bwmarrin/discordgo"
|
||||||
)
|
)
|
||||||
|
|
||||||
func main() {
|
func main() {
|
||||||
env := getEnvVariables()
|
env := env.GetEnvVariables()
|
||||||
|
|
||||||
client := tempest.CreateClient(tempest.ClientOptions{
|
dgclient, err := discordbot.New(env.Token)
|
||||||
ApplicationId: env.AppID,
|
if err != nil {
|
||||||
PublicKey: env.PublicKey,
|
panic(err)
|
||||||
Token: env.Token,
|
}
|
||||||
PreCommandExecutionHandler: func(itx tempest.CommandInteraction) *tempest.ResponseData {
|
defer dgclient.Close()
|
||||||
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 {
|
dgclient.AddHandler(handlers.OnMessageCreated)
|
||||||
panic(err)
|
|
||||||
}
|
|
||||||
|
|
||||||
client.RegisterCommand(commands.Pinned)
|
err = dgclient.Open()
|
||||||
client.RegisterCommand(commands.PsgRefreshFixtures)
|
if err != nil {
|
||||||
client.RegisterCommand(commands.PsgNextMatch)
|
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)
|
if err = client.RegisterCommand(commands.Pinned); err != nil {
|
||||||
fmt.Println("starting server at", addr)
|
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 {
|
if err = client.SyncCommands(env.ServerIDs, nil, false); err != nil {
|
||||||
panic(err)
|
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 {
|
func initialize(c tempest.Client, serverIDs []tempest.Snowflake) error {
|
||||||
if err := commands.InitPinned(c, serverIDs); err != nil {
|
if err := commands.InitPinned(c, serverIDs); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
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
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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