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:
@@ -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)
|
||||
// 未命中审计动作的请求直接放行,不产生审计记录。
|
||||
// L5:admin 类动作同样需要建 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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
// Package middleware — bodylimit.go:全局请求体大小限制。
|
||||
//
|
||||
// 安全审计 M3:此前服务对请求体无任何上限,text/plain 形态的绑定会先
|
||||
// io.ReadAll 全量进内存、multipart 大文件先完整落盘后才被策略拒绝,
|
||||
// 未认证攻击者可借大 body 消耗内存/磁盘/带宽。
|
||||
// 本中间件用 http.MaxBytesReader 按「路径类别」给出上限:
|
||||
// - 管理端(/admin/*):1MiB;
|
||||
// - 文本/元信息/取件类(/share/text|metadata|select):1MiB;
|
||||
// - 其余(含上传):maxFileSize(0=回落 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()
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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_secret;expires 为会话有效期。
|
||||
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_secret(settings 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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user