diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..28f49e0 --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +/bin/ +/llmsproxy +/config.yaml +/adapters/ +*.log \ No newline at end of file diff --git a/README.md b/README.md index f35651a..1a50352 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,139 @@ -# ModelRouter +# llmsproxy -面向多llms订阅者,部署在内网,实现一次配置多个服务共同使用的效果。支持指定模型或auto模式,按照配置的优先级选择可用模型提供服务 \ No newline at end of file +统一 OpenAI 兼容网关:把多个上游 LLM 源(DeepSeek、OpenAI、Anthropic、Gemini、 +Groq、Mistral、Ollama、KimiCode…)通过 **Lua 适配器** 做协议转换,对内网暴露 +一个标准 `OpenAI Chat Completions` 接口(`/v1/chat/completions` + `/v1/models`), +支持单次结算与 SSE 流式。 + +从 [HomeAgent](https://gitcode.com/JianFeeeee/HomeAgent) 的多源 LLM 适配层 +(`internal/agent/api/provider.go` + `internal/lua/adapters/*`)抽离而来并独立演进。 + +## 特性 + +- **多源**:一个进程内配置任意多个上游源,按请求的 `model` 自动路由。 +- **统一输出**:所有源对外都是 OpenAI 格式(含 `reasoning_content`、`tool_calls`、`usage`)。 +- **Lua 适配协议**:每个源挂一个 `.lua` 适配器,完成 `transform_request` / + `transform_response` / `transform_stream_chunk` 双向转换,协议差异全在 Lua 层。 +- **签名 / 请求头钩子**:适配器可定义 `build_headers(meta)`,在 Go 发 HTTP 前 + 动态注入/签名请求头——用于云端 API 校验调用方 app(如 KimiCode 只放行特定 + agent)。提供 `hmac_sha256_hex` / `sha256_hex` / `base64_encode` 等签名辅助。 +- **鉴权**:网关自身用 `gateway_keys` 校验客户端 Bearer key;与上游各自的 key 相互独立。 +- **流式**:SSE `chat.completion.chunk`,含角色首包与 `[DONE]` 收尾。 + +## 快速开始 + +```bash +cp config.example.yaml config.yaml # 编辑你的源与 key +GOMODCACHE=... GOPROXY=off go build -o llmsproxy ./cmd/llmsproxy +./llmsproxy -config config.yaml +``` + +```bash +# 无 key -> 401 +curl http://127.0.0.1:8080/v1/models + +# 单次 +curl -H "Authorization: Bearer sk-gw-local-0001" \ + -d '{"model":"deepseek-v4-flash","messages":[{"role":"user","content":"hi"}]}' \ + http://127.0.0.1:8080/v1/chat/completions + +# 流式 +curl -N -H "Authorization: Bearer sk-gw-local-0001" \ + -d '{"model":"deepseek-v4-flash","stream":true,"messages":[{"role":"user","content":"hi"}]}' \ + http://127.0.0.1:8080/v1/chat/completions +``` + +任何 OpenAI SDK 把 `base_url` 指到网关地址、`api_key` 用 `gateway_keys` 之一即可。 + +## 配置 + +见 [`config.example.yaml`](config.example.yaml)。核心字段: + +```yaml +listen: 127.0.0.1:8080 # 网关监听地址(建议绑内网/回环) +gateway_keys: [sk-gw-0001] # 客户端访问网关的 key;留空=不鉴权 +default_source: deepseek # model 无法路由时回落的源 +adapter_dir: adapters # Lua 适配目录,首启自动写入内置适配器 + +sources: + - name: deepseek + base_url: https://api.deepseek.com + api_key: sk-... + model: deepseek-v4-flash + adapter: deepseek + # 静态请求头(比适配器默认优先) + headers: { X-Tenant: prod } + # 传给 Lua build_headers 的透传元数据 + meta: { app_id: x, app_secret: y } + temperature: 0.7 + max_tokens: 4096 + timeout: 120s # 请求超时,默认 120s +``` + +### 模型路由 + +`/v1/chat/completions` 的 `model` 解析顺序: + +1. `source/model` 或 `source:model` 前缀 → 指定源; +2. 精确匹配某个源的 `model`; +3. 回落到 `default_source`。 + +任何 OpenAI 客户端,只要 `model` 设为某个源的 `name/任意名`,即可锁定走该源。 + +## Lua 适配器协议 + +每个适配器是一个返回 table 的 Lua 脚本(`internal/lua/adapters/.lua`), +加载时可被 `adapter_dir` 下的同名脚本覆盖。 + +```lua +return { + name = "mysrc", + version = "1.0.0", + endpoint = "/chat/completions", -- 上游路径(可被 source.endpoint 覆盖) + headers = { ["X-Static"] = "v" }, -- 静态默认请求头(build_headers 缺省时使用) + + -- 请求转换:把 OpenAI 格式 req 转成上游原生格式,返回字符串 + transform_request = function(raw_json) ... end, + + -- 响应转换:把上游原生响应转成统一格式字符串 + -- { content, reasoning_content, finish_reason, token_usage{...}, tool_calls[{...}] } + transform_response = function(raw_json) ... end, + + -- 流式分块转换:把上游 SSE data 转成 { content, done, ... },返回 "" 则跳过 + transform_stream_chunk = function(raw_chunk) ... end, + + -- [可选] 动态请求头/签名钩子 + -- meta = { url, method, body, api_key, timestamp, source={ name, meta={...} } } + build_headers = function(meta) return { ["X-App-Sign"] = sign } end, +} +``` + +内置信号辅助:`hmac_sha256_hex(key, data)`、`sha256_hex(data)`、`base64_encode(s)`、 +`tohex(s)`、`json.encode/decode`、`log(level, msg)`。 + +### 内置适配器 + +`openai` `deepseek` `anthropic` `gemini` `github` `groq` `mistral` `ollama` `kimicode`。 + +**kimicode** 是展示 `build_headers` 的样例:云端校验调用方 app,需要按 +`meta.app_secret` 对时间戳+URL+请求体哈希做 HMAC 签名并附 `X-App-Sign` 等头。 +配好 `sources[].meta.{app_id, app_secret, app_agent}` 即可。 + +## 目录 + +``` +cmd/llmsproxy # 入口 +internal/config # YAML 配置加载/校验 +internal/lua # Lua VM + AdapterCache + 内置适配器 (embed) +internal/provider # Provider(HTTP) + Registry(路由) +internal/gateway # OpenAI 兼容 HTTP 服务 + 鉴权 + SDK/流式 +internal/types # 统一格式 & OpenAI wire 类型 +``` + +## 测试 + +```bash +go test ./... +``` + +覆盖:配置校验、适配器加载/变换、签名钩子、Gateway 鉴权、SDK 结算、SSE 流式、模型路由。 \ No newline at end of file diff --git a/cmd/llmsproxy/main.go b/cmd/llmsproxy/main.go new file mode 100644 index 0000000..ea551f0 --- /dev/null +++ b/cmd/llmsproxy/main.go @@ -0,0 +1,54 @@ +// Command llmsproxy is a standalone gateway that exposes multiple upstream LLM +// sources behind a unified OpenAI-compatible HTTP API. Protocol differences are +// handled by Lua adapters (per source), optionally signing outgoing requests. +// It supports explicit model selection or AUTO routing by priority, image +// generation, multimodal payloads, a web UI for adapter/source management, and +// concurrent/queued scheduling with backoff and failover. +package main + +import ( + "flag" + "log" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "llmsproxy/internal/core" + "llmsproxy/internal/gateway" +) + +func main() { + cfgPath := flag.String("config", "config.yaml", "path to gateway config file") + flag.Parse() + + c, err := core.New(*cfgPath) + if err != nil { + log.Fatalf("[llmsproxy] core: %v", err) + } + defer c.Close() + + gw, err := gateway.New(c, c.GatewayKeys()) + if err != nil { + log.Fatalf("[llmsproxy] gateway: %v", err) + } + + srv := &http.Server{ + Addr: c.Listen(), + Handler: gw.Handler(), + ReadHeaderTimeout: 10 * time.Second, + } + go func() { + log.Printf("[llmsproxy] listening on %s (default_model=%s, models=%v, adapters=%d)", + c.Listen(), c.DefaultModel(), c.Registry().ModelList(), len(c.ListAdapters())) + if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed { + log.Fatalf("[llmsproxy] server: %v", err) + } + }() + + stop := make(chan os.Signal, 1) + signal.Notify(stop, syscall.SIGINT, syscall.SIGTERM) + <-stop + log.Printf("[llmsproxy] shutting down") +} \ No newline at end of file diff --git a/config.example.yaml b/config.example.yaml new file mode 100644 index 0000000..5001073 --- /dev/null +++ b/config.example.yaml @@ -0,0 +1,72 @@ +# llmsproxy — 统一 OpenAI 兼容网关配置 + +# 网关监听地址(默认 :8080,建议绑内网/回环) +listen: 127.0.0.1:8080 + +# 客户端访问本网关所需的 API Key(Bearer)。留空数组 = 不鉴权(仅内网)。 +gateway_keys: + - sk-gw-local-0001 + +# 默认模型选择:具体模型 id 或 AUTO(按各源模型的 priority 自动选最高可用源) +default_model: AUTO + +# Lua 适配器目录(默认 adapters/,首次启动自动写入内置适配器) +adapter_dir: adapters + +# 运行时持久化文件(WebUI 新增/编辑的源会写入此文件,重启后仍生效) +runtime_file: runtime.json + +# 全局并发上限(0 = 不限) +max_concurrent: 0 + +# ---- 上游 LLM 源列表 ---- +# 每个源对应一个 Lua 适配器(adapter);同一个适配器可被多个源复用。 +sources: + - name: deepseek + base_url: https://api.deepseek.com + api_key: sk-your-deepseek-key + adapter: deepseek + max_concurrent: 8 + models: + - id: deepseek-v4-flash + priority: 100 # 数字越大越优先被 AUTO 选中 + kind: chat + - id: deepseek-reasoner + priority: 60 + kind: chat + meta: { thinking: true } + + - name: ollama + base_url: http://127.0.0.1:11434 + api_key: "" + adapter: ollama + endpoint: /api/chat + models: + - id: llama3 + priority: 50 + kind: chat + + # 演示:KimiCode 校验调用方 app,通过 Lua build_headers 钩子做签名 + - name: kimicode + base_url: https://api.moonshot.cn + api_key: sk-your-kimi-key + adapter: kimicode + models: + - id: kimi-k2 + priority: 90 + kind: chat + meta: + app_id: your-app-id + app_secret: your-app-secret + app_agent: code-agent + api_key: sk-your-kimi-key + + # 生图示例:OpenAI 兼容生图源(flux/dall-e 等) + - name: imagegen + base_url: https://api.example.com + api_key: sk-image + adapter: openai + models: + - id: flux-1 + priority: 80 + kind: image \ No newline at end of file diff --git a/e2e/e2e_test.go b/e2e/e2e_test.go new file mode 100644 index 0000000..3afd391 --- /dev/null +++ b/e2e/e2e_test.go @@ -0,0 +1,288 @@ +// Package e2e runs the real llmsproxy binary against mocked upstreams over +// real HTTP: startup, auth, chat, streaming SSE, image generation, and +// AUTO failover when a higher-priority source fails. +package e2e + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net" + "net/http" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "testing" + "time" +) + +// mockUpstream is an OpenAI-compatible upstream that can be programmed to fail. +type mockUpstream struct { + mu sync.Mutex + hits int + fail bool + server *http.Server + baseURL string + done chan struct{} +} + +func newMockUpstream(t *testing.T) *mockUpstream { + m := &mockUpstream{done: make(chan struct{})} + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("listen: %v", err) + } + mux := http.NewServeMux() + mux.HandleFunc("/chat/completions", func(w http.ResponseWriter, r *http.Request) { + m.mu.Lock() + m.hits++ + fail := m.fail + m.mu.Unlock() + if fail { + w.WriteHeader(http.StatusInternalServerError) + fmt.Fprint(w, `{"error":"mock upstream down"}`) + return + } + body, _ := io.ReadAll(r.Body) + var req map[string]interface{} + _ = json.Unmarshal(body, &req) + if stream, _ := req["stream"].(bool); stream { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(200) + fmt.Fprintln(w, `data: {"choices":[{"delta":{"content":"Hel"}}]}`) + fmt.Fprintln(w, `data: {"choices":[{"delta":{"content":"lo"}}]}`) + fmt.Fprintln(w, `data: {"choices":[{"delta":{},"finish_reason":"stop"}]}`) + fmt.Fprintln(w, "data: [DONE]") + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(200) + fmt.Fprintf(w, `{"choices":[{"message":{"content":"pong"},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":1,"total_tokens":4}}`) + }) + mux.HandleFunc("/v1/images/generations", func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + fmt.Fprint(w, `{"created":1,"data":[{"b64_json":"QUJD"}]}`) + }) + m.server = &http.Server{Handler: mux} + go m.server.Serve(l) + m.baseURL = "http://" + l.Addr().String() + t.Cleanup(func() { + close(m.done) + _ = m.server.Close() + }) + return m +} + +func (m *mockUpstream) Hits() int { + m.mu.Lock() + defer m.mu.Unlock() + return m.hits +} + +func (m *mockUpstream) SetFail(f bool) { + m.mu.Lock() + defer m.mu.Unlock() + m.fail = f +} + +// gatewayUnderTest is the real binary, started from a temp config. +type gatewayUnderTest struct { + cmd *exec.Cmd + addr string + key string + log bytes.Buffer +} + +func buildBinary(t *testing.T) string { + t.Helper() + dir := t.TempDir() + bin := filepath.Join(dir, "llmsproxy") + out, err := exec.Command("go", "build", "-o", bin, "llmsproxy/cmd/llmsproxy").CombinedOutput() + if err != nil { + t.Fatalf("build: %v\n%s", err, out) + } + return bin +} + +func startGateway(t *testing.T, bin, listen, cfgPath string) *gatewayUnderTest { + t.Helper() + g := &gatewayUnderTest{addr: listen, key: "sk-e2e-0001"} + g.cmd = exec.Command(bin, "-config", cfgPath) + g.cmd.Stdout = &g.log + g.cmd.Stderr = &g.log + if err := g.cmd.Start(); err != nil { + t.Fatalf("start: %v", err) + } + t.Cleanup(func() { + if g.cmd.Process != nil { + _ = g.cmd.Process.Kill() + _, _ = g.cmd.Process.Wait() + } + }) + // wait for readiness + deadline := time.Now().Add(10 * time.Second) + for time.Now().Before(deadline) { + req, _ := http.NewRequest("GET", "http://"+g.addr+"/v1/models", nil) + req.Header.Set("Authorization", "Bearer "+g.key) + resp, err := http.DefaultClient.Do(req) + if err == nil { + io.Copy(io.Discard, resp.Body) + resp.Body.Close() + if resp.StatusCode == http.StatusOK { + return g + } + } + if g.cmd.ProcessState != nil { + t.Fatalf("gateway exited early:\n%s", g.log.String()) + } + time.Sleep(100 * time.Millisecond) + } + t.Fatalf("gateway did not become ready:\n%s", g.log.String()) + return nil +} + +func (g *gatewayUnderTest) do(method, path string, body string, authed bool) (*http.Response, string) { + var r io.Reader + if body != "" { + r = strings.NewReader(body) + } + req, _ := http.NewRequest(method, "http://"+g.addr+path, r) + if authed { + req.Header.Set("Authorization", "Bearer "+g.key) + } + if body != "" { + req.Header.Set("Content-Type", "application/json") + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + return nil, "" + } + data, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return resp, string(data) +} + +// writeConfig writes a temp gateway config pointing at the mock upstreams, +// listening on the given address. +func writeConfig(t *testing.T, dir, listenAddr string, upstreams map[string]*mockUpstream) string { + t.Helper() + var sb strings.Builder + sb.WriteString("listen: " + listenAddr + "\n") + sb.WriteString("gateway_keys:\n - sk-e2e-0001\n") + sb.WriteString("default_model: AUTO\n") + sb.WriteString("adapter_dir: " + filepath.Join(dir, "adapters") + "\n") + sb.WriteString("runtime_file: " + filepath.Join(dir, "runtime.json") + "\n") + sb.WriteString("sources:\n") + order := []string{"good", "image"} + for _, name := range order { + u, ok := upstreams[name] + if !ok { + continue + } + if name == "image" { + sb.WriteString(" - name: imagegen\n base_url: " + u.baseURL + "\n adapter: openai\n models:\n - id: flux-1\n kind: image\n priority: 80\n") + } else { + sb.WriteString(" - name: " + name + "\n base_url: " + u.baseURL + "\n adapter: openai\n models:\n - id: " + name + "-m\n priority: 100\n") + } + } + path := filepath.Join(dir, "config.yaml") + if err := os.WriteFile(path, []byte(sb.String()), 0o644); err != nil { + t.Fatalf("write config: %v", err) + } + return path +} + +func TestEndToEnd(t *testing.T) { + upstreams := map[string]*mockUpstream{ + "good": newMockUpstream(t), + "image": newMockUpstream(t), + } + + dir := t.TempDir() + + // pick a fixed free port for the gateway + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatalf("pick port: %v", err) + } + addr := l.Addr().String() + l.Close() + + cfg := writeConfig(t, dir, addr, upstreams) + + bin := buildBinary(t) + g := startGateway(t, bin, addr, cfg) + + // 1. no key -> 401 + resp, body := g.do("GET", "/v1/models", "", false) + if resp == nil || resp.StatusCode != http.StatusUnauthorized { + t.Fatalf("unauth status = %v body=%q", statusOf(resp), body) + } + + // 2. list models + resp, body = g.do("GET", "/v1/models", "", true) + if resp == nil || resp.StatusCode != 200 { + t.Fatalf("models status=%v body=%q", statusOf(resp), body) + } + if !strings.Contains(body, "good-m") || !strings.Contains(body, "flux-1") { + t.Fatalf("models missing entries: %s", body) + } + + // 3. single chat (explicit model) + resp, body = g.do("POST", "/v1/chat/completions", + `{"model":"good-m","messages":[{"role":"user","content":"hi"}]}`, true) + if resp == nil || resp.StatusCode != 200 { + t.Fatalf("chat status=%v body=%q", statusOf(resp), body) + } + var chat map[string]interface{} + _ = json.Unmarshal([]byte(body), &chat) + msg := chat["choices"].([]interface{})[0].(map[string]interface{})["message"].(map[string]interface{}) + if msg["content"] != "pong" { + t.Fatalf("chat content = %q", msg["content"]) + } + + // 4. streaming SSE + resp, body = g.do("POST", "/v1/chat/completions", + `{"model":"good-m","stream":true,"messages":[{"role":"user","content":"hi"}]}`, true) + if resp == nil || resp.StatusCode != 200 { + t.Fatalf("stream status=%v body=%q", statusOf(resp), body) + } + if !strings.Contains(body, "Hel") || !strings.Contains(body, "lo") || !strings.Contains(body, "[DONE]") { + t.Fatalf("stream body = %q", body) + } + + // 5. image generation + resp, body = g.do("POST", "/v1/images/generations", + `{"model":"flux-1","prompt":"a cat"}`, true) + if resp == nil || resp.StatusCode != 200 { + t.Fatalf("image status=%v body=%q", statusOf(resp), body) + } + if !strings.Contains(body, "QUJD") { + t.Fatalf("image body = %q", body) + } + + // 6. AUTO failover: make the high-priority source fail, chat should still work + before := upstreams["good"].Hits() + upstreams["good"].SetFail(true) + resp, body = g.do("POST", "/v1/chat/completions", + `{"model":"AUTO","messages":[{"role":"user","content":"hi"}]}`, true) + upstreams["good"].SetFail(false) + if resp == nil || resp.StatusCode != 200 { + t.Fatalf("failover status=%v body=%q", statusOf(resp), body) + } + // the failing upstream must have received the attempt (proving fallback happened) + if upstreams["good"].Hits() <= before { + t.Fatalf("failover did not try the failing provider (hits %d -> %d)", before, upstreams["good"].Hits()) + } +} + +func statusOf(resp *http.Response) int { + if resp == nil { + return -1 + } + return resp.StatusCode +} + diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..1383e73 --- /dev/null +++ b/go.mod @@ -0,0 +1,8 @@ +module llmsproxy + +go 1.25.0 + +require ( + github.com/yuin/gopher-lua v1.1.2 + gopkg.in/yaml.v3 v3.0.1 +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..6bc42cd --- /dev/null +++ b/go.sum @@ -0,0 +1,6 @@ +github.com/yuin/gopher-lua v1.1.2 h1:yF/FjE3hD65tBbt0VXLE13HWS9h34fdzJmrWRXwobGA= +github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..eead8a3 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,124 @@ +// Package config loads the gateway YAML configuration plus a runtime overlay +// (web UI edits) and resolves them into sources with per-model priority. +package config + +import ( + "fmt" + "os" + "time" + + "gopkg.in/yaml.v3" +) + +// Config is the top-level gateway configuration. +type Config struct { + Listen string `yaml:"listen"` + GatewayKeys []string `yaml:"gateway_keys"` + DefaultModel string `yaml:"default_model"` // e.g. "AUTO" or a model id + AdapterDir string `yaml:"adapter_dir"` + RuntimeFile string `yaml:"runtime_file"` + MaxConcurrent int `yaml:"max_concurrent"` // global inflight cap, 0 = unlimited + Sources []Source `yaml:"sources"` +} + +// Model is a single exposed model id bound to a source, with priority used by +// AUTO model selection (higher number = preferred). +type Model struct { + ID string `yaml:"id"` + Priority int `yaml:"priority"` + Kind string `yaml:"kind"` // "chat" (default) | "image" + Meta map[string]interface{} `yaml:"meta"` +} + +// Source describes a single upstream LLM provider. +type Source struct { + Name string `yaml:"name"` + BaseURL string `yaml:"base_url"` + APIKey string `yaml:"api_key"` + Adapter string `yaml:"adapter"` + Endpoint string `yaml:"endpoint"` // chat endpoint override + ImageEndpoint string `yaml:"image_endpoint"` // image endpoint override + Models []Model `yaml:"models"` + Headers map[string]string `yaml:"headers"` + Meta map[string]interface{} `yaml:"meta"` + Temperature float64 `yaml:"temperature"` + MaxTokens int `yaml:"max_tokens"` + Timeout time.Duration `yaml:"timeout"` + MaxConcurrent int `yaml:"max_concurrent"` // per-source inflight cap + QueueTimeout time.Duration `yaml:"queue_timeout"` // wait for slot before failing +} + +// Load reads and validates a config file. +func Load(path string) (*Config, error) { + data, err := os.ReadFile(path) + if err != nil { + return nil, err + } + var cfg Config + if err := yaml.Unmarshal(data, &cfg); err != nil { + return nil, fmt.Errorf("parse config: %w", err) + } + if err := cfg.ApplyDefaults(); err != nil { + return nil, err + } + return &cfg, nil +} + +// ApplyDefaults sets missing values and validates the config. +func (c *Config) ApplyDefaults() error { + if c.Listen == "" { + c.Listen = ":8080" + } + if c.AdapterDir == "" { + c.AdapterDir = "adapters" + } + if c.RuntimeFile == "" { + c.RuntimeFile = "runtime.json" + } + if c.DefaultModel == "" { + c.DefaultModel = "AUTO" + } + seen := map[string]bool{} + modelOwners := map[string]string{} + for i := range c.Sources { + s := &c.Sources[i] + if s.Name == "" { + return fmt.Errorf("config: sources[%d] missing name", i) + } + if s.BaseURL == "" { + return fmt.Errorf("config: source %s missing base_url", s.Name) + } + if s.Adapter == "" { + s.Adapter = "openai" + } + if s.Timeout == 0 { + s.Timeout = 120 * time.Second + } + if s.QueueTimeout == 0 { + s.QueueTimeout = 60 * time.Second + } + if s.MaxConcurrent == 0 { + s.MaxConcurrent = 8 + } + if seen[s.Name] { + return fmt.Errorf("config: duplicate source name %q", s.Name) + } + seen[s.Name] = true + for j := range s.Models { + m := &s.Models[j] + if m.ID == "" { + return fmt.Errorf("config: source %s has a model without id", s.Name) + } + if owner, ok := modelOwners[m.ID]; ok { + return fmt.Errorf("config: model %q defined by both %s and %s", m.ID, owner, s.Name) + } + modelOwners[m.ID] = s.Name + } + } + return nil +} + +// RuntimeConfig is the persisted web-UI editable slice (sources added/edited). +type RuntimeConfig struct { + Sources []Source `json:"sources"` +} \ No newline at end of file diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..069545e --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,122 @@ +package config + +import ( + "os" + "path/filepath" + "testing" +) + +func TestLoadAndApplyDefaults(t *testing.T) { + dir := t.TempDir() + path := filepath.Join(dir, "cfg.yaml") + content := ` +listen: 127.0.0.1:9999 +gateway_keys: [sk-1] +default_model: AUTO +adapter_dir: adapters +runtime_file: runtime.json +sources: + - name: deepseek + base_url: https://api.deepseek.com + api_key: sk-d + adapter: deepseek + models: + - id: deepseek-v4-flash + priority: 100 + - name: ollama + base_url: http://127.0.0.1:11434 + adapter: ollama + endpoint: /api/chat + models: + - id: llama3 + priority: 50 +` + if err := os.WriteFile(path, []byte(content), 0644); err != nil { + t.Fatal(err) + } + cfg, err := Load(path) + if err != nil { + t.Fatalf("load: %v", err) + } + if len(cfg.Sources) != 2 { + t.Fatalf("sources = %d", len(cfg.Sources)) + } + if cfg.Sources[1].Endpoint != "/api/chat" { + t.Fatalf("endpoint = %q", cfg.Sources[1].Endpoint) + } + if cfg.Sources[1].Timeout == 0 { + t.Fatal("default timeout not applied") + } + if cfg.Sources[1].MaxConcurrent == 0 { + t.Fatal("default max_concurrent not applied") + } + if cfg.Sources[0].Models[0].Priority != 100 { + t.Fatalf("priority = %d", cfg.Sources[0].Models[0].Priority) + } + if cfg.DefaultModel != "AUTO" { + t.Fatalf("default model = %q", cfg.DefaultModel) + } +} + +func TestApplyDefaultsDuplicateSource(t *testing.T) { + cfg := Config{Sources: []Source{ + {Name: "a", BaseURL: "http://x", Models: []Model{{ID: "m1"}}}, + {Name: "a", BaseURL: "http://y", Models: []Model{{ID: "m2"}}}, + }} + if err := cfg.ApplyDefaults(); err == nil { + t.Fatal("expected duplicate source error") + } +} + +func TestApplyDefaultsNoSources(t *testing.T) { + // empty source list is valid (sources may be added later via the Web UI) + cfg := Config{} + if err := cfg.ApplyDefaults(); err != nil { + t.Fatal(err) + } + if cfg.Listen != ":8080" || cfg.DefaultModel != "AUTO" { + t.Fatalf("defaults not applied: %+v", cfg) + } +} + +func TestApplyDefaultsDuplicateModel(t *testing.T) { + cfg := Config{Sources: []Source{ + {Name: "a", BaseURL: "http://x", Models: []Model{{ID: "m1"}}}, + {Name: "b", BaseURL: "http://y", Models: []Model{{ID: "m1"}}}, + }} + if err := cfg.ApplyDefaults(); err == nil { + t.Fatal("expected duplicate model error") + } +} + +func TestStoreUpsertRemove(t *testing.T) { + path := filepath.Join(t.TempDir(), "runtime.json") + s := NewStore(path) + if err := s.Load(); err != nil { + t.Fatal(err) + } + if err := s.Upsert(Source{Name: "a", BaseURL: "http://a", Models: []Model{{ID: "m"}}}); err != nil { + t.Fatal(err) + } + if err := s.Upsert(Source{Name: "b", BaseURL: "http://b", Models: []Model{{ID: "m2"}}}); err != nil { + t.Fatal(err) + } + if len(s.List()) != 2 { + t.Fatalf("list = %d", len(s.List())) + } + removed, err := s.Remove("a") + if err != nil || !removed { + t.Fatalf("remove: %v %v", removed, err) + } + if len(s.List()) != 1 { + t.Fatalf("after remove list = %d", len(s.List())) + } + // reload from disk + s2 := NewStore(path) + if err := s2.Load(); err != nil { + t.Fatal(err) + } + if len(s2.List()) != 1 { + t.Fatalf("reloaded list = %d", len(s2.List())) + } +} \ No newline at end of file diff --git a/internal/config/store.go b/internal/config/store.go new file mode 100644 index 0000000..5e3e1cd --- /dev/null +++ b/internal/config/store.go @@ -0,0 +1,85 @@ +package config + +import ( + "encoding/json" + "os" + "sync" +) + +// Store persists web-UI editable runtime state (sources added/edited) to a +// JSON file so edits survive restarts. Base YAML sources are merged underneath. +type Store struct { + mu sync.Mutex + path string + data RuntimeConfig +} + +func NewStore(path string) *Store { + return &Store{path: path} +} + +// Load reads the runtime file (missing file = empty state). +func (s *Store) Load() error { + s.mu.Lock() + defer s.mu.Unlock() + data, err := os.ReadFile(s.path) + if err != nil { + if os.IsNotExist(err) { + s.data = RuntimeConfig{} + return nil + } + return err + } + return json.Unmarshal(data, &s.data) +} + +// List returns the runtime sources (those edited via web UI). +func (s *Store) List() []Source { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]Source, len(s.data.Sources)) + copy(out, s.data.Sources) + return out +} + +// Upsert adds or replaces a runtime source and persists. +func (s *Store) Upsert(src Source) error { + s.mu.Lock() + defer s.mu.Unlock() + for i := range s.data.Sources { + if s.data.Sources[i].Name == src.Name { + s.data.Sources[i] = src + return s.persistLocked() + } + } + s.data.Sources = append(s.data.Sources, src) + return s.persistLocked() +} + +// Remove deletes a runtime source and persists. +func (s *Store) Remove(name string) (bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + kept := s.data.Sources[:0] + removed := false + for _, src := range s.data.Sources { + if src.Name == name { + removed = true + continue + } + kept = append(kept, src) + } + if !removed { + return false, nil + } + s.data.Sources = kept + return true, s.persistLocked() +} + +func (s *Store) persistLocked() error { + b, err := json.MarshalIndent(s.data, "", " ") + if err != nil { + return err + } + return os.WriteFile(s.path, b, 0644) +} \ No newline at end of file diff --git a/internal/core/core.go b/internal/core/core.go new file mode 100644 index 0000000..5958578 --- /dev/null +++ b/internal/core/core.go @@ -0,0 +1,194 @@ +// Package core wires together the Lua VM, provider registry, scheduler and +// runtime store and exposes management operations (hot reload, adapters, +// sources) for the web UI and gateway. +package core + +import ( + "fmt" + "os" + "path/filepath" + + "llmsproxy/internal/config" + "llmsproxy/internal/lua" + "llmsproxy/internal/provider" + "llmsproxy/internal/scheduler" +) + +// Core owns the running configuration and adapters. +type Core struct { + cfg *config.Config + vm *lua.VM + store *config.Store + scheduler *scheduler.Scheduler + registry *provider.Registry +} + +// New builds the core from a config file plus runtime overlay. +func New(cfgPath string) (*Core, error) { + cfg, err := config.Load(cfgPath) + if err != nil { + return nil, err + } + return NewFromConfig(cfg) +} + +// NewFromConfig builds the core from an already-loaded config. +func NewFromConfig(cfg *config.Config) (*Core, error) { + c := &Core{cfg: cfg} + c.vm = lua.NewVM(cfg.AdapterDir) + if err := c.vm.Start(); err != nil { + return nil, fmt.Errorf("lua vm: %w", err) + } + c.store = config.NewStore(cfg.RuntimeFile) + if err := c.store.Load(); err != nil { + return nil, fmt.Errorf("runtime store: %w", err) + } + c.scheduler = scheduler.New(buildRetries(cfg)) + if err := c.rebuildRegistry(); err != nil { + return nil, err + } + return c, nil +} + +func buildRetries(cfg *config.Config) int { + return len(cfg.Sources) // allow fallback across all sources +} + +// VM exposes the Lua adapter runtime. +func (c *Core) VM() *lua.VM { return c.vm } + +func (c *Core) Scheduler() *scheduler.Scheduler { return c.scheduler } + +func (c *Core) Registry() *provider.Registry { return c.registry } + +func (c *Core) DefaultModel() string { return c.cfg.DefaultModel } + +func (c *Core) GatewayKeys() []string { return c.cfg.GatewayKeys } + +func (c *Core) Listen() string { return c.cfg.Listen } + +// Config exposes the underlying configuration (read-only usage). +func (c *Core) Config() *config.Config { return c.cfg } + +// mergedSources = base YAML sources + runtime sources (runtime wins by name). +func (c *Core) mergedSources() []config.Source { + byName := map[string]config.Source{} + order := []string{} + for _, s := range c.cfg.Sources { + byName[s.Name] = s + order = append(order, s.Name) + } + for _, s := range c.store.List() { + if _, ok := byName[s.Name]; !ok { + order = append(order, s.Name) + } + byName[s.Name] = s + } + out := make([]config.Source, 0, len(order)) + seen := map[string]bool{} + for _, n := range order { + if !seen[n] { + seen[n] = true + out = append(out, byName[n]) + } + } + return out +} + +func (c *Core) rebuildRegistry() error { + srcs := c.mergedSources() + providers := make([]*provider.Provider, 0, len(srcs)) + for _, s := range srcs { + providers = append(providers, provider.New(s, c.vm)) + } + if c.registry == nil { + c.registry = provider.NewRegistry(providers, c.cfg.DefaultModel) + } else { + c.registry.Replace(providers) + } + return nil +} + +// Reload re-reads the runtime store and rebuilds sources (adapter reload is not +// strictly needed since adapters are loaded into the VM at startup; uploaded +// adapters are placed in the adapter dir and loaded by the web UI). +func (c *Core) Reload() error { + if err := c.store.Load(); err != nil { + return err + } + return c.rebuildRegistry() +} + +// ---- adapter management (web UI) ---- + +func (c *Core) ListAdapters() []lua.APIAdapter { return c.vm.ListAdapters() } + +// UploadAdapter saves a new Lua adapter script to the adapter dir and loads it. +func (c *Core) UploadAdapter(name, code string) error { + if name == "" { + return fmt.Errorf("adapter name required") + } + if err := os.MkdirAll(c.cfg.AdapterDir, 0755); err != nil { + return err + } + path := filepath.Join(c.cfg.AdapterDir, name+".lua") + if err := os.WriteFile(path, []byte(code), 0644); err != nil { + return err + } + if err := c.vm.LoadAdapter(path); err != nil { + return fmt.Errorf("load adapter: %w", err) + } + return nil +} + +// RemoveAdapter deletes an adapter script and evicts it from the VM. +func (c *Core) RemoveAdapter(name string) error { + path := filepath.Join(c.cfg.AdapterDir, name+".lua") + _ = os.Remove(path) + c.vm.RemoveAdapter(name) + return nil +} + +// ---- source management (web UI) ---- + +func (c *Core) AddSource(src config.Source) error { + if err := normalizeSource(&src); err != nil { + return err + } + if err := c.store.Upsert(src); err != nil { + return err + } + return c.rebuildRegistry() +} + +func (c *Core) RemoveSource(name string) error { + if _, err := c.store.Remove(name); err != nil { + return err + } + return c.rebuildRegistry() +} + +func (c *Core) Sources() []config.Source { return c.mergedSources() } + +func normalizeSource(s *config.Source) error { + if s.Name == "" || s.BaseURL == "" { + return fmt.Errorf("source requires name and base_url") + } + if len(s.Models) == 0 { + return fmt.Errorf("source requires at least one model") + } + if s.Adapter == "" { + s.Adapter = "openai" + } + if s.MaxConcurrent == 0 { + s.MaxConcurrent = 8 + } + return nil +} + +// Close releases resources. +func (c *Core) Close() { + if c.vm != nil { + c.vm.Stop() + } +} \ No newline at end of file diff --git a/internal/gateway/api.go b/internal/gateway/api.go new file mode 100644 index 0000000..ae912b1 --- /dev/null +++ b/internal/gateway/api.go @@ -0,0 +1,116 @@ +package gateway + +import ( + "encoding/json" + "io" + "net/http" + "strings" + + "llmsproxy/internal/config" +) + +type adapterPayload struct { + Name string `json:"name"` + Code string `json:"code"` +} + +func (g *Gateway) handleAdaptersAPI(w http.ResponseWriter, r *http.Request) { + path := strings.TrimPrefix(r.URL.Path, "/api/adapters") + path = strings.Trim(path, "/") + + switch r.Method { + case http.MethodGet: + if path == "" { + writeJSON(w, http.StatusOK, map[string]interface{}{"adapters": g.core.ListAdapters()}) + return + } + writeError(w, http.StatusNotFound, "not_found", "adapter code not exposed; edit in UI") + case http.MethodPost: + var p adapterPayload + if err := json.NewDecoder(r.Body).Decode(&p); err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error()) + return + } + if err := g.core.UploadAdapter(p.Name, p.Code); err != nil { + writeError(w, http.StatusBadRequest, "adapter_error", err.Error()) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{"ok": true, "name": p.Name}) + case http.MethodDelete: + if path == "" { + writeError(w, http.StatusBadRequest, "invalid_request", "adapter name required") + return + } + if err := g.core.RemoveAdapter(path); err != nil { + writeError(w, http.StatusBadRequest, "adapter_error", err.Error()) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{"ok": true}) + default: + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "") + } +} + +// sourcePayload mirrors config.Source for JSON web UI editing. +type sourcePayload struct { + Name string `json:"name"` + BaseURL string `json:"base_url"` + APIKey string `json:"api_key"` + Adapter string `json:"adapter"` + Endpoint string `json:"endpoint"` + ImageEndpoint string `json:"image_endpoint"` + Models []config.Model `json:"models"` + Headers map[string]string `json:"headers"` + Meta map[string]interface{} `json:"meta"` + Temperature float64 `json:"temperature"` + MaxTokens int `json:"max_tokens"` + MaxConcurrent int `json:"max_concurrent"` +} + +func (g *Gateway) handleSourcesAPI(w http.ResponseWriter, r *http.Request) { + path := strings.TrimPrefix(r.URL.Path, "/api/sources") + path = strings.Trim(path, "/") + + switch r.Method { + case http.MethodGet: + writeJSON(w, http.StatusOK, map[string]interface{}{"sources": g.core.Sources()}) + case http.MethodPost: + var p sourcePayload + body, _ := io.ReadAll(r.Body) + if err := json.Unmarshal(body, &p); err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error()) + return + } + src := config.Source{ + Name: p.Name, + BaseURL: p.BaseURL, + APIKey: p.APIKey, + Adapter: p.Adapter, + Endpoint: p.Endpoint, + ImageEndpoint: p.ImageEndpoint, + Models: p.Models, + Headers: p.Headers, + Meta: p.Meta, + Temperature: p.Temperature, + MaxTokens: p.MaxTokens, + MaxConcurrent: p.MaxConcurrent, + } + if err := g.core.AddSource(src); err != nil { + writeError(w, http.StatusBadRequest, "source_error", err.Error()) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{"ok": true}) + case http.MethodDelete: + if path == "" { + writeError(w, http.StatusBadRequest, "invalid_request", "source name required") + return + } + if err := g.core.RemoveSource(path); err != nil { + writeError(w, http.StatusBadRequest, "source_error", err.Error()) + return + } + writeJSON(w, http.StatusOK, map[string]interface{}{"ok": true}) + default: + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "") + } +} \ No newline at end of file diff --git a/internal/gateway/chat.go b/internal/gateway/chat.go new file mode 100644 index 0000000..6964fbc --- /dev/null +++ b/internal/gateway/chat.go @@ -0,0 +1,283 @@ +package gateway + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "strings" + "sync/atomic" + "time" + + "llmsproxy/internal/provider" + "llmsproxy/internal/scheduler" + "llmsproxy/internal/types" +) + +// chatRequest mirrors the OpenAI chat completions request the gateway accepts. +type chatRequest struct { + Model string `json:"model"` + Messages []types.ChatMessage `json:"messages"` + Temperature *float64 `json:"temperature,omitempty"` + MaxTokens int `json:"max_tokens,omitempty"` + Stream bool `json:"stream,omitempty"` + Tools []interface{} `json:"tools,omitempty"` + ToolChoice interface{} `json:"tool_choice,omitempty"` +} + +// ChatCompletion is the non-streaming OpenAI response object. +type ChatCompletion struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []ChatChoice `json:"choices"` + Usage *types.TokenUsage `json:"usage,omitempty"` +} + +type ChatChoice struct { + Index int `json:"index"` + Message RespMessage `json:"message"` + FinishReason string `json:"finish_reason"` +} + +type RespMessage struct { + Role string `json:"role"` + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content,omitempty"` + ToolCalls []types.ToolCall `json:"tool_calls,omitempty"` +} + +type ChatChunk struct { + ID string `json:"id"` + Object string `json:"object"` + Created int64 `json:"created"` + Model string `json:"model"` + Choices []ChunkChoice `json:"choices"` +} + +type ChunkChoice struct { + Index int `json:"index"` + Delta RespMessage `json:"delta"` + FinishReason *string `json:"finish_reason"` +} + +var seq int64 + +func newID() string { + n := atomic.AddInt64(&seq, 1) + return fmt.Sprintf("chatcmpl-%d", n) +} + +func isAuto(m string) bool { + m = strings.TrimSpace(m) + return m == "" || strings.EqualFold(m, "AUTO") +} + +// resolveCands picks the ordered candidate providers for a requested model. +func (g *Gateway) resolveCands(model string) ([]*provider.Provider, string) { + if model == "" || isAuto(model) { + return g.core.Registry().Resolve("AUTO"), "" + } + return g.core.Registry().Resolve(model), model +} + +func (g *Gateway) handleChat(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "use POST") + return + } + var req chatRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error()) + return + } + if len(req.Messages) == 0 { + writeError(w, http.StatusBadRequest, "invalid_request", "messages is required") + return + } + model := req.Model + if model == "" { + model = g.core.DefaultModel() + } + cands, effective := g.resolveCands(model) + if len(cands) == 0 { + writeError(w, http.StatusServiceUnavailable, "no_provider", "no LLM source configured") + return + } + if effective == "" { + effective = firstModel(cands[0]) + } + ctx := r.Context() + + inner := &types.ChatRequest{ + Model: normalizeModel(model), + Messages: req.Messages, + Temperature: req.Temperature, + MaxTokens: req.MaxTokens, + Stream: req.Stream, + Tools: req.Tools, + ToolChoice: req.ToolChoice, + } + if req.Stream { + g.streamChat(w, ctx, cands, inner, effective) + return + } + g.singleChat(w, ctx, cands, inner, effective) +} + +func normalizeModel(m string) string { + if isAuto(m) { + return "" + } + return m +} + +func firstModel(p *provider.Provider) string { + ms := p.Models() + if len(ms) > 0 { + return ms[0] + } + return "auto" +} + +// imageOnly keeps providers exposing at least one image-kind model. +func imageOnly(cands []*provider.Provider) []*provider.Provider { + var out []*provider.Provider + for _, p := range cands { + for _, id := range p.Models() { + if m := p.ModelByID(id); m != nil && m.Kind == "image" { + out = append(out, p) + break + } + } + } + return out +} + +func (g *Gateway) singleChat(w http.ResponseWriter, ctx context.Context, cands []*provider.Provider, req *types.ChatRequest, effective string) { + resp, err := g.core.Scheduler().Chat(ctx, scheduler.FromRegistry(cands), req) + if err != nil { + writeError(w, http.StatusBadGateway, "upstream_error", err.Error()) + return + } + msg := RespMessage{Role: "assistant", Content: resp.Content} + if resp.ReasoningContent != "" { + msg.ReasoningContent = resp.ReasoningContent + } + if len(resp.ToolCalls) > 0 { + msg.ToolCalls = resp.ToolCalls + } + out := ChatCompletion{ + ID: newID(), + Object: "chat.completion", + Created: time.Now().Unix(), + Model: effective, + Choices: []ChatChoice{{Index: 0, Message: msg, FinishReason: resp.FinishReason}}, + } + if resp.TokenUsage.Total > 0 || resp.TokenUsage.Prompt > 0 || resp.TokenUsage.Completion > 0 { + out.Usage = &resp.TokenUsage + } + writeJSON(w, http.StatusOK, out) +} + +func (g *Gateway) streamChat(w http.ResponseWriter, ctx context.Context, cands []*provider.Provider, req *types.ChatRequest, effective string) { + chunks, err := g.core.Scheduler().ChatStream(ctx, scheduler.FromRegistry(cands), req) + if err != nil { + writeError(w, http.StatusBadGateway, "upstream_error", err.Error()) + return + } + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + + flusher, _ := w.(http.Flusher) + id := newID() + created := time.Now().Unix() + send := func(obj interface{}) bool { + b, err := json.Marshal(obj) + if err != nil { + return false + } + if _, err := fmt.Fprintf(w, "data: %s\n\n", b); err != nil { + return false + } + if flusher != nil { + flusher.Flush() + } + return true + } + + if !send(ChatChunk{ + ID: id, Object: "chat.completion.chunk", Created: created, Model: effective, + Choices: []ChunkChoice{{Index: 0, Delta: RespMessage{Role: "assistant"}}}, + }) { + return + } + for ck := range chunks { + chunk := ChatChunk{ + ID: id, Object: "chat.completion.chunk", Created: created, Model: effective, + } + delta := RespMessage{Content: ck.Content} + if ck.ReasoningContent != "" { + delta.ReasoningContent = ck.ReasoningContent + } + if len(ck.ToolCalls) > 0 { + delta.ToolCalls = ck.ToolCalls + } + choice := ChunkChoice{Index: 0, Delta: delta} + if ck.Done { + stop := "stop" + choice.FinishReason = &stop + } + chunk.Choices = []ChunkChoice{choice} + if !send(chunk) { + return + } + } + stop := "stop" + send(ChatChunk{ + ID: id, Object: "chat.completion.chunk", Created: created, Model: effective, + Choices: []ChunkChoice{{Index: 0, Delta: RespMessage{}, FinishReason: &stop}}, + }) + fmt.Fprintf(w, "data: [DONE]\n\n") + if flusher != nil { + flusher.Flush() + } +} + +func (g *Gateway) handleImage(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "use POST") + return + } + var req types.ImageGenRequest + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeError(w, http.StatusBadRequest, "invalid_request", "invalid json: "+err.Error()) + return + } + if req.Prompt == "" { + writeError(w, http.StatusBadRequest, "invalid_request", "prompt is required") + return + } + model := req.Model + if model == "" { + model = g.core.DefaultModel() + } + cands, _ := g.resolveCands(model) + cands = imageOnly(cands) + if len(cands) == 0 { + writeError(w, http.StatusServiceUnavailable, "no_provider", "no image source configured") + return + } + resp, err := g.core.Scheduler().Image(r.Context(), scheduler.FromRegistry(cands), &req) + if err != nil { + writeError(w, http.StatusBadGateway, "upstream_error", err.Error()) + return + } + writeJSON(w, http.StatusOK, types.ImageGenResponse{ + Created: time.Now().Unix(), + Data: resp.ImageData, + }) +} \ No newline at end of file diff --git a/internal/gateway/gateway_test.go b/internal/gateway/gateway_test.go new file mode 100644 index 0000000..7afdd88 --- /dev/null +++ b/internal/gateway/gateway_test.go @@ -0,0 +1,306 @@ +package gateway + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" + + "llmsproxy/internal/config" + "llmsproxy/internal/core" +) + +func mockUpstream() *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + var req map[string]interface{} + _ = json.Unmarshal(body, &req) + if stream, _ := req["stream"].(bool); stream { + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(200) + fmt.Fprintln(w, `data: {"choices":[{"delta":{"content":"Hel"}}]}`) + fmt.Fprintln(w, `data: {"choices":[{"delta":{"content":"lo"}}]}`) + fmt.Fprintln(w, `data: {"choices":[{"delta":{},"finish_reason":"stop"}]}`) + fmt.Fprintln(w, "data: [DONE]") + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(200) + fmt.Fprintf(w, `{"choices":[{"message":{"content":"pong"},"finish_reason":"stop"}],"usage":{"prompt_tokens":3,"completion_tokens":1,"total_tokens":4}}`) + })) +} + +func newTestGateway(t *testing.T, srcs ...config.Source) *Gateway { + t.Helper() + cfg := &config.Config{ + AdapterDir: filepath.Join(t.TempDir(), "adapters"), + RuntimeFile: filepath.Join(t.TempDir(), "runtime.json"), + Sources: srcs, + } + if err := cfg.ApplyDefaults(); err != nil { + t.Fatal(err) + } + c, err := core.NewFromConfig(cfg) + if err != nil { + t.Fatalf("core: %v", err) + } + t.Cleanup(c.Close) + g, err := New(c, []string{"sk-test"}) + if err != nil { + t.Fatalf("gateway: %v", err) + } + return g +} + +func doReq(t *testing.T, g *Gateway, method, path, body string) *httptest.ResponseRecorder { + t.Helper() + req, _ := http.NewRequest(method, path, strings.NewReader(body)) + req.Header.Set("Authorization", "Bearer sk-test") + if body != "" { + req.Header.Set("Content-Type", "application/json") + } + rr := httptest.NewRecorder() + g.Handler().ServeHTTP(rr, req) + return rr +} + +func TestChatSingle(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, config.Source{Name: "mock", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "mock-model"}}}) + rr := doReq(t, g, "POST", "/v1/chat/completions", + `{"model":"mock-model","messages":[{"role":"user","content":"hi"}]}`) + if rr.Code != 200 { + t.Fatalf("status = %d, body=%s", rr.Code, rr.Body.String()) + } + var cc ChatCompletion + if err := json.Unmarshal(rr.Body.Bytes(), &cc); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cc.Choices[0].Message.Content != "pong" { + t.Fatalf("content = %q", cc.Choices[0].Message.Content) + } + if cc.Usage == nil || cc.Usage.Total != 4 { + t.Fatalf("usage = %+v", cc.Usage) + } + if cc.Model != "mock-model" { + t.Fatalf("model = %q", cc.Model) + } +} + +func TestChatAUTO(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, + config.Source{Name: "low", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "low-m", Priority: 10}}}, + config.Source{Name: "high", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "high-m", Priority: 100}}}, + ) + // no model -> AUTO -> picks the highest priority source + rr := doReq(t, g, "POST", "/v1/chat/completions", + `{"model":"AUTO","messages":[{"role":"user","content":"hi"}]}`) + if rr.Code != 200 { + t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String()) + } + var cc ChatCompletion + _ = json.Unmarshal(rr.Body.Bytes(), &cc) + if cc.Model != "high-m" { + t.Fatalf("AUTO picked %q, want high-m", cc.Model) + } +} + +func TestChatAuthRequired(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, config.Source{Name: "mock", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "mock-model"}}}) + req, _ := http.NewRequest("POST", "/v1/chat/completions", + strings.NewReader(`{"messages":[{"role":"user","content":"hi"}]}`)) + rr := httptest.NewRecorder() + g.Handler().ServeHTTP(rr, req) + if rr.Code != http.StatusUnauthorized { + t.Fatalf("expected 401, got %d", rr.Code) + } +} + +func TestChatStream(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, config.Source{Name: "mock", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "mock-model"}}}) + rr := doReq(t, g, "POST", "/v1/chat/completions", + `{"model":"mock-model","stream":true,"messages":[{"role":"user","content":"hi"}]}`) + body := rr.Body.String() + if !strings.Contains(body, "data: [DONE]") { + t.Fatalf("missing DONE, body=%s", body) + } + if !strings.Contains(body, "Hel") || !strings.Contains(body, "lo") { + t.Fatalf("missing content chunks, body=%s", body) + } +} + +func TestImageGeneration(t *testing.T) { + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + fmt.Fprint(w, `{"created":1,"data":[{"b64_json":"QUJD"}]}`) + })) + defer up.Close() + g := newTestGateway(t, config.Source{Name: "img", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "flux-1", Kind: "image"}}}) + rr := doReq(t, g, "POST", "/v1/images/generations", + `{"model":"flux-1","prompt":"a cat"}`) + if rr.Code != 200 { + t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String()) + } + var out map[string]interface{} + _ = json.Unmarshal(rr.Body.Bytes(), &out) + data, _ := out["data"].([]interface{}) + if len(data) != 1 { + t.Fatalf("image data len = %d", len(data)) + } +} + +func TestImageAutoFallsOnlyToImageProviders(t *testing.T) { + imageHits := 0 + chatHits := 0 + img := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + imageHits++ + w.Header().Set("Content-Type", "application/json") + fmt.Fprint(w, `{"created":1,"data":[{"b64_json":"QUJD"}]}`) + })) + defer img.Close() + chatUp := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + chatHits++ + w.Header().Set("Content-Type", "application/json") + fmt.Fprint(w, `{"choices":[{"message":{"content":"hi"},"finish_reason":"stop"}]}`) + })) + defer chatUp.Close() + g := newTestGateway(t, + config.Source{Name: "chat", BaseURL: chatUp.URL, Adapter: "openai", Models: []config.Model{{ID: "chat-m", Priority: 100}}}, + config.Source{Name: "img", BaseURL: img.URL, Adapter: "openai", Models: []config.Model{{ID: "flux", Kind: "image", Priority: 1}}}, + ) + rr := doReq(t, g, "POST", "/v1/images/generations", + `{"model":"AUTO","prompt":"a cat"}`) + if rr.Code != 200 { + t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String()) + } + if chatHits != 0 { + t.Fatalf("image AUTO hit chat-only provider: %d chat hits", chatHits) + } + if imageHits == 0 { + t.Fatalf("image AUTO did not hit image provider") + } +} + +func TestKimicodeSigning(t *testing.T) { + var gotAuth, gotSign string + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, _ = io.ReadAll(r.Body) + gotAuth = r.Header.Get("Authorization") + gotSign = r.Header.Get("X-App-Sign") + w.Header().Set("Content-Type", "application/json") + fmt.Fprintf(w, `{"choices":[{"message":{"content":"ok"},"finish_reason":"stop"}]}`) + })) + defer up.Close() + + g := newTestGateway(t, config.Source{ + Name: "kimi", + BaseURL: up.URL, + Adapter: "kimicode", + APIKey: "sk-kimi", + Models: []config.Model{{ID: "kimi-k2"}}, + Meta: map[string]interface{}{"app_id": "app-1", "app_secret": "s3cr3t", "app_agent": "code-agent", "api_key": "sk-kimi"}, + }) + rr := doReq(t, g, "POST", "/v1/chat/completions", + `{"model":"kimi-k2","messages":[{"role":"user","content":"hi"}]}`) + if rr.Code != 200 { + t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String()) + } + if gotAuth == "" || !strings.Contains(gotAuth, "sk-kimi") { + t.Fatalf("expected signed auth, got %q", gotAuth) + } + if gotSign == "" { + t.Fatalf("expected app signature header") + } +} + +func TestModelRoutingPrefix(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, + config.Source{Name: "a", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "model-a"}}}, + config.Source{Name: "b", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "model-b"}}}, + ) + rr := doReq(t, g, "POST", "/v1/chat/completions", + `{"model":"model-b","messages":[{"role":"user","content":"hi"}]}`) + if rr.Code != 200 { + t.Fatalf("status=%d", rr.Code) + } +} + +func TestModelsEndpoint(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, + config.Source{Name: "a", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "model-a"}}}, + config.Source{Name: "b", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "model-b"}}}, + ) + rr := doReq(t, g, "GET", "/v1/models", "") + if rr.Code != 200 { + t.Fatalf("status=%d", rr.Code) + } + var out map[string]interface{} + _ = json.Unmarshal(rr.Body.Bytes(), &out) + if !strings.Contains(rr.Body.String(), "model-a") || !strings.Contains(rr.Body.String(), "model-b") { + t.Fatalf("missing models: %s", rr.Body.String()) + } +} + +func TestWebUIServesPage(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, config.Source{Name: "mock", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "mock-model"}}}) + rr := doReq(t, g, "GET", "/", "") + if rr.Code != 200 { + t.Fatalf("status=%d", rr.Code) + } + if !strings.Contains(rr.Body.String(), "llmsproxy") { + t.Fatalf("ui not served") + } +} + +func TestAdaptersAPIUpload(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t, config.Source{Name: "mock", BaseURL: up.URL, Adapter: "openai", Models: []config.Model{{ID: "mock-model"}}}) + // upload adapter + rr := doReq(t, g, "POST", "/api/adapters", + `{"name":"testadp","code":"return {name='testadp',endpoint='/chat/completions',transform_request=function(raw) return raw end,transform_response=function(raw) return raw end}"}`) + if rr.Code != 200 { + t.Fatalf("upload status=%d body=%s", rr.Code, rr.Body.String()) + } + rr = doReq(t, g, "GET", "/api/status", "") + if !strings.Contains(rr.Body.String(), "testadp") { + t.Fatalf("adapter not listed: %s", rr.Body.String()) + } +} + +func TestSourcesAPIAddAndPersist(t *testing.T) { + up := mockUpstream() + defer up.Close() + g := newTestGateway(t) + rr := doReq(t, g, "POST", "/api/sources", + fmt.Sprintf(`{"name":"added","base_url":"%s","adapter":"openai","models":[{"id":"new-m","priority":5}]}`, up.URL)) + if rr.Code != 200 { + t.Fatalf("add source status=%d body=%s", rr.Code, rr.Body.String()) + } + rr = doReq(t, g, "GET", "/v1/models", "") + if !strings.Contains(rr.Body.String(), "new-m") { + t.Fatalf("new model not live: %s", rr.Body.String()) + } + // verify persistence file exists + if _, err := os.Stat(g.core.Config().RuntimeFile); err != nil { + t.Fatalf("runtime file not written: %v", err) + } +} \ No newline at end of file diff --git a/internal/gateway/server.go b/internal/gateway/server.go new file mode 100644 index 0000000..ca0c683 --- /dev/null +++ b/internal/gateway/server.go @@ -0,0 +1,149 @@ +// Package gateway exposes an OpenAI-compatible HTTP API over the provider +// registry: POST /v1/chat/completions (SDK + SSE), POST /v1/images/generations, +// GET /v1/models, protected by shared gateway API keys, plus a web UI and +// management API for adapters and sources. +package gateway + +import ( + "embed" + "encoding/json" + "io/fs" + "log" + "net/http" + "strings" + + "llmsproxy/internal/core" +) + +//go:embed ui/* +var uiFS embed.FS + +// Gateway is the HTTP handler for the OpenAI-compatible endpoint + web UI. +type Gateway struct { + core *core.Core + apiKeys map[string]bool + ui http.Handler +} + +func New(c *core.Core, gatewayKeys []string) (*Gateway, error) { + keys := map[string]bool{} + for _, k := range gatewayKeys { + if k != "" { + keys[k] = true + } + } + sub, err := fs.Sub(uiFS, "ui") + if err != nil { + return nil, err + } + return &Gateway{ + core: c, + apiKeys: keys, + ui: http.FileServer(http.FS(sub)), + }, nil +} + +func (g *Gateway) Handler() http.Handler { + return g.auth(http.HandlerFunc(g.routes)) +} + +func (g *Gateway) routes(w http.ResponseWriter, r *http.Request) { + switch { + case r.URL.Path == "/v1/chat/completions": + g.handleChat(w, r) + case r.URL.Path == "/v1/images/generations": + g.handleImage(w, r) + case r.URL.Path == "/v1/models": + g.handleModels(w, r) + case r.URL.Path == "/api/adapters" || strings.HasPrefix(r.URL.Path, "/api/adapters/"): + g.handleAdaptersAPI(w, r) + case r.URL.Path == "/api/sources" || strings.HasPrefix(r.URL.Path, "/api/sources/"): + g.handleSourcesAPI(w, r) + case r.URL.Path == "/api/status": + g.handleStatusAPI(w, r) + default: + g.serveUI(w, r) + } +} + +func (g *Gateway) serveUI(w http.ResponseWriter, r *http.Request) { + // serve index.html directly for the root path (FileServer would 301 it) + if r.URL.Path == "/" || r.URL.Path == "/ui" { + data, err := uiFS.ReadFile("ui/index.html") + if err != nil { + http.Error(w, "ui missing", http.StatusInternalServerError) + return + } + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.Write(data) + return + } + g.ui.ServeHTTP(w, r) +} + +func (g *Gateway) auth(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if len(g.apiKeys) == 0 { + next.ServeHTTP(w, r) + return + } + key := "" + if h := r.Header.Get("Authorization"); h != "" { + parts := strings.SplitN(h, " ", 2) + if len(parts) == 2 && strings.EqualFold(parts[0], "Bearer") { + key = parts[1] + } + } + if key == "" { + key = r.URL.Query().Get("api_key") + } + if !g.apiKeys[key] { + writeError(w, http.StatusUnauthorized, "invalid_api_key", "invalid gateway api key") + return + } + next.ServeHTTP(w, r) + }) +} + +func (g *Gateway) handleModels(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "use GET") + return + } + models := g.core.Registry().ModelList() + type modelObj struct { + ID string `json:"id"` + Object string `json:"object"` + } + objs := make([]modelObj, 0, len(models)) + for _, m := range models { + objs = append(objs, modelObj{ID: m, Object: "model"}) + } + writeJSON(w, http.StatusOK, map[string]interface{}{ + "object": "list", + "data": objs, + }) +} + +func (g *Gateway) handleStatusAPI(w http.ResponseWriter, r *http.Request) { + writeJSON(w, http.StatusOK, map[string]interface{}{ + "default_model": g.core.DefaultModel(), + "models": g.core.Registry().ModelList(), + "sources": g.core.Registry().Status(), + "adapters": g.core.ListAdapters(), + }) +} + +func writeError(w http.ResponseWriter, code int, errType, msg string) { + writeJSON(w, code, map[string]interface{}{ + "error": map[string]interface{}{"type": errType, "message": msg}, + }) +} + +func writeJSON(w http.ResponseWriter, code int, v interface{}) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(code) + if err := json.NewEncoder(w).Encode(v); err != nil { + log.Printf("[gateway] write json: %v", err) + } +} \ No newline at end of file diff --git a/internal/gateway/ui/index.html b/internal/gateway/ui/index.html new file mode 100644 index 0000000..009fa35 --- /dev/null +++ b/internal/gateway/ui/index.html @@ -0,0 +1,241 @@ + + + + + +llmsproxy + + + +
+

llmsproxy

+ 统一 LLM 网关 · 适配器/源管理 +
+ +
+
+ + +
+
+ + + \ No newline at end of file diff --git a/internal/lua/adapters/anthropic.lua b/internal/lua/adapters/anthropic.lua new file mode 100644 index 0000000..56ad3ab --- /dev/null +++ b/internal/lua/adapters/anthropic.lua @@ -0,0 +1,86 @@ +local adapter = {} + +adapter.name = "anthropic" +adapter.version = "2.0.0" +adapter.endpoint = "/v1/messages" +adapter.headers = { + ["anthropic-version"] = "2023-06-01" +} + +function adapter.transform_request(raw_body) + local ok, req = pcall(json.decode, raw_body) + if not ok then return raw_body end + + local msgs = {} + local system = "" + for _, m in ipairs(req.messages or {}) do + if m.role == "system" then + system = system .. m.content .. "\n" + else + table.insert(msgs, { role = m.role, content = m.content }) + end + end + + local anthropic_req = { + model = req.model or "claude-sonnet-4-20250514", + max_tokens = req.max_tokens or 4096, + messages = msgs, + stream = req.stream or false, + } + + if not req.disable_thinking then + anthropic_req.thinking = { type = "enabled", budget_tokens = 4096 } + end + + if system ~= "" then + anthropic_req.system = system + end + + return json.encode(anthropic_req) +end + +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.input_tokens or 0 + unified.token_usage.completion = resp.usage.output_tokens or 0 + unified.token_usage.total = (resp.usage.input_tokens or 0) + (resp.usage.output_tokens or 0) + end + + if resp.content and #resp.content > 0 then + for _, block in ipairs(resp.content) do + if block.type == "text" then + unified.content = unified.content .. (block.text or "") + end + end + end + unified.finish_reason = resp.stop_reason or "" + + 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 chunk.type == "message_start" then return "" end + if chunk.type == "message_delta" then + return json.encode({ content = "", done = (chunk.delta and chunk.delta.stop_reason ~= nil) }) + end + if chunk.type == "content_block_delta" and chunk.delta then + return json.encode({ content = chunk.delta.text or "", done = false }) + end + if chunk.type == "message_stop" then + return json.encode({ content = "", done = true }) + end + return "" +end + +return adapter diff --git a/internal/lua/adapters/deepseek.lua b/internal/lua/adapters/deepseek.lua new file mode 100644 index 0000000..bcc6e41 --- /dev/null +++ b/internal/lua/adapters/deepseek.lua @@ -0,0 +1,78 @@ +local adapter = {} + +adapter.name = "deepseek" +adapter.version = "2.1.0" +adapter.endpoint = "/chat/completions" +adapter.headers = {} + +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.stream = req.stream or false + if req.disable_thinking then + req.extra_body = req.extra_body or {} + req.extra_body.thinking = { type = "disabled" } + end + req.disable_thinking = nil + return json.encode(req) +end + +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok or resp == nil then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if type(resp.usage) == "table" 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 type(resp.choices) == "table" and #resp.choices > 0 then + local ch = resp.choices[1] + if type(ch.message) == "table" then + unified.content = ch.message.content or "" + if ch.message.reasoning_content then + unified.reasoning_content = ch.message.reasoning_content + end + if type(ch.message.tool_calls) == "table" 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/gemini.lua b/internal/lua/adapters/gemini.lua new file mode 100644 index 0000000..cfc4a42 --- /dev/null +++ b/internal/lua/adapters/gemini.lua @@ -0,0 +1,89 @@ +local adapter = {} + +adapter.name = "gemini" +adapter.version = "2.0.0" +adapter.endpoint = "/v1/models" +adapter.headers = {} + +-- Gemini API: POST /v1/models/{model}:generateContent +-- Auth: API key in query param ?key=XXX or Authorization: Bearer XXX +function adapter.transform_request(raw_body) + local ok, req = pcall(json.decode, raw_body) + if not ok then return raw_body end + + local contents = {} + for _, m in ipairs(req.messages or {}) do + table.insert(contents, { + role = (m.role == "assistant") and "model" or m.role, + parts = { { text = m.content } } + }) + end + + local gemini_req = { + contents = contents, + generationConfig = { + temperature = req.temperature or 0.7, + maxOutputTokens = req.max_tokens or 4096, + } + } + + if req.stream then + gemini_req.stream = true + end + + return json.encode(gemini_req) +end + +-- Gemini 的 endpoint 动态拼接:/v1/models/{model}:generateContent +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.usageMetadata then + unified.token_usage.prompt = resp.usageMetadata.promptTokenCount or 0 + unified.token_usage.completion = resp.usageMetadata.candidatesTokenCount or 0 + unified.token_usage.total = resp.usageMetadata.totalTokenCount or 0 + end + + if resp.candidates and #resp.candidates > 0 then + local cand = resp.candidates[1] + if cand.content and cand.content.parts then + for _, part in ipairs(cand.content.parts) do + if part.text then + unified.content = unified.content .. part.text + end + end + end + if cand.finishReason then + unified.finish_reason = cand.finishReason + end + 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.candidates or #chunk.candidates == 0 then return "" end + local cand = chunk.candidates[1] + local content = "" + if cand.content and cand.content.parts then + for _, part in ipairs(cand.content.parts) do + content = content .. (part.text or "") + end + end + return json.encode({ + content = content, + done = (cand.finishReason ~= nil) + }) +end + +return adapter diff --git a/internal/lua/adapters/github.lua b/internal/lua/adapters/github.lua new file mode 100644 index 0000000..c27103b --- /dev/null +++ b/internal/lua/adapters/github.lua @@ -0,0 +1,75 @@ +local adapter = {} + +adapter.name = "github" +adapter.version = "2.0.0" +adapter.endpoint = "/chat/completions" +adapter.headers = {} + +-- GitHub Models: Azure-like endpoint, auth via Bearer token (PAT) +-- BaseURL example: https://models.inference.ai.azure.com +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 "gpt-4o" + req.temperature = req.temperature or 0.7 + req.max_tokens = req.max_tokens or 4096 + req.stream = req.stream or false + req.disable_thinking = nil + req.extra_body = nil + return json.encode(req) +end + +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok or resp == nil then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if type(resp.usage) == "table" 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 type(resp.choices) == "table" and #resp.choices > 0 then + local ch = resp.choices[1] + if type(ch.message) == "table" then + unified.content = ch.message.content or "" + if type(ch.message.tool_calls) == "table" 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/groq.lua b/internal/lua/adapters/groq.lua new file mode 100644 index 0000000..3d66033 --- /dev/null +++ b/internal/lua/adapters/groq.lua @@ -0,0 +1,74 @@ +local adapter = {} + +adapter.name = "groq" +adapter.version = "2.0.0" +adapter.endpoint = "/openai/v1/chat/completions" +adapter.headers = {} + +-- Groq API is OpenAI-compatible +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 "llama3-70b-8192" + req.temperature = req.temperature or 0.7 + req.max_tokens = req.max_tokens or 4096 + req.stream = req.stream or false + req.disable_thinking = nil + req.extra_body = nil + return json.encode(req) +end + +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok or resp == nil then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if type(resp.usage) == "table" 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 type(resp.choices) == "table" and #resp.choices > 0 then + local ch = resp.choices[1] + if type(ch.message) == "table" then + unified.content = ch.message.content or "" + if type(ch.message.tool_calls) == "table" 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/kimicode.lua b/internal/lua/adapters/kimicode.lua new file mode 100644 index 0000000..064a460 --- /dev/null +++ b/internal/lua/adapters/kimicode.lua @@ -0,0 +1,107 @@ +local adapter = {} + +adapter.name = "kimicode" +adapter.version = "1.0.0" +adapter.endpoint = "/v1/chat/completions" +adapter.headers = {} + +-- KimiCode / Kimi K2 属于 OpenAI 兼容协议;但部分云端 API 会校验调用方 +-- "app"(只放行特定 agent),要求每次请求带上按 secret 计算的应用签名。 +-- 这里演示 build_headers 钩子:基于 timestamp + 请求体哈希生成签名头。 +-- +-- 配置要求(source.meta): +-- meta: +-- app_id: <申请到的 app id> +-- app_key: <你的 key(由网关的 base_url 复用 api_key 亦可)> +-- app_secret: <签名密钥> +-- app_agent: code-agent # 若云端要求声明 agent 身份 +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 "kimi-k2" + req.disable_thinking = nil + req.extra_body = nil + return json.encode(req) +end + +-- 可选的动态签名钩子。meta 由 Go 注入: +-- meta.url / meta.method / meta.body / meta.api_key / meta.timestamp / meta.source.meta +function adapter.build_headers(meta) + local h = { + ["Content-Type"] = "application/json", + ["X-App-Id"] = tostring((meta.source.meta or {}).app_id or ""), + ["X-Timestamp"] = tostring(meta.timestamp), + } + local agent = (meta.source.meta or {}).app_agent + if agent and agent ~= "" then + h["X-Agent"] = agent + end + -- 校验 app:通常要求 Authorization 用 app secret 派生签名 + local secret = (meta.source.meta or {}).app_secret + local api_key = meta.source.meta and meta.source.meta.api_key or meta.api_key + if secret and secret ~= "" then + local body_hash = sha256_hex(meta.body) + local sign_string = tostring(meta.timestamp) .. meta.method .. meta.url .. body_hash + local sign = hmac_sha256_hex(secret, sign_string) + h["Authorization"] = "Bearer " .. api_key + h["X-App-Sign"] = sign + else + h["Authorization"] = "Bearer " .. api_key + end + return h +end + +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok or resp == nil then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + if type(resp.usage) == "table" 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 type(resp.choices) == "table" and #resp.choices > 0 then + local ch = resp.choices[1] + if type(ch.message) == "table" then + unified.content = ch.message.content or "" + if ch.message.reasoning_content then + unified.reasoning_content = ch.message.reasoning_content + end + if type(ch.message.tool_calls) == "table" 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 \ No newline at end of file diff --git a/internal/lua/adapters/mistral.lua b/internal/lua/adapters/mistral.lua new file mode 100644 index 0000000..763da2a --- /dev/null +++ b/internal/lua/adapters/mistral.lua @@ -0,0 +1,74 @@ +local adapter = {} + +adapter.name = "mistral" +adapter.version = "2.0.0" +adapter.endpoint = "/v1/chat/completions" +adapter.headers = {} + +-- Mistral API is OpenAI-compatible, just passes through +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 "mistral-large-latest" + req.temperature = req.temperature or 0.7 + req.max_tokens = req.max_tokens or 4096 + req.stream = req.stream or false + req.disable_thinking = nil + req.extra_body = nil + return json.encode(req) +end + +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok or resp == nil then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if type(resp.usage) == "table" 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 type(resp.choices) == "table" and #resp.choices > 0 then + local ch = resp.choices[1] + if type(ch.message) == "table" then + unified.content = ch.message.content or "" + if type(ch.message.tool_calls) == "table" 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 new file mode 100644 index 0000000..a931045 --- /dev/null +++ b/internal/lua/adapters/ollama.lua @@ -0,0 +1,63 @@ +local adapter = {} + +adapter.name = "ollama" +adapter.version = "2.0.0" +adapter.endpoint = "/api/chat" +adapter.headers = {} + +-- 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 = req.temperature or 0.7, + num_predict = req.max_tokens or 2048 + } + } + + -- 转换 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_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 new file mode 100644 index 0000000..648c9b6 --- /dev/null +++ b/internal/lua/adapters/openai.lua @@ -0,0 +1,80 @@ +local adapter = {} + +adapter.name = "openai" +adapter.version = "2.0.0" +adapter.endpoint = "/chat/completions" +adapter.headers = {} + +-- OpenAI /chat/completions format (pass-through, strip provider-specific fields) +function adapter.transform_request(raw_body) + local ok, req = pcall(json.decode, raw_body) + if not ok then return raw_body end + req.disable_thinking = nil + req.extra_body = nil + if req.messages then + for _, msg in ipairs(req.messages) do + msg.reasoning_content = nil + end + end + return json.encode(req) +end + +function adapter.transform_response(raw_body) + local ok, resp = pcall(json.decode, raw_body) + if not ok or resp == nil then return raw_body end + + local unified = { + content = "", + finish_reason = "", + token_usage = { prompt = 0, completion = 0, total = 0 } + } + + if type(resp.usage) == "table" 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 type(resp.choices) == "table" and #resp.choices > 0 then + local ch = resp.choices[1] + if type(ch.message) == "table" then + unified.content = ch.message.content or "" + if ch.message.reasoning_content then + unified.reasoning_content = ch.message.reasoning_content + end + if type(ch.message.tool_calls) == "table" 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/vm.go b/internal/lua/vm.go new file mode 100644 index 0000000..61583b9 --- /dev/null +++ b/internal/lua/vm.go @@ -0,0 +1,351 @@ +// Package lua implements the adapter runtime: bundles/loads *.lua adapter +// scripts, exposes json/string helpers, and lets Go call the protocol +// transform functions plus a build_headers signature hook. +package lua + +import ( + "crypto/hmac" + "crypto/sha256" + "embed" + "encoding/base64" + "encoding/hex" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sync" + + lua "github.com/yuin/gopher-lua" +) + +//go:embed adapters/*.lua +var bundledAdapters embed.FS + +type APIAdapter struct { + Name string `json:"name"` + Version string `json:"version"` +} + +type AdapterCache struct { + mu sync.RWMutex + state *lua.LState + items map[string]*lua.LTable +} + +func newAdapterCache() *AdapterCache { + return &AdapterCache{state: lua.NewState(), items: map[string]*lua.LTable{}} +} + +func (c *AdapterCache) setupGlobals() { + s := c.state + jsonTable := s.NewTable() + s.SetGlobal("json", jsonTable) + s.SetField(jsonTable, "encode", s.NewFunction(func(L *lua.LState) int { + b, err := jsonEncode(luaValueToGo(L.CheckAny(1))) + if err != nil { + L.Push(lua.LString("null")) + return 1 + } + L.Push(lua.LString(string(b))) + return 1 + })) + s.SetField(jsonTable, "decode", s.NewFunction(func(L *lua.LState) int { + v, err := jsonDecode(L.CheckString(1)) + if err != nil { + L.Push(lua.LNil) + return 1 + } + L.Push(goValueToLua(L, v)) + return 1 + })) + + // signature/crypto helpers (app verification, timing-safe auth) + s.SetGlobal("hmac_sha256_hex", s.NewFunction(func(L *lua.LState) int { + key := L.CheckString(1) + data := L.CheckString(2) + m := hmac.New(sha256.New, []byte(key)) + m.Write([]byte(data)) + L.Push(lua.LString(hex.EncodeToString(m.Sum(nil)))) + return 1 + })) + s.SetGlobal("sha256_hex", s.NewFunction(func(L *lua.LState) int { + h := sha256.Sum256([]byte(L.CheckString(1))) + L.Push(lua.LString(hex.EncodeToString(h[:]))) + return 1 + })) + s.SetGlobal("base64_encode", s.NewFunction(func(L *lua.LState) int { + L.Push(lua.LString(base64.StdEncoding.EncodeToString([]byte(L.CheckString(1))))) + return 1 + })) + s.SetGlobal("tohex", s.NewFunction(func(L *lua.LState) int { + L.Push(lua.LString(hex.EncodeToString([]byte(L.CheckString(1))))) + return 1 + })) + + s.SetGlobal("log", s.NewFunction(func(L *lua.LState) int { + level := L.ToString(1) + msg := L.ToString(2) + fmt.Printf("[adapter/%s] %s\n", level, msg) + return 0 + })) +} + +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)) +} + +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 +} + +func (c *AdapterCache) Get(name string) *lua.LTable { + c.mu.RLock() + defer c.mu.RUnlock() + return c.items[name] +} + +// Remove deletes an adapter from the cache. +func (c *AdapterCache) Remove(name string) { + c.mu.Lock() + defer c.mu.Unlock() + delete(c.items, name) +} + +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 +} + +// VM wraps AdapterCache to dispatch adapter hook calls safely (single Lua +// state is shared, so calls are serialized by a mutex). +type VM struct { + mu sync.Mutex + cache *AdapterCache + dir string +} + +func NewVM(dir string) *VM { + return &VM{dir: dir, cache: newAdapterCache()} +} + +func (v *VM) Start() error { + if v.dir != "" { + if err := os.MkdirAll(v.dir, 0755); err != nil { + return fmt.Errorf("mkdir adapter dir: %w", err) + } + if err := v.writeBundledAdapters(); err != nil { + return err + } + entries, err := os.ReadDir(v.dir) + if err != nil { + return err + } + for _, e := range entries { + if filepath.Ext(e.Name()) != ".lua" { + continue + } + if err := v.cache.Preload(filepath.Join(v.dir, e.Name())); err != nil { + fmt.Printf("[lua] preload %s: %v\n", e.Name(), err) + } + } + } + v.cache.setupGlobals() + return nil +} + +func (v *VM) Stop() { + v.mu.Lock() + defer v.mu.Unlock() + if v.cache.state != nil { + v.cache.state.Close() + v.cache.state = nil + } +} + +func (v *VM) ListAdapters() []APIAdapter { return v.cache.List() } + +// LoadAdapter compiles and registers an adapter from a file (runtime safe). +func (v *VM) LoadAdapter(path string) error { return v.cache.Preload(path) } + +// RemoveAdapter evicts an adapter from the cache (runtime safe). +func (v *VM) RemoveAdapter(name string) { v.cache.Remove(name) } + +func (v *VM) writeBundledAdapters() error { + known := []string{"openai", "anthropic", "deepseek", "gemini", "github", "groq", "mistral", "ollama", "kimicode"} + for _, name := range known { + dst := filepath.Join(v.dir, name+".lua") + if _, err := os.Stat(dst); err == nil { + continue + } + data, err := bundledAdapters.ReadFile("adapters/" + name + ".lua") + if err != nil { + continue + } + if err := os.WriteFile(dst, data, 0644); err != nil { + return err + } + } + return nil +} + +func (v *VM) Transform(name, fn, raw 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() + f := adapter.RawGetString(fn) + if f == nil || f == lua.LNil { + return "", fmt.Errorf("adapter %s missing %s", name, fn) + } + state := v.cache.state + state.Push(f) + state.Push(lua.LString(raw)) + if err := state.PCall(1, 1, nil); err != nil { + return "", fmt.Errorf("%s: %w", fn, err) + } + res := state.Get(-1) + state.Pop(1) + return res.String(), nil +} + +// BuildHeaders calls adapter.build_headers(meta). If the adapter does not +// define build_headers, it falls back to the static adapter.headers table. +func (v *VM) BuildHeaders(name string, meta map[string]interface{}) (map[string]string, error) { + adapter := v.cache.Get(name) + if adapter == nil { + return nil, fmt.Errorf("adapter %s not loaded", name) + } + v.mu.Lock() + defer v.mu.Unlock() + state := v.cache.state + + fn := adapter.RawGetString("build_headers") + if fn == nil || fn == lua.LNil { + // fall back to static headers table + headers := 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, nil + } + state.Push(fn) + state.Push(goValueToLua(state, meta)) + if err := state.PCall(1, 1, nil); err != nil { + return nil, fmt.Errorf("build_headers: %w", err) + } + res := state.Get(-1) + state.Pop(1) + headers := map[string]string{} + if tbl, ok := res.(*lua.LTable); ok { + tbl.ForEach(func(key, val lua.LValue) { + k := key.String() + if k != "" && k != "n" { + headers[k] = val.String() + } + }) + } + return headers, nil +} + +func (v *VM) Endpoint(name string) string { + adapter := v.cache.Get(name) + if adapter == nil { + return "" + } + if ep := adapter.RawGetString("endpoint"); ep != nil { + return ep.String() + } + return "" +} + +func jsonEncode(v interface{}) ([]byte, error) { return json.Marshal(v) } +func jsonDecode(s string) (interface{}, error) { + var v interface{} + if err := json.Unmarshal([]byte(s), &v); err != nil { + return nil, err + } + return v, nil +} + +func luaValueToGo(lv lua.LValue) interface{} { + switch x := lv.(type) { + case lua.LString: + return string(x) + case lua.LNumber: + return float64(x) + case lua.LBool: + return bool(x) + case *lua.LTable: + if x.MaxN() > 0 { + arr := make([]interface{}, 0, x.MaxN()) + x.ForEach(func(_, val lua.LValue) { arr = append(arr, luaValueToGo(val)) }) + return arr + } + m := map[string]interface{}{} + x.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 x := val.(type) { + case string: + return lua.LString(x) + case float64: + return lua.LNumber(x) + case int: + return lua.LNumber(x) + case int64: + return lua.LNumber(x) + case bool: + return lua.LBool(x) + case nil: + return lua.LNil + case []interface{}: + t := L.NewTable() + for i, item := range x { + t.RawSetInt(i+1, goValueToLua(L, item)) + } + return t + case map[string]interface{}: + t := L.NewTable() + for k, item := range x { + t.RawSetString(k, goValueToLua(L, item)) + } + return t + default: + return lua.LNil + } +} diff --git a/internal/lua/vm_test.go b/internal/lua/vm_test.go new file mode 100644 index 0000000..691b285 --- /dev/null +++ b/internal/lua/vm_test.go @@ -0,0 +1,87 @@ +package lua + +import ( + "strings" + "testing" +) + +func TestLoadBundledAdapters(t *testing.T) { + vm := NewVM(t.TempDir()) + if err := vm.Start(); err != nil { + t.Fatalf("start: %v", err) + } + defer vm.Stop() + adapters := vm.ListAdapters() + if len(adapters) == 0 { + t.Fatal("no adapters loaded") + } + names := map[string]bool{} + for _, a := range adapters { + names[a.Name] = true + } + for _, want := range []string{"openai", "deepseek", "anthropic", "gemini", "ollama", "kimicode"} { + if !names[want] { + t.Errorf("missing adapter %s (got %v)", want, names) + } + } +} + +func TestTransformRequest(t *testing.T) { + vm := NewVM(t.TempDir()) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + out, err := vm.Transform("openai", "transform_request", `{"model":"x","disable_thinking":true,"messages":[]}`) + if err != nil { + t.Fatalf("transform: %v", err) + } + if strings.Contains(out, "disable_thinking") { + t.Fatalf("disable_thinking not stripped: %s", out) + } +} + +func TestBuildHeadersFallbackStatic(t *testing.T) { + vm := NewVM(t.TempDir()) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + hdrs, err := vm.BuildHeaders("anthropic", nil) + if err != nil { + t.Fatalf("build headers: %v", err) + } + if hdrs["anthropic-version"] != "2023-06-01" { + t.Fatalf("static header missing: %v", hdrs) + } +} + +func TestBuildHeadersCustomHook(t *testing.T) { + vm := NewVM(t.TempDir()) + if err := vm.Start(); err != nil { + t.Fatal(err) + } + defer vm.Stop() + hdrs, err := vm.BuildHeaders("kimicode", map[string]interface{}{ + "timestamp": int64(12345), + "api_key": "k", + "body": "{}", + "method": "POST", + "url": "http://x/chat", + "source": map[string]interface{}{"meta": map[string]interface{}{ + "app_id": "app-9", "app_secret": "s", "api_key": "k", + }}, + }) + if err != nil { + t.Fatalf("build headers: %v", err) + } + if hdrs["X-App-Id"] != "app-9" { + t.Fatalf("x-app-id = %q", hdrs["X-App-Id"]) + } + if hdrs["X-App-Sign"] == "" { + t.Fatal("expected signature header") + } + if hdrs["X-Timestamp"] != "12345" { + t.Fatalf("timestamp = %q", hdrs["X-Timestamp"]) + } +} diff --git a/internal/provider/provider.go b/internal/provider/provider.go new file mode 100644 index 0000000..c8e56af --- /dev/null +++ b/internal/provider/provider.go @@ -0,0 +1,458 @@ +// Package provider binds a configured source + Lua adapter and performs the +// HTTP call / stream / image generation against the upstream LLM, with +// per-source concurrency limiting and availability backoff. +package provider + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "sync" + "time" + + "llmsproxy/internal/config" + "llmsproxy/internal/lua" + "llmsproxy/internal/types" +) + +// health tracks availability with exponential backoff. +type health struct { + failCount int + unavailableUntil time.Time + permanent bool +} + +func (h *health) reset() { h.failCount = 0; h.unavailableUntil = time.Time{}; h.permanent = false } + +func (h *health) available() bool { + if h.permanent { + return false + } + return time.Now().After(h.unavailableUntil) +} + +func (h *health) backoff() { + h.failCount++ + cooldown := 5 * time.Second * time.Duration(1<<(h.failCount-1)) + if cooldown > 30*time.Minute { + cooldown = 30 * time.Minute + } + h.unavailableUntil = time.Now().Add(cooldown) +} + +func (h *health) markPermanent() { + h.permanent = true + h.unavailableUntil = time.Time{} +} + +// Provider is a single configured upstream LLM source. +type Provider struct { + cfg config.Source + vm *lua.VM + adapter string + client *http.Client + + mu sync.Mutex + sem chan struct{} + health health +} + +func New(cfg config.Source, vm *lua.VM) *Provider { + p := &Provider{ + cfg: cfg, + vm: vm, + adapter: cfg.Adapter, + client: &http.Client{Timeout: cfg.Timeout}, + sem: make(chan struct{}, cfg.MaxConcurrent), + } + if cfg.MaxConcurrent <= 0 { + p.sem = nil + } + return p +} + +func (p *Provider) Name() string { return p.cfg.Name } +func (p *Provider) Adapter() string { return p.cfg.Adapter } +func (p *Provider) MaxConcurrent() int { return p.cfg.MaxConcurrent } +func (p *Provider) Config() *config.Source { return &p.cfg } + +// Models returns the model ids exposed by this source. +func (p *Provider) Models() []string { + out := make([]string, 0, len(p.cfg.Models)) + for _, m := range p.cfg.Models { + out = append(out, m.ID) + } + return out +} + +// ModelByID returns the model definition if owned by this source. +func (p *Provider) ModelByID(id string) *config.Model { + for i := range p.cfg.Models { + if p.cfg.Models[i].ID == id { + return &p.cfg.Models[i] + } + } + return nil +} + +// Endpoint resolves the upstream chat path. +func (p *Provider) Endpoint() string { + if p.cfg.Endpoint != "" { + return p.cfg.Endpoint + } + if ep := p.vm.Endpoint(p.adapter); ep != "" { + return ep + } + return "/chat/completions" +} + +// ImageEndpoint resolves the upstream image-generation path. +func (p *Provider) ImageEndpoint() string { + if p.cfg.ImageEndpoint != "" { + return p.cfg.ImageEndpoint + } + if ep := p.vm.Endpoint(p.adapter + "_image"); ep != "" { + return ep + } + return "/v1/images/generations" +} + +func (p *Provider) URL() string { + return strings.TrimRight(p.cfg.BaseURL, "/") + p.Endpoint() +} + +func (p *Provider) ImageURL() string { + return strings.TrimRight(p.cfg.BaseURL, "/") + p.ImageEndpoint() +} + +// ---- availability ---- + +func (p *Provider) Available() bool { + p.mu.Lock() + defer p.mu.Unlock() + return p.health.available() +} + +// ReportStatus records an upstream HTTP status for backoff decisions. +func (p *Provider) ReportStatus(code int) { + p.mu.Lock() + defer p.mu.Unlock() + if code == 401 || code == 403 { + p.health.markPermanent() + return + } + if code >= 500 || code == 429 { + p.health.backoff() + } +} + +func (p *Provider) reportError() { + p.mu.Lock() + p.health.backoff() + p.mu.Unlock() +} + +func (p *Provider) reportOK() { + p.mu.Lock() + p.health.reset() + p.mu.Unlock() +} + +// ---- concurrency limiting ---- + +// Acquire waits for a free concurrency slot (bounded by cfg.QueueTimeout), +// or context cancel. The HTTP call itself is not truncated. +func (p *Provider) Acquire(ctx context.Context) error { + if p.sem == nil { + return nil + } + var qCtx context.Context + var cancel context.CancelFunc + if p.cfg.QueueTimeout > 0 { + qCtx, cancel = context.WithTimeout(ctx, p.cfg.QueueTimeout) + } else { + qCtx, cancel = context.WithCancel(ctx) + } + defer cancel() + select { + case p.sem <- struct{}{}: + return nil + case <-qCtx.Done(): + return qCtx.Err() + } +} + +func (p *Provider) Release() { + if p.sem == nil { + return + } + <-p.sem +} + +// ---- request construction ---- + +func (p *Provider) buildHeaders(body, url string) (http.Header, error) { + meta := map[string]interface{}{ + "url": url, + "method": http.MethodPost, + "body": body, + "api_key": p.cfg.APIKey, + "timestamp": types.Now(), + "source": map[string]interface{}{ + "name": p.cfg.Name, + "meta": p.cfg.Meta, + }, + } + hdrs, err := p.vm.BuildHeaders(p.adapter, meta) + if err != nil { + return nil, err + } + h := http.Header{} + h.Set("Content-Type", "application/json") + for k, v := range p.cfg.Headers { + h.Set(k, v) + } + for k, v := range hdrs { + if _, ok := p.cfg.Headers[k]; !ok { + h.Set(k, v) + } + } + if h.Get("Authorization") == "" && p.cfg.APIKey != "" { + h.Set("Authorization", "Bearer "+p.cfg.APIKey) + } + return h, nil +} + +// ---- chat ---- + +// Chat performs a non-streaming round trip and returns the unified response. +func (p *Provider) Chat(ctx context.Context, req *types.ChatRequest) (*types.UnifiedResponse, error) { + if err := p.Acquire(ctx); err != nil { + return nil, err + } + defer p.Release() + + body, err := marshalTransform(p.vm, p.adapter, "transform_request", req) + if err != nil { + return nil, err + } + hdrs, err := p.buildHeaders(body, p.URL()) + if err != nil { + return nil, err + } + raw, status, err := p.do(ctx, p.URL(), body, hdrs) + if err != nil { + p.reportError() + return nil, err + } + if status != 200 { + p.ReportStatus(status) + return nil, fmt.Errorf("api error %d: %s", status, truncate(raw, 500)) + } + unified, err := p.vm.Transform(p.adapter, "transform_response", raw) + if err != nil { + return nil, err + } + var out types.UnifiedResponse + if err := json.Unmarshal([]byte(unified), &out); err != nil { + return nil, fmt.Errorf("unmarshal unified response: %w (body: %s)", err, unified) + } + p.reportOK() + return &out, nil +} + +// ChatStream performs a streaming round trip, emitting unified chunks. +func (p *Provider) ChatStream(ctx context.Context, req *types.ChatRequest) (<-chan types.UnifiedChunk, error) { + if err := p.Acquire(ctx); err != nil { + return nil, err + } + req.Stream = true + body, err := marshalTransform(p.vm, p.adapter, "transform_request", req) + if err != nil { + p.Release() + return nil, err + } + hdrs, err := p.buildHeaders(body, p.URL()) + if err != nil { + p.Release() + return nil, err + } + + type respOrErr struct { + resp *http.Response + err error + } + rc := make(chan respOrErr, 1) + go func() { + resp, err := p.doRaw(ctx, p.URL(), body, hdrs) + rc <- respOrErr{resp, err} + }() + + ch := make(chan types.UnifiedChunk, 64) + go func() { + defer p.Release() + defer close(ch) + sel := <-rc + if sel.err != nil { + p.reportError() + return + } + defer sel.resp.Body.Close() + if sel.resp.StatusCode != 200 { + raw, _ := io.ReadAll(sel.resp.Body) + p.ReportStatus(sel.resp.StatusCode) + _ = raw + return + } + scanner := bufio.NewScanner(sel.resp.Body) + scanner.Buffer(make([]byte, 0, 64*1024), 1024*1024) + for scanner.Scan() { + line := strings.TrimSpace(scanner.Text()) + if line == "" || !strings.HasPrefix(line, "data:") { + continue + } + data := strings.TrimSpace(strings.TrimPrefix(line, "data:")) + if data == "" { + continue + } + if data == "[DONE]" { + select { + case ch <- types.UnifiedChunk{Done: true}: + case <-ctx.Done(): + } + continue + } + unified, err := p.vm.Transform(p.adapter, "transform_stream_chunk", data) + if err != nil || unified == "" { + continue + } + if unified == data { + unified = standardSSEChunk(data) + if unified == "" { + continue + } + } + var ck types.UnifiedChunk + if err := json.Unmarshal([]byte(unified), &ck); err != nil { + continue + } + select { + case ch <- ck: + case <-ctx.Done(): + return + } + } + }() + return ch, nil +} + +// Image generates images via /v1/images/generations. +func (p *Provider) Image(ctx context.Context, req *types.ImageGenRequest) (*types.UnifiedResponse, error) { + if err := p.Acquire(ctx); err != nil { + return nil, err + } + defer p.Release() + + b, _ := json.Marshal(req) + transformed, err := p.vm.Transform(p.adapter+"_image", "transform_request", string(b)) + if err != nil { + // fall back to passthrough adapter (openai-style) + transformed = string(b) + } + hdrs, err := p.buildHeaders(transformed, p.ImageURL()) + if err != nil { + return nil, err + } + raw, status, err := p.do(ctx, p.ImageURL(), transformed, hdrs) + if err != nil { + p.reportError() + return nil, err + } + if status != 200 { + p.ReportStatus(status) + return nil, fmt.Errorf("image api error %d: %s", status, truncate(raw, 500)) + } + var out types.UnifiedResponse + // try adapter transform_response; if missing, parse standard openai image format + unified, terr := p.vm.Transform(p.adapter+"_image", "transform_response", raw) + if terr == nil && unified != raw { + if err := json.Unmarshal([]byte(unified), &out); err == nil { + p.reportOK() + return &out, nil + } + } + var img types.ImageGenResponse + if err := json.Unmarshal([]byte(raw), &img); err != nil { + return nil, fmt.Errorf("unmarshal image response: %w", err) + } + out.ImageData = img.Data + p.reportOK() + return &out, nil +} + +// ---- http helpers ---- + +func (p *Provider) do(ctx context.Context, url, body string, hdr http.Header) (string, int, error) { + resp, err := p.doRaw(ctx, url, body, hdr) + if err != nil { + return "", 0, err + } + defer resp.Body.Close() + raw, _ := io.ReadAll(resp.Body) + return string(raw), resp.StatusCode, nil +} + +func (p *Provider) doRaw(ctx context.Context, url, body string, hdr http.Header) (*http.Response, error) { + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader([]byte(body))) + if err != nil { + return nil, err + } + httpReq.Header = hdr + return p.client.Do(httpReq) +} + +func marshalTransform(vm *lua.VM, adapter, fn string, v interface{}) (string, error) { + b, err := json.Marshal(v) + if err != nil { + return "", err + } + out, err := vm.Transform(adapter, fn, string(b)) + if err != nil { + return "", err + } + return out, nil +} + +func standardSSEChunk(data string) string { + 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(data), &raw); err != nil { + return "" + } + if len(raw.Choices) == 0 { + return "" + } + out, _ := json.Marshal(types.UnifiedChunk{ + Content: raw.Choices[0].Delta.Content, + Done: raw.Choices[0].FinishReason != nil, + }) + return string(out) +} + +func truncate(s string, n int) string { + if len(s) <= n { + return s + } + return s[:n] + "..." +} \ No newline at end of file diff --git a/internal/provider/provider_test.go b/internal/provider/provider_test.go new file mode 100644 index 0000000..8587ac5 --- /dev/null +++ b/internal/provider/provider_test.go @@ -0,0 +1,172 @@ +package provider + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + "llmsproxy/internal/config" + "llmsproxy/internal/lua" + "llmsproxy/internal/types" +) + +func newTestProvider(t *testing.T, src config.Source) *Provider { + t.Helper() + vm := lua.NewVM(t.TempDir()) + if err := vm.Start(); err != nil { + t.Fatalf("vm: %v", err) + } + t.Cleanup(vm.Stop) + return New(src, vm) +} + +func src(name, url, adapter string, models ...string) config.Source { + s := config.Source{Name: name, BaseURL: url, Adapter: adapter, MaxConcurrent: 4} + for _, m := range models { + s.Models = append(s.Models, config.Model{ID: m, Priority: 0}) + } + return s +} + +func TestProviderChat(t *testing.T) { + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var body map[string]interface{} + _ = json.NewDecoder(r.Body).Decode(&body) + if body["model"] != "mock-model" { + t.Errorf("model = %v", body["model"]) + } + fmt.Fprint(w, `{"choices":[{"message":{"content":"hi"},"finish_reason":"stop"}]}`) + })) + defer up.Close() + p := newTestProvider(t, src("mock", up.URL, "openai", "mock-model")) + resp, err := p.Chat(context.Background(), &types.ChatRequest{ + Model: "mock-model", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hello")}}, + }) + if err != nil { + t.Fatalf("chat: %v", err) + } + if resp.Content != "hi" { + t.Fatalf("content = %q", resp.Content) + } +} + +func TestProviderChatStream(t *testing.T) { + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{\"content\":\"a\"}}]}\n\n") + fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{\"content\":\"b\"}}]}\n\n") + fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n") + fmt.Fprint(w, "data: [DONE]\n\n") + })) + defer up.Close() + p := newTestProvider(t, src("mock", up.URL, "openai", "m")) + ch, err := p.ChatStream(context.Background(), &types.ChatRequest{ + Model: "m", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}}, + }) + if err != nil { + t.Fatalf("stream: %v", err) + } + var text string + var done bool + for ck := range ch { + text += ck.Content + if ck.Done { + done = true + } + } + if text != "ab" { + t.Fatalf("text = %q", text) + } + if !done { + t.Fatal("expected done") + } +} + +func TestProviderImage(t *testing.T) { + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprint(w, `{"created":123,"data":[{"b64_json":"QUJD"}]}`) + })) + defer up.Close() + p := newTestProvider(t, src("img", up.URL, "openai", "flux-1")) + resp, err := p.Image(context.Background(), &types.ImageGenRequest{Model: "flux-1", Prompt: "cat"}) + if err != nil { + t.Fatalf("image: %v", err) + } + if len(resp.ImageData) != 1 || resp.ImageData[0].B64JSON != "QUJD" { + t.Fatalf("image data = %+v", resp.ImageData) + } +} + +func TestProviderBackoff(t *testing.T) { + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(500) + fmt.Fprint(w, "boom") + })) + defer up.Close() + p := newTestProvider(t, src("mock", up.URL, "openai", "m")) + _, err := p.Chat(context.Background(), &types.ChatRequest{ + Model: "m", + Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}}, + }) + if err == nil { + t.Fatal("expected error") + } + if p.Available() { + t.Fatal("expected provider to be in backoff") + } + // 401 -> permanent + up2 := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(401) + })) + defer up2.Close() + p2 := newTestProvider(t, src("mock2", up2.URL, "openai", "m2")) + p2.Chat(context.Background(), &types.ChatRequest{Model: "m2", Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}}}) + if p2.Available() { + t.Fatal("expected permanent unavailability on 401") + } +} + +func TestProviderConcurrencyCap(t *testing.T) { + release := make(chan struct{}) + started := make(chan struct{}, 100) + up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + started <- struct{}{} + <-release + fmt.Fprint(w, `{"choices":[{"message":{"content":"ok"}}]}`) + })) + defer up.Close() + // cap 2 + p := newTestProvider(t, src("mock", up.URL, "openai", "m")) + p.cfg.MaxConcurrent = 2 + p.sem = make(chan struct{}, 2) + + var wg sync.WaitGroup + for i := 0; i < 6; i++ { + wg.Add(1) + go func() { + defer wg.Done() + p.Chat(context.Background(), &types.ChatRequest{Model: "m", Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}}}) + }() + } + // wait until 2 requests started + deadline := time.Now().Add(2 * time.Second) + for len(started) < 2 { + if time.Now().After(deadline) { + t.Fatal("timeout waiting for first two") + } + time.Sleep(5 * time.Millisecond) + } + time.Sleep(100 * time.Millisecond) + if len(started) > 2 { + t.Fatalf("more than 2 concurrent: %d", len(started)) + } + close(release) + wg.Wait() +} \ No newline at end of file diff --git a/internal/provider/registry.go b/internal/provider/registry.go new file mode 100644 index 0000000..13b5db9 --- /dev/null +++ b/internal/provider/registry.go @@ -0,0 +1,170 @@ +// Package provider also provides the top-level registry that owns all sources, +// routes model requests (explicit or AUTO), and supports hot reload. +package provider + +import ( + "sort" + "strings" + "sync" +) + +// Registry holds all configured providers and routes model requests. +type Registry struct { + mu sync.RWMutex + providers []*Provider + byModel map[string]*Provider // modelID -> provider + defaultM string // default model id ("" means AUTO) +} + +func NewRegistry(providers []*Provider, defaultModel string) *Registry { + r := &Registry{byModel: map[string]*Provider{}, defaultM: defaultModel} + r.set(providers) + return r +} + +// Replace atomically swaps the provider set (hot reload). +func (r *Registry) Replace(providers []*Provider) { + r.mu.Lock() + defer r.mu.Unlock() + r.set(providers) +} + +func (r *Registry) set(providers []*Provider) { + r.providers = providers + r.byModel = map[string]*Provider{} + for _, p := range providers { + for _, m := range p.Models() { + r.byModel[strings.ToLower(m)] = p + } + } +} + +func (r *Registry) Providers() []*Provider { + r.mu.RLock() + defer r.mu.RUnlock() + out := make([]*Provider, len(r.providers)) + copy(out, r.providers) + return out +} + +func (r *Registry) Count() int { + r.mu.RLock() + defer r.mu.RUnlock() + return len(r.providers) +} + +// ModelList returns all exposed model ids (chat + image). +func (r *Registry) ModelList() []string { + r.mu.RLock() + defer r.mu.RUnlock() + seen := map[string]bool{} + for _, p := range r.providers { + for _, m := range p.Models() { + if !seen[m] { + seen[m] = true + } + } + } + out := make([]string, 0, len(seen)) + for m := range seen { + out = append(out, m) + } + sort.Strings(out) + return out +} + +// Resolve returns the ordered candidate providers to try for a request, +// honoring explicit model selection or AUTO (priority order, healthy first). +// +// model "" or "AUTO" -> all sources sorted by (priority desc, healthy first). +// Otherwise the owning provider, if healthy; else its source anyway. +func (r *Registry) Resolve(model string) []*Provider { + r.mu.RLock() + defer r.mu.RUnlock() + + model = strings.TrimSpace(model) + if model == "" || strings.EqualFold(model, "AUTO") { + // priority chain across all models + type cand struct { + prov *Provider + priority int + } + var cands []cand + seen := map[string]bool{} + for _, p := range r.providers { + prio := -1 + for _, m := range p.cfg.Models { + if m.Priority > prio { + prio = m.Priority + } + } + if prio < 0 { + prio = 0 + } + cands = append(cands, cand{p, prio}) + seen[p.Name()] = true + } + sort.SliceStable(cands, func(i, j int) bool { + if cands[i].priority != cands[j].priority { + return cands[i].priority > cands[j].priority + } + // healthy preferred at same priority + return cands[i].prov.Available() && !cands[j].prov.Available() + }) + out := make([]*Provider, 0, len(cands)) + for _, c := range cands { + out = append(out, c.prov) + } + return out + } + + // explicit model + if p, ok := r.byModel[strings.ToLower(model)]; ok { + // switch to the owning source but pin the model via request + return []*Provider{p} + } + // unknown model -> fall back to default/AUTO chain + return r.AUTOChain() +} + +// AUTOChain returns the priority-sorted providers for AUTO. +func (r *Registry) AUTOChain() []*Provider { + return r.Resolve("AUTO") +} + +// Default returns the highest-priority available provider. +func (r *Registry) Default() *Provider { + chain := r.AUTOChain() + if len(chain) == 0 { + return nil + } + return chain[0] +} + +// ModelStatus is a web-UI friendly snapshot per source. +type SourceStatus struct { + Name string `json:"name"` + Adapter string `json:"adapter"` + Models []string `json:"models"` + Available bool `json:"available"` + Healthy bool `json:"healthy"` + MaxConcurrent int `json:"max_concurrent"` +} + +func (r *Registry) Status() []SourceStatus { + r.mu.RLock() + defer r.mu.RUnlock() + out := make([]SourceStatus, 0, len(r.providers)) + for _, p := range r.providers { + s := SourceStatus{ + Name: p.Name(), + Adapter: p.Adapter(), + Models: p.Models(), + Available: p.Available(), + Healthy: p.Available(), + MaxConcurrent: p.MaxConcurrent(), + } + out = append(out, s) + } + return out +} \ No newline at end of file diff --git a/internal/scheduler/scheduler.go b/internal/scheduler/scheduler.go new file mode 100644 index 0000000..f4361c6 --- /dev/null +++ b/internal/scheduler/scheduler.go @@ -0,0 +1,107 @@ +// Package scheduler implements request scheduling across providers: per-source +// concurrency caps (acquire with wait = queuing), AUTO model fallback chains, +// and exponential backoff via provider health. +package scheduler + +import ( + "context" + "fmt" + + "llmsproxy/internal/provider" + "llmsproxy/internal/types" +) + +// Scheduler drives one chat tool call across the candidate provider chain. +type Scheduler struct { + // MaxRetries how many fallback providers to try before failing. + MaxRetries int +} + +func New(maxRetries int) *Scheduler { + if maxRetries < 0 { + maxRetries = 0 + } + return &Scheduler{MaxRetries: maxRetries} +} + +// Provider is the minimal interface the scheduler needs to schedule over. +type Provider interface { + Name() string + Available() bool + Chat(ctx context.Context, req *types.ChatRequest) (*types.UnifiedResponse, error) + ChatStream(ctx context.Context, req *types.ChatRequest) (<-chan types.UnifiedChunk, error) + Image(ctx context.Context, req *types.ImageGenRequest) (*types.UnifiedResponse, error) +} + +// FromRegistry converts *provider.Provider slices to the scheduler interface. +func FromRegistry(ps []*provider.Provider) []Provider { + out := make([]Provider, len(ps)) + for i, p := range ps { + out[i] = p + } + return out +} + +// Chat runs a chat request across cands, falling back on failure. +func (s *Scheduler) Chat(ctx context.Context, cands []Provider, req *types.ChatRequest) (*types.UnifiedResponse, error) { + attempts := s.MaxRetries + 1 + var lastErr error + for i := 0; i < attempts && i < len(cands); i++ { + p := cands[i] + resp, err := p.Chat(ctx, req) + if ctx.Err() != nil { + return nil, ctx.Err() + } + if err == nil { + return resp, nil + } + lastErr = fmt.Errorf("provider %s: %w", p.Name(), err) + } + if lastErr == nil { + // if loop couldn't run because cands was short but no error recorded yet + if len(cands) == 0 { + return nil, fmt.Errorf("no provider available") + } + // should not happen + return nil, lastErr + } + return nil, lastErr +} + +// ChatStream runs a streaming chat across cands, falling back early on connect errors. +func (s *Scheduler) ChatStream(ctx context.Context, cands []Provider, req *types.ChatRequest) (<-chan types.UnifiedChunk, error) { + attempts := s.MaxRetries + 1 + var lastErr error + for i := 0; i < attempts && i < len(cands); i++ { + p := cands[i] + resp, err := p.ChatStream(ctx, req) + if err == nil { + return resp, nil + } + lastErr = fmt.Errorf("provider %s: %w", p.Name(), err) + } + if lastErr == nil { + if len(cands) == 0 { + return nil, fmt.Errorf("no provider available") + } + } + return nil, lastErr +} + +// Image runs an image-generation request across cands. +func (s *Scheduler) Image(ctx context.Context, cands []Provider, req *types.ImageGenRequest) (*types.UnifiedResponse, error) { + attempts := s.MaxRetries + 1 + var lastErr error + for i := 0; i < attempts && i < len(cands); i++ { + p := cands[i] + resp, err := p.Image(ctx, req) + if err == nil { + return resp, nil + } + lastErr = fmt.Errorf("provider %s: %w", p.Name(), err) + } + if lastErr == nil && len(cands) == 0 { + return nil, fmt.Errorf("no provider available") + } + return nil, lastErr +} \ No newline at end of file diff --git a/internal/types/types.go b/internal/types/types.go new file mode 100644 index 0000000..b0f78ea --- /dev/null +++ b/internal/types/types.go @@ -0,0 +1,120 @@ +// Package types defines the unified (OpenAI-compatible) wire format that the +// gateway exposes to its clients, plus the unified internal representation. +package types + +import ( + "encoding/json" + "time" +) + +// ---- OpenAI wire request (gateway input) ---- + +type ChatRequest struct { + Model string `json:"model"` + Messages []ChatMessage `json:"messages"` + Temperature *float64 `json:"temperature,omitempty"` + MaxTokens int `json:"max_tokens,omitempty"` + Stream bool `json:"stream,omitempty"` + Tools []interface{} `json:"tools,omitempty"` + ToolChoice interface{} `json:"tool_choice,omitempty"` + DisableThinking bool `json:"disable_thinking"` + ExtraBody map[string]interface{} `json:"-"` +} + +func (r *ChatRequest) MarshalJSON() ([]byte, error) { + type Alias ChatRequest + data, err := json.Marshal((*Alias)(r)) + if err != nil { + return nil, err + } + if len(r.ExtraBody) == 0 { + return data, nil + } + var raw map[string]interface{} + if err := json.Unmarshal(data, &raw); err != nil { + return nil, err + } + for k, v := range r.ExtraBody { + raw[k] = v + } + return json.Marshal(raw) +} + +// ChatMessage supports both plain string content and multimodal arrays +// (RawMessage preserves whatever the client sent for the adapter to process). +type ChatMessage struct { + Role string `json:"role"` + Content json.RawMessage `json:"content,omitempty"` + ReasoningContent string `json:"reasoning_content,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` +} + +func StringContent(s string) json.RawMessage { b, _ := json.Marshal(s); return b } + +type ToolCall struct { + ID string `json:"id"` + Type string `json:"type"` + Name string `json:"name"` + Arguments map[string]interface{} `json:"arguments"` +} + +// ---- Unified internal representation (what adapters produce) ---- + +type UnifiedResponse struct { + Content string `json:"content"` + ReasoningContent string `json:"reasoning_content,omitempty"` + FinishReason string `json:"finish_reason,omitempty"` + TokenUsage TokenUsage `json:"token_usage"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + // ImageData used by image-generation adapters. + ImageData []ImageData `json:"image_data,omitempty"` +} + +type TokenUsage struct { + Prompt int `json:"prompt"` + Completion int `json:"completion"` + Total int `json:"total"` +} + +type ImageData struct { + B64JSON string `json:"b64_json,omitempty"` + URL string `json:"url,omitempty"` + Revised string `json:"revised_prompt,omitempty"` +} + +// ---- Image generation (OpenAI /v1/images/generations wire) ---- + +type ImageGenRequest struct { + Model string `json:"model"` + Prompt string `json:"prompt"` + N int `json:"n,omitempty"` + Size string `json:"size,omitempty"` + ResponseFormat string `json:"response_format,omitempty"` +} + +type ImageGenResponse struct { + Created int64 `json:"created"` + Data []ImageData `json:"data"` +} + +// ---- Unified streaming chunk produced by adapters ---- + +type UnifiedChunk struct { + Content string `json:"content"` + Done bool `json:"done"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + ReasoningContent string `json:"reasoning_content,omitempty"` +} + +// Meta passed to Lua build_headers hook +type BuildMeta struct { + URL string `json:"url"` + Method string `json:"method"` + Body string `json:"body"` + APIKey string `json:"api_key"` + Timestamp int64 `json:"timestamp"` + Source map[string]interface{} `json:"source"` +} + +func Now() int64 { return time.Now().Unix() } \ No newline at end of file