custom llm
This commit is contained in:
Vendored
+6
@@ -30,6 +30,9 @@ func GetEnvVariables() Env {
|
||||
|
||||
GeminiToken: os.Getenv("GEMINI_API_KEY"),
|
||||
HalResponsePercent: halResPercent,
|
||||
LlmEndpoint: os.Getenv("LLM_ENDPOINT"),
|
||||
LlmToken: os.Getenv("LLM_TOKEN"),
|
||||
LlmModel: os.Getenv("LLM_MODEL"),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,4 +44,7 @@ type Env struct {
|
||||
OpenaiHalToken string
|
||||
GeminiToken string
|
||||
HalResponsePercent int
|
||||
LlmEndpoint string
|
||||
LlmToken string
|
||||
LlmModel string
|
||||
}
|
||||
|
||||
+42
-45
@@ -7,7 +7,9 @@ import (
|
||||
"regexp"
|
||||
"time"
|
||||
|
||||
"hal/env"
|
||||
"hal/gemini"
|
||||
"hal/llm"
|
||||
"hal/openai"
|
||||
|
||||
"github.com/bwmarrin/discordgo"
|
||||
@@ -32,7 +34,10 @@ var spams = []string{
|
||||
}
|
||||
|
||||
// messagesHistory []*openai.ChatMessage = []*openai.ChatMessage{}
|
||||
var geminiHistory []*genai.Content = []*genai.Content{}
|
||||
var (
|
||||
geminiHistory []*genai.Content = []*genai.Content{}
|
||||
llmHistory []string = []string{}
|
||||
)
|
||||
|
||||
var usersHistoryCount map[string]*userHistoryCount = map[string]*userHistoryCount{}
|
||||
|
||||
@@ -45,21 +50,21 @@ type userHistoryCount struct {
|
||||
type Handler struct {
|
||||
openaiClient *openai.Client
|
||||
geminiClient *gemini.Client
|
||||
llmClient *llm.Client
|
||||
}
|
||||
|
||||
func Init(openaiToken string, geminiToken string) Handler {
|
||||
func Init(env *env.Env) Handler {
|
||||
return Handler{
|
||||
openaiClient: openai.NewClient(openaiToken),
|
||||
geminiClient: gemini.NewClient(geminiToken),
|
||||
openaiClient: openai.NewClient(env.OpenaiHalToken),
|
||||
geminiClient: gemini.NewClient(env.GeminiToken),
|
||||
llmClient: llm.NewClient(env),
|
||||
}
|
||||
}
|
||||
|
||||
// Tramp: 161970441441902592
|
||||
|
||||
func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) {
|
||||
isHal := m.Author.ID == s.State.User.ID
|
||||
|
||||
addMessageToHistory(m.Message, isHal)
|
||||
// addMessageToHistory(m.Message, isHal)
|
||||
|
||||
if isHal || !containHal(m.Mentions, s.State.User.ID) {
|
||||
return
|
||||
@@ -72,7 +77,7 @@ func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCrea
|
||||
return
|
||||
}
|
||||
|
||||
llmRes, err := h.geminiClient.GenerateContent(geminiHistory)
|
||||
llmRes, err := h.llmClient.GenerateContent(cleanMessage(m.Content))
|
||||
if err != nil {
|
||||
sendResponse(
|
||||
s,
|
||||
@@ -84,48 +89,40 @@ func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCrea
|
||||
return
|
||||
}
|
||||
|
||||
response := llmRes.Text()
|
||||
|
||||
if response == "" {
|
||||
resBytes, err := llmRes.MarshalJSON()
|
||||
if err != nil {
|
||||
sendResponse(s, m.ChannelID, fmt.Sprintf("X_X: %s", err.Error()))
|
||||
if llmRes == "" {
|
||||
sendResponse(s, m.ChannelID, "X_X: Réponse vide de Hal...")
|
||||
return
|
||||
}
|
||||
|
||||
sendResponse(s, m.ChannelID, fmt.Sprintf("X_X: Réponse vide de Hal... [%s]", string(resBytes)))
|
||||
return
|
||||
}
|
||||
|
||||
sendResponse(s, m.ChannelID, response)
|
||||
sendResponse(s, m.ChannelID, llmRes)
|
||||
}
|
||||
|
||||
func addMessageToHistory(m *discordgo.Message, isHal bool) []*genai.Content {
|
||||
message := cleanMessage(m.Content)
|
||||
if message == "" {
|
||||
return geminiHistory
|
||||
}
|
||||
|
||||
var role string
|
||||
if isHal {
|
||||
role = "model"
|
||||
} else {
|
||||
role = "user"
|
||||
}
|
||||
|
||||
geminiHistory = append(geminiHistory, &genai.Content{
|
||||
Role: role,
|
||||
Parts: []*genai.Part{
|
||||
genai.NewPartFromText(message),
|
||||
},
|
||||
})
|
||||
|
||||
if len(geminiHistory) > MAX_HISTORY {
|
||||
geminiHistory = geminiHistory[1:]
|
||||
}
|
||||
|
||||
return geminiHistory
|
||||
}
|
||||
// func addMessageToHistory(m *discordgo.Message, isHal bool) []*genai.Content {
|
||||
// message := cleanMessage(m.Content)
|
||||
// if message == "" {
|
||||
// return geminiHistory
|
||||
// }
|
||||
//
|
||||
// var role string
|
||||
// if isHal {
|
||||
// role = "model"
|
||||
// } else {
|
||||
// role = "user"
|
||||
// }
|
||||
//
|
||||
// geminiHistory = append(geminiHistory, &genai.Content{
|
||||
// Role: role,
|
||||
// Parts: []*genai.Part{
|
||||
// genai.NewPartFromText(message),
|
||||
// },
|
||||
// })
|
||||
//
|
||||
// if len(geminiHistory) > MAX_HISTORY {
|
||||
// geminiHistory = geminiHistory[1:]
|
||||
// }
|
||||
//
|
||||
// return geminiHistory
|
||||
// }
|
||||
|
||||
func cleanMessage(p string) string {
|
||||
regex := regexp.MustCompile(`<@\d+>`)
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
package llm
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"hal/env"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
http http.Client
|
||||
token string
|
||||
model string
|
||||
endpoint string
|
||||
}
|
||||
|
||||
func NewClient(env *env.Env) *Client {
|
||||
return &Client{
|
||||
http: http.Client{Timeout: 10 * time.Second},
|
||||
token: env.LlmToken,
|
||||
model: env.LlmModel,
|
||||
endpoint: env.LlmEndpoint,
|
||||
}
|
||||
}
|
||||
|
||||
func (client Client) GenerateContent(message string) (string, error) {
|
||||
payload := map[string]any{
|
||||
"model": client.model,
|
||||
"stream": false,
|
||||
"prompt": message,
|
||||
}
|
||||
|
||||
payloadBytes, _ := json.Marshal(payload)
|
||||
|
||||
res, err := client.http.Post(
|
||||
client.endpoint,
|
||||
"application/json",
|
||||
bytes.NewBuffer(payloadBytes),
|
||||
)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer res.Body.Close()
|
||||
|
||||
if res.StatusCode != 200 {
|
||||
return "", fmt.Errorf("LLM returned status code: [%d]", res.StatusCode)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(res.Body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read the llm response body: %w", err)
|
||||
}
|
||||
|
||||
var llmResponse LLMResponse
|
||||
if err := json.Unmarshal(body, &llmResponse); err != nil {
|
||||
return "", fmt.Errorf("failed to unmarshal llm response: %w", err)
|
||||
}
|
||||
|
||||
return llmResponse.Response, nil
|
||||
}
|
||||
|
||||
type LLMResponse struct {
|
||||
Response string `json:"response"`
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
|
||||
"hal/env"
|
||||
"hal/handlers"
|
||||
"hal/llm"
|
||||
|
||||
discordbot "github.com/bwmarrin/discordgo"
|
||||
"github.com/joho/godotenv"
|
||||
@@ -19,7 +20,7 @@ func main() {
|
||||
fmt.Println("HAL started")
|
||||
env := env.GetEnvVariables()
|
||||
|
||||
handler := handlers.Init(env.OpenaiHalToken, env.GeminiToken)
|
||||
handler := handlers.Init(&env)
|
||||
|
||||
discordToken := fmt.Sprintf("Bot %s", env.Token)
|
||||
dgclient, err := discordbot.New(discordToken)
|
||||
@@ -35,6 +36,17 @@ func main() {
|
||||
return
|
||||
}
|
||||
|
||||
// ---
|
||||
llmClient := llm.NewClient(&env)
|
||||
res, err := llmClient.GenerateContent("Quelle est la capitale de la France?")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
|
||||
fmt.Println("### RES ###")
|
||||
fmt.Println(res)
|
||||
// ---
|
||||
|
||||
// Wait here until CTRL-C or other term signal is received.
|
||||
sc := make(chan os.Signal, 1)
|
||||
signal.Notify(sc, syscall.SIGINT, syscall.SIGTERM, os.Interrupt)
|
||||
|
||||
Reference in New Issue
Block a user