prise en charge des outils + weather

This commit is contained in:
fabien committed 2026-10-04 21:03:51 +02:00
1 parent a36f1231cd
commit 0a6ffab050
14 files changed
+623 -113

No files matched your search

+105
View File
@@ -0,0 +1,105 @@
package cmd
import (
"time"
"trankilou.fr/lassistanoque/backend/internal/adapter/auth/password"
"trankilou.fr/lassistanoque/backend/internal/adapter/database"
"trankilou.fr/lassistanoque/backend/internal/adapter/file"
"trankilou.fr/lassistanoque/backend/internal/adapter/llm"
"trankilou.fr/lassistanoque/backend/internal/adapter/security"
"trankilou.fr/lassistanoque/backend/internal/adapter/tools/weather"
"trankilou.fr/lassistanoque/backend/internal/domain"
"trankilou.fr/lassistanoque/backend/internal/service/auth"
"trankilou.fr/lassistanoque/backend/internal/service/chat"
"trankilou.fr/lassistanoque/backend/internal/service/storage"
)
var db database.Database
var store storage.StorageProvider
var userRepository domain.UserRepository
var providerRepository domain.ProviderRepository
var agentRepository domain.AgentRepository
var chatRepository domain.ChatRepository
var settingsRepository domain.SettingsRepository
var authenticators = make(map[string]auth.Authenticator)
var tokenManager auth.TokenManager
var llmengine domain.LLMEngine
func init() {
var err error
db, err := database.GetDatabase()
if err != nil {
panic(err)
}
err = db.Migrate()
if err != nil {
panic(err)
}
// repositories
userRepository = db.UserRepository()
providerRepository = db.ProviderRepository()
agentRepository = db.AgentRepository()
chatRepository = db.ChatRepository()
settingsRepository = db.SettingsRepository()
// storage
store, err = file.GetStorageProvider(db.FileRepository())
if err != nil {
panic(err)
}
// authentication
tokenManager = security.NewJwtTokenManager(12*time.Hour, 7*24*time.Hour, "lassistanoque")
pwdAuth := password.NewPasswordAuthenticator(tokenManager, GetUserRepository())
authenticators["password"] = pwdAuth
// llm
llmengine = llm.NewAnyLLMEngine()
// llm tools
chat.RegisterTool(&weather.WeatherTool{})
}
func GetStorageProvider() storage.StorageProvider {
return store
}
func GetUserRepository() domain.UserRepository {
return userRepository
}
func GetProviderRepository() domain.ProviderRepository {
return providerRepository
}
func GetAgentRepository() domain.AgentRepository {
return agentRepository
}
func GetChatRepository() domain.ChatRepository {
return chatRepository
}
func GetSettingsRepository() domain.SettingsRepository {
return settingsRepository
}
func GetAuthenticators() map[string]auth.Authenticator {
return authenticators
}
func GetTokenManager() auth.TokenManager {
return tokenManager
}
func GetLLMEngine() domain.LLMEngine {
return llmengine
}
func Close() {
db.Close()
}
+9 -44
View File
@@ -11,11 +11,6 @@ import (
"time" "time"
"github.com/spf13/cobra" "github.com/spf13/cobra"
"trankilou.fr/lassistanoque/backend/internal/adapter/auth/password"
"trankilou.fr/lassistanoque/backend/internal/adapter/database"
"trankilou.fr/lassistanoque/backend/internal/adapter/file"
"trankilou.fr/lassistanoque/backend/internal/adapter/llm"
"trankilou.fr/lassistanoque/backend/internal/adapter/security"
"trankilou.fr/lassistanoque/backend/internal/brain" "trankilou.fr/lassistanoque/backend/internal/brain"
"trankilou.fr/lassistanoque/backend/internal/http" "trankilou.fr/lassistanoque/backend/internal/http"
"trankilou.fr/lassistanoque/backend/internal/service/agent" "trankilou.fr/lassistanoque/backend/internal/service/agent"
@@ -46,60 +41,30 @@ func newServeCmd() *cobra.Command {
} }
func runServe() { func runServe() {
// database
db, err := database.GetDatabase()
if err != nil {
panic(err)
}
defer db.Close()
err = db.Migrate() defer Close()
if err != nil {
panic(err)
}
// storage
store, err := file.GetStorageProvider(db.FileRepository())
if err != nil {
panic(err)
}
// Adapters
tokenManager := security.NewJwtTokenManager(12*time.Hour, 7*24*time.Hour, "lassistanoque")
pwdAuth := password.NewPasswordAuthenticator(tokenManager, db.UserRepository())
llmEngine := llm.NewAnyLLMEngine()
repoUser := db.UserRepository()
repoProvider := db.ProviderRepository()
repoAgent := db.AgentRepository()
repoChat := db.ChatRepository()
// services // services
authService := auth.NewService( authService := auth.NewService(GetSettingsRepository(), GetUserRepository(), GetAuthenticators())
db.SettingsRepository(), userService := user.NewService(GetUserRepository())
db.UserRepository(), storageService := storage.NewService(GetStorageProvider())
map[string]auth.Authenticator{ providerService := provider.NewService(GetProviderRepository(), GetUserRepository(), GetLLMEngine())
"password": pwdAuth, agentService := agent.NewService(GetAgentRepository(), GetUserRepository())
}, chatService := chat.NewService(GetUserRepository(), GetAgentRepository(), GetProviderRepository(), GetChatRepository(), GetLLMEngine())
)
userService := user.NewService(repoUser)
storageService := storage.NewService(store)
providerService := provider.NewService(repoProvider, repoUser, llmEngine)
agentService := agent.NewService(repoAgent, repoUser)
chatService := chat.NewService(repoUser, repoAgent, repoProvider, repoChat, llmEngine)
// http server // http server
httpRouter := http.NewRouter(http.Dependencies{ httpRouter := http.NewRouter(http.Dependencies{
StorageService: storageService, StorageService: storageService,
AuthService: authService, AuthService: authService,
UserService: userService, UserService: userService,
TokenManager: tokenManager,
ProviderService: providerService, ProviderService: providerService,
AgentService: agentService, AgentService: agentService,
ChatService: chatService, ChatService: chatService,
TokenManager: GetTokenManager(),
}) })
brainRouter := brain.NewRouter(brain.Dependencies{ brainRouter := brain.NewRouter(brain.Dependencies{
ChatRepository: db.ChatRepository(), ChatRepository: GetChatRepository(),
}) })
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
@@ -67,18 +67,43 @@ func (r *TursoChatRepository) DeleteChat(userID string, teamID string, id string
) )
} }
func (r *TursoChatRepository) GetChatMessages(userID string, teamID string, chatID string) ([]*domain.Message, error) { func (r *TursoChatRepository) GetChatMessages(userID string, teamID string, chatID string, withToolCallResponses bool) ([]*domain.Message, error) {
return r.messageTable.Select( var messages []*domain.Message
orm.WithWhere( var err error
"chat_id=$1 and team_id=$2 and team_id in (select team_id from user_teams where user_id=$3)", if !withToolCallResponses {
chatID, messages, err = r.messageTable.Select(
teamID, orm.WithWhere(
userID, `chat_id=$1
), and team_id=$2
orm.WithOrder("_date_created asc"), and team_id in (select team_id from user_teams where user_id=$3)
) and (tool_call_id='' or tool_call_id is null)`,
chatID,
teamID,
userID,
),
orm.WithOrder("_date_created asc"),
)
} else {
messages, err = r.messageTable.Select(
orm.WithWhere(
"chat_id=$1 and team_id=$2 and team_id in (select team_id from user_teams where user_id=$3)",
chatID,
teamID,
userID,
),
orm.WithOrder("_date_created asc"),
)
}
if err != nil {
return nil, err
}
for _, m := range messages {
m.Json2toolCall()
}
return messages, nil
} }
func (r *TursoChatRepository) CreateChatMessage(userID string, message *domain.Message) (*domain.Message, error) { func (r *TursoChatRepository) CreateChatMessage(userID string, message *domain.Message) (*domain.Message, error) {
message.ToolCall2Json()
return r.messageTable.Insert(message) return r.messageTable.Insert(message)
} }
+75 -14
View File
@@ -3,6 +3,7 @@ package llm
import ( import (
"context" "context"
"fmt" "fmt"
"log/slog"
anyllm "github.com/mozilla-ai/any-llm-go" anyllm "github.com/mozilla-ai/any-llm-go"
"github.com/mozilla-ai/any-llm-go/providers" "github.com/mozilla-ai/any-llm-go/providers"
@@ -82,11 +83,14 @@ func (e *AnyLLMEngine) Stream(
modelID string, modelID string,
params *domain.LLMParams, params *domain.LLMParams,
messages []*domain.Message, messages []*domain.Message,
) (*domain.Message, error) { ) *domain.Message {
p, err := providerFactory(provider) p, err := providerFactory(provider)
if err != nil { if err != nil {
return nil, err return &domain.Message{
Role: string(domain.RoleAssistant),
Content: err.Error(),
}
} }
anyllmMessages := make([]anyllm.Message, 0) anyllmMessages := make([]anyllm.Message, 0)
@@ -101,18 +105,27 @@ func (e *AnyLLMEngine) Stream(
params.ReasoningEffort = "none" params.ReasoningEffort = "none"
} }
tools := make([]providers.Tool, len(params.Tools))
for i, t := range params.Tools {
tools[i] = toolConvert(t)
}
toolCalling := false
toolCalls := make([]domain.ToolCall, 0)
chunkChan, errChan := p.CompletionStream(ctx, anyllm.CompletionParams{ chunkChan, errChan := p.CompletionStream(ctx, anyllm.CompletionParams{
Model: modelID, Model: modelID,
Messages: anyllmMessages, Messages: anyllmMessages,
Stream: true, Stream: true,
ReasoningEffort: providers.ReasoningEffort(params.ReasoningEffort), ReasoningEffort: providers.ReasoningEffort(params.ReasoningEffort),
Tools: tools,
}) })
fullContent := "" fullContent := ""
role := anyllm.RoleAssistant
for chunk := range chunkChan { for chunk := range chunkChan {
if len(chunk.Choices) > 0 { if len(chunk.Choices) > 0 {
content := chunk.Choices[0].Delta.Content content := chunk.Choices[0].Delta.Content
reasoning := chunk.Choices[0].Delta.Reasoning reasoning := chunk.Choices[0].Delta.Reasoning
if params.OnChunk != nil { if params.OnChunk != nil {
@@ -120,34 +133,82 @@ func (e *AnyLLMEngine) Stream(
fullContent += content fullContent += content
params.OnChunk(&domain.Chunk{ params.OnChunk(&domain.Chunk{
Done: false, Done: false,
Role: role, Role: anyllm.RoleAssistant,
Content: content, Content: content,
Reasoning: false, Reasoning: false,
}) })
} else if reasoning != nil { } else if reasoning != nil {
params.OnChunk(&domain.Chunk{ params.OnChunk(&domain.Chunk{
Done: false, Done: false,
Role: role, Role: anyllm.RoleAssistant,
Content: reasoning.Content, Content: reasoning.Content,
Reasoning: true, Reasoning: true,
}) })
} }
} }
for _, tc := range chunk.Choices[0].Delta.ToolCalls {
toolCalling = true
slog.Info("Model is calling tool", "tool", tc.Function.Name, "arguments", tc.Function.Arguments)
extra := make(map[string]map[string]any)
for k, v := range tc.Extra {
m := map[string]any(v)
extra[k] = m
}
toolCalls = append(toolCalls, domain.ToolCall{
ID: tc.ID,
Type: tc.Type,
Function: domain.ToolCallFunction{
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
},
Extra: extra,
})
params.OnChunk(&domain.Chunk{
Done: false,
ToolName: tc.Function.Name,
Role: anyllm.RoleAssistant,
Content: tc.Function.Arguments,
Reasoning: false,
})
}
} }
} }
params.OnChunk(&domain.Chunk{ if !toolCalling {
Done: true, params.OnChunk(&domain.Chunk{
Role: role, Done: true,
Content: "", Role: anyllm.RoleAssistant,
}) Content: "",
})
}
if err := <-errChan; err != nil { if err := <-errChan; err != nil {
return nil, err return &domain.Message{
Role: string(domain.RoleAssistant),
Content: err.Error(),
}
} }
return &domain.Message{ return &domain.Message{
Role: role, Role: string(domain.RoleAssistant),
Content: fullContent, Content: fullContent,
}, nil ToolCalls: toolCalls,
}
}
func toolConvert(domainTool *domain.ToolDefinition) providers.Tool {
return providers.Tool{
Type: domainTool.Type,
Function: providers.Function{
Name: domainTool.Function.Name,
Description: domainTool.Function.Description,
Parameters: domainTool.Function.Parameters,
},
}
} }
@@ -0,0 +1,183 @@
package weather
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strings"
"trankilou.fr/lassistanoque/backend/internal/domain"
)
type WeatherTool struct {
}
type WeatherConfig struct {
}
func (t *WeatherTool) Name() string {
return "weather"
}
func (t *WeatherTool) Description() string {
return "Get the current weather for a given location"
}
func (t *WeatherTool) Definition(ctx context.Context) *domain.ToolDefinition {
return &domain.ToolDefinition{
Type: "function",
Function: &domain.ToolFunction{
Name: t.Name(),
Description: t.Description(),
Parameters: map[string]any{
"type": "object",
"properties": map[string]any{
"location": map[string]any{
"type": "string",
"description": "The city name, e.g., 'Paris' or 'New York'",
},
},
"required": []string{"location"},
},
},
}
}
func (t *WeatherTool) Execute(ctx context.Context, input []byte) ([]byte, error) {
var params struct {
Location string `json:"location"`
}
err := json.Unmarshal(input, &params)
if err != nil {
return nil, err
}
resp, err := http.Get("https://geocoding-api.open-meteo.com/v1/search?name=" + url.QueryEscape(params.Location))
if err != nil {
return nil, err
}
defer resp.Body.Close()
var location struct {
Results []struct {
Latitude float64 `json:"latitude"`
Longitude float64 `json:"longitude"`
} `json:"results"`
}
err = json.NewDecoder(resp.Body).Decode(&location)
if err != nil {
return nil, err
}
if len(location.Results) == 0 {
return nil, fmt.Errorf("location not found : %s", input)
}
lon := location.Results[0].Longitude
lat := location.Results[0].Latitude
weatherResp, err := http.Get(fmt.Sprintf("https://api.open-meteo.com/v1/forecast?latitude=%f&longitude=%f&current_weather=true&daily=weather_code,temperature_2m_max,temperature_2m_min,rain_sum,wind_speed_10m_max&hourly=temperature_2m&forecast_days=15", lat, lon))
if err != nil {
return nil, err
}
defer weatherResp.Body.Close()
var weatherData struct {
CurrentWeatherUnits struct {
Temperature string `json:"temperature"`
WindSpeed string `json:"windspeed"`
} `json:"current_weather_units"`
CurrentWeather struct {
Time string `json:"time"`
Temperature float64 `json:"temperature"`
WeatherCode int `json:"weathercode"`
} `json:"current_weather"`
DailyUnits struct {
TemperatureMax string `json:"temperature_2m_max"`
TemperatureMin string `json:"temperature_2m_min"`
Rain string `json:"rain_sum"`
WindSpeed string `json:"wind_speed_10m_max"`
} `json:"daily_units"`
Daily struct {
Time []string `json:"time"`
TemperatureMax []float64 `json:"temperature_2m_max"`
TemperatureMin []float64 `json:"temperature_2m_min"`
Rain []float64 `json:"rain_sum"`
WindSpeed []float64 `json:"wind_speed_10m_max"`
} `json:"daily"`
}
err = json.NewDecoder(weatherResp.Body).Decode(&weatherData)
if err != nil {
return nil, err
}
var builder strings.Builder
fmt.Fprintf(&builder, "Current weather in %s\n", params.Location)
fmt.Fprintf(&builder, "- Time: %s\n", weatherData.CurrentWeather.Time)
fmt.Fprintf(&builder, "- Temperature: %.2f %s\n", weatherData.CurrentWeather.Temperature, weatherData.CurrentWeatherUnits.Temperature)
fmt.Fprintf(&builder, "- Weather code: %d - %s\n", weatherData.CurrentWeather.WeatherCode, wmoCodeToDescription(weatherData.CurrentWeather.WeatherCode))
if len(weatherData.Daily.WindSpeed) > 0 {
fmt.Fprintf(&builder, "- Wind speed: %.2f %s\n", weatherData.Daily.WindSpeed[0], weatherData.DailyUnits.WindSpeed)
}
if len(weatherData.Daily.TemperatureMax) > 0 {
fmt.Fprintf(&builder, "- Max temperature: %.2f %s\n", weatherData.Daily.TemperatureMax[0], weatherData.DailyUnits.TemperatureMax)
}
if len(weatherData.Daily.TemperatureMin) > 0 {
fmt.Fprintf(&builder, "- Min temperature: %.2f %s\n", weatherData.Daily.TemperatureMin[0], weatherData.DailyUnits.TemperatureMin)
}
if len(weatherData.Daily.Rain) > 0 {
fmt.Fprintf(&builder, "- Rainfall: %.2f %s\n", weatherData.Daily.Rain[0], weatherData.DailyUnits.Rain)
}
fmt.Fprintf(&builder, "\nForecast for the next days:\n")
// Iterate over the daily data and append to the repor
for i := 0; i < len(weatherData.Daily.Time); i++ {
fmt.Fprintf(&builder, "%s: Max Temp: %.2f %s, Min Temp: %.2f %s, Rain: %.2f %s, Wind Speed: %.2f %s\n",
weatherData.Daily.Time[i],
weatherData.Daily.TemperatureMax[i], weatherData.DailyUnits.TemperatureMax,
weatherData.Daily.TemperatureMin[i], weatherData.DailyUnits.TemperatureMin,
weatherData.Daily.Rain[i], weatherData.DailyUnits.Rain,
weatherData.Daily.WindSpeed[i], weatherData.DailyUnits.WindSpeed)
}
return []byte(builder.String()), nil
}
func wmoCodeToDescription(code int) string {
switch code {
case 0:
return "Clear sky"
case 1, 2, 3:
return "Mainly clear, partly cloudy, and overcast"
case 45, 48:
return "Fog and depositing rime fog"
case 51, 53, 55:
return "Drizzle: Light, moderate, and dense intensity"
case 56, 57:
return "Freezing Drizzle: Light and dense intensity"
case 61, 63, 65:
return "Rain: Slight, moderate and heavy intensity"
case 66, 67:
return "Freezing Rain: Light and heavy intensity"
case 71, 73, 75:
return "Snowfall: Slight, moderate, and heavy intensity"
case 77:
return "Snow grains"
case 80, 81, 82:
return "Rain showers: Slight, moderate, and violent"
case 85, 86:
return "Snow showers slight and heavy"
case 95:
return "Thunderstorm: Slight and heavy intensity"
case 96, 99:
return "Thunderstorm with slight and heavy hail"
default:
return "N/A"
}
}
+54 -7
View File
@@ -2,6 +2,7 @@ package domain
import ( import (
"context" "context"
"encoding/json"
"time" "time"
) )
@@ -33,6 +34,27 @@ type Message struct {
VersionId string `db:"_version" json:"_version"` VersionId string `db:"_version" json:"_version"`
} }
func (msg *Message) ToolCall2Json() error {
if msg.ToolCalls != nil {
toolCallJson, err := json.Marshal(&msg.ToolCalls)
if err != nil {
return err
}
msg.ToolCallsJson = string(toolCallJson)
}
return nil
}
func (msg *Message) Json2toolCall() error {
if msg.ToolCallsJson != "" {
err := json.Unmarshal([]byte(msg.ToolCallsJson), &msg.ToolCalls)
if err != nil {
return err
}
}
return nil
}
type ChatWithMessages struct { type ChatWithMessages struct {
Chat *Chat `json:"chat"` Chat *Chat `json:"chat"`
Messages []*Message `json:"messages"` Messages []*Message `json:"messages"`
@@ -44,17 +66,42 @@ type ChatParams struct {
OnChunk StreamCallback OnChunk StreamCallback
} }
type Tool struct { type Tool interface {
Name() string
Description() string
Definition(ctx context.Context) *ToolDefinition
Execute(ctx context.Context, input []byte) ([]byte, error)
}
type ToolDefinition struct {
Type string
Function *ToolFunction
}
type ToolFunction struct {
Name string
Description string
Parameters map[string]any
} }
type ToolCall struct { type ToolCall struct {
ID string
Type string
Function ToolCallFunction
Extra map[string]map[string]any
}
type ToolCallFunction struct {
Name string
Arguments string
} }
type Chunk struct { type Chunk struct {
Done bool `json:"done"` Done bool `json:"done"`
Role string `json:"role"` ToolName string `json:"toolName,omitempty"`
Content string `json:"content"` Role string `json:"role,omitempty"`
Reasoning bool Content string `json:"content,omitempty"`
Reasoning bool `json:"reasoning,omitempty"`
} }
type StreamCallback func(chunk *Chunk) type StreamCallback func(chunk *Chunk)
@@ -69,7 +116,7 @@ const (
) )
type LLMParams struct { type LLMParams struct {
Tools []*Tool Tools []*ToolDefinition
OnChunk StreamCallback OnChunk StreamCallback
ReasoningEffort string ReasoningEffort string
} }
@@ -83,7 +130,7 @@ type LLMEngine interface {
modelID string, modelID string,
params *LLMParams, params *LLMParams,
messages []*Message, messages []*Message,
) (*Message, error) ) *Message
} }
type ChatRepository interface { type ChatRepository interface {
@@ -92,6 +139,6 @@ type ChatRepository interface {
CreateChat(userID string, chat *Chat) (*Chat, error) CreateChat(userID string, chat *Chat) (*Chat, error)
UpdateChat(userID string, chat *Chat) (*Chat, error) UpdateChat(userID string, chat *Chat) (*Chat, error)
DeleteChat(userID string, teamID string, id string) error DeleteChat(userID string, teamID string, id string) error
GetChatMessages(userID string, teamID string, id string) ([]*Message, error) GetChatMessages(userID string, teamID string, id string, withToolCallResponses bool) ([]*Message, error)
CreateChatMessage(userID string, message *Message) (*Message, error) CreateChatMessage(userID string, message *Message) (*Message, error)
} }
+3 -4
View File
@@ -2,7 +2,6 @@ package handlers
import ( import (
"encoding/json" "encoding/json"
"log/slog"
"net/http" "net/http"
"strconv" "strconv"
@@ -102,12 +101,11 @@ func (h *ChatHandler) NewChatMessage(c *echo.Context) error {
for { for {
chunk := <-chunkChan chunk := <-chunkChan
slog.Info("Chunk", "role", chunk.Role, "content", chunk.Content, "reasoning", chunk.Reasoning) enc.Encode(chunk)
http.NewResponseController(c.Response()).Flush()
if chunk.Done { if chunk.Done {
return nil return nil
} }
enc.Encode(chunk)
http.NewResponseController(c.Response()).Flush()
} }
} }
@@ -130,6 +128,7 @@ func (h *ChatHandler) Get(c *echo.Context) error {
teamID := c.Param("space") teamID := c.Param("space")
chatID := c.Param("id") chatID := c.Param("id")
chatWithMessages, err := h.chatService.GetChatWithMessages(userID, teamID, chatID) chatWithMessages, err := h.chatService.GetChatWithMessages(userID, teamID, chatID)
if err != nil { if err != nil {
c.Logger().Error("error getting chat with messages", "error", err) c.Logger().Error("error getting chat with messages", "error", err)
return echo.NewHTTPError(http.StatusBadRequest, "error getting chat") return echo.NewHTTPError(http.StatusBadRequest, "error getting chat")
+3 -3
View File
@@ -120,7 +120,7 @@ func (s *Service) AddChatMessage(
return nil, err return nil, err
} }
messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID) messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID, true)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -158,7 +158,7 @@ func (s *Service) GetChatWithMessages(userID string, teamID, chatID string) (*do
if err != nil { if err != nil {
return nil, err return nil, err
} }
messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID) messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID, false)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -207,7 +207,7 @@ func (s *Service) runQuery(
s.mu.Unlock() s.mu.Unlock()
go func() { go func() {
_, err := chatSession.run(ctx) err := chatSession.run(ctx)
if err != nil { if err != nil {
slog.Info("chat session error", "error", err) slog.Info("chat session error", "error", err)
} }
+68 -15
View File
@@ -21,10 +21,16 @@ type ChatSession struct {
mu sync.Mutex mu sync.Mutex
} }
func (s *ChatSession) run(ctx context.Context) (*domain.Message, error) { func (s *ChatSession) run(ctx context.Context) error {
params := &domain.LLMParams{} params := &domain.LLMParams{}
tools := make([]*domain.ToolDefinition, 0)
for _, t := range allTools {
tools = append(tools, t.Definition(ctx))
}
params.Tools = tools
s.mu.Lock() s.mu.Lock()
if len(s.subscribers) > 0 { if len(s.subscribers) > 0 {
params.OnChunk = func(chunk *domain.Chunk) { params.OnChunk = func(chunk *domain.Chunk) {
@@ -40,19 +46,66 @@ func (s *ChatSession) run(ctx context.Context) (*domain.Message, error) {
Content: s.agent.SystemPrompt, Content: s.agent.SystemPrompt,
} }
// TODO : loop with tools continue_loop := true
msg, err := s.llmEngine.Stream(
ctx,
s.provider,
s.modelID,
params,
append([]*domain.Message{systemPrompt}, s.messages...),
)
if err != nil {
return nil, err
}
msg.ChatID = s.chatID
msg.TeamID = s.teamID
return s.repoChat.CreateChatMessage(s.userID, msg) // LOOP on TOOLS
for continue_loop {
msg := s.llmEngine.Stream(
ctx,
s.provider,
s.modelID,
params,
append([]*domain.Message{systemPrompt}, s.messages...),
)
msg.ChatID = s.chatID
msg.TeamID = s.teamID
msg, err := s.repoChat.CreateChatMessage(s.userID, msg)
if err != nil {
return err
}
s.messages = append(s.messages, msg)
// TOOL
if len(msg.ToolCalls) > 0 {
continue_loop = true
for _, tc := range msg.ToolCalls {
tool := GetTool(tc.Function.Name)
output, err := tool.Execute(ctx, []byte(tc.Function.Arguments))
var toolResponse *domain.Message
if err != nil {
toolResponse = &domain.Message{
ChatID: s.chatID,
TeamID: s.teamID,
ToolCallID: tc.ID,
Role: string(domain.RoleTool),
Content: "ERROR: " + err.Error(),
}
} else {
toolResponse = &domain.Message{
ChatID: s.chatID,
TeamID: s.teamID,
ToolCallID: tc.ID,
Role: string(domain.RoleTool),
Content: string(output),
}
}
toolResponse, err = s.repoChat.CreateChatMessage(s.userID, toolResponse)
if err != nil {
return err
}
s.messages = append(s.messages, toolResponse)
}
} else {
continue_loop = false
}
// END TOOL
}
// END LOOP
return nil
} }
+14
View File
@@ -0,0 +1,14 @@
package chat
import "trankilou.fr/lassistanoque/backend/internal/domain"
var allTools = make(map[string]domain.Tool)
func RegisterTool(tool domain.Tool) error {
allTools[tool.Name()] = tool
return nil
}
func GetTool(name string) domain.Tool {
return allTools[name]
}
Binary file not shown.
+13 -1
View File
@@ -96,22 +96,34 @@ export interface Message {
id: string id: string
role: string role: string
content: string content: string
reasoning?: string
toolCallId: string toolCallId: string
toolCalls: Array<string> toolcalls: Array<ToolCall>
_date_created: string _date_created: string
_date_updated: string _date_updated: string
_version: string _version: string
} }
export interface ToolCall {
ID: string
Type: string
Function: {
Name: string
Arguments: string
}
}
export type ChatWithMessages = { export type ChatWithMessages = {
chat: Chat chat: Chat
messages: Message[] messages: Message[]
} }
export type Chunk = { export type Chunk = {
toolName: string
done: boolean done: boolean
role: string role: string
content: string content: string
reasoning: boolean
} }
export type ChatRequest = { export type ChatRequest = {
@@ -10,6 +10,7 @@
let prompt = $state("") let prompt = $state("")
let { data }: PageProps = $props(); let { data }: PageProps = $props();
let messages = $state([] as Message[]) let messages = $state([] as Message[])
let waiting = $state(false)
onMount(async ()=>{ onMount(async ()=>{
if (page.state.prompt) { if (page.state.prompt) {
@@ -26,19 +27,54 @@
const sendRequest = async(theprompt:string) => { const sendRequest = async(theprompt:string) => {
prompt = "" prompt = ""
let lastrole="" let lastrole=""
//let reasoning = false let lastToolName=""
waiting = true
for await (const chunk of chatApi.createChatMessage(data.spaceId, data.chatId, { prompt: theprompt })) { for await (const chunk of chatApi.createChatMessage(data.spaceId, data.chatId, { prompt: theprompt })) {
if (chunk.role!=lastrole) { if (waiting && chunk.role==="assistant") {
waiting = false
}
if (chunk.role!=lastrole || chunk.toolName!=lastToolName) {
const currentDate = new Date() const currentDate = new Date()
messages.push({ let m = {
id: ''+currentDate.getTime(), id: ''+currentDate.getTime(),
role: chunk.role, role: chunk.role,
content: chunk.content, content: "",
} as Message) reasoning: "",
} as Message
if (chunk.toolName) {
m.toolcalls = [{
ID: "0",
Type: "Function",
Function: {
Name: chunk.toolName,
Arguments: ""
}
}]
}
if (chunk.reasoning) {
m.reasoning = chunk.content
} else {
m.content = chunk.content
}
messages.push(m)
lastrole=chunk.role lastrole=chunk.role
lastToolName=chunk.toolName
} else { } else {
messages[messages.length-1].content += chunk.content if (chunk.reasoning) {
messages[messages.length-1].reasoning += chunk.content
} else {
messages[messages.length-1].content += chunk.content
}
} }
if (chunk.done) {
document
.querySelectorAll<HTMLDivElement>('div.reasoning:not(.hidden)')
.forEach((el) => el.classList.add('hidden'));
}
historyDiv.scrollTop = historyDiv.scrollHeight; historyDiv.scrollTop = historyDiv.scrollHeight;
} }
} }
@@ -52,11 +88,23 @@
<div id="history" class="flex-1 container overflow-auto" bind:this={historyDiv}> <div id="history" class="flex-1 container overflow-auto" bind:this={historyDiv}>
{#each messages as message (message.id)} {#each messages as message (message.id)}
<div class={["message", "rounded", message.role]}> <div class={["message", "rounded", message.role]}>
<div class="message-content"> {#if message.reasoning }
<SvelteMarkdown source={message.content} /> <div class="message-content reasoning">
<SvelteMarkdown source={message.reasoning} streaming={true}/>
</div> </div>
{/if}
{#if message.toolcalls != null && message.toolcalls.length>0}
<div class="text-green-600">⚒️ Appel de l'outil {message.toolcalls[0].Function.Name}</div>
{:else}
<div class="message-content">
<SvelteMarkdown source={message.content} streaming={true}/>
</div>
{/if}
</div> </div>
{/each} {/each}
{#if waiting}
<div class="animate-pulse">...</div>
{/if}
</div> </div>
<div id="prompt" class="container"> <div id="prompt" class="container">
<form class="request flex w-full" onsubmit={handleOnSubmit}> <form class="request flex w-full" onsubmit={handleOnSubmit}>
+5 -7
View File
@@ -275,21 +275,19 @@ h2 {
/* chat */ /* chat */
.message { .message {
@apply m-2 rounded p-2; @apply rounded p-2;
&.user { &.user {
@apply bg-sky-700 dark:text-white; @apply bg-sky-700 dark:text-white;
} }
&.assistant {
&.reasoning {
@apply text-white/50;
}
}
} }
.message-content { .message-content {
@apply text-sm; @apply text-sm;
&.reasoning {
@apply text-white/50;
}
} }
.message-content p { .message-content p {