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"), 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
View File
@@ -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
}
sendResponse(s, m.ChannelID, fmt.Sprintf("X_X: Réponse vide de Hal... [%s]", string(resBytes)))
return return
} }
sendResponse(s, m.ChannelID, response) sendResponse(s, m.ChannelID, llmRes)
} }
func addMessageToHistory(m *discordgo.Message, isHal bool) []*genai.Content { // func addMessageToHistory(m *discordgo.Message, isHal bool) []*genai.Content {
message := cleanMessage(m.Content) // message := cleanMessage(m.Content)
if message == "" { // if message == "" {
return geminiHistory // return geminiHistory
} // }
//
var role string // var role string
if isHal { // if isHal {
role = "model" // role = "model"
} else { // } else {
role = "user" // role = "user"
} // }
//
geminiHistory = append(geminiHistory, &genai.Content{ // geminiHistory = append(geminiHistory, &genai.Content{
Role: role, // Role: role,
Parts: []*genai.Part{ // Parts: []*genai.Part{
genai.NewPartFromText(message), // genai.NewPartFromText(message),
}, // },
}) // })
//
if len(geminiHistory) > MAX_HISTORY { // if len(geminiHistory) > MAX_HISTORY {
geminiHistory = geminiHistory[1:] // 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+>`)
+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/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)