feat: 十轮网关优化 - 安全加固/可观测性/性能/可靠性
- 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:
+296
-36
@@ -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://")
|
||||
|
||||
Reference in New Issue
Block a user