diff --git a/cmd/homed/main.go b/cmd/homed/main.go index 04cc3be..64c0075 100644 --- a/cmd/homed/main.go +++ b/cmd/homed/main.go @@ -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) } diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go index 4b378ef..439677f 100644 --- a/internal/agent/api/provider.go +++ b/internal/agent/api/provider.go @@ -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 { diff --git a/internal/config/registry.go b/internal/config/registry.go index 1ba35a4..110e2a3 100644 --- a/internal/config/registry.go +++ b/internal/config/registry.go @@ -185,6 +185,7 @@ var sourceFieldDefs = []struct { {"thinking_enabled", "bool", "深度思考"}, {"adapter", "string", "适配器"}, {"adapter_path", "string", "适配器路径"}, + {"max_concurrent", "int", "并发上限"}, } // registerSourceDefs 注册 core.llm.sources..* 的 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), }) } diff --git a/internal/lua/vm.go b/internal/lua/vm.go index f3d4ea9..861c149 100644 --- a/internal/lua/vm.go +++ b/internal/lua/vm.go @@ -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 } } diff --git a/internal/lua/vm_test.go b/internal/lua/vm_test.go new file mode 100644 index 0000000..763d437 --- /dev/null +++ b/internal/lua/vm_test.go @@ -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) + } +} \ No newline at end of file diff --git a/internal/sdk/llm_impl.go b/internal/sdk/llm_impl.go index 7d8c443..520c90f 100644 --- a/internal/sdk/llm_impl.go +++ b/internal/sdk/llm_impl.go @@ -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 { diff --git a/pkg/types/types.go b/pkg/types/types.go index aacec97..266ca94 100644 --- a/pkg/types/types.go +++ b/pkg/types/types.go @@ -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"` }