commit
eace6ea675
7 changed files with 178 additions and 19 deletions
|
|
@ -30,8 +30,8 @@ func (ai *DSL) Call(content string, user string, option map[string]interface{})
|
||||||
messages := []map[string]interface{}{}
|
messages := []map[string]interface{}{}
|
||||||
for _, prompt := range ai.Prompts {
|
for _, prompt := range ai.Prompts {
|
||||||
message := map[string]interface{}{"role": prompt.Role, "content": prompt.Content}
|
message := map[string]interface{}{"role": prompt.Role, "content": prompt.Content}
|
||||||
if prompt.User != "" {
|
if prompt.Name != "" {
|
||||||
message["user"] = prompt.User
|
message["name"] = prompt.Name
|
||||||
}
|
}
|
||||||
messages = append(messages, message)
|
messages = append(messages, message)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -8,20 +8,20 @@ import (
|
||||||
|
|
||||||
// DSL the connector DSL
|
// DSL the connector DSL
|
||||||
type DSL struct {
|
type DSL struct {
|
||||||
ID string `json:"-"`
|
ID string `json:"-" yaml:"-"`
|
||||||
Name string `json:"name,omitempty"`
|
Name string `json:"name,omitempty"`
|
||||||
Connector string `json:"connector"`
|
Connector string `json:"connector"`
|
||||||
Process string `json:"process,omitempty"`
|
Process string `json:"process,omitempty"`
|
||||||
Prompts []Prompt `json:"prompts"`
|
Prompts []Prompt `json:"prompts"`
|
||||||
Optional Optional `json:"optional,omitempty"`
|
Optional Optional `json:"optional,omitempty"`
|
||||||
AI AI `json:"-"`
|
AI AI `json:"-" yaml:"-"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Prompt a prompt
|
// Prompt a prompt
|
||||||
type Prompt struct {
|
type Prompt struct {
|
||||||
Role string `json:"role"`
|
Role string `json:"role"`
|
||||||
Content string `json:"content"`
|
Content string `json:"content"`
|
||||||
User string `json:"user,omitempty"`
|
Name string `json:"name,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// Optional optional
|
// Optional optional
|
||||||
|
|
|
||||||
|
|
@ -26,12 +26,12 @@ func Load(cfg config.Config) error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
err = application.Parse("neo.yml", bytes, &neo)
|
err = application.Parse("neo.yml", bytes, &setting)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
*neo = setting
|
neo = &setting
|
||||||
err = neo.newAI()
|
err = neo.newAI()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
|
|
|
||||||
24
neo/load_test.go
Normal file
24
neo/load_test.go
Normal file
|
|
@ -0,0 +1,24 @@
|
||||||
|
package neo
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestLoad(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
err := Load(config.Conf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
check(t)
|
||||||
|
}
|
||||||
|
|
||||||
|
func check(t *testing.T) {
|
||||||
|
assert.NotNil(t, neo)
|
||||||
|
}
|
||||||
21
neo/neo.go
21
neo/neo.go
|
|
@ -8,6 +8,7 @@ import (
|
||||||
"net/url"
|
"net/url"
|
||||||
|
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/google/uuid"
|
||||||
"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"
|
||||||
|
|
@ -21,7 +22,11 @@ func (neo *Neo) API(router *gin.Engine, path string, allows ...string) error {
|
||||||
|
|
||||||
prompts := []map[string]interface{}{}
|
prompts := []map[string]interface{}{}
|
||||||
for _, prompt := range neo.Prompts {
|
for _, prompt := range neo.Prompts {
|
||||||
prompts = append(prompts, map[string]interface{}{"role": prompt.Role, "content": prompt.Content, "user": prompt.User})
|
message := map[string]interface{}{"role": prompt.Role, "content": prompt.Content}
|
||||||
|
if prompt.Name != "" {
|
||||||
|
message["name"] = prompt.Name
|
||||||
|
}
|
||||||
|
prompts = append(prompts, message)
|
||||||
}
|
}
|
||||||
|
|
||||||
// set the guard
|
// set the guard
|
||||||
|
|
@ -37,7 +42,11 @@ func (neo *Neo) API(router *gin.Engine, path string, allows ...string) error {
|
||||||
router.GET(path, func(c *gin.Context) {
|
router.GET(path, func(c *gin.Context) {
|
||||||
|
|
||||||
sid := c.GetString("__sid")
|
sid := c.GetString("__sid")
|
||||||
content := c.GetString("content")
|
if sid == "" {
|
||||||
|
sid = uuid.New().String()
|
||||||
|
}
|
||||||
|
|
||||||
|
content := c.Query("content")
|
||||||
if content == "" {
|
if content == "" {
|
||||||
c.JSON(400, gin.H{"message": "content is required", "code": 400})
|
c.JSON(400, gin.H{"message": "content is required", "code": 400})
|
||||||
return
|
return
|
||||||
|
|
@ -51,7 +60,7 @@ func (neo *Neo) API(router *gin.Engine, path string, allows ...string) error {
|
||||||
}
|
}
|
||||||
|
|
||||||
messages = append(messages, history...)
|
messages = append(messages, history...)
|
||||||
messages = append(messages, map[string]interface{}{"role": "user", "content": content, "user": sid})
|
messages = append(messages, map[string]interface{}{"role": "user", "content": content, "name": sid})
|
||||||
|
|
||||||
// reply the content
|
// reply the content
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
@ -80,7 +89,7 @@ func (neo *Neo) Answer(ctx context.Context, c *gin.Context, messages []map[strin
|
||||||
close(chanError)
|
close(chanError)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
_, ex := neo.AI.ChatCompletions(messages, neo.Option, func(data []byte) int {
|
_, ex := neo.AI.ChatCompletionsWith(ctx, messages, neo.Option, func(data []byte) int {
|
||||||
chanStream <- data
|
chanStream <- data
|
||||||
return 1
|
return 1
|
||||||
})
|
})
|
||||||
|
|
@ -103,9 +112,6 @@ func (neo *Neo) Answer(ctx context.Context, c *gin.Context, messages []map[strin
|
||||||
msg = append(msg, []byte("\n")...)
|
msg = append(msg, []byte("\n")...)
|
||||||
w.Write(msg)
|
w.Write(msg)
|
||||||
return true
|
return true
|
||||||
|
|
||||||
case <-ctx.Done():
|
|
||||||
return false
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|
@ -198,6 +204,7 @@ func (neo *Neo) newAI() error {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
neo.AI = ai
|
neo.AI = ai
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return fmt.Errorf("%s connector %s not support, should be a openai", neo.ID, neo.Connector)
|
return fmt.Errorf("%s connector %s not support, should be a openai", neo.ID, neo.Connector)
|
||||||
|
|
|
||||||
128
neo/neo_test.go
Normal file
128
neo/neo_test.go
Normal file
|
|
@ -0,0 +1,128 @@
|
||||||
|
package neo
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
httpTest "github.com/yaoapp/gou/http"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/helper"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
_ "github.com/yaoapp/yao/utils"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAPI(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
// test router
|
||||||
|
router := testRouter(t)
|
||||||
|
err := neo.API(router, "/neo/chat")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// test server
|
||||||
|
host, shutdown := testServer(t, router)
|
||||||
|
defer shutdown()
|
||||||
|
|
||||||
|
// test request
|
||||||
|
url := fmt.Sprintf("%s/neo/chat?content=hello&token=%s", host, testToken(t))
|
||||||
|
res := []byte{}
|
||||||
|
req := httpTest.New(url).
|
||||||
|
WithHeader(http.Header{"Content-Type": []string{"application/json"}})
|
||||||
|
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
// send request
|
||||||
|
req.Stream(ctx, "GET", nil, func(data []byte) int {
|
||||||
|
res = append(res, data...)
|
||||||
|
return 1
|
||||||
|
})
|
||||||
|
|
||||||
|
assert.Contains(t, string(res), "[DONE]")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAPIAuth(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
router := testRouter(t)
|
||||||
|
err := neo.API(router, "/neo/chat")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
response := httptest.NewRecorder()
|
||||||
|
req, _ := http.NewRequest("GET", "/neo/chat?content=hello", nil)
|
||||||
|
assert.Panics(t, func() {
|
||||||
|
router.ServeHTTP(response, req)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func testServer(t *testing.T, router *gin.Engine) (string, func()) {
|
||||||
|
|
||||||
|
// Listen
|
||||||
|
l, err := net.Listen("tcp4", ":0")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
srv := &http.Server{Addr: ":0", Handler: router}
|
||||||
|
|
||||||
|
// start serve
|
||||||
|
go func() {
|
||||||
|
if err := srv.Serve(l); err != nil && err != http.ErrServerClosed {
|
||||||
|
fmt.Println("[TestServer] Error:", err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
|
addr := strings.Split(l.Addr().String(), ":")
|
||||||
|
if len(addr) != 2 {
|
||||||
|
t.Fatal("invalid address")
|
||||||
|
}
|
||||||
|
|
||||||
|
host := fmt.Sprintf("http://127.0.0.1:%s", addr[1])
|
||||||
|
time.Sleep(50 * time.Millisecond)
|
||||||
|
|
||||||
|
shutdown := func() {
|
||||||
|
srv.Close()
|
||||||
|
l.Close()
|
||||||
|
}
|
||||||
|
return host, shutdown
|
||||||
|
}
|
||||||
|
|
||||||
|
func testRouter(t *testing.T) *gin.Engine {
|
||||||
|
|
||||||
|
err := Load(config.Conf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
router := gin.New()
|
||||||
|
gin.SetMode(gin.ReleaseMode)
|
||||||
|
return router
|
||||||
|
}
|
||||||
|
|
||||||
|
func testToken(t *testing.T) string {
|
||||||
|
token := helper.JwtMake(1,
|
||||||
|
map[string]interface{}{
|
||||||
|
"id": 1,
|
||||||
|
"name": "Test",
|
||||||
|
},
|
||||||
|
map[string]interface{}{
|
||||||
|
"exp": 3600,
|
||||||
|
"sid": "123456",
|
||||||
|
})
|
||||||
|
return token.Token
|
||||||
|
}
|
||||||
10
neo/types.go
10
neo/types.go
|
|
@ -4,17 +4,17 @@ import "github.com/yaoapp/yao/aigc"
|
||||||
|
|
||||||
// Neo AI assistant
|
// Neo AI assistant
|
||||||
type Neo struct {
|
type Neo struct {
|
||||||
ID string `json:"-"`
|
ID string `json:"-" yaml:"-"`
|
||||||
Name string `json:"name,omitempty"`
|
Name string `json:"name,omitempty"`
|
||||||
Guard string `json:"guard,omitempty"`
|
Guard string `json:"guard,omitempty"`
|
||||||
Connector string `json:"connector"`
|
Connector string `json:"connector"`
|
||||||
ConversationSetting ConversationSetting `json:"conversation"`
|
ConversationSetting ConversationSetting `json:"conversation"`
|
||||||
Option map[string]interface{} `json:"option"`
|
Option map[string]interface{} `json:"option"`
|
||||||
Prompts []aigc.Prompt `json:"prompts"`
|
Prompts []aigc.Prompt `json:"prompts,omitempty"`
|
||||||
Allows []string `json:"allows,omitempty"`
|
Allows []string `json:"allows,omitempty"`
|
||||||
AI aigc.AI `json:"-"`
|
AI aigc.AI `json:"-" yaml:"-"`
|
||||||
Conversation Conversation `json:"-"`
|
Conversation Conversation `json:"-" yaml:"-"`
|
||||||
Command Command `json:"-"`
|
Command Command `json:"-" yaml:"-"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// ConversationSetting the conversation config
|
// ConversationSetting the conversation config
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue