da9c8334d8
- 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 全链路传播
486 lines
16 KiB
Go
486 lines
16 KiB
Go
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
|
||
}
|