FileCodeBox Go 重写版 v2.5.6(安全审计修复版)

Go 1.27.1 (Gin+GORM) + Vue 3 文件快传服务:

- 安全审计全部修复(docs/security-audit-2026-09-05.md):
  bcrypt 密码哈希与自动升级、presign 直传服务端大小/内容校验、
  全局请求体上限、依赖升级(govulncheck 0 命中)、janitor 后台清理、
  管理端审计动作落库、/admin CORS 收紧、通知内容白名单净化、
  会话默认 7 天、限流缓存故障降级、robots.txt 端点等
- 前端:取件链接复制修复(不再重复拼接提取码)、markdown 净化器加固
- Redis 支持库号(FCB_REDIS_DB / redis://…/db URL)
- 文档:docs/api/* 与 openapi.yaml 同步最新行为(robots.txt、
  提码 5 位起、chunk 32MiB 上限、admin 审计动作等)

验证:gofmt/go vet/go test 全绿;二进制端到端冒烟通过
This commit is contained in:
2026-09-05 04:22:41 +08:00
commit 9686fe887a
173 changed files with 32455 additions and 0 deletions
+261
View File
@@ -0,0 +1,261 @@
package middleware
import (
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/model"
"filecodebox/internal/response"
)
// auditHooks 审计钩子:由 API 层在响应前后填充与落库。
// 中间件负责计时与公共字段(IP/UA/设备/耗时),业务上下文通过 auditEntry 传递。
type auditEntry struct {
Entry audit.Entry
// start 请求进入审计中间件的时刻,用于计算耗时。
start time.Time
// writer 下载动作时包装的响应计数器。
writer *bytesCountWriter
// skip 为 true 表示业务 handler 显式跳过审计(AuditSkip)。
skip bool
// recorded 防止重复落库。
recorded bool
}
// bytesCountWriter 统计响应体写出字节数(用于下载审计)。
type bytesCountWriter struct {
gin.ResponseWriter
count int64
}
func (w *bytesCountWriter) Write(b []byte) (int, error) {
n, err := w.ResponseWriter.Write(b)
w.count += int64(n)
return n, err
}
func (w *bytesCountWriter) WriteString(s string) (int, error) {
n, err := w.ResponseWriter.WriteString(s)
w.count += int64(n)
return n, err
}
// Classifier 判定请求是否属于需审计的动作;返回动作名与是否命中。
type Classifier func(c *gin.Context) (action string, ok bool)
// DefaultClassifier 按任务合同的默认路由语义分类:
// - 上传:POST /share/file、/share/text、/chunk/upload*、/presign*
// - 下载:GET /share/download、/share/select、/share/metadata
// - 管理(L5):POST/PATCH/DELETE 的敏感管理操作——登录/登出、配置与密码
// 修改、存储引擎切换、文件更新/删除/策略动作
//
// API 层可传入自定义分类器覆盖。
func DefaultClassifier(c *gin.Context) (string, bool) {
path := c.FullPath()
if path == "" {
path = c.Request.URL.Path
}
p := strings.TrimRight(path, "/")
switch c.Request.Method {
case http.MethodPost, http.MethodPut:
switch {
case p == "/share/file" || p == "/share/text":
return audit.ActionUpload, true
case strings.HasPrefix(p, "/chunk/upload"):
return audit.ActionUpload, true
case strings.HasPrefix(p, "/presign"):
return audit.ActionUpload, true
}
if adminAuditActions[p] {
return audit.ActionAdmin, true
}
case http.MethodPatch, http.MethodDelete:
if adminAuditActions[p] {
return audit.ActionAdmin, true
}
case http.MethodGet:
switch p {
case "/share/download", "/share/select", "/share/metadata":
return audit.ActionDownload, true
}
}
return "", false
}
// adminAuditActions 需要审计的管理端敏感操作路由(L5)。
var adminAuditActions = map[string]bool{
"/admin/login": true,
"/admin/logout": true,
"/admin/config/update": true,
"/admin/settings/password": true,
"/admin/storage/switch": true,
"/admin/file/update": true,
"/admin/file/delete": true,
"/admin/file/batch-delete": true,
"/admin/file/batch-update": true,
"/admin/file/policy-action": true,
"/admin/file/batch-policy-action": true,
}
// Audit 审计中间件:对分类器命中的 upload/download/admin 动作写审计日志。
// handler 通过 AuditSet 填充取件码/文件名/字节数等业务字段;
// handler 未显式 AuditRecordRequest 时按 HTTP 状态兜底落库。
func Audit(service *audit.Service, classify Classifier) gin.HandlerFunc {
if classify == nil {
classify = DefaultClassifier
}
return func(c *gin.Context) {
start := time.Now()
action, ok := classify(c)
// 未命中审计动作的请求直接放行,不产生审计记录。
// L5admin 类动作同样需要建 auditEntry 并落库(登录失败/配置变更等)。
if !ok {
c.Next()
return
}
entry := audit.Entry{
Action: action,
IP: GetClientIP(c),
UserAgent: c.Request.UserAgent(),
}
info := audit.ParseUserAgent(entry.UserAgent)
entry.DeviceOS = info.OS
entry.DeviceBrowser = info.Browser
entry.DeviceType = info.Type
// 交给后续 handler 填充
state := &auditEntry{Entry: entry, start: start}
c.Set("auditEntry", state)
// 下载动作:包装 Writer 以捕获实际写出字节数(必须在 c.Next() 前替换)
if action == audit.ActionDownload {
state.writer = &bytesCountWriter{ResponseWriter: c.Writer}
c.Writer = state.writer
}
c.Next()
// 下载兜底统计:handler 未填 TransferredBytes 时取响应写出字节
if action == audit.ActionDownload && state.Entry.TransferredBytes == 0 &&
!state.recorded && !state.skip && state.writer != nil {
state.Entry.TransferredBytes = state.writer.count
}
// handler 未显式落库时兜底记录
ae, exists := c.Get("auditEntry")
if !exists {
return
}
state, isState := ae.(*auditEntry)
if !isState || state.recorded || state.skip {
return
}
state.Entry.Duration = time.Since(start)
state.Entry.Actor = resolveActor(c)
status := c.Writer.Status()
switch {
case state.Entry.Result != "":
// handler 已给出结论
case status >= 500:
state.Entry.Result = model.AuditResultFailed
case status == 401 || status == 403 || status == 423 || status == 429 || status == 428:
state.Entry.Result = model.AuditResultDenied
case status >= 400:
state.Entry.Result = model.AuditResultFailed
default:
state.Entry.Result = model.AuditResultSuccess
}
switch {
case state.Entry.ErrorMsg != "":
// handler 已给出错误信息
case c.Errors.String() != "":
state.Entry.ErrorMsg = c.Errors.String()
case status >= 400:
// 兜底:记录 HTTP 状态
state.Entry.ErrorMsg = "HTTP " + itoa64(int64(status))
}
service.Record(state.Entry)
state.recorded = true
}
}
// AuditEntry 获取当前请求的审计状态(由 Audit 中间件创建)。
func AuditEntry(c *gin.Context) *auditEntry {
if v, ok := c.Get("auditEntry"); ok {
if ae, ok := v.(*auditEntry); ok {
return ae
}
}
return nil
}
// AuditSet 填充当前请求的审计字段;仅对已启用审计的请求生效。
func AuditSet(c *gin.Context, fn func(e *audit.Entry)) {
if ae := AuditEntry(c); ae != nil && fn != nil {
fn(&ae.Entry)
}
}
// AuditRecordRequest 显式触发落库(含耗时);由 handler 在响应前调用。
func AuditRecordRequest(c *gin.Context, service *audit.Service, result, errMsg string) {
ae := AuditEntry(c)
if ae == nil || ae.recorded || ae.skip {
return
}
ae.Entry.Duration = time.Since(ae.start)
ae.Entry.Result = result
ae.Entry.ErrorMsg = errMsg
ae.Entry.Actor = resolveActor(c)
service.Record(ae.Entry)
ae.recorded = true
}
// AuditSkip 标记当前请求不写审计。
func AuditSkip(c *gin.Context) {
if ae := AuditEntry(c); ae != nil {
ae.skip = true
}
}
// resolveActor 判断请求者角色:管理员 JWT 有效 → admin,否则 guest。
func resolveActor(c *gin.Context) string {
header := c.GetHeader("Authorization")
if len(header) > 7 && header[:7] == "Bearer " {
// 仅检查声明是否有效,不重复校验签名逻辑(AdminAuth 已处理受保护路由)
if _, ok := c.Get("claims"); ok {
return audit.ActorAdmin
}
}
return audit.ActorGuest
}
// AuditRecord 显式按结果落库;duration 由中间件按起始时间计算。
func AuditRecord(c *gin.Context, service *audit.Service, result, errMsg string) {
ae := AuditEntry(c)
if ae == nil || ae.recorded || ae.skip {
return
}
AuditRecordRequest(c, service, result, errMsg)
}
// GuardNotInitialized 系统未初始化守卫:除 setup/health 外返回 428。
func GuardNotInitialized(isInit func() bool) gin.HandlerFunc {
return func(c *gin.Context) {
if isInit() {
c.Next()
return
}
path := c.Request.URL.Path
if path == "/setup" || path == "/api/v1/health" {
c.Next()
return
}
response.Fail(c, 428, "系统未初始化,请先完成初始化")
}
}
@@ -0,0 +1,89 @@
// audit_l5_test.go — L5 回归:admin 类动作(如登录失败)必须落审计。
package middleware
import (
"context"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/model"
)
type captureSink struct {
logs []model.AuditLog
}
func (s *captureSink) Save(_ context.Context, logs []model.AuditLog) error {
s.logs = append(s.logs, logs...)
return nil
}
// TestAuditRecordsAdminActions L5/admin/login 失败(401)后应产生一条
// result=denied 的 admin 审计记录(此前 skip 条件把 admin 动作整体跳过)。
func TestAuditRecordsAdminActions(t *testing.T) {
gin.SetMode(gin.TestMode)
sink := &captureSink{}
svc := audit.NewService(sink)
r := gin.New()
r.Use(Audit(svc, nil)) // DefaultClassifier
r.POST("/admin/login", func(c *gin.Context) {
c.JSON(http.StatusUnauthorized, gin.H{"code": 401})
})
r.POST("/share/text", func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"code": 200})
})
r.GET("/healthz", func(c *gin.Context) {
c.Status(http.StatusOK) // 未分类动作:不应产生审计
})
// 管理端:401 → denied
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("POST", "/admin/login", nil))
if w.Code != http.StatusUnauthorized {
t.Fatalf("login should 401, got %d", w.Code)
}
// 上传类:200 → success
w2 := httptest.NewRecorder()
r.ServeHTTP(w2, httptest.NewRequest("POST", "/share/text", nil))
// 未分类:不落库
w3 := httptest.NewRecorder()
r.ServeHTTP(w3, httptest.NewRequest("GET", "/healthz", nil))
// audit.Service 异步落库,轮询等待
var actions []string
for i := 0; i < 50; i++ {
if len(sink.logs) >= 2 {
break
}
waitMillis(20)
}
if len(sink.logs) != 2 {
t.Fatalf("应恰好 2 条审计记录, got %d", len(sink.logs))
}
for _, l := range sink.logs {
actions = append(actions, l.Action)
switch l.Action {
case audit.ActionAdmin:
if l.Result != model.AuditResultDenied {
t.Fatalf("admin 401 应记 denied, got %q", l.Result)
}
case audit.ActionUpload:
if l.Result != model.AuditResultSuccess {
t.Fatalf("upload 200 应记 success, got %q", l.Result)
}
default:
t.Fatalf("意外动作 %q", l.Action)
}
}
_ = actions
}
func waitMillis(ms int) {
time.Sleep(time.Duration(ms) * time.Millisecond)
}
+198
View File
@@ -0,0 +1,198 @@
package middleware
import (
"context"
"net/http/httptest"
"sync"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/audit"
"filecodebox/internal/model"
)
// memSink 测试用内存落库实现。
type memSink struct {
mu sync.Mutex
logs []model.AuditLog
notif chan struct{}
}
func newMemSink() *memSink { return &memSink{notif: make(chan struct{}, 16)} }
func (m *memSink) Save(_ context.Context, logs []model.AuditLog) error {
m.mu.Lock()
m.logs = append(m.logs, logs...)
m.mu.Unlock()
m.notif <- struct{}{}
return nil
}
func (m *memSink) snapshot() []model.AuditLog {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]model.AuditLog, len(m.logs))
copy(out, m.logs)
return out
}
// waitFor 等待 sink 收到 n 条记录(带超时)。
func (m *memSink) waitFor(t *testing.T, n int) []model.AuditLog {
t.Helper()
deadline := time.After(2 * time.Second)
for {
logs := m.snapshot()
if len(logs) >= n {
return logs
}
select {
case <-m.notif:
case <-deadline:
t.Fatalf("等待审计记录超时: 已收到 %d 条", len(logs))
}
}
}
func auditRouter(svc *audit.Service) *gin.Engine {
r := gin.New()
r.Use(ClientIP(nil))
r.Use(Audit(svc, nil)) // 默认分类器
// 上传路由:命中默认分类器(POST /share/file
r.POST("/share/file", func(c *gin.Context) {
// 模拟 handler 填充业务字段并显式落库
AuditSet(c, func(e *audit.Entry) {
e.FileCode = "Ab3xY"
e.FileName = "hello.zip"
e.SizeBytes = 1024
})
AuditRecordRequest(c, svc, model.AuditResultSuccess, "")
c.JSON(200, gin.H{"ok": true})
})
// 下载路由:命中默认分类器(GET /share/download
r.GET("/share/download", func(c *gin.Context) {
AuditSet(c, func(e *audit.Entry) { e.FileCode = "Xy12Z" })
c.JSON(404, gin.H{"msg": "文件已过期删除"}) // 未显式落库 → 状态码兜底
})
// 普通路由:不命中,不应产生审计
r.GET("/plain", func(c *gin.Context) { c.JSON(200, gin.H{"ok": true}) })
return r
}
func TestAuditMiddlewareRecordsUpload(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := auditRouter(svc)
w := httptest.NewRecorder()
req := httptest.NewRequest("POST", "/share/file", nil)
req.Header.Set("User-Agent", "Mozilla/5.0 (Windows NT 10.0; Win64; x64) Chrome/120.0.0.0 Safari/537.36")
r.ServeHTTP(w, req)
if w.Code != 200 {
t.Fatalf("上传应成功: %d", w.Code)
}
logs := sink.waitFor(t, 1)
e := logs[0]
if e.Action != audit.ActionUpload {
t.Errorf("action = %s", e.Action)
}
if e.FileCode != "Ab3xY" || e.FileName != "hello.zip" {
t.Errorf("file fields = %s/%s", e.FileCode, e.FileName)
}
if e.SizeBytes != 1024 {
t.Errorf("size = %d", e.SizeBytes)
}
if e.Result != model.AuditResultSuccess {
t.Errorf("result = %s", e.Result)
}
if e.DeviceOS != "Windows" || e.DeviceBrowser != "Chrome" || e.DeviceType != "desktop" {
t.Errorf("device = %s/%s/%s", e.DeviceOS, e.DeviceBrowser, e.DeviceType)
}
if e.DurationMs < 0 {
t.Errorf("duration = %d", e.DurationMs)
}
if e.Actor != audit.ActorGuest {
t.Errorf("actor = %s", e.Actor)
}
}
func TestAuditMiddlewareSkipsPlainRoutes(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := auditRouter(svc)
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/plain", nil))
if w.Code != 200 {
t.Fatalf("plain 路由应成功: %d", w.Code)
}
time.Sleep(150 * time.Millisecond)
if logs := sink.snapshot(); len(logs) != 0 {
t.Fatalf("普通路由不应产生审计记录: %v", logs)
}
}
func TestAuditFailedDownloadFallback(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := auditRouter(svc)
w := httptest.NewRecorder()
req := httptest.NewRequest("GET", "/share/download?code=xyz", nil)
req.Header.Set("User-Agent", "curl/8.4.0")
r.ServeHTTP(w, req)
logs := sink.waitFor(t, 1)
e := logs[0]
if e.Action != audit.ActionDownload {
t.Errorf("action = %s", e.Action)
}
if e.Result != model.AuditResultFailed {
t.Errorf("4xx 兜底 result = %s", e.Result)
}
if e.DeviceType != "bot" {
t.Errorf("curl 应识别为 bot: %s", e.DeviceType)
}
if e.ErrorMsg == "" {
t.Error("失败记录应包含错误信息")
}
}
func TestAuditDeniedStatusMapping(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := gin.New()
r.Use(ClientIP(nil))
r.Use(Audit(svc, nil))
r.GET("/share/select", func(c *gin.Context) { c.AbortWithStatus(429) })
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/share/select?code=abc", nil))
logs := sink.waitFor(t, 1)
if logs[0].Result != model.AuditResultDenied {
t.Errorf("429 应映射为 denied: %s", logs[0].Result)
}
}
func TestAuditDownloadBytesCounted(t *testing.T) {
sink := newMemSink()
svc := audit.NewService(sink)
r := gin.New()
r.Use(ClientIP(nil))
r.Use(Audit(svc, nil))
r.GET("/share/download", func(c *gin.Context) {
payload := []byte("0123456789abcdef") // 16 字节
c.Data(200, "application/octet-stream", payload)
// 未显式落库 → 中间件兜底;TransferredBytes 应等于写出字节
})
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/share/download?code=bytes", nil))
logs := sink.waitFor(t, 1)
if logs[0].TransferredBytes != 16 {
t.Errorf("下载字节数 = %d, want 16", logs[0].TransferredBytes)
}
if logs[0].Result != model.AuditResultSuccess {
t.Errorf("result = %s", logs[0].Result)
}
}
+30
View File
@@ -0,0 +1,30 @@
// Package middleware — bodylimit.go:全局请求体大小限制。
//
// 安全审计 M3:此前服务对请求体无任何上限,text/plain 形态的绑定会先
// io.ReadAll 全量进内存、multipart 大文件先完整落盘后才被策略拒绝,
// 未认证攻击者可借大 body 消耗内存/磁盘/带宽。
// 本中间件用 http.MaxBytesReader 按「路径类别」给出上限:
// - 管理端(/admin/*):1MiB
// - 文本/元信息/取件类(/share/text|metadata|select):1MiB
// - 其余(含上传):maxFileSize0=回落 uploadSize,仍为 0 时 64MiB 兜底)+ 2MiB 表单开销。
//
// 超限时后续读取返回错误,统一被 handler 的 bind 错误路径映射为 400。
package middleware
import (
"net/http"
"github.com/gin-gonic/gin"
)
// BodyLimit 按请求路径动态限制请求体大小(limit<=0 表示不限制)。
func BodyLimit(limitFn func(c *gin.Context) int64) gin.HandlerFunc {
return func(c *gin.Context) {
if c.Request.Body != nil && limitFn != nil {
if limit := limitFn(c); limit > 0 {
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, limit)
}
}
c.Next()
}
}
+74
View File
@@ -0,0 +1,74 @@
package middleware
import (
"net/url"
"strings"
"github.com/gin-gonic/gin"
)
// Cors 跨域中间件(L6 收紧):
// - 公开接口:维持 allow_origins=*Bearer Token 认证,无 Cookie CSRF 面);
// - 管理端(/admin/*):当请求携带 Origin 且既不同源也不在允许域名列表时,
// 不回 CORS 头(浏览器将拦截跨域读取)。防止管理端 token 泄露后
// 被任意第三方页面直接跨域调用。无 Origin 的非浏览器请求不受影响。
//
// extraAllowedOrigins:管理端额外允许的来源(如 site_domain 配置的对外域名)。
func Cors(extraAllowedOrigins ...string) gin.HandlerFunc {
allowedHosts := map[string]bool{}
for _, o := range extraAllowedOrigins {
if o == "" {
continue
}
raw := strings.TrimSpace(o)
if !strings.Contains(raw, "://") {
raw = "https://" + raw
}
if u, err := url.Parse(raw); err == nil && u.Host != "" {
allowedHosts[u.Host] = true
}
}
// adminCrossOriginBlocked 判断 /admin 请求是否应拒绝跨域:
// 仅在「带 Origin 且 Origin 既不同源也不在白名单」时为 true。
adminBlocked := func(c *gin.Context) bool {
p := c.Request.URL.Path
if p != "/admin" && !strings.HasPrefix(p, "/admin/") {
return false
}
origin := c.GetHeader("Origin")
if origin == "" {
return false
}
o, err := url.Parse(origin)
if err != nil || o.Host == "" {
return true // Origin 非法:按跨域拒绝处理
}
if o.Host == c.Request.Host || allowedHosts[o.Host] {
return false
}
return true
}
return func(c *gin.Context) {
if adminBlocked(c) {
// 不回 ACAO;预检直接 204(浏览器会因无 CORS 头拦截后续请求)
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(204)
return
}
c.Next()
return
}
c.Header("Access-Control-Allow-Origin", "*")
c.Header("Access-Control-Allow-Methods", "GET, POST, PUT, DELETE, OPTIONS, HEAD")
c.Header("Access-Control-Allow-Headers", "Authorization, Content-Type, Content-Disposition, X-Requested-With")
c.Header("Access-Control-Expose-Headers", "Content-Disposition, Content-Length")
c.Header("Access-Control-Max-Age", "86400")
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(204)
return
}
c.Next()
}
}
+99
View File
@@ -0,0 +1,99 @@
package middleware
import (
"errors"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/golang-jwt/jwt/v5"
"filecodebox/internal/response"
)
// jwtClaims 自定义声明:对齐参考实现(payload 含 is_admin 与 exp)。
type jwtClaims struct {
IsAdmin bool `json:"is_admin"`
jwt.RegisteredClaims
}
// 签发/校验相关错误。
var (
ErrTokenExpired = errors.New("token已过期")
ErrTokenInvalid = errors.New("无效的签名")
ErrNotAdmin = errors.New("未授权或授权校验失败")
)
// SignAdminToken 用 HS256 签发管理员 JWT。
// secret 为数据库 settings 中的 jwt_secretexpires 为会话有效期。
func SignAdminToken(secret string, expires time.Duration) (string, time.Time, error) {
if strings.TrimSpace(secret) == "" {
return "", time.Time{}, errors.New("JWT签名密钥未初始化")
}
expiresAt := time.Now().Add(expires)
claims := jwtClaims{
IsAdmin: true,
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(expiresAt),
IssuedAt: jwt.NewNumericDate(time.Now()),
Issuer: "filecodebox",
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
signed, err := token.SignedString([]byte(secret))
return signed, expiresAt, err
}
// VerifyAdminToken 校验管理员 JWT:签名、过期时间与 is_admin 声明。
func VerifyAdminToken(secret, token string) (*jwtClaims, error) {
if strings.TrimSpace(secret) == "" {
return nil, errors.New("JWT签名密钥未初始化")
}
parsed, err := jwt.ParseWithClaims(token, &jwtClaims{}, func(t *jwt.Token) (any, error) {
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
return nil, ErrTokenInvalid
}
return []byte(secret), nil
}, jwt.WithValidMethods([]string{"HS256"}))
if err != nil {
if errors.Is(err, jwt.ErrTokenExpired) {
return nil, ErrTokenExpired
}
return nil, ErrTokenInvalid
}
claims, ok := parsed.Claims.(*jwtClaims)
if !ok || !parsed.Valid {
return nil, ErrTokenInvalid
}
if !claims.IsAdmin {
return nil, ErrNotAdmin
}
return claims, nil
}
// SecretProvider 动态提供当前 jwt_secretsettings KV 运行时可变)。
type SecretProvider func() string
// AdminAuth 管理员鉴权中间件:校验 Authorization: Bearer <token>。
// 成功后把声明写入 gin 上下文(ctxClaims)。
func AdminAuth(secret SecretProvider) gin.HandlerFunc {
return func(c *gin.Context) {
header := c.GetHeader("Authorization")
if !strings.HasPrefix(header, "Bearer ") {
response.Fail(c, 401, "未授权或授权校验失败")
return
}
token := strings.TrimSpace(strings.TrimPrefix(header, "Bearer "))
if token == "" {
response.Fail(c, 401, "未授权或授权校验失败")
return
}
claims, err := VerifyAdminToken(secret(), token)
if err != nil {
response.Fail(c, 401, err.Error())
return
}
c.Set("claims", claims)
c.Next()
}
}
+83
View File
@@ -0,0 +1,83 @@
package middleware
import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/response"
)
const testSecret = "unit-test-secret-0123456789abcdef"
func init() { gin.SetMode(gin.TestMode) }
func TestSignAndVerifyAdminToken(t *testing.T) {
token, expiresAt, err := SignAdminToken(testSecret, time.Hour)
if err != nil {
t.Fatalf("签发失败: %v", err)
}
if expiresAt.Before(time.Now()) {
t.Fatal("过期时间不合理")
}
claims, err := VerifyAdminToken(testSecret, token)
if err != nil {
t.Fatalf("校验失败: %v", err)
}
if !claims.IsAdmin {
t.Fatal("is_admin 应为 true")
}
}
func TestVerifyTamperedToken(t *testing.T) {
token, _, _ := SignAdminToken(testSecret, time.Hour)
claims, err := VerifyAdminToken(testSecret+"-wrong", token)
if err == nil || claims != nil {
t.Fatal("密钥不匹配应校验失败")
}
// 篡改 payload
tampered := token[:len(token)-3] + "abc"
if _, err := VerifyAdminToken(testSecret, tampered); err == nil {
t.Fatal("篡改的 token 应校验失败")
}
// 非 HMAC 算法拒绝
algNone := "eyJhbGciOiJub25lIiwidHlwIjoiSldUIn0.eyJpc19hZG1pbiI6dHJ1ZX0."
if _, err := VerifyAdminToken(testSecret, algNone); err == nil {
t.Fatal("none 算法应被拒绝")
}
}
func TestAdminAuthMiddleware(t *testing.T) {
r := gin.New()
r.GET("/protected", AdminAuth(func() string { return testSecret }), func(c *gin.Context) {
response.OK(c, gin.H{"ok": true})
})
// 无 token → 401
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/protected", nil))
if w.Code != http.StatusUnauthorized {
t.Fatalf("无 token 应 401: %d", w.Code)
}
// 有效 token → 200
token, _, _ := SignAdminToken(testSecret, time.Hour)
w = httptest.NewRecorder()
req := httptest.NewRequest("GET", "/protected", nil)
req.Header.Set("Authorization", "Bearer "+token)
r.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("有效 token 应 200: %d %s", w.Code, w.Body.String())
}
// 过期 token → 401
expired, _, _ := SignAdminToken(testSecret, -time.Minute)
w = httptest.NewRecorder()
req = httptest.NewRequest("GET", "/protected", nil)
req.Header.Set("Authorization", "Bearer "+expired)
r.ServeHTTP(w, req)
if w.Code != http.StatusUnauthorized {
t.Fatalf("过期 token 应 401: %d", w.Code)
}
}
+295
View File
@@ -0,0 +1,295 @@
package middleware
import (
"context"
"errors"
"net"
"net/http"
"strings"
"sync"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/cache"
"filecodebox/internal/response"
)
// 限流类别(对齐参考 apps/base/utils.py 的 ip_limit)。
const (
LimitError = "error" // 取件错误(密码错误、取件失败)
LimitUpload = "upload" // 上传次数
LimitLogin = "login" // 管理员登录失败
LimitMeta = "metadata" // 分享元信息查询
)
// LimitRule 限流规则:window 内最多 count 次。
type LimitRule struct {
Count int // 允许次数
Window time.Duration // 时间窗口
}
// clientIP 解析客户端真实 IP:仅当直连地址属于可信代理时才采信 X-Forwarded-For / X-Real-IP。
// 语义对齐参考 apps/base/dependencies.py 的 get_client_ip。
func clientIP(c *gin.Context, trustedProxies []*net.IPNet) string {
remote := net.ParseIP(c.RemoteIP())
parse := func(s string) net.IP {
ip := net.ParseIP(strings.TrimSpace(s))
return ip
}
isTrusted := func(ip net.IP) bool {
if ip == nil {
return false
}
for _, n := range trustedProxies {
if n.Contains(ip) {
return true
}
}
return false
}
if !isTrusted(remote) {
return remote.String()
}
// X-Forwarded-For:从右往左找第一个非可信代理地址
if xff := c.GetHeader("X-Forwarded-For"); xff != "" {
parts := strings.Split(xff, ",")
for i := len(parts) - 1; i >= 0; i-- {
candidate := parse(parts[i])
if candidate == nil {
return remote.String()
}
if !isTrusted(candidate) {
return candidate.String()
}
}
return strings.TrimSpace(parts[0])
}
if xr := c.GetHeader("X-Real-IP"); xr != "" {
if ip := parse(xr); ip != nil {
return ip.String()
}
}
return remote.String()
}
// ParseTrustedProxies 把 CIDR/单 IP 字符串解析为网络列表。
func ParseTrustedProxies(items []string) []*net.IPNet {
var out []*net.IPNet
for _, item := range items {
item = strings.TrimSpace(item)
if item == "" {
continue
}
if !strings.Contains(item, "/") {
item += "/32"
if strings.Contains(item, ":") { // IPv6
item = item[:len(item)-3] + "/128"
}
}
_, network, err := net.ParseCIDR(item)
if err != nil {
continue
}
out = append(out, network)
}
return out
}
// ClientIP 中间件:解析真实 IP 并写入上下文(ctxClientIP)。
func ClientIP(trustedProxies []*net.IPNet) gin.HandlerFunc {
return func(c *gin.Context) {
c.Set("ctxClientIP", clientIP(c, trustedProxies))
c.Next()
}
}
// GetClientIP 从 gin 上下文取解析后的客户端 IP。
func GetClientIP(c *gin.Context) string {
if v, ok := c.Get("ctxClientIP"); ok {
if s, ok := v.(string); ok {
return s
}
}
if ip := net.ParseIP(c.RemoteIP()); ip != nil {
return ip.String()
}
return c.RemoteIP()
}
// RateLimiter 基于 cache.Cache 的固定窗口 IP 限流器。
// 计数语义对齐参考实现:check 通过时放行,业务方在发生"计数事件"(如失败/成功上传)后调用 Add。
//
// L9:缓存故障降级——此前 cache.Get/Incr 失败(如 Redis 宕机)时一律放行,
// 登录爆破防护随之失效。现降级为进程内固定窗口计数(单实例语义),
// 缓存恢复后自动回到共享缓存计数。降级期间计数独立于缓存,不叠加。
type RateLimiter struct {
cache cache.Cache
limits map[string]LimitRule
prefix string
fbMu sync.Mutex
fallback map[string]*fallbackEntry // 进程内降级计数
lastPrune time.Time
}
type fallbackEntry struct {
count int64
expires time.Time
}
// fallbackMaxEntries 降级计数表上限(超出即整体重置,防内存增长)。
const fallbackMaxEntries = 8192
// NewRateLimiter 构造限流器;limits 为各类别规则(来自 settings 的 errorCount/errorMinute 等)。
func NewRateLimiter(cache cache.Cache, limits map[string]LimitRule) *RateLimiter {
if limits == nil {
limits = map[string]LimitRule{}
}
return &RateLimiter{cache: cache, limits: limits, prefix: "fcb:rl", fallback: map[string]*fallbackEntry{}}
}
// SetRule 运行时更新规则(settings KV 变更后调用)。
func (r *RateLimiter) SetRule(kind string, rule LimitRule) {
r.limits[kind] = rule
}
func (r *RateLimiter) windowKey(kind, ip string, now time.Time) string {
// 固定窗口:按窗口起点分桶
bucket := now.Unix() / int64(r.limits[kind].Window/time.Second)
return r.prefix + ":" + kind + ":" + ip + ":" + itoa64(bucket)
}
// Check 只读检查该 IP 在当前窗口内是否仍被允许(不计数)。
// 对齐参考 check_ip:已用次数 >= 上限即拒绝。
// 缓存键不存在(ErrNotFound)视为 0 次;缓存故障时降级为进程内计数。
func (r *RateLimiter) Check(c *gin.Context, kind string) (bool, int64) {
rule, ok := r.limits[kind]
if !ok || rule.Count <= 0 || rule.Window <= 0 {
return true, 0
}
ip := GetClientIP(c)
ctx, cancel := context.WithTimeout(c.Request.Context(), 2*time.Second)
defer cancel()
now := time.Now()
raw, err := r.cache.Get(ctx, r.windowKey(kind, ip, now))
if err != nil {
if errors.Is(err, cache.ErrNotFound) {
return true, 0 // 键不存在:窗口内尚无计数
}
// 缓存故障:降级进程内计数判定
return r.fallbackCount(kind, ip, rule, now) < int64(rule.Count), 0
}
n := parseInt64(raw)
return n < int64(rule.Count), n
}
// Add 记录一次计数事件(对齐参考 add_ip:调用即计数,如上传成功/登录失败/取件错误)。
// 缓存故障时降级为进程内计数。
func (r *RateLimiter) Add(c *gin.Context, kind string) {
rule, ok := r.limits[kind]
if !ok || rule.Count <= 0 || rule.Window <= 0 {
return
}
ip := GetClientIP(c)
ctx, cancel := context.WithTimeout(c.Request.Context(), 2*time.Second)
defer cancel()
if _, err := r.cache.Incr(ctx, r.windowKey(kind, ip, time.Now()), rule.Window); err != nil &&
!errors.Is(err, cache.ErrNotFound) {
// Incr 正常情况下不会因键不存在失败(缺键即从 0 起);
// 其余错误视为缓存故障 → 进程内计数
r.fallbackIncr(kind, ip, rule, time.Now())
}
}
// —— 进程内降级计数(L9)——
func (r *RateLimiter) fallbackIncr(kind, ip string, rule LimitRule, now time.Time) {
key := r.windowKey(kind, ip, now)
expires := now.Add(rule.Window)
r.fbMu.Lock()
defer r.fbMu.Unlock()
r.pruneFallbackLocked(now)
if len(r.fallback) >= fallbackMaxEntries {
r.fallback = map[string]*fallbackEntry{} // 极端情况整体重置,防内存无限增长
}
e, ok := r.fallback[key]
if !ok || now.After(e.expires) {
r.fallback[key] = &fallbackEntry{count: 1, expires: expires}
return
}
e.count++
}
func (r *RateLimiter) fallbackCount(kind, ip string, rule LimitRule, now time.Time) int64 {
key := r.windowKey(kind, ip, now)
r.fbMu.Lock()
defer r.fbMu.Unlock()
e, ok := r.fallback[key]
if !ok || now.After(e.expires) {
return 0
}
return e.count
}
// pruneFallbackLocked 清理已过窗口的降级计数(低频触发:每 1000 条或 10 分钟一次)。
func (r *RateLimiter) pruneFallbackLocked(now time.Time) {
if r.lastPrune.IsZero() || len(r.fallback) >= 1024 || now.Sub(r.lastPrune) >= 10*time.Minute {
for k, e := range r.fallback {
if now.After(e.expires) {
delete(r.fallback, k)
}
}
r.lastPrune = now
}
}
// RequireRateLimit 中间件:请求进入即检查,请求完成即计数。
// 适用于"每次访问都计数"的类别(如 metadata 查询);
// 上传/登录等"仅成功/失败才计数"的场景由 handler 显式调用 Check/Add。
func (r *RateLimiter) RequireRateLimit(kind string) gin.HandlerFunc {
return func(c *gin.Context) {
allowed, _ := r.Check(c, kind)
if !allowed {
response.Fail(c, http.StatusLocked, "请求次数过多,请稍后再试")
return
}
c.Next()
r.Add(c, kind)
}
}
// parseInt64 解析十进制整数字符串,非法输入返回 0。
func parseInt64(s string) int64 {
var n int64
for _, ch := range s {
if ch < '0' || ch > '9' {
return 0
}
n = n*10 + int64(ch-'0')
}
return n
}
func itoa64(n int64) string {
if n == 0 {
return "0"
}
neg := n < 0
if neg {
n = -n
}
var buf [21]byte
i := len(buf)
for n > 0 {
i--
buf[i] = byte('0' + n%10)
n /= 10
}
if neg {
i--
buf[i] = '-'
}
return string(buf[i:])
}
@@ -0,0 +1,66 @@
package middleware
import (
"context"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/cache"
)
// failingCache 模拟缓存故障(Get/Incr 均返回非 ErrNotFound 错误)。
type failingCache struct{ cache.Cache }
func (f *failingCache) Get(_ context.Context, _ string) (string, error) {
return "", context.DeadlineExceeded
}
func (f *failingCache) Set(_ context.Context, _, _ string, _ time.Duration) error {
return context.DeadlineExceeded
}
func (f *failingCache) Incr(_ context.Context, _ string, _ time.Duration) (int64, error) {
return 0, context.DeadlineExceeded
}
func newLimiterTest(c *gin.Context, cacheImpl cache.Cache, count int) *RateLimiter {
return NewRateLimiter(cacheImpl, map[string]LimitRule{
LimitLogin: {Count: count, Window: time.Minute},
})
}
func ginTestContext() *gin.Context {
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Request = httptest.NewRequest("POST", "/admin/login", nil)
return c
}
// TestRateLimiterFallbackOnCacheFailure L9:缓存故障时限流降级为进程内计数,
// 超过上限后 Check 拒绝(此前 fail-open 会一直放行)。
func TestRateLimiterFallbackOnCacheFailure(t *testing.T) {
c := ginTestContext()
rl := newLimiterTest(c, &failingCache{cache.NewMemory()}, 3)
for i := 0; i < 3; i++ {
if ok, _ := rl.Check(c, LimitLogin); !ok {
t.Fatalf("第 %d 次检查不应拒绝", i+1)
}
rl.Add(c, LimitLogin)
}
if ok, _ := rl.Check(c, LimitLogin); ok {
t.Fatal("缓存故障降级下,超过上限后 Check 应拒绝(fail-close")
}
}
// TestRateLimiterNormalCacheCounting 正常缓存路径行为不变。
func TestRateLimiterNormalCacheCounting(t *testing.T) {
c := ginTestContext()
rl := newLimiterTest(c, cache.NewMemory(), 2)
for i := 0; i < 2; i++ {
rl.Add(c, LimitLogin)
}
if ok, _ := rl.Check(c, LimitLogin); ok {
t.Fatal("达到上限后 Check 应拒绝")
}
}
@@ -0,0 +1,126 @@
package middleware
import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/gin-gonic/gin"
"filecodebox/internal/cache"
)
func rateLimitRouter(rl *RateLimiter, kind string) *gin.Engine {
r := gin.New()
r.Use(ClientIP(nil))
r.GET("/limited", rl.RequireRateLimit(kind), func(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true})
})
return r
}
func TestRateLimiterBlocksAfterCount(t *testing.T) {
mem := cache.NewMemory()
defer mem.Close()
rl := NewRateLimiter(mem, map[string]LimitRule{
LimitMeta: {Count: 3, Window: time.Minute},
})
r := rateLimitRouter(rl, LimitMeta)
// 前 3 次通过
for i := 0; i < 3; i++ {
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/limited", nil))
if w.Code != http.StatusOK {
t.Fatalf("第 %d 次应通过: %d", i+1, w.Code)
}
}
// 第 4 次 423
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("GET", "/limited", nil))
if w.Code != http.StatusLocked {
t.Fatalf("超限应 423: %d", w.Code)
}
}
func TestRateLimiterAddAfterSuccess(t *testing.T) {
// 模拟 upload 语义:Check 放行 + handler 成功后 Add
mem := cache.NewMemory()
defer mem.Close()
rl := NewRateLimiter(mem, map[string]LimitRule{
LimitUpload: {Count: 2, Window: time.Minute},
})
r := gin.New()
r.Use(ClientIP(nil))
r.POST("/upload", func(c *gin.Context) {
if allowed, _ := rl.Check(c, LimitUpload); !allowed {
c.JSON(http.StatusLocked, gin.H{"err": "too many"})
return
}
rl.Add(c, LimitUpload) // 成功上传计数
c.JSON(http.StatusOK, gin.H{"ok": true})
})
for i := 0; i < 2; i++ {
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("POST", "/upload", nil))
if w.Code != http.StatusOK {
t.Fatalf("第 %d 次上传应通过", i+1)
}
}
w := httptest.NewRecorder()
r.ServeHTTP(w, httptest.NewRequest("POST", "/upload", nil))
if w.Code != http.StatusLocked {
t.Fatalf("第 3 次上传应被拒绝: %d", w.Code)
}
}
func TestParseTrustedProxies(t *testing.T) {
nets := ParseTrustedProxies([]string{"10.0.0.0/8", "192.168.1.1", "", "bad-input"})
if len(nets) != 2 {
t.Fatalf("应解析出 2 个可信网段: %d", len(nets))
}
}
func TestClientIPFromTrustedProxy(t *testing.T) {
// 对齐参考语义:仅当直连地址可信时才解析 XFF,
// 且从右往左返回第一个"非可信代理"地址(该地址即最近可信代理看到的客户端)。
nets := ParseTrustedProxies([]string{"127.0.0.0/8", "10.0.0.0/8"})
r := gin.New()
r.Use(ClientIP(nets))
var seen string
r.GET("/ip", func(c *gin.Context) { seen = GetClientIP(c) })
req := httptest.NewRequest("GET", "/ip", nil)
req.RemoteAddr = "127.0.0.1:5000"
req.Header.Set("X-Forwarded-For", "203.0.113.9, 10.0.0.1")
r.ServeHTTP(httptest.NewRecorder(), req)
if seen != "203.0.113.9" {
t.Fatalf("多级可信代理下应取最左非可信地址: %s", seen)
}
// 仅直连可信:XFF 右起第一个(10.0.0.1)非可信 → 取它
nets2 := ParseTrustedProxies([]string{"127.0.0.0/8"})
r2 := gin.New()
r2.Use(ClientIP(nets2))
var seen2 string
r2.GET("/ip", func(c *gin.Context) { seen2 = GetClientIP(c) })
req2 := httptest.NewRequest("GET", "/ip", nil)
req2.RemoteAddr = "127.0.0.1:5000"
req2.Header.Set("X-Forwarded-For", "203.0.113.9, 10.0.0.1")
r2.ServeHTTP(httptest.NewRecorder(), req2)
if seen2 != "10.0.0.1" {
t.Fatalf("右起第一个非可信地址应为 10.0.0.1: %s", seen2)
}
// 非可信直连:忽略伪造头
seen = ""
req = httptest.NewRequest("GET", "/ip", nil)
req.RemoteAddr = "8.8.8.8:1234"
req.Header.Set("X-Forwarded-For", "1.2.3.4")
r.ServeHTTP(httptest.NewRecorder(), req)
if seen != "8.8.8.8" {
t.Fatalf("非可信直连应忽略 XFF: %s", seen)
}
}