mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 17:38:10 +00:00
lua: 吸收 llmsproxy 适配器高级特性(worker 池/静态预提取/动态签名钩子)
- 适配器 worker 池化:单 LState+全局锁(串行瓶颈)→ 每 adapter 一个 gopher-lua LState 池,按使用该 adapter 的源并发上限求和配置池大小,并发 transform 互不阻塞 - staticInfo 预提取:name/version/endpoint/headers 加载期编译缓存,Endpoint/Headers 读缓存不占 worker;加载即预编译首个 worker - build_headers 动态钩子 + hmac/sha256/base64/tohex 全局:签名型上游(kimicode 等)可接入 - provider applyAdapterHeaders 接入动态头(url/method/body/api_key/timestamp/source 元数据), 未定义时回落静态 headers,缺省补 Authorization - LLMSource.MaxConcurrent + core.llm.sources.<name>.max_concurrent,注册时汇总 VM.ConfigureConcurrency - 新增 Lua VM 测试(load/transform/build_headers/并发) 验证: go test ./... 27 包 0 失败;Windows 交叉编译通过;部署后 9 adapter 全部预加载
This commit is contained in:
@ -276,6 +276,7 @@ func main() {
|
||||
baseAPIKey := apiKey
|
||||
|
||||
providerMgr := agentAPI.NewProviderManager()
|
||||
adapterConcurrency := map[string]int{}
|
||||
for _, src := range cfg.LLM.Sources {
|
||||
if !agentAPI.IsValidSourceConfig(src.Name, src.BaseURL, src.Model, src.Adapter) {
|
||||
log.Printf("[homed] skip invalid llm source %q (base_url=%q model=%q adapter=%q)", src.Name, src.BaseURL, src.Model, src.Adapter)
|
||||
@ -292,9 +293,14 @@ func main() {
|
||||
Temperature: cfg.LLM.Temperature,
|
||||
MaxTokens: cfg.LLM.MaxTokens,
|
||||
ContextWindow: src.ContextWindow,
|
||||
MaxConcurrent: src.MaxConcurrent,
|
||||
}, luaVM, src.Name, src.Adapter)
|
||||
providerMgr.Register(src.Name, luaProvider)
|
||||
if src.Adapter != "" {
|
||||
adapterConcurrency[src.Adapter] += src.MaxConcurrent
|
||||
}
|
||||
}
|
||||
luaVM.ConfigureConcurrency(adapterConcurrency)
|
||||
if cfg.LLM.Provider != "" {
|
||||
providerMgr.SetDefault(cfg.LLM.Provider)
|
||||
}
|
||||
|
||||
@ -201,6 +201,7 @@ type BaseConfig struct {
|
||||
Temperature float64 `json:"temperature"`
|
||||
MaxTokens int `json:"max_tokens"`
|
||||
ContextWindow int `json:"context_window"`
|
||||
MaxConcurrent int `json:"max_concurrent"`
|
||||
}
|
||||
|
||||
// LuaAdaptedProvider 使用 Lua 脚本做请求/响应变换,直接发起 HTTP 调用
|
||||
@ -272,11 +273,7 @@ func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (
|
||||
return nil, fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
httpReq.Header.Set("Content-Type", "application/json")
|
||||
httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey)
|
||||
|
||||
for k, v := range p.vm.GetAdapterHeaders(p.adapter) {
|
||||
httpReq.Header.Set(k, v)
|
||||
}
|
||||
p.applyAdapterHeaders(httpReq, url, transformedBody)
|
||||
|
||||
resp, err := p.client.Do(httpReq)
|
||||
if err != nil {
|
||||
@ -315,6 +312,31 @@ func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// applyAdapterHeaders 优先调用 adapter.build_headers(meta) 动态签名钩子,
|
||||
// 未定义时回落到静态 adapter.headers,最后确保带 Authorization。
|
||||
func (p *LuaAdaptedProvider) applyAdapterHeaders(httpReq *http.Request, url, body string) {
|
||||
meta := map[string]interface{}{
|
||||
"url": url,
|
||||
"method": http.MethodPost,
|
||||
"body": body,
|
||||
"api_key": p.cfg.APIKey,
|
||||
"timestamp": time.Now().Unix(),
|
||||
"source": map[string]interface{}{
|
||||
"name": p.name,
|
||||
},
|
||||
}
|
||||
hdrs, err := p.vm.BuildHeaders(p.adapter, meta)
|
||||
if err != nil {
|
||||
hdrs = p.vm.GetAdapterHeaders(p.adapter)
|
||||
}
|
||||
for k, v := range hdrs {
|
||||
httpReq.Header.Set(k, v)
|
||||
}
|
||||
if httpReq.Header.Get("Authorization") == "" && p.cfg.APIKey != "" {
|
||||
httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey)
|
||||
}
|
||||
}
|
||||
|
||||
func parseOpenAICompatibleResponse(raw []byte) (*CompletionResponse, error) {
|
||||
var resp struct {
|
||||
Usage struct {
|
||||
|
||||
@ -185,6 +185,7 @@ var sourceFieldDefs = []struct {
|
||||
{"thinking_enabled", "bool", "深度思考"},
|
||||
{"adapter", "string", "适配器"},
|
||||
{"adapter_path", "string", "适配器路径"},
|
||||
{"max_concurrent", "int", "并发上限"},
|
||||
}
|
||||
|
||||
// registerSourceDefs 注册 core.llm.sources.<name>.* 的 ConfigDef
|
||||
@ -814,6 +815,7 @@ func (r *ConfigRegistry) ToConfig() *types.Config {
|
||||
Adapter: read(p+".adapter", ""),
|
||||
AdapterPath: read(p+".adapter_path", ""),
|
||||
ContextWindow: readInt(p+".context_window", 0),
|
||||
MaxConcurrent: readInt(p+".max_concurrent", 8),
|
||||
ThinkingEnabled: readBool(p+".thinking_enabled", false),
|
||||
})
|
||||
}
|
||||
|
||||
@ -1,13 +1,19 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"embed"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@ -18,37 +24,458 @@ import (
|
||||
//go:embed adapters/*.lua
|
||||
var bundledAdapters embed.FS
|
||||
|
||||
// adapterGlobal 是每个 worker state 中保存适配器表的保留全局名。
|
||||
const adapterGlobal = "__ha_adapter"
|
||||
|
||||
type APIAdapter struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
}
|
||||
|
||||
// AdapterCache 预加载容器:启动时编译全部脚本到内存,运行时只读缓存,无文件 I/O
|
||||
type AdapterCache struct {
|
||||
mu sync.RWMutex
|
||||
state *lua.LState
|
||||
items map[string]*lua.LTable
|
||||
// worker 封装一个独立的 gopher-lua 解释器。LState 非并发安全,
|
||||
// 同一时刻只能被一个 goroutine 使用(从池中取出)。
|
||||
type worker struct {
|
||||
L *lua.LState
|
||||
}
|
||||
|
||||
func newAdapterCache() *AdapterCache {
|
||||
return &AdapterCache{
|
||||
state: lua.NewState(),
|
||||
items: make(map[string]*lua.LTable),
|
||||
func (w *worker) close() {
|
||||
if w.L != nil {
|
||||
w.L.Close()
|
||||
w.L = nil
|
||||
}
|
||||
}
|
||||
|
||||
func (c *AdapterCache) setupGlobals() {
|
||||
s := c.state
|
||||
s.SetGlobal("log", s.NewFunction(func(L *lua.LState) int {
|
||||
// staticInfo 缓存适配器脚本在加载时提取出的不可变字段,
|
||||
// Endpoint/Headers 等读取不需要占用池 worker。
|
||||
type staticInfo struct {
|
||||
name string
|
||||
version string
|
||||
endpoint string
|
||||
headers map[string]string
|
||||
}
|
||||
|
||||
// adapterPool 管理单个适配器的 worker 池。空闲 worker 存放在 idle 切片,
|
||||
// 按 target 惰性创建;同一 worker 同一时刻只被一个 goroutine 使用。
|
||||
type adapterPool struct {
|
||||
name string
|
||||
script string
|
||||
static staticInfo
|
||||
|
||||
mu sync.Mutex
|
||||
cond *sync.Cond
|
||||
idle []*worker
|
||||
created int
|
||||
target int
|
||||
closed bool
|
||||
}
|
||||
|
||||
func newAdapterPool(name, script string, static staticInfo) *adapterPool {
|
||||
p := &adapterPool{name: name, script: script, static: static, target: 1}
|
||||
p.cond = sync.NewCond(&p.mu)
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *adapterPool) setTarget(n int) {
|
||||
p.mu.Lock()
|
||||
p.target = n
|
||||
p.cond.Broadcast()
|
||||
p.mu.Unlock()
|
||||
}
|
||||
|
||||
// boot 创建全新 gopher-lua 状态:注册共享全局、执行适配器脚本、
|
||||
// 把返回表存到保留全局。compile Lua 的成本较高,尽量复用池。
|
||||
func (p *adapterPool) boot() (*worker, error) {
|
||||
L := lua.NewState()
|
||||
setupGlobals(L)
|
||||
if err := L.DoString(p.script); err != nil {
|
||||
L.Close()
|
||||
return nil, fmt.Errorf("compile adapter %s: %w", p.name, err)
|
||||
}
|
||||
tbl, ok := L.Get(-1).(*lua.LTable)
|
||||
L.Pop(1)
|
||||
if !ok {
|
||||
L.Close()
|
||||
return nil, fmt.Errorf("adapter %s must return a table", p.name)
|
||||
}
|
||||
L.SetGlobal(adapterGlobal, tbl)
|
||||
return &worker{L: L}, nil
|
||||
}
|
||||
|
||||
func (p *adapterPool) acquire() (*worker, error) {
|
||||
p.mu.Lock()
|
||||
for {
|
||||
if p.closed {
|
||||
p.mu.Unlock()
|
||||
return nil, fmt.Errorf("adapter %s pool closed", p.name)
|
||||
}
|
||||
if n := len(p.idle); n > 0 {
|
||||
w := p.idle[n-1]
|
||||
p.idle = p.idle[:n-1]
|
||||
p.mu.Unlock()
|
||||
return w, nil
|
||||
}
|
||||
if p.created < p.target {
|
||||
p.created++
|
||||
break
|
||||
}
|
||||
p.cond.Wait()
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
w, err := p.boot()
|
||||
if err != nil {
|
||||
p.mu.Lock()
|
||||
p.created--
|
||||
p.cond.Signal()
|
||||
p.mu.Unlock()
|
||||
return nil, err
|
||||
}
|
||||
return w, nil
|
||||
}
|
||||
|
||||
func (p *adapterPool) release(w *worker) {
|
||||
p.mu.Lock()
|
||||
if p.closed {
|
||||
p.mu.Unlock()
|
||||
w.close()
|
||||
return
|
||||
}
|
||||
p.idle = append(p.idle, w)
|
||||
p.cond.Signal()
|
||||
p.mu.Unlock()
|
||||
}
|
||||
|
||||
func (p *adapterPool) shutdown() {
|
||||
p.mu.Lock()
|
||||
p.closed = true
|
||||
idle := p.idle
|
||||
p.idle = nil
|
||||
p.cond.Broadcast()
|
||||
p.mu.Unlock()
|
||||
for _, w := range idle {
|
||||
w.close()
|
||||
}
|
||||
}
|
||||
|
||||
// VM 聚合各适配器的 worker 池。所有导出方法并发安全。
|
||||
type VM struct {
|
||||
mu sync.RWMutex
|
||||
dir string
|
||||
pools map[string]*adapterPool
|
||||
}
|
||||
|
||||
func NewVM(dir string) *VM {
|
||||
return &VM{dir: dir, pools: map[string]*adapterPool{}}
|
||||
}
|
||||
|
||||
func (v *VM) AdapterDir() string { return v.dir }
|
||||
|
||||
func (v *VM) Start() error {
|
||||
if err := os.MkdirAll(v.dir, 0755); err != nil {
|
||||
return fmt.Errorf("mkdir adapter dir: %w", err)
|
||||
}
|
||||
if err := v.writeBundledAdapters(); err != nil {
|
||||
return fmt.Errorf("write bundled adapters: %w", err)
|
||||
}
|
||||
|
||||
entries, err := os.ReadDir(v.dir)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if filepath.Ext(entry.Name()) != ".lua" {
|
||||
continue
|
||||
}
|
||||
path := filepath.Join(v.dir, entry.Name())
|
||||
if err := v.LoadAdapter(path); err != nil {
|
||||
fmt.Printf("[lua] preload %s: %v\n", entry.Name(), err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v *VM) Stop() {
|
||||
v.mu.Lock()
|
||||
pools := v.pools
|
||||
v.pools = map[string]*adapterPool{}
|
||||
v.mu.Unlock()
|
||||
for _, p := range pools {
|
||||
p.shutdown()
|
||||
}
|
||||
}
|
||||
|
||||
// ConfigureConcurrency sets each adapter's worker target to the sum of every
|
||||
// source's max concurrent calls that share the adapter. Clamped >= 1.
|
||||
func (v *VM) ConfigureConcurrency(adapterConcurrency map[string]int) {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
for name, n := range adapterConcurrency {
|
||||
if p, ok := v.pools[name]; ok {
|
||||
if n < 1 {
|
||||
n = 1
|
||||
}
|
||||
p.setTarget(n)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (v *VM) LoadAdapter(path string) error {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read adapter: %w", err)
|
||||
}
|
||||
name := strings.TrimSuffix(filepath.Base(path), filepath.Ext(path))
|
||||
return v.LoadAdapterSource(name, string(data))
|
||||
}
|
||||
|
||||
func (v *VM) LoadAdapterSource(name, code string) error {
|
||||
static, err := inspectScript(name, code)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
target := 1
|
||||
v.mu.RLock()
|
||||
if p, ok := v.pools[static.name]; ok {
|
||||
target = p.target
|
||||
}
|
||||
v.mu.RUnlock()
|
||||
|
||||
v.mu.Lock()
|
||||
if p, ok := v.pools[static.name]; ok {
|
||||
p.shutdown()
|
||||
}
|
||||
p := newAdapterPool(static.name, code, *static)
|
||||
p.setTarget(target)
|
||||
v.pools[static.name] = p
|
||||
v.mu.Unlock()
|
||||
|
||||
// 立即 boot 一个 worker,把编译成本放到加载阶段而不是首个请求。
|
||||
w, err := p.boot()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.release(w)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v *VM) RemoveAdapter(name string) {
|
||||
v.mu.Lock()
|
||||
p := v.pools[name]
|
||||
delete(v.pools, name)
|
||||
v.mu.Unlock()
|
||||
if p != nil {
|
||||
p.shutdown()
|
||||
}
|
||||
}
|
||||
|
||||
func (v *VM) ListAdapters() []APIAdapter {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
list := make([]APIAdapter, 0, len(v.pools))
|
||||
for name, p := range v.pools {
|
||||
list = append(list, APIAdapter{Name: name, Version: p.static.version})
|
||||
}
|
||||
sort.Slice(list, func(i, j int) bool { return list[i].Name < list[j].Name })
|
||||
return list
|
||||
}
|
||||
|
||||
func (v *VM) pool(name string) *adapterPool {
|
||||
v.mu.RLock()
|
||||
defer v.mu.RUnlock()
|
||||
return v.pools[name]
|
||||
}
|
||||
|
||||
// callMethod 在池 worker 上执行 adapter 的 fn(raw) 并返回字符串结果。
|
||||
func callMethod(p *adapterPool, fn, raw string) (string, error) {
|
||||
w, err := p.acquire()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer p.release(w)
|
||||
L := w.L
|
||||
adapter := L.GetGlobal(adapterGlobal)
|
||||
tbl, ok := adapter.(*lua.LTable)
|
||||
if !ok {
|
||||
return "", fmt.Errorf("adapter %s has no table", p.name)
|
||||
}
|
||||
f := tbl.RawGetString(fn)
|
||||
if _, ok := f.(*lua.LFunction); !ok {
|
||||
return "", fmt.Errorf("adapter %s missing %s", p.name, fn)
|
||||
}
|
||||
L.Push(f)
|
||||
L.Push(lua.LString(raw))
|
||||
if err := L.PCall(1, 1, nil); err != nil {
|
||||
return "", fmt.Errorf("%s: %w", fn, err)
|
||||
}
|
||||
out := L.Get(-1)
|
||||
L.Pop(1)
|
||||
return out.String(), nil
|
||||
}
|
||||
|
||||
func (v *VM) CallTransformRequest(name, rawJSON string) (string, error) {
|
||||
p := v.pool(name)
|
||||
if p == nil {
|
||||
return "", fmt.Errorf("adapter %s not loaded", name)
|
||||
}
|
||||
return callMethod(p, "transform_request", rawJSON)
|
||||
}
|
||||
|
||||
func (v *VM) CallTransformResponse(name, rawJSON string) (string, error) {
|
||||
p := v.pool(name)
|
||||
if p == nil {
|
||||
return rawJSON, nil
|
||||
}
|
||||
out, err := callMethod(p, "transform_response", rawJSON)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (v *VM) CallTransformStreamChunk(name, rawLine string) (string, error) {
|
||||
p := v.pool(name)
|
||||
if p == nil {
|
||||
return rawLine, nil
|
||||
}
|
||||
out, err := callMethod(p, "transform_stream_chunk", rawLine)
|
||||
if err != nil {
|
||||
return rawLine, nil
|
||||
}
|
||||
if out == "" {
|
||||
return "", nil
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
func (v *VM) GetAdapterEndpoint(name string) string {
|
||||
p := v.pool(name)
|
||||
if p == nil {
|
||||
return ""
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return p.static.endpoint
|
||||
}
|
||||
|
||||
func (v *VM) GetAdapterHeaders(name string) map[string]string {
|
||||
p := v.pool(name)
|
||||
if p == nil {
|
||||
return nil
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
hdr := make(map[string]string, len(p.static.headers))
|
||||
for k, h := range p.static.headers {
|
||||
hdr[k] = h
|
||||
}
|
||||
return hdr
|
||||
}
|
||||
|
||||
// BuildHeaders executes adapter.build_headers(meta); falls back to the static
|
||||
// adapter.headers table when no dynamic hook is defined.
|
||||
func (v *VM) BuildHeaders(name string, meta map[string]interface{}) (map[string]string, error) {
|
||||
p := v.pool(name)
|
||||
if p == nil {
|
||||
return nil, fmt.Errorf("adapter %s not loaded", name)
|
||||
}
|
||||
w, err := p.acquire()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer p.release(w)
|
||||
L := w.L
|
||||
adapter := L.GetGlobal(adapterGlobal)
|
||||
tbl, ok := adapter.(*lua.LTable)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("adapter %s has no table", name)
|
||||
}
|
||||
f := tbl.RawGetString("build_headers")
|
||||
if _, ok := f.(*lua.LFunction); !ok {
|
||||
// no dynamic hook -> static headers
|
||||
p.mu.Lock()
|
||||
hdr := make(map[string]string, len(p.static.headers))
|
||||
for k, h := range p.static.headers {
|
||||
hdr[k] = h
|
||||
}
|
||||
p.mu.Unlock()
|
||||
return hdr, nil
|
||||
}
|
||||
L.Push(f)
|
||||
L.Push(goValueToLua(L, meta))
|
||||
if err := L.PCall(1, 1, nil); err != nil {
|
||||
return nil, fmt.Errorf("build_headers: %w", err)
|
||||
}
|
||||
res := L.Get(-1)
|
||||
L.Pop(1)
|
||||
tbl2, ok := res.(*lua.LTable)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("build_headers returned non-table")
|
||||
}
|
||||
headers := make(map[string]string)
|
||||
tbl2.ForEach(func(key, val lua.LValue) {
|
||||
headers[key.String()] = val.String()
|
||||
})
|
||||
return headers, nil
|
||||
}
|
||||
|
||||
// ---- static inspection (compile-once at load time) ----
|
||||
|
||||
func inspectScript(name, code string) (*staticInfo, error) {
|
||||
L := lua.NewState()
|
||||
setupGlobals(L)
|
||||
if err := L.DoString(code); err != nil {
|
||||
L.Close()
|
||||
return nil, fmt.Errorf("compile adapter: %w", err)
|
||||
}
|
||||
tbl, ok := L.Get(-1).(*lua.LTable)
|
||||
L.Pop(1)
|
||||
if !ok {
|
||||
L.Close()
|
||||
return nil, fmt.Errorf("adapter script must return a table")
|
||||
}
|
||||
info := &staticInfo{name: name, headers: map[string]string{}}
|
||||
if n := readString(tbl, "name"); n != "" {
|
||||
info.name = n
|
||||
}
|
||||
info.version = readString(tbl, "version")
|
||||
info.endpoint = readString(tbl, "endpoint")
|
||||
if ht, ok := tbl.RawGetString("headers").(*lua.LTable); ok {
|
||||
ht.ForEach(func(k, val lua.LValue) {
|
||||
info.headers[k.String()] = val.String()
|
||||
})
|
||||
}
|
||||
L.Close()
|
||||
return info, nil
|
||||
}
|
||||
|
||||
func readString(tbl *lua.LTable, field string) string {
|
||||
v := tbl.RawGetString(field)
|
||||
switch x := v.(type) {
|
||||
case lua.LString:
|
||||
return string(x)
|
||||
case lua.LNumber:
|
||||
return strconv.FormatFloat(float64(x), 'f', -1, 64)
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
// ---- shared globals installed into every worker state ----
|
||||
|
||||
func setupGlobals(L *lua.LState) {
|
||||
L.SetGlobal("log", L.NewFunction(func(L *lua.LState) int {
|
||||
level := L.ToString(1)
|
||||
msg := L.ToString(2)
|
||||
fmt.Printf("[lua/%s] %s\n", level, msg)
|
||||
return 0
|
||||
}))
|
||||
|
||||
jsonTable := s.NewTable()
|
||||
s.SetGlobal("json", jsonTable)
|
||||
s.SetField(jsonTable, "encode", s.NewFunction(func(L *lua.LState) int {
|
||||
jsonTable := L.NewTable()
|
||||
L.SetGlobal("json", jsonTable)
|
||||
L.SetField(jsonTable, "encode", L.NewFunction(func(L *lua.LState) int {
|
||||
val := L.CheckAny(1)
|
||||
goVal := luaValueToGo(val)
|
||||
b, err := json.Marshal(goVal)
|
||||
@ -59,7 +486,7 @@ func (c *AdapterCache) setupGlobals() {
|
||||
L.Push(lua.LString(string(b)))
|
||||
return 1
|
||||
}))
|
||||
s.SetField(jsonTable, "decode", s.NewFunction(func(L *lua.LState) int {
|
||||
L.SetField(jsonTable, "decode", L.NewFunction(func(L *lua.LState) int {
|
||||
str := L.CheckString(1)
|
||||
var val interface{}
|
||||
if err := json.Unmarshal([]byte(str), &val); err != nil {
|
||||
@ -70,7 +497,7 @@ func (c *AdapterCache) setupGlobals() {
|
||||
return 1
|
||||
}))
|
||||
|
||||
s.SetGlobal("http_get", s.NewFunction(func(L *lua.LState) int {
|
||||
L.SetGlobal("http_get", L.NewFunction(func(L *lua.LState) int {
|
||||
url := L.ToString(1)
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp, err := client.Get(url)
|
||||
@ -86,7 +513,7 @@ func (c *AdapterCache) setupGlobals() {
|
||||
return 1
|
||||
}))
|
||||
|
||||
s.SetGlobal("http_post", s.NewFunction(func(L *lua.LState) int {
|
||||
L.SetGlobal("http_post", L.NewFunction(func(L *lua.LState) int {
|
||||
url := L.ToString(1)
|
||||
body := L.ToString(2)
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
@ -102,258 +529,39 @@ func (c *AdapterCache) setupGlobals() {
|
||||
L.Push(lua.LString(string(result)))
|
||||
return 1
|
||||
}))
|
||||
}
|
||||
|
||||
// Preload 编译单个 Lua 适配器脚本并注入缓存
|
||||
func (c *AdapterCache) Preload(path string) error {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read adapter: %w", err)
|
||||
}
|
||||
return c.PreloadSource(filepath.Base(path), string(data))
|
||||
}
|
||||
|
||||
// PreloadSource 从源码字符串编译适配器并注入缓存
|
||||
func (c *AdapterCache) PreloadSource(name, code string) error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if err := c.state.DoString(code); err != nil {
|
||||
return fmt.Errorf("compile adapter: %w", err)
|
||||
}
|
||||
|
||||
tbl, ok := c.state.Get(-1).(*lua.LTable)
|
||||
c.state.Pop(1)
|
||||
if !ok {
|
||||
return fmt.Errorf("adapter script must return a table")
|
||||
}
|
||||
|
||||
if n := tbl.RawGetString("name"); n != nil && n.String() != "" {
|
||||
name = n.String()
|
||||
}
|
||||
|
||||
c.items[name] = tbl
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get 运行时从缓存读取已编译的适配器表(无文件 I/O)
|
||||
func (c *AdapterCache) Get(name string) *lua.LTable {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.items[name]
|
||||
}
|
||||
|
||||
// Remove 从缓存移除适配器
|
||||
func (c *AdapterCache) Remove(name string) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
delete(c.items, name)
|
||||
}
|
||||
|
||||
// List 返回缓存中所有适配器摘要
|
||||
func (c *AdapterCache) List() []APIAdapter {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
list := make([]APIAdapter, 0, len(c.items))
|
||||
for name, tbl := range c.items {
|
||||
a := APIAdapter{Name: name}
|
||||
if v := tbl.RawGetString("version"); v != nil {
|
||||
a.Version = v.String()
|
||||
}
|
||||
list = append(list, a)
|
||||
}
|
||||
return list
|
||||
}
|
||||
|
||||
// Close 释放 Lua 状态
|
||||
func (c *AdapterCache) Close() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.state != nil {
|
||||
c.state.Close()
|
||||
c.state = nil
|
||||
}
|
||||
c.items = nil
|
||||
}
|
||||
|
||||
// VM 运行时虚拟机,封装 AdapterCache 提供适配器调用
|
||||
type VM struct {
|
||||
mu sync.Mutex
|
||||
cache *AdapterCache
|
||||
adapterDir string
|
||||
}
|
||||
|
||||
func NewVM(adapterDir string) *VM {
|
||||
return &VM{
|
||||
adapterDir: adapterDir,
|
||||
cache: newAdapterCache(),
|
||||
}
|
||||
}
|
||||
|
||||
func (v *VM) AdapterDir() string { return v.adapterDir }
|
||||
func (v *VM) Cache() *AdapterCache { return v.cache }
|
||||
|
||||
func (v *VM) Start() error {
|
||||
if err := os.MkdirAll(v.adapterDir, 0755); err != nil {
|
||||
return fmt.Errorf("mkdir adapter dir: %w", err)
|
||||
}
|
||||
if err := v.writeBundledAdapters(); err != nil {
|
||||
return fmt.Errorf("write bundled adapters: %w", err)
|
||||
}
|
||||
|
||||
v.cache.setupGlobals()
|
||||
|
||||
// 预加载:扫描适配器目录,全部编译到缓存
|
||||
entries, err := os.ReadDir(v.adapterDir)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if filepath.Ext(entry.Name()) != ".lua" {
|
||||
continue
|
||||
}
|
||||
path := filepath.Join(v.adapterDir, entry.Name())
|
||||
if err := v.cache.Preload(path); err != nil {
|
||||
fmt.Printf("[lua] preload %s: %v\n", entry.Name(), err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (v *VM) Stop() {
|
||||
v.cache.Close()
|
||||
}
|
||||
|
||||
// LoadAdapter 对外接口:从文件加载并编译适配器到缓存(运行时安全,不影响其他适配器)
|
||||
func (v *VM) LoadAdapter(path string) error {
|
||||
return v.cache.Preload(path)
|
||||
}
|
||||
|
||||
// RemoveAdapter 对外接口:从缓存移除适配器(运行时安全)
|
||||
func (v *VM) RemoveAdapter(name string) {
|
||||
v.cache.Remove(name)
|
||||
}
|
||||
|
||||
func (v *VM) ListAdapters() []APIAdapter {
|
||||
return v.cache.List()
|
||||
}
|
||||
|
||||
func (v *VM) CallTransformRequest(name, rawJSON string) (string, error) {
|
||||
adapter := v.cache.Get(name)
|
||||
if adapter == nil {
|
||||
return "", fmt.Errorf("adapter %s not loaded", name)
|
||||
}
|
||||
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
fn := adapter.RawGetString("transform_request")
|
||||
if fn == nil {
|
||||
return "", fmt.Errorf("adapter %s missing transform_request", name)
|
||||
}
|
||||
|
||||
state := v.cache.state
|
||||
state.Push(fn)
|
||||
state.Push(lua.LString(rawJSON))
|
||||
if err := state.PCall(1, 1, nil); err != nil {
|
||||
return "", fmt.Errorf("transform_request: %w", err)
|
||||
}
|
||||
result := state.Get(-1)
|
||||
state.Pop(1)
|
||||
return result.String(), nil
|
||||
}
|
||||
|
||||
func (v *VM) CallTransformResponse(name, rawJSON string) (string, error) {
|
||||
adapter := v.cache.Get(name)
|
||||
if adapter == nil {
|
||||
return rawJSON, nil
|
||||
}
|
||||
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
fn := adapter.RawGetString("transform_response")
|
||||
if fn == nil {
|
||||
return rawJSON, nil
|
||||
}
|
||||
|
||||
state := v.cache.state
|
||||
state.Push(fn)
|
||||
state.Push(lua.LString(rawJSON))
|
||||
if err := state.PCall(1, 1, nil); err != nil {
|
||||
return "", fmt.Errorf("transform_response: %w", err)
|
||||
}
|
||||
result := state.Get(-1)
|
||||
state.Pop(1)
|
||||
return result.String(), nil
|
||||
}
|
||||
|
||||
func (v *VM) CallTransformStreamChunk(name, rawLine string) (string, error) {
|
||||
adapter := v.cache.Get(name)
|
||||
if adapter == nil {
|
||||
return rawLine, nil
|
||||
}
|
||||
|
||||
v.mu.Lock()
|
||||
defer v.mu.Unlock()
|
||||
|
||||
fn := adapter.RawGetString("transform_stream_chunk")
|
||||
if fn == nil {
|
||||
return rawLine, nil
|
||||
}
|
||||
|
||||
state := v.cache.state
|
||||
state.Push(fn)
|
||||
state.Push(lua.LString(rawLine))
|
||||
if err := state.PCall(1, 1, nil); err != nil {
|
||||
return "", fmt.Errorf("transform_stream_chunk: %w", err)
|
||||
}
|
||||
result := state.Get(-1)
|
||||
state.Pop(1)
|
||||
if result.String() == "" {
|
||||
return "", nil
|
||||
}
|
||||
return result.String(), nil
|
||||
}
|
||||
|
||||
func (v *VM) GetAdapterEndpoint(name string) string {
|
||||
adapter := v.cache.Get(name)
|
||||
if adapter == nil {
|
||||
return ""
|
||||
}
|
||||
if ep := adapter.RawGetString("endpoint"); ep != nil {
|
||||
return ep.String()
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (v *VM) GetAdapterHeaders(name string) map[string]string {
|
||||
adapter := v.cache.Get(name)
|
||||
if adapter == nil {
|
||||
return nil
|
||||
}
|
||||
headers := make(map[string]string)
|
||||
if ht := adapter.RawGetString("headers"); ht != nil {
|
||||
if tbl, ok := ht.(*lua.LTable); ok {
|
||||
tbl.ForEach(func(key, val lua.LValue) {
|
||||
headers[key.String()] = val.String()
|
||||
})
|
||||
}
|
||||
}
|
||||
return headers
|
||||
// 签名辅助(kimicode 等上游签名型 adapter 需要)
|
||||
L.SetGlobal("hmac_sha256_hex", L.NewFunction(func(L *lua.LState) int {
|
||||
key := L.ToString(1)
|
||||
data := L.ToString(2)
|
||||
m := hmac.New(sha256.New, []byte(key))
|
||||
m.Write([]byte(data))
|
||||
L.Push(lua.LString(hex.EncodeToString(m.Sum(nil))))
|
||||
return 1
|
||||
}))
|
||||
L.SetGlobal("sha256_hex", L.NewFunction(func(L *lua.LState) int {
|
||||
h := sha256.Sum256([]byte(L.ToString(1)))
|
||||
L.Push(lua.LString(hex.EncodeToString(h[:])))
|
||||
return 1
|
||||
}))
|
||||
L.SetGlobal("base64_encode", L.NewFunction(func(L *lua.LState) int {
|
||||
L.Push(lua.LString(base64.StdEncoding.EncodeToString([]byte(L.ToString(1)))))
|
||||
return 1
|
||||
}))
|
||||
L.SetGlobal("tohex", L.NewFunction(func(L *lua.LState) int {
|
||||
L.Push(lua.LString(hex.EncodeToString([]byte(L.ToString(1)))))
|
||||
return 1
|
||||
}))
|
||||
}
|
||||
|
||||
func (v *VM) writeBundledAdapters() error {
|
||||
known := []string{
|
||||
"openai", "anthropic", "deepseek", "gemini",
|
||||
"github", "groq", "mistral", "ollama",
|
||||
"github", "groq", "mistral", "ollama", "kimicode",
|
||||
}
|
||||
for _, name := range known {
|
||||
srcPath := "adapters/" + name + ".lua"
|
||||
dstPath := filepath.Join(v.adapterDir, name+".lua")
|
||||
dstPath := filepath.Join(v.dir, name+".lua")
|
||||
if _, err := os.Stat(dstPath); err == nil {
|
||||
continue
|
||||
}
|
||||
@ -422,6 +630,13 @@ func goValueToLua(L *lua.LState, val interface{}) lua.LValue {
|
||||
}
|
||||
return tbl
|
||||
default:
|
||||
b, jerr := json.Marshal(v)
|
||||
if jerr == nil {
|
||||
var iv interface{}
|
||||
if json.Unmarshal(b, &iv) == nil {
|
||||
return goValueToLua(L, iv)
|
||||
}
|
||||
}
|
||||
return lua.LNil
|
||||
}
|
||||
}
|
||||
|
||||
137
internal/lua/vm_test.go
Normal file
137
internal/lua/vm_test.go
Normal file
@ -0,0 +1,137 @@
|
||||
package lua
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestVMLoadAndTransform(t *testing.T) {
|
||||
vm := NewVM(t.TempDir())
|
||||
if err := vm.Start(); err != nil {
|
||||
t.Fatalf("start vm: %v", err)
|
||||
}
|
||||
defer vm.Stop()
|
||||
|
||||
out, err := vm.CallTransformRequest("openai", `{"model":"x"}`)
|
||||
if err != nil {
|
||||
t.Fatalf("transform_request: %v", err)
|
||||
}
|
||||
if !strings.Contains(out, `"model":"x"`) {
|
||||
t.Fatalf("unexpected output: %s", out)
|
||||
}
|
||||
|
||||
ep := vm.GetAdapterEndpoint("openai")
|
||||
if ep == "" {
|
||||
t.Fatal("openai endpoint empty")
|
||||
}
|
||||
hdrs := vm.GetAdapterHeaders("openai")
|
||||
if hdrs == nil {
|
||||
t.Fatal("openai headers nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestVMLoadCustomAdapter(t *testing.T) {
|
||||
vm := NewVM(t.TempDir())
|
||||
if err := vm.Start(); err != nil {
|
||||
t.Fatalf("start vm: %v", err)
|
||||
}
|
||||
defer vm.Stop()
|
||||
|
||||
code := `
|
||||
local adapter = {}
|
||||
adapter.name = "testa"
|
||||
adapter.version = "1.0"
|
||||
adapter.endpoint = "/x"
|
||||
adapter.headers = { ["X-A"] = "1" }
|
||||
function adapter.transform_request(raw) return "REQ" end
|
||||
function adapter.build_headers(meta) return { ["X-Key"] = "k" } end
|
||||
function adapter.transform_response(raw) return "RESP" end
|
||||
function adapter.transform_stream_chunk(raw) return "CHUNK" end
|
||||
return adapter
|
||||
`
|
||||
if err := vm.LoadAdapterSource("testa", code); err != nil {
|
||||
t.Fatalf("load adapter: %v", err)
|
||||
}
|
||||
|
||||
if out, err := vm.CallTransformRequest("testa", "x"); err != nil || out != "REQ" {
|
||||
t.Fatalf("req=%q err=%v", out, err)
|
||||
}
|
||||
if out, err := vm.CallTransformResponse("testa", "x"); err != nil || out != "RESP" {
|
||||
t.Fatalf("resp=%q err=%v", out, err)
|
||||
}
|
||||
if out, err := vm.CallTransformStreamChunk("testa", "x"); err != nil || out != "CHUNK" {
|
||||
t.Fatalf("chunk=%q err=%v", out, err)
|
||||
}
|
||||
|
||||
meta := map[string]interface{}{"url": "http://x", "api_key": "k1"}
|
||||
hdrs, err := vm.BuildHeaders("testa", meta)
|
||||
if err != nil {
|
||||
t.Fatalf("build_headers: %v", err)
|
||||
}
|
||||
if hdrs["X-Key"] != "k" {
|
||||
t.Fatalf("missing dynamic header: %v", hdrs)
|
||||
}
|
||||
|
||||
if ep := vm.GetAdapterEndpoint("testa"); ep != "/x" {
|
||||
t.Fatalf("endpoint=%q", ep)
|
||||
}
|
||||
if st := vm.GetAdapterHeaders("testa"); st["X-A"] != "1" {
|
||||
t.Fatalf("static headers=%v", st)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVMLoadStaticHeaderFallback(t *testing.T) {
|
||||
vm := NewVM(t.TempDir())
|
||||
if err := vm.Start(); err != nil {
|
||||
t.Fatalf("start vm: %v", err)
|
||||
}
|
||||
defer vm.Stop()
|
||||
|
||||
code := `
|
||||
local adapter = {}
|
||||
adapter.name = "statics"
|
||||
adapter.headers = { ["X-S"] = "s" }
|
||||
function adapter.transform_request(raw) return raw end
|
||||
return adapter
|
||||
`
|
||||
if err := vm.LoadAdapterSource("statics", code); err != nil {
|
||||
t.Fatalf("load: %v", err)
|
||||
}
|
||||
hdrs, err := vm.BuildHeaders("statics", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("build_headers: %v", err)
|
||||
}
|
||||
if hdrs["X-S"] != "s" {
|
||||
t.Fatalf("expected static fallback, got %v", hdrs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVMConcurrentCalls(t *testing.T) {
|
||||
vm := NewVM(t.TempDir())
|
||||
if err := vm.Start(); err != nil {
|
||||
t.Fatalf("start vm: %v", err)
|
||||
}
|
||||
defer vm.Stop()
|
||||
|
||||
var wg sync.WaitGroup
|
||||
errs := make(chan error, 40)
|
||||
for i := 0; i < 20; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for j := 0; j < 10; j++ {
|
||||
_, err := vm.CallTransformRequest("openai", `{"model":"x"}`)
|
||||
if err != nil {
|
||||
errs <- err
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
for err := range errs {
|
||||
t.Fatalf("concurrent call: %v", err)
|
||||
}
|
||||
}
|
||||
@ -109,6 +109,7 @@ func (l *llmImpl) ReloadFromConfig() error {
|
||||
return nil
|
||||
}
|
||||
l.mgr.Reset()
|
||||
adapterConcurrency := map[string]int{}
|
||||
for _, src := range cfg.LLM.Sources {
|
||||
if !agentAPI.IsValidSourceConfig(src.Name, src.BaseURL, src.Model, src.Adapter) {
|
||||
continue
|
||||
@ -124,8 +125,15 @@ func (l *llmImpl) ReloadFromConfig() error {
|
||||
Temperature: cfg.LLM.Temperature,
|
||||
MaxTokens: cfg.LLM.MaxTokens,
|
||||
ContextWindow: src.ContextWindow,
|
||||
MaxConcurrent: src.MaxConcurrent,
|
||||
}, l.lua, src.Name, src.Adapter)
|
||||
l.mgr.Register(src.Name, provider)
|
||||
if src.Adapter != "" {
|
||||
adapterConcurrency[src.Adapter] += src.MaxConcurrent
|
||||
}
|
||||
}
|
||||
if l.lua != nil {
|
||||
l.lua.ConfigureConcurrency(adapterConcurrency)
|
||||
}
|
||||
if cfg.LLM.Provider != "" {
|
||||
if l.mgr.Get(cfg.LLM.Provider) != nil {
|
||||
|
||||
@ -104,6 +104,7 @@ type LLMSource struct {
|
||||
Adapter string `json:"adapter"`
|
||||
AdapterPath string `json:"adapter_path,omitempty"`
|
||||
ContextWindow int `json:"context_window,omitempty"`
|
||||
MaxConcurrent int `json:"max_concurrent,omitempty"`
|
||||
ThinkingEnabled bool `json:"thinking_enabled,omitempty"`
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user