Files
yuqianhe 3e27017e42 feat: add auth, db, gateway, email, ratelimit, geoblock modules and new handlers
- 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
2026-06-26 21:32:51 +09:00

369 lines
11 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"
"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"
}
}