Merge pull request #779 from trheyi/main
Refactor Neo message handling to use gin.ResponseWriter and improve e…
This commit is contained in:
commit
13e94aebed
3 changed files with 200 additions and 160 deletions
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
303
neo/neo.go
303
neo/neo.go
|
|
@ -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
41
neo/process.go
Normal 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
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue