- New CLI commands: key, user - New internal modules: auth (middleware, password, user), db (db, logs, products), email, gateway (gateway, hooks, middlewares), handler (admin, auth, dev, log, product, settings), middleware/geoblock, ratelimit - New demo frontend page - Updated config, server, analytics, stats, and existing frontend pages
369 lines
11 KiB
Go
369 lines
11 KiB
Go
package server
|
||
|
||
import (
|
||
"context"
|
||
"fmt"
|
||
"io"
|
||
"log"
|
||
"net/http"
|
||
"os"
|
||
"os/signal"
|
||
"path/filepath"
|
||
"strings"
|
||
"sync"
|
||
"syscall"
|
||
"time"
|
||
|
||
"api-server/docs"
|
||
"api-server/internal/analytics"
|
||
"api-server/internal/auth"
|
||
"api-server/internal/config"
|
||
"api-server/internal/db"
|
||
"api-server/internal/email"
|
||
"api-server/internal/gateway"
|
||
"api-server/internal/handler"
|
||
"api-server/internal/middleware"
|
||
"api-server/internal/ratelimit"
|
||
"api-server/web"
|
||
|
||
"github.com/gin-gonic/gin"
|
||
swaggerFiles "github.com/swaggo/files"
|
||
ginSwagger "github.com/swaggo/gin-swagger"
|
||
)
|
||
|
||
// Server HTTP 服务
|
||
type Server struct {
|
||
cfg *config.Config
|
||
engine *gin.Engine
|
||
stats *analytics.Engine
|
||
geoip *analytics.GeoIP
|
||
srv *http.Server
|
||
database *db.DB
|
||
keyLimit *ratelimit.Limiter
|
||
ipLimit *ratelimit.Limiter
|
||
}
|
||
|
||
// New 创建服务实例
|
||
func New(cfg *config.Config) (*Server, error) {
|
||
gin.SetMode(cfg.Server.Mode)
|
||
|
||
// 初始化 JWT 密钥(优先从环境变量 JWT_SECRET 加载)
|
||
auth.InitJWTSecret()
|
||
|
||
// 初始化 IP 地理位置
|
||
geoip, err := analytics.NewGeoIP(cfg.GeoIP.DBPathV4, cfg.GeoIP.DBPathV6)
|
||
if err != nil {
|
||
log.Printf("[server] GeoIP 加载失败: %v(已禁用)", err)
|
||
geoip = nil
|
||
}
|
||
|
||
// 初始化统计引擎
|
||
statsEngine := analytics.NewEngine(geoip)
|
||
|
||
// 初始化数据库
|
||
database, err := db.Open(db.Config{
|
||
Driver: db.Driver(cfg.Database.Driver),
|
||
DSN: cfg.Database.DSN,
|
||
})
|
||
if err != nil {
|
||
return nil, fmt.Errorf("数据库初始化失败: %w", err)
|
||
}
|
||
|
||
// 初始化邮件发送
|
||
settings, _ := database.GetSystemSettings()
|
||
emailSender := email.New(email.Config{
|
||
Host: settings["smtp_host"],
|
||
Port: settings["smtp_port"],
|
||
User: settings["smtp_user"],
|
||
Pass: settings["smtp_pass"],
|
||
From: settings["smtp_from"],
|
||
})
|
||
|
||
// 初始化限流器
|
||
keyLimit := ratelimit.New(cfg.RateLimit.PerKeyRPS, cfg.RateLimit.PerKeyBurst)
|
||
ipLimit := ratelimit.New(cfg.RateLimit.PerIPRPS, cfg.RateLimit.PerIPBurst)
|
||
|
||
// 初始化 gin
|
||
router := gin.New()
|
||
router.Use(gin.LoggerWithFormatter(func(param gin.LogFormatterParams) string {
|
||
return fmt.Sprintf("[%s] %s %s %d %s %s\n",
|
||
param.TimeStamp.Format("15:04:05"), param.ClientIP,
|
||
param.Method, param.StatusCode, param.Latency, param.Path)
|
||
}))
|
||
router.Use(gin.Recovery())
|
||
router.Use(middleware.Analytics(statsEngine))
|
||
geoBlocker := middleware.NewGeoBlocker(database, geoip)
|
||
router.Use(geoBlocker.Handler())
|
||
|
||
// ── 网关层:动态路由中间件 ────────────────────
|
||
// 在路由注册前注入,拦截匹配 DB 中 api_products 的请求
|
||
gatewayCache := gateway.NewMemoryCache()
|
||
gatewayBreakers := &sync.Map{}
|
||
gatewayHooks := gateway.DefaultHookChain(database, gatewayCache, gatewayBreakers)
|
||
router.Use(gateway.DynamicRoute(database, gatewayHooks))
|
||
|
||
// ── 创建所有处理器 ────────────────────────────
|
||
base := handler.New(cfg)
|
||
authHandler := handler.NewAuthHandler(database, cfg.JWT.CookieSecure)
|
||
devHandler := handler.NewDevHandler(database)
|
||
adminHandler := handler.NewAdminHandler(database)
|
||
extractHandler := handler.NewExtractHandler()
|
||
contentHandler := handler.NewContentHandler()
|
||
statsHandler := handler.NewStats(statsEngine)
|
||
proxyHandler := handler.NewProxyHandler()
|
||
productHandler := handler.NewAPIProductHandler(database)
|
||
logHandler := handler.NewLogHandler(database)
|
||
settingsHandler := handler.NewSettingsHandler(database, emailSender)
|
||
|
||
// ── 路由注册辅助 ──────────────────────────────
|
||
// registerAllRoutes 在给定 Group 上注册所有路由(消除 /api 和 /api/v1 重复)
|
||
registerAllRoutes := func(r *gin.RouterGroup, biz *gin.RouterGroup) {
|
||
r.GET("/health", base.Health)
|
||
r.GET("/ping", base.Ping)
|
||
authHandler.RegisterRoutes(r)
|
||
devHandler.RegisterRoutes(r)
|
||
adminHandler.RegisterRoutes(r)
|
||
settingsHandler.RegisterRoutes(r)
|
||
|
||
extractHandler.RegisterRoutes(biz)
|
||
contentHandler.RegisterRoutes(biz)
|
||
statsHandler.RegisterRoutes(biz)
|
||
proxyHandler.RegisterRoutes(biz)
|
||
}
|
||
|
||
// ── API 路由 ──────────────────────────────────
|
||
api := router.Group("/api")
|
||
|
||
// 业务 API 中间件链
|
||
biz := api.Group("")
|
||
if cfg.Auth.Enabled {
|
||
biz.Use(auth.RequireAPIKey(database))
|
||
} else {
|
||
biz.Use(auth.OptionalAPIKey(database))
|
||
}
|
||
if cfg.RateLimit.Enabled {
|
||
biz.Use(authRateLimitIP(ipLimit))
|
||
}
|
||
|
||
registerAllRoutes(api, biz)
|
||
|
||
// API 版本化:/api/v1 复用同一路由注册
|
||
v1 := router.Group("/api/v1")
|
||
v1Biz := v1.Group("")
|
||
if cfg.Auth.Enabled {
|
||
v1Biz.Use(auth.RequireAPIKey(database))
|
||
} else {
|
||
v1Biz.Use(auth.OptionalAPIKey(database))
|
||
}
|
||
if cfg.RateLimit.Enabled {
|
||
v1Biz.Use(authRateLimitIP(ipLimit))
|
||
}
|
||
|
||
registerAllRoutes(v1, v1Biz)
|
||
|
||
// ── API 产品抽象层路由 ───────────────────────
|
||
// 公开:API 目录(无需认证)
|
||
productHandler.RegisterPublicRoutes(api)
|
||
|
||
// 开发者订阅端点(已由 devHandler 的 /developer 组提供 auth 中间件,这里复用)
|
||
devGroup := api.Group("/developer")
|
||
devGroup.Use(auth.RequireAuth(database))
|
||
productHandler.RegisterDevRoutes(devGroup)
|
||
// 调用日志
|
||
devGroup.GET("/logs", logHandler.ListMyLogs)
|
||
|
||
// 管理员端点(已由 adminHandler 的 /admin 组提供 auth+admin 中间件,这里复用)
|
||
adminGroup := api.Group("/admin")
|
||
adminGroup.Use(auth.RequireAuth(database), auth.RequireAdmin())
|
||
productHandler.RegisterAdminRoutes(adminGroup)
|
||
// 调用日志 + 运营统计
|
||
adminGroup.GET("/logs", logHandler.ListAllLogs)
|
||
adminGroup.GET("/logs/summary", logHandler.LogSummary)
|
||
adminGroup.GET("/product-stats", logHandler.ProductStats)
|
||
adminGroup.GET("/subscriptions", logHandler.ListAllSubs)
|
||
adminGroup.PUT("/subscriptions/:id/quota", logHandler.SetSubQuota)
|
||
|
||
// Swagger
|
||
docs.SwaggerInfo.Host = fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.Port)
|
||
router.GET("/swagger/*any", ginSwagger.WrapHandler(swaggerFiles.Handler))
|
||
|
||
// ── 前端路由(SPA 壳)──────────────────────────────────
|
||
noCache := func(c *gin.Context) {
|
||
c.Header("Cache-Control", "no-store")
|
||
}
|
||
|
||
// SPA 入口:所有前端页面由 index.html 承载
|
||
router.GET("/", func(c *gin.Context) {
|
||
noCache(c)
|
||
serveStaticFile(c, "index.html")
|
||
})
|
||
|
||
// Demo 独立页面(需登录)
|
||
router.GET("/demo/index.html", auth.RequireAuth(database), func(c *gin.Context) {
|
||
noCache(c)
|
||
serveStaticFile(c, "demo/index.html")
|
||
})
|
||
|
||
// 向后兼容
|
||
router.GET("/video", func(c *gin.Context) { c.Redirect(http.StatusFound, "/demo") })
|
||
router.GET("/video/", func(c *gin.Context) { c.Redirect(http.StatusFound, "/demo") })
|
||
router.GET("/developers", func(c *gin.Context) { c.Redirect(http.StatusFound, "/") })
|
||
|
||
// NoRoute: 静态文件 + SPA fallback
|
||
router.NoRoute(func(c *gin.Context) {
|
||
path := c.Request.URL.Path
|
||
|
||
if strings.HasPrefix(path, "/api") {
|
||
c.String(http.StatusNotFound, "Not Found")
|
||
return
|
||
}
|
||
|
||
// 静态资源
|
||
trimmed := strings.TrimPrefix(path, "/")
|
||
if strings.HasSuffix(trimmed, ".js") || strings.HasSuffix(trimmed, ".css") || strings.HasSuffix(trimmed, ".png") || strings.HasSuffix(trimmed, ".svg") || strings.HasSuffix(trimmed, ".ico") {
|
||
data, err := readStaticFile(trimmed)
|
||
if err == nil {
|
||
noCache(c)
|
||
c.Data(http.StatusOK, mimeByExt(trimmed), data)
|
||
return
|
||
}
|
||
}
|
||
|
||
// 兜底: SPA
|
||
noCache(c)
|
||
serveStaticFile(c, "index.html")
|
||
})
|
||
|
||
srv := &http.Server{
|
||
Addr: fmt.Sprintf("%s:%d", cfg.Server.Host, cfg.Server.Port),
|
||
Handler: router,
|
||
ReadTimeout: cfg.Server.ReadTimeout,
|
||
WriteTimeout: cfg.Server.WriteTimeout,
|
||
}
|
||
|
||
log.Printf("📋 首页: http://%s:%d/", cfg.Server.Host, cfg.Server.Port)
|
||
log.Printf("🔑 工作台: http://%s:%d/dashboard", cfg.Server.Host, cfg.Server.Port)
|
||
log.Printf("🎬 Demo: http://%s:%d/demo", cfg.Server.Host, cfg.Server.Port)
|
||
log.Printf("⚙️ 控制台: http://%s:%d/console", cfg.Server.Host, cfg.Server.Port)
|
||
|
||
return &Server{
|
||
cfg: cfg,
|
||
engine: router,
|
||
stats: statsEngine,
|
||
geoip: geoip,
|
||
srv: srv,
|
||
database: database,
|
||
keyLimit: keyLimit,
|
||
ipLimit: ipLimit,
|
||
}, nil
|
||
}
|
||
|
||
// Run 启动服务
|
||
func (s *Server) Run() error {
|
||
go func() {
|
||
ticker := time.NewTicker(5 * time.Minute)
|
||
defer ticker.Stop()
|
||
for range ticker.C {
|
||
s.stats.LogStats()
|
||
}
|
||
}()
|
||
|
||
// 定时发布到期的 API 产品
|
||
go func() {
|
||
pubTicker := time.NewTicker(30 * time.Second)
|
||
defer pubTicker.Stop()
|
||
for range pubTicker.C {
|
||
if n, err := s.database.ScheduledPublishProducts(); err == nil && n > 0 {
|
||
log.Printf("[publish] 自动发布了 %d 个 API 产品", n)
|
||
}
|
||
}
|
||
}()
|
||
|
||
go func() {
|
||
log.Printf("🚀 服务启动: http://%s:%d", s.cfg.Server.Host, s.cfg.Server.Port)
|
||
if err := s.srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||
log.Fatalf("服务启动失败: %v", err)
|
||
}
|
||
}()
|
||
|
||
quit := make(chan os.Signal, 1)
|
||
signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM)
|
||
<-quit
|
||
log.Println("正在关闭...")
|
||
|
||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||
defer cancel()
|
||
|
||
if err := s.srv.Shutdown(ctx); err != nil {
|
||
return fmt.Errorf("关闭失败: %w", err)
|
||
}
|
||
|
||
s.stats.Stop()
|
||
if s.geoip != nil {
|
||
s.geoip.Close()
|
||
}
|
||
if s.database != nil {
|
||
s.database.Close()
|
||
}
|
||
|
||
log.Println("服务已关闭")
|
||
return nil
|
||
}
|
||
|
||
// ── 中间件 ──────────────────────────────────────────────
|
||
|
||
func authRateLimitIP(ipLimit *ratelimit.Limiter) gin.HandlerFunc {
|
||
return func(c *gin.Context) {
|
||
ip := c.ClientIP()
|
||
if !ipLimit.Allow(ip) {
|
||
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{"code": 429, "msg": "请求过于频繁"})
|
||
return
|
||
}
|
||
c.Next()
|
||
}
|
||
}
|
||
|
||
// ── 静态文件 ──────────────────────────────────────────────
|
||
|
||
func readStaticFile(path string) ([]byte, error) {
|
||
f, err := web.StaticFS.Open(path)
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
defer f.Close()
|
||
return io.ReadAll(f)
|
||
}
|
||
|
||
func serveStaticFile(c *gin.Context, path string) {
|
||
data, err := readStaticFile(path)
|
||
if err != nil {
|
||
c.String(http.StatusNotFound, "Not Found")
|
||
return
|
||
}
|
||
c.Data(http.StatusOK, mimeByExt(path), data)
|
||
}
|
||
|
||
func mimeByExt(path string) string {
|
||
ext := strings.ToLower(filepath.Ext(path))
|
||
switch ext {
|
||
case ".html", ".htm":
|
||
return "text/html; charset=utf-8"
|
||
case ".css":
|
||
return "text/css; charset=utf-8"
|
||
case ".js":
|
||
return "application/javascript; charset=utf-8"
|
||
case ".json":
|
||
return "application/json; charset=utf-8"
|
||
case ".png":
|
||
return "image/png"
|
||
case ".jpg", ".jpeg":
|
||
return "image/jpeg"
|
||
case ".svg":
|
||
return "image/svg+xml"
|
||
case ".ico":
|
||
return "image/x-icon"
|
||
default:
|
||
return "application/octet-stream"
|
||
}
|
||
}
|