diff --git a/neo/api.go b/neo/api.go
index d4d34913..71e866da 100644
--- a/neo/api.go
+++ b/neo/api.go
@@ -230,6 +230,7 @@ func (neo *DSL) handleChat(c *gin.Context) {
// Set the context with validated chat_id
ctx, cancel := chatctx.NewWithCancel(sid, chatID, c.Query("context"))
defer cancel()
+ defer ctx.Release() // Release the context after the request is done
neo.Answer(ctx, content, c)
}
diff --git a/neo/assistant/api.go b/neo/assistant/api.go
index 59a00e44..64ffc10a 100644
--- a/neo/assistant/api.go
+++ b/neo/assistant/api.go
@@ -7,10 +7,12 @@ import (
"os"
"strings"
+ "github.com/fatih/color"
"github.com/gin-gonic/gin"
jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/gou/fs"
- "github.com/yaoapp/kun/utils"
+ "github.com/yaoapp/kun/exception"
+ "github.com/yaoapp/kun/log"
chatctx "github.com/yaoapp/yao/neo/context"
chatMessage "github.com/yaoapp/yao/neo/message"
)
@@ -46,17 +48,37 @@ func GetByConnector(connector string, name string) (*Assistant, error) {
}
// Execute implements the execute functionality
-func (ast *Assistant) Execute(c *gin.Context, ctx chatctx.Context, input string, options map[string]interface{}) error {
+func (ast *Assistant) Execute(c *gin.Context, ctx chatctx.Context, input interface{}, options map[string]interface{}, callback ...interface{}) (interface{}, error) {
contents := chatMessage.NewContents()
messages, err := ast.withHistory(ctx, input)
if err != nil {
- return err
+ return nil, err
}
- return ast.execute(c, ctx, messages, options, contents)
+ return ast.execute(c, ctx, messages, options, contents, callback...)
}
// Execute implements the execute functionality
-func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatMessage.Message, userOptions map[string]interface{}, contents *chatMessage.Contents) error {
+func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, userInput interface{}, userOptions map[string]interface{}, contents *chatMessage.Contents, callback ...interface{}) (interface{}, error) {
+
+ var input []chatMessage.Message
+
+ switch v := userInput.(type) {
+ case string:
+ input = []chatMessage.Message{{Role: "user", Text: v}}
+
+ case []interface{}:
+ raw, err := jsoniter.Marshal(v)
+ if err != nil {
+ return nil, fmt.Errorf("marshal input error: %s", err.Error())
+ }
+ err = jsoniter.Unmarshal(raw, &input)
+ if err != nil {
+ return nil, fmt.Errorf("unmarshal input error: %s", err.Error())
+ }
+
+ case []chatMessage.Message:
+ input = v
+ }
if contents == nil {
contents = chatMessage.NewContents()
@@ -68,14 +90,14 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatM
ctx.Version = ast.vision
// Run init hook
- res, err := ast.HookInit(c, ctx, input, options, contents)
+ res, err := ast.HookCreate(c, ctx, input, options, contents)
if err != nil {
chatMessage.New().
Assistant(ast.ID, ast.Name, ast.Avatar).
Error(err).
Done().
Write(c.Writer)
- return err
+ return nil, err
}
// Update options if provided
@@ -103,14 +125,14 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatM
Error(err).
Done().
Write(c.Writer)
- return err
+ return nil, err
}
// Reset Message Contents
last := input[len(input)-1]
input, err = newAst.withHistory(ctx, last)
if err != nil {
- return err
+ return nil, err
}
// Reset options
@@ -123,15 +145,15 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatM
// Update assistant id
ctx.AssistantID = res.AssistantID
- return newAst.handleChatStream(c, ctx, input, options, contents)
+ return newAst.handleChatStream(c, ctx, input, options, contents, callback...)
}
// Only proceed with chat stream if no specific next action was handled
- return ast.handleChatStream(c, ctx, input, options, contents)
+ return ast.handleChatStream(c, ctx, input, options, contents, callback...)
}
// Execute the next action
-func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *chatMessage.Contents) error {
+func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *chatMessage.Contents, callback ...interface{}) (interface{}, error) {
switch next.Action {
// It's not used, because the process could be executed in the hook script
@@ -168,26 +190,34 @@ func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *c
case "assistant":
if next.Payload == nil {
- return fmt.Errorf("payload is required")
+ return nil, fmt.Errorf("payload is required")
}
// Get assistant id
id, ok := next.Payload["assistant_id"].(string)
if !ok {
- return fmt.Errorf("assistant id should be string")
+ return nil, fmt.Errorf("assistant id should be string")
}
// Get assistant
assistant, err := Get(id)
if err != nil {
- return fmt.Errorf("get assistant error: %s", err.Error())
+ return nil, fmt.Errorf("get assistant error: %s", err.Error())
}
// Input
input := chatMessage.Message{}
_, has := next.Payload["input"]
if !has {
- return fmt.Errorf("input is required")
+ return nil, fmt.Errorf("input is required")
+ }
+
+ // Retry mode
+ retry := false
+ _, has = next.Payload["retry"]
+ if has {
+ retry = next.Payload["retry"].(bool)
+ ctx.Retry = retry
}
switch v := next.Payload["input"].(type) {
@@ -195,14 +225,14 @@ func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *c
messages := chatMessage.Message{}
err := jsoniter.UnmarshalFromString(v, &messages)
if err != nil {
- return fmt.Errorf("unmarshal input error: %s", err.Error())
+ return nil, fmt.Errorf("unmarshal input error: %s", err.Error())
}
input = messages
case map[string]interface{}:
msg, err := chatMessage.NewMap(v)
if err != nil {
- return fmt.Errorf("unmarshal input error: %s", err.Error())
+ return nil, fmt.Errorf("unmarshal input error: %s", err.Error())
}
input = *msg
@@ -213,7 +243,7 @@ func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *c
input = v
default:
- return fmt.Errorf("input should be string or []chatMessage.Message")
+ return nil, fmt.Errorf("input should be string or []chatMessage.Message")
}
// Options
@@ -229,30 +259,32 @@ func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *c
messages, err := assistant.withHistory(ctx, input)
if err != nil {
- return fmt.Errorf("with history error: %s", err.Error())
+ return nil, fmt.Errorf("with history error: %s", err.Error())
}
// Create a new Text
// Send loading message and mark as new
- msg := chatMessage.New().Map(map[string]interface{}{
- "new": true,
- "role": "assistant",
- "type": "loading",
- "props": map[string]interface{}{"placeholder": "Calling " + assistant.Name},
- })
- msg.Assistant(assistant.ID, assistant.Name, assistant.Avatar)
- msg.Write(c.Writer)
+ if !ctx.Silent {
+ msg := chatMessage.New().Map(map[string]interface{}{
+ "new": true,
+ "role": "assistant",
+ "type": "loading",
+ "props": map[string]interface{}{"placeholder": "Calling " + assistant.Name},
+ })
+ msg.Assistant(assistant.ID, assistant.Name, assistant.Avatar)
+ msg.Write(c.Writer)
+ }
newContents := chatMessage.NewContents()
// Update the context id
ctx.AssistantID = assistant.ID
- return assistant.execute(c, ctx, messages, options, newContents)
+ return assistant.execute(c, ctx, messages, options, newContents, callback...)
case "exit":
- return nil
+ return nil, nil
default:
- return fmt.Errorf("unknown action: %s", next.Action)
+ return nil, fmt.Errorf("unknown action: %s", next.Action)
}
}
@@ -281,26 +313,35 @@ func (ast *Assistant) Call(c *gin.Context, payload APIPayload) (interface{}, err
}
// handleChatStream manages the streaming chat interaction with the AI
-func (ast *Assistant) handleChatStream(c *gin.Context, ctx chatctx.Context, messages []chatMessage.Message, options map[string]interface{}, contents *chatMessage.Contents) error {
+func (ast *Assistant) handleChatStream(c *gin.Context, ctx chatctx.Context, messages []chatMessage.Message, options map[string]interface{}, contents *chatMessage.Contents, callback ...interface{}) (interface{}, error) {
clientBreak := make(chan bool, 1)
done := make(chan bool, 1)
+ var result interface{} = nil
+ var err error = nil
- // Chat with AI in background
+ requestCtx := c.Request.Context()
go func() {
- err := ast.streamChat(c, ctx, messages, options, clientBreak, done, contents)
+ var res interface{} = nil
+ res, err = ast.streamChat(c, ctx, messages, options, clientBreak, contents, callback...)
if err != nil {
chatMessage.New().Error(err).Done().Write(c.Writer)
+ err = fmt.Errorf("stream chat error %s", err.Error())
}
+ result = res
done <- true
}()
// Wait for completion or client disconnect
select {
case <-done:
- return nil
- case <-c.Writer.CloseNotify():
+ if err != nil {
+ return nil, err
+ }
+ return result, nil
+
+ case <-requestCtx.Done():
clientBreak <- true
- return nil
+ return nil, nil
}
}
@@ -311,18 +352,27 @@ func (ast *Assistant) streamChat(
messages []chatMessage.Message,
options map[string]interface{},
clientBreak chan bool,
- done chan bool,
- contents *chatMessage.Contents) error {
+ contents *chatMessage.Contents,
+ callback ...interface{},
+) (interface{}, error) {
+
+ var cb interface{}
+ if len(callback) > 0 {
+ cb = callback[0]
+ }
errorRaw := ""
isFirst := true
isFirstThink := true
isThinking := false
- isFirstTool := true
- isTool := false
+ toolsCount := 0
currentMessageID := ""
+ var retry error = nil
+ var result interface{} = nil // To save the result
+ var content string = "" // To save the content
err := ast.Chat(c.Request.Context(), messages, options, func(data []byte) int {
+
select {
case <-clientBreak:
return 0 // break
@@ -338,6 +388,10 @@ func (ast *Assistant) streamChat(
return 1 // continue
}
+ // Retry mode
+ msg.Retry = ctx.Retry // Retry mode
+ msg.Silent = ctx.Silent // Silent mode
+
// Handle error
if msg.Type == "error" {
value := msg.String()
@@ -348,7 +402,10 @@ func (ast *Assistant) streamChat(
value = res.Error
}
}
- chatMessage.New().Error(value).Done().Write(c.Writer)
+ newMsg := chatMessage.New().Error(value).Done()
+ newMsg.Retry = ctx.Retry
+ newMsg.Silent = ctx.Silent
+ newMsg.Callback(cb).Write(c.Writer)
return 0 // break
}
@@ -365,8 +422,11 @@ func (ast *Assistant) streamChat(
if isThinking && msg.Type != "think" {
// add the think close tag
end := chatMessage.New().Map(map[string]interface{}{"text": "\n\n", "type": "think", "delta": true})
- end.Write(c.Writer)
end.ID = currentMessageID
+ end.Retry = ctx.Retry
+ end.Silent = ctx.Silent
+
+ end.Callback(cb).Write(c.Writer)
end.AppendTo(contents)
contents.UpdateType("think", map[string]interface{}{"text": contents.Text()}, currentMessageID)
isThinking = false
@@ -376,30 +436,34 @@ func (ast *Assistant) streamChat(
contents.ClearToken()
}
- // for native tool_calls response
+ // for native tool_calls response, keep the first tool_calls_native message
if msg.Type == "tool_calls_native" {
- if isFirstTool {
- msg.Text = "\n\n" + msg.Text // add the tool_calls begin tag
- isFirstTool = false
- isTool = true
- }
- }
- // for tool response
- if isTool && msg.Type != "tool_calls_native" {
-
- if msg.IsDone {
- end := chatMessage.New().Map(map[string]interface{}{"text": "}\n\n", "type": "tool", "delta": true})
- end.Write(c.Writer)
- end.ID = currentMessageID
- end.AppendTo(contents)
- contents.UpdateType("tool", map[string]interface{}{"text": contents.Text()}, currentMessageID)
- isTool = false
- } else {
- msg.Text = "\n\n" + msg.Text // add the tool_calls close tag
+ if toolsCount > 1 {
+ msg.Text = "" // clear the text
+ msg.Type = "text"
+ msg.IsNew = false
+ return 1 // continue
}
- isTool = false
+ if msg.IsBeginTool {
+
+ if toolsCount == 1 {
+ msg.IsNew = false
+ msg.Text = "\n\n" // add the tool_calls close tag
+ }
+
+ if toolsCount == 0 {
+ msg.Text = "\n\n" + msg.Text // add the tool_calls begin tag
+ }
+
+ toolsCount++
+
+ }
+
+ if msg.IsEndTool {
+ msg.Text = msg.Text + "\n\n" // add the tool_calls close tag
+ }
}
delta := msg.String()
@@ -461,6 +525,11 @@ func (ast *Assistant) streamChat(
msgType = "tool"
}
+ // Add the text content to the content
+ if msgType == "text" || msgType == "" {
+ content += msg.Text // Save the content
+ }
+
output := chatMessage.New().Map(map[string]interface{}{
"text": delta,
"type": msgType,
@@ -468,11 +537,13 @@ func (ast *Assistant) streamChat(
"delta": true,
})
+ output.Retry = ctx.Retry // Retry mode
+ output.Silent = ctx.Silent // Silent mode
if isFirst {
output.Assistant(ast.ID, ast.Name, ast.Avatar)
isFirst = false
}
- output.Write(c.Writer)
+ output.Callback(cb).Write(c.Writer)
}
// Complete the stream
@@ -489,7 +560,10 @@ func (ast *Assistant) streamChat(
"type": "text",
"delta": true,
"done": true,
+ "retry": ctx.Retry,
+ "silent": ctx.Silent,
}).
+ Callback(cb).
Write(c.Writer)
}
@@ -499,8 +573,7 @@ func (ast *Assistant) streamChat(
// Some error occurred in the hook, return the error
if hookErr != nil {
- chatMessage.New().Error(hookErr.Error()).Done().Write(c.Writer)
- done <- true
+ retry = hookErr
return 0 // break
}
@@ -509,21 +582,32 @@ func (ast *Assistant) streamChat(
// If the hook is successful, execute the next action
if res != nil && res.Next != nil {
- err := res.Next.Execute(c, ctx, contents)
+ _, err := res.Next.Execute(c, ctx, contents, cb)
if err != nil {
- chatMessage.New().Error(err.Error()).Done().Write(c.Writer)
+ chatMessage.New().Error(err.Error()).Done().Callback(cb).Write(c.Writer)
}
- done <- true
return 0 // break
}
+ // if the result is not nil, save the result
+ if res != nil && res.Result != nil {
+ result = res.Result
+ }
+
// The default output
output := chatMessage.New().Done()
if res != nil && res.Output != nil {
output = chatMessage.New().Map(map[string]interface{}{"text": res.Output, "done": true})
+ output.Retry = ctx.Retry
+ output.Silent = ctx.Silent
}
- output.Write(c.Writer)
- done <- true
+
+ // has result
+ if res != nil && res.Result != nil && cb != nil {
+ output.Result = res.Result // Add the result to the output message
+ }
+
+ output.Callback(cb).Write(c.Writer)
return 0 // break
}
@@ -531,21 +615,97 @@ func (ast *Assistant) streamChat(
}
})
+ // retry
+ if retry != nil {
+
+ // Update the retry times
+ ctx.RetryTimes = ctx.RetryTimes + 1 // Increment the retry times
+ ctx.Retry = true // Set the retry mode
+
+ // Hook retry
+ promptAny, retryErr := ast.HookRetry(c, ctx, messages, contents, exception.Trim(retry))
+ if retryErr != nil {
+ color.Red("%s, try to fix the error %d times, but failed with %s", exception.Trim(retry), ctx.RetryTimes, exception.Trim(retryErr))
+ chatMessage.New().Error(retry.Error()).Done().Callback(cb).Write(c.Writer)
+ return nil, retry
+ }
+
+ if promptAny == nil {
+ chatMessage.New().Error(retry.Error()).Done().Callback(cb).Write(c.Writer)
+ return nil, retry
+ }
+
+ var prompt string = ""
+ switch v := promptAny.(type) {
+ case NextAction:
+ result, err := v.Execute(c, ctx, contents, cb)
+ if err != nil {
+ chatMessage.New().Error(err.Error()).Done().Callback(cb).Write(c.Writer)
+ return nil, retry
+ }
+ return result, nil
+
+ case string:
+ prompt = v
+ }
+
+ // Add the prompt to the messages
+ retryMessages, retryErr := ast.retryMessages(messages, prompt)
+ if retryErr != nil {
+ color.Red("%s, try to fix the error %d times, but failed with %s", exception.Trim(retry), ctx.RetryTimes, exception.Trim(retryErr))
+ chatMessage.New().Error(retry.Error()).Done().Callback(cb).Write(c.Writer)
+ return nil, retry
+ }
+
+ // Retry the chat
+ retryContents := chatMessage.NewContents()
+ return ast.execute(c, ctx, retryMessages, options, retryContents, cb)
+ }
+
// Handle error
if err != nil {
- return err
+ return nil, err
}
// raw error
if errorRaw != "" {
msg, err := chatMessage.NewStringError(errorRaw)
if err != nil {
- return fmt.Errorf("error: %s", err.Error())
+ return nil, fmt.Errorf("stream chat error %s", err.Error())
}
- msg.Done().Write(c.Writer)
+ msg.Retry = ctx.Retry
+ msg.Silent = ctx.Silent
+ msg.Done().Callback(cb).Write(c.Writer)
}
- return nil
+ // If the result is not nil, return the result
+ if result != nil {
+ return result, nil
+ }
+
+ // Return the content
+ return strings.TrimSpace(content), nil
+}
+
+func (ast *Assistant) retryMessages(messages []chatMessage.Message, prompt string) ([]chatMessage.Message, error) {
+
+ // Get the last user message
+ var lastIndex int
+ for i := len(messages) - 1; i >= 0; i-- {
+ if messages[i].Role == "user" {
+ messages[i].Text = prompt
+ lastIndex = i
+ break
+ }
+ }
+
+ if lastIndex == 0 {
+ return nil, fmt.Errorf("no user message found")
+ }
+
+ // Remove the messages after the last user message
+ messages = messages[:lastIndex+1]
+ return messages, nil
}
// saveChatHistory saves the chat history if storage is available
@@ -696,12 +856,25 @@ func (ast *Assistant) withPrompts(messages []chatMessage.Message) []chatMessage.
func (ast *Assistant) withHistory(ctx chatctx.Context, input interface{}) ([]chatMessage.Message, error) {
- var userMessage *chatMessage.Message = chatMessage.New()
+ var userMessage *chatMessage.Message
+ var inputMessages []*chatMessage.Message
switch v := input.(type) {
case string:
- userMessage.Map(map[string]interface{}{"role": "user", "content": v})
+ userMessage = chatMessage.New().Map(map[string]interface{}{"role": "user", "content": v})
+
case map[string]interface{}:
- userMessage.Map(v)
+ userMessage = chatMessage.New().Map(v)
+
+ case []interface{}:
+ raw, err := jsoniter.Marshal(v)
+ if err != nil {
+ return nil, fmt.Errorf("marshal input error: %s", err.Error())
+ }
+ err = jsoniter.Unmarshal(raw, &inputMessages)
+ if err != nil {
+ return nil, fmt.Errorf("unmarshal input error: %s", err.Error())
+ }
+
case chatMessage.Message:
userMessage = &v
case *chatMessage.Message:
@@ -711,7 +884,6 @@ func (ast *Assistant) withHistory(ctx chatctx.Context, input interface{}) ([]cha
}
messages := []chatMessage.Message{}
-
if storage != nil {
history, err := storage.GetHistory(ctx.Sid, ctx.ChatID)
if err != nil {
@@ -732,7 +904,19 @@ func (ast *Assistant) withHistory(ctx chatctx.Context, input interface{}) ([]cha
messages = ast.withPrompts(messages)
// Add user message
- messages = append(messages, *userMessage)
+ if userMessage != nil {
+ messages = append(messages, *userMessage)
+ }
+
+ // Add input messages
+ if len(inputMessages) > 0 {
+ for _, msg := range inputMessages {
+ if msg == nil || msg.Role == "" {
+ continue
+ }
+ messages = append(messages, *msg)
+ }
+ }
return messages, nil
}
@@ -755,13 +939,144 @@ func (ast *Assistant) Chat(ctx context.Context, messages []chatMessage.Message,
return nil
}
-func (ast *Assistant) requestMessages(ctx context.Context, messages []chatMessage.Message) ([]map[string]interface{}, error) {
+// formatMessages processes messages to ensure they meet the required standards:
+// 1. Filters out duplicate messages with identical content, role, and name
+// 2. Moves system messages to the beginning while preserving the order of other messages
+// 3. Ensures the first non-system message is a user message (removes leading assistant messages)
+// 4. Ensures the last message is a user message (removes trailing assistant messages)
+// 5. Merges consecutive assistant messages from the same assistant
+func formatMessages(messages []map[string]interface{}) []map[string]interface{} {
+ // Filter out duplicate messages with identical content, role, and name
+ filteredMessages := []map[string]interface{}{}
+ seen := make(map[string]bool)
+ for _, msg := range messages {
+ // Create a unique key for each message based on role, content, and name
+ role := msg["role"].(string)
+ content := fmt.Sprintf("%v", msg["content"]) // Convert to string regardless of type
+
+ // Get name if it exists
+ name := ""
+ if nameVal, exists := msg["name"]; exists {
+ name = fmt.Sprintf("%v", nameVal)
+ }
+
+ // Create a unique key for this message
+ key := fmt.Sprintf("%s:%s:%s", role, content, name)
+
+ // If we haven't seen this message before, add it to filtered messages
+ if !seen[key] {
+ filteredMessages = append(filteredMessages, msg)
+ seen[key] = true
+ }
+ }
+
+ // Separate system messages while preserving the order of other messages
+ systemMessages := []map[string]interface{}{}
+ otherMessages := []map[string]interface{}{}
+
+ for _, msg := range filteredMessages {
+ if msg["role"].(string) == "system" {
+ systemMessages = append(systemMessages, msg)
+ } else {
+ otherMessages = append(otherMessages, msg)
+ }
+ }
+
+ // Ensure the first non-system message is a user message
+ // If there are no user messages or the first message is not a user message, remove leading assistant messages
+ validOtherMessages := []map[string]interface{}{}
+ foundUserMessage := false
+
+ for _, msg := range otherMessages {
+ if msg["role"].(string) == "user" {
+ foundUserMessage = true
+ validOtherMessages = append(validOtherMessages, msg)
+ } else if foundUserMessage {
+ // Only keep assistant messages that come after a user message
+ validOtherMessages = append(validOtherMessages, msg)
+ }
+ // Skip assistant messages that come before any user message
+ }
+
+ // If no valid messages remain, return just the system messages
+ if len(validOtherMessages) == 0 {
+ return systemMessages
+ }
+
+ // Ensure the last message is a user message
+ // Remove any trailing assistant messages
+ lastUserIndex := -1
+ for i := len(validOtherMessages) - 1; i >= 0; i-- {
+ if validOtherMessages[i]["role"].(string) == "user" {
+ lastUserIndex = i
+ break
+ }
+ }
+
+ // If we found a user message, trim any assistant messages after it
+ if lastUserIndex >= 0 && lastUserIndex < len(validOtherMessages)-1 {
+ validOtherMessages = validOtherMessages[:lastUserIndex+1]
+ }
+
+ // If there are no user messages left after filtering, return just the system messages
+ if len(validOtherMessages) == 0 {
+ return systemMessages
+ }
+
+ // Combine system messages first, followed by other valid messages in their original order
+ orderedMessages := append(systemMessages, validOtherMessages...)
+
+ // Merge consecutive assistant messages
+ mergedMessages := []map[string]interface{}{}
+ var lastMessage map[string]interface{}
+
+ for _, msg := range orderedMessages {
+ // If this is the first message, just add it
+ if lastMessage == nil {
+ mergedMessages = append(mergedMessages, msg)
+ lastMessage = msg
+ continue
+ }
+
+ // If both current and last messages are from assistant, check if they can be merged
+ if msg["role"].(string) == "assistant" && lastMessage["role"].(string) == "assistant" {
+ // Get name information
+ nameVal, hasName := msg["name"]
+
+ // Prepare name prefix for the content
+ namePrefix := ""
+ if hasName {
+ namePrefix = fmt.Sprintf("[%v]: ", nameVal)
+ }
+
+ // Merge the content, including name information if available
+ lastContent := fmt.Sprintf("%v", lastMessage["content"])
+ content := fmt.Sprintf("%v", msg["content"])
+
+ // Add the name prefix to the content
+ if namePrefix != "" {
+ content = namePrefix + content
+ }
+
+ // Merge the messages
+ lastMessage["content"] = lastContent + "\n" + content
+ continue
+ }
+
+ // If we can't merge, add as a new message
+ mergedMessages = append(mergedMessages, msg)
+ lastMessage = msg
+ }
+
+ return mergedMessages
+}
+
+func (ast *Assistant) requestMessages(ctx context.Context, messages []chatMessage.Message) ([]map[string]interface{}, error) {
newMessages := []map[string]interface{}{}
length := len(messages)
for index, message := range messages {
-
// Ignore the tool, think, error
if message.Type == "tool" || message.Type == "think" || message.Type == "error" {
continue
@@ -824,14 +1139,18 @@ func (ast *Assistant) requestMessages(ctx context.Context, messages []chatMessag
newMessages = append(newMessages, newMessage)
}
+ // Process messages to standardize format, filter duplicates, and merge consecutive assistant messages
+ processedMessages := formatMessages(newMessages)
+
// For debug environment, print the request messages
if os.Getenv("YAO_AGENT_PRINT_REQUEST_MESSAGES") == "true" {
- fmt.Println("--- REQUEST_MESSAGES -----------------------------")
- utils.Dump(newMessages)
- fmt.Println("--- END REQUEST_MESSAGES -----------------------------")
+ for _, message := range processedMessages {
+ raw, _ := jsoniter.MarshalToString(message)
+ log.Trace("[Request Message] %s", raw)
+ }
}
- return newMessages, nil
+ return processedMessages, nil
}
func (ast *Assistant) withAttachments(ctx context.Context, msg *chatMessage.Message) ([]map[string]interface{}, error) {
diff --git a/neo/assistant/assistant.go b/neo/assistant/assistant.go
index e320f91e..7af06d0b 100644
--- a/neo/assistant/assistant.go
+++ b/neo/assistant/assistant.go
@@ -3,12 +3,15 @@ package assistant
import (
"context"
"fmt"
+ "path"
"time"
"github.com/fatih/color"
jsoniter "github.com/json-iterator/go"
+ "github.com/yaoapp/gou/fs"
"github.com/yaoapp/gou/rag/driver"
"github.com/yaoapp/kun/log"
+ sui "github.com/yaoapp/yao/sui/core"
)
// Save save the assistant
@@ -156,6 +159,28 @@ func (ast *Assistant) Validate() error {
return nil
}
+// Assets get the assets content
+func (ast *Assistant) Assets(name string, data sui.Data) (string, error) {
+
+ app, err := fs.Get("app")
+ if err != nil {
+ return "", err
+ }
+
+ root := path.Join(ast.Path, "assets", name)
+ raw, err := app.ReadFile(root)
+ if err != nil {
+ return "", err
+ }
+
+ if data != nil {
+ content, _ := data.Replace(string(raw))
+ return content, nil
+ }
+
+ return string(raw), nil
+}
+
// Clone creates a deep copy of the assistant
func (ast *Assistant) Clone() *Assistant {
if ast == nil {
diff --git a/neo/assistant/call.go b/neo/assistant/call.go
new file mode 100644
index 00000000..cc17194e
--- /dev/null
+++ b/neo/assistant/call.go
@@ -0,0 +1,701 @@
+package assistant
+
+import (
+ "context"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/fatih/color"
+ "github.com/google/uuid"
+ jsoniter "github.com/json-iterator/go"
+ "github.com/yaoapp/gou/runtime/v8/bridge"
+ "github.com/yaoapp/kun/exception"
+ "github.com/yaoapp/kun/log"
+ chatctx "github.com/yaoapp/yao/neo/context"
+ chatMessage "github.com/yaoapp/yao/neo/message"
+ sui "github.com/yaoapp/yao/sui/core"
+ "rogchap.com/v8go"
+)
+
+// objectCall is the object for the call function
+type objectCall struct{}
+
+// OptionsCall is the options for the call function
+type OptionsCall struct {
+ Retry OptionsCallRetry `json:"retry,omitempty"` // Retry options
+ Options map[string]interface{} `json:"options,omitempty"` // LLM API options
+ Silent bool `json:"silent,omitempty"` // Silent mode, default is true
+}
+
+// OptionsCallRetry is the retry options for the call function
+type OptionsCallRetry struct {
+ Times int `json:"times,omitempty"` // Retry times, default is 3
+ Delay int `json:"delay,omitempty"` // Retry delay, default is 200
+ DelayMax int `json:"delay_max,omitempty"` // Retry delay max, default is 5000
+ Prompt string `json:"prompt,omitempty"` // Retry prompt, default is "Please fix the error. \n {{ error }}"
+}
+
+// allowedEvents is the allowed events for the call function
+var allowedEvents = map[string]bool{
+ "done": true,
+ "retry": true,
+ "message": true,
+}
+
+var callProps = []string{
+ "assistant_id",
+ "input",
+ "options",
+ "retry_times",
+}
+
+// jsNewPlan create a plan object and return it
+func jsCall(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+ args := info.Args()
+ if len(args) < 2 {
+ return bridge.JsException(info.Context(), "Run requires at least two arguments")
+ }
+
+ options := v8go.Undefined(info.Context().Isolate())
+ if len(args) > 2 {
+ options = args[2]
+ }
+
+ // Export the object
+ obj := &objectCall{}
+ objectTmpl := obj.ExportObject(info)
+ this, err := objectTmpl.NewInstance(info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ // Copy global properties
+ global := info.This()
+ for _, prop := range objectProperties {
+ if !global.Has(prop) {
+ continue
+ }
+ value, err := global.Get(prop)
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get property %s: %s", prop, err.Error()))
+ }
+ this.Set(prop, value)
+ }
+
+ this.Set("assistant_id", args[0])
+ this.Set("input", args[1])
+ this.Set("options", options)
+ this.Set("retry_times", int32(1))
+ return this.Value
+}
+
+// ExportObject Export as a FS Object
+func (obj *objectCall) ExportObject(info *v8go.FunctionCallbackInfo) *v8go.ObjectTemplate {
+ tmpl := v8go.NewObjectTemplate(info.Context().Isolate())
+ tmpl.Set("On", v8go.NewFunctionTemplate(info.Context().Isolate(), obj.on)) // On the call
+ tmpl.Set("Run", v8go.NewFunctionTemplate(info.Context().Isolate(), obj.run)) // Run the call
+ return tmpl
+}
+
+// on bind the callback to the call object
+func (obj *objectCall) on(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+ args := info.Args()
+ if len(args) < 2 {
+ return bridge.JsException(info.Context(), "On requires at least one argument")
+ }
+
+ if !args[0].IsString() {
+ return bridge.JsException(info.Context(), "The first argument should be a string")
+ }
+
+ name := args[0].String()
+ if !allowedEvents[name] {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Invalid event %s", name))
+ }
+
+ cb := args[1]
+ if !cb.IsFunction() {
+ return bridge.JsException(info.Context(), fmt.Sprintf("The second argument should be a function for event %s", name))
+ }
+
+ this := info.This()
+ this.Set(fmt.Sprintf("on_%s", name), cb)
+ return this.Value
+}
+
+// run run the call
+func (obj *objectCall) run(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+ this := info.This()
+ args := info.Args()
+
+ global, err := getGlobal(info.Context(), this)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ goArgs := []interface{}{}
+ jsArgs := []v8go.Valuer{}
+ if len(args) > 0 {
+ for _, arg := range args {
+ v, err := bridge.GoValue(arg, info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+ goArgs = append(goArgs, v)
+ jsArgs = append(jsArgs, arg)
+ }
+ }
+
+ // Get the assistant id
+ jsAssistantID, err := this.Get("assistant_id")
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ assistantID := jsAssistantID.String()
+
+ // Get the input
+ jsInput, err := this.Get("input")
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the input: %s", err.Error()))
+ }
+
+ input, err := bridge.GoValue(jsInput, info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to unmarshal the input: %s", err.Error()))
+ }
+
+ // Get the retry input
+ if this.Has("retry_input") {
+
+ jsRetryInput, err := this.Get("retry_input")
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the retry input: %s", err.Error()))
+ }
+
+ input, err = bridge.GoValue(jsRetryInput, info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to unmarshal the retry input: %s", err.Error()))
+ }
+ }
+
+ // Options
+ options := OptionsCall{
+ Retry: OptionsCallRetry{
+ Times: 3,
+ Delay: 200,
+ DelayMax: 1000,
+ Prompt: "{{ input }}\n**Answer is not correct, please try again.**\nError:\n{{ error }} \nAssistant's last answer:\n{{ output }}",
+ },
+ Silent: true,
+ Options: map[string]interface{}{}, // LLM API options
+ }
+
+ // Get the options
+ if this.Has("options") {
+ jsOptions, err := this.Get("options")
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the options: %s", err.Error()))
+ }
+
+ // Check if the options is undefined
+ if !jsOptions.IsUndefined() {
+ err = bridge.Unmarshal(jsOptions, &options)
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to unmarshal the options: %s", err.Error()))
+ }
+ }
+ }
+
+ // Get the assistant
+ newAst, err := Get(assistantID)
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the assistant: %s", err.Error()))
+ }
+
+ // Get the message event ( it will be used for the message event )
+ eventMessage := ""
+ goCallProps := map[string]interface{}{}
+ if this.Has("on_message") {
+ jsEventMessage, err := this.Get("on_message")
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the message: %s", err.Error()))
+ }
+ eventMessage = jsEventMessage.String()
+
+ for _, prop := range callProps {
+ if this.Has(prop) {
+ value, err := this.Get(prop)
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the %s property: %s", prop, err.Error()))
+ }
+ goValue, err := bridge.GoValue(value, info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the %s property: %s", prop, err.Error()))
+ }
+ goCallProps[prop] = goValue
+ }
+ }
+ }
+
+ // Update the chat context
+ var chatCtx chatctx.Context = global.ChatContext
+ chatCtx.AssistantID = assistantID
+ chatCtx.ChatID = fmt.Sprintf("call_%s", uuid.New().String()) // New chat id
+ chatCtx.Silent = options.Silent // Check the silent mode
+
+ // Define the callback function
+ var cb func(msg *chatMessage.Message) = nil
+ var output = []chatMessage.Message{}
+ cb = func(msg *chatMessage.Message) {
+ output = append(output, *msg)
+ if eventMessage != "" {
+ err := obj.triggerAnonymous(chatCtx, global, goCallProps, eventMessage, goArgs, msg)
+ if err != nil {
+ color.Red("Failed to trigger the message event: %s", err.Error())
+ log.Error("Failed to trigger the message event: %s", err.Error())
+ return
+ }
+ }
+ }
+
+ // Execute the assistant
+ result, err := newAst.Execute(global.GinContext, chatCtx, input, options.Options, cb) // Execute the assistant
+ if err != nil {
+ result, err = obj.retry(jsArgs, err, input, output, info, options)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+ }
+
+ // Copy props
+ for name, value := range goCallProps {
+ info.Context().Global().Set(name, value)
+ }
+
+ // Trigger the done event
+ doneResult, err := obj.trigger(info, "done", jsArgs...)
+ if err != nil {
+ result, err = obj.retry(jsArgs, err, input, output, info, options)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+ }
+
+ // Return the done result
+ if doneResult != nil && !doneResult.IsUndefined() {
+ return doneResult
+ }
+
+ // Return Value
+ switch v := result.(type) {
+ case *v8go.Value:
+ return v
+ case error:
+ return bridge.JsException(info.Context(), v.Error())
+ }
+
+ // Return Value
+ jsResult, err := bridge.JsValue(info.Context(), result)
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("Failed to get the result: %s", err.Error()))
+ }
+ return jsResult
+}
+
+func (obj *objectCall) retry(jsArgs []v8go.Valuer, err error, input interface{}, output []chatMessage.Message, info *v8go.FunctionCallbackInfo, options OptionsCall) (*v8go.Value, error) {
+
+ // Retry times, if not set, return the error
+ if options.Retry.Times <= 0 {
+ return nil, err
+ }
+
+ this := info.This()
+ errmsg := exception.Trim(err)
+
+ // Get current retry times
+ jsTimes, retryErr := this.Get("retry_times")
+ if retryErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to get the retry times: %s", errmsg, retryErr.Error())
+ }
+
+ times := int(jsTimes.Int32())
+ if times > options.Retry.Times {
+ return nil, fmt.Errorf("%s occurred, max retry times reached", errmsg)
+ }
+
+ // Update the retry times
+ times = times + 1
+ this.Set("retry_times", int32(times))
+
+ // Message content
+ content := ""
+ for _, msg := range output {
+ if msg.Type == "text" && msg.IsDelta {
+ content += msg.Text
+ }
+ }
+
+ // Delay
+ delay := options.Retry.Delay * int(times)
+ if delay > options.Retry.DelayMax {
+ delay = options.Retry.DelayMax
+ }
+
+ // Retry delay (millisecond)
+ if delay > 0 {
+ time.Sleep(time.Duration(delay) * time.Millisecond)
+ }
+
+ // Format the input
+ var lastUserMessage *chatMessage.Message = nil
+ var inputMessages []*chatMessage.Message = nil
+ var lastUserMessageIndex int = 0
+ switch v := input.(type) {
+ case string:
+ lastUserMessage = &chatMessage.Message{Type: "text", Text: v, Role: "user"}
+ inputMessages = []*chatMessage.Message{lastUserMessage}
+
+ case []interface{}:
+ // Get the last user message
+ raw, parseErr := jsoniter.Marshal(v)
+ if parseErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to marshal the input: %s", errmsg, parseErr.Error())
+ }
+
+ parseErr = jsoniter.Unmarshal(raw, &inputMessages)
+ if parseErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to unmarshal the input: %s", errmsg, parseErr.Error())
+ }
+
+ // Get the last user message
+ for i := len(inputMessages) - 1; i >= 0; i-- {
+ if inputMessages[i].Type == "text" && inputMessages[i].Role == "user" {
+ lastUserMessage = inputMessages[i]
+ lastUserMessageIndex = i
+ break
+ }
+ }
+
+ case *chatMessage.Message:
+ lastUserMessage = v
+ inputMessages = []*chatMessage.Message{lastUserMessage}
+
+ case map[string]interface{}:
+ text, ok := v["text"].(string)
+ if !ok {
+ return nil, fmt.Errorf("%s occurred but failed to get the text", errmsg)
+ }
+
+ if v["role"] != "user" {
+ return nil, fmt.Errorf("%s occurred but the role is not user", errmsg)
+ }
+
+ lastUserMessage = &chatMessage.Message{Type: "text", Text: text, Role: "user"}
+ inputMessages = []*chatMessage.Message{lastUserMessage}
+ }
+
+ // Get the prompt from the options
+ promptTmpl := options.Retry.Prompt
+ data := sui.Data{
+ "error": errmsg,
+ "output": strings.TrimSpace(content),
+ "input": lastUserMessage.Text,
+ }
+ prompt, _ := data.Replace(promptTmpl)
+
+ // Custom retry prompt by hooking the retry event
+ if this.Has("on_retry") {
+ info.Context().Global().Set("error", errmsg) // Set error
+ jsDelay, _ := bridge.JsValue(info.Context(), delay)
+ jsPrompt, _ := bridge.JsValue(info.Context(), prompt)
+ newPrompt, retryErr := obj.trigger(info, "retry", jsTimes, jsDelay, jsPrompt)
+ if retryErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to trigger the retry event: %s", errmsg, retryErr.Error())
+ }
+ // Update the prompt
+ if newPrompt.IsString() {
+ prompt = newPrompt.String()
+ }
+ }
+
+ // Generate the new input with the prompt
+ // Update the input
+ inputMessages[lastUserMessageIndex].Text = prompt
+ jsInput, inputErr := bridge.JsValue(info.Context(), inputMessages)
+ if inputErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to update the input: %s", errmsg, inputErr.Error())
+ }
+ // Update the input
+ this.Set("retry_input", jsInput)
+
+ // Call the run function
+ run, funcErr := this.Get("Run")
+ if funcErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to get the run function: %s", errmsg, funcErr.Error())
+ }
+
+ if !run.IsFunction() {
+ return nil, fmt.Errorf("%s occurred but the run function is not a function", errmsg)
+ }
+
+ fn, fnErr := run.AsFunction()
+ if fnErr != nil {
+ return nil, fmt.Errorf("%s occurred but failed to get the run function: %s", errmsg, fnErr.Error())
+ }
+
+ // Call the run function
+ result, resErr := fn.Call(this, jsArgs...)
+ if resErr != nil {
+ return nil, fmt.Errorf("%s (%d)", exception.Trim(resErr), times-1)
+ }
+
+ return result, nil
+}
+
+func (obj *objectCall) triggerAnonymous(chatCtx chatctx.Context, global *GlobalVariables, goCallProps map[string]interface{}, source string, bindArgs []interface{}, fnArgs ...interface{}) error {
+
+ ctx, err := global.Assistant.Script.NewContext(global.ChatContext.Sid, nil)
+ if err != nil {
+ return err
+ }
+ defer ctx.Close()
+
+ // Update Context
+ global.Assistant.InitObject(ctx, global.GinContext, chatCtx, global.Contents)
+
+ // Copy props
+ for k, v := range goCallProps {
+ ctx.WithGlobal(k, v)
+ }
+
+ // Add the args
+ ctx.WithGlobal("args", bindArgs)
+ _, err = ctx.CallAnonymousWith(context.Background(), source, fnArgs...)
+ if err != nil {
+ return err
+ }
+ return nil
+
+}
+
+// trigger trigger the callback
+func (obj *objectCall) trigger(info *v8go.FunctionCallbackInfo, name string, fnArgs ...v8go.Valuer) (*v8go.Value, error) {
+ // Try to get the callback
+ this := info.This()
+ if this.Has(fmt.Sprintf("on_%s", name)) {
+ event, err := this.Get(fmt.Sprintf("on_%s", name))
+ if err != nil {
+ return nil, err
+ }
+
+ if event.IsFunction() {
+
+ cb, err := event.AsFunction()
+ if err != nil {
+ return nil, err
+ }
+
+ result, err := cb.Call(this, fnArgs...)
+ if err != nil {
+ return nil, err
+ }
+ return result, nil
+ }
+ }
+
+ return nil, nil
+}
+
+// jsCallBackup is the backup function for the call function
+// func jsCallBackup(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+// // Get the args
+// args := info.Args()
+// if len(args) < 2 {
+// return bridge.JsException(info.Context(), "Run requires at least two arguments")
+// }
+
+// // Get the assistant id
+// assistantID := args[0].String()
+
+// // Get the assistant
+// newAst, err := Get(assistantID)
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+
+// // Get the input
+// input := args[1].String()
+
+// // Get the global variables
+// global, err := global(info)
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+
+// // Update Context
+// chatContext := global.ChatContext
+// chatContext.AssistantID = assistantID
+// chatContext.ChatID = fmt.Sprintf("chat_%s", uuid.New().String()) // New chat id
+// chatContext.Silent = true // Silent mode
+
+// var cb func(msg *chatMessage.Message) = nil
+// if len(args) > 2 {
+
+// // Rest args
+// var jsArgs *v8go.Value
+// goArgs := []interface{}{}
+// if len(args) > 3 {
+// jsArgs = args[3]
+// if jsArgs != nil {
+// if jsArgs.IsArray() {
+// v, err := bridge.GoValue(jsArgs, info.Context())
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+// arr, ok := v.([]interface{})
+// if !ok {
+// return bridge.JsException(info.Context(), "Invalid arguments")
+// }
+// goArgs = arr
+// } else {
+// v, err := bridge.GoValue(jsArgs, info.Context())
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+// goArgs = []interface{}{v}
+// }
+// }
+// }
+
+// // Parse the callback
+// funcType := "method"
+// name := ""
+// userArgs := []interface{}{}
+// if args[2].IsFunction() {
+// funcType = "anonymous"
+// } else {
+// goValue, err := bridge.GoValue(args[2], info.Context())
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+// switch v := goValue.(type) {
+// case string:
+// name = v
+// case map[string]interface{}:
+// if fname, ok := v["name"].(string); ok {
+// name = fname
+// }
+// if args, ok := v["args"].([]interface{}); ok {
+// userArgs = args
+// }
+// }
+
+// if strings.Contains(name, ".") {
+// funcType = "process"
+// }
+// }
+
+// switch funcType {
+// case "anonymous":
+// source := args[2].String()
+// cb = func(msg *chatMessage.Message) {
+// cbArgs := []interface{}{msg}
+// cbArgs = append(cbArgs, goArgs...)
+// ctx, err := global.Assistant.Script.NewContext(global.ChatContext.Sid, nil)
+// if err != nil {
+// fmt.Println("Failed to create context", err.Error())
+// return
+// }
+// defer ctx.Close()
+
+// global.Assistant.InitObject(ctx, global.GinContext, chatContext, global.Contents)
+// _, err = ctx.CallAnonymousWith(context.Background(), source, cbArgs...)
+// if err != nil {
+// log.Error("Failed to call the method: %s", err.Error())
+// color.Red("Failed to call the method: %s", err.Error())
+// return
+// }
+// }
+// break
+
+// case "process":
+
+// cb = func(msg *chatMessage.Message) {
+// cbArgs := []interface{}{}
+// cbArgs = append(cbArgs, msg)
+// cbArgs = append(cbArgs, userArgs...)
+// p, err := process.Of(name, cbArgs...)
+// if err != nil {
+// log.Error("Failed to get the process: %s", err.Error())
+// color.Red("Failed to get the process: %s", err.Error())
+// return
+// }
+// err = p.Execute()
+// if err != nil {
+// log.Error("Failed to execute the process: %s", err.Error())
+// color.Red("Failed to execute the process: %s", err.Error())
+// return
+// }
+// defer p.Release()
+// }
+
+// case "method":
+
+// cb = func(msg *chatMessage.Message) {
+// cbArgs := []interface{}{}
+// cbArgs = append(cbArgs, msg)
+// cbArgs = append(cbArgs, userArgs...)
+// ctx, err := global.Assistant.Script.NewContext(global.ChatContext.Sid, nil)
+// if err != nil {
+// return
+// }
+// defer ctx.Close()
+
+// global.Assistant.InitObject(ctx, global.GinContext, global.ChatContext, global.Contents)
+// _, err = ctx.CallWith(context.Background(), name, cbArgs...)
+// if err != nil {
+// log.Error("Failed to call the method: %s", err.Error())
+// color.Red("Failed to call the method: %s", err.Error())
+// return
+// }
+// }
+// }
+
+// }
+
+// // Parse the options
+// options := map[string]interface{}{}
+// if len(args) > 4 {
+// optionsRaw, err := bridge.GoValue(args[4], info.Context())
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+
+// // Parse the options
+// if optionsRaw != nil {
+// switch v := optionsRaw.(type) {
+// case string:
+// err := jsoniter.UnmarshalFromString(v, &options)
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+// case map[string]interface{}:
+// options = v
+// default:
+// return bridge.JsException(info.Context(), "Invalid options")
+// }
+// }
+// }
+
+// err = newAst.Execute(global.GinContext, chatContext, input, options, cb) // Execute the assistant
+// if err != nil {
+// return bridge.JsException(info.Context(), err.Error())
+// }
+// return nil
+// }
diff --git a/neo/assistant/hooks.go b/neo/assistant/hooks.go
index ad61278b..31aaec2c 100644
--- a/neo/assistant/hooks.go
+++ b/neo/assistant/hooks.go
@@ -9,18 +9,17 @@ import (
"github.com/gin-gonic/gin"
jsoniter "github.com/json-iterator/go"
- "github.com/yaoapp/gou/runtime/v8/bridge"
+ "github.com/yaoapp/kun/log"
chatctx "github.com/yaoapp/yao/neo/context"
"github.com/yaoapp/yao/neo/message"
chatMessage "github.com/yaoapp/yao/neo/message"
- "rogchap.com/v8go"
)
-// HookInit initialize the assistant
-func (ast *Assistant) HookInit(c *gin.Context, context chatctx.Context, input []message.Message, options map[string]interface{}, contents *message.Contents) (*ResHookInit, error) {
+// HookCreate create a new assistant
+func (ast *Assistant) HookCreate(c *gin.Context, context chatctx.Context, input []chatMessage.Message, options map[string]interface{}, contents *chatMessage.Contents) (*ResHookInit, error) {
// Create timeout context
ctx := ast.createBackgroundContext()
- v, err := ast.call(ctx, "Init", c, contents, context, input, options)
+ v, err := ast.call(ctx, "Create", c, contents, context, input, options)
if err != nil {
if err.Error() == HookErrorMethodNotFound {
return nil, nil
@@ -135,6 +134,46 @@ func (ast *Assistant) HookStream(c *gin.Context, context chatctx.Context, input
return response, nil
}
+// HookRetry Handle retry of assistant response
+func (ast *Assistant) HookRetry(c *gin.Context, context chatctx.Context, input []message.Message, contents *chatMessage.Contents, errmsg string) (interface{}, error) {
+ ctx := ast.createBackgroundContext()
+ output := []message.Data{}
+ if len(input) < 1 {
+ return "", fmt.Errorf("no input")
+ }
+
+ var lastInput message.Message = input[len(input)-1]
+ for _, data := range contents.Data {
+ if data.Type == "think" {
+ continue
+ }
+ output = append(output, data)
+ }
+
+ v, err := ast.call(ctx, "Retry", c, contents, context, lastInput.String(), output, errmsg)
+ if err != nil {
+ if err.Error() == HookErrorMethodNotFound {
+ return "", nil
+ }
+ return "", err
+ }
+
+ switch v := v.(type) {
+ case string:
+ return v, nil
+ case map[string]interface{}:
+ var next NextAction
+ raw, _ := jsoniter.MarshalToString(v)
+ err := jsoniter.UnmarshalFromString(raw, &next)
+ if err != nil {
+ return "", err
+ }
+ return next, nil
+ }
+
+ return "", nil
+}
+
// HookDone Handle completion of assistant response
func (ast *Assistant) HookDone(c *gin.Context, context chatctx.Context, input []message.Message, contents *chatMessage.Contents) (*ResHookDone, error) {
// Create timeout context
@@ -168,9 +207,7 @@ func (ast *Assistant) HookDone(c *gin.Context, context chatctx.Context, input []
text = content[:endIndex]
text = strings.TrimSpace(text)
if os.Getenv("YAO_AGENT_PRINT_TOOL_CALL") == "true" {
- fmt.Println("---- EXTRACTED TOOL CALL ----")
- fmt.Println(text)
- fmt.Println("---- END EXTRACTED TOOL CALL ----")
+ log.Trace("[TOOL CALL] %s", text)
}
}
}
@@ -220,6 +257,11 @@ func (ast *Assistant) HookDone(c *gin.Context, context chatctx.Context, input []
response.Output = vv
}
+ // has result
+ if res, has := v["result"]; has {
+ response.Result = res
+ }
+
if res, ok := v["next"].(map[string]interface{}); ok {
response.Next = &NextAction{}
if name, ok := res["action"].(string); ok {
@@ -308,53 +350,8 @@ func (ast *Assistant) call(ctx context.Context, method string, c *gin.Context, c
}
defer scriptCtx.Close()
- // Add sendMessage function to the script context
- scriptCtx.WithFunction("SendMessage", func(info *v8go.FunctionCallbackInfo) *v8go.Value {
-
- // Get the message
- args := info.Args()
- if len(args) < 1 {
- return bridge.JsException(info.Context(), "SendMessage requires at least one argument")
- }
-
- input, err := bridge.GoValue(args[0], info.Context())
- if err != nil {
- return bridge.JsException(info.Context(), err.Error())
- }
-
- // Save history by default
- saveHistory := true
- if len(args) > 1 && args[1].IsBoolean() {
- saveHistory = args[1].Boolean()
- }
-
- switch v := input.(type) {
- case string:
- // Check if the message is json
- msg, err := message.NewString(v)
- if err != nil {
- return bridge.JsException(info.Context(), err.Error())
- }
-
- // Append the message to the contents
- if saveHistory {
- msg.AppendTo(contents)
- }
- msg.Write(c.Writer)
- return nil
-
- case map[string]interface{}:
- msg := message.New().Map(v)
- if saveHistory {
- msg.AppendTo(contents)
- }
- msg.Write(c.Writer)
- return nil
-
- default:
- return bridge.JsException(info.Context(), "SendMessage requires a string or a map")
- }
- })
+ // Initialize the object, add the global variables, methods to the script context
+ ast.InitObject(scriptCtx, c, context, contents)
// Check if the method exists
if !scriptCtx.Global().Has(method) {
@@ -362,7 +359,6 @@ func (ast *Assistant) call(ctx context.Context, method string, c *gin.Context, c
}
// Call the method directly in the current thread
- args = append([]interface{}{context.Map()}, args...)
if scriptCtx != nil {
return scriptCtx.CallWith(ctx, method, args...)
}
diff --git a/neo/assistant/load_test.go b/neo/assistant/load_test.go
index dd5088c8..1e55d31f 100644
--- a/neo/assistant/load_test.go
+++ b/neo/assistant/load_test.go
@@ -371,12 +371,18 @@ func (m *mockStore) GetAssistants(filter store.AssistantFilter) (*store.Assistan
return nil, nil
}
func (m *mockStore) GetChat(id string, chatID string) (*store.ChatInfo, error) { return nil, nil }
+func (m *mockStore) GetChatWithFilter(id string, chatID string, filter store.ChatFilter) (*store.ChatInfo, error) {
+ return nil, nil
+}
func (m *mockStore) GetChats(id string, filter store.ChatFilter) (*store.ChatGroupResponse, error) {
return nil, nil
}
func (m *mockStore) GetHistory(id string, chatID string) ([]map[string]interface{}, error) {
return nil, nil
}
+func (m *mockStore) GetHistoryWithFilter(id string, chatID string, filter store.ChatFilter) ([]map[string]interface{}, error) {
+ return nil, nil
+}
func (m *mockStore) SaveAssistant(assistant map[string]interface{}) (interface{}, error) {
return nil, nil
}
diff --git a/neo/assistant/object.go b/neo/assistant/object.go
new file mode 100644
index 00000000..ea6e1c9b
--- /dev/null
+++ b/neo/assistant/object.go
@@ -0,0 +1,370 @@
+package assistant
+
+import (
+ "fmt"
+ "strings"
+
+ "github.com/gin-gonic/gin"
+ v8 "github.com/yaoapp/gou/runtime/v8"
+ "github.com/yaoapp/gou/runtime/v8/bridge"
+ chatctx "github.com/yaoapp/yao/neo/context"
+ "github.com/yaoapp/yao/neo/message"
+ chatMessage "github.com/yaoapp/yao/neo/message"
+ sui "github.com/yaoapp/yao/sui/core"
+ "rogchap.com/v8go"
+)
+
+// objectProperties is the properties of the assistant object
+var objectProperties = []string{
+ "__yao_agent_global",
+ "assistant",
+ "context",
+ "Plan",
+ "Send",
+ "Call",
+ "Assets",
+ "Set",
+ "Get",
+ "Del",
+ "Clear",
+}
+
+// GlobalVariables is the global variables for the assistant
+type GlobalVariables struct {
+ Assistant *Assistant
+ Contents *chatMessage.Contents
+ GinContext *gin.Context
+ ChatContext chatctx.Context
+}
+
+// JsValue return the javascript value of the global variables
+func (global *GlobalVariables) JsValue(ctx *v8go.Context) (*v8go.Value, error) {
+ return v8go.NewExternal(ctx.Isolate(), global)
+}
+
+// InitObject add the global variables and methods to the script context
+func (ast *Assistant) InitObject(v8ctx *v8.Context, c *gin.Context, context chatctx.Context, contents *chatMessage.Contents) {
+
+ // Add global variables to the script context
+ global := &GlobalVariables{
+ Assistant: ast,
+ Contents: contents,
+ GinContext: c,
+ ChatContext: context,
+ }
+
+ // Add global variables to the script context
+ v8ctx.WithGlobal("__yao_agent_global", global)
+
+ // Add assistant to the script context
+ v8ctx.WithGlobal("assistant", ast.Map())
+ v8ctx.WithGlobal("context", context.Map())
+
+ // Add methods to the script contexts
+ v8ctx.WithFunction("Send", jsSend)
+ v8ctx.WithFunction("Assets", jsAssets)
+ v8ctx.WithFunction("MakeCall", jsCall) // Create a new call object
+ v8ctx.WithFunction("MakePlan", jsPlan) // Create a new plan object
+
+ // Shared space methods
+ v8ctx.WithFunction("Set", jsSet)
+ v8ctx.WithFunction("Get", jsGet)
+ v8ctx.WithFunction("Del", jsDel)
+ v8ctx.WithFunction("Clear", jsClear)
+
+ // Template methods
+ v8ctx.WithFunction("Replace", jsReplace)
+}
+
+// jsSet function, set a value to the shared space
+func jsSet(info *v8go.FunctionCallbackInfo) *v8go.Value {
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ if global.ChatContext.SharedSpace == nil {
+ return bridge.JsException(info.Context(), "Shared space is not set")
+ }
+
+ args := info.Args()
+ if len(args) < 2 {
+ return bridge.JsException(info.Context(), "Set requires at least two arguments")
+ }
+
+ if !args[0].IsString() {
+ return bridge.JsException(info.Context(), "Set requires a valid key")
+ }
+
+ // Validate the key
+ key := args[0].String()
+ if key == "" {
+ return bridge.JsException(info.Context(), "Set requires a valid key")
+ }
+
+ // Validate the value
+ value, err := bridge.GoValue(args[1], info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ // Set the value
+ err = global.ChatContext.SharedSpace.Set(key, value)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ return nil
+}
+
+// jsGet function, get a value from the shared space
+func jsGet(info *v8go.FunctionCallbackInfo) *v8go.Value {
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ if global.ChatContext.SharedSpace == nil {
+ return bridge.JsException(info.Context(), "Shared space is not set")
+ }
+
+ args := info.Args()
+ if len(args) < 1 {
+ return bridge.JsException(info.Context(), "Get requires at least one argument")
+ }
+
+ if !args[0].IsString() {
+ return bridge.JsException(info.Context(), "Get requires a valid key")
+ }
+
+ // Get the key
+ key := args[0].String()
+ if key == "" {
+ return bridge.JsException(info.Context(), "Get requires a valid key")
+ }
+
+ // Get the value
+ value, err := global.ChatContext.SharedSpace.Get(key)
+ if err != nil {
+ // If the key is not found, return null
+ if strings.Contains(err.Error(), "not found") {
+ return v8go.Null(info.Context().Isolate())
+ }
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ jsValue, err := bridge.JsValue(info.Context(), value)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ return jsValue
+}
+
+// jsDel function, delete a value from the shared space
+func jsDel(info *v8go.FunctionCallbackInfo) *v8go.Value {
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ if global.ChatContext.SharedSpace == nil {
+ return bridge.JsException(info.Context(), "Shared space is not set")
+ }
+
+ args := info.Args()
+ if len(args) < 1 {
+ return bridge.JsException(info.Context(), "Get requires at least one argument")
+ }
+
+ if !args[0].IsString() {
+ return bridge.JsException(info.Context(), "Get requires a valid key")
+ }
+
+ // Get the key
+ key := args[0].String()
+ if key == "" {
+ return bridge.JsException(info.Context(), "Get requires a valid key")
+ }
+
+ err = global.ChatContext.SharedSpace.Delete(key)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ return nil
+}
+
+func jsClear(info *v8go.FunctionCallbackInfo) *v8go.Value {
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ if global.ChatContext.SharedSpace == nil {
+ return bridge.JsException(info.Context(), "Shared space is not set")
+ }
+
+ err = global.ChatContext.SharedSpace.Clear()
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ return nil
+}
+
+// jsAssets function, get the assets content
+func jsAssets(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ // Get the message
+ args := info.Args()
+ if len(args) < 1 {
+ return bridge.JsException(info.Context(), "Assets requires at least one argument")
+ }
+
+ // Get the name
+ name := args[0].String()
+
+ data := map[string]interface{}{}
+ if len(args) > 1 {
+ raw, err := bridge.GoValue(args[1], info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ v, ok := raw.(map[string]interface{})
+ if !ok {
+ return bridge.JsException(info.Context(), "Assets requires a map")
+ }
+ data = v
+ }
+
+ content, err := global.Assistant.Assets(name, data)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ jsContent, err := bridge.JsValue(info.Context(), content)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ return jsContent
+}
+
+// jsSend function, send a message to the http stream connection
+func jsSend(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+ // Get the message
+ args := info.Args()
+ if len(args) < 1 {
+ return bridge.JsException(info.Context(), "SendMessage requires at least one argument")
+ }
+
+ input, err := bridge.GoValue(args[0], info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ // Save history by default
+ saveHistory := true
+ if len(args) > 1 && args[1].IsBoolean() {
+ saveHistory = args[1].Boolean()
+ }
+
+ switch v := input.(type) {
+ case string:
+ // Check if the message is json
+ msg, err := message.NewString(v)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ // Append the message to the contents
+ if saveHistory {
+ msg.AppendTo(global.Contents)
+ }
+ msg.Write(global.GinContext.Writer)
+ return nil
+
+ case map[string]interface{}:
+ msg := message.New().Map(v)
+ if saveHistory {
+ msg.AppendTo(global.Contents)
+ }
+ msg.Write(global.GinContext.Writer)
+ return nil
+
+ default:
+ return bridge.JsException(info.Context(), "Send requires a string or a map")
+ }
+}
+
+func jsReplace(info *v8go.FunctionCallbackInfo) *v8go.Value {
+ args := info.Args()
+ if len(args) < 2 {
+ return bridge.JsException(info.Context(), "Replace requires at least two arguments")
+ }
+
+ if !args[0].IsString() {
+ return bridge.JsException(info.Context(), "the first argument must be a string")
+ }
+ tmpl := args[0].String()
+
+ raw, err := bridge.GoValue(args[1], info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ data, ok := raw.(map[string]interface{})
+ if !ok {
+ return bridge.JsException(info.Context(), "the second argument must be a map")
+ }
+
+ replaced, _ := sui.Data(data).Replace(tmpl)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ jsReplaced, err := bridge.JsValue(info.Context(), replaced)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ return jsReplaced
+}
+
+// global get the global variables
+func global(info *v8go.FunctionCallbackInfo) (global *GlobalVariables, err error) {
+ return getGlobal(info.Context(), info.This())
+}
+
+func getGlobal(ctx *v8go.Context, obj *v8go.Object) (global *GlobalVariables, err error) {
+ jsGlobal, err := obj.Get("__yao_agent_global")
+ if err != nil {
+ return nil, err
+ }
+
+ // Convert to go interface
+ goGlobal, err := bridge.GoValue(jsGlobal, ctx)
+ if err != nil {
+ return nil, err
+ }
+
+ global, ok := goGlobal.(*GlobalVariables)
+ if !ok {
+ return nil, fmt.Errorf("global is not a valid GlobalVariables. %#v", goGlobal)
+ }
+
+ return global, nil
+}
diff --git a/neo/assistant/plan.go b/neo/assistant/plan.go
new file mode 100644
index 00000000..8f5794f4
--- /dev/null
+++ b/neo/assistant/plan.go
@@ -0,0 +1,140 @@
+package assistant
+
+import (
+ "context"
+ "fmt"
+
+ "github.com/fatih/color"
+ "github.com/yaoapp/gou/runtime/v8/bridge"
+ v8plan "github.com/yaoapp/gou/runtime/v8/objects/plan"
+ "rogchap.com/v8go"
+)
+
+// TaskFn is the task function
+func TaskFn(plan_id string, task_id string, source bool, method string, args ...interface{}) (interface{}, error) {
+
+ if !source {
+ return v8plan.DefaultTaskFn(plan_id, task_id, source, method, args...)
+ }
+
+ // Data
+ plan, err := v8plan.GetPlan(plan_id)
+ if err != nil {
+ return nil, err
+ }
+
+ global, ok := plan.Data().(*GlobalVariables)
+ if !ok {
+ return nil, fmt.Errorf("plan data is not a GlobalVariables")
+ }
+
+ if global.Assistant == nil {
+ return nil, fmt.Errorf("assistant is not set")
+ }
+
+ if global.Assistant.Script == nil {
+ return nil, fmt.Errorf("script is not set")
+ }
+
+ scriptCtx, err := global.Assistant.Script.NewContext(global.ChatContext.Sid, nil)
+ if err != nil {
+ return nil, err
+ }
+ defer scriptCtx.Close()
+
+ // Initialize the object
+ global.Assistant.InitObject(scriptCtx, global.GinContext, global.ChatContext, global.Contents)
+
+ fnargs := []interface{}{plan_id, task_id}
+ fnargs = append(fnargs, args...)
+
+ // Execute the anonymous function
+ return scriptCtx.CallAnonymousWith(context.Background(), method, fnargs...)
+
+}
+
+// SubscribeFn is the default subscribe function
+func SubscribeFn(plan_id string, key string, value interface{}, source bool, method string, args ...interface{}) {
+
+ if !source {
+ v8plan.DefaultSubscribeFn(plan_id, key, value, source, method, args...)
+ return
+ }
+
+ // Data
+ plan, err := v8plan.GetPlan(plan_id)
+ if err != nil {
+ color.Red("Subscribe Failed to get the plan: %s", err.Error())
+ return
+ }
+
+ global, ok := plan.Data().(*GlobalVariables)
+ if !ok {
+ color.Red("Subscribe Failed: plan data is not a GlobalVariables")
+ return
+ }
+
+ if global.Assistant == nil {
+ color.Red("Subscribe Failed: assistant is not set")
+ return
+ }
+
+ if global.Assistant.Script == nil {
+ color.Red("Subscribe Failed: script is not set")
+ return
+ }
+
+ scriptCtx, err := global.Assistant.Script.NewContext(global.ChatContext.Sid, nil)
+ if err != nil {
+ color.Red("Subscribe Failed: Failed to create the script context: %s", err.Error())
+ return
+ }
+ defer scriptCtx.Close()
+
+ fnargs := []interface{}{plan_id, key, value}
+ fnargs = append(fnargs, args...)
+
+ // Initialize the object
+ global.Assistant.InitObject(scriptCtx, global.GinContext, global.ChatContext, global.Contents)
+ _, err = scriptCtx.CallAnonymousWith(context.Background(), method, fnargs...)
+ if err != nil {
+ return
+ }
+}
+
+// jsNewPlan create a plan object and return it
+func jsPlan(info *v8go.FunctionCallbackInfo) *v8go.Value {
+
+ global, err := global(info)
+ if err != nil {
+ return bridge.JsException(info.Context(), err.Error())
+ }
+
+ obj := newPlanObject()
+
+ args := info.Args()
+ if len(args) < 1 {
+ return bridge.JsException(info.Context(), "the first parameter should be a string")
+ }
+
+ if !args[0].IsString() {
+ return bridge.JsException(info.Context(), "the first parameter should be a string")
+ }
+
+ id := args[0].String()
+ objectTmpl := obj.ExportObject(info.Context().Isolate())
+ plan, err := objectTmpl.NewInstance(info.Context())
+ if err != nil {
+ return bridge.JsException(info.Context(), fmt.Sprintf("failed to create plan object %s", err.Error()))
+ }
+
+ return obj.NewInstance(id, plan, global)
+}
+
+func newPlanObject() *v8plan.Object {
+ obj := v8plan.New(v8plan.Options{
+ TaskFn: TaskFn,
+ SubscribeFn: SubscribeFn,
+ })
+ return obj
+}
diff --git a/neo/assistant/types.go b/neo/assistant/types.go
index cad8fd16..036cf5bc 100644
--- a/neo/assistant/types.go
+++ b/neo/assistant/types.go
@@ -26,7 +26,7 @@ type API interface {
ReadBase64(ctx context.Context, fileID string) (string, error)
GetPlaceholder() *Placeholder
- Execute(c *gin.Context, ctx chatctx.Context, input string, options map[string]interface{}) error
+ Execute(c *gin.Context, ctx chatctx.Context, input interface{}, options map[string]interface{}, callback ...interface{}) (interface{}, error)
Call(c *gin.Context, payload APIPayload) (interface{}, error)
}
@@ -58,6 +58,7 @@ type ResHookDone struct {
Next *NextAction `json:"next,omitempty"`
Input []message.Message `json:"input,omitempty"`
Output []message.Data `json:"output,omitempty"`
+ Result any `json:"result,omitempty"`
}
// ResHookFail the response of the fail hook
diff --git a/neo/context/context.go b/neo/context/context.go
index 686d3949..77d1c43f 100644
--- a/neo/context/context.go
+++ b/neo/context/context.go
@@ -5,6 +5,7 @@ import (
"time"
jsoniter "github.com/json-iterator/go"
+ "github.com/yaoapp/gou/plan"
"github.com/yaoapp/kun/log"
)
@@ -14,16 +15,20 @@ type Context struct {
Sid string `json:"sid" yaml:"-"` // Session ID
ChatID string `json:"chat_id,omitempty"` // Chat ID, use to select chat
AssistantID string `json:"assistant_id,omitempty"` // Assistant ID, use to select assistant
- Stack string `json:"stack,omitempty"`
- Path string `json:"pathname,omitempty"`
+ Stack string `json:"stack,omitempty"` // will be removed in the future
+ Path string `json:"pathname,omitempty"` // wiil be rename to path
FormData map[string]interface{} `json:"formdata,omitempty"`
Field *Field `json:"field,omitempty"`
Namespace string `json:"namespace,omitempty"`
Config map[string]interface{} `json:"config,omitempty"`
Signal interface{} `json:"signal,omitempty"`
+ Silent bool `json:"silent,omitempty"` // Silent mode
+ Retry bool `json:"retry,omitempty"` // Retry mode
+ RetryTimes uint8 `json:"retry_times,omitempty"` // Retry times
Upload *FileUpload `json:"upload,omitempty"`
Version bool `json:"version,omitempty"` // Version support
RAG bool `json:"rag,omitempty"` // RAG support
+ SharedSpace plan.Space `json:"-"` // Shared space
}
// Field the context field
@@ -45,7 +50,8 @@ type FileUpload struct {
// New create a new context
func New(sid, cid, payload string) Context {
- ctx := Context{Context: context.Background(), Sid: sid, ChatID: cid}
+
+ ctx := Context{Context: context.Background(), Sid: sid, ChatID: cid, SharedSpace: plan.NewMemorySharedSpace()}
if payload == "" {
return ctx
}
@@ -54,6 +60,7 @@ func New(sid, cid, payload string) Context {
if err != nil {
log.Error("%s", err.Error())
}
+
return ctx
}
@@ -83,6 +90,13 @@ func WithTimeout(parent Context, timeout time.Duration) (Context, context.Cancel
return parent, cancel
}
+// Release the context
+func (ctx *Context) Release() {
+ ctx.SharedSpace.Clear()
+ ctx.SharedSpace = nil
+ ctx = nil
+}
+
// Map the context to a map
func (ctx *Context) Map() map[string]interface{} {
data := map[string]interface{}{
@@ -100,6 +114,20 @@ func (ctx *Context) Map() map[string]interface{} {
if ctx.Stack != "" {
data["stack"] = ctx.Stack
}
+
+ // Silent mode
+ if ctx.Silent {
+ data["silent"] = ctx.Silent
+ }
+
+ // Retry mode
+ if ctx.Retry {
+ data["retry"] = ctx.Retry
+ }
+
+ // Retry times
+ data["retry_times"] = ctx.RetryTimes
+
if ctx.Path != "" {
data["pathname"] = ctx.Path
}
diff --git a/neo/message/message.go b/neo/message/message.go
index 19e8fbe0..b57147a2 100644
--- a/neo/message/message.go
+++ b/neo/message/message.go
@@ -4,6 +4,7 @@ import (
"fmt"
"os"
"strings"
+ "sync"
"github.com/fatih/color"
"github.com/gin-gonic/gin"
@@ -15,6 +16,8 @@ import (
"github.com/yaoapp/yao/openai"
)
+var locker = sync.Mutex{}
+
// Message the message
type Message struct {
ID string `json:"id,omitempty"` // id for the message
@@ -35,6 +38,12 @@ type Message struct {
Data map[string]interface{} `json:"-"` // data for the message
Pending bool `json:"-"` // pending for the message
Hidden bool `json:"hidden,omitempty"` // hidden for the message (not show in the UI and history)
+ Retry bool `json:"retry,omitempty"` // retry for the message
+ Silent bool `json:"silent,omitempty"` // silent for the message (not show in the UI and history)
+ IsTool bool `json:"-"` // is tool for the message for native tool_calls
+ IsBeginTool bool `json:"-"` // is new tool for the message for native tool_calls
+ IsEndTool bool `json:"-"` // is end tool for the message for native tool_calls
+ Result any `json:"result,omitempty"` // result for the message
}
// Mention represents a mention
@@ -187,10 +196,9 @@ func NewAny(content interface{}) (*Message, error) {
// NewOpenAI create a new message from OpenAI response
func NewOpenAI(data []byte, isThinking bool) *Message {
- // For Debug
// For debug environment, print the response data
if os.Getenv("YAO_AGENT_PRINT_RESPONSE_DATA") == "true" {
- fmt.Printf("%s\n", string(data))
+ log.Trace("[Response Data] %s", string(data))
}
if data == nil || len(data) == 0 {
@@ -220,19 +228,26 @@ func NewOpenAI(data []byte, isThinking bool) *Message {
}
// Tool calls
- if len(chunk.Choices[0].Delta.ToolCalls) > 0 {
+ if len(chunk.Choices[0].Delta.ToolCalls) > 0 || chunk.Choices[0].FinishReason == "tool_calls" {
msg.Type = "tool_calls_native"
- id := chunk.Choices[0].Delta.ToolCalls[0].ID
- function := chunk.Choices[0].Delta.ToolCalls[0].Function.Name
- arguments := chunk.Choices[0].Delta.ToolCalls[0].Function.Arguments
- text := arguments
- if id != "" {
- text = fmt.Sprintf(`{"id": "%s", "function": "%s", "arguments": %s`, id, function, arguments)
- msg.IsNew = true // mark as a new message
+ text := ""
+ if len(chunk.Choices[0].Delta.ToolCalls) > 0 {
+ id := chunk.Choices[0].Delta.ToolCalls[0].ID
+ function := chunk.Choices[0].Delta.ToolCalls[0].Function.Name
+ arguments := chunk.Choices[0].Delta.ToolCalls[0].Function.Arguments
+ text = arguments
+ if id != "" {
+ msg.IsBeginTool = true
+ msg.IsNew = true // mark as a new message
+ text = fmt.Sprintf(`{"id": "%s", "function": "%s", "arguments": %s`, id, function, arguments)
+ }
+ }
+
+ if chunk.Choices[0].FinishReason == "tool_calls" {
+ msg.IsEndTool = true
}
msg.Text = text
- msg.IsDone = chunk.Choices[0].FinishReason == "tool_calls" // is done when tool calls are finished
return msg
}
@@ -692,15 +707,61 @@ func (m *Message) Bind(data map[string]interface{}) *Message {
return m
}
+// Callback callback the message
+func (m *Message) Callback(fn interface{}) *Message {
+ if fn != nil {
+ switch v := fn.(type) {
+ case func(msg *Message):
+ if v == nil {
+ break
+ }
+ v(m)
+ break
+
+ case func():
+ if v == nil {
+ break
+ }
+ v()
+ break
+
+ default:
+ fmt.Println("no match callback")
+ break
+ }
+ }
+ return m
+}
+
// Write writes the message to response writer
func (m *Message) Write(w gin.ResponseWriter) bool {
+
+ // Sync write to response writer
+ locker.Lock()
+ defer locker.Unlock()
+
defer func() {
if r := recover(); r != nil {
+
+ // Ignore if done is true
+ if m.IsDone {
+ return
+ }
+
message := "Write Response Exception: (if client close the connection, it's normal) \n %s\n\n"
color.Red(message, r)
+
+ // Print the message
+ raw, _ := jsoniter.MarshalToString(m)
+ color.White("Message:\n %s", raw)
}
}()
+ // Ignore silent messages
+ if m.Silent {
+ return true
+ }
+
data, err := jsoniter.Marshal(m)
if err != nil {
log.Error("%s", err.Error())
diff --git a/neo/neo.go b/neo/neo.go
index 76d94bce..5bdfe438 100644
--- a/neo/neo.go
+++ b/neo/neo.go
@@ -22,7 +22,8 @@ func (neo *DSL) Answer(ctx chatctx.Context, question string, c *gin.Context) err
return err
}
}
- return ast.Execute(c, ctx, question, nil)
+ _, err = ast.Execute(c, ctx, question, nil)
+ return err
}
// Select select an assistant
diff --git a/neo/store/mongo.go b/neo/store/mongo.go
index e94b4278..b6fc9146 100644
--- a/neo/store/mongo.go
+++ b/neo/store/mongo.go
@@ -18,11 +18,21 @@ func (m *Mongo) GetChat(sid string, cid string) (*ChatInfo, error) {
return &ChatInfo{}, nil
}
+// GetChatWithFilter retrieves a single chat's information with filter options
+func (m *Mongo) GetChatWithFilter(sid string, cid string, filter ChatFilter) (*ChatInfo, error) {
+ return &ChatInfo{}, nil
+}
+
// GetHistory retrieves chat history
func (m *Mongo) GetHistory(sid string, cid string) ([]map[string]interface{}, error) {
return []map[string]interface{}{}, nil
}
+// GetHistoryWithFilter retrieves chat history with filter options
+func (m *Mongo) GetHistoryWithFilter(sid string, cid string, filter ChatFilter) ([]map[string]interface{}, error) {
+ return []map[string]interface{}{}, nil
+}
+
// SaveHistory saves chat history
func (m *Mongo) SaveHistory(sid string, messages []map[string]interface{}, cid string, context map[string]interface{}) error {
return nil
diff --git a/neo/store/redis.go b/neo/store/redis.go
index 50206e1a..a5176b27 100644
--- a/neo/store/redis.go
+++ b/neo/store/redis.go
@@ -18,11 +18,21 @@ func (r *Redis) GetChat(sid string, cid string) (*ChatInfo, error) {
return &ChatInfo{}, nil
}
+// GetChatWithFilter retrieves a single chat's information with filter options
+func (r *Redis) GetChatWithFilter(sid string, cid string, filter ChatFilter) (*ChatInfo, error) {
+ return &ChatInfo{}, nil
+}
+
// GetHistory retrieves chat history
func (r *Redis) GetHistory(sid string, cid string) ([]map[string]interface{}, error) {
return []map[string]interface{}{}, nil
}
+// GetHistoryWithFilter retrieves chat history with filter options
+func (r *Redis) GetHistoryWithFilter(sid string, cid string, filter ChatFilter) ([]map[string]interface{}, error) {
+ return []map[string]interface{}{}, nil
+}
+
// SaveHistory saves chat history
func (r *Redis) SaveHistory(sid string, messages []map[string]interface{}, cid string, context map[string]interface{}) error {
return nil
diff --git a/neo/store/types.go b/neo/store/types.go
index d4dcdba3..22cab473 100644
--- a/neo/store/types.go
+++ b/neo/store/types.go
@@ -24,6 +24,7 @@ type ChatFilter struct {
Page int `json:"page,omitempty"` // Page number, starting from 1
PageSize int `json:"pagesize,omitempty"` // Number of items per page
Order string `json:"order,omitempty"` // Sort order: desc/asc
+ Silent *bool `json:"silent,omitempty"` // Include silent messages (default: false)
}
// ChatGroup represents the chat group structure
@@ -86,12 +87,26 @@ type Store interface {
// Returns: Chat information and potential error
GetChat(sid string, cid string) (*ChatInfo, error)
+ // GetChatWithFilter retrieves a single chat's information with filter options
+ // sid: Session ID
+ // cid: Chat ID
+ // filter: Filter conditions
+ // Returns: Chat information and potential error
+ GetChatWithFilter(sid string, cid string, filter ChatFilter) (*ChatInfo, error)
+
// GetHistory retrieves chat history
// sid: Session ID
// cid: Chat ID
// Returns: History record list and potential error
GetHistory(sid string, cid string) ([]map[string]interface{}, error)
+ // GetHistoryWithFilter retrieves chat history with filter options
+ // sid: Session ID
+ // cid: Chat ID
+ // filter: Filter conditions
+ // Returns: History record list and potential error
+ GetHistoryWithFilter(sid string, cid string, filter ChatFilter) ([]map[string]interface{}, error)
+
// SaveHistory saves chat history
// sid: Session ID
// messages: Message list
diff --git a/neo/store/xun.go b/neo/store/xun.go
index 6f3b4c97..ceae2f0a 100644
--- a/neo/store/xun.go
+++ b/neo/store/xun.go
@@ -3,7 +3,6 @@ package store
import (
"fmt"
"math"
- "strings"
"time"
"github.com/google/uuid"
@@ -145,6 +144,7 @@ func (conv *Xun) initHistoryTable() error {
table.String("assistant_name", 200).Null()
table.String("assistant_avatar", 200).Null()
table.JSON("mentions").Null()
+ table.Boolean("silent").SetDefault(false).Index()
table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index()
table.TimestampTz("updated_at").Null().Index()
table.TimestampTz("expired_at").Null().Index()
@@ -162,7 +162,7 @@ func (conv *Xun) initHistoryTable() error {
return err
}
- fields := []string{"id", "sid", "cid", "uid", "role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "created_at", "updated_at", "expired_at"}
+ fields := []string{"id", "sid", "cid", "uid", "role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "silent", "created_at", "updated_at", "expired_at"}
for _, field := range fields {
if !tab.HasColumn(field) {
return fmt.Errorf("%s is required", field)
@@ -187,6 +187,7 @@ func (conv *Xun) initChatTable() error {
table.String("title", 200).Null()
table.String("assistant_id", 200).Null().Index()
table.String("sid", 255).Index()
+ table.Boolean("silent").SetDefault(false).Index()
table.TimestampTz("created_at").SetDefaultRaw("NOW()").Index()
table.TimestampTz("updated_at").Null().Index()
})
@@ -203,7 +204,7 @@ func (conv *Xun) initChatTable() error {
return err
}
- fields := []string{"id", "chat_id", "title", "assistant_id", "sid", "created_at", "updated_at"}
+ fields := []string{"id", "chat_id", "title", "assistant_id", "sid", "silent", "created_at", "updated_at"}
for _, field := range fields {
if !tab.HasColumn(field) {
return fmt.Errorf("%s is required", field)
@@ -319,53 +320,90 @@ func (conv *Xun) UpdateChatTitle(sid string, cid string, title string) error {
// GetChats get the chat list with grouping by date
func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, error) {
+ // Default behavior: exclude silent chats
+ if filter.Silent == nil {
+ silentFalse := false
+ filter.Silent = &silentFalse
+ }
+
+ return conv.getChatsWithFilter(sid, filter)
+}
+
+// getChatsWithFilter get the chats with filter options
+func (conv *Xun) getChatsWithFilter(sid string, filter ChatFilter) (*ChatGroupResponse, error) {
userID, err := conv.getUserID(sid)
if err != nil {
return nil, err
}
- // Set defaults
- if filter.PageSize <= 0 {
- filter.PageSize = 100
- }
+ // Set default values
if filter.Page <= 0 {
filter.Page = 1
}
+ if filter.PageSize <= 0 {
+ filter.PageSize = 20
+ }
if filter.Order == "" {
filter.Order = "desc"
}
- // Build base query
- qb := conv.newQueryChat().
- Select("chat_id", "title", "assistant_id", "created_at", "updated_at").
- Where("sid", userID).
- Where("chat_id", "!=", "")
+ // Get total count
+ qbCount := conv.newQueryChat().
+ Where("sid", userID)
- // Add keyword filter
- if filter.Keywords != "" {
- keyword := strings.TrimSpace(filter.Keywords)
- if keyword != "" {
- qb.Where("title", "like", "%"+keyword+"%")
+ // Apply silent filter if provided
+ if filter.Silent != nil {
+ if *filter.Silent {
+ // Include all chats (both silent and non-silent)
+ } else {
+ // Only include non-silent chats
+ qbCount.Where("silent", false)
}
}
- // Get total count
- total, err := qb.Clone().Count()
+ // Apply keyword filter if provided
+ if filter.Keywords != "" {
+ qbCount.Where("title", "like", fmt.Sprintf("%%%s%%", filter.Keywords))
+ }
+
+ total, err := qbCount.Count()
if err != nil {
return nil, err
}
- // Calculate pagination
- offset := (filter.Page - 1) * filter.PageSize
+ // Calculate last page
lastPage := int(math.Ceil(float64(total) / float64(filter.PageSize)))
+ if lastPage < 1 {
+ lastPage = 1
+ }
- // Get paginated results
- rows, err := qb.
- OrderBy("updated_at", filter.Order).
- OrderBy("created_at", filter.Order).
+ // Get chats with pagination
+ qb := conv.newQueryChat().
+ Select("chat_id", "title", "assistant_id", "silent", "created_at", "updated_at").
+ Where("sid", userID)
+
+ // Apply silent filter if provided
+ if filter.Silent != nil {
+ if *filter.Silent {
+ // Include all chats (both silent and non-silent)
+ } else {
+ // Only include non-silent chats
+ qb.Where("silent", false)
+ }
+ }
+
+ // Apply keyword filter if provided
+ if filter.Keywords != "" {
+ qb.Where("title", "like", fmt.Sprintf("%%%s%%", filter.Keywords))
+ }
+
+ // Apply pagination
+ offset := (filter.Page - 1) * filter.PageSize
+ qb.OrderBy("updated_at", filter.Order).
Offset(offset).
- Limit(filter.PageSize).
- Get()
+ Limit(filter.PageSize)
+
+ rows, err := qb.Get()
if err != nil {
return nil, err
}
@@ -385,16 +423,16 @@ func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, er
"Even Earlier": {},
}
- // Get assistant details for all chats
+ // Collect assistant IDs to fetch their details
assistantIDs := []interface{}{}
- assistantMap := make(map[string]map[string]interface{})
-
for _, row := range rows {
if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
assistantIDs = append(assistantIDs, assistantID)
}
}
+ // Fetch assistant details
+ assistantMap := map[string]map[string]interface{}{}
if len(assistantIDs) > 0 {
assistants, err := conv.query.New().
Table(conv.getAssistantTable()).
@@ -425,6 +463,7 @@ func (conv *Xun) GetChats(sid string, filter ChatFilter) (*ChatGroupResponse, er
"chat_id": chatID,
"title": row.Get("title"),
"assistant_id": row.Get("assistant_id"),
+ "silent": row.Get("silent"),
}
// Add assistant details if available
@@ -502,11 +541,14 @@ func (conv *Xun) GetHistory(sid string, cid string) ([]map[string]interface{}, e
}
qb := conv.newQuery().
- Select("role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "uid", "created_at", "updated_at").
+ Select("role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "uid", "silent", "created_at", "updated_at").
Where("sid", userID).
Where("cid", cid).
OrderBy("id", "desc")
+ // By default, exclude silent messages
+ qb.Where("silent", false)
+
if conv.setting.TTL > 0 {
qb.Where("expired_at", ">", time.Now())
}
@@ -533,6 +575,7 @@ func (conv *Xun) GetHistory(sid string, cid string) ([]map[string]interface{}, e
"assistant_avatar": row.Get("assistant_avatar"),
"mentions": row.Get("mentions"),
"uid": row.Get("uid"),
+ "silent": row.Get("silent"),
"created_at": row.Get("created_at"),
"updated_at": row.Get("updated_at"),
}
@@ -562,6 +605,23 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid
}
}
+ // Get silent flag from context
+ var silent bool = false
+ if context != nil {
+ if silentVal, ok := context["silent"]; ok {
+ switch v := silentVal.(type) {
+ case bool:
+ silent = v
+ case string:
+ silent = v == "true" || v == "1" || v == "yes"
+ case int:
+ silent = v != 0
+ case float64:
+ silent = v != 0
+ }
+ }
+ }
+
// First ensure chat record exists
exists, err := conv.newQueryChat().
Where("chat_id", cid).
@@ -579,6 +639,7 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid
"chat_id": cid,
"sid": userID,
"assistant_id": assistantID,
+ "silent": silent,
"created_at": time.Now(),
})
@@ -586,17 +647,16 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid
return err
}
} else {
- // Update assistant_id if it exists
- if assistantID != nil {
- _, err = conv.newQueryChat().
- Where("chat_id", cid).
- Where("sid", userID).
- Update(map[string]interface{}{
- "assistant_id": assistantID,
- })
- if err != nil {
- return err
- }
+ // Update assistant_id and silent if needed
+ _, err = conv.newQueryChat().
+ Where("chat_id", cid).
+ Where("sid", userID).
+ Update(map[string]interface{}{
+ "assistant_id": assistantID,
+ "silent": silent,
+ })
+ if err != nil {
+ return err
}
}
@@ -650,6 +710,7 @@ func (conv *Xun) SaveHistory(sid string, messages []map[string]interface{}, cid
"assistant_id": nil,
"assistant_name": nil,
"assistant_avatar": nil,
+ "silent": silent,
"created_at": now,
"updated_at": nil,
"expired_at": expiredAt,
@@ -736,7 +797,7 @@ func (conv *Xun) GetChat(sid string, cid string) (*ChatInfo, error) {
}
}
- // Get chat history
+ // Get chat history with default filter (silent=false)
history, err := conv.GetHistory(sid, cid)
if err != nil {
return nil, err
@@ -748,6 +809,64 @@ func (conv *Xun) GetChat(sid string, cid string) (*ChatInfo, error) {
}, nil
}
+// GetChatWithFilter get the chat info and its history with filter options
+func (conv *Xun) GetChatWithFilter(sid string, cid string, filter ChatFilter) (*ChatInfo, error) {
+ userID, err := conv.getUserID(sid)
+ if err != nil {
+ return nil, err
+ }
+
+ // Get chat info
+ qb := conv.newQueryChat().
+ Select("chat_id", "title", "assistant_id").
+ Where("sid", userID).
+ Where("chat_id", cid)
+
+ row, err := qb.First()
+ if err != nil {
+ return nil, err
+ }
+
+ // Return nil if chat_id is nil (means no chat found)
+ if row.Get("chat_id") == nil {
+ return nil, nil
+ }
+
+ chat := map[string]interface{}{
+ "chat_id": row.Get("chat_id"),
+ "title": row.Get("title"),
+ "assistant_id": row.Get("assistant_id"),
+ }
+
+ // Get assistant details if assistant_id exists
+ if assistantID := row.Get("assistant_id"); assistantID != nil && assistantID != "" {
+ assistant, err := conv.query.New().
+ Table(conv.getAssistantTable()).
+ Select("name", "avatar").
+ Where("assistant_id", assistantID).
+ First()
+ if err != nil {
+ return nil, err
+ }
+
+ if assistant != nil {
+ chat["assistant_name"] = assistant.Get("name")
+ chat["assistant_avatar"] = assistant.Get("avatar")
+ }
+ }
+
+ // Get chat history with filter
+ history, err := conv.GetHistoryWithFilter(sid, cid, filter)
+ if err != nil {
+ return nil, err
+ }
+
+ return &ChatInfo{
+ Chat: chat,
+ History: history,
+ }, nil
+}
+
// DeleteChat deletes a specific chat and its history
func (conv *Xun) DeleteChat(sid string, cid string) error {
userID, err := conv.getUserID(sid)
@@ -1164,3 +1283,74 @@ func (conv *Xun) GetAssistantTags() ([]string, error) {
}
return tags, nil
}
+
+// GetHistoryWithFilter get the history with filter options
+func (conv *Xun) GetHistoryWithFilter(sid string, cid string, filter ChatFilter) ([]map[string]interface{}, error) {
+ userID, err := conv.getUserID(sid)
+ if err != nil {
+ return nil, err
+ }
+
+ qb := conv.newQuery().
+ Select("role", "name", "content", "context", "assistant_id", "assistant_name", "assistant_avatar", "mentions", "uid", "silent", "created_at", "updated_at").
+ Where("sid", userID).
+ Where("cid", cid).
+ OrderBy("id", "desc")
+
+ // Apply silent filter if provided, otherwise exclude silent messages by default
+ if filter.Silent != nil {
+ if *filter.Silent {
+ // Include all messages (both silent and non-silent)
+ } else {
+ // Only include non-silent messages
+ qb.Where("silent", false)
+ }
+ } else {
+ // Default behavior: exclude silent messages
+ qb.Where("silent", false)
+ }
+
+ if conv.setting.TTL > 0 {
+ qb.Where("expired_at", ">", time.Now())
+ }
+
+ limit := 20
+ if conv.setting.MaxSize > 0 {
+ limit = conv.setting.MaxSize
+ }
+ if filter.PageSize > 0 {
+ limit = filter.PageSize
+ }
+
+ // Apply pagination if provided
+ if filter.Page > 0 {
+ offset := (filter.Page - 1) * limit
+ qb.Offset(offset)
+ }
+
+ rows, err := qb.Limit(limit).Get()
+ if err != nil {
+ return nil, err
+ }
+
+ res := []map[string]interface{}{}
+ for _, row := range rows {
+ message := map[string]interface{}{
+ "role": row.Get("role"),
+ "name": row.Get("name"),
+ "content": row.Get("content"),
+ "context": row.Get("context"),
+ "assistant_id": row.Get("assistant_id"),
+ "assistant_name": row.Get("assistant_name"),
+ "assistant_avatar": row.Get("assistant_avatar"),
+ "mentions": row.Get("mentions"),
+ "uid": row.Get("uid"),
+ "silent": row.Get("silent"),
+ "created_at": row.Get("created_at"),
+ "updated_at": row.Get("updated_at"),
+ }
+ res = append([]map[string]interface{}{message}, res...)
+ }
+
+ return res, nil
+}
diff --git a/neo/store/xun_test.go b/neo/store/xun_test.go
index e1a4d7d7..e1d1f1bd 100644
--- a/neo/store/xun_test.go
+++ b/neo/store/xun_test.go
@@ -996,3 +996,252 @@ func TestGetAssistantTags(t *testing.T) {
}
}
}
+
+func TestXunSaveAndGetHistoryWithSilent(t *testing.T) {
+ test.Prepare(t, config.Conf)
+ defer test.Clean()
+ defer capsule.Schema().DropTableIfExists("__unit_test_conversation_history")
+ defer capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
+
+ err := capsule.Schema().DropTableIfExists("__unit_test_conversation_history")
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ err = capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ store, err := NewXun(Setting{
+ Connector: "default",
+ Prefix: "__unit_test_conversation_",
+ TTL: 3600,
+ })
+
+ // save the history with silent messages
+ sid := "123456"
+ cid := "silent_test"
+
+ // First save regular messages
+ messages := []map[string]interface{}{
+ {"role": "user", "name": "user1", "content": "hello"},
+ {"role": "assistant", "name": "assistant1", "content": "Hi! How can I help you?"},
+ }
+ context := map[string]interface{}{
+ "assistant_id": "test-assistant-1",
+ }
+ err = store.SaveHistory(sid, messages, cid, context)
+ assert.Nil(t, err)
+
+ // Then save silent messages
+ silentMessages := []map[string]interface{}{
+ {"role": "user", "name": "user1", "content": "silent message"},
+ {"role": "assistant", "name": "assistant1", "content": "This is a silent response"},
+ }
+ silentContext := map[string]interface{}{
+ "assistant_id": "test-assistant-1",
+ "silent": true,
+ }
+ err = store.SaveHistory(sid, silentMessages, cid, silentContext)
+ assert.Nil(t, err)
+
+ // Get history without filter (should only return non-silent messages)
+ data, err := store.GetHistory(sid, cid)
+ if err != nil {
+ t.Fatal(err)
+ }
+ assert.Equal(t, 2, len(data))
+ for _, msg := range data {
+ // Check if silent is false, handling different types
+ isSilent := false
+ switch v := msg["silent"].(type) {
+ case bool:
+ isSilent = v
+ case int:
+ isSilent = v != 0
+ case int64:
+ isSilent = v != 0
+ case float64:
+ isSilent = v != 0
+ }
+ assert.False(t, isSilent, "message should not be silent")
+ }
+
+ // Get history with silent=true filter (should return all messages)
+ silentTrue := true
+ filter := ChatFilter{
+ Silent: &silentTrue,
+ }
+ allData, err := store.GetHistoryWithFilter(sid, cid, filter)
+ if err != nil {
+ t.Fatal(err)
+ }
+ assert.Equal(t, 4, len(allData))
+
+ // Count silent messages
+ silentCount := 0
+ for _, msg := range allData {
+ // Check if silent is true, handling different types
+ isSilent := false
+ switch v := msg["silent"].(type) {
+ case bool:
+ isSilent = v
+ case int:
+ isSilent = v != 0
+ case int64:
+ isSilent = v != 0
+ case float64:
+ isSilent = v != 0
+ }
+ if isSilent {
+ silentCount++
+ }
+ }
+ assert.Equal(t, 2, silentCount)
+
+ // Get chat with filter (should include silent messages)
+ chat, err := store.GetChatWithFilter(sid, cid, filter)
+ assert.Nil(t, err)
+ assert.Equal(t, 4, len(chat.History))
+
+ // Get chat without filter (should exclude silent messages)
+ chatNoSilent, err := store.GetChat(sid, cid)
+ assert.Nil(t, err)
+ assert.Equal(t, 2, len(chatNoSilent.History))
+}
+
+func TestXunGetChatsWithSilent(t *testing.T) {
+ test.Prepare(t, config.Conf)
+ defer test.Clean()
+ defer capsule.Schema().DropTableIfExists("__unit_test_conversation_history")
+ defer capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
+ defer capsule.Schema().DropTableIfExists("__unit_test_conversation_assistant")
+
+ // Drop tables before test
+ err := capsule.Schema().DropTableIfExists("__unit_test_conversation_history")
+ if err != nil {
+ t.Fatal(err)
+ }
+ err = capsule.Schema().DropTableIfExists("__unit_test_conversation_chat")
+ if err != nil {
+ t.Fatal(err)
+ }
+ err = capsule.Schema().DropTableIfExists("__unit_test_conversation_assistant")
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ store, err := NewXun(Setting{
+ Connector: "default",
+ Prefix: "__unit_test_conversation_",
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+
+ // Create test assistant
+ assistant := map[string]interface{}{
+ "assistant_id": "test-assistant-1",
+ "name": "Test Assistant 1",
+ "avatar": "avatar1.png",
+ "type": "assistant",
+ "connector": "test",
+ }
+ _, err = store.SaveAssistant(assistant)
+ assert.Nil(t, err)
+
+ // Save some test chats
+ sid := "test_user"
+ messages := []map[string]interface{}{
+ {"role": "user", "content": "test message"},
+ }
+
+ // Create regular chats
+ for i := 0; i < 3; i++ {
+ chatID := fmt.Sprintf("regular_chat_%d", i)
+ title := fmt.Sprintf("Regular Chat %d", i)
+ context := map[string]interface{}{
+ "assistant_id": "test-assistant-1",
+ "silent": false,
+ }
+
+ // Save history to create the chat
+ err = store.SaveHistory(sid, messages, chatID, context)
+ assert.Nil(t, err)
+
+ // Update the chat title
+ err = store.UpdateChatTitle(sid, chatID, title)
+ assert.Nil(t, err)
+ }
+
+ // Create silent chats
+ for i := 0; i < 2; i++ {
+ chatID := fmt.Sprintf("silent_chat_%d", i)
+ title := fmt.Sprintf("Silent Chat %d", i)
+ context := map[string]interface{}{
+ "assistant_id": "test-assistant-1",
+ "silent": true,
+ }
+
+ // Save history to create the chat
+ err = store.SaveHistory(sid, messages, chatID, context)
+ assert.Nil(t, err)
+
+ // Update the chat title
+ err = store.UpdateChatTitle(sid, chatID, title)
+ assert.Nil(t, err)
+ }
+
+ // Test GetChats with default filter (should exclude silent chats)
+ defaultFilter := ChatFilter{
+ PageSize: 10,
+ Order: "desc",
+ }
+ defaultGroups, err := store.GetChats(sid, defaultFilter)
+ assert.Nil(t, err)
+ assert.NotNil(t, defaultGroups)
+
+ // Count total chats in all groups
+ totalDefaultChats := 0
+ for _, group := range defaultGroups.Groups {
+ totalDefaultChats += len(group.Chats)
+ }
+ assert.Equal(t, 3, totalDefaultChats, "Default filter should only return non-silent chats")
+
+ // Test GetChats with silent=true filter (should include all chats)
+ silentTrue := true
+ silentFilter := ChatFilter{
+ PageSize: 10,
+ Order: "desc",
+ Silent: &silentTrue,
+ }
+ silentGroups, err := store.GetChats(sid, silentFilter)
+ assert.Nil(t, err)
+ assert.NotNil(t, silentGroups)
+
+ // Count total chats in all groups
+ totalSilentChats := 0
+ for _, group := range silentGroups.Groups {
+ totalSilentChats += len(group.Chats)
+ }
+ assert.Equal(t, 5, totalSilentChats, "Silent filter should return all chats")
+
+ // Test GetChats with silent=false filter (should only include non-silent chats)
+ silentFalse := false
+ nonSilentFilter := ChatFilter{
+ PageSize: 10,
+ Order: "desc",
+ Silent: &silentFalse,
+ }
+ nonSilentGroups, err := store.GetChats(sid, nonSilentFilter)
+ assert.Nil(t, err)
+ assert.NotNil(t, nonSilentGroups)
+
+ // Count total chats in all groups
+ totalNonSilentChats := 0
+ for _, group := range nonSilentGroups.Groups {
+ totalNonSilentChats += len(group.Chats)
+ }
+ assert.Equal(t, 3, totalNonSilentChats, "Non-silent filter should only return non-silent chats")
+}
diff --git a/script/script.go b/script/script.go
index f983e843..18b5fe48 100644
--- a/script/script.go
+++ b/script/script.go
@@ -25,6 +25,27 @@ func Load(cfg config.Config) error {
return err
}
+ // Load assistants - Move to the neo assistant package
+ // err = application.App.Walk("assistants", func(root, file string, isdir bool) error {
+ // if isdir {
+ // return nil
+ // }
+
+ // // Keep the src.index only
+ // if !strings.HasSuffix(file, "src/index.ts") {
+ // return nil
+ // }
+
+ // id := fmt.Sprintf("assistants.%s", share.ID(root, file))
+ // id = strings.TrimSuffix(id, ".src.index")
+ // _, err := v8.Load(file, id)
+ // return err
+ // }, exts...)
+
+ // if err != nil {
+ // return err
+ // }
+
return application.App.Walk("services", func(root, file string, isdir bool) error {
if isdir {
return nil
diff --git a/sui/api/request.go b/sui/api/request.go
index 8ca48a40..6b45d5ef 100644
--- a/sui/api/request.go
+++ b/sui/api/request.go
@@ -141,7 +141,7 @@ func (r *Request) Render() (string, int, error) {
if c.Data != "" {
err = r.Request.ExecStringMerge(data, c.Data)
if err != nil {
- return "", 500, fmt.Errorf("data error, please re-complie the page. %s", err.Error())
+ return "", 500, fmt.Errorf("data merge error, please re-complie the page. %s", err.Error())
}
}
diff --git a/sui/core/request.go b/sui/core/request.go
index 43d8f71b..3382d447 100644
--- a/sui/core/request.go
+++ b/sui/core/request.go
@@ -218,7 +218,12 @@ func (r *Request) execValue(value interface{}) (interface{}, error) {
}
if strings.HasPrefix(v, "$") {
- return r.call(strings.TrimLeft(v, "$"))
+ res, err := r.call(strings.TrimLeft(v, "$"))
+ if err != nil {
+ log.Error("[Request] Exec value:%s, %s", v, err.Error())
+ return nil, nil
+ }
+ return res, nil
}
return v, nil