Files
rill/main.go
T
kevin 84768e94ed Swagger 文档改为按权限分组
- @Tags 由功能维度改为 public/user/admin 权限维度
- main.go 增加全局 tag 声明与权限说明(需放在 @securitydefinitions 之前,否则会被解析器吞掉)
- 重新生成 docs/,公开 4 个、需登录 7 个、管理员 11 个接口
2026-09-21 16:24:19 +08:00

183 lines
5.1 KiB
Go

package main
import (
"context"
"errors"
"flag"
"fmt"
"io/fs"
"log/slog"
"net"
"net/http"
"os"
"os/signal"
"path/filepath"
"strconv"
"strings"
"syscall"
"time"
"github.com/gin-contrib/cors"
"github.com/gin-gonic/gin"
"rill/docs"
"rill/internal/api"
"rill/internal/config"
"rill/internal/database"
)
//go:generate go tool swag init -g main.go -o docs --parseInternal
// @title Rill API
// @version 1.0
// @description Rill server HTTP API documentation. All endpoints are prefixed with api.prefix (default /api); requests and responses are JSON.
// @description Swagger UI: {prefix}/swagger/index.html; OpenAPI JSON: {prefix}/swagger/doc.json.
// @description Endpoints are grouped by required permission: public needs no authentication, user needs a Bearer JWT obtained from /auth/login, admin needs an admin user.
// @BasePath /api
// @tag.name public
// @tag.description No authentication required
// @tag.name user
// @tag.description Requires a logged-in user (Bearer JWT)
// @tag.name admin
// @tag.description Requires an admin user (Bearer JWT + admin group)
// @securityDefinitions.apikey BearerAuth
// @in header
// @name Authorization
// @description Bearer JWT, format: Bearer {token}, obtained from /auth/login
func main() {
configPath := flag.String("c", "data/config.yaml", "配置文件路径(不存在时自动生成)")
flag.Parse()
cfg, err := config.LoadConfig(*configPath)
if err != nil {
slog.Error("加载配置失败", "err", err)
os.Exit(1)
}
slog.SetDefault(slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{Level: cfg.LogLevel()})))
// 接口文档路径跟随 API 前缀
docs.SwaggerInfo.BasePath = cfg.API.Prefix
//设置gin运行模式
gin.SetMode(cfg.Server.Mode)
//启动gin服务
r := gin.New()
r.Use(gin.Recovery())
if cfg.Log.AccessLog {
r.Use(gin.Logger())
}
if cfg.API.CORS.Enabled {
r.Use(cors.New(cors.Config{
AllowOrigins: cfg.API.CORS.AllowOrigins,
AllowMethods: cfg.API.CORS.AllowMethods,
AllowHeaders: cfg.API.CORS.AllowHeaders,
AllowCredentials: cfg.API.CORS.AllowCredentials,
MaxAge: cfg.CORSMaxAge(),
}))
}
// 初始化数据库并执行迁移
db, err := database.Open(cfg)
if err != nil {
slog.Error("初始化数据库失败", "err", err)
os.Exit(1)
}
defer func() {
if err := database.Close(db); err != nil {
slog.Warn("关闭数据库失败", "err", err)
}
}()
migrateCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
err = database.Migrate(migrateCtx, db)
cancel()
if err != nil {
slog.Error("数据库迁移失败", "err", err)
os.Exit(1)
}
// API 路由
api.RegisterRoutes(r.Group(cfg.API.Prefix), db, cfg)
// 静态文件服务
fs := http.FileServer(http.Dir(cfg.Static.Dir))
indexFile := filepath.Join(cfg.Static.Dir, "index.html")
// 中间件处理路由
r.Use(func(c *gin.Context) {
if strings.HasPrefix(c.Request.URL.Path, cfg.API.Prefix) {
c.Next() // 继续处理API请求
return
}
// 处理静态文件;未命中的无扩展名路径回退到 index.html,支持前端 history 路由
cleanPath := filepath.Clean(c.Request.URL.Path)
if _, err := os.Stat(filepath.Join(cfg.Static.Dir, cleanPath)); err != nil && filepath.Ext(cleanPath) == "" {
c.File(indexFile)
c.Abort()
return
}
fs.ServeHTTP(c.Writer, c.Request)
c.Abort()
})
if err := serve(cfg, r); err != nil {
slog.Error("服务运行失败", "err", err)
os.Exit(1)
}
}
// serve 依据配置监听 TCP 与 unix socket(可同时启用),并在收到退出信号后优雅关闭。
func serve(cfg *config.Config, handler http.Handler) error {
srv := &http.Server{Handler: handler}
ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
errCh := make(chan error, 2)
listen := func(network, addr string) error {
ln, err := net.Listen(network, addr)
if err != nil {
return fmt.Errorf("监听 %s %s 失败: %w", network, addr, err)
}
slog.Info("服务监听中", "network", network, "addr", addr, "static", cfg.Static.Dir, "mode", cfg.Server.Mode)
go func() {
if err := srv.Serve(ln); err != nil && !errors.Is(err, http.ErrServerClosed) {
errCh <- err
}
}()
return nil
}
if cfg.Server.TCPEnabled() {
addr := net.JoinHostPort(cfg.Server.Host, strconv.Itoa(cfg.Server.Port))
if err := listen("tcp", addr); err != nil {
return err
}
}
if cfg.Server.SockEnabled() {
sock := strings.TrimSpace(cfg.Server.Sock)
if err := os.Remove(sock); err != nil && !errors.Is(err, fs.ErrNotExist) {
return fmt.Errorf("清理残留 socket %s 失败: %w", sock, err)
}
if err := listen("unix", sock); err != nil {
return err
}
defer os.Remove(sock)
}
select {
case err := <-errCh:
return err
case <-ctx.Done():
stop()
slog.Info("收到退出信号,正在关闭服务")
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := srv.Shutdown(shutdownCtx); err != nil {
return fmt.Errorf("关闭服务失败: %w", err)
}
}
return nil
}