feat: 十轮网关优化 - 安全加固/可观测性/性能/可靠性
CI / lint (push) Has been cancelled
CI / test (push) Has been cancelled
CI / build (push) Has been cancelled
CI / security-scan (push) Has been cancelled

- SSE Keepalive Ping (15s心跳防止代理断连)
- Timing HTTP 头 (X-Timing-Queue/Inference/Total-Ms)
- Adapter Request-ID 传播到后端
- Session 清理日志回调
- Server 安全加固 (ReadHeaderTimeout/MaxHeaderBytes 防 slowloris)
- Usage Tracker 数据保留清理 (retentionDays + 定期清理)
- Config Reload 后 Adapter Registry 更新 (RegisterIfAbsent + RWMutex)
- Rate Limiter 空闲 Bucket 清理 (30分钟过期)
- Shutdown Drain 超时可配置 (ShutdownDrainSeconds)
- Config 模型字段校验增强 (provider/endpoint/actual_model)
- Auth 过期 Key 自动清理 (5分钟扫描)
- Admin API Rate Limiting
- Adapter Health Check 独立超时 (每个 adapter 3s)
- TCP 连接阶段超时 (DialContext 5s + KeepAlive 30s)
- 幂等键缓存、审计日志、Gzip 中间件、CORS Expose Headers
- Backpressure 响应头、熔断器 Prometheus 指标
- 连接池优化、Trace-ID 全链路传播
This commit is contained in:
selfrelease
2026-08-03 15:43:11 +08:00
parent e7e98271d4
commit da9c8334d8
27 changed files with 3002 additions and 309 deletions
+333
View File
@@ -0,0 +1,333 @@
package server
import (
"encoding/json"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/edgeai/gateway/internal/adapter"
"github.com/edgeai/gateway/internal/auth"
"github.com/edgeai/gateway/internal/config"
"github.com/edgeai/gateway/internal/handler"
"github.com/edgeai/gateway/internal/middleware"
"github.com/edgeai/gateway/internal/observability"
)
// handleAdminKeys 处理 /v1/admin/keysGET 列出所有 KeyPOST 创建新 Key)。
func (s *Server) handleAdminKeys(w http.ResponseWriter, r *http.Request) {
requestID := middleware.GetRequestID(r.Context())
switch r.Method {
case http.MethodGet:
keys, err := s.auth.ListKeys()
if err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
return
}
handler.WriteJSON(w, http.StatusOK, map[string]any{"keys": keys})
case http.MethodPost:
var req struct {
AppID string `json:"app_id"`
TenantID string `json:"tenant_id"`
Name string `json:"name"`
AllowedModels []string `json:"allowed_models"`
AllowedPriorities []int `json:"allowed_priorities"`
IsAdmin bool `json:"is_admin"`
ExpiresAt string `json:"expires_at"` // RFC3339 格式,空表示永不过期
}
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid JSON body", requestID))
return
}
if req.AppID == "" {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "app_id is required", requestID))
return
}
if req.Name == "" {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "name is required", requestID))
return
}
apiKey := auth.GenerateAPIKey()
identity := &auth.AppIdentity{
AppID: req.AppID,
TenantID: req.TenantID,
Name: req.Name,
AllowedModels: req.AllowedModels,
AllowedPriorities: req.AllowedPriorities,
IsAdmin: req.IsAdmin,
}
if req.ExpiresAt != "" {
if t, err := time.Parse(time.RFC3339, req.ExpiresAt); err == nil {
identity.ExpiresAt = &t
}
}
if err := s.auth.AddKey(apiKey, identity); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
return
}
// 审计日志
actor := ""
if id := auth.GetAppIdentityFromRequest(r); id != nil {
actor = id.AppID
}
s.audit.RecordFromRequest(r, actor, "create_key", "api_key", req.AppID, http.StatusCreated, map[string]interface{}{
"app_id": req.AppID,
"name": req.Name,
"is_admin": req.IsAdmin,
"allowed_models": req.AllowedModels,
})
handler.WriteJSON(w, http.StatusCreated, map[string]any{
"api_key": apiKey,
"app_id": req.AppID,
"name": req.Name,
"is_admin": req.IsAdmin,
"allowed_models": req.AllowedModels,
"expires_at": req.ExpiresAt,
"message": "请妥善保存此 API Key,之后将无法再次查看",
})
default:
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
}
}
// handleAdminKeyByID 处理 /v1/admin/keys/{id}DELETE 删除,PATCH 禁用,PUT 轮换)。
func (s *Server) handleAdminKeyByID(w http.ResponseWriter, r *http.Request) {
requestID := middleware.GetRequestID(r.Context())
idStr := strings.TrimPrefix(r.URL.Path, "/v1/admin/keys/")
id, err := strconv.ParseInt(idStr, 10, 64)
if err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid key id", requestID))
return
}
switch r.Method {
case http.MethodDelete:
if err := s.auth.DeleteKey(id); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
return
}
// 审计日志
actor := ""
if ident := auth.GetAppIdentityFromRequest(r); ident != nil {
actor = ident.AppID
}
s.audit.RecordFromRequest(r, actor, "delete_key", "api_key", strconv.FormatInt(id, 10), http.StatusOK, nil)
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "deleted"})
case http.MethodPatch:
if err := s.auth.DisableKey(id); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
return
}
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "disabled"})
case http.MethodPut:
// Key 轮换:生成新 Key,禁用旧 Key
newKey, err := s.auth.RotateKey(id)
if err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
return
}
// 审计日志
actor := ""
if ident := auth.GetAppIdentityFromRequest(r); ident != nil {
actor = ident.AppID
}
s.audit.RecordFromRequest(r, actor, "rotate_key", "api_key", strconv.FormatInt(id, 10), http.StatusOK, nil)
handler.WriteJSON(w, http.StatusOK, map[string]any{
"status": "rotated",
"api_key": newKey,
"message": "旧 Key 已禁用,请妥善保存新 Key",
})
default:
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
}
}
// handleAdminUsage 处理 /v1/admin/usageGET 返回所有 app 的用量统计)。
func (s *Server) handleAdminUsage(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
return
}
usage := s.usageTracker.GetAll()
handler.WriteJSON(w, http.StatusOK, map[string]any{"usage": usage})
}
// handleAdminUsageByApp 处理 /v1/admin/usage/{app_id}GET 返回指定 app 的用量统计,支持 ?history=hours 查询历史)。
func (s *Server) handleAdminUsageByApp(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
return
}
appID := strings.TrimPrefix(r.URL.Path, "/v1/admin/usage/")
if appID == "" {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "app_id is required"))
return
}
// 支持 ?history=24 查询历史用量
if historyHours := r.URL.Query().Get("history"); historyHours != "" {
hours := 24
fmt.Sscanf(historyHours, "%d", &hours)
history, err := s.usageTracker.GetHistory(appID, hours)
if err != nil {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInternalError, err.Error()))
return
}
handler.WriteJSON(w, http.StatusOK, map[string]any{
"app_id": appID,
"hours": hours,
"history": history,
})
return
}
usage := s.usageTracker.Get(appID)
handler.WriteJSON(w, http.StatusOK, usage)
}
// handleAdminRateLimit 处理 /v1/admin/ratelimit/{app_id}GET 查看限流,PUT 设置自定义限流)。
func (s *Server) handleAdminRateLimit(w http.ResponseWriter, r *http.Request) {
requestID := middleware.GetRequestID(r.Context())
appID := strings.TrimPrefix(r.URL.Path, "/v1/admin/ratelimit/")
if appID == "" {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "app_id is required", requestID))
return
}
switch r.Method {
case http.MethodGet:
burst, ratePerMin := s.rateLimiter.GetLimit(appID)
handler.WriteJSON(w, http.StatusOK, map[string]any{
"app_id": appID,
"burst": burst,
"rate_per_minute": ratePerMin,
})
case http.MethodPut:
var req struct {
Burst int `json:"burst"`
RatePerMinute int `json:"rate_per_minute"`
}
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "invalid JSON body", requestID))
return
}
if req.Burst <= 0 || req.RatePerMinute <= 0 {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "burst and rate_per_minute must be positive", requestID))
return
}
refillPerSec := float64(req.RatePerMinute) / 60.0
s.rateLimiter.SetLimit(appID, float64(req.Burst), refillPerSec)
handler.WriteJSON(w, http.StatusOK, map[string]any{
"app_id": appID,
"burst": req.Burst,
"rate_per_minute": req.RatePerMinute,
"status": "updated",
})
default:
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
}
}
// handleAdminCircuitBreaker 处理 /v1/admin/circuit-breakerGET 返回熔断器状态)。
func (s *Server) handleAdminCircuitBreaker(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
return
}
stats := s.breaker.Stats()
handler.WriteJSON(w, http.StatusOK, stats)
}
// handleAdminScheduler 处理 /v1/admin/schedulerGET 返回调度器状态)。
func (s *Server) handleAdminScheduler(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
return
}
stats := map[string]any{
"running_tasks": s.scheduler.RunningCount(),
"queued_tasks": s.scheduler.QueueLength(),
"max_running": s.cfg.Scheduler.MaxRunningTasks,
"max_queued": s.cfg.Scheduler.MaxQueuedTasks,
"priority_aging_seconds": s.cfg.Scheduler.PriorityAgingSeconds,
"fairness": s.cfg.Scheduler.Fairness,
}
handler.WriteJSON(w, http.StatusOK, stats)
}
// handleAdminConfigReload 处理 /v1/admin/config/reloadPOST 触发配置热重载)。
func (s *Server) handleAdminConfigReload(w http.ResponseWriter, r *http.Request) {
requestID := middleware.GetRequestID(r.Context())
if r.Method != http.MethodPost {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
return
}
// 重新加载配置文件
cfgPath := config.ConfigPath()
newCfg, err := config.Load(cfgPath)
if err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, fmt.Sprintf("reload config failed: %v", err), requestID))
return
}
// 热更新模型映射
s.modelMap.Update(newCfg)
// 注册新增 adapter(已有 adapter 不重复注册)
for _, mc := range newCfg.Models {
provider := mc.Provider
endpoint := mc.Endpoint
switch provider {
case "ollama":
s.registry.RegisterIfAbsent(provider, func() adapter.ModelAdapter {
return adapter.NewOllamaAdapter(endpoint)
})
case "vllm":
s.registry.RegisterIfAbsent(provider, func() adapter.ModelAdapter {
return adapter.NewVLLMAdapter(endpoint)
})
}
}
// 更新 Server 持有的配置引用
s.cfg = newCfg
s.logger.Info("config hot-reloaded",
observability.F().
Event("config_reload").
Set("config_path", cfgPath))
// 审计日志
actor := ""
if ident := auth.GetAppIdentityFromRequest(r); ident != nil {
actor = ident.AppID
}
s.audit.RecordFromRequest(r, actor, "config_reload", "config", cfgPath, http.StatusOK, map[string]interface{}{
"models": len(newCfg.Models),
})
handler.WriteJSON(w, http.StatusOK, map[string]any{
"status": "reloaded",
"models": len(newCfg.Models),
"message": "配置已热重载,模型映射已更新",
})
}
+376 -25
View File
@@ -5,12 +5,14 @@ import (
"encoding/json"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/edgeai/gateway/internal/adapter"
"github.com/edgeai/gateway/internal/auth"
"github.com/edgeai/gateway/internal/config"
"github.com/edgeai/gateway/internal/connector"
ctxasm "github.com/edgeai/gateway/internal/context"
"github.com/edgeai/gateway/internal/handler"
"github.com/edgeai/gateway/internal/middleware"
@@ -47,12 +49,39 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
return
}
// 参数范围校验
if req.Temperature != nil && (*req.Temperature < 0 || *req.Temperature > 2) {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "temperature must be between 0 and 2", requestID))
return
}
if req.TopP != nil && (*req.TopP < 0 || *req.TopP > 1) {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "top_p must be between 0 and 1", requestID))
return
}
if req.MaxOutputTokens < 0 {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "max_output_tokens must be non-negative", requestID))
return
}
if len(req.Messages) > s.cfg.Context.MaxSessionMessages {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrContextTooLarge, fmt.Sprintf("messages count exceeds max %d", s.cfg.Context.MaxSessionMessages), requestID))
return
}
// 幂等性检查(仅非流式请求支持幂等)
if !req.Stream && req.IdempotencyKey != "" && identity != nil {
if s.checkIdempotency(w, req.IdempotencyKey, identity.AppID) {
return
}
}
// Check model permission
if identity != nil && !auth.CheckModelPermission(identity, req.Model) {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrPermissionDenied, "model not allowed for this application", requestID))
return
}
// Per-app 限流由全局 RateLimit 中间件处理,此处不再重复检查
// Resolve logical model
target, err := s.modelMap.Resolve(req.Model)
if err != nil {
@@ -60,6 +89,13 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
return
}
// 校验 max_output_tokens 不超过模型配置上限
if req.MaxOutputTokens > 0 && target.MaxOutputTokens > 0 && req.MaxOutputTokens > target.MaxOutputTokens {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest,
fmt.Sprintf("max_output_tokens %d exceeds model limit %d", req.MaxOutputTokens, target.MaxOutputTokens), requestID))
return
}
// Get adapter
adapterInst, err := s.registry.Get(target.Provider)
if err != nil {
@@ -75,6 +111,7 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
// Backpressure check: reject low-priority requests under high load
s.backpressure.Update(s.scheduler.RunningCount(), s.scheduler.QueueLength())
w.Header().Set("X-Backpressure-Level", fmt.Sprintf("%d", s.backpressure.Level()))
if !s.backpressure.ShouldAccept(priority) {
reason := s.backpressure.RejectReason(priority)
s.metrics.IncRequest("backpressure_rejected")
@@ -85,9 +122,19 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
// Circuit breaker check
if !s.breaker.AllowRequest() {
s.metrics.IncRequest("circuit_open")
s.metrics.SetBreakerState(1) // open
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, "circuit breaker open, please retry later", requestID))
return
}
// 同步熔断器状态到 Prometheus 指标
switch s.breaker.State() {
case "open":
s.metrics.SetBreakerState(1)
case "half_open":
s.metrics.SetBreakerState(2)
default:
s.metrics.SetBreakerState(0)
}
// Create task
taskID := uuid.New().String()
@@ -103,9 +150,21 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
s.metrics.SetQueueLength(s.scheduler.QueueLength())
// 解析 per-request 超时覆盖
var overrides map[string]int
if req.Timeouts != nil {
overrides = map[string]int{
"queue_ms": req.Timeouts.QueueMs,
"first_token_ms": req.Timeouts.FirstTokenMs,
"inference_ms": req.Timeouts.InferenceMs,
"total_ms": req.Timeouts.TotalMs,
}
}
resolvedTimeouts := s.timeoutMgr.ResolveTimeouts(&s.cfg.Timeouts, overrides)
// Wait for task to be dequeued
queueStart := time.Now()
ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultQueueMs)*time.Millisecond)
ctx, cancel := context.WithTimeout(r.Context(), time.Duration(resolvedTimeouts.QueueMs)*time.Millisecond)
defer cancel()
dequeued, err := s.scheduler.GetNext(ctx)
@@ -127,12 +186,25 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
// Build adapter request — assemble context if session_id provided
messages := req.Messages
if req.SessionID != "" {
sess, err := s.sessions.Get(req.SessionID)
if err == nil && sess != nil {
// 会话访问隔离:校验 session 属于当前 app
var sess *session.Session
var sessErr error
if identity != nil {
sess, sessErr = s.sessions.GetForApp(req.SessionID, identity.AppID)
} else {
sess, sessErr = s.sessions.Get(req.SessionID)
}
if sessErr == nil && sess != nil {
// Load session history and assemble context
history := s.loadSessionHistory(sess)
if len(history) > 0 {
assembler := ctxasm.NewAssembler(&s.cfg.Context)
// 注入摘要器,启用真实 LLM 摘要(受配置控制)
if s.cfg.Context.EnableLLMSummary {
summaryGen := ctxasm.NewAdapterSummaryGenerator(adapterInst, target.ActualModel, s.cfg.Context.SummaryMaxTokens)
summarizer := ctxasm.NewSummarizer(summaryGen, time.Duration(s.cfg.Context.SummaryTimeoutSeconds)*time.Second)
assembler.SetSummarizer(summarizer)
}
result := assembler.Assemble(history, req.Messages, target.ContextWindow, target.MaxOutputTokens, req.ContextPolicy)
messages = result.Messages
}
@@ -155,13 +227,20 @@ func (s *Server) handleChatCompletions(w http.ResponseWriter, r *http.Request) {
}
if req.Stream {
s.handleStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model)
s.handleStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts)
} else {
s.handleNonStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model)
// 非流式请求:如果带有幂等键,捕获响应用于缓存
if req.IdempotencyKey != "" && identity != nil {
crw := &captureResponseWriter{ResponseWriter: w}
s.handleNonStreaming(crw, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts)
s.storeIdempotencyResult(req.IdempotencyKey, identity.AppID, crw.statusCode, crw.body)
} else {
s.handleNonStreaming(w, r, adapterInst, adapterReq, dequeued, requestID, target, req.Model, req.SessionID, req.Messages, identity.AppID, resolvedTimeouts)
}
}
}
func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, _ *router.ModelTarget, logicalModel string) {
func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, target *router.ModelTarget, logicalModel string, sessionID string, originalMessages []api.Message, appID string, tc *connector.TimeoutConfig) {
sse := handler.NewSSEWriter(w)
if sse == nil {
s.scheduler.Complete(tk.ID)
@@ -172,21 +251,120 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter
tk.Transition(task.StateStreaming)
inferenceStart := time.Now()
ch, err := adapterInst.ChatCompletionStream(r.Context(), req)
// 流式超时控制:使用 per-request 解析后的 inference timeout
streamCtx, streamCancel := context.WithTimeout(r.Context(), time.Duration(tc.InferenceMs)*time.Millisecond)
defer streamCancel()
ch, err := adapterInst.ChatCompletionStream(streamCtx, req)
if err != nil {
s.scheduler.Complete(tk.ID)
tk.Transition(task.StateFailed)
s.metrics.IncTask("failed")
// 发送 SSE 错误事件
sse.WriteChunk(map[string]any{
"id": requestID,
"object": "chat.completion.chunk",
"model": logicalModel,
"choices": []map[string]any{
{
"index": 0,
"delta": map[string]any{},
"finish_reason": "error",
},
},
"error": map[string]string{
"code": handler.ErrModelUnavailable,
"message": err.Error(),
},
})
sse.WriteDone()
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrModelUnavailable, err.Error(), requestID))
return
}
inputTokens, outputTokens, err := handler.StreamChatCompletion(sse, ch, requestID, tk.ID, logicalModel)
inputTokens, outputTokens, fullContent, err := handler.StreamChatCompletion(sse, ch, requestID, tk.ID, logicalModel, func() {
tk.FirstTokenMs = int(time.Since(inferenceStart) / time.Millisecond)
s.metrics.ObserveFirstTokenLatency(int64(tk.FirstTokenMs))
})
if err != nil {
s.logger.Error("streaming error", observability.F().Event("stream_error").TaskID(tk.ID).Reason(err.Error()))
tk.Transition(task.StateFailed)
s.metrics.IncTask("failed")
s.breaker.RecordError()
// 检查是否客户端断开
if r.Context().Err() != nil {
s.logger.Info("client disconnected during streaming",
observability.F().Event("client_disconnect").TaskID(tk.ID))
tk.Cancel("client disconnected")
s.metrics.IncCancellation()
} else {
s.logger.Error("streaming error", observability.F().Event("stream_error").TaskID(tk.ID).Reason(err.Error()))
// 尝试非流式 fallback 降级
fallbackResp, fbErr := s.tryOverloadFallback(r.Context(), req, target, logicalModel)
if fbErr == nil && fallbackResp != nil {
tk.Degraded = true
tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond)
tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond)
s.metrics.ObserveGatewayLatency(int64(tk.TotalMs))
tk.Transition(task.StateCompleted)
s.metrics.IncTask("completed")
s.metrics.IncRequest("stream_fallback")
s.metrics.IncDegraded()
s.metrics.AddTokens(fallbackResp.InputTokens, fallbackResp.OutputTokens)
s.breaker.RecordSuccess()
s.scheduler.Complete(tk.ID)
s.metrics.SetRunningTasks(s.scheduler.RunningCount())
s.metrics.SetQueueLength(s.scheduler.QueueLength())
s.usageTracker.Record(appID, fallbackResp.InputTokens, fallbackResp.OutputTokens, false)
// 将降级结果作为单个 SSE chunk 发送
sse.WriteChunk(map[string]any{
"id": requestID,
"object": "chat.completion.chunk",
"model": logicalModel,
"choices": []map[string]any{
{
"index": 0,
"delta": map[string]any{
"content": fallbackResp.Content,
},
"finish_reason": fallbackResp.FinishReason,
},
},
"usage": map[string]int{
"input_tokens": fallbackResp.InputTokens,
"output_tokens": fallbackResp.OutputTokens,
"total_tokens": fallbackResp.InputTokens + fallbackResp.OutputTokens,
},
"degraded": true,
})
sse.WriteDone()
if sessionID != "" && fallbackResp.Content != "" {
s.saveSessionMessages(sessionID, originalMessages, fallbackResp.Content)
}
return
}
// fallback 也失败,发送 SSE 错误事件
sse.WriteChunk(map[string]any{
"id": requestID,
"object": "chat.completion.chunk",
"model": logicalModel,
"choices": []map[string]any{
{
"index": 0,
"delta": map[string]any{},
"finish_reason": "error",
},
},
"error": map[string]string{
"code": handler.ErrInferenceTimeout,
"message": err.Error(),
},
})
sse.WriteDone()
tk.Transition(task.StateFailed)
s.metrics.IncTask("failed")
s.breaker.RecordError()
}
} else {
tk.Transition(task.StateCompleted)
s.metrics.IncTask("completed")
@@ -196,15 +374,30 @@ func (s *Server) handleStreaming(w http.ResponseWriter, r *http.Request, adapter
tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond)
tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond)
// 记录延迟指标
s.metrics.ObserveGatewayLatency(int64(tk.TotalMs))
s.metrics.AddTokens(inputTokens, outputTokens)
s.scheduler.Complete(tk.ID)
s.metrics.SetRunningTasks(s.scheduler.RunningCount())
s.metrics.SetQueueLength(s.scheduler.QueueLength())
s.metrics.IncRequest("stream_ok")
if err != nil {
s.metrics.IncRequest("stream_error")
} else {
s.metrics.IncRequest("stream_ok")
}
// 记录 per-app 用量
s.usageTracker.Record(appID, inputTokens, outputTokens, err != nil)
// 保存对话到 session
if sessionID != "" && fullContent != "" {
s.saveSessionMessages(sessionID, originalMessages, fullContent)
}
}
func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, target *router.ModelTarget, logicalModel string) {
ctx, cancel := context.WithTimeout(r.Context(), time.Duration(s.cfg.Timeouts.DefaultInferenceMs)*time.Millisecond)
func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adapterInst adapter.ModelAdapter, req *adapter.ChatRequest, tk *task.Task, requestID string, target *router.ModelTarget, logicalModel string, sessionID string, originalMessages []api.Message, appID string, tc *connector.TimeoutConfig) {
ctx, cancel := context.WithTimeout(r.Context(), time.Duration(tc.InferenceMs)*time.Millisecond)
defer cancel()
inferenceStart := time.Now()
@@ -228,9 +421,16 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap
s.metrics.SetRunningTasks(s.scheduler.RunningCount())
s.metrics.SetQueueLength(s.scheduler.QueueLength())
// 记录 per-app 用量(降级请求)
s.usageTracker.Record(appID, fallbackResp.InputTokens, fallbackResp.OutputTokens, false)
chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, fallbackResp)
chatResp.Degraded = true
chatResp.Timing = s.buildTiming(tk)
// 通过 HTTP 头暴露关键 timing 指标
w.Header().Set("X-Timing-Queue-Ms", strconv.Itoa(tk.QueueMs))
w.Header().Set("X-Timing-Inference-Ms", strconv.Itoa(tk.InferenceMs))
w.Header().Set("X-Timing-Total-Ms", strconv.Itoa(tk.TotalMs))
handler.WriteJSON(w, http.StatusOK, chatResp)
return
}
@@ -253,6 +453,9 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap
tk.InferenceMs = int(time.Since(inferenceStart) / time.Millisecond)
tk.TotalMs = int(time.Since(tk.CreatedAt) / time.Millisecond)
// 记录延迟指标
s.metrics.ObserveGatewayLatency(int64(tk.TotalMs))
tk.Transition(task.StateCompleted)
s.metrics.IncTask("completed")
s.metrics.IncRequest("ok")
@@ -262,9 +465,21 @@ func (s *Server) handleNonStreaming(w http.ResponseWriter, r *http.Request, adap
s.metrics.SetRunningTasks(s.scheduler.RunningCount())
s.metrics.SetQueueLength(s.scheduler.QueueLength())
// 记录 per-app 用量
s.usageTracker.Record(appID, resp.InputTokens, resp.OutputTokens, false)
chatResp := handler.BuildChatResponse(requestID, tk.ID, logicalModel, resp)
chatResp.Timing = s.buildTiming(tk)
// 通过 HTTP 头暴露关键 timing 指标,便于客户端和监控系统采集
w.Header().Set("X-Timing-Queue-Ms", strconv.Itoa(tk.QueueMs))
w.Header().Set("X-Timing-Inference-Ms", strconv.Itoa(tk.InferenceMs))
w.Header().Set("X-Timing-Total-Ms", strconv.Itoa(tk.TotalMs))
handler.WriteJSON(w, http.StatusOK, chatResp)
// 保存对话到 session
if sessionID != "" {
s.saveSessionMessages(sessionID, originalMessages, resp.Content)
}
}
// tryOverloadFallback attempts fallback strategies when the primary model fails.
@@ -314,10 +529,18 @@ func (s *Server) handleModels(w http.ResponseWriter, r *http.Request) {
models := s.modelMap.List()
data := make([]api.ModelInfo, len(models))
for i, m := range models {
target, err := s.modelMap.Resolve(m)
if err != nil {
data[i] = api.ModelInfo{ID: m, Object: "model", OwnedBy: "edgeai-gateway"}
continue
}
data[i] = api.ModelInfo{
ID: m,
Object: "model",
OwnedBy: "edgeai-gateway",
ID: m,
Object: "model",
OwnedBy: "edgeai-gateway",
Provider: target.Provider,
ContextWindow: target.ContextWindow,
MaxOutputTokens: target.MaxOutputTokens,
}
}
@@ -366,8 +589,46 @@ func (s *Server) handleSessions(w http.ResponseWriter, r *http.Request) {
handler.WriteJSON(w, http.StatusCreated, resp)
case http.MethodGet:
// List sessions (simplified: return empty for now)
handler.WriteJSON(w, http.StatusOK, map[string]any{"sessions": []any{}})
// 列出当前 app 的会话(支持分页)
appID := ""
if identity != nil {
appID = identity.AppID
}
page := 1
pageSize := 20
if v := r.URL.Query().Get("page"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
page = n
}
}
if v := r.URL.Query().Get("page_size"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 {
pageSize = n
}
}
offset := (page - 1) * pageSize
sessions, err := s.sessions.ListByApp(appID, pageSize, offset)
if err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
return
}
// 脱敏:不返回 messages 全文,只返回元信息
items := make([]map[string]any, 0, len(sessions))
for _, sess := range sessions {
items = append(items, map[string]any{
"session_id": sess.ID,
"application_id": sess.ApplicationID,
"user_id": sess.UserID,
"message_count": len(sess.Messages),
"created_at": sess.CreatedAt.Format(time.RFC3339),
"last_active": sess.LastActive.Format(time.RFC3339),
})
}
handler.WriteJSON(w, http.StatusOK, map[string]any{
"sessions": items,
"page": page,
"page_size": pageSize,
})
default:
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
@@ -376,11 +637,19 @@ func (s *Server) handleSessions(w http.ResponseWriter, r *http.Request) {
func (s *Server) handleSessionByID(w http.ResponseWriter, r *http.Request) {
requestID := middleware.GetRequestID(r.Context())
identity := auth.GetAppIdentityFromRequest(r)
sessionID := strings.TrimPrefix(r.URL.Path, "/v1/sessions/")
switch r.Method {
case http.MethodGet:
sess, err := s.sessions.Get(sessionID)
// 会话访问隔离:校验 session 属于当前 app
var sess *session.Session
var err error
if identity != nil {
sess, err = s.sessions.GetForApp(sessionID, identity.AppID)
} else {
sess, err = s.sessions.Get(sessionID)
}
if err != nil || sess == nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, "session not found", requestID))
return
@@ -388,9 +657,17 @@ func (s *Server) handleSessionByID(w http.ResponseWriter, r *http.Request) {
handler.WriteJSON(w, http.StatusOK, sess)
case http.MethodDelete:
if err := s.sessions.Delete(sessionID); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
return
// 会话访问隔离:校验 session 属于当前 app
if identity != nil {
if err := s.sessions.DeleteForApp(sessionID, identity.AppID); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInvalidRequest, err.Error(), requestID))
return
}
} else {
if err := s.sessions.Delete(sessionID); err != nil {
handler.WriteError(w, handler.NewGatewayErrorWithID(handler.ErrInternalError, err.Error(), requestID))
return
}
}
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "deleted"})
@@ -407,6 +684,28 @@ func (s *Server) loadSessionHistory(sess *session.Session) []api.Message {
return sess.Messages
}
// saveSessionMessages 将用户消息和助手回复保存到会话历史。
// 当配置了 MaxSessionMessages 时,自动裁剪最旧消息以防止无限增长。
func (s *Server) saveSessionMessages(sessionID string, userMessages []api.Message, assistantContent string) {
if sessionID == "" {
return
}
maxMsgs := s.cfg.Context.MaxSessionMessages
for _, msg := range userMessages {
if err := s.sessions.AddMessageWithLimit(sessionID, msg, maxMsgs); err != nil {
s.logger.Error("failed to save user message to session",
observability.F().Event("session_save_error").Reason(err.Error()))
}
}
if assistantContent != "" {
assistantMsg := api.Message{Role: "assistant", Content: assistantContent}
if err := s.sessions.AddMessageWithLimit(sessionID, assistantMsg, maxMsgs); err != nil {
s.logger.Error("failed to save assistant message to session",
observability.F().Event("session_save_error").Reason(err.Error()))
}
}
}
// buildTiming constructs Timing metadata from a task.
func (s *Server) buildTiming(tk *task.Task) *api.Timing {
return &api.Timing{
@@ -422,7 +721,59 @@ func (s *Server) handleTasks(w http.ResponseWriter, r *http.Request) {
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "method not allowed"))
return
}
handler.WriteJSON(w, http.StatusOK, map[string]any{"tasks": []any{}})
// 分页参数
page := 1
pageSize := 20
if v := r.URL.Query().Get("page"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
page = n
}
}
if v := r.URL.Query().Get("page_size"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 && n <= 100 {
pageSize = n
}
}
offset := (page - 1) * pageSize
// 状态过滤
status := r.URL.Query().Get("status")
tasks, total := s.scheduler.ListTasks(status, pageSize, offset)
items := make([]map[string]any, 0, len(tasks))
for _, tk := range tasks {
item := map[string]any{
"task_id": tk.ID,
"request_id": tk.RequestID,
"session_id": tk.SessionID,
"state": string(tk.GetState()),
"priority": int(tk.Priority),
"logical_model": tk.LogicalModel,
"stream": tk.Stream,
"created_at": tk.CreatedAt.Format(time.RFC3339),
"degraded": tk.Degraded,
}
if tk.StartedAt != nil {
item["started_at"] = tk.StartedAt.Format(time.RFC3339)
}
if tk.CompletedAt != nil {
item["completed_at"] = tk.CompletedAt.Format(time.RFC3339)
}
if tk.ErrorMessage != "" {
item["error_message"] = tk.ErrorMessage
}
item["timing"] = s.buildTiming(tk)
items = append(items, item)
}
handler.WriteJSON(w, http.StatusOK, map[string]any{
"tasks": items,
"total": total,
"page": page,
"page_size": pageSize,
})
}
func (s *Server) handleTaskByID(w http.ResponseWriter, r *http.Request) {
+147
View File
@@ -0,0 +1,147 @@
package server
import (
"net/http"
"time"
"github.com/edgeai/gateway/internal/handler"
)
// idempotencyEntry 幂等键缓存条目。
// 记录请求处理状态和响应,用于在重复请求时返回缓存结果。
type idempotencyEntry struct {
status int
response []byte
createdAt time.Time
inFlight bool // 正在处理中
}
const (
// idempotencyTTL 幂等键缓存存活时间。
idempotencyTTL = 10 * time.Minute
// idempotencyCleanupInterval 清理间隔。
idempotencyCleanupInterval = 5 * time.Minute
)
// checkIdempotency 检查幂等键,如果重复请求则返回缓存的响应。
// 如果是首次请求,返回 nil 表示可以继续处理。
// 如果是正在处理中的重复请求,返回 409 Conflict。
func (s *Server) checkIdempotency(w http.ResponseWriter, key, appID string) bool {
if key == "" {
return false
}
cacheKey := appID + ":" + key
s.idempotencyMu.Lock()
entry, exists := s.idempotencyCache[cacheKey]
if exists {
if entry.inFlight {
// 正在处理中,返回 409
s.idempotencyMu.Unlock()
handler.WriteError(w, handler.NewGatewayError(handler.ErrInvalidRequest, "duplicate idempotency key, request already in progress"))
return true
}
// 检查是否过期
if time.Since(entry.createdAt) > idempotencyTTL {
delete(s.idempotencyCache, cacheKey)
} else {
// 返回缓存的响应
s.idempotencyMu.Unlock()
w.Header().Set("Content-Type", "application/json")
w.Header().Set("X-Idempotent-Replay", "true")
w.WriteHeader(entry.status)
w.Write(entry.response)
return true
}
}
// 标记为正在处理中
s.idempotencyCache[cacheKey] = &idempotencyEntry{
inFlight: true,
createdAt: time.Now(),
}
s.idempotencyMu.Unlock()
return false
}
// storeIdempotencyResult 存储幂等请求的响应结果。
func (s *Server) storeIdempotencyResult(key, appID string, status int, response []byte) {
if key == "" {
return
}
cacheKey := appID + ":" + key
s.idempotencyMu.Lock()
s.idempotencyCache[cacheKey] = &idempotencyEntry{
status: status,
response: response,
createdAt: time.Now(),
inFlight: false,
}
s.idempotencyMu.Unlock()
}
// cleanupIdempotencyCache 清理过期的幂等键缓存。
func (s *Server) cleanupIdempotencyCache() {
s.idempotencyMu.Lock()
defer s.idempotencyMu.Unlock()
now := time.Now()
for k, entry := range s.idempotencyCache {
if now.Sub(entry.createdAt) > idempotencyTTL {
delete(s.idempotencyCache, k)
}
}
}
// startIdempotencyCleanup 启动后台清理任务,定期清理过期的幂等键。
// 返回停止函数。
func (s *Server) startIdempotencyCleanup() func() {
ticker := time.NewTicker(idempotencyCleanupInterval)
stopCh := make(chan struct{})
go func() {
for {
select {
case <-ticker.C:
s.cleanupIdempotencyCache()
case <-stopCh:
ticker.Stop()
return
}
}
}()
return func() {
close(stopCh)
}
}
// captureResponseWriter 捕获响应状态和响应体,用于幂等性缓存。
type captureResponseWriter struct {
http.ResponseWriter
statusCode int
body []byte
}
func (crw *captureResponseWriter) WriteHeader(code int) {
crw.statusCode = code
crw.ResponseWriter.WriteHeader(code)
}
func (crw *captureResponseWriter) Write(b []byte) (int, error) {
if crw.statusCode == 0 {
crw.statusCode = 200
}
crw.body = append(crw.body, b...)
return crw.ResponseWriter.Write(b)
}
func (crw *captureResponseWriter) Flush() {
if f, ok := crw.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
+296 -36
View File
@@ -7,33 +7,50 @@ import (
"os"
"path/filepath"
"strings"
"sync"
"time"
"github.com/edgeai/gateway/internal/adapter"
"github.com/edgeai/gateway/internal/auth"
"github.com/edgeai/gateway/internal/config"
"github.com/edgeai/gateway/internal/connector"
"github.com/edgeai/gateway/internal/handler"
"github.com/edgeai/gateway/internal/middleware"
"github.com/edgeai/gateway/internal/observability"
"github.com/edgeai/gateway/internal/ratelimit"
"github.com/edgeai/gateway/internal/router"
"github.com/edgeai/gateway/internal/scheduler"
"github.com/edgeai/gateway/internal/session"
"github.com/edgeai/gateway/internal/usage"
)
// Server is the main HTTP server for the AI gateway.
type Server struct {
cfg *config.Config
logger *observability.Logger
metrics *observability.Metrics
HTTPSrv *http.Server
auth *auth.Authenticator
registry *adapter.Registry
modelMap *router.LogicalModelMapping
scheduler *scheduler.Scheduler
sessions *session.Store
breaker *scheduler.CircuitBreaker
backpressure *scheduler.BackpressureManager
overload *router.OverloadResolver
cfg *config.Config
logger *observability.Logger
metrics *observability.Metrics
HTTPSrv *http.Server
auth *auth.Authenticator
registry *adapter.Registry
modelMap *router.LogicalModelMapping
scheduler *scheduler.Scheduler
sessions *session.Store
breaker *scheduler.CircuitBreaker
backpressure *scheduler.BackpressureManager
overload *router.OverloadResolver
rateLimiter *ratelimit.Limiter
usageTracker *usage.Tracker
audit *observability.AuditLogger
timeoutMgr *connector.TimeoutManager
sessionCleanup func() // 停止会话清理任务的函数
adminSrv *http.Server // Admin API 独立端口服务器
healthCache map[string]bool // adapter 健康缓存
healthCacheTime time.Time // 缓存时间
healthCacheMu sync.RWMutex // 缓存锁
idempotencyCache map[string]*idempotencyEntry // 幂等键缓存
idempotencyMu sync.Mutex // 幂等缓存锁
}
// New creates a new Server instance with all components wired.
@@ -105,36 +122,115 @@ func New(cfg *config.Config, logger *observability.Logger) (*Server, error) {
// Initialize metrics
metrics := observability.NewMetrics()
s := &Server{
cfg: cfg,
logger: logger,
metrics: metrics,
auth: authenticator,
registry: registry,
modelMap: modelMap,
scheduler: sched,
sessions: sessionStore,
breaker: breaker,
backpressure: bp,
overload: overloadResolver,
// Initialize per-app rate limiter (令牌桶)
refillPerSec := float64(cfg.Auth.RateLimitPerMinute) / 60.0
rateLimiter := ratelimit.NewLimiter(float64(cfg.Auth.RateLimitBurst), refillPerSec)
// Initialize per-app usage tracker (SQLite 持久化)
usageDBPath := filepath.Join(filepath.Dir(dbPath), "usage.db")
usageTracker, err := usage.NewTracker(time.Duration(cfg.Auth.UsageWindowMinutes)*time.Minute, usageDBPath, cfg.Observability.AuditRetentionDays)
if err != nil {
return nil, fmt.Errorf("init usage tracker: %w", err)
}
timeoutMgr := connector.NewTimeoutManager(&cfg.Timeouts)
auditLogger := observability.NewAuditLogger(filepath.Dir(dbPath), logger)
s := &Server{
cfg: cfg,
logger: logger,
metrics: metrics,
auth: authenticator,
registry: registry,
modelMap: modelMap,
scheduler: sched,
sessions: sessionStore,
breaker: breaker,
backpressure: bp,
overload: overloadResolver,
rateLimiter: rateLimiter,
usageTracker: usageTracker,
audit: auditLogger,
timeoutMgr: timeoutMgr,
healthCache: make(map[string]bool),
idempotencyCache: make(map[string]*idempotencyEntry),
}
// 启动会话 TTL 自动清理任务(带日志回调)
s.sessionCleanup = sessionStore.StartCleanupTask(cfg.Context.SessionIdleTTLMinutes, 10, func(deleted int64) {
logger.Info("session cleanup completed",
observability.F().
Event("session_cleanup").
Set("deleted_sessions", deleted))
})
// 启动幂等键缓存清理任务
idempotencyCleanup := s.startIdempotencyCleanup()
_ = idempotencyCleanup // 进程退出时自动回收
mux := http.NewServeMux()
s.registerRoutes(mux)
// Apply middleware chain (order: Recovery → Logging → RequestID → BodyLimit → Auth → handler)
// Apply middleware chain (order: Recovery → Logging → CORS → Auth → RateLimit → RequestID → BodyLimit → handler)
h := middleware.RequestID(mux)
h = middleware.BodyLimit(cfg.Server.MaxRequestBodyMB)(h)
h = middleware.RateLimit(func(appID string) (bool, int, int) {
info := s.rateLimiter.AllowWithInfo(appID)
return info.Allowed, info.Limit, info.Remaining
}, func(r *http.Request) string {
if id := auth.GetAppIdentityFromRequest(r); id != nil {
return id.AppID
}
return ""
}, map[string]bool{"/health": true, "/ready": true, "/metrics": true})(h)
h = s.auth.Middleware(h)
h = middleware.CORS(cfg.Server.CORSAllowedOrigins)(h)
h = middleware.Gzip(h)
h = middleware.Logging(logger)(h)
h = middleware.Recovery(logger)(h)
s.HTTPSrv = &http.Server{
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.Port),
Handler: h,
ReadTimeout: 30 * time.Second,
WriteTimeout: 0, // no write timeout for SSE
IdleTimeout: 120 * time.Second,
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.Port),
Handler: h,
ReadTimeout: 30 * time.Second,
ReadHeaderTimeout: 10 * time.Second, // 防止 slowloris 攻击
WriteTimeout: 0, // no write timeout for SSE
IdleTimeout: 120 * time.Second,
MaxHeaderBytes: 1 << 20, // 1MB 限制请求头大小
}
// Admin API 独立端口(如果配置了 admin_port)
if cfg.Server.AdminPort > 0 && cfg.Server.AdminPort != cfg.Server.Port {
adminMux := http.NewServeMux()
s.registerAdminRoutes(adminMux)
// Admin 中间件链(与主服务器相同,但额外加 CORS)
adminH := middleware.RequestID(adminMux)
adminH = middleware.BodyLimit(cfg.Server.MaxRequestBodyMB)(adminH)
adminH = middleware.RateLimit(func(appID string) (bool, int, int) {
info := s.rateLimiter.AllowWithInfo(appID)
return info.Allowed, info.Limit, info.Remaining
}, func(r *http.Request) string {
if id := auth.GetAppIdentityFromRequest(r); id != nil {
return id.AppID
}
return ""
}, nil)(adminH)
adminH = s.auth.Middleware(adminH)
adminH = middleware.CORS(cfg.Server.CORSAllowedOrigins)(adminH)
adminH = middleware.Logging(logger)(adminH)
adminH = middleware.Recovery(logger)(adminH)
s.adminSrv = &http.Server{
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.AdminPort),
Handler: adminH,
ReadTimeout: 30 * time.Second,
ReadHeaderTimeout: 10 * time.Second, // 防止 slowloris 攻击
WriteTimeout: 30 * time.Second,
IdleTimeout: 120 * time.Second,
MaxHeaderBytes: 1 << 20, // 1MB 限制请求头大小
}
}
return s, nil
@@ -161,53 +257,182 @@ func (s *Server) registerRoutes(mux *http.ServeMux) {
mux.HandleFunc("/v1/tasks/", s.handleTaskByID)
}
// registerAdminRoutes 注册管理 API 路由到独立的 mux(运行在 admin_port)。
func (s *Server) registerAdminRoutes(mux *http.ServeMux) {
// Admin API: API Key 管理(需要 admin 权限)
mux.Handle("/v1/admin/keys", auth.RequireAdmin(http.HandlerFunc(s.handleAdminKeys)))
mux.Handle("/v1/admin/keys/", auth.RequireAdmin(http.HandlerFunc(s.handleAdminKeyByID)))
// Admin API: 用量统计(需要 admin 权限)
mux.Handle("/v1/admin/usage", auth.RequireAdmin(http.HandlerFunc(s.handleAdminUsage)))
mux.Handle("/v1/admin/usage/", auth.RequireAdmin(http.HandlerFunc(s.handleAdminUsageByApp)))
// Admin API: per-app 限流管理(需要 admin 权限)
mux.Handle("/v1/admin/ratelimit/", auth.RequireAdmin(http.HandlerFunc(s.handleAdminRateLimit)))
// Admin API: 熔断器状态(需要 admin 权限)
mux.Handle("/v1/admin/circuit-breaker", auth.RequireAdmin(http.HandlerFunc(s.handleAdminCircuitBreaker)))
// Admin API: 调度器状态(需要 admin 权限)
mux.Handle("/v1/admin/scheduler", auth.RequireAdmin(http.HandlerFunc(s.handleAdminScheduler)))
// Admin API: 配置热重载(需要 admin 权限)
mux.Handle("/v1/admin/config/reload", auth.RequireAdmin(http.HandlerFunc(s.handleAdminConfigReload)))
// Admin API: 健康检查(复用主服务器逻辑)
mux.HandleFunc("/health", s.handleHealth)
mux.HandleFunc("/ready", s.handleReady)
}
// Authenticator returns the authenticator instance (for testing/management).
func (s *Server) Authenticator() *auth.Authenticator {
return s.auth
}
// Start begins listening for HTTP requests.
// 如果配置了 admin_port,同时启动 Admin API 服务器。
// 启动时对所有已注册 adapter 执行健康预检,不可达的 adapter 记录警告但不阻止启动。
func (s *Server) Start() error {
s.logger.Info("http server starting", observability.F().
Event("server_start").
Set("addr", s.HTTPSrv.Addr))
// Adapter 启动健康预检
s.checkAdaptersHealth()
// 启动 Admin API 服务器(独立 goroutine
if s.adminSrv != nil {
s.logger.Info("admin api server starting", observability.F().
Event("admin_server_start").
Set("addr", s.adminSrv.Addr))
go func() {
if err := s.adminSrv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
s.logger.Error("admin api server error", observability.F().
Event("admin_server_error").
Reason(err.Error()))
}
}()
}
return s.HTTPSrv.ListenAndServe()
}
// checkAdaptersHealth 对所有已注册 adapter 执行健康预检。
// 不可达的 adapter 记录警告但不阻止启动,允许部分降级运行。
// 每个 adapter 使用独立超时,避免单个 adapter 阻塞拖累其他。
func (s *Server) checkAdaptersHealth() {
for _, name := range s.registry.Names() {
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
a, _ := s.registry.Get(name)
if err := a.HealthCheck(ctx); err != nil {
s.logger.Warn("adapter health check failed on startup",
observability.F().
Event("adapter_health_check_failed").
Set("adapter", name).
Reason(err.Error()))
} else {
s.logger.Info("adapter health check passed",
observability.F().
Event("adapter_health_check_ok").
Set("adapter", name))
}
cancel()
}
}
// Shutdown gracefully shuts down the server.
// 等待最多 30 秒让进行中的请求完成,记录未完成任务状态。
func (s *Server) Shutdown() error {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
// 记录关机时的任务状态
runningCount := s.scheduler.RunningCount()
queueLength := s.scheduler.QueueLength()
s.logger.Info("server shutting down",
observability.F().
Event("server_shutdown").
Set("running_tasks", runningCount).
Set("queued_tasks", queueLength))
// 等待运行中任务完成(超时时间从配置读取)
drainSeconds := s.cfg.Timeouts.ShutdownDrainSeconds
if drainSeconds <= 0 {
drainSeconds = 10
}
if runningCount > 0 {
s.logger.Info("waiting for running tasks to complete",
observability.F().
Event("shutdown_drain").
Set("running_tasks", runningCount).
Set("drain_timeout_seconds", drainSeconds))
for i := 0; i < drainSeconds*10 && s.scheduler.RunningCount() > 0; i++ {
time.Sleep(100 * time.Millisecond)
}
remaining := s.scheduler.RunningCount()
if remaining > 0 {
s.logger.Warn("shutdown: tasks still running after drain timeout",
observability.F().
Event("shutdown_drain_timeout").
Set("remaining_tasks", remaining))
} else {
s.logger.Info("all running tasks completed",
observability.F().Event("shutdown_drain_complete"))
}
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
s.scheduler.Stop()
if s.sessionCleanup != nil {
s.sessionCleanup()
}
if s.sessions != nil {
s.sessions.Close()
}
if s.auth != nil {
s.auth.Close()
}
if s.usageTracker != nil {
s.usageTracker.Close()
}
if s.audit != nil {
s.audit.Close()
}
if s.adminSrv != nil {
s.adminSrv.Shutdown(ctx)
}
return s.HTTPSrv.Shutdown(ctx)
}
func (s *Server) handleHealth(w http.ResponseWriter, r *http.Request) {
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "ok"})
// Liveness 探针:进程存活即返回 200
handler.WriteJSON(w, http.StatusOK, map[string]string{"status": "alive"})
}
func (s *Server) handleReady(w http.ResponseWriter, r *http.Request) {
// Readiness 探针:检查 adapter 健康状态 + DB 连通性
// 使用 5 秒 TTL 缓存,避免高频探针请求打满 adapter
w.Header().Set("Content-Type", "application/json")
ready := true
reasons := []string{}
for _, name := range s.registry.Names() {
a, _ := s.registry.Get(name)
if err := a.HealthCheck(r.Context()); err != nil {
// 检查所有 adapter 健康状态(带缓存)
adapterHealth := s.getCachedAdapterHealth(r.Context())
for name, ok := range adapterHealth {
if !ok {
ready = false
reasons = append(reasons, fmt.Sprintf("%s: %v", name, err))
reasons = append(reasons, fmt.Sprintf("adapter %s: unhealthy", name))
}
}
// 检查 scheduler 是否正常
if s.scheduler.RunningCount() >= s.cfg.Scheduler.MaxRunningTasks {
ready = false
reasons = append(reasons, "scheduler: at max capacity")
}
if ready {
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{"status":"ready"}`))
@@ -217,6 +442,41 @@ func (s *Server) handleReady(w http.ResponseWriter, r *http.Request) {
}
}
// getCachedAdapterHealth 返回 adapter 健康状态,使用 5 秒 TTL 缓存。
// 缓存过期后异步刷新,首次请求或缓存过期时同步调用 adapter HealthCheck。
func (s *Server) getCachedAdapterHealth(ctx context.Context) map[string]bool {
const healthCacheTTL = 5 * time.Second
s.healthCacheMu.RLock()
if time.Since(s.healthCacheTime) < healthCacheTTL && len(s.healthCache) > 0 {
result := make(map[string]bool, len(s.healthCache))
for k, v := range s.healthCache {
result[k] = v
}
s.healthCacheMu.RUnlock()
return result
}
s.healthCacheMu.RUnlock()
// 缓存过期,同步执行健康检查
result := make(map[string]bool)
for _, name := range s.registry.Names() {
a, _ := s.registry.Get(name)
if err := a.HealthCheck(ctx); err != nil {
result[name] = false
} else {
result[name] = true
}
}
s.healthCacheMu.Lock()
s.healthCache = result
s.healthCacheTime = time.Now()
s.healthCacheMu.Unlock()
return result
}
func extractDBPath(connStr string) string {
if strings.HasPrefix(connStr, "sqlite://") {
return strings.TrimPrefix(connStr, "sqlite://")