Merge pull request #779 from trheyi/main

Refactor Neo message handling to use gin.ResponseWriter and improve e…
This commit is contained in:
Max 2024-11-07 08:48:43 +08:00 committed by GitHub
commit 13e94aebed
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 200 additions and 160 deletions

View file

@ -1,9 +1,9 @@
package message package message
import ( import (
"io"
"strings" "strings"
"github.com/gin-gonic/gin"
jsoniter "github.com/json-iterator/go" jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/gou/helper" "github.com/yaoapp/gou/helper"
"github.com/yaoapp/kun/log" "github.com/yaoapp/kun/log"
@ -47,6 +47,10 @@ func NewOpenAI(data []byte) *JSON {
msg.Done = true msg.Done = true
break break
case strings.Contains(text, `"finish_reason":"stop"`):
msg.Done = true
break
default: default:
msg.Error = text msg.Error = text
} }
@ -189,7 +193,13 @@ func (json *JSON) IsDone() bool {
} }
// Write the message // Write the message
func (json *JSON) Write(w io.Writer) bool { func (json *JSON) Write(w gin.ResponseWriter) bool {
defer func() {
if r := recover(); r != nil {
log.Error("Write JSON Message Error: %s", r)
}
}()
data, err := jsoniter.Marshal(json.Message) data, err := jsoniter.Marshal(json.Message)
if err != nil { if err != nil {
@ -205,7 +215,7 @@ func (json *JSON) Write(w io.Writer) bool {
log.Error("%s", err.Error()) log.Error("%s", err.Error())
return false return false
} }
w.Flush()
return true return true
} }

View file

@ -2,14 +2,12 @@ package neo
import ( import (
"fmt" "fmt"
"io"
"net/url" "net/url"
"strings" "strings"
"github.com/fatih/color" "github.com/fatih/color"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
"github.com/google/uuid" "github.com/google/uuid"
jsoniter "github.com/json-iterator/go"
"github.com/yaoapp/gou/api" "github.com/yaoapp/gou/api"
"github.com/yaoapp/gou/connector" "github.com/yaoapp/gou/connector"
"github.com/yaoapp/gou/process" "github.com/yaoapp/gou/process"
@ -172,150 +170,167 @@ func (neo *DSL) API(router *gin.Engine, path string) error {
return nil return nil
} }
// Answer the message // Answer reply the message
func (neo *DSL) Answer(ctx command.Context, question string, answer Answer) error { func (neo *DSL) Answer(ctx command.Context, question string, c *gin.Context) error {
chanStream := make(chan *message.JSON, 1)
chanError := make(chan error, 1)
content := []byte{}
errorMsg := []byte{}
// get the chat messages // get the chat messages
messages, err := neo.chatMessages(ctx, question) messages, err := neo.chatMessages(ctx, question)
if err != nil { if err != nil {
return err return err
} }
// check the command clientBreak := make(chan bool, 1)
cmd, isCommand := neo.matchCommand(ctx, messages) done := make(chan bool, 1)
content := []byte{}
// Execute the command or chat with AI in the background
go func() { go func() {
defer func() {
close(chanStream)
close(chanError)
}()
// execute the command // chat with AI
c.Header("Content-Type", "text/event-stream;charset=utf-8")
c.Header("Cache-Control", "no-cache")
c.Header("Connection", "keep-alive")
// check the command
cmd, isCommand := neo.matchCommand(ctx, messages)
if isCommand { if isCommand {
// execute the command
req, err := cmd.NewRequest(ctx, neo.Conversation) req, err := cmd.NewRequest(ctx, neo.Conversation)
if err != nil { if err != nil {
chanError <- err log.Error("Command with AI error: %s", err.Error())
done <- true
return return
} }
log.Trace("Command with AI: question: %s messages:%v", question, messages)
err = req.Run(messages, func(msg *message.JSON) int { err = req.Run(messages, func(msg *message.JSON) int {
chanStream <- msg err := neo.send(ctx, msg, messages, content, c)
if err != nil {
c.Status(500)
return 0 // break
}
// Complete the stream
if msg.IsDone() {
return 0 // break
}
return 1 return 1
}) })
if err != nil { if err != nil {
chanError <- err c.Status(500)
log.Error("Command with AI error: %s", err.Error())
} }
return return
} }
// chat with AI
log.Trace("Chat with AI: question:%s messages:%v", question, messages)
_, ex := neo.AI.ChatCompletionsWith(ctx, messages, neo.Option, func(data []byte) int { _, ex := neo.AI.ChatCompletionsWith(ctx, messages, neo.Option, func(data []byte) int {
chanStream <- message.NewOpenAI(data)
return 1 select {
case <-clientBreak:
return 0 // break
default:
msg := message.NewOpenAI(data)
if msg == nil {
return 1 // continue success
}
if msg.Error != "" {
neo.send(ctx, msg, messages, content, c)
return 0 // break
}
content = msg.Append(content)
err := neo.send(ctx, msg, messages, content, c)
if err != nil {
c.Status(500)
return 0 // break
}
// Complete the stream
if msg.IsDone() {
done <- true
return 0 // break
}
return 1 // continue success
}
}) })
// Throw the error
if ex != nil { if ex != nil {
chanError <- fmt.Errorf("AI chat error: %s", ex.Message) log.Error("Neo chat error: %s", ex.Message)
c.Status(200)
done <- true
return
} }
defer neo.saveHistory(ctx.Sid, content, messages) // save the history
neo.saveHistory(ctx.Sid, content, messages)
c.Status(200)
// Complete the stream
done <- true
}() }()
answer.Header("Content-Type", "text/event-stream;charset=utf-8") select {
ok := answer.Stream(func(w io.Writer) bool { case <-done:
select { return nil
case err := <-chanError: case <-c.Writer.CloseNotify():
if err != nil { clientBreak <- true
message.New().Text(err.Error()).Write(w)
}
if len(errorMsg) > 0 {
var errData openai.ErrorMessage
err := jsoniter.Unmarshal(errorMsg, &errData)
if err == nil {
msg := errData.Error.Message
if msg == "" {
msg = fmt.Sprintf("OpenAI error: %s", errData.Error.Code)
}
message.New().Text(msg).Write(w)
message.New().Done().Write(w)
return false
}
message.New().Text(string(errorMsg)).Write(w)
message.New().Done().Write(w)
return false
}
message.New().Done().Write(w)
return false
case msg := <-chanStream:
if msg == nil {
return true
}
if msg.Error != "" {
errorMsg = append(errorMsg, []byte(msg.Error)...)
return true
}
content = msg.Append(content)
err := neo.write(msg, w, ctx, messages, content)
if err != nil {
log.Warn("Neo write process msg: %v error: %s", msg, err.Error())
msg.Write(w)
}
return !msg.IsDone()
// case <-ctx.Done():
// if err := ctx.Err(); err != nil {
// message.New().Text(err.Error()).Write(w)
// }
// if len(errorMsg) > 0 {
// var errData openai.ErrorMessage
// err := jsoniter.Unmarshal(errorMsg, &errData)
// if err == nil {
// msg := errData.Error.Message
// if msg == "" {
// msg = fmt.Sprintf("OpenAI error: %s", errData.Error.Code)
// }
// message.New().Text(msg).Write(w)
// message.New().Done().Write(w)
// return false
// }
// message.New().Text(string(errorMsg)).Write(w)
// message.New().Done().Write(w)
// return false
// }
// message.New().Done().Write(w)
// return false
}
})
if !ok {
answer.Status(500)
return nil return nil
} }
answer.Status(200) }
return nil
// Send send the message to the stream
func (neo *DSL) send(ctx command.Context, msg *message.JSON, messages []map[string]interface{}, content []byte, c *gin.Context) error {
w := c.Writer
// Directly write the message
if neo.Write == "" {
ok := msg.Write(c.Writer)
if !ok {
return fmt.Errorf("Stream write error")
}
return nil
}
// Execute the custom write hook get the response
args := []interface{}{ctx, messages, msg, string(content), w}
p, err := process.Of(neo.Write, args...)
if err != nil {
msg.Write(w)
return fmt.Errorf("Stream write error: %s", err.Error())
}
err = p.WithSID(ctx.Sid).Execute()
if err != nil {
log.Error("Neo custom write error: %s", err.Error())
msg.Write(w)
return nil
}
defer p.Release()
res := p.Value()
if res == nil {
return fmt.Errorf("Neo custom write return null")
}
// Send the custom write response to the stream
if messages, ok := res.([]interface{}); ok {
for _, new := range messages {
if v, ok := new.(map[string]interface{}); ok {
newMsg := message.New().Map(v)
newMsg.Write(w)
}
}
return nil
}
return fmt.Errorf("Neo should return an array of response")
} }
// prompts get the prompts // prompts get the prompts
@ -332,47 +347,6 @@ func (neo *DSL) prompts() []map[string]interface{} {
return prompts return prompts
} }
// after the after hook
func (neo *DSL) write(msg *message.JSON, w io.Writer, ctx command.Context, messages []map[string]interface{}, content []byte) error {
if neo.Write == "" {
msg.Write(w)
return nil
}
args := []interface{}{ctx, messages, msg, string(content)}
p, err := process.Of(neo.Write, args...)
if err != nil {
log.Error("Neo custom write process error: %s", err.Error())
msg.Write(w)
return nil
}
err = p.WithSID(ctx.Sid).Execute()
if err != nil {
log.Error("Neo custom write error: %s", err.Error())
msg.Write(w)
return nil
}
defer p.Release()
res := p.Value()
if res == nil {
return fmt.Errorf("Neo custom write return null")
}
if messages, ok := res.([]interface{}); ok {
for _, new := range messages {
if v, ok := new.(map[string]interface{}); ok {
newMsg := message.New().Map(v)
newMsg.Write(w)
}
}
return nil
}
return fmt.Errorf("Neo custom write return not map")
}
// prepare the messages // prepare the messages
func (neo *DSL) prepare(ctx command.Context, messages []map[string]interface{}) []map[string]interface{} { func (neo *DSL) prepare(ctx command.Context, messages []map[string]interface{}) []map[string]interface{} {
if neo.Prepare == "" { if neo.Prepare == "" {
@ -493,20 +467,17 @@ func (neo *DSL) getCorsHandlers(router *gin.Engine, path string) ([]gin.HandlerF
allowsMap[allow] = true allowsMap[allow] = true
} }
router.OPTIONS(path, func(c *gin.Context) { c.AbortWithStatus(204) }) router.OPTIONS(path+"/history", neo.optionsHandler)
router.OPTIONS(path+"/commands", func(c *gin.Context) { c.AbortWithStatus(204) }) router.OPTIONS(path+"/commands", neo.optionsHandler)
router.OPTIONS(path+"/history", func(c *gin.Context) { c.AbortWithStatus(204) })
return []gin.HandlerFunc{ return []gin.HandlerFunc{
func(c *gin.Context) { func(c *gin.Context) {
referer := c.Request.Referer() referer := neo.getOrigin(c)
if referer != "" { if referer != "" {
if !api.IsAllowed(c, allowsMap) { if !api.IsAllowed(c, allowsMap) {
c.JSON(403, gin.H{"message": referer + " not allowed", "code": 403}) c.JSON(403, gin.H{"message": referer + " not allowed", "code": 403})
c.Abort() c.Abort()
return return
} }
url, _ := url.Parse(referer) url, _ := url.Parse(referer)
referer = fmt.Sprintf("%s://%s", url.Scheme, url.Host) referer = fmt.Sprintf("%s://%s", url.Scheme, url.Host)
c.Writer.Header().Set("Access-Control-Allow-Origin", referer) c.Writer.Header().Set("Access-Control-Allow-Origin", referer)
@ -519,6 +490,24 @@ func (neo *DSL) getCorsHandlers(router *gin.Engine, path string) ([]gin.HandlerF
}, nil }, nil
} }
func (neo *DSL) optionsHandler(c *gin.Context) {
origin := neo.getOrigin(c)
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET")
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Authorization")
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
c.AbortWithStatus(204)
}
func (neo *DSL) getOrigin(c *gin.Context) string {
referer := c.Request.Referer()
origin := c.Request.Header.Get("Origin")
if origin == "" {
origin = referer
}
return origin
}
func (neo *DSL) getGuardHandlers() ([]gin.HandlerFunc, error) { func (neo *DSL) getGuardHandlers() ([]gin.HandlerFunc, error) {
if neo.Guard == "" { if neo.Guard == "" {

41
neo/process.go Normal file
View file

@ -0,0 +1,41 @@
package neo
import (
"github.com/gin-gonic/gin"
"github.com/yaoapp/gou/process"
"github.com/yaoapp/kun/exception"
"github.com/yaoapp/yao/neo/message"
)
func init() {
process.RegisterGroup("neo", map[string]process.Handler{
"write": ProcessWrite,
})
}
// ProcessWrite process the write request
func ProcessWrite(process *process.Process) interface{} {
process.ValidateArgNums(2)
w, ok := process.Args[0].(gin.ResponseWriter)
if !ok {
exception.New("The first argument must be a io.Writer", 400).Throw()
return nil
}
data, ok := process.Args[1].([]interface{})
if !ok {
exception.New("The second argument must be a Array", 400).Throw()
return nil
}
for _, new := range data {
if v, ok := new.(map[string]interface{}); ok {
newMsg := message.New().Map(v)
newMsg.Write(w)
}
}
return nil
}