diff --git a/backend/cmd/dip.go b/backend/cmd/dip.go new file mode 100644 index 0000000..746aecd --- /dev/null +++ b/backend/cmd/dip.go @@ -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() +} diff --git a/backend/cmd/serve.go b/backend/cmd/serve.go index d05d022..44d3dcf 100644 --- a/backend/cmd/serve.go +++ b/backend/cmd/serve.go @@ -11,11 +11,6 @@ import ( "time" "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/http" "trankilou.fr/lassistanoque/backend/internal/service/agent" @@ -46,60 +41,30 @@ func newServeCmd() *cobra.Command { } func runServe() { - // database - db, err := database.GetDatabase() - if err != nil { - panic(err) - } - defer db.Close() - err = db.Migrate() - 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() + defer Close() // services - authService := auth.NewService( - db.SettingsRepository(), - db.UserRepository(), - map[string]auth.Authenticator{ - "password": pwdAuth, - }, - ) - 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) + authService := auth.NewService(GetSettingsRepository(), GetUserRepository(), GetAuthenticators()) + userService := user.NewService(GetUserRepository()) + storageService := storage.NewService(GetStorageProvider()) + providerService := provider.NewService(GetProviderRepository(), GetUserRepository(), GetLLMEngine()) + agentService := agent.NewService(GetAgentRepository(), GetUserRepository()) + chatService := chat.NewService(GetUserRepository(), GetAgentRepository(), GetProviderRepository(), GetChatRepository(), GetLLMEngine()) // http server httpRouter := http.NewRouter(http.Dependencies{ StorageService: storageService, AuthService: authService, UserService: userService, - TokenManager: tokenManager, ProviderService: providerService, AgentService: agentService, ChatService: chatService, + TokenManager: GetTokenManager(), }) brainRouter := brain.NewRouter(brain.Dependencies{ - ChatRepository: db.ChatRepository(), + ChatRepository: GetChatRepository(), }) ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) diff --git a/backend/internal/adapter/database/turso/chat_repository.go b/backend/internal/adapter/database/turso/chat_repository.go index 25994a5..6f2e109 100644 --- a/backend/internal/adapter/database/turso/chat_repository.go +++ b/backend/internal/adapter/database/turso/chat_repository.go @@ -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) { - return 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"), - ) +func (r *TursoChatRepository) GetChatMessages(userID string, teamID string, chatID string, withToolCallResponses bool) ([]*domain.Message, error) { + var messages []*domain.Message + var err error + if !withToolCallResponses { + 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) + 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) { + message.ToolCall2Json() return r.messageTable.Insert(message) } diff --git a/backend/internal/adapter/llm/anyllm.go b/backend/internal/adapter/llm/anyllm.go index 7d30ddd..f39c0ee 100644 --- a/backend/internal/adapter/llm/anyllm.go +++ b/backend/internal/adapter/llm/anyllm.go @@ -3,6 +3,7 @@ package llm import ( "context" "fmt" + "log/slog" anyllm "github.com/mozilla-ai/any-llm-go" "github.com/mozilla-ai/any-llm-go/providers" @@ -82,11 +83,14 @@ func (e *AnyLLMEngine) Stream( modelID string, params *domain.LLMParams, messages []*domain.Message, -) (*domain.Message, error) { +) *domain.Message { p, err := providerFactory(provider) if err != nil { - return nil, err + return &domain.Message{ + Role: string(domain.RoleAssistant), + Content: err.Error(), + } } anyllmMessages := make([]anyllm.Message, 0) @@ -101,18 +105,27 @@ func (e *AnyLLMEngine) Stream( 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{ Model: modelID, Messages: anyllmMessages, Stream: true, ReasoningEffort: providers.ReasoningEffort(params.ReasoningEffort), + Tools: tools, }) fullContent := "" - role := anyllm.RoleAssistant for chunk := range chunkChan { if len(chunk.Choices) > 0 { + content := chunk.Choices[0].Delta.Content reasoning := chunk.Choices[0].Delta.Reasoning if params.OnChunk != nil { @@ -120,34 +133,82 @@ func (e *AnyLLMEngine) Stream( fullContent += content params.OnChunk(&domain.Chunk{ Done: false, - Role: role, + Role: anyllm.RoleAssistant, Content: content, Reasoning: false, }) } else if reasoning != nil { params.OnChunk(&domain.Chunk{ Done: false, - Role: role, + Role: anyllm.RoleAssistant, Content: reasoning.Content, 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{ - Done: true, - Role: role, - Content: "", - }) + if !toolCalling { + params.OnChunk(&domain.Chunk{ + Done: true, + Role: anyllm.RoleAssistant, + Content: "", + }) + } if err := <-errChan; err != nil { - return nil, err + return &domain.Message{ + Role: string(domain.RoleAssistant), + Content: err.Error(), + } } return &domain.Message{ - Role: role, - Content: fullContent, - }, nil + Role: string(domain.RoleAssistant), + Content: fullContent, + 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, + }, + } } diff --git a/backend/internal/adapter/tools/weather/tool.go b/backend/internal/adapter/tools/weather/tool.go new file mode 100644 index 0000000..33e0e21 --- /dev/null +++ b/backend/internal/adapter/tools/weather/tool.go @@ -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, ¶ms) + 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¤t_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" + } +} diff --git a/backend/internal/domain/llm.go b/backend/internal/domain/llm.go index bb304f1..8ba76bd 100644 --- a/backend/internal/domain/llm.go +++ b/backend/internal/domain/llm.go @@ -2,6 +2,7 @@ package domain import ( "context" + "encoding/json" "time" ) @@ -33,6 +34,27 @@ type Message struct { 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 { Chat *Chat `json:"chat"` Messages []*Message `json:"messages"` @@ -44,17 +66,42 @@ type ChatParams struct { 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 { + ID string + Type string + Function ToolCallFunction + Extra map[string]map[string]any +} + +type ToolCallFunction struct { + Name string + Arguments string } type Chunk struct { Done bool `json:"done"` - Role string `json:"role"` - Content string `json:"content"` - Reasoning bool + ToolName string `json:"toolName,omitempty"` + Role string `json:"role,omitempty"` + Content string `json:"content,omitempty"` + Reasoning bool `json:"reasoning,omitempty"` } type StreamCallback func(chunk *Chunk) @@ -69,7 +116,7 @@ const ( ) type LLMParams struct { - Tools []*Tool + Tools []*ToolDefinition OnChunk StreamCallback ReasoningEffort string } @@ -83,7 +130,7 @@ type LLMEngine interface { modelID string, params *LLMParams, messages []*Message, - ) (*Message, error) + ) *Message } type ChatRepository interface { @@ -92,6 +139,6 @@ type ChatRepository interface { CreateChat(userID string, chat *Chat) (*Chat, error) UpdateChat(userID string, chat *Chat) (*Chat, 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) } diff --git a/backend/internal/http/handlers/chat.go b/backend/internal/http/handlers/chat.go index 63a4667..d35aff7 100644 --- a/backend/internal/http/handlers/chat.go +++ b/backend/internal/http/handlers/chat.go @@ -2,7 +2,6 @@ package handlers import ( "encoding/json" - "log/slog" "net/http" "strconv" @@ -102,12 +101,11 @@ func (h *ChatHandler) NewChatMessage(c *echo.Context) error { for { 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 { 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") chatID := c.Param("id") chatWithMessages, err := h.chatService.GetChatWithMessages(userID, teamID, chatID) + if err != nil { c.Logger().Error("error getting chat with messages", "error", err) return echo.NewHTTPError(http.StatusBadRequest, "error getting chat") diff --git a/backend/internal/service/chat/service.go b/backend/internal/service/chat/service.go index ab706b3..0b8ed5e 100644 --- a/backend/internal/service/chat/service.go +++ b/backend/internal/service/chat/service.go @@ -120,7 +120,7 @@ func (s *Service) AddChatMessage( return nil, err } - messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID) + messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID, true) if err != nil { return nil, err } @@ -158,7 +158,7 @@ func (s *Service) GetChatWithMessages(userID string, teamID, chatID string) (*do if err != nil { return nil, err } - messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID) + messages, err := s.repoChat.GetChatMessages(userID, teamID, chatID, false) if err != nil { return nil, err } @@ -207,7 +207,7 @@ func (s *Service) runQuery( s.mu.Unlock() go func() { - _, err := chatSession.run(ctx) + err := chatSession.run(ctx) if err != nil { slog.Info("chat session error", "error", err) } diff --git a/backend/internal/service/chat/session.go b/backend/internal/service/chat/session.go index a3cb928..7d08761 100644 --- a/backend/internal/service/chat/session.go +++ b/backend/internal/service/chat/session.go @@ -21,10 +21,16 @@ type ChatSession struct { mu sync.Mutex } -func (s *ChatSession) run(ctx context.Context) (*domain.Message, error) { +func (s *ChatSession) run(ctx context.Context) error { params := &domain.LLMParams{} + tools := make([]*domain.ToolDefinition, 0) + for _, t := range allTools { + tools = append(tools, t.Definition(ctx)) + } + params.Tools = tools + s.mu.Lock() if len(s.subscribers) > 0 { params.OnChunk = func(chunk *domain.Chunk) { @@ -40,19 +46,66 @@ func (s *ChatSession) run(ctx context.Context) (*domain.Message, error) { Content: s.agent.SystemPrompt, } - // TODO : loop with tools - 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 + continue_loop := true - 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 } diff --git a/backend/internal/service/chat/tools.go b/backend/internal/service/chat/tools.go new file mode 100644 index 0000000..2924385 --- /dev/null +++ b/backend/internal/service/chat/tools.go @@ -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] +} diff --git a/backend/lassistanoque.db-wal b/backend/lassistanoque.db-wal index 80ba96b..eef948f 100644 Binary files a/backend/lassistanoque.db-wal and b/backend/lassistanoque.db-wal differ diff --git a/backend/web/src/lib/types/api.ts b/backend/web/src/lib/types/api.ts index 7f3aca6..b5fe81b 100644 --- a/backend/web/src/lib/types/api.ts +++ b/backend/web/src/lib/types/api.ts @@ -96,22 +96,34 @@ export interface Message { id: string role: string content: string + reasoning?: string toolCallId: string - toolCalls: Array + toolcalls: Array _date_created: string _date_updated: string _version: string } +export interface ToolCall { + ID: string + Type: string + Function: { + Name: string + Arguments: string + } +} + export type ChatWithMessages = { chat: Chat messages: Message[] } export type Chunk = { + toolName: string done: boolean role: string content: string + reasoning: boolean } export type ChatRequest = { diff --git a/backend/web/src/routes/(authenticated)/[space]/(chat)/[conversation]/+page.svelte b/backend/web/src/routes/(authenticated)/[space]/(chat)/[conversation]/+page.svelte index 757f6f6..b181b6c 100644 --- a/backend/web/src/routes/(authenticated)/[space]/(chat)/[conversation]/+page.svelte +++ b/backend/web/src/routes/(authenticated)/[space]/(chat)/[conversation]/+page.svelte @@ -10,6 +10,7 @@ let prompt = $state("") let { data }: PageProps = $props(); let messages = $state([] as Message[]) + let waiting = $state(false) onMount(async ()=>{ if (page.state.prompt) { @@ -26,19 +27,54 @@ const sendRequest = async(theprompt:string) => { prompt = "" let lastrole="" - //let reasoning = false + let lastToolName="" + waiting = true + 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() - messages.push({ + let m = { id: ''+currentDate.getTime(), role: chunk.role, - content: chunk.content, - } as Message) + content: "", + 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 + lastToolName=chunk.toolName } 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('div.reasoning:not(.hidden)') + .forEach((el) => el.classList.add('hidden')); + } + historyDiv.scrollTop = historyDiv.scrollHeight; } } @@ -52,11 +88,23 @@
{#each messages as message (message.id)}
-
- + {#if message.reasoning } +
+
+ {/if} + {#if message.toolcalls != null && message.toolcalls.length>0} +
⚒️ Appel de l'outil {message.toolcalls[0].Function.Name}
+ {:else} +
+ +
+ {/if}
{/each} + {#if waiting} +
...
+ {/if}
diff --git a/backend/web/src/routes/layout.css b/backend/web/src/routes/layout.css index dcaa4a7..aee9c9f 100644 --- a/backend/web/src/routes/layout.css +++ b/backend/web/src/routes/layout.css @@ -275,21 +275,19 @@ h2 { /* chat */ .message { - @apply m-2 rounded p-2; + @apply rounded p-2; &.user { @apply bg-sky-700 dark:text-white; } - - &.assistant { - &.reasoning { - @apply text-white/50; - } - } } .message-content { @apply text-sm; + + &.reasoning { + @apply text-white/50; + } } .message-content p {