custom llm
This commit is contained in:
Vendored
+6
@@ -30,6 +30,9 @@ func GetEnvVariables() Env {
|
|||||||
|
|
||||||
GeminiToken: os.Getenv("GEMINI_API_KEY"),
|
GeminiToken: os.Getenv("GEMINI_API_KEY"),
|
||||||
HalResponsePercent: halResPercent,
|
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
|
OpenaiHalToken string
|
||||||
GeminiToken string
|
GeminiToken string
|
||||||
HalResponsePercent int
|
HalResponsePercent int
|
||||||
|
LlmEndpoint string
|
||||||
|
LlmToken string
|
||||||
|
LlmModel string
|
||||||
}
|
}
|
||||||
|
|||||||
+42
-45
@@ -7,7 +7,9 @@ import (
|
|||||||
"regexp"
|
"regexp"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
"hal/env"
|
||||||
"hal/gemini"
|
"hal/gemini"
|
||||||
|
"hal/llm"
|
||||||
"hal/openai"
|
"hal/openai"
|
||||||
|
|
||||||
"github.com/bwmarrin/discordgo"
|
"github.com/bwmarrin/discordgo"
|
||||||
@@ -32,7 +34,10 @@ var spams = []string{
|
|||||||
}
|
}
|
||||||
|
|
||||||
// messagesHistory []*openai.ChatMessage = []*openai.ChatMessage{}
|
// 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{}
|
var usersHistoryCount map[string]*userHistoryCount = map[string]*userHistoryCount{}
|
||||||
|
|
||||||
@@ -45,21 +50,21 @@ type userHistoryCount struct {
|
|||||||
type Handler struct {
|
type Handler struct {
|
||||||
openaiClient *openai.Client
|
openaiClient *openai.Client
|
||||||
geminiClient *gemini.Client
|
geminiClient *gemini.Client
|
||||||
|
llmClient *llm.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
func Init(openaiToken string, geminiToken string) Handler {
|
func Init(env *env.Env) Handler {
|
||||||
return Handler{
|
return Handler{
|
||||||
openaiClient: openai.NewClient(openaiToken),
|
openaiClient: openai.NewClient(env.OpenaiHalToken),
|
||||||
geminiClient: gemini.NewClient(geminiToken),
|
geminiClient: gemini.NewClient(env.GeminiToken),
|
||||||
|
llmClient: llm.NewClient(env),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Tramp: 161970441441902592
|
|
||||||
|
|
||||||
func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) {
|
func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCreate) {
|
||||||
isHal := m.Author.ID == s.State.User.ID
|
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) {
|
if isHal || !containHal(m.Mentions, s.State.User.ID) {
|
||||||
return
|
return
|
||||||
@@ -72,7 +77,7 @@ func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCrea
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
llmRes, err := h.geminiClient.GenerateContent(geminiHistory)
|
llmRes, err := h.llmClient.GenerateContent(cleanMessage(m.Content))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
sendResponse(
|
sendResponse(
|
||||||
s,
|
s,
|
||||||
@@ -84,48 +89,40 @@ func (h Handler) OnMessageCreated(s *discordgo.Session, m *discordgo.MessageCrea
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
response := llmRes.Text()
|
if llmRes == "" {
|
||||||
|
sendResponse(s, m.ChannelID, "X_X: Réponse vide de Hal...")
|
||||||
if response == "" {
|
|
||||||
resBytes, err := llmRes.MarshalJSON()
|
|
||||||
if err != nil {
|
|
||||||
sendResponse(s, m.ChannelID, fmt.Sprintf("X_X: %s", err.Error()))
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
sendResponse(s, m.ChannelID, fmt.Sprintf("X_X: Réponse vide de Hal... [%s]", string(resBytes)))
|
sendResponse(s, m.ChannelID, llmRes)
|
||||||
return
|
|
||||||
}
|
}
|
||||||
|
|
||||||
sendResponse(s, m.ChannelID, response)
|
// func addMessageToHistory(m *discordgo.Message, isHal bool) []*genai.Content {
|
||||||
}
|
// message := cleanMessage(m.Content)
|
||||||
|
// if message == "" {
|
||||||
func addMessageToHistory(m *discordgo.Message, isHal bool) []*genai.Content {
|
// return geminiHistory
|
||||||
message := cleanMessage(m.Content)
|
// }
|
||||||
if message == "" {
|
//
|
||||||
return geminiHistory
|
// var role string
|
||||||
}
|
// if isHal {
|
||||||
|
// role = "model"
|
||||||
var role string
|
// } else {
|
||||||
if isHal {
|
// role = "user"
|
||||||
role = "model"
|
// }
|
||||||
} else {
|
//
|
||||||
role = "user"
|
// geminiHistory = append(geminiHistory, &genai.Content{
|
||||||
}
|
// Role: role,
|
||||||
|
// Parts: []*genai.Part{
|
||||||
geminiHistory = append(geminiHistory, &genai.Content{
|
// genai.NewPartFromText(message),
|
||||||
Role: role,
|
// },
|
||||||
Parts: []*genai.Part{
|
// })
|
||||||
genai.NewPartFromText(message),
|
//
|
||||||
},
|
// if len(geminiHistory) > MAX_HISTORY {
|
||||||
})
|
// geminiHistory = geminiHistory[1:]
|
||||||
|
// }
|
||||||
if len(geminiHistory) > MAX_HISTORY {
|
//
|
||||||
geminiHistory = geminiHistory[1:]
|
// return geminiHistory
|
||||||
}
|
// }
|
||||||
|
|
||||||
return geminiHistory
|
|
||||||
}
|
|
||||||
|
|
||||||
func cleanMessage(p string) string {
|
func cleanMessage(p string) string {
|
||||||
regex := regexp.MustCompile(`<@\d+>`)
|
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/env"
|
||||||
"hal/handlers"
|
"hal/handlers"
|
||||||
|
"hal/llm"
|
||||||
|
|
||||||
discordbot "github.com/bwmarrin/discordgo"
|
discordbot "github.com/bwmarrin/discordgo"
|
||||||
"github.com/joho/godotenv"
|
"github.com/joho/godotenv"
|
||||||
@@ -19,7 +20,7 @@ func main() {
|
|||||||
fmt.Println("HAL started")
|
fmt.Println("HAL started")
|
||||||
env := env.GetEnvVariables()
|
env := env.GetEnvVariables()
|
||||||
|
|
||||||
handler := handlers.Init(env.OpenaiHalToken, env.GeminiToken)
|
handler := handlers.Init(&env)
|
||||||
|
|
||||||
discordToken := fmt.Sprintf("Bot %s", env.Token)
|
discordToken := fmt.Sprintf("Bot %s", env.Token)
|
||||||
dgclient, err := discordbot.New(discordToken)
|
dgclient, err := discordbot.New(discordToken)
|
||||||
@@ -35,6 +36,17 @@ func main() {
|
|||||||
return
|
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.
|
// Wait here until CTRL-C or other term signal is received.
|
||||||
sc := make(chan os.Signal, 1)
|
sc := make(chan os.Signal, 1)
|
||||||
signal.Notify(sc, syscall.SIGINT, syscall.SIGTERM, os.Interrupt)
|
signal.Notify(sc, syscall.SIGINT, syscall.SIGTERM, os.Interrupt)
|
||||||
|
|||||||
Reference in New Issue
Block a user