Merge pull request #817 from trheyi/main

Added function support and action execution capabilities to Neo API assistant, improving its flexibility and robustness. Key updates include function management, enhanced message handling, and improved error validation.
This commit is contained in:
Max 2025-01-15 11:26:14 +08:00 committed by GitHub
commit cad8aa2c1a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 309 additions and 23 deletions

View file

@ -8,6 +8,7 @@ import (
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/yaoapp/gou/fs" "github.com/yaoapp/gou/fs"
"github.com/yaoapp/gou/process"
chatctx "github.com/yaoapp/yao/neo/context" chatctx "github.com/yaoapp/yao/neo/context"
"github.com/yaoapp/yao/neo/message" "github.com/yaoapp/yao/neo/message"
chatMessage "github.com/yaoapp/yao/neo/message" chatMessage "github.com/yaoapp/yao/neo/message"
@ -69,11 +70,7 @@ func (ast *Assistant) Execute(c *gin.Context, ctx chatctx.Context, input string,
// Handle next action // Handle next action
if res.Next != nil { if res.Next != nil {
switch res.Next.Action { return res.Next.Execute(c, ctx)
case "exit":
return nil
// Add other actions here if needed
}
} }
// Update options if provided // Update options if provided
@ -90,15 +87,87 @@ func (ast *Assistant) Execute(c *gin.Context, ctx chatctx.Context, input string,
return ast.handleChatStream(c, ctx, messages, options) return ast.handleChatStream(c, ctx, messages, options)
} }
// Execute the next action
func (next *NextAction) Execute(c *gin.Context, ctx chatctx.Context) error {
switch next.Action {
case "process":
if next.Payload == nil {
return fmt.Errorf("payload is required")
}
name, ok := next.Payload["name"].(string)
if !ok {
return fmt.Errorf("process name should be string")
}
args := []interface{}{}
if v, ok := next.Payload["args"].([]interface{}); ok {
args = v
}
// Add context and writer to args
args = append(args, ctx, c.Writer)
p, err := process.Of(name, args...)
if err != nil {
return fmt.Errorf("get process error: %s", err.Error())
}
err = p.Execute()
if err != nil {
return fmt.Errorf("execute process error: %s", err.Error())
}
defer p.Release()
return nil
case "assistant":
if next.Payload == nil {
return 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")
}
// Get assistant
assistant, err := Get(id)
if err != nil {
return fmt.Errorf("get assistant error: %s", err.Error())
}
// Input
input, ok := next.Payload["input"].(string)
if !ok {
return fmt.Errorf("input should be string")
}
// Options
options := map[string]interface{}{}
if v, ok := next.Payload["options"].(map[string]interface{}); ok {
options = v
}
return assistant.Execute(c, ctx, input, options)
case "exit":
return nil
default:
return fmt.Errorf("unknown action: %s", next.Action)
}
}
// handleChatStream manages the streaming chat interaction with the AI // handleChatStream manages the streaming chat interaction with the AI
func (ast *Assistant) handleChatStream(c *gin.Context, ctx chatctx.Context, messages []message.Message, options map[string]interface{}) error { func (ast *Assistant) handleChatStream(c *gin.Context, ctx chatctx.Context, messages []message.Message, options map[string]interface{}) error {
clientBreak := make(chan bool, 1) clientBreak := make(chan bool, 1)
done := make(chan bool, 1) done := make(chan bool, 1)
content := []byte{} content := message.NewContent("text")
// Chat with AI in background // Chat with AI in background
go func() { go func() {
err := ast.streamChat(c, ctx, messages, options, clientBreak, done, &content) err := ast.streamChat(c, ctx, messages, options, clientBreak, done, content)
if err != nil { if err != nil {
chatMessage.New().Error(err).Done().Write(c.Writer) chatMessage.New().Error(err).Done().Write(c.Writer)
} }
@ -118,8 +187,14 @@ func (ast *Assistant) handleChatStream(c *gin.Context, ctx chatctx.Context, mess
} }
// streamChat handles the streaming chat interaction // streamChat handles the streaming chat interaction
func (ast *Assistant) streamChat(c *gin.Context, ctx chatctx.Context, messages []message.Message, options map[string]interface{}, func (ast *Assistant) streamChat(
clientBreak chan bool, done chan bool, content *[]byte) error { c *gin.Context,
ctx chatctx.Context,
messages []message.Message,
options map[string]interface{},
clientBreak chan bool,
done chan bool,
content *message.Content) error {
return ast.Chat(c.Request.Context(), messages, options, func(data []byte) int { return ast.Chat(c.Request.Context(), messages, options, func(data []byte) int {
select { select {
@ -135,7 +210,7 @@ func (ast *Assistant) streamChat(c *gin.Context, ctx chatctx.Context, messages [
// Handle error // Handle error
if msg.Type == "error" { if msg.Type == "error" {
value := msg.String() value := msg.String()
res, hookErr := ast.HookFail(c, ctx, messages, string(*content), fmt.Errorf("%s", value)) res, hookErr := ast.HookFail(c, ctx, messages, content.String(), fmt.Errorf("%s", value))
if hookErr == nil && res != nil && (res.Output != "" || res.Error != "") { if hookErr == nil && res != nil && (res.Output != "" || res.Error != "") {
value = res.Output value = res.Output
if res.Error != "" { if res.Error != "" {
@ -146,20 +221,41 @@ func (ast *Assistant) streamChat(c *gin.Context, ctx chatctx.Context, messages [
return 0 // break return 0 // break
} }
// Handle tool call
if msg.Type == "tool_calls" {
content.SetType("function") // Set type to function
// Set id
if id, ok := msg.Props["id"].(string); ok && id != "" {
content.SetID(id)
}
// Set name
if name, ok := msg.Props["name"].(string); ok && name != "" {
content.SetName(name)
}
}
// Append content and send message // Append content and send message
*content = msg.Append(*content)
value := msg.String() value := msg.String()
content.Append(value)
if value != "" { if value != "" {
// Handle stream // Handle stream
res, err := ast.HookStream(c, ctx, messages, string(*content)) res, err := ast.HookStream(c, ctx, messages, content.String(), msg.Type == "tool_calls")
if err == nil && res != nil { if err == nil && res != nil {
if res.Output != "" { if res.Output != "" {
value = res.Output value = res.Output
} }
if res.Next != nil && res.Next.Action == "exit" {
if res.Next != nil {
err = res.Next.Execute(c, ctx)
if err != nil {
chatMessage.New().Error(err.Error()).Done().Write(c.Writer)
}
done <- true done <- true
return 0 // break return 0 // break
} }
if res.Silent { if res.Silent {
return 1 // continue return 1 // continue
} }
@ -180,7 +276,8 @@ func (ast *Assistant) streamChat(c *gin.Context, ctx chatctx.Context, messages [
// } // }
// Call HookDone // Call HookDone
res, hookErr := ast.HookDone(c, ctx, messages, string(*content)) content.SetStatus(message.ContentStatusDone)
res, hookErr := ast.HookDone(c, ctx, messages, content.String(), msg.Type == "tool_calls")
if hookErr == nil && res != nil { if hookErr == nil && res != nil {
if res.Output != "" { if res.Output != "" {
chatMessage.New(). chatMessage.New().
@ -190,10 +287,16 @@ func (ast *Assistant) streamChat(c *gin.Context, ctx chatctx.Context, messages [
}). }).
Write(c.Writer) Write(c.Writer)
} }
if res.Next != nil && res.Next.Action == "exit" {
if res.Next != nil {
err := res.Next.Execute(c, ctx)
if err != nil {
chatMessage.New().Error(err.Error()).Done().Write(c.Writer)
}
done <- true done <- true
return 0 // break return 0 // break
} }
} else if value != "" { } else if value != "" {
chatMessage.New(). chatMessage.New().
Map(map[string]interface{}{ Map(map[string]interface{}{
@ -213,13 +316,13 @@ func (ast *Assistant) streamChat(c *gin.Context, ctx chatctx.Context, messages [
} }
// saveChatHistory saves the chat history if storage is available // saveChatHistory saves the chat history if storage is available
func (ast *Assistant) saveChatHistory(ctx chatctx.Context, messages []message.Message, content []byte) { func (ast *Assistant) saveChatHistory(ctx chatctx.Context, messages []message.Message, content *message.Content) {
if len(content) > 0 && ctx.Sid != "" && len(messages) > 0 { if len(content.Bytes) > 0 && ctx.Sid != "" && len(messages) > 0 {
storage.SaveHistory( storage.SaveHistory(
ctx.Sid, ctx.Sid,
[]map[string]interface{}{ []map[string]interface{}{
{"role": "user", "content": messages[len(messages)-1].Content(), "name": ctx.Sid}, {"role": "user", "content": messages[len(messages)-1].Content(), "name": ctx.Sid},
{"role": "assistant", "content": string(content), "name": ctx.Sid}, {"role": "assistant", "content": content.String(), "name": ctx.Sid},
}, },
ctx.ChatID, ctx.ChatID,
nil, nil,
@ -237,6 +340,15 @@ func (ast *Assistant) withOptions(options map[string]interface{}) map[string]int
options[key] = value options[key] = value
} }
} }
// Add functions
if ast.Functions != nil {
options["tools"] = ast.Functions
if options["tool_choice"] == nil {
options["tool_choice"] = "auto"
}
}
return options return options
} }

View file

@ -132,6 +132,7 @@ func (ast *Assistant) Map() map[string]interface{} {
"description": ast.Description, "description": ast.Description,
"options": ast.Options, "options": ast.Options,
"prompts": ast.Prompts, "prompts": ast.Prompts,
"functions": ast.Functions,
"tags": ast.Tags, "tags": ast.Tags,
"mentionable": ast.Mentionable, "mentionable": ast.Mentionable,
"automated": ast.Automated, "automated": ast.Automated,

View file

@ -57,13 +57,13 @@ func (ast *Assistant) HookInit(c *gin.Context, context chatctx.Context, input []
} }
// HookStream Handle streaming response from LLM // HookStream Handle streaming response from LLM
func (ast *Assistant) HookStream(c *gin.Context, context chatctx.Context, input []message.Message, output string) (*ResHookStream, error) { func (ast *Assistant) HookStream(c *gin.Context, context chatctx.Context, input []message.Message, output string, toolcall bool) (*ResHookStream, error) {
// Create timeout context // Create timeout context
ctx, cancel := ast.createTimeoutContext(c) ctx, cancel := ast.createTimeoutContext(c)
defer cancel() defer cancel()
v, err := ast.call(ctx, "Stream", context, input, output, c.Writer) v, err := ast.call(ctx, "Stream", context, input, output, toolcall, c.Writer)
if err != nil { if err != nil {
if err.Error() == HookErrorMethodNotFound { if err.Error() == HookErrorMethodNotFound {
return nil, nil return nil, nil
@ -100,12 +100,12 @@ func (ast *Assistant) HookStream(c *gin.Context, context chatctx.Context, input
} }
// HookDone Handle completion of assistant response // HookDone Handle completion of assistant response
func (ast *Assistant) HookDone(c *gin.Context, context chatctx.Context, input []message.Message, output string) (*ResHookDone, error) { func (ast *Assistant) HookDone(c *gin.Context, context chatctx.Context, input []message.Message, output string, toolcall bool) (*ResHookDone, error) {
// Create timeout context // Create timeout context
ctx, cancel := ast.createTimeoutContext(c) ctx, cancel := ast.createTimeoutContext(c)
defer cancel() defer cancel()
v, err := ast.call(ctx, "Done", context, input, output, c.Writer) v, err := ast.call(ctx, "Done", context, input, output, toolcall, c.Writer)
if err != nil { if err != nil {
if err.Error() == HookErrorMethodNotFound { if err.Error() == HookErrorMethodNotFound {
return nil, nil return nil, nil

View file

@ -246,6 +246,16 @@ func LoadPath(path string) (*Assistant, error) {
} }
// load functions // load functions
functionsfile := filepath.Join(path, "functions.json")
if has, _ := app.Exists(functionsfile); has {
functions, ts, err := loadFunctions(functionsfile)
if err != nil {
return nil, err
}
data["functions"] = functions
updatedAt = max(updatedAt, ts)
data["updated_at"] = updatedAt
}
// load flow // load flow
@ -340,6 +350,25 @@ func loadMap(data map[string]interface{}) (*Assistant, error) {
assistant.Prompts = prompts assistant.Prompts = prompts
} }
// functions
if funcs, has := data["functions"]; has {
switch vv := funcs.(type) {
case []Function:
assistant.Functions = vv
default:
raw, err := jsoniter.Marshal(vv)
if err != nil {
return nil, err
}
var functions []Function
err = jsoniter.Unmarshal(raw, &functions)
if err != nil {
return nil, err
}
assistant.Functions = functions
}
}
// script // script
if data["script"] != nil { if data["script"] != nil {
switch v := data["script"].(type) { switch v := data["script"].(type) {
@ -382,6 +411,32 @@ func loadMap(data map[string]interface{}) (*Assistant, error) {
return assistant, nil return assistant, nil
} }
func loadFunctions(file string) ([]Function, int64, error) {
app, err := fs.Get("app")
if err != nil {
return nil, 0, err
}
ts, err := app.ModTime(file)
if err != nil {
return nil, 0, err
}
raw, err := app.ReadFile(file)
if err != nil {
return nil, 0, err
}
var functions []Function
err = jsoniter.Unmarshal(raw, &functions)
if err != nil {
return nil, 0, err
}
return functions, ts.UnixNano(), nil
}
func loadPrompts(file string, root string) (string, int64, error) { func loadPrompts(file string, root string) (string, int64, error) {
app, err := fs.Get("app") app, err := fs.Get("app")

View file

@ -85,6 +85,16 @@ type Prompt struct {
Name string `json:"name,omitempty"` Name string `json:"name,omitempty"`
} }
// Function a function
type Function struct {
Type string `json:"type"`
Function struct {
Name string `json:"name"`
Description string `json:"description"`
Parameters map[string]interface{} `json:"parameters"`
} `json:"function"`
}
// QueryParam the assistant query param // QueryParam the assistant query param
type QueryParam struct { type QueryParam struct {
Limit uint `json:"limit"` Limit uint `json:"limit"`
@ -110,6 +120,7 @@ type Assistant struct {
Automated bool `json:"automated,omitempty"` // Whether this assistant is automated Automated bool `json:"automated,omitempty"` // Whether this assistant is automated
Options map[string]interface{} `json:"options,omitempty"` // AI Options Options map[string]interface{} `json:"options,omitempty"` // AI Options
Prompts []Prompt `json:"prompts,omitempty"` // AI Prompts Prompts []Prompt `json:"prompts,omitempty"` // AI Prompts
Functions []Function `json:"functions,omitempty"` // Assistant Functions
Flows []map[string]interface{} `json:"flows,omitempty"` // Assistant Flows Flows []map[string]interface{} `json:"flows,omitempty"` // Assistant Flows
Script *v8.Script `json:"-" yaml:"-"` // Assistant Script Script *v8.Script `json:"-" yaml:"-"` // Assistant Script
CreatedAt int64 `json:"created_at"` // Creation timestamp CreatedAt int64 `json:"created_at"` // Creation timestamp

67
neo/message/content.go Normal file
View file

@ -0,0 +1,67 @@
package message
import "fmt"
const (
// ContentStatusPending the content status pending
ContentStatusPending = iota
// ContentStatusDone the content status done
ContentStatusDone
// ContentStatusError the content status error
ContentStatusError
)
// Content the content
type Content struct {
ID string `json:"id"`
Name string `json:"name"`
Bytes []byte `json:"bytes"`
Type string `json:"type"` // text, function, error
Status uint8 `json:"status"` // 0: pending, 1: done
}
// NewContent create a new content
func NewContent(typ string) *Content {
if typ == "" {
typ = "text"
}
return &Content{
Bytes: []byte{},
Type: typ,
Status: ContentStatusPending,
}
}
// String the content string
func (c *Content) String() string {
if c.Type == "function" {
return fmt.Sprintf(`{"id":"%s","type": "function", "function": {"name": "%s", "arguments": "%s"}}`, c.ID, c.Name, c.Bytes)
}
return string(c.Bytes)
}
// SetID set the content id
func (c *Content) SetID(id string) {
c.ID = id
}
// SetName set the content name
func (c *Content) SetName(name string) {
c.Name = name
}
// SetType set the content type
func (c *Content) SetType(typ string) {
c.Type = typ
}
// Append append the content
func (c *Content) Append(data string) {
c.Bytes = append(c.Bytes, []byte(data)...)
}
// SetStatus set the content status
func (c *Content) SetStatus(status uint8) {
c.Status = status
}

View file

@ -49,7 +49,7 @@ type Action struct {
// New create a new message // New create a new message
func New() *Message { func New() *Message {
return &Message{Actions: []Action{}} return &Message{Actions: []Action{}, Props: map[string]interface{}{}}
} }
// NewString create a new message from string // NewString create a new message from string
@ -75,6 +75,21 @@ func NewOpenAI(data []byte) *Message {
data = []byte(strings.TrimPrefix(text, "data: ")) data = []byte(strings.TrimPrefix(text, "data: "))
switch { switch {
case strings.Contains(text, `"delta":{`) && strings.Contains(text, `"tool_calls"`):
var toolCalls openai.ToolCalls
if err := jsoniter.Unmarshal(data, &toolCalls); err != nil {
msg.Text = err.Error() + "\n" + string(data)
return msg
}
msg.Type = "tool_calls"
if len(toolCalls.Choices) > 0 && len(toolCalls.Choices[0].Delta.ToolCalls) > 0 {
msg.Props["id"] = toolCalls.Choices[0].Delta.ToolCalls[0].ID
msg.Props["name"] = toolCalls.Choices[0].Delta.ToolCalls[0].Function.Name
msg.Text = toolCalls.Choices[0].Delta.ToolCalls[0].Function.Arguments
}
case strings.Contains(text, `"delta":{`) && strings.Contains(text, `"content":`): case strings.Contains(text, `"delta":{`) && strings.Contains(text, `"content":`):
var message openai.Message var message openai.Message
if err := jsoniter.Unmarshal(data, &message); err != nil { if err := jsoniter.Unmarshal(data, &message); err != nil {
@ -92,6 +107,9 @@ func NewOpenAI(data []byte) *Message {
case strings.Contains(text, `"finish_reason":"stop"`): case strings.Contains(text, `"finish_reason":"stop"`):
msg.IsDone = true msg.IsDone = true
case strings.Contains(text, `"finish_reason":"tool_calls"`):
msg.IsDone = true
default: default:
str := strings.TrimPrefix(strings.Trim(string(data), "\""), "data: ") str := strings.TrimPrefix(strings.Trim(string(data), "\""), "data: ")
msg.Type = "error" msg.Type = "error"

View file

@ -16,6 +16,28 @@ type Message struct {
} `json:"choices,omitempty"` } `json:"choices,omitempty"`
} }
// ToolCalls is the response from OpenAI
type ToolCalls struct {
ID string `json:"id,omitempty"`
Object string `json:"object,omitempty"`
Created int64 `json:"created,omitempty"`
Model string `json:"model,omitempty"`
Choices []struct {
Delta struct {
ToolCalls []struct {
ID string `json:"id,omitempty"`
Type string `json:"type,omitempty"`
Function struct {
Name string `json:"name,omitempty"`
Arguments string `json:"arguments,omitempty"`
} `json:"function,omitempty"`
} `json:"tool_calls,omitempty"`
} `json:"delta,omitempty"`
Index int `json:"index,omitempty"`
FinishReason string `json:"finish_reason,omitempty"`
} `json:"choices,omitempty"`
}
// ErrorMessage is the error response from OpenAI // ErrorMessage is the error response from OpenAI
type ErrorMessage struct { type ErrorMessage struct {
Error Error `json:"error,omitempty"` Error Error `json:"error,omitempty"`