diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..0ff38f4 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,6 @@ +.git +.github +docs +*.md +tmp/ +data/ diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 8621ae8..610aff5 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -28,5 +28,6 @@ jobs: - name: Create release uses: softprops/action-gh-release@v2 with: + name: "Strait ${{ github.ref_name }}" files: build/* generate_release_notes: true diff --git a/.gitignore b/.gitignore index defe4d7..bc11b29 100644 --- a/.gitignore +++ b/.gitignore @@ -19,6 +19,7 @@ # 环境变量 .env .env.local +k8s/secret.yaml # 依赖缓存 vendor/ diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..37f5bb3 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,17 @@ +# 构建阶段 +FROM golang:1.26-alpine AS builder +WORKDIR /app +COPY go.mod go.sum ./ +RUN go mod download +COPY . . +RUN CGO_ENABLED=0 go build -o strait ./cmd/strait/ + +# 运行阶段 +FROM alpine:3.20 +RUN apk add --no-cache ca-certificates +WORKDIR /app +COPY --from=builder /app/strait . +COPY --from=builder /app/configs ./configs + +EXPOSE 8080 +CMD ["./strait"] \ No newline at end of file diff --git a/api/plugin.go b/api/plugin.go index 4da1878..649e598 100644 --- a/api/plugin.go +++ b/api/plugin.go @@ -29,6 +29,21 @@ type ChatAdapter interface { SendChatStream(ctx context.Context, payload *ChatRequest, route *RouteDecision) (<-chan *StreamChunk, error) } +type Guard interface { + Plugin + Guard(pctx *PipelineContext) error +} + +type PreProcessor interface { + Plugin + PreProcess(pctx *PipelineContext) error +} + +type PostProcessor interface { + Plugin + PostProcess(pctx *PipelineContext) error +} + // Constructor 插件构造函数 type Constructor func() Plugin diff --git a/api/types.go b/api/types.go index cb9ada4..54282f8 100644 --- a/api/types.go +++ b/api/types.go @@ -89,9 +89,10 @@ type ToolCallFunction struct { // RouteDecision 路由决策 type RouteDecision struct { - Protocol string `json:"protocol"` // 调用协议:openAI / anthropic / deepseek / ollama - BaseURL string `json:"base_url"` // 调用地址 - APIKey string `json:"api_key"` // 认证密钥 + Protocol string `json:"protocol"` // 调用协议:openAI / anthropic / deepseek / ollama + BaseURL string `json:"base_url"` // 调用地址 + APIKey string `json:"api_key"` // 认证密钥 + Model string `json:"model,omitempty"` // 目标模型名称 } // Subject 认证后的调用方信息 @@ -99,6 +100,17 @@ type Subject struct { ID string `json:"id"` // 调用方唯一标识 } +// ─ 管线 ─ + +// PipelineContext 管线上下文,贯穿 guard → preprocess → route → adapter → postprocess 全流程 +type PipelineContext struct { + Context context.Context // 请求上下文 + Request any // 请求体:*ChatRequest / *EmbeddingRequest / ... + Response any // 响应体:*ChatResponse / *EmbeddingResponse / ... + Route *RouteDecision // 路由决策 + Meta map[string]any // 插件间传递的元数据 +} + // ─ Context ─ type ctxKey string diff --git a/cmd/strait/main.go b/cmd/strait/main.go index 8b71af4..c30fc81 100644 --- a/cmd/strait/main.go +++ b/cmd/strait/main.go @@ -1,31 +1,59 @@ -// Strait — AI 代理网关入口。 +// main Strait — AI 代理网关入口。 package main import ( "context" - "log" + "errors" + "fmt" + "log/slog" "net/http" + "os" + "os/signal" + "syscall" + "time" + "strait/internal/app" "strait/internal/config" - "strait/internal/hotreload" - - "strait/internal/app" + "strait/internal/metrics" "strait/internal/plugin" "strait/internal/router" _ "strait/plugins/adapter-ollama" _ "strait/plugins/adapter-openai" _ "strait/plugins/auth-static-token" + _ "strait/plugins/prompt-injector" ) +const banner = ` +███████╗████████╗██████╗ █████╗ ██╗████████╗ +██╔════╝╚══██╔══╝██╔══██╗██╔══██╗██║╚══██╔══╝ +███████╗ ██║ ██████╔╝███████║██║ ██║ +╚════██║ ██║ ██╔══██╗██╔══██║██║ ██║ +███████║ ██║ ██║ ██║██║ ██║██║ ██║ +╚══════╝ ╚═╝ ╚═╝ ╚═╝╚═╝ ╚═╝╚═╝ ╚═╝ v0.2 +` + func main() { + slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{ + Level: slog.LevelInfo, + }))) + loader := plugin.NewLoader(config.PluginsPath) m, err := loader.Build() if err != nil { - log.Fatal(err) + slog.Error("startup failed", "error", err) + os.Exit(1) + } + + if os.Getenv("STRAIT_BANNER") != "false" { + fmt.Print(banner) + fmt.Println(" Plugins:") + fmt.Print(m.Summary()) + fmt.Println() } scheduler := plugin.NewScheduler(m) + met := metrics.NewMetrics() reload := func(filename string) error { if filename == "plugins.yaml" { @@ -49,7 +77,30 @@ func main() { _ = watcher.Start(context.Background()) }() - server := app.NewServer(scheduler) - log.Println("strait listening on :8080") - log.Fatal(http.ListenAndServe(":8080", server.Handler())) + // 启动 HTTP 服务 + server := app.NewServer(scheduler, met) + srv := &http.Server{ + Addr: ":8080", + Handler: server.Handler(), + } + + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) + defer stop() + + go func() { + slog.Info("strait listening on :8080") + if err := srv.ListenAndServe(); err != nil && !errors.Is(err, http.ErrServerClosed) { + slog.Error("server crashed", "error", err) + os.Exit(1) + } + }() + + <-ctx.Done() + slog.Info("shutting down...") + shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if err := srv.Shutdown(shutdownCtx); err != nil { + slog.Error("shutdown failed", "error", err) + } + slog.Info("strait stopped") } diff --git a/configs/plugins.yaml b/configs/plugins.yaml index 5f359d5..f8aa2b8 100644 --- a/configs/plugins.yaml +++ b/configs/plugins.yaml @@ -9,4 +9,8 @@ plugins: token: sk-admin-init subject: admin - id: adapter-ollama - type: adapter \ No newline at end of file + type: adapter + - id: prompt-injector + type: preprocessor + config: + system_prompt: "You are a helpful assistant." \ No newline at end of file diff --git a/configs/routes.yaml b/configs/routes.yaml index 6af6468..80baa4b 100644 --- a/configs/routes.yaml +++ b/configs/routes.yaml @@ -2,9 +2,19 @@ routes: - id: deepseek-chat match: model: deepseek-chat + strategy: weight target: provider: deepseek-main - model: deepseek-chat + model: deepseek-reasoner + targets: + - provider: deepseek-main + model: deepseek-chat + priority: 1 + weight: 3 + - provider: ollama-local + model: qwen2.5:0.5b + priority: 1 + weight: 1 - id: deepseek-reasoner match: @@ -16,6 +26,9 @@ routes: - id: ollama-qwen match: model: qwen2.5:0.5b - target: - provider: ollama-local - model: qwen2.5:0.5b + strategy: priority + targets: + - provider: ollama-local + model: qwen2.5:0.5b + priority: 1 + weight: 1 diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..5d0236c --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,8 @@ +services: + strait: + build: . + ports: + - "8080:8080" + volumes: + - ./configs:/app/configs + restart: unless-stopped \ No newline at end of file diff --git a/docs/deployment.md b/docs/deployment.md index b625f4a..16c9299 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -1,20 +1,97 @@ # 部署指南 -## MySQL +## Docker + +### 构建镜像 ```bash -docker-compose up -d +docker build -t strait . +``` + +### 运行 + +```bash +docker run -d \ + -p 8080:8080 \ + -e DEEPSEEK_API_KEY=your-key \ + -v $(pwd)/configs:/app/configs \ + strait ``` -## Docker Compose +### Docker Compose ```bash +# 设置环境变量 +export DEEPSEEK_API_KEY=your-key + +# 启动 docker-compose up -d ``` +`docker-compose.yml` 会自动挂载 `configs/` 目录并注入环境变量。 + +--- + ## Kubernetes +### 1. 创建 Secret(存放 API Key) + +参考 `k8s/secret.example.yaml`: + +```bash +kubectl create secret generic strait-secret \ + --from-literal=DEEPSEEK_API_KEY=your-key +``` + +### 2. 创建 ConfigMap(挂载配置文件) + +```bash +kubectl create configmap strait-config \ + --from-file=configs/ +``` + +### 3. 部署 + ```bash -helm repo add strait https://yourorg.github.io/helm-charts -helm install strait strait/strait +kubectl apply -f k8s/service.yaml +kubectl apply -f k8s/deployment.yaml ``` + +### 4. 验证 + +```bash +# 检查 Pod 状态 +kubectl get pods -l app=strait + +# 健康检查 +kubectl port-forward svc/strait 8080:8080 +curl http://localhost:8080/health +``` + +### 探针说明 + +| 探针 | 路径 | 说明 | +|------|------|------| +| liveness | `/health` | 存活检测,失败则重启容器 | +| readiness | `/ready` | 就绪检测,失败则从 Service 摘除 | + +--- + +## 监控 + +### Prometheus Metrics + +服务暴露 `/metrics` 端点,返回 Prometheus 格式的请求计数: + +```bash +curl http://localhost:8080/metrics +``` + +--- + +## 配置热重载 + +修改 `configs/` 目录下的 YAML 文件后自动重载,无需重启: + +- `providers.yaml` / `routes.yaml` → 路由重载 +- `plugins.yaml` → 插件重载 diff --git a/docs/plugins/built-in.md b/docs/plugins/built-in.md index 364df3a..9febf96 100644 --- a/docs/plugins/built-in.md +++ b/docs/plugins/built-in.md @@ -15,7 +15,7 @@ Strait 提供 4 个内置插件,启动即用。 ## router-yaml -根据 `routes.yaml` 中的路由规则,将请求匹配到对应的 provider。 +根据 `routes.yaml` 中的路由规则,将请求匹配到对应的 provider。支持多目标路由、优先级/权重策略选择。 ```yaml # plugins.yaml @@ -24,6 +24,8 @@ plugins: type: router ``` +### 单目标(兼容旧格式) + ```yaml # routes.yaml routes: @@ -35,8 +37,43 @@ routes: model: deepseek-chat # 实际请求上游的模型名 ``` +### 多目标路由 + +一个路由规则可以指向多个 provider,通过策略选择最终目标: + +```yaml +routes: + - id: deepseek-chat + match: + model: deepseek-chat + strategy: priority # 策略:priority(优先级)/ weight(权重) + targets: + - provider: deepseek-main # 优先级数字越小越优先 + model: deepseek-chat + priority: 1 + weight: 3 # 同优先级内按权重随机 + - provider: deepseek-backup + model: deepseek-chat + priority: 2 + weight: 1 +``` + +### 路由策略 + +| 策略 | 说明 | +|------|------| +| `priority` | 按优先级选择,数字越小越优先(默认) | +| `weight` | 按权重随机选择,权重越大被选中概率越高 | + +### 模型映射 + +`targets` 中的 `model` 字段可指定目标 provider 的实际模型名,实现请求模型到上游模型的映射。例如请求 `model: gpt-4` 可映射到 `deepseek-chat`。 + +### 特性 + - **无需额外配置** — 直接声明 `type: router` 即可加载 - **支持热重载** — 修改 `routes.yaml` 后自动生效 +- **兼容旧格式** — 单目标 `target` 字段仍然有效 --- diff --git a/docs/roadmap.md b/docs/roadmap.md index 6a07ce2..92a5389 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -21,10 +21,10 @@ | **P1 插件加载器** | ✅ | Loader + 自动注册 + 配置驱动 + 启动日志 | | **P2 生产可用性** | ✅ | 鉴权 + 错误模型 + 热重载 + 统一内部模型 + OpenAI 输出格式 | | **P3 Agent 协议** | ✅ | Function Calling + Tool Use + Tool Calls 响应 | -| **P4 多供应商** | 🚧 | Ollama 适配 ✅ / Anthropic + Gemini 适配 ⏳ / 负载均衡 + 熔断 ⏳ | -| **P5 管道扩展** | ⏳ | PreProcessor / PostProcessor 接口 + 模型分组路由 | +| **P4 多供应商** | 🚧 | Ollama 适配 ✅ / 负载均衡(优先级/权重)✅ / Anthropic + Gemini 适配 ⏳ / 熔断 ⏳ | +| **P5 管道扩展** | 🚧 | Guard / PreProcessor / PostProcessor 接口 ✅ + 模型分组路由 ⏳ | | **P6 持久化** | ⏳ | SQLite + Repository 层 + 管理 API CRUD + Playground | -| **P7 生产部署** | ⏳ | Docker + 优雅关闭 + Prometheus + 限流 | +| **P7 生产部署** | 🚧 | Docker ✅ + K8s 部署 ✅ + Prometheus /metrics ✅ + 优雅关闭 ✅ + 限流 ⏳ | | **P8 扩展协议** | ⏳ | Embedding 透传 / MCP 端点 / Rerank 适配 | ### 内核关键里程碑 @@ -45,10 +45,11 @@ Strait 所有组成部分都是插件——Router、Adapter、Authenticator、Gu | 插件 | 类型 | 状态 | 说明 | |------|------|------|------| -| router-yaml | Router | ✅ | YAML 配置路由匹配 | +| router-yaml | Router | ✅ | YAML 配置路由匹配 + 多目标 + 优先级/权重策略 | | adapter-openai | Adapter | ✅ | OpenAI 兼容协议适配 | | adapter-ollama | Adapter | ✅ | Ollama 本地模型适配 | | auth-static-token | Authenticator | ✅ | 静态 Token 鉴权 | +| prompt-injector | PreProcessor | ✅ | 系统提示词自动注入 | ### 开发中 @@ -57,7 +58,6 @@ Strait 所有组成部分都是插件——Router、Adapter、Authenticator、Gu | adapter-anthropic | Adapter | P4 | Claude 协议适配 | | adapter-gemini | Adapter | P4 | Gemini 协议适配 | | rate-limiter | Guard | P5 | 令牌桶 / 滑动窗口限流 | -| prompt-injector | PreProcessor | P5 | 系统提示词自动注入 | | audit-logger | PostProcessor | P5 | 请求-响应审计记录 | | cost-tracker | PostProcessor | P5 | Token 用量和成本统计 | diff --git a/internal/app/server.go b/internal/app/server.go index 4415e6d..d8b9d01 100644 --- a/internal/app/server.go +++ b/internal/app/server.go @@ -10,6 +10,8 @@ import ( "strings" "time" + "strait/internal/metrics" + "strait/api" "strait/internal/plugin" ) @@ -18,11 +20,12 @@ import ( // (辅助理解)相当于 Java 的 @RestController + @Autowired type Server struct { scheduler *plugin.Scheduler + metrics *metrics.Metrics } // NewServer 创建 HTTP 服务 -func NewServer(s *plugin.Scheduler) *Server { - return &Server{scheduler: s} +func NewServer(s *plugin.Scheduler, m *metrics.Metrics) *Server { + return &Server{scheduler: s, metrics: m} } // Handler 注册路由,返回 http.Handler @@ -31,6 +34,12 @@ func (s *Server) Handler() http.Handler { mux.HandleFunc("GET /health", s.healthHandler) mux.HandleFunc("GET /ready", s.readyHandler) mux.HandleFunc("POST /v1/chat/completions", s.chatHandler) + mux.HandleFunc("GET /metrics", func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/plain; version=0.0.4") + if _, err := fmt.Fprint(w, s.metrics.Handler()); err != nil { + slog.Error("write metrics failed", "error", err) + } + }) return mux } @@ -45,6 +54,7 @@ func (s *Server) readyHandler(w http.ResponseWriter, _ *http.Request) { func (s *Server) chatHandler(w http.ResponseWriter, r *http.Request) { start := time.Now() reqID := generateReqID() + s.metrics.IncRequests("POST", "/v1/chat/completions") var req api.ChatRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { diff --git a/internal/hotreload/watcher.go b/internal/hotreload/watcher.go index a8d1943..a629967 100644 --- a/internal/hotreload/watcher.go +++ b/internal/hotreload/watcher.go @@ -50,15 +50,14 @@ func (w *Watcher) Start(ctx context.Context) error { } timer = time.AfterFunc(500*time.Millisecond, func() { if err := w.reload(filepath.Base(name)); err != nil { - slog.Error("HotReload failed", "error", err) + slog.Error("hot reload failed", "error", err) } else { slog.Info("config reloaded") } }) case err := <-watcher.Errors: - slog.Error("HotReload watcher error", "error", err) - + slog.Error("hot reload watcher error", "error", err) case <-ctx.Done(): return nil } diff --git a/internal/metrics/metrics.go b/internal/metrics/metrics.go new file mode 100644 index 0000000..7c870e3 --- /dev/null +++ b/internal/metrics/metrics.go @@ -0,0 +1,39 @@ +package metrics + +import ( + "fmt" + "strings" + "sync" +) + +type Metrics struct { + mu sync.Mutex // 请求锁 + requests map[string]int64 // 请求次数(key: "method|path" -> count) +} + +func NewMetrics() *Metrics { + return &Metrics{ + requests: make(map[string]int64), + } +} + +func (m *Metrics) IncRequests(method, path string) { + m.mu.Lock() + defer m.mu.Unlock() + key := method + "|" + path + m.requests[key]++ +} + +func (m *Metrics) Handler() string { + m.mu.Lock() + defer m.mu.Unlock() + var b strings.Builder + b.WriteString("# HELP strait_requests_total Total HTTP requests\n") + b.WriteString("# TYPE strait_requests_total counter\n") + for key, count := range m.requests { + parts := strings.SplitN(key, "|", 2) + b.WriteString(fmt.Sprintf("strait_requests_total{method=\"%s\",path=\"%s\"} %d\n", + parts[0], parts[1], count)) + } + return b.String() +} diff --git a/internal/plugin/loader.go b/internal/plugin/loader.go index aebef47..cd6c94f 100644 --- a/internal/plugin/loader.go +++ b/internal/plugin/loader.go @@ -32,6 +32,7 @@ func NewLoader(path string) *Loader { return &Loader{configPath: path} } +// Build 构建插件管理器 func (l *Loader) Build() (*Manager, error) { slog.Info("starting plugin loader", "config_path", l.configPath) @@ -51,7 +52,10 @@ func (l *Loader) Build() (*Manager, error) { var router api.Router for _, entry := range cfg.Plugins { if err := l.loadPlugin(entry, m, &router); err != nil { - return nil, err + if isCriticalType(entry.Type) { + return nil, err + } + slog.Warn("optional plugin skipped", "id", entry.ID, "type", entry.Type, "error", err) } } @@ -59,37 +63,71 @@ func (l *Loader) Build() (*Manager, error) { m.SetRouter(router) } - slog.Info( - "all plugins loaded successfully", - "router_loaded", router != nil, - "authenticator_count", len(m.authenticators), - "adapter_count", len(m.adapters), - ) return m, nil } +// isCriticalType 判断插件类型是否为关键类型,加载失败必须退出 +func isCriticalType(t string) bool { + switch t { + case "router", "adapter", "authenticator": + return true + default: + return false + } +} + +// loadPlugin 加载插件 func (l *Loader) loadPlugin(entry pluginEntry, m *Manager, router *api.Router) error { - slog.Info("loading plugin", "id", entry.ID, "type", entry.Type) + slog.Debug("loading plugin", "id", entry.ID, "type", entry.Type) p, ok := api.CreatePlugin(entry.ID) if !ok { return fmt.Errorf("plugin not registered: %s", entry.ID) } - // 初始化插件 if err := p.Init(entry.Config); err != nil { return fmt.Errorf("plugin %s init failed: %w", entry.ID, err) } - // 根据类型进行分类 + if err := registerPlugin(entry, p, m, router); err != nil { + return err + } + return nil +} + +// registerPlugin 注册插件 +func registerPlugin(entry pluginEntry, p api.Plugin, m *Manager, router *api.Router) error { switch entry.Type { + case "authenticator": + auth, ok := p.(api.Authenticator) + if !ok { + return fmt.Errorf("plugin %s is not an Authenticator", entry.ID) + } + m.AddAuthenticator(auth) + case "guard": + guard, ok := p.(api.Guard) + if !ok { + return fmt.Errorf("plugin %s is not a Guard", entry.ID) + } + m.AddGuard(guard) + case "preprocessor": + pre, ok := p.(api.PreProcessor) + if !ok { + return fmt.Errorf("plugin %s is not a PreProcessor", entry.ID) + } + m.AddPreProcessor(pre) case "router": r, ok := p.(api.Router) if !ok { return fmt.Errorf("plugin %s is not a Router", entry.ID) } *router = r - slog.Info("router plugin loaded", "id", entry.ID) + case "postprocessor": + post, ok := p.(api.PostProcessor) + if !ok { + return fmt.Errorf("plugin %s is not a PostProcessor", entry.ID) + } + m.AddPostProcessor(post) case "adapter": a, ok := p.(api.ChatAdapter) if !ok { @@ -98,16 +136,11 @@ func (l *Loader) loadPlugin(entry pluginEntry, m *Manager, router *api.Router) e if err := m.AddAdapter(a); err != nil { return err } - slog.Info("adapter plugin loaded", "id", entry.ID, "protocol", a.Protocol()) - case "authenticator": - auth, ok := p.(api.Authenticator) - if !ok { - return fmt.Errorf("plugin %s is not an Authenticator", entry.ID) - } - m.AddAuthenticator(auth) - slog.Info("authenticator plugin loaded", "id", entry.ID) + slog.Debug("adapter plugin loaded", "id", entry.ID, "protocol", a.Protocol()) + return nil default: return fmt.Errorf("unknown plugin type: %s", entry.Type) } + slog.Debug(entry.Type+" plugin loaded", "id", entry.ID) return nil } diff --git a/internal/plugin/manager.go b/internal/plugin/manager.go index 2e69d7d..f0bc149 100644 --- a/internal/plugin/manager.go +++ b/internal/plugin/manager.go @@ -12,8 +12,12 @@ import ( type Manager struct { router api.Router // 路由 registry *Registry // 插件注册表 - adapters map[string]api.ChatAdapter // 适配器 authenticators []api.Authenticator // 认证器 + guards []api.Guard // 守卫 + preprocessors []api.PreProcessor // 预处理器 + adapters map[string]api.ChatAdapter // 适配器 + postprocessors []api.PostProcessor // 后置处理器 + Warnings []string // 启动警告 } func NewManager() *Manager { @@ -27,6 +31,8 @@ func (m *Manager) SetRouter(r api.Router) { m.router = r } func (m *Manager) AddAuthenticator(a api.Authenticator) { m.authenticators = append(m.authenticators, a) } +func (m *Manager) AddGuard(g api.Guard) { m.guards = append(m.guards, g) } +func (m *Manager) AddPreProcessor(p api.PreProcessor) { m.preprocessors = append(m.preprocessors, p) } func (m *Manager) AddAdapter(a api.ChatAdapter) error { protocol := strings.ToLower(a.Protocol()) @@ -37,6 +43,10 @@ func (m *Manager) AddAdapter(a api.ChatAdapter) error { return nil } +func (m *Manager) AddPostProcessor(p api.PostProcessor) { + m.postprocessors = append(m.postprocessors, p) +} + func (m *Manager) Router() api.Router { return m.router } @@ -44,8 +54,34 @@ func (m *Manager) Router() api.Router { func (m *Manager) Authenticators() []api.Authenticator { return m.authenticators } - +func (m *Manager) Guards() []api.Guard { return m.guards } +func (m *Manager) PreProcessors() []api.PreProcessor { return m.preprocessors } func (m *Manager) Adapter(protocol string) (api.ChatAdapter, bool) { a, ok := m.adapters[strings.ToLower(protocol)] return a, ok } +func (m *Manager) PostProcessors() []api.PostProcessor { return m.postprocessors } + +// Summary 返回已加载插件的摘要信息 +func (m *Manager) Summary() string { + var b strings.Builder + if m.router != nil { + fmt.Fprintf(&b, " ● %s (router)\n", m.router.ID()) + } + for _, a := range m.adapters { + fmt.Fprintf(&b, " ● %s (adapter)\n", a.ID()) + } + for _, a := range m.authenticators { + fmt.Fprintf(&b, " ● %s (authenticator)\n", a.ID()) + } + for _, g := range m.guards { + fmt.Fprintf(&b, " ● %s (guard)\n", g.ID()) + } + for _, p := range m.preprocessors { + fmt.Fprintf(&b, " ● %s (preprocessor)\n", p.ID()) + } + for _, p := range m.postprocessors { + fmt.Fprintf(&b, " ● %s (postprocessor)\n", p.ID()) + } + return b.String() +} diff --git a/internal/plugin/scheduler.go b/internal/plugin/scheduler.go index b26950a..63e4580 100644 --- a/internal/plugin/scheduler.go +++ b/internal/plugin/scheduler.go @@ -43,43 +43,99 @@ func (s *Scheduler) executeAuth(ctx context.Context) error { } } -func (s *Scheduler) executePipeline(ctx context.Context, model string) (*api.RouteDecision, error) { +func (s *Scheduler) executePipeline(ctx context.Context, model string, req any) (*api.RouteDecision, string, error) { + // auth if err := s.executeAuth(ctx); err != nil { - return nil, err + return nil, "", err + } + + pctx := &api.PipelineContext{Context: ctx, Request: req, Meta: make(map[string]any)} + + // guard + for _, g := range s.manager.Load().Guards() { + if err := g.Guard(pctx); err != nil { + return nil, "", err + } + } + + // preprocess + for _, p := range s.manager.Load().PreProcessors() { + if err := p.PreProcess(pctx); err != nil { + return nil, "", err + } + } + + // route + decision, err := s.manager.Load().Router().Route(ctx, model) + if err != nil { + return nil, "", err } - return s.manager.Load().Router().Route(ctx, model) + pctx.Route = decision + + if decision.Model != "" { + return decision, decision.Model, nil + } + return decision, model, nil } // ExecuteChat 执行完整 chat 管线:鉴权 → 路由 → 协议适配。 func (s *Scheduler) ExecuteChat(ctx context.Context, payload *api.ChatRequest) (*api.ChatResponse, error) { - decision, err := s.executePipeline(ctx, payload.Model) + decision, actualModel, a, err := s.resolveAdapter(ctx, payload.Model, payload) if err != nil { return nil, err } - a, ok := s.manager.Load().Adapter(decision.Protocol) - if !ok { - return nil, &api.PluginError{ - Code: "adapter_not_found", - Message: fmt.Sprintf("adapter not found: %s", decision.Protocol), - Retryable: false, - } + payload.Model = actualModel + resp, err := a.SendChat(ctx, payload, decision) + if err != nil { + return nil, err + } + if err := s.executePostProcess(ctx, payload, resp, decision); err != nil { + return nil, err } - return a.SendChat(ctx, payload, decision) + return resp, nil } // ExecuteChatStream 执行流式 chat 管线,返回响应 channel。 func (s *Scheduler) ExecuteChatStream(ctx context.Context, payload *api.ChatRequest) (<-chan *api.StreamChunk, error) { - decision, err := s.executePipeline(ctx, payload.Model) + decision, actualModel, a, err := s.resolveAdapter(ctx, payload.Model, payload) if err != nil { return nil, err } + payload.Model = actualModel + return a.SendChatStream(ctx, payload, decision) +} + +func (s *Scheduler) resolveAdapter(ctx context.Context, model string, req any) (*api.RouteDecision, string, api.ChatAdapter, + error, +) { + decision, actualModel, err := s.executePipeline(ctx, model, req) + if err != nil { + return nil, "", nil, err + } a, ok := s.manager.Load().Adapter(decision.Protocol) if !ok { - return nil, &api.PluginError{ + return nil, "", nil, &api.PluginError{ Code: "adapter_not_found", Message: fmt.Sprintf("adapter not found: %s", decision.Protocol), Retryable: false, } } - return a.SendChatStream(ctx, payload, decision) + return decision, actualModel, a, nil +} + +func (s *Scheduler) executePostProcess(ctx context.Context, req any, resp any, decision *api.RouteDecision) error { + postprocessors := s.manager.Load().PostProcessors() + if len(postprocessors) == 0 { + return nil + } + pctx := &api.PipelineContext{ + Context: ctx, Request: req, Response: resp, Route: decision, + Meta: make(map[string]any), + } + for _, p := range postprocessors { + if err := p.PostProcess(pctx); err != nil { + return err + } + } + return nil } diff --git a/internal/plugin/scheduler_test.go b/internal/plugin/scheduler_test.go new file mode 100644 index 0000000..44b7c5b --- /dev/null +++ b/internal/plugin/scheduler_test.go @@ -0,0 +1,166 @@ +package plugin + +import ( + "context" + "testing" + + "strait/api" + "strait/internal/testutil" +) + +func newTestManager(router api.Router, adapter api.ChatAdapter, auths ...api.Authenticator) *Manager { + m := NewManager() + m.SetRouter(router) + if adapter != nil { + _ = m.AddAdapter(adapter) + } + for _, a := range auths { + m.AddAuthenticator(a) + } + return m +} + +// ── ExecuteChat ── + +func TestExecuteChat_Success(t *testing.T) { + router := testutil.NewMockRouter("mock", "http://test", "sk-test") + adapter := testutil.NewMockAdapter("mock", "hello") + m := newTestManager(router, adapter) + s := NewScheduler(m) + + payload := &api.ChatRequest{Model: "test-model", Messages: []api.Message{{Role: "user", Content: "hi"}}} + resp, err := s.ExecuteChat(context.Background(), payload) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.Choices[0].Content != "hello" { + t.Fatalf("expected 'hello', got '%s'", resp.Choices[0].Content) + } +} + +func TestExecuteChat_AuthFailed(t *testing.T) { + router := testutil.NewMockRouter("mock", "http://test", "sk-test") + adapter := testutil.NewMockAdapter("mock", "hello") + auth := testutil.NewMockAuthenticator(map[string]string{"valid-token": "admin"}) + m := newTestManager(router, adapter, auth) + s := NewScheduler(m) + + ctx := api.WithAuthToken(context.Background(), "wrong-token") + payload := &api.ChatRequest{Model: "test-model"} + _, err := s.ExecuteChat(ctx, payload) + if err == nil { + t.Fatal("expected auth error") + } +} + +func TestExecuteChat_AuthSuccess(t *testing.T) { + router := testutil.NewMockRouter("mock", "http://test", "sk-test") + adapter := testutil.NewMockAdapter("mock", "ok") + auth := testutil.NewMockAuthenticator(map[string]string{"valid-token": "admin"}) + m := newTestManager(router, adapter, auth) + s := NewScheduler(m) + + ctx := api.WithAuthToken(context.Background(), "valid-token") + payload := &api.ChatRequest{Model: "test-model"} + _, err := s.ExecuteChat(ctx, payload) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestExecuteChat_NoAdapter(t *testing.T) { + router := testutil.NewMockRouter("unknown-protocol", "http://test", "sk-test") + m := newTestManager(router, nil) + s := NewScheduler(m) + + payload := &api.ChatRequest{Model: "test-model"} + _, err := s.ExecuteChat(context.Background(), payload) + if err == nil { + t.Fatal("expected adapter_not_found error") + } +} + +func TestExecuteChat_RouterError(t *testing.T) { + router := testutil.NewMockRouterWithError(&api.PluginError{Code: "NO_ROUTE", Message: "no route"}) + adapter := testutil.NewMockAdapter("mock", "hello") + m := newTestManager(router, adapter) + s := NewScheduler(m) + + payload := &api.ChatRequest{Model: "unknown"} + _, err := s.ExecuteChat(context.Background(), payload) + if err == nil { + t.Fatal("expected router error") + } +} + +func TestExecuteChat_ModelOverride(t *testing.T) { + router := &testutil.MockRouter{ + Decision: &api.RouteDecision{ + Protocol: "mock", + BaseURL: "http://test", + APIKey: "sk-test", + Model: "overridden-model", + }, + } + adapter := testutil.NewMockAdapter("mock", "ok") + m := newTestManager(router, adapter) + s := NewScheduler(m) + + payload := &api.ChatRequest{Model: "original-model"} + _, err := s.ExecuteChat(context.Background(), payload) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if payload.Model != "overridden-model" { + t.Fatalf("expected model override to 'overridden-model', got '%s'", payload.Model) + } +} + +// ── ExecuteChatStream ── + +func TestExecuteChatStream_Success(t *testing.T) { + router := testutil.NewMockRouter("mock", "http://test", "sk-test") + adapter := testutil.NewMockAdapter("mock", "stream-ok") + m := newTestManager(router, adapter) + s := NewScheduler(m) + + payload := &api.ChatRequest{Model: "test-model", Stream: true} + ch, err := s.ExecuteChatStream(context.Background(), payload) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + chunk, ok := <-ch + if !ok { + t.Fatal("expected at least one chunk") + } + if chunk.Choices[0].Content != "mock" { + t.Fatalf("expected 'mock', got '%s'", chunk.Choices[0].Content) + } +} + +// ── ReloadManager ── + +func TestReloadManager(t *testing.T) { + router1 := testutil.NewMockRouter("mock", "http://v1", "sk-v1") + adapter1 := testutil.NewMockAdapter("mock", "v1") + m1 := newTestManager(router1, adapter1) + s := NewScheduler(m1) + + // v1 + payload := &api.ChatRequest{Model: "test"} + resp, _ := s.ExecuteChat(context.Background(), payload) + if resp.Choices[0].Content != "v1" { + t.Fatalf("expected v1, got %s", resp.Choices[0].Content) + } + + // reload to v2 + router2 := testutil.NewMockRouter("mock", "http://v2", "sk-v2") + adapter2 := testutil.NewMockAdapter("mock", "v2") + m2 := newTestManager(router2, adapter2) + s.ReloadManager(m2) + + resp, _ = s.ExecuteChat(context.Background(), payload) + if resp.Choices[0].Content != "v2" { + t.Fatalf("expected v2 after reload, got %s", resp.Choices[0].Content) + } +} diff --git a/internal/router/router.go b/internal/router/router.go index d76d395..9bf0225 100644 --- a/internal/router/router.go +++ b/internal/router/router.go @@ -4,7 +4,10 @@ import ( "context" "errors" "fmt" + "log/slog" + "math/rand" "os" + "strings" "sync" "strait/internal/config" @@ -23,22 +26,26 @@ type providerYAML struct { Models []string `yaml:"models"` // 支持的模型名称列表 } +// targetYAML 定义路由指向的目标提供者与模型 +type targetYAML struct { + Provider string `yaml:"provider"` // 目标提供者 ID + Model string `yaml:"model"` // 目标模型名称 + Priority int `yaml:"priority"` // 路由优先级 + Weight int `yaml:"weight"` // 路由权重 +} + // routeYAML 定义路由规则的 YAML 配置结构 type routeYAML struct { - ID string `yaml:"id"` // 路由唯一标识 - Match matchYAML `yaml:"match"` // 匹配表达式(模型名或通配符) - Target targetYAML `yaml:"target"` // 路由目标 + ID string `yaml:"id"` // 路由唯一标识 + Match matchYAML `yaml:"match"` // 匹配表达式(模型名或通配符) + Target targetYAML `yaml:"target"` // 路由目标 + Targets []targetYAML `yaml:"targets"` // 路由目标列表 + Strategy string `yaml:"strategy"` // 路由策略(priority / weight) } // matchYAML 定义匹配规则的 YAML 配置结构 type matchYAML struct { - Model string `yaml:"model"` // 匹配的模型名称 -} - -// targetYAML 定义路由指向的目标提供者与模型 -type targetYAML struct { - Provider string `yaml:"provider"` // 目标提供者 ID - Model string `yaml:"model"` // 目标模型名称 + Model string `yaml:"model"` // 请求匹配的模型名称 } // Router 定义了路由器结构 @@ -46,8 +53,8 @@ type Router struct { mu sync.RWMutex // 读写锁 providers map[string]providerYAML // 提供者配置映射 routes []routeYAML // 路由规则列表 - providersPath string - routesPath string + providersPath string // 提供者配置文件路径 + routesPath string // 路由规则文件路径 } func (r *Router) ID() string { return "router-yaml" } @@ -97,9 +104,21 @@ func (r *Router) loadConfig() (map[string]providerYAML, []routeYAML, error) { // 4. 校验路由关联有效性 for _, route := range rs.Routes { - providerID := route.Target.Provider - if _, exists := r.providers[providerID]; !exists { - return nil, nil, fmt.Errorf("route %s references unknown provider: %s", route.ID, providerID) + targets := route.Targets + if len(route.Targets) > 0 && route.Target.Provider != "" { + slog.Warn("both target and targets defined, using targets", "route", route.ID) + } + if len(targets) == 0 { + if route.Target.Provider == "" { + return nil, nil, fmt.Errorf("route %s has no target", route.ID) + } + targets = []targetYAML{route.Target} + } + for _, target := range targets { + providerID := target.Provider + if _, exists := r.providers[providerID]; !exists { + return nil, nil, fmt.Errorf("route %s references unknown provider: %s", route.ID, providerID) + } } } return r.providers, rs.Routes, nil @@ -134,37 +153,108 @@ func (r *Router) Reload() error { func (r *Router) Route(_ context.Context, model string) (*api.RouteDecision, error) { r.mu.RLock() - var provider *providerYAML + defer r.mu.RUnlock() + for _, rt := range r.routes { if rt.Match.Model == model { - if p, ok := r.providers[rt.Target.Provider]; ok { - provider = &p - } - break + return r.resolveRoute(rt, model) } } - r.mu.RUnlock() - if provider == nil { - return nil, &api.PluginError{ - Code: "NO_ROUTE", - Message: fmt.Sprintf("no route for model: %s", model), - Retryable: false, + for _, rt := range r.routes { + if strings.Contains(rt.Match.Model, "*") && matchModel(rt.Match.Model, model) { + return r.resolveRoute(rt, model) } } - apiKey := os.Getenv(provider.APIKeyEnv) - if apiKey == "" { - return nil, &api.PluginError{ - Code: "NO_ROUTE", - Message: fmt.Sprintf("provider %s API key not set in env %s", provider.ID, provider.APIKeyEnv), - Retryable: false, + return nil, &api.PluginError{ + Code: "NO_ROUTE", + Message: fmt.Sprintf("no route for model: %s", model), + Retryable: false, + } +} + +func (r *Router) resolveRoute(rt routeYAML, model string) (*api.RouteDecision, error) { + // 获取 targets 列表 + targets := rt.Targets + if len(targets) == 0 { + targets = []targetYAML{rt.Target} + } + + // 按 strategy 选择 + selected := r.selectTarget(targets, rt.Strategy) + + if p, ok := r.providers[selected.Provider]; ok { + apiKey := os.Getenv(p.APIKeyEnv) + if apiKey == "" { + return nil, &api.PluginError{ + Code: "NO_ROUTE", + Message: fmt.Sprintf("provider %s API key not set in env %s", p.ID, p.APIKeyEnv), + Retryable: false, + } + } + + slog.Info("route selected", "route", rt.ID, "model", model, "provider", selected.Provider, "strategy", rt.Strategy) + + return &api.RouteDecision{ + Protocol: p.Protocol, + BaseURL: p.BaseURL, + APIKey: apiKey, + Model: selected.Model, + }, nil + } + + return nil, &api.PluginError{ + Code: "NO_ROUTE", + Message: fmt.Sprintf("no route for model: %s", model), + Retryable: false, + } +} + +func matchModel(pattern, model string) bool { + if pattern == "*" { + return true + } + if strings.HasSuffix(pattern, "*") { + return strings.HasPrefix(model, strings.TrimSuffix(pattern, "*")) + } + return pattern == model +} + +func (r *Router) selectTarget(targets []targetYAML, strategy string) targetYAML { + if len(targets) == 1 { + return targets[0] + } + + switch strategy { + case "priority": + // 按优先级选择 + best := targets[0] + for _, t := range targets { + if t.Priority < best.Priority { + best = t + } } + return best + case "weight": + // 按权重选择 + return r.selectByWeight(targets) + default: + return r.selectTarget(targets, "priority") } +} - return &api.RouteDecision{ - Protocol: provider.Protocol, - BaseURL: provider.BaseURL, - APIKey: apiKey, - }, nil +func (r *Router) selectByWeight(targets []targetYAML) targetYAML { + totalWeight := 0 + for _, t := range targets { + totalWeight += t.Weight + } + randWeight := rand.Intn(totalWeight) + for _, t := range targets { + randWeight -= t.Weight + if randWeight < 0 { + return t + } + } + return targets[0] } diff --git a/internal/router/router_test.go b/internal/router/router_test.go new file mode 100644 index 0000000..b478c1f --- /dev/null +++ b/internal/router/router_test.go @@ -0,0 +1,211 @@ +package router + +import ( + "context" + "os" + "testing" + + "strait/api" +) + +// ── matchModel ── + +func TestMatchModel_Exact(t *testing.T) { + if !matchModel("deepseek-chat", "deepseek-chat") { + t.Fatal("expected exact match") + } + if matchModel("deepseek-chat", "gpt-4") { + t.Fatal("expected no match") + } +} + +func TestMatchModel_Wildcard(t *testing.T) { + if !matchModel("deepseek-*", "deepseek-chat") { + t.Fatal("expected wildcard match") + } + if !matchModel("deepseek-*", "deepseek-reasoner") { + t.Fatal("expected wildcard match") + } + if matchModel("deepseek-*", "gpt-4") { + t.Fatal("expected wildcard no match") + } +} + +func TestMatchModel_MatchAll(t *testing.T) { + if !matchModel("*", "anything") { + t.Fatal("expected * to match anything") + } +} + +// ── selectTarget ── + +func TestSelectTarget_Single(t *testing.T) { + r := &Router{} + targets := []targetYAML{{Provider: "a", Priority: 1}} + got := r.selectTarget(targets, "priority") + if got.Provider != "a" { + t.Fatalf("expected a, got %s", got.Provider) + } +} + +func TestSelectTarget_Priority(t *testing.T) { + r := &Router{} + targets := []targetYAML{ + {Provider: "low", Priority: 2}, + {Provider: "high", Priority: 1}, + {Provider: "mid", Priority: 3}, + } + got := r.selectTarget(targets, "priority") + if got.Provider != "high" { + t.Fatalf("expected high, got %s", got.Provider) + } +} + +func TestSelectTarget_DefaultStrategy(t *testing.T) { + r := &Router{} + targets := []targetYAML{ + {Provider: "b", Priority: 2}, + {Provider: "a", Priority: 1}, + } + got := r.selectTarget(targets, "") + if got.Provider != "a" { + t.Fatalf("expected a (default to priority), got %s", got.Provider) + } +} + +func TestSelectTarget_Weight(t *testing.T) { + r := &Router{} + targets := []targetYAML{ + {Provider: "heavy", Weight: 90}, + {Provider: "light", Weight: 10}, + } + counts := map[string]int{} + for i := 0; i < 1000; i++ { + got := r.selectTarget(targets, "weight") + counts[got.Provider]++ + } + if counts["heavy"] < 700 || counts["heavy"] > 950 { + t.Fatalf("expected ~90%% heavy, got heavy=%d light=%d", counts["heavy"], counts["light"]) + } +} + +// ── Route (端到端,绕过文件 I/O) ── + +func newTestRouter() *Router { + return &Router{ + providers: map[string]providerYAML{ + "deepseek-main": { + ID: "deepseek-main", + Protocol: "openai", + BaseURL: "https://api.deepseek.com/v1", + APIKeyEnv: "TEST_DEEPSEEK_KEY", + }, + "ollama-local": { + ID: "ollama-local", + Protocol: "ollama", + BaseURL: "http://localhost:11434", + APIKeyEnv: "TEST_OLLAMA_KEY", + }, + }, + routes: []routeYAML{ + { + ID: "exact-route", + Match: matchYAML{Model: "deepseek-chat"}, + Strategy: "priority", + Targets: []targetYAML{ + {Provider: "deepseek-main", Model: "deepseek-chat", Priority: 1, Weight: 1}, + }, + }, + { + ID: "wildcard-route", + Match: matchYAML{Model: "ollama-*"}, + Strategy: "priority", + Targets: []targetYAML{ + {Provider: "ollama-local", Model: "qwen2.5:0.5b", Priority: 1, Weight: 1}, + }, + }, + { + ID: "legacy-route", + Match: matchYAML{Model: "deepseek-reasoner"}, + Target: targetYAML{ + Provider: "deepseek-main", + Model: "deepseek-reasoner", + }, + }, + }, + } +} + +func TestRoute_ExactMatch(t *testing.T) { + r := newTestRouter() + os.Setenv("TEST_DEEPSEEK_KEY", "sk-test") + defer os.Unsetenv("TEST_DEEPSEEK_KEY") + + decision, err := r.Route(context.Background(), "deepseek-chat") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if decision.Protocol != "openai" { + t.Fatalf("expected openai, got %s", decision.Protocol) + } + if decision.Model != "deepseek-chat" { + t.Fatalf("expected deepseek-chat, got %s", decision.Model) + } +} + +func TestRoute_WildcardMatch(t *testing.T) { + r := newTestRouter() + os.Setenv("TEST_OLLAMA_KEY", "sk-test") + defer os.Unsetenv("TEST_OLLAMA_KEY") + + decision, err := r.Route(context.Background(), "ollama-llama3") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if decision.Protocol != "ollama" { + t.Fatalf("expected ollama, got %s", decision.Protocol) + } + if decision.Model != "qwen2.5:0.5b" { + t.Fatalf("expected qwen2.5:0.5b, got %s", decision.Model) + } +} + +func TestRoute_NoMatch(t *testing.T) { + r := newTestRouter() + + _, err := r.Route(context.Background(), "unknown-model") + if err == nil { + t.Fatal("expected error for unknown model") + } + pe, ok := err.(*api.PluginError) + if !ok { + t.Fatalf("expected PluginError, got %T", err) + } + if pe.Code != "NO_ROUTE" { + t.Fatalf("expected NO_ROUTE, got %s", pe.Code) + } +} + +func TestRoute_MissingAPIKey(t *testing.T) { + r := newTestRouter() + os.Unsetenv("TEST_DEEPSEEK_KEY") + + _, err := r.Route(context.Background(), "deepseek-chat") + if err == nil { + t.Fatal("expected error for missing API key") + } +} + +func TestRoute_LegacyTarget(t *testing.T) { + r := newTestRouter() + os.Setenv("TEST_DEEPSEEK_KEY", "sk-test") + defer os.Unsetenv("TEST_DEEPSEEK_KEY") + + decision, err := r.Route(context.Background(), "deepseek-reasoner") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if decision.Model != "deepseek-reasoner" { + t.Fatalf("expected deepseek-reasoner, got %s", decision.Model) + } +} diff --git a/internal/testutil/mock.go b/internal/testutil/mock.go new file mode 100644 index 0000000..3918eb4 --- /dev/null +++ b/internal/testutil/mock.go @@ -0,0 +1,107 @@ +package testutil + +import ( + "context" + "fmt" + + "strait/api" +) + +// ── MockRouter ── + +// MockRouter 测试用路由插件,返回预设的 RouteDecision +type MockRouter struct { + Decision *api.RouteDecision + Err error +} + +func (m *MockRouter) ID() string { return "mock-router" } +func (m *MockRouter) Init(_ map[string]any) error { return nil } +func (m *MockRouter) Route(_ context.Context, _ string) (*api.RouteDecision, error) { + return m.Decision, m.Err +} + +// NewMockRouter 创建返回指定 provider 的 MockRouter +func NewMockRouter(protocol, baseURL, apiKey string) *MockRouter { + return &MockRouter{ + Decision: &api.RouteDecision{ + Protocol: protocol, + BaseURL: baseURL, + APIKey: apiKey, + }, + } +} + +// NewMockRouterWithError 创建返回错误的 MockRouter +func NewMockRouterWithError(err error) *MockRouter { + return &MockRouter{Err: err} +} + +// ── MockAdapter ── + +// MockAdapter 测试用 Chat 适配器,返回预设响应 +type MockAdapter struct { + ProtocolName string + Response *api.ChatResponse + Err error +} + +func (m *MockAdapter) ID() string { return "mock-adapter" } +func (m *MockAdapter) Init(_ map[string]any) error { return nil } +func (m *MockAdapter) Protocol() string { return m.ProtocolName } + +func (m *MockAdapter) SendChat(_ context.Context, _ *api.ChatRequest, _ *api.RouteDecision) (*api.ChatResponse, error) { + if m.Err != nil { + return nil, m.Err + } + return m.Response, nil +} + +func (m *MockAdapter) SendChatStream(_ context.Context, _ *api.ChatRequest, _ *api.RouteDecision) (<-chan *api.StreamChunk, error) { + if m.Err != nil { + return nil, m.Err + } + ch := make(chan *api.StreamChunk, 1) + ch <- &api.StreamChunk{ + Choices: []api.Choice{{Role: "assistant", Content: "mock", FinishReason: "stop"}}, + Model: m.Response.Model, + } + close(ch) + return ch, nil +} + +// NewMockAdapter 创建返回简单文本响应的 MockAdapter +func NewMockAdapter(protocol, content string) *MockAdapter { + return &MockAdapter{ + ProtocolName: protocol, + Response: &api.ChatResponse{ + ID: "mock-req-001", + Model: "mock-model", + Choices: []api.Choice{ + {Index: 0, Role: "assistant", Content: content, FinishReason: "stop"}, + }, + }, + } +} + +// ── MockAuthenticator ── + +// MockAuthenticator 测试用认证插件 +type MockAuthenticator struct { + ValidTokens map[string]string // token → subject +} + +func (m *MockAuthenticator) ID() string { return "mock-auth" } +func (m *MockAuthenticator) Init(_ map[string]any) error { return nil } + +func (m *MockAuthenticator) Authenticate(_ context.Context, token string) (*api.Subject, error) { + if subject, ok := m.ValidTokens[token]; ok { + return &api.Subject{ID: subject}, nil + } + return nil, fmt.Errorf("invalid token: %s", token) +} + +// NewMockAuthenticator 创建预设 token 的 MockAuthenticator +func NewMockAuthenticator(tokens map[string]string) *MockAuthenticator { + return &MockAuthenticator{ValidTokens: tokens} +} diff --git a/k8s/configmap.yaml b/k8s/configmap.yaml new file mode 100644 index 0000000..275620b --- /dev/null +++ b/k8s/configmap.yaml @@ -0,0 +1,57 @@ +apiVersion: v1 +kind: ConfigMap +metadata: + name: strait-config +data: + plugins.yaml: | + plugins: + - id: router-yaml + type: router + - id: adapter-openai + type: adapter + - id: auth-static-token + type: authenticator + config: + token: sk-admin-init + subject: admin + - id: adapter-ollama + type: adapter + providers.yaml : | + # 参考提供者配置 + providers: + - id: deepseek-main + protocol: openai + base_url: https://api.deepseek.com/v1 + api_key_env: DEEPSEEK_API_KEY + models: + - deepseek-chat + - deepseek-reasoner + - id: ollama-local + protocol: ollama + base_url: http://localhost:11434 + api_key_env: DEEPSEEK_API_KEY + models: + - qwen2.5:0.5b + routes.yaml : | + # 参考路由配置 + routes: + - id: deepseek-chat + match: + model: deepseek-chat + target: + provider: deepseek-main + model: deepseek-chat + + - id: deepseek-reasoner + match: + model: deepseek-reasoner + target: + provider: deepseek-main + model: deepseek-reasoner + + - id: ollama-qwen + match: + model: qwen2.5:0.5b + target: + provider: ollama-local + model: qwen2.5:0.5b \ No newline at end of file diff --git a/k8s/deployment.yaml b/k8s/deployment.yaml new file mode 100644 index 0000000..e5fdb9a --- /dev/null +++ b/k8s/deployment.yaml @@ -0,0 +1,41 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: strait +spec: + replicas: 1 + selector: + matchLabels: + app: strait + template: + metadata: + labels: + app: strait + spec: + containers: + - name: strait + image: strait:latest + envFrom: + - secretRef: + name: strait-secret + ports: + - containerPort: 8080 + livenessProbe: + httpGet: + path: /health + port: 8080 + initialDelaySeconds: 5 + periodSeconds: 10 + readinessProbe: + httpGet: + path: /ready + port: 8080 + initialDelaySeconds: 3 + periodSeconds: 5 + volumeMounts: + - name: config + mountPath: /app/configs + volumes: + - name: config + configMap: + name: strait-config \ No newline at end of file diff --git a/k8s/secret.example.yaml b/k8s/secret.example.yaml new file mode 100644 index 0000000..dcfab45 --- /dev/null +++ b/k8s/secret.example.yaml @@ -0,0 +1,7 @@ +apiVersion: v1 +kind: Secret +metadata: + name: strait-secret +type: Opaque +stringData: + DEEPSEEK_API_KEY: "your-api-key-here" diff --git a/k8s/service.yaml b/k8s/service.yaml new file mode 100644 index 0000000..20b7e5c --- /dev/null +++ b/k8s/service.yaml @@ -0,0 +1,11 @@ +apiVersion: v1 +kind: Service +metadata: + name: strait +spec: + selector: + app: strait + ports: + - port: 8080 + targetPort: 8080 + type: ClusterIP diff --git a/plugins/prompt-injector/injector.go b/plugins/prompt-injector/injector.go new file mode 100644 index 0000000..00e56da --- /dev/null +++ b/plugins/prompt-injector/injector.go @@ -0,0 +1,34 @@ +// Package prompt_injector 提供提示语注入功能(目前实现较为简单,后续可考虑使用更复杂的注入方式) +package prompt_injector + +import ( + "log/slog" + + "strait/api" +) + +type PromptInjector struct { + SystemPrompt string +} + +func (p *PromptInjector) ID() string { return "prompt-injector" } +func init() { + api.Register("prompt-injector", func() api.Plugin { return &PromptInjector{} }) +} + +func (p *PromptInjector) Init(cfg map[string]any) error { + if v, ok := cfg["system_prompt"].(string); ok { + p.SystemPrompt = v + } + return nil +} + +func (p *PromptInjector) PreProcess(pctx *api.PipelineContext) error { + req, ok := pctx.Request.(*api.ChatRequest) + if !ok { + return nil + } + req.Messages = append([]api.Message{{Role: "system", Content: p.SystemPrompt}}, req.Messages...) + slog.Info("injecting system prompt", "prompt", p.SystemPrompt, "messages_before", len(req.Messages)) + return nil +}