From d83c326b26065bf576221046e243c55f954cb42b Mon Sep 17 00:00:00 2001 From: unintendedfraud Date: Thu, 25 Dec 2025 17:58:19 +0100 Subject: [PATCH] custom llm --- env/main.go | 6 +++ handlers/message_created.go | 87 ++++++++++++++++++------------------- llm/client.go | 68 +++++++++++++++++++++++++++++ main.go | 14 +++++- 4 files changed, 129 insertions(+), 46 deletions(-) create mode 100644 llm/client.go diff --git a/env/main.go b/env/main.go index 58d2e70..59a4862 100644 --- a/env/main.go +++ b/env/main.go @@ -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 } diff --git a/handlers/message_created.go b/handlers/message_created.go index 810e790..61e17fa 100644 --- a/handlers/message_created.go +++ b/handlers/message_created.go @@ -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())) - return - } - - sendResponse(s, m.ChannelID, fmt.Sprintf("X_X: Réponse vide de Hal... [%s]", string(resBytes))) + if llmRes == "" { + sendResponse(s, m.ChannelID, "X_X: Réponse vide de Hal...") 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+>`) diff --git a/llm/client.go b/llm/client.go new file mode 100644 index 0000000..dbc249c --- /dev/null +++ b/llm/client.go @@ -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"` +} diff --git a/main.go b/main.go index 540a18e..a98b245 100644 --- a/main.go +++ b/main.go @@ -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)