[add] aigc widget done
This commit is contained in:
parent
7e480633d4
commit
894c574db4
10 changed files with 415 additions and 6 deletions
113
aigc/aigc.go
Normal file
113
aigc/aigc.go
Normal file
|
|
@ -0,0 +1,113 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
jsoniter "github.com/json-iterator/go"
|
||||||
|
"github.com/yaoapp/gou/connector"
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/kun/exception"
|
||||||
|
"github.com/yaoapp/yao/openai"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Autopilots the loaded autopilots
|
||||||
|
var Autopilots = []string{}
|
||||||
|
|
||||||
|
// AIGCs the loaded AIGCs
|
||||||
|
var AIGCs = map[string]*DSL{}
|
||||||
|
|
||||||
|
// Select select the AIGC
|
||||||
|
func Select(id string) (*DSL, error) {
|
||||||
|
if AIGCs[id] == nil {
|
||||||
|
return nil, fmt.Errorf("aigc %s not found", id)
|
||||||
|
}
|
||||||
|
return AIGCs[id], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Call the AIGC
|
||||||
|
func (ai *DSL) Call(content string, user string, option map[string]interface{}) (interface{}, *exception.Exception) {
|
||||||
|
|
||||||
|
messages := []map[string]interface{}{}
|
||||||
|
for _, prompt := range ai.Prompts {
|
||||||
|
message := map[string]interface{}{"role": prompt.Role, "content": prompt.Content}
|
||||||
|
if prompt.User != "" {
|
||||||
|
message["user"] = prompt.User
|
||||||
|
}
|
||||||
|
messages = append(messages, message)
|
||||||
|
}
|
||||||
|
|
||||||
|
// add the user message
|
||||||
|
message := map[string]interface{}{"role": "user", "content": content}
|
||||||
|
if user != "" {
|
||||||
|
message["user"] = user
|
||||||
|
}
|
||||||
|
messages = append(messages, message)
|
||||||
|
|
||||||
|
bytes, err := jsoniter.Marshal(messages)
|
||||||
|
if err != nil {
|
||||||
|
return nil, exception.New(err.Error(), 400)
|
||||||
|
}
|
||||||
|
|
||||||
|
token, err := ai.AI.Tiktoken(string(bytes))
|
||||||
|
if err != nil {
|
||||||
|
return nil, exception.New(err.Error(), 400)
|
||||||
|
}
|
||||||
|
|
||||||
|
if token > ai.AI.MaxToken() {
|
||||||
|
return nil, exception.New("token limit exceeded", 400)
|
||||||
|
}
|
||||||
|
|
||||||
|
// call the AI
|
||||||
|
res, ex := ai.AI.ChatCompletions(messages, option, nil)
|
||||||
|
if ex != nil {
|
||||||
|
return nil, ex
|
||||||
|
}
|
||||||
|
|
||||||
|
resText, ex := ai.AI.GetContent(res)
|
||||||
|
if ex != nil {
|
||||||
|
return nil, ex
|
||||||
|
}
|
||||||
|
|
||||||
|
if ai.Process == "" {
|
||||||
|
return resText, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var param interface{} = resText
|
||||||
|
if ai.Optional.JSON {
|
||||||
|
err = jsoniter.Unmarshal([]byte(resText), ¶m)
|
||||||
|
if err != nil {
|
||||||
|
return nil, exception.New("%s parse error: %s", 400, resText, err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
p, err := process.Of(ai.Process, param)
|
||||||
|
if err != nil {
|
||||||
|
return nil, exception.New(err.Error(), 400)
|
||||||
|
}
|
||||||
|
|
||||||
|
resProcess, err := p.Exec()
|
||||||
|
if err != nil {
|
||||||
|
return nil, exception.New(err.Error(), 500)
|
||||||
|
}
|
||||||
|
|
||||||
|
return resProcess, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewAI create a new AI
|
||||||
|
func (ai *DSL) newAI() (AI, error) {
|
||||||
|
|
||||||
|
if ai.Connector == "" {
|
||||||
|
return nil, fmt.Errorf("%s connector is required", ai.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
conn, err := connector.Select(ai.Connector)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if conn.Is(connector.OPENAI) {
|
||||||
|
return openai.New(ai.Connector)
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil, fmt.Errorf("%s connector %s not support, should be a openai", ai.ID, ai.Connector)
|
||||||
|
}
|
||||||
57
aigc/aigc_test.go
Normal file
57
aigc/aigc_test.go
Normal file
|
|
@ -0,0 +1,57 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestCall(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
prepare(t)
|
||||||
|
|
||||||
|
aigc, err := Select("translate")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
content, ex := aigc.Call("你好哇", "", nil)
|
||||||
|
if ex != nil {
|
||||||
|
t.Fatal(ex.Message)
|
||||||
|
}
|
||||||
|
assert.Equal(t, "Hello", content)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCallWithProcess(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
prepare(t)
|
||||||
|
|
||||||
|
aigc, err := Select("draw")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
args, ex := aigc.Call("帮我画一只小白兔,要有白色的耳朵. 画布高度 256,宽度 256", "", nil)
|
||||||
|
if ex != nil {
|
||||||
|
t.Fatal(ex.Message)
|
||||||
|
}
|
||||||
|
|
||||||
|
data, ok := args.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("args is not map[string]interface{}")
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Equal(t, float64(256), data["height"])
|
||||||
|
assert.Equal(t, float64(256), data["width"])
|
||||||
|
}
|
||||||
|
|
||||||
|
func prepare(t *testing.T) {
|
||||||
|
err := Load(config.Conf)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
69
aigc/load.go
Normal file
69
aigc/load.go
Normal file
|
|
@ -0,0 +1,69 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/application"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/share"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Load load AIGC
|
||||||
|
func Load(cfg config.Config) error {
|
||||||
|
exts := []string{"*.ai.yml", "*.ai.yaml"}
|
||||||
|
return application.App.Walk("aigcs", func(root, file string, isdir bool) error {
|
||||||
|
if isdir {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
id := share.ID(root, file)
|
||||||
|
_, err := LoadFile(file, id)
|
||||||
|
return err
|
||||||
|
}, exts...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoadFile load AIGC by file
|
||||||
|
func LoadFile(file string, id string) (*DSL, error) {
|
||||||
|
|
||||||
|
data, err := application.App.Read(file)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return LoadSource(data, file, id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LoadSource load AIGC
|
||||||
|
func LoadSource(data []byte, file, id string) (*DSL, error) {
|
||||||
|
|
||||||
|
dsl := DSL{
|
||||||
|
ID: id,
|
||||||
|
Optional: Optional{
|
||||||
|
Autopilot: false,
|
||||||
|
JSON: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
err := application.Parse(file, data, &dsl)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
if dsl.Prompts == nil || len(dsl.Prompts) == 0 {
|
||||||
|
return nil, fmt.Errorf("%s prompts is required", id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// create AI interface
|
||||||
|
dsl.AI, err = dsl.newAI()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
// add to autopilots
|
||||||
|
if dsl.Optional.Autopilot {
|
||||||
|
Autopilots = append(Autopilots, id)
|
||||||
|
}
|
||||||
|
|
||||||
|
// add to AIGCs
|
||||||
|
AIGCs[id] = &dsl
|
||||||
|
return AIGCs[id], nil
|
||||||
|
}
|
||||||
28
aigc/load_test.go
Normal file
28
aigc/load_test.go
Normal file
|
|
@ -0,0 +1,28 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
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()
|
||||||
|
|
||||||
|
Load(config.Conf)
|
||||||
|
check(t)
|
||||||
|
}
|
||||||
|
|
||||||
|
func check(t *testing.T) {
|
||||||
|
ids := map[string]bool{}
|
||||||
|
for id := range AIGCs {
|
||||||
|
ids[id] = true
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.True(t, ids["translate"])
|
||||||
|
assert.True(t, ids["draw"])
|
||||||
|
assert.Equal(t, 2, len(Autopilots))
|
||||||
|
}
|
||||||
40
aigc/process.go
Normal file
40
aigc/process.go
Normal file
|
|
@ -0,0 +1,40 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/kun/exception"
|
||||||
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
process.Register("aigcs", processAigcs)
|
||||||
|
}
|
||||||
|
|
||||||
|
// processScripts
|
||||||
|
func processAigcs(process *process.Process) interface{} {
|
||||||
|
|
||||||
|
process.ValidateArgNums(1)
|
||||||
|
aigc, err := Select(process.ID)
|
||||||
|
if err != nil {
|
||||||
|
exception.New("scripts.%s not loaded", 404, process.ID).Throw()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
content := process.ArgsString(0)
|
||||||
|
user := ""
|
||||||
|
|
||||||
|
var option map[string]interface{} = nil
|
||||||
|
if process.NumOfArgs() > 1 {
|
||||||
|
user = process.ArgsString(1)
|
||||||
|
}
|
||||||
|
|
||||||
|
if process.NumOfArgs() > 2 {
|
||||||
|
option = process.ArgsMap(2)
|
||||||
|
}
|
||||||
|
|
||||||
|
res, ex := aigc.Call(content, user, option)
|
||||||
|
if ex != nil {
|
||||||
|
ex.Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
return res
|
||||||
|
}
|
||||||
20
aigc/process_test.go
Normal file
20
aigc/process_test.go
Normal file
|
|
@ -0,0 +1,20 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestProcessAigcs(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
prepare(t)
|
||||||
|
|
||||||
|
args := []interface{}{"你好"}
|
||||||
|
res := process.New("aigcs.translate", args...).Run()
|
||||||
|
assert.Contains(t, res, "Hello")
|
||||||
|
}
|
||||||
36
aigc/types.go
Normal file
36
aigc/types.go
Normal file
|
|
@ -0,0 +1,36 @@
|
||||||
|
package aigc
|
||||||
|
|
||||||
|
import "github.com/yaoapp/kun/exception"
|
||||||
|
|
||||||
|
// DSL the connector DSL
|
||||||
|
type DSL struct {
|
||||||
|
ID string `json:"-"`
|
||||||
|
Name string `json:"name,omitempty"`
|
||||||
|
Connector string `json:"connector"`
|
||||||
|
Process string `json:"process,omitempty"`
|
||||||
|
Prompts []Prompt `json:"prompts"`
|
||||||
|
Optional Optional `json:"optional,omitempty"`
|
||||||
|
AI AI `json:"-"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Prompt a prompt
|
||||||
|
type Prompt struct {
|
||||||
|
Role string `json:"role"`
|
||||||
|
Content string `json:"content"`
|
||||||
|
User string `json:"user,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// Optional optional
|
||||||
|
type Optional struct {
|
||||||
|
Autopilot bool `json:"autopilot,omitempty"`
|
||||||
|
JSON bool `json:"json,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// AI the AI interface
|
||||||
|
type AI interface {
|
||||||
|
ChatCompletions(messages []map[string]interface{}, option map[string]interface{}, cb func(data []byte) int) (interface{}, *exception.Exception)
|
||||||
|
GetContent(response interface{}) (string, *exception.Exception)
|
||||||
|
Embeddings(input interface{}, user string) (interface{}, *exception.Exception)
|
||||||
|
Tiktoken(input string) (int, error)
|
||||||
|
MaxToken() int
|
||||||
|
}
|
||||||
1
main.go
1
main.go
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"github.com/yaoapp/yao/cmd"
|
"github.com/yaoapp/yao/cmd"
|
||||||
|
|
||||||
_ "github.com/yaoapp/gou/encoding"
|
_ "github.com/yaoapp/gou/encoding"
|
||||||
|
_ "github.com/yaoapp/yao/aigc"
|
||||||
_ "github.com/yaoapp/yao/crypto"
|
_ "github.com/yaoapp/yao/crypto"
|
||||||
_ "github.com/yaoapp/yao/helper"
|
_ "github.com/yaoapp/yao/helper"
|
||||||
_ "github.com/yaoapp/yao/openai"
|
_ "github.com/yaoapp/yao/openai"
|
||||||
|
|
|
||||||
|
|
@ -22,9 +22,10 @@ func Tiktoken(model string, input string) (int, error) {
|
||||||
|
|
||||||
// OpenAI struct
|
// OpenAI struct
|
||||||
type OpenAI struct {
|
type OpenAI struct {
|
||||||
key string
|
key string
|
||||||
model string
|
model string
|
||||||
host string
|
host string
|
||||||
|
maxToken int
|
||||||
}
|
}
|
||||||
|
|
||||||
// New create a new OpenAI instance by connector id
|
// New create a new OpenAI instance by connector id
|
||||||
|
|
@ -40,9 +41,10 @@ func New(id string) (*OpenAI, error) {
|
||||||
|
|
||||||
setting := c.Setting()
|
setting := c.Setting()
|
||||||
return &OpenAI{
|
return &OpenAI{
|
||||||
key: setting["key"].(string),
|
key: setting["key"].(string),
|
||||||
model: setting["model"].(string),
|
model: setting["model"].(string),
|
||||||
host: setting["host"].(string),
|
host: setting["host"].(string),
|
||||||
|
maxToken: 2048,
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -191,6 +193,36 @@ func (openai OpenAI) Tiktoken(input string) (int, error) {
|
||||||
return len(token), nil
|
return len(token), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// MaxToken get max number of tokens
|
||||||
|
func (openai OpenAI) MaxToken() int {
|
||||||
|
return openai.maxToken
|
||||||
|
}
|
||||||
|
|
||||||
|
// GetContent get the content of chat completions
|
||||||
|
func (openai OpenAI) GetContent(response interface{}) (string, *exception.Exception) {
|
||||||
|
if response == nil {
|
||||||
|
return "", exception.New("response is nil", 500)
|
||||||
|
}
|
||||||
|
|
||||||
|
if data, ok := response.(map[string]interface{}); ok {
|
||||||
|
if choices, ok := data["choices"].([]interface{}); ok {
|
||||||
|
if len(choices) == 0 {
|
||||||
|
return "", exception.New("choices is null, %v", 500, response)
|
||||||
|
}
|
||||||
|
|
||||||
|
if choice, ok := choices[0].(map[string]interface{}); ok {
|
||||||
|
if message, ok := choice["message"].(map[string]interface{}); ok {
|
||||||
|
if content, ok := message["content"].(string); ok {
|
||||||
|
return content, nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return "", exception.New("response format error, %#v", 500, response)
|
||||||
|
}
|
||||||
|
|
||||||
// post post request
|
// post post request
|
||||||
func (openai OpenAI) post(path string, payload map[string]interface{}) (interface{}, *exception.Exception) {
|
func (openai OpenAI) post(path string, payload map[string]interface{}) (interface{}, *exception.Exception) {
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ import (
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
"github.com/yaoapp/gou/api"
|
"github.com/yaoapp/gou/api"
|
||||||
"github.com/yaoapp/gou/application"
|
"github.com/yaoapp/gou/application"
|
||||||
|
"github.com/yaoapp/gou/connector"
|
||||||
"github.com/yaoapp/gou/model"
|
"github.com/yaoapp/gou/model"
|
||||||
"github.com/yaoapp/gou/query"
|
"github.com/yaoapp/gou/query"
|
||||||
"github.com/yaoapp/gou/query/gou"
|
"github.com/yaoapp/gou/query/gou"
|
||||||
|
|
@ -162,6 +163,7 @@ func load(t *testing.T, cfg config.Config) {
|
||||||
loadFS(t, cfg)
|
loadFS(t, cfg)
|
||||||
loadScript(t, cfg)
|
loadScript(t, cfg)
|
||||||
loadModel(t, cfg)
|
loadModel(t, cfg)
|
||||||
|
loadConnector(t, cfg)
|
||||||
loadQuery(t, cfg)
|
loadQuery(t, cfg)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -172,6 +174,17 @@ func loadFS(t *testing.T, cfg config.Config) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func loadConnector(t *testing.T, cfg config.Config) {
|
||||||
|
exts := []string{"*.yao", "*.json", "*.jsonc"}
|
||||||
|
application.App.Walk("connectors", func(root, file string, isdir bool) error {
|
||||||
|
if isdir {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
_, err := connector.Load(file, share.ID(root, file))
|
||||||
|
return err
|
||||||
|
}, exts...)
|
||||||
|
}
|
||||||
|
|
||||||
func loadScript(t *testing.T, cfg config.Config) {
|
func loadScript(t *testing.T, cfg config.Config) {
|
||||||
exts := []string{"*.js"}
|
exts := []string{"*.js"}
|
||||||
err := application.App.Walk("scripts", func(root, file string, isdir bool) error {
|
err := application.App.Walk("scripts", func(root, file string, isdir bool) error {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue