[add] aigc widget done

This commit is contained in:
Max 2023-04-27 22:31:51 +08:00
parent 7e480633d4
commit 894c574db4
10 changed files with 415 additions and 6 deletions

113
aigc/aigc.go Normal file
View 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), &param)
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
View 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
View 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
View 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
View 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
View 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
View 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
}

View file

@ -4,6 +4,7 @@ import (
"github.com/yaoapp/yao/cmd"
_ "github.com/yaoapp/gou/encoding"
_ "github.com/yaoapp/yao/aigc"
_ "github.com/yaoapp/yao/crypto"
_ "github.com/yaoapp/yao/helper"
_ "github.com/yaoapp/yao/openai"

View file

@ -25,6 +25,7 @@ type OpenAI struct {
key string
model string
host string
maxToken int
}
// New create a new OpenAI instance by connector id
@ -43,6 +44,7 @@ func New(id string) (*OpenAI, error) {
key: setting["key"].(string),
model: setting["model"].(string),
host: setting["host"].(string),
maxToken: 2048,
}, nil
}
@ -191,6 +193,36 @@ func (openai OpenAI) Tiktoken(input string) (int, error) {
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
func (openai OpenAI) post(path string, payload map[string]interface{}) (interface{}, *exception.Exception) {

View file

@ -11,6 +11,7 @@ import (
"github.com/gin-gonic/gin"
"github.com/yaoapp/gou/api"
"github.com/yaoapp/gou/application"
"github.com/yaoapp/gou/connector"
"github.com/yaoapp/gou/model"
"github.com/yaoapp/gou/query"
"github.com/yaoapp/gou/query/gou"
@ -162,6 +163,7 @@ func load(t *testing.T, cfg config.Config) {
loadFS(t, cfg)
loadScript(t, cfg)
loadModel(t, cfg)
loadConnector(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) {
exts := []string{"*.js"}
err := application.App.Walk("scripts", func(root, file string, isdir bool) error {