custom llm

This commit is contained in:
unintendedfraud
2025-12-25 17:58:19 +01:00
parent 8702b87d66
commit d83c326b26
4 changed files with 129 additions and 46 deletions
Vendored
+6
View File
@@ -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
View File
@@ -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, llmRes)
}
sendResponse(s, m.ChannelID, response)
}
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+>`)
+68
View File
@@ -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"`
}
+13 -1
View File
@@ -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)