diff --git a/internal/agent/api/provider.go b/internal/agent/api/provider.go index 338d604..4d764b3 100644 --- a/internal/agent/api/provider.go +++ b/internal/agent/api/provider.go @@ -1,6 +1,7 @@ package api import ( + "bufio" "context" "encoding/json" "fmt" @@ -81,10 +82,11 @@ func (r *CompletionRequest) MarshalJSON() ([]byte, error) { } type CompletionResponse struct { - Content string `json:"content"` - FinishReason string `json:"finish_reason,omitempty"` - TokenUsage TokenUsage `json:"token_usage,omitempty"` - ToolCalls []ToolCall `json:"tool_calls,omitempty"` + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content,omitempty"` + FinishReason string `json:"finish_reason,omitempty"` + TokenUsage TokenUsage `json:"token_usage,omitempty"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` } type TokenUsage struct { @@ -349,68 +351,202 @@ func (p *OllamaProvider) ChatStream(ctx context.Context, req *CompletionRequest) return ch, nil } +// LuaAdaptedProvider 使用 Lua 脚本做请求/响应变换,直接发起 HTTP 调用 +// 不再包裹其他 Provider,协议差异全部在 Lua 层处理 type LuaAdaptedProvider struct { - name string - base Provider - vm *luaVM.VM - adapter string + name string + cfg BaseConfig + vm *luaVM.VM + adapter string + client *http.Client } -func NewLuaAdaptedProvider(base Provider, vm *luaVM.VM, adapter string) *LuaAdaptedProvider { +func NewLuaAdaptedProvider(cfg BaseConfig, vm *luaVM.VM, adapter string) *LuaAdaptedProvider { + if cfg.Temperature == 0 { + cfg.Temperature = 0.7 + } + if cfg.MaxTokens == 0 { + cfg.MaxTokens = 4096 + } return &LuaAdaptedProvider{ name: fmt.Sprintf("lua_%s", adapter), - base: base, + cfg: cfg, vm: vm, adapter: adapter, + client: &http.Client{Timeout: 120 * time.Second}, } } func (p *LuaAdaptedProvider) Name() string { return p.name } func (p *LuaAdaptedProvider) Chat(ctx context.Context, req *CompletionRequest) (*CompletionResponse, error) { - inputMap := map[string]interface{}{ - "model": req.Model, - "messages": messagesToMap(req.Messages), - "temperature": req.Temperature, - "max_tokens": req.MaxTokens, - "stream": false, + if req.Model == "" { + req.Model = p.cfg.Model } - transformed, err := p.vm.CallTransform(p.adapter, inputMap) + rawReq, _ := json.Marshal(req) + + transformedBody, err := p.vm.CallTransformRequest(p.adapter, string(rawReq)) if err != nil { - return nil, fmt.Errorf("lua transform: %w", err) + return nil, fmt.Errorf("lua transform_request: %w", err) } - transformedReq := &CompletionRequest{ - Model: getString(transformed, "model"), - Temperature: getFloat(transformed, "temperature"), - MaxTokens: int(getFloat(transformed, "max_tokens")), - Stream: false, + endpoint := p.vm.GetAdapterEndpoint(p.adapter) + if endpoint == "" { + endpoint = "/chat/completions" } + url := strings.TrimRight(p.cfg.BaseURL, "/") + endpoint - if msgs, ok := transformed["messages"].([]interface{}); ok { - for _, m := range msgs { - if mm, ok := m.(map[string]interface{}); ok { - transformedReq.Messages = append(transformedReq.Messages, Message{ - Role: getString(mm, "role"), - Content: getString(mm, "content"), - }) - } - } - } - - resp, err := p.base.Chat(ctx, transformedReq) + httpReq, err := http.NewRequestWithContext(ctx, "POST", url, strings.NewReader(transformedBody)) if err != nil { - return nil, err + 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) } - return resp, nil + resp, err := p.client.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("api call: %w", err) + } + defer resp.Body.Close() + + rawResp, err := io.ReadAll(resp.Body) + if err != nil { + return nil, fmt.Errorf("read response: %w", err) + } + + if resp.StatusCode != 200 { + return nil, fmt.Errorf("api error %d: %s", resp.StatusCode, string(rawResp)) + } + + unifiedJSON, err := p.vm.CallTransformResponse(p.adapter, string(rawResp)) + if err != nil { + return nil, fmt.Errorf("lua transform_response: %w", err) + } + + var result CompletionResponse + if err := json.Unmarshal([]byte(unifiedJSON), &result); err != nil { + return nil, fmt.Errorf("unmarshal unified response: %w (body: %s)", err, unifiedJSON) + } + + return &result, nil } func (p *LuaAdaptedProvider) ChatStream(ctx context.Context, req *CompletionRequest) (<-chan StreamChunk, error) { - return p.base.ChatStream(ctx, req) + if req.Model == "" { + req.Model = p.cfg.Model + } + req.Stream = true + rawReq, _ := json.Marshal(req) + + transformedBody, err := p.vm.CallTransformRequest(p.adapter, string(rawReq)) + if err != nil { + return nil, fmt.Errorf("lua transform_request (stream): %w", err) + } + + endpoint := p.vm.GetAdapterEndpoint(p.adapter) + if endpoint == "" { + endpoint = "/chat/completions" + } + url := strings.TrimRight(p.cfg.BaseURL, "/") + endpoint + + httpReq, err := http.NewRequestWithContext(ctx, "POST", url, strings.NewReader(transformedBody)) + if err != nil { + return nil, fmt.Errorf("create stream 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) + } + + resp, err := p.client.Do(httpReq) + if err != nil { + return nil, fmt.Errorf("stream api: %w", err) + } + + ch := make(chan StreamChunk, 64) + go func() { + defer resp.Body.Close() + defer close(ch) + + scanner := NewSSEScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + if line == "" { + continue + } + + // 尝试用 Lua 变换流块(如果 adapter 定义了 transform_stream_chunk) + unified, err := p.vm.CallTransformStreamChunk(p.adapter, line) + if err != nil || unified == line { + // 无流变换函数或变换透传,尝试标准 SSE 解析 + var raw struct { + Choices []struct { + Delta struct { + Content string `json:"content"` + } `json:"delta"` + FinishReason *string `json:"finish_reason"` + } `json:"choices"` + } + if err := json.Unmarshal([]byte(unified), &raw); err != nil { + continue + } + if len(raw.Choices) > 0 { + ch <- StreamChunk{ + Content: raw.Choices[0].Delta.Content, + Done: raw.Choices[0].FinishReason != nil, + } + } + continue + } + + // Lua 返回了变换后的统一格式 + var chunk StreamChunk + if err := json.Unmarshal([]byte(unified), &chunk); err == nil { + ch <- chunk + } + } + }() + + return ch, nil } +// SSEScanner 读取 SSE 格式的流(data: ...) +type SSEScanner struct { + reader *bufio.Reader + pending string +} + +func NewSSEScanner(r io.Reader) *SSEScanner { + return &SSEScanner{reader: bufio.NewReader(r)} +} + +func (s *SSEScanner) Scan() bool { + s.pending = "" + for { + line, err := s.reader.ReadString('\n') + if err != nil { + return false + } + line = strings.TrimRight(line, "\r\n") + if strings.HasPrefix(line, "data: ") { + s.pending = strings.TrimPrefix(line, "data: ") + if s.pending == "[DONE]" { + return false + } + return true + } + } +} + +func (s *SSEScanner) Text() string { return s.pending } + type ProviderManager struct { mu sync.RWMutex providers map[string]Provider diff --git a/internal/lua/adapters/deepseek.lua b/internal/lua/adapters/deepseek.lua index 3f0362f..0675dd1 100644 --- a/internal/lua/adapters/deepseek.lua +++ b/internal/lua/adapters/deepseek.lua @@ -1,22 +1,74 @@ local adapter = {} adapter.name = "deepseek" -adapter.version = "1.0.0" +adapter.version = "2.0.0" +adapter.endpoint = "/chat/completions" +adapter.headers = {} -function adapter.transform_request(input) - local messages = input.messages or {} - local result = { - model = input.model or "deepseek-chat", - messages = messages, - temperature = input.temperature or 0.0, - max_tokens = input.max_tokens or 4096, - stream = input.stream or false - } - return result +-- DeepSeek 格式与 OpenAI 兼容,只需要强制 temperature=0(禁用 thinking) +function adapter.transform_request(raw_body) + local ok, req = pcall(json.decode, raw_body) + if not ok then return raw_body end + req.model = req.model or "deepseek-chat" + req.temperature = 0.0 + req.stream = req.stream or false + return json.encode(req) end -function adapter.transform_response(raw) - return raw +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if resp.usage then + unified.token_usage.prompt = resp.usage.prompt_tokens or 0 + unified.token_usage.completion = resp.usage.completion_tokens or 0 + unified.token_usage.total = resp.usage.total_tokens or 0 + end + + if resp.choices and #resp.choices > 0 then + local ch = resp.choices[1] + if ch.message then + unified.content = ch.message.content or "" + if ch.message.reasoning_content then + unified.reasoning_content = ch.message.reasoning_content + end + if ch.message.tool_calls then + local tcs = {} + for _, tc in ipairs(ch.message.tool_calls) do + local args_ok, args = pcall(json.decode, tc["function"].arguments) + if not args_ok then args = {} end + table.insert(tcs, { + id = tc.id, + type = tc.type or "function", + name = tc["function"].name, + arguments = args + }) + end + unified.tool_calls = tcs + end + end + unified.finish_reason = ch.finish_reason or "" + end + + return json.encode(unified) +end + +function adapter.transform_stream_chunk(raw_chunk) + local ok, chunk = pcall(json.decode, raw_chunk) + if not ok then return "" end + if not chunk.choices or #chunk.choices == 0 then return "" end + local delta = chunk.choices[1].delta or {} + local fr = chunk.choices[1].finish_reason + return json.encode({ + content = delta.content or "", + done = (fr ~= nil) + }) end return adapter diff --git a/internal/lua/adapters/ollama.lua b/internal/lua/adapters/ollama.lua index ff7ce94..a931045 100644 --- a/internal/lua/adapters/ollama.lua +++ b/internal/lua/adapters/ollama.lua @@ -1,24 +1,63 @@ local adapter = {} adapter.name = "ollama" -adapter.version = "1.0.0" +adapter.version = "2.0.0" +adapter.endpoint = "/api/chat" +adapter.headers = {} -function adapter.transform_request(input) - local messages = input.messages or {} - local result = { - model = input.model or "llama3", - messages = messages, - stream = input.stream or false, +-- Ollama API 格式:{ model, messages, stream, options:{temperature,num_predict} } +function adapter.transform_request(raw_body) + local ok, req = pcall(json.decode, raw_body) + if not ok then return raw_body end + + local ollama_req = { + model = req.model or "llama3", + stream = req.stream or false, options = { - temperature = input.temperature or 0.7, - num_predict = input.max_tokens or 2048 + temperature = req.temperature or 0.7, + num_predict = req.max_tokens or 2048 } } - return result + + -- 转换 messages 格式(Ollama 兼容 OpenAI 的 messages 格式) + if req.messages then + local msgs = {} + for _, m in ipairs(req.messages) do + table.insert(msgs, { role = m.role, content = m.content }) + end + ollama_req.messages = msgs + end + + return json.encode(ollama_req) end -function adapter.transform_response(raw) - return raw +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok then return raw_body end + + local unified = { + content = "", + finish_reason = resp.done_reason or "", + tool_calls = {}, + usage = { prompt = 0, completion = 0, total = 0 } + } + + if resp.message then + unified.content = resp.message.content or "" + end + + return json.encode(unified) +end + +function adapter.transform_stream_chunk(raw_chunk) + local ok, chunk = pcall(json.decode, raw_chunk) + if not ok then return "" end + if not chunk.message then return "" end + + return json.encode({ + content = chunk.message.content or "", + done = chunk.done or false + }) end return adapter diff --git a/internal/lua/adapters/openai.lua b/internal/lua/adapters/openai.lua index ef4a9d3..288662f 100644 --- a/internal/lua/adapters/openai.lua +++ b/internal/lua/adapters/openai.lua @@ -1,22 +1,76 @@ local adapter = {} adapter.name = "openai" -adapter.version = "1.0.0" +adapter.version = "2.0.0" +adapter.endpoint = "/chat/completions" +adapter.headers = {} -function adapter.transform_request(input) - local messages = input.messages or {} - local result = { - model = input.model or "gpt-4", - messages = messages, - temperature = input.temperature or 0.7, - max_tokens = input.max_tokens or 2048, - stream = input.stream or false - } - return result +-- raw_body: JSON string as received from Go (CompletionRequest marshalled) +-- return: transformed JSON string to send to API +function adapter.transform_request(raw_body) + return raw_body end -function adapter.transform_response(raw) - return raw +-- raw_body: JSON string from HTTP response body +-- return: unified JSON string in CompletionResponse format +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if resp.usage then + unified.token_usage.prompt = resp.usage.prompt_tokens or 0 + unified.token_usage.completion = resp.usage.completion_tokens or 0 + unified.token_usage.total = resp.usage.total_tokens or 0 + end + + if resp.choices and #resp.choices > 0 then + local ch = resp.choices[1] + if ch.message then + unified.content = ch.message.content or "" + if ch.message.reasoning_content then + unified.reasoning_content = ch.message.reasoning_content + end + if ch.message.tool_calls then + local tcs = {} + for _, tc in ipairs(ch.message.tool_calls) do + local args_ok, args = pcall(json.decode, tc["function"].arguments) + if not args_ok then args = {} end + table.insert(tcs, { + id = tc.id, + type = tc.type or "function", + name = tc["function"].name, + arguments = args + }) + end + unified.tool_calls = tcs + end + end + unified.finish_reason = ch.finish_reason or "" + end + + return json.encode(unified) +end + +-- raw_chunk: single SSE data line (after "data: " prefix) +-- return: unified chunk JSON string, or "" to skip +function adapter.transform_stream_chunk(raw_chunk) + local ok, chunk = pcall(json.decode, raw_chunk) + if not ok then return "" end + + if not chunk.choices or #chunk.choices == 0 then return "" end + local delta = chunk.choices[1].delta or {} + local fr = chunk.choices[1].finish_reason + + return json.encode({ + content = delta.content or "", + done = (fr ~= nil) + }) end return adapter diff --git a/internal/lua/vm.go b/internal/lua/vm.go index 0b0efcc..caf971d 100644 --- a/internal/lua/vm.go +++ b/internal/lua/vm.go @@ -2,6 +2,7 @@ package lua import ( "embed" + "encoding/json" "fmt" "os" "path/filepath" @@ -53,9 +54,27 @@ func (v *VM) Start() error { return 0 })) - v.state.SetGlobal("json_encode", v.state.NewFunction(func(L *lua.LState) int { + jsonTable := v.state.NewTable() + v.state.SetGlobal("json", jsonTable) + v.state.SetField(jsonTable, "encode", v.state.NewFunction(func(L *lua.LState) int { val := L.CheckAny(1) - L.Push(lua.LString(fmt.Sprintf("%v", val))) + 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 + })) + v.state.SetField(jsonTable, "decode", v.state.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 })) @@ -143,13 +162,13 @@ func (v *VM) LoadAdapter(path string) error { return nil } -func (v *VM) CallTransform(name string, input map[string]interface{}) (map[string]interface{}, error) { +func (v *VM) CallTransformRequest(name, rawJSON string) (string, error) { v.mu.Lock() adapter, ok := v.loaded[name] v.mu.Unlock() if !ok { - return nil, fmt.Errorf("adapter %s not loaded", name) + return "", fmt.Errorf("adapter %s not loaded", name) } v.mu.Lock() @@ -157,35 +176,29 @@ func (v *VM) CallTransform(name string, input map[string]interface{}) (map[strin fn := adapter.RawGetString("transform_request") if fn == nil { - return nil, fmt.Errorf("adapter %s missing transform_request", name) + return "", fmt.Errorf("adapter %s missing transform_request", name) } - inputTable := mapToTable(v.state, input) v.state.Push(fn) - v.state.Push(inputTable) + v.state.Push(lua.LString(rawJSON)) if err := v.state.PCall(1, 1, nil); err != nil { - return nil, fmt.Errorf("transform_request: %w", err) + return "", fmt.Errorf("transform_request: %w", err) } result := v.state.Get(-1) v.state.Pop(1) - resultTable, ok := result.(*lua.LTable) - if !ok { - return nil, fmt.Errorf("transform_request must return a table") - } - - return tableToMap(resultTable), nil + return result.String(), nil } -func (v *VM) CallResponseTransform(name string, raw []byte) ([]byte, error) { +func (v *VM) CallTransformResponse(name, rawJSON string) (string, error) { v.mu.Lock() adapter, ok := v.loaded[name] v.mu.Unlock() if !ok { - return raw, nil + return rawJSON, nil } v.mu.Lock() @@ -193,20 +206,94 @@ func (v *VM) CallResponseTransform(name string, raw []byte) ([]byte, error) { fn := adapter.RawGetString("transform_response") if fn == nil { - return raw, nil + return rawJSON, nil } v.state.Push(fn) - v.state.Push(lua.LString(string(raw))) + v.state.Push(lua.LString(rawJSON)) if err := v.state.PCall(1, 1, nil); err != nil { - return nil, fmt.Errorf("transform_response: %w", err) + return "", fmt.Errorf("transform_response: %w", err) } result := v.state.Get(-1) v.state.Pop(1) - return []byte(result.String()), nil + return result.String(), nil +} + +func (v *VM) CallTransformStreamChunk(name, rawLine string) (string, error) { + v.mu.Lock() + adapter, ok := v.loaded[name] + v.mu.Unlock() + + if !ok { + return rawLine, nil + } + + v.mu.Lock() + defer v.mu.Unlock() + + fn := adapter.RawGetString("transform_stream_chunk") + if fn == nil { + return rawLine, nil + } + + v.state.Push(fn) + v.state.Push(lua.LString(rawLine)) + + if err := v.state.PCall(1, 1, nil); err != nil { + return "", fmt.Errorf("transform_stream_chunk: %w", err) + } + + result := v.state.Get(-1) + v.state.Pop(1) + + if result.String() == "" { + return "", nil + } + return result.String(), nil +} + +func (v *VM) GetAdapterEndpoint(name string) string { + v.mu.Lock() + adapter, ok := v.loaded[name] + v.mu.Unlock() + + if !ok { + return "" + } + + v.mu.Lock() + defer v.mu.Unlock() + + if ep := adapter.RawGetString("endpoint"); ep != nil { + return ep.String() + } + return "" +} + +func (v *VM) GetAdapterHeaders(name string) map[string]string { + v.mu.Lock() + adapter, ok := v.loaded[name] + v.mu.Unlock() + + if !ok { + return nil + } + + v.mu.Lock() + defer v.mu.Unlock() + + 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 } func (v *VM) ListAdapters() []APIAdapter { @@ -237,35 +324,61 @@ func (v *VM) ReloadAll() error { return v.Start() } -func mapToTable(L *lua.LState, m map[string]interface{}) *lua.LTable { - tbl := L.NewTable() - for k, v := range m { - switch val := v.(type) { - case string: - tbl.RawSetString(k, lua.LString(val)) - case float64: - tbl.RawSetString(k, lua.LNumber(val)) - case int: - tbl.RawSetString(k, lua.LNumber(val)) - case bool: - tbl.RawSetString(k, lua.LBool(val)) - case map[string]interface{}: - tbl.RawSetString(k, mapToTable(L, val)) - case []interface{}: - arr := L.NewTable() - for i, item := range val { - if m, ok := item.(map[string]interface{}); ok { - arr.RawSetInt(i+1, mapToTable(L, m)) - } else { - arr.RawSetInt(i+1, lua.LString(fmt.Sprintf("%v", item))) - } - } - tbl.RawSetString(k, arr) - default: - tbl.RawSetString(k, lua.LString(fmt.Sprintf("%v", v))) +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: + return lua.LNil } - return tbl } func (v *VM) writeBundledAdapters() error { @@ -296,32 +409,4 @@ func (v *VM) writeBundledAdapters() error { return nil } -func tableToMap(tbl *lua.LTable) map[string]interface{} { - result := make(map[string]interface{}) - tbl.ForEach(func(key lua.LValue, val lua.LValue) { - k := key.String() - switch v := val.(type) { - case lua.LString: - result[k] = string(v) - case lua.LNumber: - result[k] = float64(v) - case lua.LBool: - result[k] = bool(v) - case *lua.LTable: - if v.MaxN() == 0 { - result[k] = tableToMap(v) - } else { - var arr []interface{} - v.ForEach(func(_, item lua.LValue) { - if tbl, ok := item.(*lua.LTable); ok { - arr = append(arr, tableToMap(tbl)) - } else { - arr = append(arr, item.String()) - } - }) - result[k] = arr - } - } - }) - return result -} +