[add] Neo (70%)
This commit is contained in:
parent
eace6ea675
commit
7ee7de96c9
8 changed files with 89 additions and 26 deletions
|
|
@ -20,6 +20,7 @@ import (
|
|||
"github.com/yaoapp/yao/i18n"
|
||||
"github.com/yaoapp/yao/importer"
|
||||
"github.com/yaoapp/yao/model"
|
||||
"github.com/yaoapp/yao/neo"
|
||||
"github.com/yaoapp/yao/pack"
|
||||
"github.com/yaoapp/yao/plugin"
|
||||
"github.com/yaoapp/yao/query"
|
||||
|
|
@ -178,11 +179,18 @@ func Load(cfg config.Config) (err error) {
|
|||
printErr(cfg.Mode, "Widget", err)
|
||||
}
|
||||
|
||||
// Load AIGC
|
||||
err = aigc.Load(cfg)
|
||||
if err != nil {
|
||||
printErr(cfg.Mode, "AIGC", err)
|
||||
}
|
||||
|
||||
// Load Neo
|
||||
err = neo.Load(cfg)
|
||||
if err != nil {
|
||||
printErr(cfg.Mode, "AIGC", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
|
|
|||
13
neo/load.go
13
neo/load.go
|
|
@ -8,12 +8,13 @@ import (
|
|||
"github.com/yaoapp/yao/config"
|
||||
)
|
||||
|
||||
var neo *Neo
|
||||
// Neo the neo AI assistant
|
||||
var Neo *DSL
|
||||
|
||||
// Load load AIGC
|
||||
func Load(cfg config.Config) error {
|
||||
|
||||
setting := Neo{
|
||||
setting := DSL{
|
||||
ID: "neo",
|
||||
Prompts: []aigc.Prompt{},
|
||||
Option: map[string]interface{}{},
|
||||
|
|
@ -31,13 +32,13 @@ func Load(cfg config.Config) error {
|
|||
return err
|
||||
}
|
||||
|
||||
neo = &setting
|
||||
err = neo.newAI()
|
||||
Neo = &setting
|
||||
err = Neo.newAI()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = neo.newConversation()
|
||||
err = Neo.newConversation()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
|
@ -46,4 +47,4 @@ func Load(cfg config.Config) error {
|
|||
}
|
||||
|
||||
// LoadCommands load the commands
|
||||
func (neo *Neo) LoadCommands() {}
|
||||
func (neo *DSL) LoadCommands() {}
|
||||
|
|
|
|||
|
|
@ -20,5 +20,5 @@ func TestLoad(t *testing.T) {
|
|||
}
|
||||
|
||||
func check(t *testing.T) {
|
||||
assert.NotNil(t, neo)
|
||||
assert.NotNil(t, Neo)
|
||||
}
|
||||
|
|
|
|||
61
neo/neo.go
61
neo/neo.go
|
|
@ -6,9 +6,11 @@ import (
|
|||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"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"
|
||||
|
|
@ -18,7 +20,7 @@ import (
|
|||
)
|
||||
|
||||
// API is a method on the Neo type
|
||||
func (neo *Neo) API(router *gin.Engine, path string, allows ...string) error {
|
||||
func (neo *DSL) API(router *gin.Engine, path string) error {
|
||||
|
||||
prompts := []map[string]interface{}{}
|
||||
for _, prompt := range neo.Prompts {
|
||||
|
|
@ -36,7 +38,7 @@ func (neo *Neo) API(router *gin.Engine, path string, allows ...string) error {
|
|||
}
|
||||
|
||||
// Cross-Domain
|
||||
neo.crossDomain(router, path, allows...)
|
||||
neo.crossDomain(router, path)
|
||||
|
||||
// api router
|
||||
router.GET(path, func(c *gin.Context) {
|
||||
|
|
@ -78,7 +80,7 @@ func (neo *Neo) API(router *gin.Engine, path string, allows ...string) error {
|
|||
}
|
||||
|
||||
// Answer the message
|
||||
func (neo *Neo) Answer(ctx context.Context, c *gin.Context, messages []map[string]interface{}) error {
|
||||
func (neo *DSL) Answer(ctx context.Context, c *gin.Context, messages []map[string]interface{}) error {
|
||||
|
||||
chanStream := make(chan []byte, 1)
|
||||
chanError := make(chan error, 1)
|
||||
|
|
@ -104,13 +106,39 @@ func (neo *Neo) Answer(ctx context.Context, c *gin.Context, messages []map[strin
|
|||
select {
|
||||
case err := <-chanError:
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, err.Error())
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"message": err.Error(), "code": 500})
|
||||
}
|
||||
return false
|
||||
|
||||
case msg := <-chanStream:
|
||||
msg = append(msg, []byte("\n")...)
|
||||
w.Write(msg)
|
||||
if msg != nil && len(msg) > 0 {
|
||||
|
||||
if strings.Contains(string(msg), `"delta":{"content"`) {
|
||||
msg = []byte(strings.TrimPrefix(string(msg), "data: "))
|
||||
var message openai.Message
|
||||
err := jsoniter.Unmarshal(msg, &message)
|
||||
|
||||
if err != nil {
|
||||
data, _ := jsoniter.Marshal(map[string]interface{}{"text": err.Error()})
|
||||
w.Write([]byte(fmt.Sprintf("data: %s\n\n", data)))
|
||||
return true
|
||||
}
|
||||
|
||||
if len(message.Choices) > 0 {
|
||||
content := message.Choices[0].Delta.Content
|
||||
data, _ := jsoniter.Marshal(map[string]interface{}{"text": content})
|
||||
w.Write([]byte(fmt.Sprintf("data: %s\n\n", data)))
|
||||
return true
|
||||
}
|
||||
|
||||
} else if strings.Contains(string(msg), `[DONE]`) {
|
||||
w.Write([]byte(fmt.Sprintf("data: %s\n\n", `{"done":true}`)))
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// msg = append(msg, []byte("\n")...)
|
||||
// w.Write(msg)
|
||||
return true
|
||||
}
|
||||
})
|
||||
|
|
@ -124,14 +152,16 @@ func (neo *Neo) Answer(ctx context.Context, c *gin.Context, messages []map[strin
|
|||
return nil
|
||||
}
|
||||
|
||||
func (neo *Neo) crossDomain(router *gin.Engine, path string, allows ...string) {
|
||||
func (neo *DSL) crossDomain(router *gin.Engine, path string) {
|
||||
|
||||
if len(allows) == 0 {
|
||||
if len(neo.Allows) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
allowsMap := map[string]bool{}
|
||||
for _, allow := range allows {
|
||||
for _, allow := range neo.Allows {
|
||||
allow = strings.TrimPrefix(allow, "http://")
|
||||
allow = strings.TrimPrefix(allow, "https://")
|
||||
allowsMap[allow] = true
|
||||
}
|
||||
|
||||
|
|
@ -140,7 +170,8 @@ func (neo *Neo) crossDomain(router *gin.Engine, path string, allows ...string) {
|
|||
if referer != "" {
|
||||
|
||||
if !api.IsAllowed(c, allowsMap) {
|
||||
c.AbortWithStatus(403)
|
||||
c.JSON(403, gin.H{"message": referer + " not allowed", "code": 403})
|
||||
c.Abort()
|
||||
return
|
||||
}
|
||||
|
||||
|
|
@ -150,14 +181,14 @@ func (neo *Neo) crossDomain(router *gin.Engine, path string, allows ...string) {
|
|||
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With")
|
||||
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT")
|
||||
c.AbortWithStatus(204)
|
||||
c.Next()
|
||||
}
|
||||
})
|
||||
|
||||
router.OPTIONS(path, func(c *gin.Context) { c.Status(200) })
|
||||
router.OPTIONS(path, func(c *gin.Context) { c.AbortWithStatus(204) })
|
||||
}
|
||||
|
||||
func (neo *Neo) setGuard(router *gin.Engine) error {
|
||||
func (neo *DSL) setGuard(router *gin.Engine) error {
|
||||
|
||||
if neo.Guard == "" {
|
||||
router.Use(func(c *gin.Context) {
|
||||
|
|
@ -187,7 +218,7 @@ func (neo *Neo) setGuard(router *gin.Engine) error {
|
|||
}
|
||||
|
||||
// NewAI create a new AI
|
||||
func (neo *Neo) newAI() error {
|
||||
func (neo *DSL) newAI() error {
|
||||
|
||||
if neo.Connector == "" {
|
||||
return fmt.Errorf("%s connector is required", neo.ID)
|
||||
|
|
@ -211,7 +242,7 @@ func (neo *Neo) newAI() error {
|
|||
}
|
||||
|
||||
// newConversation create a new conversation
|
||||
func (neo *Neo) newConversation() error {
|
||||
func (neo *DSL) newConversation() error {
|
||||
|
||||
if neo.ConversationSetting.Connector == "default" || neo.ConversationSetting.Connector == "" {
|
||||
neo.Conversation = conversation.NewXun()
|
||||
|
|
|
|||
|
|
@ -25,7 +25,7 @@ func TestAPI(t *testing.T) {
|
|||
|
||||
// test router
|
||||
router := testRouter(t)
|
||||
err := neo.API(router, "/neo/chat")
|
||||
err := Neo.API(router, "/neo/chat")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
|
@ -57,7 +57,7 @@ func TestAPIAuth(t *testing.T) {
|
|||
defer test.Clean()
|
||||
|
||||
router := testRouter(t)
|
||||
err := neo.API(router, "/neo/chat")
|
||||
err := Neo.API(router, "/neo/chat")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ package neo
|
|||
|
||||
import "github.com/yaoapp/yao/aigc"
|
||||
|
||||
// Neo AI assistant
|
||||
type Neo struct {
|
||||
// DSL AI assistant
|
||||
type DSL struct {
|
||||
ID string `json:"-" yaml:"-"`
|
||||
Name string `json:"name,omitempty"`
|
||||
Guard string `json:"guard,omitempty"`
|
||||
|
|
|
|||
17
openai/types.go
Normal file
17
openai/types.go
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
package openai
|
||||
|
||||
// Message is the response from OpenAI
|
||||
// {"id":"chatcmpl-7Atx502nGBuYcvoZfIaWU4FREI1mT","object":"chat.completion.chunk","created":1682832715,"model":"gpt-3.5-turbo-0301","choices":[{"delta":{"content":"Hello"},"index":0,"finish_reason":null}]}
|
||||
type Message 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 {
|
||||
Content string `json:"content,omitempty"`
|
||||
} `json:"delta,omitempty"`
|
||||
Index int `json:"index,omitempty"`
|
||||
FinishReason string `json:"finish_reason,omitempty"`
|
||||
} `json:"choices,omitempty"`
|
||||
}
|
||||
|
|
@ -7,6 +7,7 @@ import (
|
|||
"github.com/yaoapp/gou/api"
|
||||
"github.com/yaoapp/gou/server/http"
|
||||
"github.com/yaoapp/yao/config"
|
||||
"github.com/yaoapp/yao/neo"
|
||||
"github.com/yaoapp/yao/share"
|
||||
)
|
||||
|
||||
|
|
@ -35,6 +36,11 @@ func Start(cfg config.Config) (*http.Server, error) {
|
|||
Timeout: 5 * time.Second,
|
||||
}).With(Middlewares...)
|
||||
|
||||
// Neo API
|
||||
if neo.Neo != nil {
|
||||
neo.Neo.API(router, "/api/__yao/neo")
|
||||
}
|
||||
|
||||
go func() {
|
||||
err = srv.Start()
|
||||
}()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue