Files
selfrelease da9c8334d8
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
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 全链路传播
2026-08-03 15:43:11 +08:00

486 lines
16 KiB
Go
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}