package server import ( "context" "fmt" "net/http" "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 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. func New(cfg *config.Config, logger *observability.Logger) (*Server, error) { // Ensure data directory exists dbPath := extractDBPath(cfg.Storage.SessionDB) if dbPath != "" { os.MkdirAll(filepath.Dir(dbPath), 0755) } // Initialize auth authPath := filepath.Join(filepath.Dir(dbPath), "auth.db") authenticator, err := auth.NewAuthenticator(authPath, logger) if err != nil { return nil, fmt.Errorf("init auth: %w", err) } // Initialize session store sessionStore, err := session.NewStore(dbPath) if err != nil { return nil, fmt.Errorf("init session store: %w", err) } // Initialize adapter registry registry := adapter.NewRegistry() // Initialize logical model mapping modelMap := router.NewLogicalModelMapping(cfg) // Initialize overload resolver overloadResolver := router.NewOverloadResolver(&cfg.Routing, modelMap) // Register adapters for each unique endpoint registered := make(map[string]bool) for _, mc := range cfg.Models { key := mc.Provider + "|" + mc.Endpoint if !registered[key] { switch mc.Provider { case "ollama": registry.Register(mc.Provider, adapter.NewOllamaAdapter(mc.Endpoint)) case "vllm": registry.Register(mc.Provider, adapter.NewVLLMAdapter(mc.Endpoint)) } registered[key] = true } } // Initialize scheduler sched := scheduler.NewScheduler(&cfg.Scheduler, logger) // Initialize circuit breaker breaker := scheduler.NewCircuitBreaker( cfg.CircuitBreaker.ErrorRateThreshold, cfg.CircuitBreaker.MinRequests, cfg.CircuitBreaker.WindowSeconds, cfg.CircuitBreaker.OpenDurationSeconds, cfg.CircuitBreaker.HalfOpenMaxRequests, ) // Initialize backpressure manager bp := scheduler.NewBackpressureManager( cfg.Scheduler.MaxRunningTasks, cfg.Scheduler.MaxQueuedTasks, cfg.Backpressure.Level1Threshold, cfg.Backpressure.Level2Threshold, cfg.Backpressure.Level3Threshold, ) // Initialize metrics metrics := observability.NewMetrics() // 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 → 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, 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 } func (s *Server) registerRoutes(mux *http.ServeMux) { // Health and readiness mux.HandleFunc("/health", s.handleHealth) mux.HandleFunc("/ready", s.handleReady) // Metrics mux.HandleFunc(s.cfg.Observability.MetricsPath, s.metrics.Handler()) // OpenAI-compatible API mux.HandleFunc("/v1/chat/completions", s.handleChatCompletions) mux.HandleFunc("/v1/models", s.handleModels) // Session management mux.HandleFunc("/v1/sessions", s.handleSessions) mux.HandleFunc("/v1/sessions/", s.handleSessionByID) // Task management mux.HandleFunc("/v1/tasks", s.handleTasks) 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 { // 记录关机时的任务状态 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) { // 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{} // 检查所有 adapter 健康状态(带缓存) adapterHealth := s.getCachedAdapterHealth(r.Context()) for name, ok := range adapterHealth { if !ok { ready = false 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"}`)) } else { w.WriteHeader(http.StatusServiceUnavailable) fmt.Fprintf(w, `{"status":"not_ready","reasons":["%s"]}`, strings.Join(reasons, `","`)) } } // 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://") } return connStr }