diff --git a/neo/assistant/api.go b/neo/assistant/api.go index 422f08e2..92f2530f 100644 --- a/neo/assistant/api.go +++ b/neo/assistant/api.go @@ -10,6 +10,7 @@ import ( jsoniter "github.com/json-iterator/go" "github.com/yaoapp/gou/fs" "github.com/yaoapp/gou/process" + "github.com/yaoapp/kun/utils" chatctx "github.com/yaoapp/yao/neo/context" chatMessage "github.com/yaoapp/yao/neo/message" ) @@ -47,16 +48,15 @@ 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 { contents := chatMessage.NewContents() - return ast.execute(c, ctx, input, options, contents) -} - -// Execute implements the execute functionality -func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input string, options map[string]interface{}, contents *chatMessage.Contents) error { - messages, err := ast.withHistory(ctx, input) if err != nil { return err } + return ast.execute(c, ctx, messages, options, contents) +} + +// Execute implements the execute functionality +func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input []chatMessage.Message, options map[string]interface{}, contents *chatMessage.Contents) error { if contents == nil { contents = chatMessage.NewContents() @@ -68,7 +68,7 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input string, ctx.Version = ast.vision // Run init hook - res, err := ast.HookInit(c, ctx, messages, options, contents) + res, err := ast.HookInit(c, ctx, input, options, contents) if err != nil { chatMessage.New(). Assistant(ast.ID, ast.Name, ast.Avatar). @@ -104,11 +104,11 @@ func (ast *Assistant) execute(c *gin.Context, ctx chatctx.Context, input string, // messages if res != nil && res.Input != nil { - messages = res.Input + input = res.Input } // Only proceed with chat stream if no specific next action was handled - return ast.handleChatStream(c, ctx, messages, options, contents) + return ast.handleChatStream(c, ctx, input, options, contents) } // Execute the next action @@ -163,9 +163,36 @@ func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *c } // Input - input, ok := next.Payload["input"].(string) - if !ok { - return fmt.Errorf("input should be string") + input := chatMessage.Message{} + _, has := next.Payload["input"] + if !has { + return fmt.Errorf("input is required") + } + + switch v := next.Payload["input"].(type) { + case string: + messages := chatMessage.Message{} + err := jsoniter.UnmarshalFromString(v, &messages) + if err != nil { + return 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()) + } + input = *msg + + case *chatMessage.Message: + input = *v + + case chatMessage.Message: + input = v + + default: + return fmt.Errorf("input should be string or []chatMessage.Message") } // Options @@ -173,7 +200,17 @@ func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context, contents *c if v, ok := next.Payload["options"].(map[string]interface{}); ok { options = v } - return assistant.execute(c, ctx, input, options, contents) + + messages, err := assistant.withHistory(ctx, input) + if err != nil { + return fmt.Errorf("with history error: %s", err.Error()) + } + + fmt.Println("---messages ---") + utils.Dump(messages) + fmt.Println(`chatID: `, ctx.ChatID) + + return assistant.execute(c, ctx, messages, options, contents) case "exit": return nil @@ -220,6 +257,7 @@ func (ast *Assistant) handleChatStream(c *gin.Context, ctx chatctx.Context, mess } ast.saveChatHistory(ctx, messages, contents) + fmt.Printf("saveChatHistory %v\n", ctx.ChatID) done <- true }() @@ -243,7 +281,8 @@ func (ast *Assistant) streamChat( done chan bool, contents *chatMessage.Contents) error { - return ast.Chat(c.Request.Context(), messages, options, func(data []byte) int { + errorRaw := "" + err := ast.Chat(c.Request.Context(), messages, options, func(data []byte) int { select { case <-clientBreak: return 0 // break @@ -254,6 +293,11 @@ func (ast *Assistant) streamChat( return 1 // continue } + if msg.Pending { + errorRaw += msg.Text + return 1 // continue + } + // Handle error if msg.Type == "error" { value := msg.String() @@ -350,6 +394,22 @@ func (ast *Assistant) streamChat( return 1 // continue } }) + + // Handle error + if err != nil { + return err + } + + // raw error + if errorRaw != "" { + msg, err := chatMessage.NewStringError(errorRaw) + if err != nil { + return fmt.Errorf("error: %s", err.Error()) + } + msg.Done().Write(c.Writer) + } + + return nil } // saveChatHistory saves the chat history if storage is available @@ -372,6 +432,11 @@ func (ast *Assistant) saveChatHistory(ctx chatctx.Context, messages []chatMessag }, } + // contents + fmt.Println("---contents ---") + utils.Dump(contents) + fmt.Println("---contents end ---") + // Add mentions if userMessage.Mentions != nil { data[0]["mentions"] = userMessage.Mentions @@ -417,7 +482,22 @@ func (ast *Assistant) withPrompts(messages []chatMessage.Message) []chatMessage. return messages } -func (ast *Assistant) withHistory(ctx chatctx.Context, input string) ([]chatMessage.Message, error) { +func (ast *Assistant) withHistory(ctx chatctx.Context, input interface{}) ([]chatMessage.Message, error) { + + var userMessage *chatMessage.Message = chatMessage.New() + switch v := input.(type) { + case string: + userMessage.Map(map[string]interface{}{"role": "user", "content": v}) + case map[string]interface{}: + userMessage.Map(v) + case chatMessage.Message: + userMessage = &v + case *chatMessage.Message: + userMessage = v + default: + return nil, fmt.Errorf("unknown input type: %T", input) + } + messages := []chatMessage.Message{} messages = ast.withPrompts(messages) if storage != nil { @@ -437,7 +517,7 @@ func (ast *Assistant) withHistory(ctx chatctx.Context, input string) ([]chatMess } // Add user message - messages = append(messages, *chatMessage.New().Map(map[string]interface{}{"role": "user", "content": input, "name": ctx.Sid})) + messages = append(messages, *userMessage) return messages, nil } diff --git a/neo/assistant/load.go b/neo/assistant/load.go index 56f166bf..2811fa82 100644 --- a/neo/assistant/load.go +++ b/neo/assistant/load.go @@ -2,6 +2,7 @@ package assistant import ( "fmt" + "os" "path/filepath" "regexp" "strings" @@ -9,6 +10,7 @@ import ( jsoniter "github.com/json-iterator/go" "github.com/spf13/cast" + "github.com/yaoapp/gou/application" "github.com/yaoapp/gou/fs" "github.com/yaoapp/gou/rag/driver" v8 "github.com/yaoapp/gou/runtime/v8" @@ -195,8 +197,8 @@ func LoadStore(id string) (*Assistant, error) { return assistant, nil } -// LoadPath load assistant from path -func LoadPath(path string) (*Assistant, error) { +// loadPackage loads and parses the package.yao file +func loadPackage(path string) (map[string]interface{}, error) { app, err := fs.Get("app") if err != nil { return nil, err @@ -207,19 +209,44 @@ func LoadPath(path string) (*Assistant, error) { return nil, fmt.Errorf("package.yao not found in %s", path) } - pkg, err := app.ReadFile(pkgfile) + pkgraw, err := app.ReadFile(pkgfile) if err != nil { return nil, err } - id := strings.ReplaceAll(strings.TrimPrefix(path, "/assistants/"), "/", ".") var data map[string]interface{} - err = jsoniter.Unmarshal(pkg, &data) + err = application.Parse(pkgfile, pkgraw, &data) + if err != nil { + return nil, err + } + + // Process connector environment variable + if connector, ok := data["connector"].(string); ok { + if strings.HasPrefix(connector, "$ENV.") { + envKey := strings.TrimPrefix(connector, "$ENV.") + if envValue := os.Getenv(envKey); envValue != "" { + data["connector"] = envValue + } + } + } + + return data, nil +} + +// LoadPath load assistant from path +func LoadPath(path string) (*Assistant, error) { + app, err := fs.Get("app") + if err != nil { + return nil, err + } + + data, err := loadPackage(path) if err != nil { return nil, err } // assistant_id + id := strings.ReplaceAll(strings.TrimPrefix(path, "/assistants/"), "/", ".") data["assistant_id"] = id data["type"] = "assistant" data["path"] = path diff --git a/neo/message/message.go b/neo/message/message.go index 5de1d157..80e0001c 100644 --- a/neo/message/message.go +++ b/neo/message/message.go @@ -29,7 +29,8 @@ type Message struct { AssistantName string `json:"assistant_name,omitempty"` // assistant_name (for assistant role = assistant ) AssistantAvatar string `json:"assistant_avatar,omitempty"` // assistant_avatar (for assistant role = assistant ) Mentions []Mention `json:"menions,omitempty"` // Mentions for the message ( for user role = user ) - Data map[string]interface{} `json:"-"` + Data map[string]interface{} `json:"-"` // data for the message + Pending bool `json:"-"` // pending for the message } // Mention represents a mention @@ -144,6 +145,38 @@ func NewString(content string) (*Message, error) { return &Message{Text: content}, nil } +// NewStringError create a new message from string error +func NewStringError(content string) (*Message, error) { + if strings.HasPrefix(content, "{") && strings.HasSuffix(content, "}") { + var msg = New() + var errorMessage openai.ErrorMessage + if err := jsoniter.UnmarshalFromString(content, &errorMessage); err != nil { + msg.Text = err.Error() + "\n" + content + return msg, nil + } + msg.Type = "error" + msg.Text = errorMessage.Error.Message + return msg, nil + } + return &Message{Text: content}, nil +} + +// NewMap create a new message from map +func NewMap(content map[string]interface{}) (*Message, error) { + return New().Map(content), nil +} + +// NewAny create a new message from any content +func NewAny(content interface{}) (*Message, error) { + switch v := content.(type) { + case string: + return NewString(v) + case map[string]interface{}: + return NewMap(v) + } + return nil, fmt.Errorf("unknown content type: %T", content) +} + // NewOpenAI create a new message from OpenAI response func NewOpenAI(data []byte) *Message { if data == nil || len(data) == 0 { @@ -199,6 +232,11 @@ func NewOpenAI(data []byte) *Message { case strings.Contains(text, `"finish_reason":"tool_calls"`): msg.IsDone = true + // Not a data message + case !strings.Contains(text, `data: `): + msg.Pending = true + msg.Text = text + default: str := strings.TrimPrefix(strings.Trim(string(data), "\""), "data: ") msg.Type = "error" diff --git a/neo/neo.go b/neo/neo.go index 67d57caf..a13a85ea 100644 --- a/neo/neo.go +++ b/neo/neo.go @@ -123,6 +123,7 @@ func (neo *DSL) GenerateWithAI(ctx chatctx.Context, input string, messageType st } } + errorRaw := "" err := ast.Chat(c.Request.Context(), msgList, neo.Option, func(data []byte) int { select { case <-clientBreak: @@ -134,6 +135,11 @@ func (neo *DSL) GenerateWithAI(ctx chatctx.Context, input string, messageType st return 1 // continue } + if msg.Pending { + errorRaw += msg.Text + return 1 // continue + } + // Handle error if msg.Type == "error" { fail <- fmt.Errorf("%s", msg.Text) @@ -175,6 +181,14 @@ func (neo *DSL) GenerateWithAI(ctx chatctx.Context, input string, messageType st } } + if errorRaw != "" { + msg, err := message.NewStringError(errorRaw) + if err != nil { + log.Error("Error parsing error message: %s", err.Error()) + } + msg.Write(c.Writer) + } + done <- true }()