[add] openai.* process
This commit is contained in:
parent
a0d5ca35b3
commit
538c6c063a
2 changed files with 160 additions and 3 deletions
|
|
@ -1,12 +1,20 @@
|
||||||
package openai
|
package openai
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/yaoapp/gou/http"
|
||||||
"github.com/yaoapp/gou/process"
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/gou/runtime/v8/bridge"
|
||||||
"github.com/yaoapp/kun/exception"
|
"github.com/yaoapp/kun/exception"
|
||||||
|
"github.com/yaoapp/kun/log"
|
||||||
)
|
)
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
process.Register("yao.openai.tiktoken", ProcessTiktoken)
|
process.RegisterGroup("openai", map[string]process.Handler{
|
||||||
|
"tiktoken": ProcessTiktoken,
|
||||||
|
"chat.completions": ProcessChatCompletions,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// ProcessTiktoken get number of tokens
|
// ProcessTiktoken get number of tokens
|
||||||
|
|
@ -20,3 +28,88 @@ func ProcessTiktoken(process *process.Process) interface{} {
|
||||||
}
|
}
|
||||||
return nums
|
return nums
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ProcessChatCompletions get number of tokens
|
||||||
|
func ProcessChatCompletions(process *process.Process) interface{} {
|
||||||
|
|
||||||
|
process.ValidateArgNums(2)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
model := process.ArgsString(0)
|
||||||
|
messages := []map[string]interface{}{}
|
||||||
|
intput := process.ArgsArray(1)
|
||||||
|
for idx, v := range intput {
|
||||||
|
message, ok := v.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
exception.New("ChatCompletions input must be array of map, index %d", 400, idx).Throw()
|
||||||
|
}
|
||||||
|
messages = append(messages, message)
|
||||||
|
}
|
||||||
|
|
||||||
|
ai, err := New(model)
|
||||||
|
if err != nil {
|
||||||
|
exception.New("ChatCompletions error: %s", 400, err).Throw()
|
||||||
|
}
|
||||||
|
|
||||||
|
options := map[string]interface{}{}
|
||||||
|
if process.NumOfArgs() > 2 {
|
||||||
|
if opts, ok := process.Args[2].(map[string]interface{}); ok {
|
||||||
|
options = opts
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if process.NumOfArgs() == 3 {
|
||||||
|
data, ex := ai.ChatCompletionsWith(ctx, messages, options, nil)
|
||||||
|
if ex != nil {
|
||||||
|
ex.Throw()
|
||||||
|
}
|
||||||
|
return data
|
||||||
|
}
|
||||||
|
|
||||||
|
if process.NumOfArgs() == 4 {
|
||||||
|
|
||||||
|
switch cb := process.Args[3].(type) {
|
||||||
|
case func(data []byte) int:
|
||||||
|
res, ex := ai.ChatCompletionsWith(ctx, messages, options, cb)
|
||||||
|
if ex != nil {
|
||||||
|
ex.Throw()
|
||||||
|
}
|
||||||
|
return res
|
||||||
|
|
||||||
|
case bridge.FunctionT:
|
||||||
|
res, ex := ai.ChatCompletionsWith(ctx, messages, options, func(data []byte) int {
|
||||||
|
|
||||||
|
v, err := cb.Call(string(data))
|
||||||
|
if err != nil {
|
||||||
|
log.Error("Call callback function error: %s", err.Error())
|
||||||
|
return http.HandlerReturnError
|
||||||
|
}
|
||||||
|
|
||||||
|
ret, ok := v.(int)
|
||||||
|
if !ok {
|
||||||
|
log.Error("Callback function must return int")
|
||||||
|
return http.HandlerReturnError
|
||||||
|
}
|
||||||
|
|
||||||
|
return ret
|
||||||
|
})
|
||||||
|
|
||||||
|
if ex != nil {
|
||||||
|
ex.Throw()
|
||||||
|
}
|
||||||
|
return res
|
||||||
|
|
||||||
|
default:
|
||||||
|
exception.New("ChatCompletions error: invalid callback arguments", 400).Throw()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
res, ex := ai.ChatCompletionsWith(ctx, messages, options, nil)
|
||||||
|
if ex != nil {
|
||||||
|
ex.Throw()
|
||||||
|
}
|
||||||
|
return res
|
||||||
|
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,15 +5,79 @@ import (
|
||||||
|
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/yaoapp/gou/process"
|
"github.com/yaoapp/gou/process"
|
||||||
|
"github.com/yaoapp/yao/config"
|
||||||
|
"github.com/yaoapp/yao/test"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestProcessTiktoken(t *testing.T) {
|
func TestProcessTiktoken(t *testing.T) {
|
||||||
// Hash
|
// Hash
|
||||||
args := []interface{}{"gpt-3.5-turbo", "hello world"}
|
args := []interface{}{"gpt-3.5-turbo", "hello world"}
|
||||||
res := process.New("yao.openai.Tiktoken", args...).Run()
|
res := process.New("openai.Tiktoken", args...).Run()
|
||||||
assert.Equal(t, 2, res)
|
assert.Equal(t, 2, res)
|
||||||
|
|
||||||
args = []interface{}{"gpt-3.5-turbo", "你好世界!"}
|
args = []interface{}{"gpt-3.5-turbo", "你好世界!"}
|
||||||
res = process.New("yao.openai.Tiktoken", args...).Run()
|
res = process.New("openai.Tiktoken", args...).Run()
|
||||||
assert.Equal(t, 6, res)
|
assert.Equal(t, 6, res)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestProcessChatCompletions(t *testing.T) {
|
||||||
|
test.Prepare(t, config.Conf)
|
||||||
|
defer test.Clean()
|
||||||
|
|
||||||
|
args := []interface{}{"gpt-3_5-turbo", []map[string]interface{}{{"role": "user", "content": "hello"}}}
|
||||||
|
res := process.New("openai.chat.Completions", args...).Run()
|
||||||
|
data, ok := res.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("ChatCompletions return type error")
|
||||||
|
}
|
||||||
|
assert.NotEmpty(t, data["id"])
|
||||||
|
|
||||||
|
// With options
|
||||||
|
args = []interface{}{
|
||||||
|
"gpt-3_5-turbo",
|
||||||
|
[]map[string]interface{}{{"role": "user", "content": "hello"}},
|
||||||
|
map[string]interface{}{"max_tokens": 2},
|
||||||
|
}
|
||||||
|
res = process.New("openai.chat.Completions", args...).Run()
|
||||||
|
data, ok = res.(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("ChatCompletions return type error")
|
||||||
|
}
|
||||||
|
|
||||||
|
usage, ok := data["usage"].(map[string]interface{})
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("ChatCompletions return type error")
|
||||||
|
}
|
||||||
|
assert.Equal(t, 2, int(usage["completion_tokens"].(float64)))
|
||||||
|
|
||||||
|
// With callback
|
||||||
|
content := []byte{}
|
||||||
|
args = []interface{}{
|
||||||
|
"gpt-3_5-turbo",
|
||||||
|
[]map[string]interface{}{{"role": "user", "content": "hello"}},
|
||||||
|
nil,
|
||||||
|
func(data []byte) int {
|
||||||
|
|
||||||
|
content = append(content, data...)
|
||||||
|
if len(data) == 0 {
|
||||||
|
res = append(content, []byte("\n")...)
|
||||||
|
}
|
||||||
|
|
||||||
|
if string(data) == "data: [DONE]" {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
|
||||||
|
return 1
|
||||||
|
},
|
||||||
|
}
|
||||||
|
res = process.New("openai.chat.Completions", args...).Run()
|
||||||
|
assert.Contains(t, string(content), "[DONE]")
|
||||||
|
|
||||||
|
// With JS Callback
|
||||||
|
res, err := process.New("scripts.openai.TestProcessChatCompletions").Exec()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
assert.Contains(t, res, "[DONE]")
|
||||||
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue