From 122542d87e8a0223bf166afba69b08737ed5dd63 Mon Sep 17 00:00:00 2001 From: Yukk1o <1984867293@qq.com> Date: Wed, 27 May 2026 21:37:45 +0800 Subject: [PATCH 1/8] =?UTF-8?q?fix(ci):=20=E6=B7=BB=E5=8A=A0=20Windows=20?= =?UTF-8?q?=E6=9E=84=E5=BB=BA=20+=20release=20=E6=A0=87=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .github/workflows/release.yml | 1 + 1 file changed, 1 insertion(+) 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 From b46133b0d8f8ff0e661628a2ee9881d31e8ce8ad Mon Sep 17 00:00:00 2001 From: Yukk1o <69411148+Yukk1o@users.noreply.github.com> Date: Thu, 28 May 2026 13:46:25 +0800 Subject: [PATCH 2/8] =?UTF-8?q?feat(deploy):=20=E4=BA=91=E5=8E=9F=E7=94=9F?= =?UTF-8?q?=E9=83=A8=E7=BD=B2=E6=94=AF=E6=8C=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## 变更内容 - Dockerfile: golang:1.25-alpine 多阶段构建,最终镜像 ~10MB - .dockerignore: 排除 .git/docs/vendor 等无关文件 - docker-compose.yml: 端口映射 8080,configs 目录挂载 - K8s deployment/service/configmap/secret 完整部署配置 - /metrics: Prometheus 指标端点 + HTTP 请求计数 ## 测试 - [x] Docker 构建成功 - [x] /health 返回正常 - [x] /ready 返回正常 - [x] /metrics 返回 Prometheus 格式指标 - [x] K8s 部署测试通过" --- .dockerignore | 6 ++++ .gitignore | 1 + Dockerfile | 17 +++++++++++ cmd/strait/main.go | 5 +++- docker-compose.yml | 8 ++++++ internal/app/server.go | 14 +++++++-- internal/metrics/metrics.go | 39 +++++++++++++++++++++++++ k8s/configmap.yaml | 57 +++++++++++++++++++++++++++++++++++++ k8s/deployment.yaml | 41 ++++++++++++++++++++++++++ k8s/secret.example.yaml | 7 +++++ k8s/service.yaml | 11 +++++++ 11 files changed, 203 insertions(+), 3 deletions(-) create mode 100644 .dockerignore create mode 100644 Dockerfile create mode 100644 docker-compose.yml create mode 100644 internal/metrics/metrics.go create mode 100644 k8s/configmap.yaml create mode 100644 k8s/deployment.yaml create mode 100644 k8s/secret.example.yaml create mode 100644 k8s/service.yaml 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/.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/cmd/strait/main.go b/cmd/strait/main.go index 8b71af4..b2349c4 100644 --- a/cmd/strait/main.go +++ b/cmd/strait/main.go @@ -6,6 +6,8 @@ import ( "log" "net/http" + "strait/internal/metrics" + "strait/internal/config" "strait/internal/hotreload" @@ -26,6 +28,7 @@ func main() { } scheduler := plugin.NewScheduler(m) + met := metrics.NewMetrics() reload := func(filename string) error { if filename == "plugins.yaml" { @@ -49,7 +52,7 @@ func main() { _ = watcher.Start(context.Background()) }() - server := app.NewServer(scheduler) + server := app.NewServer(scheduler, met) log.Println("strait listening on :8080") log.Fatal(http.ListenAndServe(":8080", server.Handler())) } 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/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/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/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 From d708a254688d67e8d79c082bf892b0c397050084 Mon Sep 17 00:00:00 2001 From: Yukk1o <69411148+Yukk1o@users.noreply.github.com> Date: Thu, 28 May 2026 19:44:21 +0800 Subject: [PATCH 3/8] =?UTF-8?q?feat(router):=20=E8=B7=AF=E7=94=B1=E5=A2=9E?= =?UTF-8?q?=E5=BC=BA=20=E2=80=94=20=E5=A4=9A=E7=9B=AE=E6=A0=87/=E6=9D=83?= =?UTF-8?q?=E9=87=8D/=E9=80=9A=E9=85=8D=E7=AC=A6/=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E6=98=A0=E5=B0=84?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Summary 路由增强:多目标路由 + 优先级/权重策略 + 通配符匹配 + 模型映射 ## 改动 **路由核心** - routeYAML 支持 targets 列表和 strategy 字段 - targetYAML 新增 Priority / Weight 字段 - 新增 selectTarget / selectByWeight / matchModel / resolveRoute 方法 - Route() 支持精确匹配 + 通配符两遍扫描 - RouteDecision 新增可选 Model 字段,支持目标模型映射 **管线重构** - executePipeline 统一处理模型覆盖,返回三值 - 提取 resolveAdapter 消除 ExecuteChat / ExecuteChatStream 重复 **日志** - 全项目 log → slog 迁移(main.go / watcher.go / router.go) - 路由选择结构化日志输出 **配置** - routes.yaml 支持 targets + strategy 新格式 - 旧格式 target(单数)向后兼容 - target 和 targets 同时存在时输出 WARN 日志 **文档** - roadmap 更新 P4/P7 状态 - built-in.md 更新 router-yaml 配置说明 - deployment.md 补充 Docker / K8s / Metrics 部署指南 --- api/types.go | 7 +- cmd/strait/main.go | 20 +++-- configs/routes.yaml | 21 ++++- docs/deployment.md | 87 ++++++++++++++++-- docs/plugins/built-in.md | 39 +++++++- docs/roadmap.md | 6 +- internal/hotreload/watcher.go | 5 +- internal/plugin/scheduler.go | 39 ++++---- internal/router/router.go | 164 ++++++++++++++++++++++++++-------- 9 files changed, 311 insertions(+), 77 deletions(-) diff --git a/api/types.go b/api/types.go index cb9ada4..942b3a2 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 认证后的调用方信息 diff --git a/cmd/strait/main.go b/cmd/strait/main.go index b2349c4..9df3d99 100644 --- a/cmd/strait/main.go +++ b/cmd/strait/main.go @@ -1,10 +1,11 @@ -// Strait — AI 代理网关入口。 +// main Strait — AI 代理网关入口。 package main import ( "context" - "log" + "log/slog" "net/http" + "os" "strait/internal/metrics" @@ -21,10 +22,14 @@ import ( ) 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) } scheduler := plugin.NewScheduler(m) @@ -52,7 +57,10 @@ func main() { _ = watcher.Start(context.Background()) }() - server := app.NewServer(scheduler, met) - log.Println("strait listening on :8080") - log.Fatal(http.ListenAndServe(":8080", server.Handler())) + server := app.NewServer(scheduler) + slog.Info("strait listening on :8080") + if err := http.ListenAndServe(":8080", server.Handler()); err != nil { + slog.Error("server crashed", "error", err) + os.Exit(1) + } } 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/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..80fb45f 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 适配 ⏳ / 负载均衡 + 熔断 ⏳ | +| **P4 多供应商** | 🚧 | Ollama 适配 ✅ / 负载均衡(优先级/权重)✅ / Anthropic + Gemini 适配 ⏳ / 熔断 ⏳ | | **P5 管道扩展** | ⏳ | PreProcessor / PostProcessor 接口 + 模型分组路由 | | **P6 持久化** | ⏳ | SQLite + Repository 层 + 管理 API CRUD + Playground | -| **P7 生产部署** | ⏳ | Docker + 优雅关闭 + Prometheus + 限流 | +| **P7 生产部署** | 🚧 | Docker ✅ + K8s 部署 ✅ + Prometheus /metrics ✅ + 优雅关闭 ⏳ + 限流 ⏳ | | **P8 扩展协议** | ⏳ | Embedding 透传 / MCP 端点 / Rerank 适配 | ### 内核关键里程碑 @@ -45,7 +45,7 @@ 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 鉴权 | 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/plugin/scheduler.go b/internal/plugin/scheduler.go index b26950a..6712cf2 100644 --- a/internal/plugin/scheduler.go +++ b/internal/plugin/scheduler.go @@ -43,43 +43,52 @@ 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) (*api.RouteDecision, string, error) { if err := s.executeAuth(ctx); err != nil { - return nil, err + return nil, "", err + } + decision, err := s.manager.Load().Router().Route(ctx, model) + if err != nil { + return nil, "", err } - return s.manager.Load().Router().Route(ctx, model) + 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) 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 return a.SendChat(ctx, payload, decision) } // 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) if err != nil { return nil, err } + payload.Model = actualModel + return a.SendChatStream(ctx, payload, decision) +} + +func (s *Scheduler) resolveAdapter(ctx context.Context, model string) (*api.RouteDecision, string, api.ChatAdapter, error) { + decision, actualModel, err := s.executePipeline(ctx, model) + 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 } 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] } From 24bfaf5905fa94599a1799929a87186e4941106b Mon Sep 17 00:00:00 2001 From: Yukk1o <1984867293@qq.com> Date: Thu, 28 May 2026 19:55:09 +0800 Subject: [PATCH 4/8] =?UTF-8?q?fix(server):=20=E4=BF=AE=E5=A4=8D=E6=9C=8D?= =?UTF-8?q?=E5=8A=A1=E5=99=A8=E5=90=AF=E5=8A=A8=E6=97=B6=E5=8F=82=E6=95=B0?= =?UTF-8?q?=E4=BC=A0=E9=80=92=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 修正 NewServer 函数调用,新增 met 参数 - 解决服务器初始化时缺少依赖传递的问题 - 确保监控组件正确集成至服务器实例 - 增强服务器稳定性,防止启动异常 --- cmd/strait/main.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cmd/strait/main.go b/cmd/strait/main.go index 9df3d99..be4ee96 100644 --- a/cmd/strait/main.go +++ b/cmd/strait/main.go @@ -57,7 +57,7 @@ func main() { _ = watcher.Start(context.Background()) }() - server := app.NewServer(scheduler) + server := app.NewServer(scheduler, met) slog.Info("strait listening on :8080") if err := http.ListenAndServe(":8080", server.Handler()); err != nil { slog.Error("server crashed", "error", err) From 14d415a3eebf5696bf3bc6d267d38fa586fcff08 Mon Sep 17 00:00:00 2001 From: Yukk1o <1984867293@qq.com> Date: Thu, 28 May 2026 20:07:11 +0800 Subject: [PATCH 5/8] =?UTF-8?q?test(plugin):=20=E4=B8=BA=20Scheduler=20?= =?UTF-8?q?=E6=B7=BB=E5=8A=A0=E5=85=A8=E9=9D=A2=E5=8D=95=E5=85=83=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=E8=A6=86=E7=9B=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 添加 ExecuteChat 不同场景的测试,包括成功、认证失败、认证成功、无适配器、路由错误和模型覆盖 - 添加 ExecuteChatStream 流式接口的成功测试 - 添加 Scheduler 的 ReloadManager 重新加载管理器后的行为测试 - 添加 router 包的单元测试,包括模型匹配、目标选择策略和路由端到端测试 - 实现 testutil 包中 MockRouter、MockAdapter 和 MockAuthenticator 的模拟实现,便于测试依赖注入 - 确保所有测试用例覆盖主要逻辑分支及异常情况,提高代码质量和稳定性 --- internal/plugin/scheduler_test.go | 166 +++++++++++++++++++++++ internal/router/router_test.go | 211 ++++++++++++++++++++++++++++++ internal/testutil/mock.go | 107 +++++++++++++++ 3 files changed, 484 insertions(+) create mode 100644 internal/plugin/scheduler_test.go create mode 100644 internal/router/router_test.go create mode 100644 internal/testutil/mock.go 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_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} +} From dbe85c65a6a7a51af7a421645b3efb3c98aa6bf6 Mon Sep 17 00:00:00 2001 From: Yukk1o <1984867293@qq.com> Date: Thu, 28 May 2026 20:21:35 +0800 Subject: [PATCH 6/8] =?UTF-8?q?feat(server):=20=E4=BC=98=E5=8C=96HTTP?= =?UTF-8?q?=E6=9C=8D=E5=8A=A1=E5=90=AF=E5=8A=A8=E4=B8=8E=E4=BC=98=E9=9B=85?= =?UTF-8?q?=E5=85=B3=E9=97=AD=E6=9C=BA=E5=88=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 使用http.Server结构体启动HTTP服务 - 监听系统中断信号实现优雅关闭 - 在关闭时设置10秒的超时上下文 - 添加日志记录服务启动、关闭及异常信息 - 启动服务与信号监听在独立协程中运行 --- cmd/strait/main.go | 32 ++++++++++++++++++++++++++++---- 1 file changed, 28 insertions(+), 4 deletions(-) diff --git a/cmd/strait/main.go b/cmd/strait/main.go index be4ee96..c7b7890 100644 --- a/cmd/strait/main.go +++ b/cmd/strait/main.go @@ -3,9 +3,13 @@ package main import ( "context" + "errors" "log/slog" "net/http" "os" + "os/signal" + "syscall" + "time" "strait/internal/metrics" @@ -57,10 +61,30 @@ func main() { _ = watcher.Start(context.Background()) }() + // 启动 HTTP 服务 server := app.NewServer(scheduler, met) - slog.Info("strait listening on :8080") - if err := http.ListenAndServe(":8080", server.Handler()); err != nil { - slog.Error("server crashed", "error", err) - os.Exit(1) + 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") } From 4abf1d66a3f891eb28cf43bd3b46d514fa8e8f72 Mon Sep 17 00:00:00 2001 From: Yukk1o <1984867293@qq.com> Date: Thu, 28 May 2026 22:10:09 +0800 Subject: [PATCH 7/8] =?UTF-8?q?feat(plugin):=20=E5=A2=9E=E5=8A=A0=E5=AE=8C?= =?UTF-8?q?=E6=95=B4=E7=9A=84=E6=8F=92=E4=BB=B6=E7=AE=A1=E9=81=93=E6=94=AF?= =?UTF-8?q?=E6=8C=81=E4=B8=8E=E6=8F=90=E7=A4=BA=E6=B3=A8=E5=85=A5=E5=8A=9F?= =?UTF-8?q?=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 Guard、PreProcessor、PostProcessor 三个插件接口,拓展插件类型支持 - 定义 PipelineContext 贯穿插件管线各阶段信息传递 - 插件管理器支持守卫、预处理器与后置处理器插件的注册与管理 - 插件加载器新增对 guard、preprocessor、postprocessor 类型插件的加载与日志区分 - 调度器执行流程改为鉴权 → 守卫 → 预处理 → 路由 → 适配 → 后置处理的管线顺序 - 适配器调用前后新增预处理与后置处理钩子,支持插件对请求和响应的修改 - 在启动命令行输出插件加载摘要,显示包括新插件类型的状态信息 - 新增提示注入插件 prompt-injector,支持在请求消息最前面注入系统提示语 - 配置文件示例添加 prompt-injector 插件的定义及默认系统提示配置 --- api/plugin.go | 15 +++++++ api/types.go | 11 +++++ cmd/strait/main.go | 26 ++++++++--- configs/plugins.yaml | 6 ++- docs/roadmap.md | 6 +-- internal/plugin/loader.go | 68 +++++++++++++++++++++-------- internal/plugin/manager.go | 40 ++++++++++++++++- internal/plugin/scheduler.go | 59 ++++++++++++++++++++++--- plugins/prompt-injector/injector.go | 34 +++++++++++++++ 9 files changed, 231 insertions(+), 34 deletions(-) create mode 100644 plugins/prompt-injector/injector.go 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 942b3a2..54282f8 100644 --- a/api/types.go +++ b/api/types.go @@ -100,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 c7b7890..c30fc81 100644 --- a/cmd/strait/main.go +++ b/cmd/strait/main.go @@ -4,6 +4,7 @@ package main import ( "context" "errors" + "fmt" "log/slog" "net/http" "os" @@ -11,20 +12,27 @@ import ( "syscall" "time" - "strait/internal/metrics" - + "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, @@ -34,6 +42,14 @@ func main() { m, err := loader.Build() if err != nil { 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) 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/docs/roadmap.md b/docs/roadmap.md index 80fb45f..92a5389 100644 --- a/docs/roadmap.md +++ b/docs/roadmap.md @@ -22,9 +22,9 @@ | **P2 生产可用性** | ✅ | 鉴权 + 错误模型 + 热重载 + 统一内部模型 + OpenAI 输出格式 | | **P3 Agent 协议** | ✅ | Function Calling + Tool Use + Tool Calls 响应 | | **P4 多供应商** | 🚧 | Ollama 适配 ✅ / 负载均衡(优先级/权重)✅ / Anthropic + Gemini 适配 ⏳ / 熔断 ⏳ | -| **P5 管道扩展** | ⏳ | PreProcessor / PostProcessor 接口 + 模型分组路由 | +| **P5 管道扩展** | 🚧 | Guard / PreProcessor / PostProcessor 接口 ✅ + 模型分组路由 ⏳ | | **P6 持久化** | ⏳ | SQLite + Repository 层 + 管理 API CRUD + Playground | -| **P7 生产部署** | 🚧 | Docker ✅ + K8s 部署 ✅ + Prometheus /metrics ✅ + 优雅关闭 ⏳ + 限流 ⏳ | +| **P7 生产部署** | 🚧 | Docker ✅ + K8s 部署 ✅ + Prometheus /metrics ✅ + 优雅关闭 ✅ + 限流 ⏳ | | **P8 扩展协议** | ⏳ | Embedding 透传 / MCP 端点 / Rerank 适配 | ### 内核关键里程碑 @@ -49,6 +49,7 @@ Strait 所有组成部分都是插件——Router、Adapter、Authenticator、Gu | 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/plugin/loader.go b/internal/plugin/loader.go index aebef47..b5ffa1c 100644 --- a/internal/plugin/loader.go +++ b/internal/plugin/loader.go @@ -51,7 +51,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,17 +62,21 @@ 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 + } +} + 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 { @@ -83,14 +90,48 @@ func (l *Loader) loadPlugin(entry pluginEntry, m *Manager, router *api.Router) e // 根据类型进行分类 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) + slog.Debug("authenticator plugin loaded", "id", entry.ID) + case "guard": + // 守卫插件 + guard, ok := p.(api.Guard) + if !ok { + return fmt.Errorf("plugin %s is not a Guard", entry.ID) + } + m.AddGuard(guard) + slog.Debug("guard plugin loaded", "id", entry.ID) + case "preprocessor": + // 预处理器插件 + pre, ok := p.(api.PreProcessor) + if !ok { + return fmt.Errorf("plugin %s is not a PreProcessor", entry.ID) + } + m.AddPreProcessor(pre) + slog.Debug("preprocessor plugin loaded", "id", entry.ID) 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) + slog.Debug("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) + slog.Debug("postprocessor plugin loaded", "id", entry.ID) case "adapter": + // 适配器插件 a, ok := p.(api.ChatAdapter) if !ok { return fmt.Errorf("plugin %s is not a ChatAdapter", entry.ID) @@ -98,14 +139,7 @@ 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()) default: return fmt.Errorf("unknown plugin type: %s", entry.Type) } 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 6712cf2..63e4580 100644 --- a/internal/plugin/scheduler.go +++ b/internal/plugin/scheduler.go @@ -43,14 +43,35 @@ func (s *Scheduler) executeAuth(ctx context.Context) error { } } -func (s *Scheduler) executePipeline(ctx context.Context, model string) (*api.RouteDecision, string, 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 } + + 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 } + pctx.Route = decision + if decision.Model != "" { return decision, decision.Model, nil } @@ -59,17 +80,24 @@ func (s *Scheduler) executePipeline(ctx context.Context, model string) (*api.Rou // ExecuteChat 执行完整 chat 管线:鉴权 → 路由 → 协议适配。 func (s *Scheduler) ExecuteChat(ctx context.Context, payload *api.ChatRequest) (*api.ChatResponse, error) { - decision, actualModel, a, err := s.resolveAdapter(ctx, payload.Model) + decision, actualModel, a, err := s.resolveAdapter(ctx, payload.Model, payload) if err != nil { return nil, err } payload.Model = actualModel - return a.SendChat(ctx, payload, decision) + 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 resp, nil } // ExecuteChatStream 执行流式 chat 管线,返回响应 channel。 func (s *Scheduler) ExecuteChatStream(ctx context.Context, payload *api.ChatRequest) (<-chan *api.StreamChunk, error) { - decision, actualModel, a, err := s.resolveAdapter(ctx, payload.Model) + decision, actualModel, a, err := s.resolveAdapter(ctx, payload.Model, payload) if err != nil { return nil, err } @@ -77,8 +105,10 @@ func (s *Scheduler) ExecuteChatStream(ctx context.Context, payload *api.ChatRequ return a.SendChatStream(ctx, payload, decision) } -func (s *Scheduler) resolveAdapter(ctx context.Context, model string) (*api.RouteDecision, string, api.ChatAdapter, error) { - decision, actualModel, err := s.executePipeline(ctx, model) +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 } @@ -92,3 +122,20 @@ func (s *Scheduler) resolveAdapter(ctx context.Context, model string) (*api.Rout } 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/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 +} From dd0ae9e1b157b28ca4f51b15d5fdd1be171b6576 Mon Sep 17 00:00:00 2001 From: Yukk1o <1984867293@qq.com> Date: Thu, 28 May 2026 23:47:53 +0800 Subject: [PATCH 8/8] =?UTF-8?q?refactor(plugin):=20=E4=BC=98=E5=8C=96?= =?UTF-8?q?=E6=8F=92=E4=BB=B6=E5=8A=A0=E8=BD=BD=E5=92=8C=E6=B3=A8=E5=86=8C?= =?UTF-8?q?=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 将插件加载拆分为 loadPlugin 和 registerPlugin 两个函数,将 loadPlugin 的圈复杂度从 16 降到 8 - registerPlugin 根据插件类型注册到管理器中 - 增加类型断言错误处理,确保插件类型正确 - 删除重复或无效的调试日志代码 - 统一插件加载成功时的日志输出格式 --- internal/plugin/loader.go | 25 ++++++++++++------------- 1 file changed, 12 insertions(+), 13 deletions(-) diff --git a/internal/plugin/loader.go b/internal/plugin/loader.go index b5ffa1c..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) @@ -75,6 +76,7 @@ func isCriticalType(t string) bool { } } +// loadPlugin 加载插件 func (l *Loader) loadPlugin(entry pluginEntry, m *Manager, router *api.Router) error { slog.Debug("loading plugin", "id", entry.ID, "type", entry.Type) @@ -83,55 +85,50 @@ func (l *Loader) loadPlugin(entry pluginEntry, m *Manager, router *api.Router) e 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) - slog.Debug("authenticator plugin loaded", "id", entry.ID) case "guard": - // 守卫插件 guard, ok := p.(api.Guard) if !ok { return fmt.Errorf("plugin %s is not a Guard", entry.ID) } m.AddGuard(guard) - slog.Debug("guard plugin loaded", "id", entry.ID) case "preprocessor": - // 预处理器插件 pre, ok := p.(api.PreProcessor) if !ok { return fmt.Errorf("plugin %s is not a PreProcessor", entry.ID) } m.AddPreProcessor(pre) - slog.Debug("preprocessor plugin loaded", "id", entry.ID) case "router": - // 路由器插件 r, ok := p.(api.Router) if !ok { return fmt.Errorf("plugin %s is not a Router", entry.ID) } *router = r - slog.Debug("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) - slog.Debug("postprocessor plugin loaded", "id", entry.ID) case "adapter": - // 适配器插件 a, ok := p.(api.ChatAdapter) if !ok { return fmt.Errorf("plugin %s is not a ChatAdapter", entry.ID) @@ -140,8 +137,10 @@ func (l *Loader) loadPlugin(entry pluginEntry, m *Manager, router *api.Router) e return err } 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 }