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:
root
2026-08-10 11:29:57 +08:00
parent 6d64543918
commit f0dacef281
7 changed files with 657 additions and 266 deletions

View File

@ -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)
}

View File

@ -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 {

View File

@ -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),
})
}

View File

@ -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
View 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)
}
}

View File

@ -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 {

View File

@ -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"`
}