diff --git a/neo/message/json.go b/neo/message/json.go index 0828f8b8..4ada8f6c 100644 --- a/neo/message/json.go +++ b/neo/message/json.go @@ -48,7 +48,7 @@ func NewOpenAI(data []byte) *JSON { break default: - return nil + msg.Error = string(data) } return &JSON{msg} diff --git a/neo/message/types.go b/neo/message/types.go index a06ecb47..45d3b26a 100644 --- a/neo/message/types.go +++ b/neo/message/types.go @@ -3,6 +3,7 @@ package message // Message the message type Message struct { Text string `json:"text,omitempty"` + Error string `json:"error,omitempty"` Done bool `json:"done,omitempty"` Confirm bool `json:"confirm,omitempty"` Command *Command `json:"command,omitempty"` diff --git a/neo/neo.go b/neo/neo.go index 87faeec8..4ef96edd 100644 --- a/neo/neo.go +++ b/neo/neo.go @@ -9,6 +9,7 @@ import ( "github.com/fatih/color" "github.com/gin-gonic/gin" "github.com/google/uuid" + jsoniter "github.com/json-iterator/go" "github.com/yaoapp/gou/api" "github.com/yaoapp/gou/connector" "github.com/yaoapp/gou/process" @@ -142,6 +143,7 @@ func (neo *DSL) Answer(ctx command.Context, question string, answer Answer) erro chanStream := make(chan *message.JSON, 1) chanError := make(chan error, 1) content := []byte{} + errorMsg := []byte{} // get the chat messages messages, err := neo.chatMessages(ctx, question) @@ -200,6 +202,25 @@ func (neo *DSL) Answer(ctx command.Context, question string, answer Answer) erro 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 @@ -207,6 +228,12 @@ func (neo *DSL) Answer(ctx command.Context, question string, answer Answer) erro if msg == nil { return true } + + if msg.Error != "" { + errorMsg = append(errorMsg, []byte(msg.Error)...) + return true + } + msg.Write(w) content = msg.Append(content) return !msg.IsDone() @@ -215,6 +242,26 @@ func (neo *DSL) Answer(ctx command.Context, question string, answer Answer) erro 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 } diff --git a/openai/types.go b/openai/types.go index 5c2bbb83..12e2bdbd 100644 --- a/openai/types.go +++ b/openai/types.go @@ -15,3 +15,16 @@ type Message struct { FinishReason string `json:"finish_reason,omitempty"` } `json:"choices,omitempty"` } + +// ErrorMessage is the error response from OpenAI +type ErrorMessage struct { + Error Error `json:"error,omitempty"` +} + +// Error is the error response from OpenAI +type Error struct { + Message string `json:"message,omitempty"` + Type string `json:"type,omitempty"` + Param interface{} `json:"param,omitempty"` + Code string `json:"code,omitempty"` +}