Files
HomeAgent/internal/lua/vm.go
JianFeeeee 7d6c0bb90b feat(provider): complete ChatStream with llmsproxy-grade streaming
Rewrite LuaAdaptedProvider.ChatStream to match the maturity of
llmsproxy's streaming implementation:

HTTP layer:
  - Dedicated stream HTTP client with no overall timeout (SSE must not
    be cut by the 180s Chat timeout); only a 30s dial timeout
  - Uses applyAdapterHeaders (supports build_headers dynamic signing
    hook), matching the non-streaming Chat path

Non-200 response handling:
  - New TransformError Lua hook (adapter.transform_error) for per-source
    protocol knowledge in error messages
  - Safe fallback truncation of raw error bodies (prevents HTML dump
    leakage to clients)

SSE parsing enhancements:
  - parseOpenAICompatibleStreamChunkFull: handles token usage in the
    final chunk (prompt_tokens/prompt, total_tokens/total dual keys),
    prompt cache detail fields, and empty-string finish_reason filtering
    (sensenova sends "" on every chunk)
  - Replaced old SSEScanner with bufio.Scanner (larger buffer, fewer
    allocations)

Stream integrity:
  - errorOnlyChunk detection: holds back the first chunk to reject
    degenerate streams (e.g. zen free pool's finish_reason:"network_error"
    with empty content) before any byte reaches the caller
  - [DONE] dedup: adapters that already emit a terminating done chunk
    with the real finish_reason don't get a second reason-less done
  - Clean EOF sends a final Done:true if no done was seen

Struct changes:
  - StreamChunk: added FinishReason and Usage fields for callers
  - LuaAdaptedProvider: added streamClient (lazy) + streamMu

Tested: curl against llmsproxy SSE confirms reasoning_content parsing
is correct (delta.reasoning_content), usage chunk handling works, and
[DONE] termination is properly emitted.
2026-08-25 08:30:53 +08:00

680 lines
16 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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"
lua "github.com/yuin/gopher-lua"
)
//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"`
}
// worker 封装一个独立的 gopher-lua 解释器。LState 非并发安全,
// 同一时刻只能被一个 goroutine 使用(从池中取出)。
type worker struct {
L *lua.LState
}
func (w *worker) close() {
if w.L != nil {
w.L.Close()
w.L = nil
}
}
// 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
}
// TransformError 执行可选的 adapter.transform_error(status, body) 钩子。
// 返回 reason: 适配器提取的错误原因ok=false 表示适配器未定义此钩子。
func (v *VM) TransformError(name string, status int, body string) (reason string, ok bool, err error) {
p := v.pool(name)
if p == nil {
return "", false, nil
}
w, err := p.acquire()
if err != nil {
return "", false, err
}
defer p.release(w)
L := w.L
adapter := L.GetGlobal(adapterGlobal)
tbl, ok2 := adapter.(*lua.LTable)
if !ok2 {
return "", false, nil
}
f := tbl.RawGetString("transform_error")
if _, ok := f.(*lua.LFunction); !ok {
return "", false, nil
}
L.Push(f)
L.Push(lua.LNumber(status))
L.Push(lua.LString(body))
if err := L.PCall(2, 1, nil); err != nil {
return "", false, fmt.Errorf("transform_error: %w", err)
}
res := L.Get(-1)
L.Pop(1)
if res.Type() != lua.LTString {
return "", false, nil
}
return res.String(), true, 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 := 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)
if err != nil {
L.Push(lua.LString("null"))
return 1
}
L.Push(lua.LString(string(b)))
return 1
}))
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 {
L.Push(lua.LNil)
return 1
}
L.Push(goValueToLua(L, val))
return 1
}))
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)
if err != nil {
errJSON, _ := json.Marshal(map[string]interface{}{"url": url, "error": err.Error()})
L.Push(lua.LString(string(errJSON)))
return 1
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
result, _ := json.Marshal(map[string]interface{}{"url": url, "status": resp.StatusCode, "body": string(body)})
L.Push(lua.LString(string(result)))
return 1
}))
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}
resp, err := client.Post(url, "application/json", strings.NewReader(body))
if err != nil {
errJSON, _ := json.Marshal(map[string]interface{}{"url": url, "error": err.Error()})
L.Push(lua.LString(string(errJSON)))
return 1
}
defer resp.Body.Close()
respBody, _ := io.ReadAll(resp.Body)
result, _ := json.Marshal(map[string]interface{}{"url": url, "body": string(respBody), "status": resp.StatusCode})
L.Push(lua.LString(string(result)))
return 1
}))
// 签名辅助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", "kimicode",
"server",
}
for _, name := range known {
srcPath := "adapters/" + name + ".lua"
dstPath := filepath.Join(v.dir, name+".lua")
if _, err := os.Stat(dstPath); err == nil {
continue
}
data, err := bundledAdapters.ReadFile(srcPath)
if err != nil {
continue
}
if err := os.WriteFile(dstPath, data, 0644); err != nil {
return fmt.Errorf("write %s: %w", name+".lua", err)
}
fmt.Printf("[lua] installed bundled adapter: %s\n", name+".lua")
}
return nil
}
func luaValueToGo(lv lua.LValue) interface{} {
switch v := lv.(type) {
case lua.LString:
return string(v)
case lua.LNumber:
return float64(v)
case lua.LBool:
return bool(v)
case *lua.LTable:
if v.MaxN() > 0 {
arr := make([]interface{}, 0, v.MaxN())
v.ForEach(func(_, val lua.LValue) {
arr = append(arr, luaValueToGo(val))
})
return arr
}
m := make(map[string]interface{})
v.ForEach(func(key, val lua.LValue) {
m[key.String()] = luaValueToGo(val)
})
return m
default:
return nil
}
}
func goValueToLua(L *lua.LState, val interface{}) lua.LValue {
switch v := val.(type) {
case string:
return lua.LString(v)
case float64:
return lua.LNumber(v)
case int:
return lua.LNumber(v)
case int64:
return lua.LNumber(v)
case bool:
return lua.LBool(v)
case nil:
return lua.LNil
case []interface{}:
tbl := L.NewTable()
for i, item := range v {
tbl.RawSetInt(i+1, goValueToLua(L, item))
}
return tbl
case map[string]interface{}:
tbl := L.NewTable()
for k, item := range v {
tbl.RawSetString(k, goValueToLua(L, item))
}
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
}
}